简介:这是基于BERT模型的深度学习中文文本分类Python项目,面向计算机、人工智能等相关专业的学生或算法入门者,旨在解决中文新闻文本的分类问题;项目具备完整的训练、评估、预测与接口调用链路,适合作为毕业设计、课程设计或实战练手项目。压缩包共18个文件,以11个Python源码文件为主,并配有TXT格式的2万条新闻训练与测试数据、JSON标签映射、Jupyter Notebook演示文档和说明文件,整体体积约1008KB,结构清晰,便于按模块学习与二次开发。目前已有350人学习下载,代码经测试可正常运行。项目中包含数据预处理、BERT模型构建、训练器、预测器、评估指标等核心模块,并提供HTTP接口及客户端调用脚本,可体验从训练到部署的完整流程;附带Notebook和Shell脚本,适合快速复现实验并在此基础上扩展,也便于深入理解BERT文本分类的微调机制。
1. 这个标题把训练、数据和部署串成了闭环:到底值不值得动手
拿到这个项目标题,我第一反应不是 BERT 有多强,而是那 20000 条新闻数据。做深度学习文本分类的人都清楚,跑通一个 MNIST 手写识别说明不了任何问题,真实世界里的中文文本分类,大部分时间都耗在数据清洗、类别平衡、模型调参和接口部署这些杂活上。这个项目把 BERT 模型的完整训练链路、20000 条新闻数据集、HTTP 接口打在同一个包里,正好补上了从算法 demo 到可调用服务之间那块让人头疼的空白。适合三类人:拿它做毕业论文的在校生,想快速给系统接入文本分类能力的后端工程师,以及刚入门深度学习、想找一个不掺水的中文实战案例的开发者。下面对它的核心链路逐一拆解,从为什么选 BERT,到数据和参数怎么设,再到 HTTP 接口怎么包。
2. 为什么是 BERT:从词向量到预训练模型,中文文本分类的选型逻辑
2.1 从 TF-IDF 到 BERT:文本分类在学的“表示”到底是什么
文本分类的问题定义很简单:给定一段文本,映射到一个或多个类别标签。难的是“映射”这两个字背后的特征表达。传统机器学习模型的做法是用 TF-IDF 或词袋把文本变成高维稀疏向量,再丢给 SVM、朴素贝叶斯或逻辑回归。这套做法在短文本、类别少、数据量小的场景下仍然能打,但核心缺陷很明显:词序信息几乎丢失。“中国队战胜了日本队”和“日本队战胜了中国队”在词袋视角下可能只差一个位置,向量距离非常近,分类边界很难拉开。
Word2Vec 比词袋进了一步,把词变成稠密向量,词汇之间的语义相似度开始有意义。但句子表示通常是对词向量做平均或加权平均,仍然是“词袋 + 稠密向量”的思路,上下文交互没有被建模。TextCNN 用多个卷积核去捕捉 n-gram 局部特征,对短文本效果好,训练快,但卷积核的感受野有限,长距离依赖要么靠加深层数,要么靠膨胀卷积,代价都不小。文本分类的本质,是在学一句话里哪些词在什么语境下对类别判断起决定作用,这需要一个能建模上下文的表示方法。
正因如此,BERT 这类预训练模型出现后,文本分类的基线一下被抬高了一大截。BERT 不再手工设计特征,而是通过海量语料预训练出“懂语法、懂常识、懂上下文”的通用表示,下游任务只需微调。分类任务的准确率在绝大多数中文数据集上明显超过传统机器学习模型,同时还能省掉大量特征工程的重复劳动。
2.2 BERT 的输入输出机制:为什么 CLS 向量能当分类特征
Devlin、Chang、Lee 等人在 2018 年发表的 BERT 论文里,核心是用 Transformer 的 Encoder 堆出深层双向编码结构。base 版本是 12 层 Transformer、hidden_size 为 768、12 个注意力头,参数量大约 1.1 亿。这里的“双向”是关键:每个 token 在每一层都会同时融合左边和右边的信息,这和传统的从左到右语言模型有本质区别。对中文文本来说,同一个“苹果”在不同语境里指水果还是手机品牌,只有看到上下文才能判断,双向编码天然适合这个任务。
输入侧,BERT 会把文本拼成这样的结构:开头放一个 [CLS] 标记,中间是切分后的 token 序列,句子之间用 [SEP] 分隔。对单句分类来说,你只需要把原始文本交给 tokenizer,它自动补好 [CLS] 和 [SEP],再加上 segment embedding 和 position embedding。输出侧,每个 token 位置都会得到一个 768 维向量,而 [CLS] 位置的向量被设计为整句话的聚合表示。微调分类任务时,把这个向量接一个全连接层,再过一个 softmax,就得到每个类别的概率分布。
为什么偏偏是 [CLS] 而不是把所有 token 的向量做平均?因为预训练阶段 [CLS] 被训练成了“聚合整个序列信息”的角色,微调时它的输出自然携带了句级语义。实践中直接取最后一层 [CLS] 向量接分类头就是最标准、最稳定的做法。要注意的是,如果你把 [CLS] 向量换成平均池化,有时效果反而更好,但这是后话,初学阶段不要为了剪枝而剪枝,先跑通标准做法。
2.3 选型对比:TextCNN、BERT、大模型,成本与效果差在哪
| 方案 | 训练成本 | 推理延迟(CPU) | 中文分类效果 | 典型场景 |
|---|---|---|---|---|
| TF-IDF + SVM | 秒级 | 毫秒级 | 中等偏下 | 小规模、类别少、对延迟极敏感 |
| TextCNN | 分钟级 | 毫秒级 | 中等 | 数据量不大、长文本、资源受限 |
| BERT-base | 单卡数十分钟 | 几十到几百毫秒 | 高 | 大部分中文分类业务,性价比最高 |
| 大模型(ChatGLM 等) | 微调成本高 | 秒级 | 高,但抖动明显 | 少样本、零样本、需要推理解释 |
从这张表能看出,BERT-base 处在“效果和成本都很合理”的位置上。TextCNN 的优势在速度,如果你有 10 万条以上已标注数据且文本规律性强,TextCNN 微调后也许只比 BERT 低两三个点,但部署成本低一个量级。反过来,基于大模型做文本分类已经是新趋势,尤其适合标注样本很少的场景,大模型的 zero-shot 能力可以直接跳过训练环节,不过服务化成本和 GPU 资源不是所有团队都能承担的。
我一般会这样选:数据量在几千到几万条、标注规范、类别固定,直接用 BERT-base 微调;如果上线后 CPU 资源紧张且延迟要求小于 20 毫秒,再考虑蒸馏到 TextCNN 或转 ONNX。不要一上来就追大模型,BERT 微调这套流程的理解深度,决定了你后面踩坑时能不能快速定位问题。
3. 20000 条新闻数据的处理与训练:从数据体检到参数设置
3.1 数据体检:编码、类别分布、文本长度,三个必查项
拿到数据集先别急着写训练脚本。中文文本数据最常见的坑全部集中在第一步:文件编码是 GBK 还是 UTF-8,标签列是不是有空值,类别是不是严重不均衡。20000 条新闻听起来不少,但如果某一个类别占了 80%,模型学到的基本就是“猜那个大类”,整体准确率看着很高,实际没有泛化能力。用 pandas 做一次快速体检:
import pandas as pd df = pd.read_csv("news.csv", encoding="utf-8") print(df.head()) print(df["label"].value_counts()) df["text_len"] = df["text"].astype(str).map(len) print(df["text_len"].describe())这里的 value_counts 一眼就能看出类别分布是否均衡。如果某个类别只有几十条,模型基本学不好,后面要重点观察这个类别的召回率。text_len.describe() 输出的是文本长度的分位数,重点看 75% 和 max 两个值:如果 75% 的文本长度在 200 字以内,但 max 到了 5000,说明存在严重的超长尾,训练时要么截断要么过滤。编码问题通常表现为读出来全是乱码,或者报 UnicodeDecodeError,遇到这种情况把 encoding 参数换成 gbk 再试。
数据切分建议按 8:1:1 划分训练、验证、测试集。注意用 sklearn 的 train_test_split 时要设置 stratify 参数,按类别比例分层采样,否则切分后小类别的样本可能全跑进训练集,验证集里根本看不到它。数据量只有 20000 条时,随机切分和分层切分的差异会被放大,分层是必须的。
3.2 数据加载器与 Tokenizer 配合:Dataset 类的写法
PyTorch 训练 BERT 的标准姿势是自定义 Dataset 类,在getitem里完成 tokenizer 编码。常见做法是在初始化时先加载 bert-base-chinese 的 tokenizer,然后对每条文本做编码。这样做的好处是内存占用小,每条样本实时转成 input_ids;缺点是每个 epoch 都要重复编码,数据量大时拖慢训练。先看代码:
from torch.utils.data import Dataset from transformers import BertTokenizer import torch tokenizer = BertTokenizer.from_pretrained("bert-base-chinese") class NewsDataset(Dataset): def __init__(self, texts, labels, max_len): self.texts = texts self.labels = labels self.max_len = max_len def __len__(self): return len(self.texts) def __getitem__(self, idx): encoded = tokenizer( self.texts[idx], truncation=True, padding="max_length", max_length=self.max_len, return_tensors="pt", ) return { "input_ids": encoded["input_ids"].squeeze(0), "attention_mask": encoded["attention_mask"].squeeze(0), "label": torch.tensor(self.labels[idx], dtype=torch.long), }注意几个参数。truncation=True 表示超长文本直接截断到 max_length,padding="max_length" 表示不够长的补 0,让一个 batch 里的样本长度完全一致,这样才好拼成矩阵。max_length 的选择很关键,建议先用数据体检里 describe() 的 90% 分位数作为初始值,后面会专门讲这个参数的坑。squeeze(0) 是因为 tokenizer 返回的是形状 [1, seq_len] 的 tensor,batch 维度在这个场景是多余的。
如果你的显存足够而 CPU 编码成为瓶颈,可以在离线阶段把全部文本编码成 numpy 数组存盘,训练时直接读数组。20000 条新闻编码后的文件大约几百 MB,还在可接受范围。对初学阶段,上面这个 Dataset 写法最直观,也最容易调试。
3.3 训练参数设置:batch_size、学习率、max_len 与 epochs 的配合
BERT 微调有一套约定俗成的参数区间,不是越大越好,也不是越小越好。下面的参数表是我在类似规模中文分类任务上的常用起点:
| 参数 | 建议值 | 说明 |
|---|---|---|
| batch_size | 16(12G 显存) / 32(24G 显存) | 太大容易 OOM,太小收敛慢 |
| learning_rate | 2e-5 ~ 5e-5 | BERT 微调不建议超过 5e-5 |
| max_len | 128 或 256 | 按文本长度分布定,别盲目用 512 |
| epochs | 3 ~ 5 | 20000 条数据通常 3 轮内收敛 |
| warmup_ratio | 0.1 | 前 10% 的 step 学习率线性上升 |
| weight_decay | 0.01 | 对非 bias 和 LayerNorm 参数生效 |
训练循环本体倒不复杂,关键在优化器和调度器的配合:
from transformers import AdamW, get_linear_schedule_with_warmup optimizer = AdamW(model.parameters(), lr=2e-5, weight_decay=0.01) total_steps = len(train_loader) * epochs scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=int(total_steps * 0.1), num_training_steps=total_steps, ) for epoch in range(epochs): model.train() for batch in train_loader: input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) labels = batch["label"].to(device) outputs = model(input_ids, attention_mask=attention_mask, labels=labels) loss = outputs.loss loss.backward() optimizer.step() scheduler.step() optimizer.zero_grad()warmup 的作用是避免训练刚开始时模型参数剧烈震荡;BERT 的预训练参数已经在一个很好的位置,微调时学习率过大容易把学到的通用语义冲掉。zero_grad 放在 backward 之前还是之后都行,放在 step 之后是最常见的写法,注意不要漏掉,否则梯度会跨 batch 累加。训练过程中我习惯每个 epoch 都保存一次 checkpoint,并且记录验证集 loss,如果连续两个 epoch 验证集 loss 不再下降,就提前停掉。
3.4 模型评估:准确率之外,F1 和混淆矩阵才能反映真实问题
训练结束后,评估流程要覆盖验证集和测试集。准确率是大家最直观的指标,但在类别不均衡时它很容易骗人:如果某个类别占了一半样本,模型全猜这个类别准确率也有 50%。所以 F1 和混淆矩阵是必看的。用 sklearn 的 classification_report 一把梭:
from sklearn.metrics import classification_report, confusion_matrix import torch model.eval() preds, true_labels = [], [] with torch.no_grad(): for batch in val_loader: input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) logits = model(input_ids, attention_mask=attention_mask).logits preds.extend(torch.argmax(logits, dim=-1).cpu().numpy()) true_labels.extend(batch["label"].numpy()) print(classification_report(true_labels, preds, target_names=label_names))torch.no_grad() 一定要包住推理循环,否则模型会为推理过程构建计算图,显存直接翻倍还多。输出里重点看每个类别的 recall:如果某一个类别的召回率明显低于整体水平,基本就是样本量太少或者文本特征和别的类别太接近。混淆矩阵能进一步告诉你它被分到了哪个类别,比如“娱乐”新闻频繁被分到“体育”,那就去看训练集里这两类样本的标注是否本身就有歧义。
4. BERT 文本分类的常见问题与避坑指南:五个高频翻车点
4.1 显存不足:不只是调小 batch_size 这一条路
现象是训练跑到第二个 batch 直接报 CUDA out of memory。新手的第一反应是 batch_size 从 16 改成 8,发现还是爆,再改 4,勉强能跑但慢得离谱。原因往往不是 batch_size 本身,而是 max_len 太大。BERT 的显存占用和序列长度的平方相关,512 长度的显存消耗远大于 128,哪怕 batch_size 只有 4,照样爆显存。解决的优先级是这样:先用 nvidia-smi 确认没有其他进程占显存,然后把 max_len 压到 128,最后才是调 batch_size。如果这三步都不够,还有一个不动代码的办法:用梯度累积。
梯度累积的思路上,每 N 个 batch 的梯度累加后再更新一次参数,等价于把 batch_size 扩大了 N 倍。代码上只需要在 optimizer.zero_grad() 的时机上做文章:每个 batch 做 loss.backward(),但只在 step 计数达到 N 的倍数时才调 optimizer.step() 和 scheduler.step()。注意 batch 的损失要除以 N,否则学习率的实际效果变大。这个技巧在显存只有 6G 的笔记本上尤其好用。
4.2 中文文本乱码与 [UNK]:编码问题藏在第一步
现象是训练 loss 下降正常,但预测阶段大量文本被 tokenizer 切成了 [UNK],分类结果基本靠猜。查了一圈才发现是数据文件编码不是 UTF-8,读进来以后每个字符都变成了替换符。另一个隐蔽场景:同一个 CSV 文件里大部分是 UTF-8,但用户从某个旧系统导出时混入了几行 GBK,pandas 按 UTF-8 读会直接报错,按 GBK 读又出现乱码。
解决的办法是在数据体检阶段就锁死编码。用 Python 的 chardet 判断文件编码:
import chardet with open("news.csv", "rb") as f: raw = f.read(100000) print(chardet.detect(raw))如果是 GBK,用 iconv 转成 UTF-8:
iconv -f gbk -t utf-8 news.csv > news_utf8.csv转换之后再跑一遍 3.1 的体检脚本,确认 text_len 分布正常、没有 [UNK] 堆积。这里最容易翻车的点是:pandas 读文件时第几行开始乱码并不会报错,而是静默把内容替换成乱码字符,等你发现时数据已经脏了,所以清洗步骤别省。
4.3 loss 不降、过拟合:先分清是数据问题还是学习率问题
训练时 loss 纹丝不动是最折磨人的现象之一。造成它的原因有几种:一是数据加载器返回的 label 全是同一个值,相当于模型在学一个常数,此时去看 train_loader 里的 batch 内容;二是 tokenizer 加载错了,比如用了不是中文的预训练模型,分词全变成 [UNK],模型根本没有有效输入;三是学习率设置不合理,BERT 微调用 1e-3 这种在 CNN 上常用的学习率,loss 必然剧烈震荡甚至直接发散。
反过来,验证集 loss 到第三个 epoch 开始回升,训练 loss 还在降,这是典型的过拟合信号。20000 条数据量不大,BERT-base 参数量 1.1 亿,微调阶段在验证集上出现过拟合很正常。先加早停机制,patience 设为 2;再降学习率到 1e-5 或 2e-5 重跑;最后才考虑加 dropout 或数据增强。不要一上来就加正则项,BERT 微调对权重衰减很敏感,weight_decay 从 0.01 改成 0.1 的效果经常适得其反。
4.4 模型加载慢与设备不匹配:checkpoint 的保存和加载规范
现象是接口每次重启后第一个请求要等 10 秒以上,或者本地 CPU 能加载的模型放到服务器 GPU 上直接报 device mismatch。根因是 checkpoint 保存时没有统一约定设备,或者模型类每次都重新初始化。正确做法是保存时用 model.save_pretrained 和 tokenizer.save_pretrained,这是一个目录,包含模型权重和配置文件:
model.save_pretrained("./bert-news-model") tokenizer.save_pretrained("./bert-news-model")加载时用 from_pretrained 一次读入,并显式指定设备映射:
model = BertForSequenceClassification.from_pretrained("./bert-news-model") model.to(device)服务器上如果只有 CPU,加载时加 map_location="cpu";如果 GPU 卡号变了,用 torch.load 时指定 map_location={"cuda:0": "cuda:1"}。还有一类踩坑:训练时用的是多卡 DataParallel,保存下来的 state_dict 键名带 module. 前缀,加载到单卡模型时报 size mismatch。遇到这种情况,加载后把键名里多余的 module. 去掉即可,通常一行代码解决。
4.5 类别不均衡:准确率虚高时,去看混淆矩阵
现象是测试集准确率 95%,看起来模型很优秀,但每种类别单独看,某个类别的 F1 只有 0.2。这是文本分类最常见的“假成功”。原因通常是数据集中“体育”占 60%,“房产”占 2%,模型把所有样本都预测成“体育”就能拿到高准确率。光调网络结构解决不了这个,必须从数据和训练策略下手。
第一步是在 loss 上加类别权重,PyTorch 的 CrossEntropyLoss 直接支持 weight 参数,把小类别的 loss 放大。第二步是评估指标改用 macro-F1,它把所有类别的 F1 取平均,大类别无法掩盖小类别的问题。第三步是数据层面做欠采样或过采样,20000 条数据量不算大,把小类别的样本复制几份或者用回译做增强都是可行方案。这里最容易犯的错是只看整体准确率就宣布项目完成,一定要坚持用 classification_report 打印每个类别的指标。
5. HTTP 接口化:把训练好的 BERT 模型包成文本分类服务
5.1 Flask 还是 FastAPI:文本分类服务的框架选型
模型训练完,下一步是把它变成可以给别人调用的 HTTP 服务。框架选择上,Flask 和 FastAPI 是目前的两个主流答案。做一个简单对比:
| 对比项 | Flask | FastAPI |
|---|---|---|
| 上手难度 | 极低 | 低,需要理解类型注解 |
| 自动文档 | 无 | 自带 /docs,可在线调试 |
| 异步支持 | 需要额外插件 | 原生 async |
| 请求参数校验 | 手写 | Pydantic 自动校验 |
| AI 推理服务生态 | 老项目多 | 新项目主流 |
如果是从零开始写这个标题要求的接口服务,我推荐 FastAPI。原因很具体:模型推理接口最烦的就是调用方传了一个空字符串或者不是 String 的类型,FastAPI 的 Pydantic 模型能直接在请求入口拦截;/docs 页面方便你快速验证接口;异步支持虽然对 GPU 推理意义有限,但对后续接入异步调用方没有障碍。如果你所在的团队老项目已经全是 Flask,那也没必要强行迁移,Flask + threading 同样能完成任务。
5.2 最小可用接口:从模型加载到 POST /predict
一个可用的分类接口至少包含四部分:模型和 tokenizer 加载、请求结构定义、推理逻辑、错误处理。下面这个脚本是 FASTAPI 的最小实现,可以直接跑在服务器上:
from fastapi import FastAPI, HTTPException from pydantic import BaseModel import torch from transformers import BertTokenizer, BertForSequenceClassification app = FastAPI() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = BertForSequenceClassification.from_pretrained("./bert-news-model") tokenizer = BertTokenizer.from_pretrained("./bert-news-model") model.to(device) model.eval() id_to_label = {0: "体育", 1: "财经", 2: "娱乐", 3: "科技"} class PredictRequest(BaseModel): text: str class PredictResponse(BaseModel): label: str confidence: float @app.post("/predict", response_model=PredictResponse) def predict(req: PredictRequest): if not req.text.strip(): raise HTTPException(status_code=400, detail="text is empty") encoded = tokenizer( req.text, truncation=True, padding="max_length", max_length=128, return_tensors="pt", ).to(device) with torch.no_grad(): logits = model(**encoded).logits prob = torch.softmax(logits, dim=-1) label_id = torch.argmax(prob, dim=-1).item() confidence = prob[0, label_id].item() return {"label": id_to_label[label_id], "confidence": confidence}模块导入时就把模型加载到内存,这是避免冷启动的关键。model.eval() 一定要加,因为 BERT 里有 dropout,不切到 eval 模式时每次推理的 dropout mask 不同,同一个文本两次调用结果会有细微差异。id_to_label 这个映射最好在保存模型时在 config.json 里记录类别名,别硬编码在脚本里,否则类别一多就难维护。启动方式在命令行执行 uvicorn main:app --host 0.0.0.0 --port 8000。
5.3 请求与返回设计:参数校验、错误码、超时处理
接口设计直接决定调用方的接入成本。请求体用 JSON,字段名要直观,text 就是待分类文本。返回体里 label 和 confidence 是必有字段,confidence 保留小数点后四位即可。有一点容易被忽视:空字符串、纯空格、超长文本,这三类请求一定要在接口层拦截,而不是让 tokenizer 和模型去处理,否则返回的 500 错误会让调用方无从排查。超长文本的默认策略是截断到 max_length,但更合理的做法是返回 422 并提示“文本超长,请控制在 500 字以内”。
超时方面,BERT 在 CPU 上单条推理可能耗时 200 到 500 毫秒,GPU 上 20 到 50 毫秒。接口层建议把读超时设成 3 秒,写超时 5 秒。如果你的服务跑在 CPU 上,要提前在 API 文档里写明预期延迟,否则调用方用 1 秒超时来调,结果每次报错。加一个可选的 latency_ms 字段对排查问题很有帮助,前端可以直接看到哪一段耗时异常。
5.4 并发与性能优化:加锁、批处理与 ONNX 导出
模型服务上线后,第一个并发测试就可能翻车。PyTorch 的模型在 CUDA 上做推理时不是线程安全的,多个请求同时进入同一个 model.forward,轻则结果错乱,重则直接报非法内存访问。最简单的办法是给推理过程加一个线程锁:
import threading infer_lock = threading.Lock() @app.post("/predict") def predict(req: PredictRequest): with infer_lock: encoded = tokenizer(...).to(device) with torch.no_grad(): logits = model(**encoded).logits ...加锁之后并发会串行化,QPS 上不去,但对小规模内部服务完全够用。如果 QPS 要求超过 10,两个方向可以考虑:一是把请求攒成 batch 推理,几十条一起输入模型,GPU 利用率能拉满,但要自己实现请求队列;二是把模型转成 ONNX,用 onnxruntime-gpu 推理,单条延迟通常能降到 PyTorch 的 50% 左右,导出命令很简单:
python -m transformers.onnx --model=./bert-news-model ./bert-news-onnxONNX 导出后要注意验证输出和 PyTorch 是否一致,个别算子可能因为版本原因不被支持,此时查一下 opset 版本或者把模型降级用 PyTorch 跑。这些优化不要一开始就全上,先把加了锁的版本跑稳,再根据压测结果决定要不要做批量推理和 ONNX 转换。
6. 进阶方向:多标签、增量训练与轻量化落地
模型跑通、接口上线,只代表这个标题要求的闭环完成了。如果手里有余力,最值得做的三个方向是多标签分类、增量训练和轻量化部署。
多标签分类在新闻场景很常见,一篇新闻可能同时属于“财经”和“政策”。改造其实很小:输出层换成 sigmoid 而不是 softmax,损失函数换成 BCEWithLogitsLoss,评估时对每个类别单独算 F1 再平均。代码层面对应的改动只有模型输出层的激活函数和 loss 计算方式,但需要重新处理数据集,把原来的单标签 label 改成 multi-hot 向量。
增量训练解决的是“新类别来了怎么办”的问题。常见做法是直接在原有 checkpoint 上继续训练,但学习率要降到 1e-5 以下,否则模型会快速遗忘旧类别知识。更稳妥的方案是固定 BERT 底层层参数,只训练分类头和顶层 Transformer。
轻量化落地是生产环境的刚需。20000 条数据训练的 BERT-base 可能 90% 的效果都能被一个蒸馏后的 TextCNN 继承,而推理延迟从几百毫秒降到几毫秒。先在 BERT 上得到伪标签,再用这些伪标签训练一个简单的分类模型,是成本最低的模型蒸馏路径。
我自己最早跑 BERT 文本分类时,翻车最狠的一次是把 max_length 随手设成 512,20000 条新闻的训练时间直接翻了两倍多,后来用数据体检一看,90% 的文本都在 200 字以内。这个教训让我养成习惯:所有参数都从数据分布推导,而不是从别人的配置里复制。希望这一篇能帮你在 BERT 文本分类这条路上少走几步弯路。
本文还有配套的精品资源,点击获取