
简介一套基于Pytorch的中文文本分类知识蒸馏项目主要面向自然语言处理模型压缩与加速场景核心是将Hugging Face预训练的BERT-base-Chinese作为教师模型把其logits输出中的知识蒸馏到轻量级BiLSTM学生模型上。压缩包共43个文件包含22个Python源码、9个pkl数据或模型文件、5个txt说明、4个json配置以及sh启动脚本等总体积约63.85MB目录明确划分为data、config、models、processor、checkpoints数据加载、模型定义、训练入口与结果保存一目了然。项目使用THUCNews新闻数据集共10类BiLSTM部分以单字为输入并内置约5000字的词汇表配置脚本可分别控制训练、验证、测试、预测流程也方便切换梯度累加、混合精度训练、对抗训练等策略便于对比知识蒸馏与不同训练技巧的效果。多个入口脚本覆盖基础蒸馏、对抗样本训练和混合精度等实验组合适合NLP入门者及希望将大模型能力迁移到轻量级模型上的开发者参考。目前已有312人学习可直接用于中文文本分类模型压缩的复现与扩展。1. 知识蒸馏在中文文本分类上的落地价值从 BERT 到 BiLSTM 的轻量化实践中文文本分类这块BERT 效果好是公认的但真要上线就卡在显存和时延上。知识蒸馏的思路很直接把 BERT 当成老师把它输出的 logits 当作软标签教一个轻量的 BiLSTM 学生模型学会同样的分类能力。这份基于 Pytorch 的中文文本分类蒸馏工程训练主线就是拿bert-base-chinese做 teacher在 THUCNews 的 10 分类任务上蒸馏到 BiLSTM工程里还顺带集成了梯度累加、混合精度训练、对抗训练三个实验开关全部通过 config 控制不用改代码就能跑。适合正在做人工智能课程大作业、毕业设计或者想评估轻量模型上线中文短文本分类可行性的人。看懂这个工程的目录设计和训练链路把数据换成你自己的语料就能复现出一套完整的蒸馏文本分类方案。2. 蒸馏链路拆解温度、软标签与 KD Loss 的落地配置2.1 为什么蒸馏 logits 而不是 hidden states知识蒸馏有两个主流方向一是对齐 BERT 的中间层输出比如 hidden states 或 attention map这种方式对维度对齐要求高teacher 和 student 的隐层尺寸不一致时还要额外加投影层调起来很麻烦二是只对齐最终分类 logits也就是这份工程采用的方式。后者不需要关心 teacher 中间层怎么表示语义只需要在训练时让 student 的输出概率分布去拟合 teacher 的输出概率分布。选 logits 蒸馏的另一个原因是中文文本分类任务的特性类别少、分类头简单信息集中在最后一层特征里。对 10 分类的 THUCNews 来说teacher 的 logits 已经承载了足够的类间区分信息student 学到一个在概率分布上足够接近的结果在验证集上就能拿到不错的准确率。工程里的 teacher 模型用bert-base-chinese结构是 BERT 编码器 分类头输出 10 类 logitsstudent 的 BiLSTM 也是输出 10 类 logits两者在最后一维上天然对齐不需要做额外映射。2.2 温度 T 与软标签的配合这里要理解温度的作用。直接用 teacher 的硬 logits 做训练学生学到的是哪一类得分最高这个结果把 logits 除以温度 T 后再做 softmax分布会被拉平低概率类别的差异也能被学生捕捉到。温度越高分布越平滑知识迁移得越充分但过高的温度会把类别差异全部抹平学生反而学不到东西。我一般把 T 设在 2 到 6 之间先在验证集上用 2 和 4 各跑一轮看 student 的 loss 下降趋势再定。KD Loss 的常见组合是两项加和KL 散度计算 student 和 teacher 在温度缩放后的概率分布差异同时交叉熵计算 student 在真实标签上的损失。KL 项乘上 T 的平方是为了抵消温度缩放带来的梯度尺度变化这一项在代码里忘了乘是新手常犯的错会导致学生学到的东西偏弱。alpha 控制两项权重工程里默认是 0.7 左右也就是七成注意力放在模仿老师、三成保留对真实标签的拟合。2.3 入口脚本与 config 的参数控制工程里有多个入口脚本最核心的是kd_main.py它负责蒸馏主流程main.py可以单独训练某一边的模型。其余入口分别对应附加实验main_with_apex.py开混合精度main_with_gradient_accumulation.py开梯度累加main_with_attack.py开对抗训练。每个入口脚本对应config/目录下的一个配置文件默认蒸馏跑config.py加实验就换成对应的 config。conda activate kd_env # 先检查数据与模型是否存在THUCNews 原始数据放在 data/THUCNews 下 python kd_main.py --config config/config.py # 如果要带对抗训练换成 config_with_attack.py python kd_main.py --config config/config_with_attack.py脚本逻辑上不复杂config.py里集中了全部可调参数训练、验证、测试、预测的模式控制也在里面。换实验只需改--config指向的文件入口脚本本身不用动。config 里的关键参数有三类模型结构参数bert 的 max_len、bilstm 的 hidden_size、dropout、训练参数batch_size、learning_rate、epoch、蒸馏参数temperature、alpha。run.sh里把环境激活、数据路径检查和启动命令串成一条流水线复现时先看这个脚本就能了解作者的启动顺序。参数设置方面BERT 侧和 BiLSTM 侧的学习率一般不共用BERT 微调用 1e-5 到 3e-5BiLSTM 字向量训练用 1e-3 量级。batch_size 受显存限制时方案优先考虑梯度累加而不是强行调小 BERT 的 batch后面第 4 章会讲原因。epoch 不必太多蒸馏过程里 teacher 是冻结的student 一般 5 到 8 轮就能收敛跑太多轮 BiLSTM 反而会对训练集过拟合。3. 数据与处理器THUCNews 双格式输入与 5000 字表的设计3.1 THUCNews 的任务设定工程使用的是 THUCNews 子集共 10 类涵盖体育、娱乐、家居、房产、教育、时尚、时政、游戏、科技、财经。类别之间区分度参差不齐比如科技和游戏的相关性较高像这类容易混淆的类别对蒸馏后的 BiLSTM 来说就是难点所在。数据目录里还能看到 IFLYTEK 数据那是另一个中文长文本分类集类别数更多可以作为蒸馏迁移的备选数据集。数据预处理时processor.py负责最基础的读取和清洗工作比如切分训练集、验证集、测试集统一文本编码格式。kd_processor.py则是在基础处理之上做双格式编码因为蒸馏时需要同时喂两个模型BERT 需要 tokenizer 编出来的 input_ids 和 attention_maskBiLSTM 需要字符级 id 序列。这个双编码逻辑是整个数据处理里最容易出错的地方两个模型对 padding 和截断的容忍度完全不一样。3.2 5000 字词表背后的选型逻辑BiLSTM 侧的输入采用单字切分配一个整理好的 5000 字词汇表。为什么只保留 5000 字中文常用字集中在 3000 到 4000 的范围5000 能覆盖绝大多数文本内容同时把 embedding 矩阵控制在可接受的大小。词表越大embedding 层参数越多BiLSTM 本身轻量化的优势就被削弱了。def build_char_input(text, char2id, max_len128): # 按字切分oov 用 id1 兜底padding 用 id0 char_ids [char2id.get(ch, 1) for ch in text] if len(char_ids) max_len: char_ids char_ids[:max_len] else: char_ids char_ids [0] * (max_len - len(char_ids)) return char_ids这套逻辑里OOV 字符统一映射到 id 1而不是直接丢弃避免某些生僻字导致整句长度错乱。max_len 这里取 128对 THUCNews 这类标题级别的短文本足够如果换成 IFLYTEK 那种长文本建议提到 256 并配合截断策略。词汇表不是一次性生成的而是在整个训练集上统计字频、取前 5000 获得的替换自己的数据时这一步要重新跑。3.3 kd_processor 的双流编码与 DataLoader 配合蒸馏训练时一个 batch 里同时包含两套特征BERT 的 input_ids、attention_mask、token_type_ids以及 BiLSTM 的 char_ids。kd_processor.py的核心逻辑就是从同一段中文文本出发分别走两条编码路径最终输出的 batch 结构同时适配两个模型的 forward。def encode_for_kd(text, tokenizer, char2id, max_len): # BERT 侧用 bert-base-chinese 自带的 tokenizer 编码 bert_enc tokenizer.encode_plus( text, max_lengthmax_len, truncationTrue, paddingmax_length, return_tensorspt ) # BiLSTM 侧按字切分并映射到 5000 词表 char_ids [char2id.get(ch, 1) for ch in text] char_ids (char_ids [0] * max_len)[:max_len] return { bert_input_ids: bert_enc[input_ids], bert_attention_mask: bert_enc[attention_mask], char_ids: torch.tensor(char_ids) }这里有一个值得留意的设计细节两套特征用的 max_len 完全一致为的是在 batch 维度上简单对齐。BERT 的 tokenizer 会把中文文本切出现更多 token 吗不会中文子词切分不改变字符数量级但个别标点符号会合并所以 BERT 分支建议多预留 8 到 16 个 token 的余量。工程里如果同一个超参数同时控制两侧迁移到新数据集时要多留意这一点。4. 模型与三种实验模式BiLSTM 结构、梯度累加、混合精度与对抗训练4.1 三个模型文件的定位差异models/目录下有三个模型文件bertForClassification.py、bilstmForClassification.py、lstmForClassification.py。在蒸馏场景里BERT 是 teacherBiLSTM 是 student单向 LSTM 则是对照组用来评估双向结构在蒸馏任务里的增益。bilstm 模型结构不复杂CharEmbedding 从 5000 词表里查字向量双向 LSTM 编码序列把正反向最后一步的隐状态拼接起来接 Dropout 和线性分类头输出 10 类 logits。注意它对 BERT 侧没有依赖推理时可以完全脱离 transformer 库加载这也是蒸馏的主要收益线上服务不再需要拉起 BERT 模型显存占用和推理时延都会明显下降。4.2 梯度累加显存不够时等价放大 batch size蒸馏时 teacher 是 BERT base输入 batch 太大会直接爆显存但又不想缩小有效 batch size 影响分布稳定性。梯度累加的思路是攒够 n 个 batch 的梯度后统一更新一次参数效果等价于用 n 倍大小的 batch 训练。config_with_gradient_accumulation.py里控制累加步数入口脚本是main_with_gradient_accumulation.py。optimizer.zero_grad() for step, batch in enumerate(dataloader): loss compute_kd_loss(batch) / accumulation_steps loss.backward() if (step 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这里有个细节loss 要除以 accumulation_steps 再做反向传播因为梯度是逐 batch 累加的不归一化的话实际学习率会被放大 n 倍模型会不稳定。累加步数一般设 2 到 8显存紧张时优先调这个参数不要优先调小 BERT 的 batch_size因为过小的 batch 会让 BERT 侧的 BatchNorm 和 attention 统计不稳定蒸馏效果会跟着波动。4.3 混合精度训练如何用 apex 跑 O1 模式混合精度这块用的是 NVIDIA apex入口脚本main_with_apex.py配置在config_with_apex.py。apex 的 O1 模式会把大部分算子切成 FP16少数对精度敏感的算子保持 FP32是性价比最均衡的选项。O2 模式更激进把所有 FP32 网络参数都转 FP16但蒸馏场景下 teacher 和 student 都有 softmax 与 KL 散度计算O2 的数值精度问题会让 loss 波动加大。from apex import amp model, optimizer amp.initialize( modelmodel, optimizeroptimizer, opt_levelO1, loss_scaledynamic )使用混合精度时loss scaling 建议保持 dynamic。蒸馏损失中 KL 项的数值范围本身就比普通分类损失小Fixed loss scale 设置不当会导致下溢dynamic 模式会自动调整缩放因子。另外要注意验证和测试阶段不要开 amp前向推理用完整 FP32 精度否则结果会和训练时的指标有出入。4.4 对抗训练的引入位置与扰动方式main_with_attack.py与utils/attack_utils.py还有一份拼写略有出入的attak_utils.py后面避坑章节会细说把对抗训练接进蒸馏流程。做法上属于经典 FGM 家族在 embedding 层加一个小扰动让模型对微小输入变化更稳健。对抗样本生成后student 在原始样本和对抗样本上的损失都参与梯度计算。对抗训练的扰动幅度要压得很小epsilon 一般控制在 0.01 到 0.05 量级。扰动幅度太大student 学到的特征会被破坏蒸馏训练根本收敛不了。工程里对抗训练是附加实验不是蒸馏主线所以 config 中可以单独控制是否开启攻击模块。建议先跑通不带攻击的蒸馏确认 student 基线稳定后再叠加对抗训练观察泛化能力变化。5. 避坑与排查环境依赖、文件命名与蒸馏收敛的五个典型问题5.1 环境与依赖相关的坑现象直接跑kd_main.py启动时报错找不到bert-base-chinese的预训练权重程序卡在下载阶段迟迟不继续。原因工程使用的是 Hugging Face 上的bert-base-chinese首次运行需要联网下载权重。很多情况下网络环境访问 Hugging Face 不稳定下载会中断或超时而工程本身没有内置权重文件。解决先从可用的镜像站或已下载好权重的机器上把bert-base-chinese目录完整拷下来放到~/.cache/huggingface/transformers/下在kd_processor.py或 config 里显式指定本地路径直接在代码里传入bert-base-chinese的本地绝对路径。这样既不依赖运行时网络也避免每次启动都重复检查缓存。现象安装 apex 后跑main_with_apex.py时报 CUDA error或者 O1 模式下 loss 变成 NaN训练中断。原因apex 与 PyTorch 的版本强绑定版本过新的 PyTorch 对旧版 apex 不兼容FP16 算子内部调用出错。O1 模式把自己的模型包装到了不合适的计算图上问题会在第一个训练 step 暴露出来。解决先把训练入口切回普通模式跑通确认模型和数据处理本身没问题再回来处理 apex。安装 apex 时用源码编译方式并严格按照官方 Readme 指定的 PyTorch 版本对应关系来配。pytorch 与 CUDA 的版本也要对齐nvidia-smi显示的 CUDA 版本和torch.version.cuda不一致时GPU 算子会莫名报出奇怪的形状错误。5.2 数据处理与训练过程相关的坑现象蒸馏训练时每个 step 都报错loss.backward()回传时维度不匹配张量形状一个来自 BERT一个来自 BiLSTM对不上。原因两个模型对同一段文本计算出的序列长度不一致。BERT 的 tokenizer 遇到全角标点时可能合并或拆分 token导致最终序列长度少于按字切分的 char_ids 长度截断时又因为 max_len 不够会分别截到不同位置。解决编码阶段单独打印两边的 shape逐条对比前 3 个样本的 input_ids 和 char_ids 长度。统一把 BERT 分支的 max_len 设得比 char_ids 长比如 char 侧 128、BERT 侧 136padding 后取对齐长度。最省事的方案是处理完立刻断言两边 batch 维一致不通过就停下看是哪一条文本触发的。现象蒸馏 loss 一直下降但 student 在验证集上的准确率迟迟上不去甚至不如完全从头训练的 BiLSTM。原因温度 T 和 alpha 设置不匹配。T 过大时 teacher 的软标签趋近均匀分布student 学不到类间倾向alpha 过大时则完全忽略真实标签student 脱离任务本身。更隐蔽的原因是有些实现把 KL 项该乘的 T^2 漏掉了导致蒸馏项梯度严重偏小student 根本没从 teacher 那里学到东西。解决做三组快速对照温度分别取 1、2、4、8alpha 固定 0.7 跑三个 epoch看 student 在验证集上的准确率峰值。温度取 1 相当于直接用硬 logits 训练这个基线用来判断是否真的学进去了。另外检查一下 KL 项在代码里有没有补上 T^2这步花两分钟确认能省一整天的排查时间。现象from utils.attak_utils import ...这个导入路径在本地环境时好时坏换一台机器克隆项目后就报 ModuleNotFoundError。原因utils/目录里有两份功能几乎一致的代码一份叫attack_utils.py一份叫attak_utils.py原文里的拼写不一致主入口main_with_attack.py引用的到底是哪一份取决于作者当时最后保存的是哪个文件。直接用 IDE 的自动导入补全路径匹配到的可能是另一份代码里却还是旧的名字。解决改代码是长期方案这里有个临时技巧在utils/目录下加一层兼容层把两份名字都暴露出来。更简单的做法是全局搜索import语句统一改成实际存在的文件名。改完后顺手检查utils/__init__.py是否导入了同名的类__init__.py里的导入顺序也会影响最终生效的符号。现象叠加对抗训练后loss 曲线在前面几十个 step 急剧上升然后进入平台期验证集准确率反而比不加攻击低了不少。原因攻击参数 epsilon 设得过大对抗样本已经超出原始样本所在的特征流形student 在强化对扰动鲁棒性的同时失去了对正常文本的判别力另一种常见情况是攻击只在 student 分支做但扰动却加在了 BERT 共享的 embedding 上teacher 也被迫在噪声上学了错误分布。解决先把 epsilon 往小调一个数量级确认 loss 回落了再逐步放大检查attack_utils.py里扰动加在哪个 tensor 上teacher 分支的输入如果也被加了扰动要隔离开。对抗训练作为附加实验应当排在蒸馏主线之后不要一上来就带着攻击调温度参数那会把问题搅在一起很难定位。6. 验证蒸馏效果的进阶手段从单条预测到迁移私有数据6.1 蒸馏是否有效的自检方法训练完成后用一份从前没参与训练的数据做单条推理是快速判断 student 是否学到知识的最直接方式。但要注意不要在边界样本上做一两个案例就下结论正确做法是统计一组同分布样本的预测分布是否接近 teacher 的输出分布而不是只对比预测类别。from models.bilstmForClassification import BiLSTMClassifier model BiLSTMClassifier.load_from_checkpoint(checkpoints/bilstm_distilled.pt) model.eval() text 苹果发布新款手机售价超过万元 inputs build_char_input(text, char2id, max_len128) with torch.no_grad(): logits model(torch.tensor([inputs])) pred int(logits.argmax(dim-1)[0]) print(预测类别:, id_to_label[pred])6.2 迁移到自己的中文数据集时的修改清单先把词表换成目标语料的字频前 5000然后检查标签集合大小把分类头的线性层输出改成对应的类别数第三处理 max_len长文本数据集建议增到 256并验证 BERT 分支和 BiLSTM 分支的序列长度对齐最后跑一遍从kd_main.py到验证的完整流程。迁移后最常被忽略的是 IFLYTEK 数据格式与 THUCNews 不一致字段分隔符和标签编码方式不同需要先调整processor.py的加载逻辑。从那以后我每做一次蒸馏实验都会强制自己先跑一条从零训练的 BiLSTM 基线再跑蒸馏版最后才叠加对抗和混合精度。没有基线的正确率做参照student 蒸馏后到底提升多少就说不清。温度、alpha、epsilon 三个参数每次只调一个记录验证集指标变化这个习惯帮我避免了大量调参翻车希望帮到你。本文还有配套的精品资源点击获取