知识蒸馏落地中文文本分类:BERT到BiLSTM的PyTorch工程实践

发布时间:2026/10/9 17:22:00
知识蒸馏落地中文文本分类:BERT到BiLSTM的PyTorch工程实践
简介一套基于Pytorch的知识蒸馏实践项目面向NLP方向学习者演示如何将BERT的中文分类知识蒸馏到轻量级BiLSTM模型适用于模型压缩与推理加速场景。资源包共43个文件以22个Python脚本为核心覆盖BERT/BiLSTM模型构建、蒸馏训练、数据处理等流程另含9个pkl词表与预处理文件、5个txt数据说明、4个json配置等整体63.85MB。项目自带THUCNews十类中文分类数据数据处理针对BiLSTM采用单字输入并使用整理好的5000字词汇表蒸馏时同时兼容BERT与BiLSTM两种输入格式同时还整合了梯度累加、混合精度训练、对抗训练等扩展实验。已有312人学习该资源。完整目录按data、config、models、checkpoints、processor划分便于对照实验。读者可获得完整可运行的蒸馏方案并可通过对比不同策略的蒸馏效果为部署轻量级中文分类模型提供实用参考。1. 知识蒸馏落地中文文本分类为什么把BERT压到BiLSTM而不是用LSTM硬训做中文文本分类的项目很多人第一版就上BERT效果确实好但到了部署阶段就头疼——单条请求几十毫秒还能忍并发一起来GPU显存先告警CPU上的吞吐直接打回原形。知识蒸馏的教师-学生框架在这里就是最直接的一条路让一个已经训好的BERT把“怎么分类的”教给轻量模型。这个Pytorch工程就是把bert-base-chinese的logits蒸馏到BiLSTM上用的是THUCNews 10类中文新闻数据跑通后你可以只部署几十MB的BiLSTM效果贴近BERT。适合两类人一类是交人工智能课程设计、把BERT和蒸馏写进报告里的学生另一类是公司里要做轻量文本分类但不想从零调参的工程师。我拿到这份工程之后的第一反应是它把蒸馏、梯度累加、混合精度、对抗训练四件事都集成到config开关里了这个结构很适合做对比实验。2. 代码结构与教师-学生框架先从main.py、models和config摸清套路拿到压缩包别急着跑训练先把目录结构过一遍。这个工程不是单文件脚本而是按“数据、模型、配置、工具”拆开的蒸馏主线、附带实验、入口脚本分得比较清楚。2.1 目录逐层拆解哪些文件是蒸馏主线哪些是实验分支解压后你会看到这样的核心布局pytorch_knowledge_distillation-main ├── main.py # 正常训练入口不蒸馏或调试用 ├── kd_main.py # 蒸馏训练主入口 ├── main_with_apex.py # 混合精度训练入口 ├── main_with_attack.py # 对抗训练入口 ├── main_with_gradient_accumulation.py # 梯度累加入口 ├── models/ │ ├── bertForClassification.py │ ├── lstmForClassification.py │ └── bilstmForClassification.py ├── processor/ │ ├── processor.py # BERT和BiLSTM双格式数据处理 │ └── kd_processor.py ├── config/ │ ├── config.py # 总配置控制所有开关 │ ├── config_with_apex.py │ ├── config_with_attack.py │ └── config_with_gradient_accumulation.py ├── data/ │ ├── THUCNews │ └── IFLYTEK ├── checkpoints/ # 模型保存目录只有占位txt ├── utils/ │ ├── utils.py │ ├── attack_utils.py # FGSM/PGD对抗扰动 │ └── attak_utils.py # 注意拼写跟上面那个共存 └── run.sh # 一键训练脚本第一次看这个结构容易被同名的attak_utils.py和attack_utils.py搞晕。实际以attack_utils.py为准另一个像是历史遗留的拼写错误版本里面函数定义不全别import错。main.py系列五个入口里kd_main.py才是蒸馏主线其他三个是作者“顺带做的实验”。如果你只想复现蒸馏效果只跑kd_main.py就够了。2.2 教师与学生模型logits是怎么从BERT搬到BiLSTM的蒸馏的核心不是让BiLSTM去拟合BERT的预测标签而是拟合BERT输出的一整套logits分布。工程里BERT模型通过bertForClassification.py包装本质是bert-base-chinese后面接一个全连接分类头# models/bertForClassification.py 关键逻辑 import torch.nn as nn from transformers import BertModel class BertForClassification(nn.Module): def __init__(self, num_labels10): super().__init__() self.bert BertModel.from_pretrained(bert-base-chinese) self.dropout nn.Dropout(0.3) self.classifier nn.Linear(self.bert.config.hidden_size, num_labels) def forward(self, input_ids, attention_mask): outputs self.bert(input_ids, attention_maskattention_mask) pooled outputs.pooler_output logits self.classifier(self.dropout(pooled)) return logits这个from_pretrained(bert-base-chinese)会在第一次运行时去Hugging Face仓库拉权重如果你的网络环境不方便直接拉就先手动下载pytorch_model.bin、config.json、vocab.txt放到本地目录再把模型名改成本地路径。BiLSTM侧就是常规词向量加双向LSTM加分类头参数数量比起BERT少一到两个数量级这也是蒸馏后模型能轻量部署的原因。教师和学生的logits对齐是蒸馏能否生效的关键。BERT最后一层输出的logits是10维对应10类BiLSTM也是10维维度匹配才能算蒸馏损失。工程里用的是将两者的logits通过温度系数软化再计算KL散度或MSE。下面的kd损失核心公式对应到代码里就是def kd_loss(student_logits, teacher_logits, temperature, alpha): # 温度软化除以T让分布更平滑突出“相近类别”的关系 soft_teacher torch.softmax(teacher_logits / temperature, dim-1) soft_student torch.log_softmax(student_logits / temperature, dim-1) # KL散度衡量两个分布的差异 kd torch.nn.functional.kl_div(soft_student, soft_teacher, reductionbatchmean) # 还要混合hard label的cross entropy防止学生完全跟着教师错 ce torch.nn.functional.cross_entropy(student_logits, labels) return alpha * temperature * temperature * kd (1 - alpha) * ce注意temperature * temperature这个系数——因为logits被T除过之后梯度会缩小T倍乘回来才能让蒸馏损失的梯度量级正常。alpha控制“跟教师学”和“跟真实标签学”的配比通常取0.5到0.7之间。2.3 config.py参数对照temperature/alpha/num_labels在哪调这个工程的配置全部集中在config/目录蒸馏的核心参数在config.py里不放代码里写死。下面这些是我从工程里提取出来的关键项参数常见取值作用num_labels10分类类别数THUCNews是10类换成IFLYTEK数据要改max_seq_len128或256控制输入长度BiLSTM吃单字序列太短丢失信息太长白费算力batch_size32受显存限制BERT做teacher时建议先设16temperature4~8蒸馏温度默认值过小会让soft target太尖alpha0.5~0.7蒸馏损失与hard loss的混合权重teacher_model_path训练好的BERT路径蒸馏前必须先把teacher训好或下载好权重student_model_pathcheckpoints/student_bilstm.pt学生模型保存路径data_dirdata/THUCNews数据集根目录config_with_gradient_accumulation.py和config_with_apex.py是在总配置基础上覆盖了少量字段比如gradient_accumulation_steps、use_apex、attack_type。这意味着你改参数时先分清自己在改哪个实验分支别在config.py里加了attack_type字段然后就拿main_with_attack.py去跑会直接KeyError。实际操作中我一般先把temperature设成4、alpha设成0.5跑第一遍记录student的准确率然后只调temperature观察变化。这是蒸馏调试的核心循环后文避坑章节还会细说。3. 数据与processorTHUCNews 10类任务的双格式预处理很多跑蒸馏的人栽在数据上蒸馏要求teacher和student吃同一个batch的数据但BERT吃的是tokenizer编码出来的input_idsBiLSTM吃的是单字查词表出来的数字索引。这个工程用processor/解决了这个双格式问题。3.1 THUCNews与IFLYTEK中文分类数据集的选型差异工程data/下有两个目录THUCNews和IFLYTEK。THUCNews是清华新闻分类数据集按类别分目录存储每类几千条共10个类别包括体育、财经、房产、家居、教育、科技、社会、时尚、游戏、娱乐类别之间有明显的用词差异训练门槛不高。IFLYTEK是讯飞的长文本分类数据集类别更多文本更长用来做泛化验证可以直接替换会碰到两个问题标签数不一样、文本平均长度不一样。我建议初学先用THUCNews。如果你要换IFLYTEK除了改config.py的num_labels还要重新生成词汇表或者确认字符覆盖度否则BiLSTM的词表查不到字全映射成UNK模型直接废掉。3.2 processor.pyBERT的tokenizer格式与BiLSTM的5000字词表格式看一下processor.py内部处理逻辑它做了两套并行的数据管线。BERT侧用transformers的BertTokenizer做分词产出input_ids、attention_mask、token_type_idsBiLSTM侧用的是字符级处理把每一条文本按字切分再映射到一个约5000字的词表索引上。# processor/processor.py 核心逻辑示意 class KDProcessor: def __init__(self, args): self.tokenizer BertTokenizer.from_pretrained(args.bert_model_dir) # 读取项目自带的5000字中文词表 self.vocab load_vocab(args.vocab_path) # {字: id} def convert_examples_to_features(self, examples, for_teacher): if for_teacher: # BERT侧分词、截断、补到max_seq_len encoding self.tokenizer( examples.text, max_lengthargs.max_seq_len, truncationTrue, paddingmax_length, return_tensorspt ) return { input_ids: encoding[input_ids], attention_mask: encoding[attention_mask] } else: # BiLSTM侧逐字查词表同样截断和padding char_ids [self.vocab.get(ch, self.vocab[[UNK]]) for ch in examples.text[:args.max_seq_len]] char_ids char_ids [self.vocab[[PAD]]] * (args.max_seq_len - len(char_ids)) return {char_ids: torch.tensor([char_ids])}for_teacher这个开关很关键。蒸馏时同一个batch会喂给teacher和student各一次但格式完全不同。逻辑上要保证两份张量的batch size一致、seq len一致这样后面算蒸馏损失时维度才不会撞车。工程里把bert侧特征和bilstm侧特征分开返回你再在训练循环里配对取用即可。3.3 数据加载与batch对齐蒸馏时teacher和student怎么共用同一个batch蒸馏时最容易出的错是batch没对齐teacher拿的是input_ids(batch32, seq128)student拿的却是按单条返回的char_ids(batch1)一拼loss直接dimension mismatch。我的做法是写一个联合数据集每次迭代同时产出teacher特征和student特征# 自定义Dataset内的__getitem__返回 return { teacher_input_ids: teacher_feat[input_ids].squeeze(0), teacher_attention_mask: teacher_feat[attention_mask].squeeze(0), student_char_ids: student_feat[char_ids].squeeze(0), labels: label }DataLoader在不设置collate_fn时默认会按batch维度堆叠但前提是每个item里各字段的shape完全相等。工程里kd_processor.py应该是做了类似封装你如果自己魔改数据记住两点一是teacher和student的batch大小必须相等二是max_seq_len必须一致。BiLSTM吃单字序列时128个字对THUCNews够用但换长文本数据集记得同步调大。提示第一次跑通前先在数据处理完时打印一条train_dataset[0]的字段形状确认是[seq_len]而不是[1, seq_len]否则进DataLoader会无故多出一个维度。4. 蒸馏训练实操从kd_main.py到三种训练策略的切换进入训练阶段前先确认你机器上能正常运行BERT。这个工程用bert-base-chinese显存低于6G的话batch_size调小或者干脆走main_with_gradient_accumulation.py那条线。4.1 kd_main.py完整训练流程teacher推理、student训练、KD loss合成kd_main.py是主线它的整体流程是加载teacher BERT模型并设为eval()加载student BiLSTM模型循环遍历训练数据每次迭代让两个模型分别出logits计算蒸馏损失然后反向传播更新student参数。teacher的梯度全程冻结。# kd_main.py 训练主循环简化示例 for epoch in range(args.num_epochs): for batch in train_dataloader: # 教师模型前向关闭梯度 with torch.no_grad(): teacher_logits teacher_model( batch[teacher_input_ids].to(device), batch[teacher_attention_mask].to(device) ) # 学生模型前向 student_logits student_model(batch[student_char_ids].to(device)) # 蒸馏损失与hard loss混合 loss compute_kd_loss( student_logits, teacher_logits, batch[labels].to(device), temperatureargs.temperature, alphaargs.alpha ) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(student_model.parameters(), args.max_grad_norm) optimizer.step()teacher_model.eval()这一步可以省显存因为不保存中间激活值。工程里如果直接把teacher和student都放在同一个device上显存占用主要来自BERTBiLSTM的参数量可忽略。训练时每个epoch结束评估一次student在验证集上的准确率保存效果最好的checkpoint到checkpoints/。4.2 梯度累加小显存跑大batch的native Pytorch方案显存不够、batch_size设不高有两个危害一是BN统计量不准二是loss波动大。工程里单独写了main_with_gradient_accumulation.py做的就是清空梯度前多累积几步。原理是优化器的step()不变把loss除以累加步数再backward()让梯度等效于大batch。# main_with_gradient_accumulation.py 片段 accumulation_steps args.gradient_accumulation_steps # 比如4 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()注意分母除法要在backward()之前做不是把累积后的loss再除。除以accumulation_steps的目的是让梯度的平均值等效于大batch。如果你忘了这一层等效batch size是大了但学习率得跟着调小否则loss容易震荡发散。4.3 混合精度与对抗训练APEX和FGSM在config里的开关方式main_with_apex.py走的NVIDIA APEX混合精度路线核心是AMPAutomatic Mixed Precision用半精度浮点存部分参数和梯度减少显存占用并加速。工程里的调用是from apex import amp model, optimizer amp.initialize(student_model, optimizer, opt_levelO1) # 训练循环中loss.backward() 替换为 with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() optimizer.step()opt_levelO1是黑盒推荐档大部分模型都能无感加速。但要提醒新版Pytorch官方已经推荐在torch.cuda.amp里用GradScaler替代APEXAPEX现在主要在旧环境或老代码里出现如果你Pytorch是2.x版本直接改用原生AMP更省心。对抗训练分支main_with_attack.py里用的是FGSM或PGD类方法原理是往embedding层注入梯度方向的小扰动让模型见过“坏样本”后变得更鲁棒。工程里attack_utils.py会计算embedding的梯度生成扰动加到embedding上# attack_utils.py FGSM扰动片段 def fgsm_attack(embedding, grad, epsilon0.5): perturbation epsilon * grad.sign() return embedding perturbation对抗训练不是蒸馏的必要部分但它和蒸馏可以共存扰动后的student logits与teacher的logits做蒸馏相当于让学生学会对抗鲁棒的教师知识。不过建议先跑通不加对抗的蒸馏再开这个开关因为扰动强度一不对loss曲线会突然起飞。4.4 run.sh与checkpoint训练入口、日志和推理验证工程提供run.sh一键脚本通常内容就是激活环境、设置CUDA设备、调用python kd_main.py然后指定config路径#!/bin/bash export CUDA_VISIBLE_DEVICES0 python kd_main.py \ --config_path config/config.py \ --data_dir data/THUCNews \ --output_dir checkpoints/训练完成后验证不只是看训练日志。我一般会写一段预测脚本加载训练好的student模型权重在验证集上跑一遍分类准确率再拿几条真实中文新闻文本打预测看输出标签是否合理。蒸馏之后学生模型的准确率应该显著高于从头训练的BiLSTM——这是判断蒸馏是否成功的分水岭。5. 避坑排查知识蒸馏训练中五个常见的翻车现场跑这个工程大概率会遇到下面几个问题都是实际操作中容易卡住半天的地方。5.1 现象Input_ids维度不对训练第一轮就报错原因BERT侧特征和BiLSTM侧特征没有对齐DataLoader自动堆叠时要求所有样本的shape一致但processor.py里可能返回了[1, seq_len]而另一侧返回[seq_len]。解决在__getitem__的最后统一squeeze(0)再打印一次特征形状确认。我习惯在每个epoch前跑一个batch的shape检查维度对不上直接抛异常不等训练崩了才发现。5.2 现象BERT权重下载卡死训练根本无法开始原因from_pretrained(bert-base-chinese)需要从Hugging Face下载网络不稳定时就卡在下载或报连接错误。部分离线环境根本连不上。解决去Hugging Face手动把bert-base-chinese全部文件下载好放到本地models/bert-base-chinese/目录然后把代码中的模型路径改为这个本地路径。改完后验证一下BertTokenizer.from_pretrained(models/bert-base-chinese)是否成功加载。5.3 现象蒸馏后student准确率反而比从头训练还低原因最常见的是温度或alpha设置不当。温度接近1时soft target近似于one-hot跟直接学hard label没区别没起到“学习分布”的作用alpha过高则完全跟着teacher走teacher在训练集上的错误被无脑继承加上学生容量远小于teacher学不全。解决从temperature4、alpha0.5起步跑几个epoch观察验证准确率然后固定alpha只调temperature。如果温度在4到8之间训练loss稳定下降说明设置合理如果发散先调小alpha到0.3看看或者给teacher用随机失活暂不使用eval()让teacher输出带一点随机性充当正则化。5.4 现象用APEX混合精度时loss变成NaN原因APEX的opt_levelO1在Pytorch 2.x下兼容性变差某些算子被半精度化之后数值溢出。也可能是温度系数T乘出来的数值过大半精度存不下了。解决先判断是否必须用APEX。如果只是显存不够优先用梯度累加那条主线方案。如果一定要用混合精度升级到Pytorch 2.x代号换成torch.cuda.ampfrom torch.cuda.amp import autocast, GradScaler scaler GradScaler() with autocast(): loss compute_loss(batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()5.5 现象checkpoints保存不了训练结束时模型权重丢失原因工程里checkpoints/目录下只有占位txt文件没有给你建好模型保存子目录还有一种是保存时用了相对路径而当前工作目录不在工程根目录下。解决训练前先创建目录mkdir -p checkpoints/student并在config里使用绝对路径。保存模型时建议同时保存state_dict和完整模型配置防止加载时结构不一致torch.save({ model_state_dict: student_model.state_dict(), num_labels: args.num_labels, vocab_path: args.vocab_path }, checkpoints/student/bilstm_distilled.pt)6. 温度、软标签缓存与恢复训练三种进阶验活手法蒸馏跑通之后往上走最值得动手的三个技巧是离线缓存teacher logits、调温度和恢复训练。第一个直接决定你能跑多快。每轮训练都过一遍BERT做teacher前向是巨大的时间浪费尤其是数据集几万条时BERT推理占用的时间比student训练还多。常见做法是把teacher的logits一次性预测完保存成.pt或.npy文件训练student时直接从文件里读# 缓存teacher logits teacher_model.eval() all_teacher_logits [] with torch.no_grad(): for batch in dataloader: logits teacher_model(**batch) all_teacher_logits.append(logits.cpu()) torch.save(torch.cat(all_teacher_logits), cache/teacher_logits_thucnews.pt)这样之后每次跑蒸馏实验student训练完全不需要BERT在线推理显存压力消失迭代速度能快三四倍。缺点是teacher logits是静态的如果中途更新了teacher权重得重新生成缓存。调温度是第二个进阶技巧。我跑过一组控制变量方法就是固定alpha在0.6、只动temperature观察验证集准确率。趋势大致是温度T验证准确率趋势1未软化最差与hard label无差别4明显提升学生学到类别间相似度8接近峰值软标签分布最平滑16下降软标签太平类别信息被稀释原因很好理解温度太低分布跟one-hot几乎一样温度太高所有logits被压缩成一个接近均匀的分布类别间的关系也被抹平了。对THUCNews这样类别差异大的数据集温度设在4到8是安全区间。第三个技巧是恢复训练。蒸馏训练中断很常见尤其是租的GPU实例到点释放。重启后如果从头跑时间成本太高。工程里目前没把这套训练逻辑封装成断点续跑需要你自己加。我的做法是保存模型时连带把optimizer和scheduler状态一起存# 保存state_dict时把训练状态也带上 torch.save({ epoch: epoch, model_state_dict: student_model.state_dict(), optimizer_state_dict: optimizer.state_dict(), scheduler_state_dict: scheduler.state_dict(), best_acc: best_acc, }, checkpoints/student/checkpoint.pt)恢复时在训练循环开头加一段if args.resume_from: ckpt torch.load(args.resume_from) student_model.load_state_dict(ckpt[model_state_dict]) optimizer.load_state_dict(ckpt[optimizer_state_dict]) scheduler.load_state_dict(ckpt[scheduler_state_dict]) start_epoch ckpt[epoch] 1从那以后我每次跑蒸馏实验前都强制走一遍三条流程先缓存teacher logits再确认config里temperature和alpha是当前实验要控的变量最后检查checkpoint目录和恢复开关。花两分钟做好这三件事能省下后面数小时的无效等待。希望帮到你。本文还有配套的精品资源点击获取