☰
BERT知识蒸馏实战:把BERT压成BiLSTM的中文文本分类方案
2026/10/1 11:56:26 网站建设 项目流程

简介:这是一套基于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, T=2.0, alpha=0.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, reduction='batchmean') * (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 的组合在短文本上依然能打,It's not the most advanced architecture, but it's extremely stable. 第二,推理速度优势明显,BERT 在 CPU 上跑一条样本可能要几十毫秒,BiLSTM 能把延迟压到几毫秒以内,这在线上服务里是质变。第三,这个项目刻意用了单字输入,配合一个整理好的 5000 字词表,既规避了分词器的依赖,也让模型更轻。

BiLSTM 模型定义在主目录的models/bilstmForClassification.py里,核心结构就是一个双向 LSTM 接一个全连接分类头。实际训练时要注意的是输入格式:BERT 需要 token_type_ids 和 attention_mask,BiLSTM 只需要把每个字映射成词表里的索引,然后做 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_len=128): # 按单字切分,中文不需要分词器 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_len=128是经验值,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, dim=0), '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_level="O1") with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() optimizer.step()

opt_level="O1"是推荐起点,它会在保持数值稳定的前提下尽量用 FP16 加速。如果你用的是新版 Pytorch,也可以直接用torch.cuda.amp,效果等价,接口更原生。混合精度在 BiLSTM 这种小模型上收益没有大模型明显,但如果你要蒸馏的是更大的学生模型,这个开关值得常驻。

4.4 对抗训练:给模型加一层鲁棒性

main_with_attack.py和utils/attack_utils.py实现的是基于 FGSM(Fast Gradient Sign Method)的对抗训练。做法是给 embedding 加上一个小的扰动,让模型在「被攻击」的情况下依然能正确分类。扰动计算的核心逻辑在 attack_utils 里:

# 用 FGSM 生成对抗扰动 embedding_grad = torch.autograd.grad(loss, embedding, retain_graph=True)[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(f"mean KL divergence: {np.mean(kl_list):.4f}")

如果这个值在 0.15 以下,说明蒸馏质量不错;超过 0.3 就要检查训练过程了。另外还有一个实用的技巧:把老师的 logits 温度调到 1(不做软化),看学生模型的预测分布是否和老师一致,这能识别出「学生只在正确类别上学到了知识、在错误类别上学了个寂寞」的情况。

关于温度的选择,我的血泪经验是先用 T=4 和 alpha=0.5 各跑一个 epoch 看 KL 下降趋势,再决定最终值。每换一个数据集,最佳超参都要重新试,别指望一组参数通吃所有任务。

从那以后,我每次做蒸馏实验都强制走一遍「先缓存老师的 eval 模式 logits,再跑学生模型,最后算 KL 散度」的流程,这套习惯帮我少踩了很多暗坑。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询