# tmf_v3 **Repository Path**: WorldEating/tmf_v3 ## Basic Information - **Project Name**: tmf_v3 - **Description**: 从 v2 的 FastText 升级到深度学习模型(BERT / RoBERTa 等),以获得更高的分类精度。 - **Primary Language**: Unknown - **License**: Not specified - **Default Branch**: master - **Homepage**: None - **GVP Project**: No ## Statistics - **Stars**: 0 - **Forks**: 0 - **Created**: 2026-06-27 - **Last Updated**: 2026-06-30 ## Categories & Tags **Categories**: Uncategorized **Tags**: None ## README # TMF v3 — 中文文本分类服务 (Deep Learning) 基于 **Transformers** 的中文短文本分类项目,支持离线训练流水线、在线推理服务和批量分类,附带 Streamlit 可视化交互界面。 > 从 v2 的 FastText 升级到深度学习模型(BERT / RoBERTa 等),以获得更高的分类精度。 **10 个分类类别**:财经 `finance` · 房产 `realty` · 股票 `stocks` · 教育 `education` · 科技 `science` · 社会 `society` · 政治 `politics` · 体育 `sports` · 游戏 `game` · 娱乐 `entertainment` **v3 达成**:测试集 Acc **94.75%** / F1 **94.75%** | GPU 推理 p50 **5.8ms** | FP16 吞吐 **1053 条/s** --- ## 与 v2 的区别 | 维度 | v2 (FastText) | v3 (Transformers) | |------|---------------|-------------------| | 模型架构 | FastText + n-gram | chinese-roberta-wwm-ext (102M) | | 上下文理解 | 词袋模型,无语序 | 双向注意力,完整上下文 | | F1 | 91.18% | **94.75%** | | 推理延迟 | 0.04 ms (CPU) | 5.8ms (GPU) / 77ms (CPU) | | 模型体积 | 447 MB (量化) | 391 MB (float32) / 195 MB (FP16) | | 训练时间 | ~5 min (autotune) | 55 min (GPU, 180k 条) | | 分词方式 | jieba 分词 | BERT Tokenizer (BPE) | --- ## 项目结构 ``` tmf_v3/ ├── wsgi.py # Flask 服务入口 ├── run_pipeline.py # 离线训练流水线编排 ├── app/ # Web 应用层 │ ├── __init__.py # Flask 工厂函数 + 请求日志中间件 + CORS + 异常兜底 │ ├── config.py # 环境配置(开发 / 生产) │ ├── extensions.py # 深度学习模型单例(启动时加载一次) │ └── v3/ │ ├── __init__.py │ └── predict.py # /api/v3/predict(单条 Top-K)+ /predict_batch(批量) ├── src/ # 机器学习层(不依赖 Flask) │ ├── config.py # 路径与常量集中管理 │ ├── logger.py # 统一日志模块(控制台 + 文件双通道,按天轮转) │ ├── data_pre.py # 数据预处理(清洗 → Tokenization → Dataset) │ ├── model_train.py # 深度学习模型训练 + 验证集评估 │ ├── model_eval.py # 测试集评估(全量 + 抽样压测 + 延迟统计) │ ├── model_compress.py # 模型量化/蒸馏/剪枝对比 │ ├── data_eda.py # 探索性数据分析 │ └── utils.py # 工具函数 ├── front/ │ └── tmf_app.py # Streamlit 界面(浅色主题 · 分类光谱 · 10类示例) ├── .streamlit/ │ └── config.toml # Streamlit 主题配置(冷白 + 深紫) ├── tests/ │ ├── test_api.py # API 接口测试 │ ├── test_data_pre.py # 文本清洗 / 过滤 / 标签加载 │ ├── test_model_train.py # 分类报告 / 混淆矩阵 / 错分分析 │ └── test_model_eval.py # 延迟统计 / 模型推理 ├── data/ │ ├── stopwords.txt # 停用词表(749 个) │ ├── tmf_class.txt # 类别标签(10 个) │ ├── sample_batch.csv # 批量分类测试样本 │ ├── raw/ # 原始 TSV 数据 │ └── pre/ # 预处理后数据(Git 不纳管) ├── models/ │ └── .gitkeep # 模型目录(Git 不纳管) ├── logs/ # 运行日志(按天轮转,保留 7 天) ├── pytest.ini # pytest 标记注册 ├── requirements.txt # 依赖清单 └── README.md ``` --- ## 快速开始 ### 环境要求 - Python 3.12+ - PyTorch 2.12+ (推荐 CUDA 13.2) - Conda 环境 `tmf_v3` ```bash # 安装依赖 pip install -r requirements.txt ``` ### 离线训练 ```bash # 一键流水线(5 步可独立开关,训练默认关闭) python run_pipeline.py # 或分步执行 python src/data_pre.py # Step 1: 数据预处理 → Tokenization + Dataset python src/model_train.py # Step 2: BERT 微调训练 + 验证集评估 python src/model_eval.py # Step 3: 测试集评估 + 延迟/吞吐压测 # Step 4 & 5 在 model_compress.py 中: python src/model_compress.py # Step 4: 知识蒸馏 (教师→学生小模型) # Step 5: FP16/INT8/蒸馏对比 ``` `run_pipeline.py` 顶部配置项: ```python # 流水线开关(训练默认关闭,已完成) ENABLE_PREPROCESS = False ENABLE_TRAIN = False ENABLE_EVALUATE = False ENABLE_DISTILL = False # 知识蒸馏 — 教师模型指导小模型训练 ENABLE_COMPRESS = True # FP16/INT8/蒸馏压缩对比 ``` 训练参数在 `src/config.py` 中统一管理: ```python batch_size: int = 8 # 环境变量 TMF_BATCH_SIZE learning_rate: float = 2e-5 # 环境变量 TMF_LR num_epochs: int = 3 # 环境变量 TMF_EPOCHS max_length: int = 128 # 环境变量 TMF_MAX_LENGTH ``` ### 启动在线服务 ```bash # 终端 1:启动 Flask API(端口 5000) python wsgi.py # 终端 2:启动 Streamlit 界面(端口 8501) streamlit run front/tmf_app.py ``` ### 切换模型 API 默认自动使用 FP16 模型 (`models/text-clf-model-fp16/`),可通过环境变量切换: ```bash # 使用 float32 原始模型 TMF_MODEL_PATH=models/text-clf-model python wsgi.py # 使用指定路径 TMF_MODEL_PATH=models/text-clf-model-int8 python wsgi.py ``` 模型/Tokenizer 启动时加载一次,整个应用生命周期内复用。 --- ## API 接口 ### 单条预测 `POST /api/v3/predict` ```bash curl -X POST http://127.0.0.1:5000/api/v3/predict \ -H "Content-Type: application/json" \ -d '{"text": "中超联赛战火重燃 北京国安工体迎战武汉三镇", "top_k": 3}' ``` 响应: ```json { "label": "sports", "probability": 0.9512, "code": 0, "message": "ok", "top3": [ {"label": "sports", "probability": 0.9512}, {"label": "game", "probability": 0.0321}, {"label": "entertainment", "probability": 0.0123} ] } ``` ### 批量分类 `POST /api/v3/predict_batch` ```bash curl -F "file=@data/sample_batch.csv" http://127.0.0.1:5000/api/v3/predict_batch ``` ### 健康检查 `GET /` ```bash curl http://127.0.0.1:5000/ # → {"code": 0, "message": "ok"} ``` --- ## 运行测试 ```bash pytest tests/ -v # 全部用例(53 个) pytest tests/ -m "not slow and not api" # 快速单元测试(34 个,无需服务) pytest tests/ -m slow -v # 模型推理测试(需模型文件) pytest tests/ -m api -v # API 接口测试(需启动 Flask) pytest -m smoke -v # 冒烟测试 ``` --- ## 技术架构 ``` 离线训练流水线 原始 TSV → data_pre.py → Tokenization + Dataset → model_train.py → model_eval.py → 模型 │ ▼ 在线推理服务 wsgi.py → app/__init__.py → app/extensions.py ← 启动加载 │ │ ▼ ▼ (只加载一次) /api/v3/predict model + tokenizer + labels /api/v3/predict_batch (应用级单例,启动时挂载) │ ▼ front/tmf_app.py (Streamlit 单条推理 + 批量分类) ``` | 层级 | 职责 | 说明 | |------|------|------| | `src/` | 深度学习 | 预处理、训练、评估,不依赖 Web 框架 | | `app/` | Web 服务 | 工厂模式 Flask,蓝图版本化路由,请求日志中间件 | | `front/` | 用户界面 | Streamlit 消费 API,单条推理 + 批量分类两 Tab | | `tests/` | 测试 | pytest + parametrize,smoke / api 标签 | --- ## 迁移计划 (v2 → v3) — ✅ 全部完成 - [x] **数据预处理**:jieba 分词 → HuggingFace Tokenizer - [x] **模型训练**:FastText → BERT/RoBERTa + Trainer API - [x] **模型评估**:测试集评估 + GPU/CPU 延迟压测 - [x] **API 适配**:`/api/v3/*` 端点,兼容 v2 响应格式 - [x] **前端适配**:Streamlit 指向 v3 API - [x] **模型压缩**:FP16 (195MB) / INT8 动态量化 --- ## 版本历史 | 版本 | 日期 | 更新内容 | |------|------|----------| | `v3.2.0` | 2026-06-30 | 知识蒸馏(教师 BERT-base → 学生 rbt3)、流水线 5 步编排、压缩对比含蒸馏模型 | | `v3.1.0` | 2026-06-26 | BERT 微调完成: Acc 94.75%/F1 94.75%, FP16 压缩, Flask API, Streamlit | | `v2.1.0` | 2026-06-26 | 批量分类接口 + Streamlit 批量 Tab | | `v2.0.0` | 2026-06-25 | FastText 迁移:移除 sklearn/ONNX,autotune 自动调参 | --- ## License 仅供学习交流使用。