多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

BERT知识蒸馏实战:把BERT压成BiLSTM的中文文本分类方案

BERT知识蒸馏实战:把BERT压成BiLSTM的中文文本分类方案 简介这是一套基于Pytorch实现的中文文本分类知识蒸馏项目面向有一定深度学习基础、希望将大规模预训练模型压缩至轻量级模型的开发者和研究者。项目核心是将Hugging Face的bert-base-chinese蒸馏到BiLSTM上通过迁移教师模型logits中的知识让轻量模型在保持精度的同时大幅降低推理成本。资源包含43个文件核心为22个Python脚本覆盖模型定义、训练/验证/测试、数据处理与配置管理9个pkl为预处理数据或词表5个txt含说明或占位文件4个json为配置另有shell脚本便于一键运行整体压缩包63.85MB。除基础蒸馏外还附带梯度累加、混合精度训练、对抗训练等对比实验并提供了完整目录结构data使用THUCNews十类数据models存放bert与bilstm代码processor负责格式转换checkpoints保存模型。目前已有312人学习适合想系统掌握知识蒸馏实战、并扩展训练技巧的读者。1. 基于 Pytorch 的知识蒸馏实战把 BERT 压成 BiLSTM中文文本分类不掉点知识蒸馏这两年早就不只是论文里的概念了工程上最常见的诉求就是「把 BERT 的能力塞进一个小模型」。这个项目实践恰好就是干这件事的用 Hugging Face 上的 bert-base-chinese 训练一个中文文本分类模型然后把它的 logits 知识蒸馏到一个 BiLSTM 上。数据集用的是 THUCNews共 10 类。整个项目代码结构清晰蒸馏主流程、梯度累加、混合精度apex、对抗训练这些实验都给你分开写了配置文件想复现哪条路线直接改 config 即可。适合两类人一类是刚入门知识蒸馏、想跑通一条完整 baseline 的另一类是想在工程里落地小模型但希望尽量保住大模型精度的。2. 知识蒸馏的原理与项目架构为什么偏偏选 BiLSTM 当学生模型2.1 蒸馏的本质让学生的 logits 去拟合老师的 logits知识蒸馏的核心思路说白了就是「抄答案」——老师模型BERT在 Softmax 之前输出的 logits 向量里不光有正确类别的信息还有错误类别的相对概率关系。比如一条新闻BERT 可能给「体育」打了 8.2 分、给「娱乐」打了 3.1 分这个 8.2 和 3.1 的差距本身就是一种知识它告诉学生模型「这两个类别在语义上有点接近」。项目里蒸馏的目标函数由两部分组成一部分是学生模型BiLSTM在真实标签上的交叉熵损失另一部分是学生 logits 和老师 logits 之间的 KL 散度。中间用温度参数 T 来软化概率分布T 越大分布越平滑类别间的相对关系暴露得越充分。常见做法是 T 取 2 到 8 之间这个项目默认配置里用的是 2后面你想调大观察效果可以改 config。损失函数的一个常见写法如下def distillation_loss(student_logits, teacher_logits, labels, T2.0, alpha0.7): # student_logits / teacher_logits: [batch_size, num_classes] # labels: [batch_size]真实标签 soft_teacher F.softmax(teacher_logits / T, dim-1) soft_student F.log_softmax(student_logits / T, dim-1) kd_loss F.kl_div(soft_student, soft_teacher, reductionbatchmean) * (T * T) ce_loss F.cross_entropy(student_logits, labels) return alpha * ce_loss (1.0 - alpha) * kd_loss这段代码里alpha控制的是真实标签损失和蒸馏损失的配比0.7 表示更依赖真实标签蒸馏信号作为正则。T * T是 KL 散度对温度梯度的补偿因为软化后的 logits 数值变小了梯度也跟着变小乘回去能保证训练步长不缩水。2.2 学生模型为什么是 BiLSTM 而不是别的选 BiLSTM 当学生模型有几个现实理由。第一中文文本分类里单字输入 BiLSTM 的组合在短文本上依然能打Its not the most advanced architecture, but its extremely stable. 第二推理速度优势明显BERT 在 CPU 上跑一条样本可能要几十毫秒BiLSTM 能把延迟压到几毫秒以内这在线上服务里是质变。第三这个项目刻意用了单字输入配合一个整理好的 5000 字词表既规避了分词器的依赖也让模型更轻。BiLSTM 模型定义在主目录的models/bilstmForClassification.py里核心结构就是一个双向 LSTM 接一个全连接分类头。实际训练时要注意的是输入格式BERT 需要 token_type_ids 和 attention_maskBiLSTM 只需要把每个字映射成词表里的索引然后做 padding。2.3 项目目录结构六个模块各管一摊这块直接看目录结构就能明白作者的设计思路。config目录下四个配置文件分别对应基础训练、apex 混合精度、对抗训练、梯度累加四条实验线models目录下有 BERT、LSTM、BiLSTM 三个模型文件processor目录负责数据格式转换BERT 和 BiLSTM 各自的预处理逻辑是分开的utils目录里的attack_utils.py是给对抗训练准备的main.py是标准蒸馏主入口main_with_apex.py、main_with_attack.py、main_with_gradient_accumulation.py是三个变体。3. 数据与预处理BERT 和 BiLSTM 的输入格式怎么统一3.1 THUCNews 十分类任务的数据加载项目用的是 THUCNews 数据集10 个类别数据目录下有原始文本和标签。加载的时候核心是把文本转成模型能吃的 ID 序列。BERT 侧直接用BertTokenizer处理BiLSTM 侧走的是自定义词表映射。# processor/kd_processor.py 核心逻辑示意 def encode_for_bilstm(text, word2idx, max_len128): # 按单字切分中文不需要分词器 tokens list(text.strip().replace( , ))[:max_len] ids [word2idx.get(w, 1) for w in tokens] # 1 是 UNK 的索引 mask [1] * len(ids) # padding 到固定长度 ids [0] * (max_len - len(ids)) mask [0] * (max_len - len(mask)) return torch.tensor(ids), torch.tensor(mask)这里的max_len128是经验值THUCNews 的新闻文本普遍不长128 个字足够覆盖绝大多数样本。如果跑其他长文本数据集这个值要按长度分布去调而不是盲目加大——BiLSTM 对长序列的梯度传播会有衰减硬拉长度反而可能掉点。BERT 侧的编码直接用tokenizer.encode_plus就行返回input_ids、token_type_ids、attention_mask三个字段。蒸馏时一个 batch 需要同时拿到老师BERT和学生BiLSTM的输入所以 processor 里会在同一个 batch 内并行处理两份特征这算是知识蒸馏工程实现里比较典型的一个设计点。3.2 5000 字词表一个小而实用的细节项目特意强调「整理好的 5000 字的词汇表」这其实是针对中文场景的一个优化。BERT 的词表是 2 万多个 WordPiece 片段而 BiLSTM 用单字输入时常用汉字也就几千个。5000 字的覆盖规模在 THUCNews 上能覆盖 95% 以上的输入字符剩下的用 UNK 兜底。这个设计的直接收益是模型参数量的缩减词表 5000、embedding 维度 128 的话embedding 层参数才 64 万。对比 BERT 的 embedding 层动辄上千万参数学生模型的体积优势非常明显。训练时如果遇到大量 UNK说明词表覆盖不够需要把训练集里出现频率最高的字重新统计一遍。3.3 老师 logits 的缓存策略蒸馏训练有个效率问题如果每个 epoch 都让 BERT 重新前向一遍训练时间会翻好几倍。常见做法是先把训练集和验证集全部过一遍 BERT把 logits 存成文件或内存张量之后训练 BiLSTM 的时候直接读缓存。# 伪代码先离线生成 teacher logits teacher_model.eval() all_logits [] with torch.no_grad(): for batch in teacher_dataloader: logits teacher_model(**batch) all_logits.append(logits.cpu()) torch.save(torch.cat(all_logits, dim0), teacher_logits.pt)离线缓存 logits 的做法虽然会占一些磁盘空间10 万条样本、10 类float32 大概是 40MB但换来的是蒸馏训练时不再需要加载 BERT 模型。这个项目的主流程里是实时跑 BERT 的如果你的训练集很大建议改成缓存模式能省不少时间。4. 主训练流程与四个变体梯度累加、混合精度、对抗训练怎么选4.1 标准蒸馏主流程kd_main.py 跑通 baseline项目的标准入口是kd_main.py流程可以拆成四步加载配置、初始化两个模型和 optimizer、循环训练、每个 epoch 结束跑验证集。核心逻辑如下for epoch in range(config.epochs): for batch in train_dataloader: bert_inputs {k: v.cuda() for k, v in batch[bert].items()} lstm_inputs {k: v.cuda() for k, v in batch[lstm].items()} labels batch[label].cuda() with torch.no_grad(): teacher_logits teacher_model(**bert_inputs) student_logits student_model(lstm_inputs[input_ids], lstm_inputs[mask]) loss distillation_loss(student_logits, teacher_logits, labels) loss.backward() optimizer.step() scheduler.step() optimizer.zero_grad()这里最关键的一点是老师模型必须挂在torch.no_grad()下因为 BERT 的参数不参与梯度更新。学生模型的输入只有input_ids和mask没有 token_type_ids因为 BiLSTM 不需要区分句子对。4.2 梯度累加小 batch 训练大模型的后悔药main_with_gradient_accumulation.py解决的是显存不够的问题。BERT 哪怕只是跑前向显存占用也不小如果 GPU 只有 6GB一个 batch 塞 16 条样本可能就爆了。梯度累加的思路是把一个大的 batch 拆成多个 micro batch梯度攒够了再统一更新参数。accumulation_steps 4 optimizer.zero_grad() for step, batch in enumerate(train_dataloader): loss compute_loss(batch) / accumulation_steps loss.backward() if (step 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()注意loss / accumulation_steps这步如果不除梯度就会变成原来的 accumulation_steps 倍学习率等于被放大模型大概率直接发散。这是新手最容易踩的坑之一。4.3 混合精度APEZ训练速度与显存的双赢main_with_apex.py用的是 NVIDIA 的 APEX 库做混合精度训练。原理是让一部分操作走 FP16、一部分走 FP32减少显存占用和计算时间。核心加入的代码就几行from apex import amp model, optimizer amp.initialize(student_model, optimizer, opt_levelO1) with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() optimizer.step()opt_levelO1是推荐起点它会在保持数值稳定的前提下尽量用 FP16 加速。如果你用的是新版 Pytorch也可以直接用torch.cuda.amp效果等价接口更原生。混合精度在 BiLSTM 这种小模型上收益没有大模型明显但如果你要蒸馏的是更大的学生模型这个开关值得常驻。4.4 对抗训练给模型加一层鲁棒性main_with_attack.py和utils/attack_utils.py实现的是基于 FGSMFast Gradient Sign Method的对抗训练。做法是给 embedding 加上一个小的扰动让模型在「被攻击」的情况下依然能正确分类。扰动计算的核心逻辑在 attack_utils 里# 用 FGSM 生成对抗扰动 embedding_grad torch.autograd.grad(loss, embedding, retain_graphTrue)[0] perturbation config.epsilon * torch.sign(embedding_grad.detach()) embedding_adv embedding perturbation这里的epsilon控制扰动幅度项目里一般设 0.5 到 1.0。epsilon 太大扰动直接破坏语义模型训练不起来太小则起不到正则效果。对抗训练对蒸馏的实际收益是提升学生模型的稳定性特别是在输入有轻微噪声的场景下掉点幅度会更小。5. 避坑与常见问题排查蒸馏训练中我踩过的四个典型坑5.1 温度 T 和 alpha 同时调大损失直接 NaN现象训练几步之后 loss 变成 NaN模型输出全是一个固定向量。原因T设成 10 以上KL 散度乘上T * T之后梯度爆炸alpha设成 0.9 以上时交叉熵权重过高学生模型拟合噪声。解决T 控制在 2 到 6 之间alpha 控制在 0.5 到 0.8 之间。如果必须用大 T需要同步调小学习率。我一般会在训练启动后打印前几个 step 的 loss 值做 sanity check超过 15 基本就是温度或 alpha 设置出了问题。5.2 BERT 的 [CLS] 向量和 BiLSTM 的最后一层输出维度对不上现象拼接蒸馏损失时报维度不匹配的错误。原因BERT 的分类头输出的是[batch_size, num_classes]BiLSTM 的最后一个 hidden state 经过全连接层后也是[batch_size, num_classes]正常不会出问题。但如果改了models/bertForClassification.py里的num_labels而没同步改 BiLSTM 的分类头两边数字就会不一致。解决修改类别数时同步检查三个模型文件里的num_labels参数确保全部一致。5.3 离线缓存 teacher logits 时用了训练模式的 BERT现象蒸馏效果奇差学生模型的准确率跟在瞎猜一样。原因BERT 里有 Dropout训练模式下会随机丢弃部分神经元生成的 logits 带有随机性。用这种 logits 当老师等于每次教给学生的答案都不一样学生模型直接被教懵。解决缓存 logits 时务必加model.eval()并且包在torch.no_grad()里。这一个坑的翻车概率极高我认识的人里至少有一半在这上面栽过。5.4 词表里没有覆盖的字符全部映射到 UNK导致新闻分类准确率暴跌现象训练 loss 正常下降但验证集准确率只有 60% 出头。原因THUCNews 里有很多标点符号和数字如果 5000 词表里没收录这些字符全部变成 UNK。新闻文本里的数字和标点往往有语义信息比如「5G」「2024」全变 UNK 等于信息丢失。解决构建词表时把标点、数字、常见英文单词都单独收进去。更稳妥的做法是在训练前统计一遍训练集的字符频率取 Top 5000 而不是直接用一个固定词表。6. 进阶验证如何确认蒸馏真的学到了老师的行为先说明一个判断逻辑学生模型测试集准确率高不代表蒸馏成功。准确率只能说明「分类结果接近」不能说明「模型行为接近」。要验证蒸馏质量需要对比学生和老师在 logits 层面的分布一致性。我最常做的一个验证是随机抽 500 条验证集样本分别用老师和学生跑出 logits计算两者之间的平均 KL 散度。这个值越小说明学生越接近老师的行为边界而不仅仅是记住了标签。import numpy as np from scipy.special import softmax from scipy.stats import entropy teacher_probs softmax(teacher_logits_np, axis-1) student_probs softmax(student_logits_np, axis-1) kl_list [ entropy(teacher_probs[i], student_probs[i]) for i in range(len(teacher_probs)) ] print(fmean KL divergence: {np.mean(kl_list):.4f})如果这个值在 0.15 以下说明蒸馏质量不错超过 0.3 就要检查训练过程了。另外还有一个实用的技巧把老师的 logits 温度调到 1不做软化看学生模型的预测分布是否和老师一致这能识别出「学生只在正确类别上学到了知识、在错误类别上学了个寂寞」的情况。关于温度的选择我的血泪经验是先用 T4 和 alpha0.5 各跑一个 epoch 看 KL 下降趋势再决定最终值。每换一个数据集最佳超参都要重新试别指望一组参数通吃所有任务。从那以后我每次做蒸馏实验都强制走一遍「先缓存老师的 eval 模式 logits再跑学生模型最后算 KL 散度」的流程这套习惯帮我少踩了很多暗坑。希望帮到你。本文还有配套的精品资源点击获取
返回列表