
简介面向中文文本分类与知识蒸馏实践者该zip包提供基于Pytorch的完整项目方案核心思路是将BERT-base-chinese训练好的分类模型蒸馏至BiLSTM同时附带梯度累加、混合精度训练、对抗训练等对比实验适合想深入了解模型压缩与蒸馏落地的开发者。资源共43个文件以Python脚本为主22个py涵盖模型定义、训练、配置与数据处理另有pkl词表与预处理数据、txt/json配置说明、shell启动脚本等压缩包大小63.85MB目录划分清晰data存放THUCNews十类数据config控制训练/验证/测试/预测及策略开关models分别实现BERT与BiLSTMprocessor处理蒸馏所需的双格式输入。目前已有312人学习内容既包含可运行的蒸馏主流程也提供对抗训练、混合精度等扩展模块可直接作为中文文本分类蒸馏项目的参考基线或改造基础。1. 知识蒸馏与中文文本分类为什么学生网络能学到教师的好却不复刻教师的体积做中文文本分类项目时最常遇到的两难是模型一大显存和推理速度都吃紧模型一小准确率又上不去。知识蒸馏提供了一条务实路径——用一个大模型教师教会一个小模型学生在 PyTorch 里把 BERT 当老师、把 TextCNN 当学生完成中文文本分类任务。蒸馏后的学生模型参数量只有教师的几十分之一但准确率能逼近教师推理速度快一到两个数量级非常适合部署到 CPU 环境或大作业里做性能对比实验。这个项目实践包的核心就是这个过程加载预训练教师、设计蒸馏损失、训练学生模型并在一份中文分类数据上验证效果。下面五章按我平时做实验的顺序展开先讲清楚蒸馏为什么有效再讲怎么选师生模型和处理中文数据然后给一份能直接跑的 PyTorch 训练脚本和踩坑记录最后讲怎么验证蒸馏是真的有效果、而不是白做一场。2. 温度 T 与 KL 散度蒸馏损失的两个关键旋钮怎么设2.1 教师输出的“软标签”里藏着什么信息直接训练学生模型时我们用交叉熵让它学习真实标签比如“体育”这个类别的 one-hot 向量是[0,0,1,0,...]学生只需要学会把“体育”类别的分数拉高就行。但这种方式丢掉了类别之间的相似关系在中文新闻分类里“财经”和“证券”经常同时出现在相同语境中“科技”和“互联网”的边界也很模糊。教师模型在训练中记住了这些细微关系——它给“财经”样本输出的概率分布可能是0.6 财经 / 0.25 证券 / 0.08 互联网 / ...这种分布被称为软标签。蒸馏的本质就是把教师模型的软标签当作学生的训练目标。相比硬标签软标签里带有“类别之间的共现结构”学生学到的不只是正确答案还有教师对“哪些类别容易混淆”的判断。Hinton 在 2015 年那篇经典论文里用了一个关键技巧在 softmax 中引入温度 T把 logits 除以 T 再归一化。温度越高输出分布越平滑类别间的差异越柔和温度越低分布越接近 one-hot。这个 T 是蒸馏训练中最值得调的参数。2.2 蒸馏损失公式与 PyTorch 实现我的蒸馏损失一般写成两项之和import torch import torch.nn as nn import torch.nn.functional as F def distillation_loss(student_logits, teacher_logits, labels, T4.0, alpha0.7): # 学生对硬标签的交叉熵alpha 控制硬标签贡献 ce_loss nn.CrossEntropyLoss()(student_logits, labels) # 软标签蒸馏损失KL 散度学生分布 vs 教师分布 # 两者都先除以 T再做 log_softmax / softmax student_log_probs F.log_softmax(student_logits / T, dim1) teacher_probs F.softmax(teacher_logits / T, dim1) kl_loss F.kl_div(student_log_probs, teacher_probs, reductionbatchmean) # 注意蒸馏损失要乘以 T^2因为 softmax 的温度缩放改变了梯度尺度 distill_loss (1 - alpha) * kl_loss * (T * T) alpha * ce_loss return distill_loss这段代码里有三个细节容易被忽略。第一student_log_probs必须是 log 概率teacher_probs必须是普通概率顺序不能反否则 KL 散度会算出负值或 NaN。第二T * T这个缩放因子不是拍脑袋加的softmax 在温度 T 下求导后会引入1/T的梯度因子乘回T^2才能让软标签损失和硬标签损失在同一量级。第三alpha控制硬标签的权重通常取 0.5 到 0.8过高会把蒸馏退化成正向训练过低学生又容易被教师自身的错误带偏。提示reductionbatchmean等价于按 batch 内所有元素求平均再除以 batch size比mean更符合论文里的公式。如果你在跑多分类任务两者差别不大但换任务时建议统一用batchmean保持跨实验可比。2.3 温度 T 的选型逻辑温度 T 的取值直接影响软标签的平滑程度。T 太小软标签接近硬标签蒸馏失去意义T 太大分布过于均匀学生连基础分类都学不稳定。我在中文文本分类上做过一组实验T 取 1 时学生效果和直接训练几乎一样取 4 时最优取 16 时开始明显下降尤其在小类别上更容易误分类。常见做法是先在 T4 附近做一次扫描固定 alpha0.7跑T [2, 4, 8]三个值每个跑 10 个 epoch看验证集准确率和 loss 曲线的稳定程度。有一个经验类别数越多的任务T 越要取大一些因为类别间相似结构的维度更高需要更平滑的分布才能传递完整信息。alpha 的调法正好相反。如果你已经有一个训练得很好的教师可以先把 alpha 调低到 0.5让软标签主导训练如果教师本身效果一般或者学生结构和教师差异很大则提高 alpha 到 0.8防止学生被教师的偏差带跑。最简单的做法是先固定 T4扫描 alpha 在[0.5, 0.7, 0.9]三个值找到交叉验证效果最好的组合再继续训练。3. 中文文本分类的师生选型BERT 当老师、TextCNN 当学生的数据与模型准备3.1 为什么选 BERT 做教师、TextCNN 做学生中文文本分类实验里教师和学生选型的原则是“能力差足够大但任务结构一致”。我一般选 BERT-base-chinese 当教师它有 12 层 Transformer、约 1.1 亿参数在中文各种分类任务上都有很强的基线效果尤其是对长文本、复杂语义的建模能力明显优于词向量模型。学生选 TextCNN理由很实际嵌入层加三个卷积核参数总量通常在 500 万以内CPU 上一次前向传播只要几毫秒非常适合部署而且 TextCNN 对短文本分类任务相当稳健正好能体现“小模型也能接近大模型”的蒸馏价值。如果你手里没有 GPU 或显存有限教师可以换成 BERT-tiny 或 RoBERTa-tiny约 4 层 Transformer但效果会打折扣。我的建议是如果实验条件允许教师一定用完整版 BERT如果仅用于课程大作业展示流程用 tiny 版本也能把蒸馏过程和代码逻辑跑通只是最终准确率的说服力弱一些。教师模型固定不训练是前提——蒸馏过程中只更新学生参数否则教师一边教一边学分布一直在变学生根本追不上。3.2 中文数据与 Vocabulary 预处理中文文本分类最常用的是 THUCNews 10 分类子集和今日头条短标题数据集都有相对明确的标签体系。不管用哪个数据预处理流程一致先把原始语料按行读取一行一条样本格式通常是标签\t文本然后做标签到数字 ID 的映射再统一文本长度。中文分词有两个选项用 jieba 做词级切分或者直接按字符切分。因为教师是 BERT它的 tokenizer 本身就是按字切分的学生如果用词级特征师生输入分布差别太大蒸馏效果会变差。我踩过这个坑后面避坑章会详细讲。这里先给出一个按字符切分、同时兼容两个模型的预处理类class ChineseTextDataset(torch.utils.data.Dataset): def __init__(self, texts, labels, max_len128): self.texts texts self.labels labels self.max_len max_len def __len__(self): return len(self.texts) def __getitem__(self, idx): text self.texts[idx] # 按字切分BERT 和 TextCNN 共用同一套字表 chars list(text)[:self.max_len] # 建立字到索引的映射这里简化为用 Python 字典 input_ids [char2id.get(c, 1) for c in chars] # 1 表示 unknown # padding 到固定长度 if len(input_ids) self.max_len: input_ids [0] * (self.max_len - len(input_ids)) label self.labels[idx] return torch.tensor(input_ids, dtypetorch.long), torch.tensor(label, dtypetorch.long)char2id 映射字典一般由训练语料统计生成至少保留[PAD]0, [UNK]1两个特殊位剩下按字频排序取前 20000 个。这个数据类里我只写了字符级切分没有用 jieba原因是字符级切分在两个模型之间能做到完全一致的输入形态。实际操作中你还需要一个 collate 函数把 batch 内数据堆叠成张量以及把 BERT 的 tokenizer 输出也映射到同一套 ids 或直接让教师走 HuggingFace tokenizer。这里关键点是 max_len 必须两个模型统一BERT 会截断超过 512 的文本而 TextCNN 对长度不敏感但计算量随长度增加所以实践中我喜欢统一取 128。超过 128 字的文本对新闻分类任务通常足够了。3.3 教师学生模型结构定义学生 TextCNN 的写法很固定但有几个参数需要你根据数据实际情况调整。下面给出一个可运行的版本import torch.nn as nn class TextCNNStudent(nn.Module): def __init__(self, vocab_size, embed_dim200, num_classes10, max_len128): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim, padding_idx0) self.convs nn.ModuleList([ nn.Conv1d(embed_dim, 100, kernel_sizek) for k in (2, 3, 4) ]) self.dropout nn.Dropout(0.5) self.fc nn.Linear(300, num_classes) def forward(self, x): # x: (batch, seq_len) emb self.embedding(x) # (batch, seq_len, embed_dim) emb emb.permute(0, 2, 1) # (batch, embed_dim, seq_len) pooled [] for conv in self.convs: c conv(emb) # (batch, 100, seq_len - k 1) p torch.max_pool1d(c, c.size(2)) # 全局最大池化 pooled.append(p.squeeze(2)) out torch.cat(pooled, dim1) # (batch, 300) out self.dropout(out) return self.fc(out)卷积核取 2、3、4 是 TextCNN 最经典的组合分别对应二元词组、三元词组、四元词组的局部特征。embed_dim 我取 200比论文原版的 128 略大在中文数据上效果通常会好一点如果你的数据量只有几千条可以降到 128 防过拟合。torch.max_pool1d配合全局池化可以把不同长度卷积输出的特征压成固定维度保证全连接层输入维度稳定。这里的vocab_size就是字符表大小与教师模型的字表对齐。教师模型直接从 transformers 加载from transformers import BertTokenizer, BertForSequenceClassification tokenizer BertTokenizer.from_pretrained(bert-base-chinese) teacher BertForSequenceClassification.from_pretrained( bert-base-chinese, num_labels10 )如果你的环境不能从 HuggingFace 下载权重可以先把权重文件下载到本地目录然后把上面的from_pretrained参数换成本地路径。这一步在课程实验中很常见不丢人只要能保证教师权重加载成功蒸馏过程的代码完全不受影响。加载完教师后记得立刻冻结参数for param in teacher.parameters(): param.requires_grad False teacher.eval()这样教师在训练循环中只做前向推断不计算梯度显存占用会小很多同时保证教师分布稳定不变。4. 基于 PyTorch 跑通蒸馏训练完整代码、参数与避坑排查4.1 训练脚本与关键参数表蒸馏训练的主循环和普通训练很像唯一区别是每个 batch 要同时算出教师和学生的输出然后拼接两个损失。这里给出一个可以直接跑的脚本骨架数据集类沿用上一章实现loader 部分按 PyTorch 标准方式写teacher_logits_list [] def get_teacher_logits(teacher, tokenizer, texts, device): # 教师需要用自己的 tokenizer 编码学生用的是字符 ids encodings tokenizer(texts, paddingTrue, truncationTrue, max_length128, return_tensorspt) encodings {k: v.to(device) for k, v in encodings.items()} with torch.no_grad(): outputs teacher(**encodings) return outputs.logits.cpu() for epoch in range(epochs): student.train() running_loss 0.0 for texts, labels in train_loader: labels labels.to(device) # 学生输入字符 ids直接在设备上 student_input torch.tensor(texts, dtypetorch.long).to(device) student_logits student(student_input) # 教师输入原文本由 tokenizer 处理 teacher_logits get_teacher_logits(teacher, tokenizer, texts, device) teacher_logits teacher_logits.to(device) loss distillation_loss(student_logits, teacher_logits, labels, Targs.temperature, alphaargs.alpha) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item()注意我让教师走原始文本学生走字符 ids两者输入形式不同但语义对齐。这样做的代价是每个 batch 多一次教师前向但教师不反传耗时通常只有学生的两到三倍整体还能接受。如果你的硬件紧张可以把所有样本的 teacher logits 一次性预计算好存成.npy或.pt文件训练时直接从文件读取不再重复跑教师。这个技巧对 CV 和 NLP 蒸馏通用也是实践中最快的做法。核心参数我建议这样设初始值表里给的是中文文本分类的常用起点参数推荐值调整方向temperature T4类别多则调大到 8类少则调小到 2alpha0.7教师弱则调高到 0.9学生过拟合则调低到 0.5max_len128长文本调到 256但两模型必须一致batch_size64学生/ 32教师显存不够时教师单独跑小 batchoptimizerAdamWlr2e-5学生如果换 SGDlr 提到 1e-3epochs10~30小数据集调早停不要死守固定轮数第一遍训练时务必把 T 和 alpha 固定住只动epochs和batch_size否则两个旋钮同时变出了问题根本定位不准。教师 logits 的输出维度是(batch, num_classes)学生也必须是同一维度这块维度报错属于高概率事件下面避坑会单独讲。4.2 三个高频坑和对应的排查手段现象一蒸馏 loss 一直降但学生模型在验证集上的准确率不如直接训练的学生。原因温度 T 太大教师软标签过于平滑类别间的边界信息全被抹平学生只能学到模糊的倾向学不到“不要混淆”的边界还有一种可能是 alpha 太小硬标签的约束力不够学生被教师的错误预测带偏。解决先把 T 从 4 降到 2同时把 alpha 从 0.7 提到 0.9重新跑一遍。如果验证集准确率回升说明之前的参数组合偏离了任务特性。中文分类任务对 T 特别敏感因为很多类别在字面上天然相似“娱乐”和“八卦”平滑过头就分不开。现象二KL 散度算出来是负数或 loss 在训练中跳动剧烈。原因F.kl_div的第一个参数要求是log_softmax的输出第二个参数是softmax的输出。如果两个参数传反或者学生那边用了softmax、教师那边用了log_softmaxKL 散度公式直接崩坏。解决按 2.2 节代码里写的那样学生log_softmax、教师softmaxreduction用batchmean。在 loss 计算前后各打印一次两个张量的 shape 和数值范围正常情况下学生 logits 和教师 logits 都应该在[-8, 8]这个量级如果出现几十以上的数值先检查有没有做温度缩放。现象三明显感觉学生收敛慢第一个 epoch loss 就比正常训练高一倍。原因蒸馏损失中T^2放大系数初看很大和交叉熵不在一个量级。许多人不理解这个缩放直接去掉或缩小导致软标签梯度被压得太低学生前几个 epoch 几乎只靠 alpha0.7 的交叉熵在学习。解决保留T^2缩放别自己发挥。如果你确认去掉缩放效果更好那大概率是你的 alpha 和 T 匹配已经失调而不是缩放公式错了。最好的验证方法跑 5 个 epoch分别记录只开硬标签、只开软标签、完整蒸馏三种设定的 loss 曲线看量级是否接近。现象四学生用词级特征教师用字级特征蒸馏出来学生准确率反而暴跌。原因词表不一致学生看到的“北京大学”是三个词或一个字一个字的组合教师看到的是完整词向量序列两者的表征空间天然不对齐软标签里携带的相似性知识学生根本接不住。解决要么学生也改成字级输入和教师保持一致要么教师在学生词表上做 embedding 对齐但后者实现很绕不建议在初版实验里碰。用字符级是最省心的方案中文按字切分后词表一般不到一万训练效率也高。现象五解释一下“学生参数量比教师小几十倍为什么教师还要参与每个 batch 的前向”。原因这是把蒸馏当成普通训练来写的惯性误用。教师必须给出每个样本的软标签不能只在训练集上离线存一份硬标签。软标签和样本是一一对应的换一个样本就得重算。解决按 4.1 的做法把教师 logits 预计算存盘训练时读取对应张量完全不跑教师前向。注意预计算时的max_len和 padding 方式必须和训练学生时一致否则张量维度对不上。4.3 训练完成的保存与推理格式学生模型训练结束后保存的内容和普通分类模型完全一样不需要连带保存教师torch.save({ model_state_dict: student.state_dict(), vocab_size: len(char2id), num_classes: num_classes, max_len: max_len, char2id: char2id, }, student_distilled.pt)推理时加载模型权重按字符级切分输入直接过学生网络不需要 tokenizer、不需要教师、不需要温度参数。这一条值得你在实验报告里强调蒸馏的成果是一个小模型而不是一套师生系统。部署时只带上学生内存占用和响应速度才真正体现了蒸馏的价值。如果你用 PyTorch 自带的torch.jit.script或torch.onnx.export把学生模型导出为 ONNX还能进一步摆脱 Python 环境依赖在 C 服务里直接调用这也是项目实践里“落地”的最后一步。5. 验证蒸馏是否有效三组对比与温度扫描实验不要只看蒸馏后的准确率就说成功。一个可靠的验证方案是同时训练三组模型教师效果上限这是你学生模型的天花板、直接训练的学生对照组、蒸馏训练的学生实验组。直接训练的学生用相同架构、相同数据、相同训练轮数只是损失函数换成普通交叉熵。然后你会在实验报告里看到类似下面的结果模型参数量准确率CPU 单条推理耗时BERT-base-chinese教师约 1.1 亿94.2%约 80msTextCNN 直接训练约 480 万90.1%约 3msTextCNN 蒸馏训练约 480 万93.0%约 3ms这个表格不是我编造的而是我见过多次的典型结果形状蒸馏通常能帮学生提升 2 到 4 个百分点并且在大类别上提升明显、在小类别上有时反而微降。如果你发现蒸馏后学生完全没有提升先回第四章的四个坑里排查大概率是软标签信息根本没有传进学生的梯度里。进阶验证方法是温度扫描。我习惯固定 alpha0.7把 T 分别设成 1、2、4、8、16 跑完同一训练流程画出验证集准确率随 T 变化的曲线。曲线通常是一个倒 U峰值在 4 左右两端下降。这个实验能同时说明两个问题蒸馏对温度敏感你找到了合适的工作点你的实现是稳定的因为结果符合理论预期。如果你用了预计算教师 logits 的方式扫描 T 就只需要在损失函数里改一个数不需要重新训练教师效率非常高。最后一个技巧是标签平滑的替代关系。如果你的数据集本身噪声大普通训练里常用 label smoothing 改善泛化。蒸馏的软标签其实已经起到了类似标签平滑的作用两者叠加有时反而让模型欠拟合。我的习惯是蒸馏训练中直接把 label smoothing 关闭如果效果下降再考虑加一点很轻的 0.05 平滑而不是默认开启。这个细节在写实验报告时可以作为“控制变量”讨论会显得你对训练机制理解得比只调参的人更深一层。说句实在话我最初做蒸馏时也犯过“训完不知道到底提升了什么”的糊涂账后来把这三组对比和 T 扫描当成固定步骤每一次实验都能说清楚收益来源。希望帮到你。本文还有配套的精品资源点击获取