ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Bert/ERNIE中文短文本分类工程包:数据、训练与避坑指南

Bert/ERNIE中文短文本分类工程包:数据、训练与避坑指南 简介面向自然语言处理初学者与文本分类实践者这份资源聚焦中文短文本分类任务提供从 Bert、ERNIE 模型原理到代码实现的一站式方案并额外附带 CNN、RNN、RCNN、DPCNN 等分类网络变体方便对比不同架构的模型效果。压缩包内含 54 个文件以 Python 脚本为主涵盖模型定义、训练评估、数据处理等模块并配有编译文件、文本数据说明、文档及许可文件整体大小仅 6.11MB便于快速下载与离线使用。资源自带 THUCNews 中文数据集可直接用于训练与验证配套文档细致讲解了数据预处理、模型加载、分类器构建、训练优化与评估的完整流程并给出不同模型的调用示例与调参建议。已有 219 人学习下载无论课程设计还是项目预研都能借助这份资源快速搭建实验环境掌握预训练模型在中文短文本分类中的落地应用方法。1. 用 Bert/ERNIE 做中文短文本分类这份带数据集的工程包拆给你看如果你以为中文短文本分类的瓶颈在模型选型那大概率会在数据上翻车。我拆过不少这类项目最后发现 Bert 和 ERNIE 在绝大多数中文短文本任务上的差距远没有网上说的那么玄真正拉开效果差距的是数据集怎么清洗、长度怎么截断、标签怎么映射。这份「使用 BertERNIE 进行中文短文本分类(附数据集).zip」就是把这两条线打包好的完整工程自带一份能够直接训练的中文短文本数据集Bert、ERNIE 两条基线都能一键跑起来训练、评估、预测代码都是开箱即用的。适合两类人一是刚接触预训练模型的 NLP 新手想看到一条最朴素的落地路径二是要在业务里快速出一个文本分类 baseline 的从业者拿它跑通之后换自己的数据就行。下面按数据、训练、评估、排错四层拆开讲。2. 拆包看结构数据、训练、推理三层边界在哪里2.1 工程的目录与职责划分压缩包解开之后常见结构长这样。数据、预训练权重、源码、配置分开换数据或者换模型不用动整条链路。bert-ernie-text-cls/ ├── data/ # 自带数据集 │ ├── train.csv │ ├── dev.csv │ └── test.csv ├── pretrained/ # 预训练模型权重目录 │ ├── bert-base-chinese/ │ └── ernie-1.0-base-zh/ ├── src/ │ ├── preprocess.py # 清洗与标签映射 │ ├── dataset.py # Dataset 与 DataLoader │ ├── train.py # 训练与评估 │ └── predict.py # 单条/批量推理 └── config.py # 超参数集中管理这个分层是这类工程最常见的设计也是我推荐的边界划分方式。data 目录只管原始数据src/dataset.py 负责把文本和标签转成模型能吃的张量train.py 只关心训练循环predict.py 独立承担推理。好处很直接你要换成自己的业务数据时只需要保证新数据格式和 train.csv 一致再在 preprocess.py 里改标签映射逻辑训练层基本不用碰。config.py 把超参数集中放是有道理的。短文本分类要调的东西说多不多说少不少max_len、batch_size、learning_rate、epochs、early stopping 的 patience全部集中在一个文件里跑对比实验时改一处就行。很多新手喜欢把参数散落在各个脚本里结果调完一个参数忘了另一个复现结果全靠玄学。2.2 先跑通默认配置最小训练闭环拿到工程的第一步我建议先不动任何配置把自带的默认流程跑通。# 用默认配置跑 Bert数据与模型路径都在 config.py 里写死 python src/train.py --model bert # 跑通之后换成 ERNIE 再出一版 python src/train.py --model ernietrain.py 的 --model 参数控制在两个预训练模型之间切换默认值一般是 bert。第一次跑的时候不需要理解每一行代码只需要确认三件事训练循环正常启动、loss 在逐步下降、dev 集的评估指标打印出来。数据集比较小的话纯 CPU 也能跑几十条到几百条的短文本训练几分钟就能出结果如果数据量上了万建议直接上 GPU否则一个 epoch 就要等很久。跑通训练之后顺手用一下推理脚本保证保存下来的模型真的能加载python src/predict.py --text 发票什么时候能开能正常打印出类别说明最小闭环已经通了。后面改数据、改模型、调参数都在这条闭环上迭代不会出现改了一处代码整个流程跑不起来的情况。2.3 数据集的构成与预处理动手前先看分布自带数据集的格式一般不会太复杂最常见的是两列text 和 label。第一件事不是急着写模型而是先把数据分布摸清楚。head -5 data/train.csvtext,label 这家店的售后客服响应很快,服务态度 发票什么时候能开,发票问题 产品包装破损严重,物流问题 退货怎么操作,售后服务文本是短文本标签是类别名类别数量不多。接着看文本长度的分布这一步直接决定 max_len 设多少。python -c import pandas as pd df pd.read_csv(data/train.csv) print(df[text].str.len().describe(percentiles[.5,.9,.99])) count 20000.000000 mean 28.354000 std 19.802039 50% 24.000000 90% 51.000000 99% 79.000000 max 189.000000这类输出是决定超参数最重要的依据。p90 是 51意味着 90% 的样本在 51 个字以内如果你的 max_len 设成 32超过一成的样本会被硬生生截掉关键信息丢了模型当然学不好。我一般的做法是取 p90 到 p99 之间的值这份数据 64 就是合理的选择既不需要 128 浪费显存也不会截掉太多有效内容。标签那一列也要检查是否有空格、换行符混进去类别名带着 \n 会直接导致后续 label2id 映射出现莫名其妙的错误。3. 把中文短文本喂给预训练模型Tokenizer、长度与数据管道3.1 短文本分类为什么绕不开字粒度中文和英文最大的区别是没有天然的空格分隔分词本身就是一层误差来源。Bert 官方的中文模型 bert-base-chinese 采用的是以字为主的词表绝大多数情况下一个汉字就是一个 token这对短文本尤其友好。ERNIE 走的也是字级词表和 Bert 的差异在预训练阶段的 mask 策略。Bert 随机 mask 单字而 ERNIE 会 mask 短语和实体相当于在预训练时就让模型见过「开发票」「退货流程」这类完整语义单元。所以 ERNIE 在新闻标题、搜索 query 这类含专有名词多的短文本上往往比标准 Bert 更容易抓住关键信息。实践里也确实是这么回事Bert 更强调通用语义理解ERNIE 对实体类表达更敏感。这不是说 ERNIE 一定更好而是告诉你做对比实验时这两条线各有价值不能只跑一个就下结论。短文本本身的困境在于长度短、上下文少、口语化严重、语义密度高。模型可用的线索就那么几十个字预处理阶段每丢掉一个有效 token都是在削减模型的判断依据。3.2 截断与 Paddingmax_len 的临界点怎么找准备好了数据集下一步是把文本转成 input_ids 和 attention_mask。这里最关键的参数就是 max_len。from transformers import BertTokenizer # 使用工程里本地化的权重目录避免训练时再去远程拉取 tokenizer BertTokenizer.from_pretrained(pretrained/bert-base-chinese) texts [ 发票什么时候能开, 这家店的售后客服响应很快, ] enc tokenizer( texts, max_length64, # 由长度分布决定常见取值 64 或 128 truncationTrue, # 超过 max_len 的部分截断 paddingmax_length, # 不足 max_len 的部分补 [PAD] return_tensorspt, # 返回 PyTorch 张量 ) print(enc[input_ids].shape) # torch.Size([2, 64]) print(enc[attention_mask][0]) # [1,1,...,0,0]这个调用的四个参数每一个都有明确作用。truncationTrue 保证超长文本被截到 64不会因为样本太长导致 batch 内 tensor 形状不一致。paddingmax_length 让所有样本都补到同一个长度这样 DataLoader 在组 batch 时不需要额外写 collate_fn。return_tensorspt 直接返回 torch tensor省掉手动转换。要注意的是输入序列里自动加入了 [CLS] 和 [SEP]所以真实可用的文本长度是 max_len 减 2。max_len64 实际给文本留了 62 个 token 的位置。这也是为什么前面强调用 p90 来定长度——把 padding 占位也算进去62 个 token 足够覆盖这份数据 90% 的样本。ERNIE 的 tokenizer 也是同一套接口只需要把加载路径换成 ernie 权重目录。两者的 vocab.txt 差异不影响代码层面的理解。3.3 从 CSV 到 Dataset分层抽样与标签映射文本转 token 之后还要把原始 CSV 包装成 PyTorch 的 Dataset。这里面最容易出错的是标签处理和数据集划分我一般会这样写import pandas as pd from sklearn.model_selection import train_test_split df pd.read_csv(data/train.csv, encodingutf-8-sig) df[text] df[text].astype(str).str.strip() df[label] df[label].astype(str).str.strip() # 类别排序后固定映射顺序保证训练和推理用同一套 labels sorted(df[label].unique()) label2id {label: i for i, label in enumerate(labels)} id2label {i: label for label, i in label2id.items()} X df[text].tolist() y df[label].map(label2id).tolist() # stratify 按类别比例抽样避免划分后某个类别在 dev 集里消失 train_texts, dev_texts, train_labels, dev_labels train_test_split( X, y, test_size0.1, stratifyy, random_state42 )encodingutf-8-sig 是处理从 Excel 导出的 CSV 时的习惯能自动吃掉开头的 BOM 头。text 和 label 都做了 strip这是为了防 \xa0、\u3000 这类不可见字符混进标签名。stratifyy 必须加短文本分类的标签分布往往不均衡不做分层抽样的话小类别很可能整个划分进训练集dev 集里一个样本都没有评估指标虚高却不代表真实水平。Dataset 类的定义也比较固定每一行文本做一次 tokenizer 调用import torch from torch.utils.data import Dataset class TextClsDataset(Dataset): def __init__(self, texts, labels, tokenizer, max_len): self.texts texts self.labels labels self.tokenizer tokenizer self.max_len max_len def __len__(self): return len(self.texts) def __getitem__(self, idx): enc self.tokenizer( self.texts[idx], max_lengthself.max_len, truncationTrue, paddingmax_length, return_tensorspt, ) return { input_ids: enc[input_ids].squeeze(0), # [max_len] attention_mask: enc[attention_mask].squeeze(0), labels: torch.tensor(self.labels[idx], dtypetorch.long), }单条文本经过 tokenizer 返回的形状是[1, max_len]squeeze(0) 去掉 batch 维变成[max_len]这样 DataLoader 在组 batch 时自然会叠成[batch_size, max_len]。DataLoader 不需要额外 collate_fn因为 paddingmax_length 已经把所有样本拉齐了。这种做法在样本量不大、max_len 不超过 128 的情况下非常省事显存稍微多占一点但代码可读性高很多。4. BERT 与 ERNIE 的中文模型实操加载、训练与超参数调优4.1 模型加载BERT 和 ERNIE 为什么不冲突模型加载这一步短文本分类的标准做法是直接用 transformers 的序列分类接口。BERT 和 ERNIE 在代码层面几乎一致因为 ERNIE 的 transformer 结构本身就是从 Bert 演化来的。from transformers import BertForSequenceClassification num_labels len(label2id) model BertForSequenceClassification.from_pretrained( pretrained/bert-base-chinese, # 本地权重目录不依赖远程 num_labelsnum_labels, )ERNIE 的加载方式也一样只是路径换成 ERNIE 权重目录。使用这个接口的前提是 ERNIE 权重已经被整理成 transformers 能识别的格式目录下需要 vocab.txt、config.json、pytorch_model.bin 三个文件。常见做法是跑一次转换脚本把百度原始 checkpoint 映射到 Bert 结构上之后加载时直接 from_pretrained 这个目录就行。模型对比上两个 base 模型的规模基本相当模型层数隐层维度词表特点对短文本的倾向bert-base-chinese12768以字为主通用语义强稳定适合通用场景ernie-1.0-base-zh12768预训练加入短语/实体 mask对实体、品牌词更敏感num_labels 必须和 label2id 的长度一致。这个参数决定模型分类头的输出维度如果训练时是 10 个类别推理时却用默认的 2加载就会报形状不匹配。4.2 训练循环2e-5 到 3e-5 之间的安全区训练循环是整个工程里最需要细心的一段核心超参数就那么几个但每一个都能直接影响结果。from transformers import AdamW, get_linear_schedule_with_warmup optimizer AdamW(model.parameters(), lr2e-5) total_steps len(train_loader) * epochs scheduler get_linear_schedule_with_warmup( optimizer, num_warmup_stepsint(0.1 * total_steps), num_training_stepstotal_steps, ) best_f1 0.0 patience 0 for epoch in range(epochs): model.train() for batch in train_loader: optimizer.zero_grad() outputs model( input_idsbatch[input_ids].to(device), attention_maskbatch[attention_mask].to(device), labelsbatch[labels].to(device), ) outputs.loss.backward() optimizer.step() scheduler.step() # 每个 epoch 结束在 dev 上评估做早停 dev_f1 evaluate(model, dev_loader) if dev_f1 best_f1: best_f1 dev_f1 torch.save(model.state_dict(), checkpoint/model.pt) patience 0 else: patience 1 if patience 2: print(fearly stop at epoch {epoch}) break中文短文本分类的微调学习率我一般锁定在 2e-5 到 3e-5 之间。预训练模型已经收敛过一轮再拿大的学习率去微调很容易把学到的语义知识冲掉表现就是训练集 loss 降得很快、dev 集指标不涨反跌。warmup 比例 0.1 的含义是前 10% 的 step 里学习率从零线性升到 2e-5这样能避免模型在最开始用大步长把参数推到不理想的位置。early stopping 的 patience 设在 2意思是连续两个 epoch dev 指标没有刷新就停防止在小数据集上过拟合。保存模型时只保存 model.state_dict() 而不是整个 trainer 对象这个习惯能省掉很多坑。checkpoint 里只有参数和结构信息加载时和新的 num_labels 配置对齐即可。如果一个 epoch 的训练数据量很大batch_size 又受显存限制上不去可以保留参数更新频率不变、累积几步再更新一次梯度这是常见的妥协方案。4.3 评估短文本分类不要只看 Accuracy短文本分类数据集的标签分布经常是偏的某一个类别可能占了一半样本。这时候 accuracy 会非常具有欺骗性看起来 0.92 很高实际上小类别全部没预测出来。from sklearn.metrics import classification_report, confusion_matrix print(classification_report(all_labels, all_preds, target_namesall_classes, digits3))classification_report 会输出每个类别的 precision、recall、f1-score以及整体的 macro avg 和 weighted avg。我判断一个模型能不能上线主要看 macro-F1它把每个类别的权重拉平了某个类别被忽略时分数立刻跳水。weighted-F1 则更接近业务真实分布适合按样本量评估整体表现。只看 accuracy 的教训我吃过一次后面会在避坑章里展开。另外训练结束可以顺手看一眼混淆矩阵哪两个类别互相混最多一目了然。比如「发票问题」和「售后服务」经常混说明这两个类别在语义上本来就有重叠需要回去看数据标注口径是否清晰。5. 避坑指南中文短文本分类里最常翻车的五个点5.1 卡在权重下载离线加载与词表编码的坑现象第一次跑训练脚本日志停在一行 Downloading... 上几十分钟不动最后报连接超时。第二次跑干脆连不上整个流程卡死在模型加载阶段。原因transformers 的 from_pretrained 默认会去远程仓库拉权重。网络环境一波动大文件下载很容易中断而且中断之后没有断点续传反复重试浪费时间。解决提前在有条件的环境里把 bert-base-chinese 和 ernie 权重完整下载到本地再复制到工程的 pretrained 目录训练时用本地路径加载。判断是否加载成功可以看这条日志有没有出现python -c from transformers import BertTokenizer, BertForSequenceClassification m BertForSequenceClassification.from_pretrained(pretrained/bert-base-chinese, num_labels2) print(local load ok) 跑出 local load ok 再进训练流程。从那以后我拿到任何预训练模型工程第一件事永远是先把权重目录备齐绝不把「在线拉取」当作理所当然。现象数据集读进来之后类别名打印出来带 \xa0 和 \u3000或者第一个类别名前面多了看不见的字符导致 label2id 映射对不上训练时直接报 KeyError。原因CSV 文件编码混杂常见的是 UTF-8 带 BOM、GBK、UTF-16 混用加上文本里掺了全角空格和不可见控制符。解决读取时统一用 encodingutf-8-sig文本列和标签列都做 strip再补一个正则把 \xa0、\u3000 全部替换成普通空格import re def clean_text(s): s s.replace(\xa0, ).replace(\u3000, ) s re.sub(r\s, , s) return s.strip()这个清洗函数放在 preprocess.py 里所有文本进入 tokenizer 之前先过一遍。短文本分类的数据本来就是几十个字的规模清洗成本很低但漏掉一个全角空格就可能让某个类别的文本全部变成异常输入。5.2 长度、标签不均衡与 checkpoint 加载的三个典型翻车现象同一份数据max_len32 跑出来的 F1 比 max_len64 低三到五个点反过来设成 128效果也没有提升训练时间却翻倍。原因短文本的平均长度只有 24但 p90 在 51 左右。max_len32 时超过一成的样本被截掉了尾部关键信息比如「申请发票需要提供订单号和收件邮箱」这类句子后半段全是关键内容。max_len128 又引入了过多 padding小数据集上 padding 比例太高会让模型学到无意义的 [PAD] 模式。解决按 2.3 节的方法统计长度分布取 p90 到 p99 之间的值再用 train/dev 各跑一次对比以 dev 指标为准选长度。短文本分类没有固定最优长度只有「和你的数据分布匹配」的长度。我后来做任何数据集的第一版都会先打印长度分布把 p50、p90、p99 三个数写进实验记录防止后续调参时忘了当初为什么选 64。现象训练日志显示 accuracy 0.91dev 集上看起来一切正常但打开分类报告发现某个小类别的 F1 是 0.0所有该类样本全被预测成了大类。原因类别分布极端模型用交叉熵损失优化时把样本量大的类别学得越来越好小类别因为出现次数太少梯度贡献被淹没整个类别被模型忽略。解决两个手段配合。一是在损失函数上做类别加权用 sklearn 计算权重from sklearn.utils.class_weight import compute_class_weight import numpy as np class_weights compute_class_weight( class_weightbalanced, classesnp.array(list(label2id.values())), ynp.array(y), )然后把 class_weights 传给损失函数让训练时少数类的错误产生更大的梯度。二是评估指标不要只看 accuracy以 macro-F1 为准。0.91 accuracy 加 0.55 macro-F1 的组合说明模型在大类上过拟合严重这种模型在真实业务里基本不可用因为小类别往往才是用户投诉最严重的部分。现象训练好的模型预测阶段加载报错提示 state_dict 中的键不匹配某个 tensor 形状对不上。原因保存模型时把 optimizer 的 state_dict 也一并存进去了或者加载时 num_labels 与训练时不一致。optimizer 里包含每个参数的动量等额外信息和模型结构不是一一对应的num_labels 不一致则直接改变分类头的形状。解决训练循环里只保存 model.state_dict()并同时保存 config 信息torch.save( { model_state: model.state_dict(), num_labels: num_labels, id2label: id2label, }, checkpoint/model.pt, )加载时先从 checkpoint 里读出 num_labels 和 id2label再实例化模型保证结构完全对齐。从那以后我每次保存模型都会把 id2label 一起存进去推理脚本只用这一个文件就能还原完整分类体系不会再出现「预测输出一个数字编号但不知道对应哪个类别」的尴尬。6. 验证模型真的能用混淆矩阵、Bad Case 分析与五折对比单次划分数据集跑出来的指标有随机性。短文本数据量小、分布偏换一次随机种子结果可能波动一到两个点。所以验证阶段我的固定动作是先看混淆矩阵找混淆密集的类别对再做 bad case 抽样逐条读原文最后用五折交叉验证给出稳定的对比结论。bad case 抽样脚本一般这么写每次从预测错误的样本里随机抽 20 条打印出来import random wrong [ (true, pred, text) for true, pred, text in zip(all_labels, all_preds, all_texts) if true ! pred ] sample random.sample(wrong, min(20, len(wrong))) for true, pred, text in sample: print(f真实: {id2label[true]:8s} 预测: {id2label[pred]:8s} 文本: {text})逐条读这 20 条比看任何指标都直观。有一次我以为模型效果不错抽出来却发现「投诉催促」全部被分到「咨询」原因是两条文本里都有「什么时候能好」这种相似表达模型学会了抓表面词没学会区分语气和诉求。这种问题只能靠人眼从 bad case 里看出来指标数字不会告诉你。交叉验证用 StratifiedKFold 就能做五折对短文本场景足够稳定from sklearn.model_selection import StratifiedKFold skf StratifiedKFold(n_splits5, shuffleTrue, random_state42)每一折独立训练、独立评估最后取 macro-F1 的均值和标准差。Bert 和 ERNIE 各跑五折对比表格大概长这样模型Fold1Fold2Fold3Fold4Fold5均值Bert0.9120.9050.9180.9010.9100.909ERNIE0.9210.9170.9250.9110.9190.919标准差超过一个点说明模型对数据划分敏感这时优先考虑回去查标签噪声和类别分布而不是继续堆模型复杂度。有一回我拿着 0.93 的 accuracy 去汇报结果 bad case 里全是把「退货怎么操作」分到「物流问题」的从那以后每次训练完我都强制走一遍混淆矩阵加 bad case 抽样再决定要不要把这个模型交出去。这份带数据集的工程包正好可以让你把这条验证链路完整跑一遍Bert、ERNIE 两条线的对比也能直接出数希望帮到你。本文还有配套的精品资源点击获取
返回列表