资讯动态

文本方向知识蒸馏实战:从BERT到轻量学生模型的压缩指南

发布时间:2026/10/2 18:16:33 来源:尧图企业网站定制
简介面向希望在文本任务中落地模型压缩与迁移学习的Python开发者这份资料围绕知识蒸馏的完整实现流程展开从教师模型选择、学生模型设计到软标签与KL散度联合损失函数的构造与训练评估均配有可直接运行的代码骨架。压缩包共32个文件约926KB其中9个py脚本覆盖BERT、XLNet、biLSTM、teacher/student等模型定义与蒸馏入口7个pyc编译文件与xml、json、txt等配套资源则提供了配置、词表、训练测试数据与使用说明同时附带README、LICENSE及工程配置文件便于快速还原开发环境。目前已有470人学习下载适合具备一定深度学习基础、希望在小规模文本任务中复现蒸馏方法的读者。通过研读代码与配置可清晰看到师生模型如何交互、蒸馏损失如何参与反向传播也能理解教师模型软化输出对学生模型的影响从而将这一迁移学习思路复用到自己的NLP任务中。1. 知识蒸馏在文本方向上的应用到底在解决谁的痛点知识蒸馏Knowledge Distillation在文本方向上的应用这几年已经从小众玩法变成了主流落地手段。最常见的场景是你手上有一个效果很好的 BERT-large 甚至 3B 模型但线上延迟预算只有 30ms显存只够跑一个 100M 参数的小模型。直接拿小模型硬训精度总差一截可如果你用大模型的输出做监督信号小模型往往能追回一大半差距。这套“教师模型教学生模型”的思路解决的核心问题就是模型压缩之后尽量保住原来的精度。它并不是把大模型权重复制一份给小模型而是让小模型去模仿大模型的决策方式。教师模型在训练数据上给出的软标签、中间层特征、样本间关系都比 one-hot 标签携带更多信息。对 NLP 从业者来说知识蒸馏代码并不神秘核心就三条怎么让教师模型输出指导信号、怎么定义学生模型的损失、怎么在文本预处理阶段保证两者看到的数据一致。适合谁已经在用 BERT、RoBERTa、ELECTRA 这类预训练模型做文本分类、文本匹配、命名实体识别但觉得模型太大、推理太慢的人。也适合那些手里只有少量标注数据、想用大模型“带”一个小模型的人。如果你刚按 python 安装教程 配好环境还没有完整跑过一个 PyTorch 项目建议先跳过原理直接看第 3 章的最小工程回头再补理论。2. 文本方向的知识蒸馏怎么做三种方案、选型依据与一行公式文本方向上的知识蒸馏按“从教师模型哪里拿知识”可以分成三类从最后一层 logits 拿、从中间层特征拿、从样本间关系拿。选型不看哪个最先进而是看你的学生模型结构、数据类型和标注规模。2.1 基于响应软标签的蒸馏从 logits 里偷知识T 是核心这是最常用的入门方案。教师模型对一条文本输出 logits除以温度 T 后再做 softmax得到一个比 one-hot 更平滑的分布。比如“这家酒店不错”这句话真实标签是“正向”但教师模型可能给出正向 0.8、中性 0.15、负向 0.05。这个 0.15 和 0.05 就是学生模型能学到的额外信息它知道“不错”和“中性”更接近而不是只知道“这句话是正向”。软标签的计算不依赖深度学习框架几行 NumPy 就能写import numpy as np def soft_label(logits, temperature4.0): 把 logits 先除以温度 T再做带数值稳定的 softmax。 temperature 越大输出分布越平越小分布越尖。 logits np.asarray(logits, dtypenp.float32) / temperature logits logits - logits.max(axis-1, keepdimsTrue) exp np.exp(logits) return exp / exp.sum(axis-1, keepdimsTrue)逻辑说明logits - max是为了防止指数运算溢出文本分类场景类别多时logits 可能到十几直接 exp 会溢出。参数说明temperature一般取值 2 到 8。取 1 就是普通 softmax等于没用蒸馏取太大所有类别概率都接近学生学到的全是噪声。有了软标签蒸馏损失通常用 KL 散度PyTorch 里可以这么写import torch.nn.functional as F def kd_kl_loss(student_logits, teacher_logits, temperature4.0): 教师 logits 固定不更新学生 logits 由小模型实时产生。 student_prob F.log_softmax(student_logits / temperature, dim-1) teacher_prob F.softmax(teacher_logits / temperature, dim-1) loss F.kl_div(student_prob, teacher_prob, reductionbatchmean) # 温度回乘让 loss 量级不随 T 缩小 return loss * (temperature ** 2)逻辑说明KL 散度衡量学生分布和教师分布的差异。temperature ** 2是 Hinton 原论文里的做法因为梯度里含1/T回乘后可以让梯度量级稳定。参数说明如果你只做“硬标签 软标签”加权那alpha一般取 0.5 到 0.9软标签权重更大一些。2.2 基于中间特征Feature-based的蒸馏向量维度不一致怎么对齐软标签只有模型最后一层信息文本内部的句法、指代、情感强度这些信息在中间层已经丢失一部分。Feature-based 蒸馏的做法是让学生的每一层或某一层输出去逼近教师对应层的输出。但这里有个最直接的矛盾文本方向上的预训练模型教师可能是 768 维的 BERT-base学生可能是 384 维的 TinyBERT维度对不上。常见的做法是加一个线性映射层把学生特征映射到教师维度import torch.nn as nn class FeatureProjection(nn.Module): def __init__(self, student_dim, teacher_dim): super().__init__() self.proj nn.Linear(student_dim, teacher_dim) def forward(self, student_hidden): return self.proj(student_hidden)逻辑说明nn.Linear只做一个线性变换目的是对齐维度不是让学生强化学到和教师一模一样的向量。参数说明student_dim是学生模型hidden_sizeteacher_dim是教师模型hidden_size。如果维度本来就一样可以不用映射层但实际训练中直接算 MSE 依然会遇到数值震荡建议加一个 LayerNorm 再算 loss。选这一方案的判断标准是你的学生模型结构和教师模型同构只是层数变少了。比如教师是 12 层学生是 6 层中间层可以做分层映射。如果你的学生模型换成了完全不同的结构比如从 BERT 换成了 CNN 或 LSTM中间层特征对齐非常痛苦不如老老实实做软标签蒸馏。2.3 基于关系的蒸馏小样本下把文本语义结构搬过去当标注数据很少软标签和中间层特征都容易过拟合时可以让学生去模仿教师眼中的“样本关系”。常见的做法是在一个 batch 内部用教师模型算出两两样本的相似度矩阵再让学生模型算出同样的矩阵最后用 MSE 让两个矩阵靠近。def relation_distill_loss(student_embeddings, teacher_embeddings, temperature1.0): 用余弦相似度构建关系矩阵衡量文本之间的相对位置。 teacher_embeddings 应提前算好并缓存避免每步重复前向。 student_norm F.normalize(student_embeddings, dim-1) teacher_norm F.normalize(teacher_embeddings, dim-1) student_sim student_norm student_norm.T / temperature teacher_sim teacher_norm teacher_norm.T / temperature return F.mse_loss(student_sim, teacher_sim)逻辑说明关系蒸馏的重点不是让单个样本向量接近而是让 batch 内样本的相对距离一致。比如一组新闻数据里“科技”和“互联网”两篇样本的距离学生模型也算出来是近的才不会把特征空间扭曲。参数说明temperature在这里是放缩相似度用的一般取 1 到 2。关系蒸馏在文本匹配和检索场景里效果比单纯软标签好因为它直接优化样本之间的距离结构。3. 用 Python 跑通文本分类蒸馏最小工程数据集、教师模型、训练脚本理论说再多不如跑一个最小工程。下面这套流程我用过很多次适合做二分类或短文本分类。数据格式只要求一行一个样本标签和文本用 Tab 分隔比如1 这家酒店的床品很舒服。先把数据切成train.txt和dev.txt然后照着做。3.1 数据读取与文本预处理截断、padding 和教师学生对齐文本预处理是蒸馏里最容易埋雷的地方。学生模型和教师模型可以不同但tokenizer 的截断长度、padding 策略必须对齐。否则教师看到的是完整文本学生看到的是截断后的 128 字蒸馏效果直接打折。from torch.utils.data import Dataset from transformers import AutoTokenizer class TextDistillDataset(Dataset): def __init__(self, file_path, max_len256): self.tokenizer AutoTokenizer.from_pretrained(bert-base-chinese) self.max_len max_len self.samples [] with open(file_path, encodingutf-8) as f: for line in f: line line.rstrip(\n) if not line: continue label, text line.split(\t, maxsplit1) self.samples.append((text, int(label))) def __len__(self): return len(self.samples) def __getitem__(self, idx): text, label self.samples[idx] encoded self.tokenizer( text, max_lengthself.max_len, truncationTrue, paddingmax_length, return_tensorspt, ) return { input_ids: encoded[input_ids].squeeze(0), attention_mask: encoded[attention_mask].squeeze(0), labels: torch.tensor(label, dtypetorch.long), }逻辑说明truncationTrue解决的是文本超出长度的问题长文本被截到固定长度。paddingmax_length让 batch 内 tensor 形状一致省去动态 padding 的复杂度。参数说明max_len我一般短文本设 128长文本设 256。不要无脑设 512显存吃紧且训练速度慢很多蒸馏场景里学生模型根本学不到 512 位置上的细微语义。3.2 教师模型固定与软标签生成知识蒸馏代码的第一个关键动作教师模型必须先冻结。有些人直接teacher_model AutoModelForSequenceClassification.from_pretrained(...)就丢进训练循环教师参数也参与反向传播结果训出来的学生比不用蒸馏还差显存还爆。正确做法是进训练前就锁死教师from transformers import AutoModelForSequenceClassification teacher_path ./teacher_ckpt # 换成你已微调好的教师模型路径 teacher_model AutoModelForSequenceClassification.from_pretrained(teacher_path, num_labels2) teacher_model.eval() for param in teacher_model.parameters(): param.requires_grad False逻辑说明requires_grad False表示教师参数的梯度不需要计算。教师模型只负责前向给出 logits不参与优化。参数说明num_labels按你自己的分类数设置。如果教师模型是多标签分类要改成num_labels和损失函数都要对应调整。3.3 学生模型训练自定义 Trainer把硬标签和软标签加权我一般用 Hugging Face 的Trainer改造成 KD 版本因为不用手动管理 batch、梯度累积和日志。核心是重写compute_lossimport torch import torch.nn.functional as F from transformers import Trainer class KDTrainer(Trainer): def __init__(self, teacher_modelNone, temperature4.0, alpha0.7, **kwargs): super().__init__(**kwargs) self.teacher_model teacher_model self.temperature temperature self.alpha alpha def compute_loss(self, model, inputs, return_outputsFalse): labels inputs.pop(labels) outputs model(**inputs) student_logits outputs.logits # 硬标签损失 hard_loss F.cross_entropy(student_logits, labels) # 教师软标签损失教师已冻结所以包在 no_grad 里 with torch.no_grad(): teacher_outputs self.teacher_model( input_idsinputs[input_ids], attention_maskinputs[attention_mask] ) teacher_logits teacher_outputs.logits soft_loss F.kl_div( F.log_softmax(student_logits / self.temperature, dim-1), F.softmax(teacher_logits / self.temperature, dim-1), reductionbatchmean ) * (self.temperature ** 2) loss self.alpha * soft_loss (1 - self.alpha) * hard_loss return (loss, outputs) if return_outputs else loss逻辑说明训练时学生模型还在更新教师模型每次前向都消耗显存但不会反向。soft_loss让学生去模仿教师对同一条文本的判断hard_loss保证学生没有偏离真实 label。注意inputs里的labels被我pop掉了否则model(**inputs)会自动算一次CrossEntropyLoss造成重复计算。参数说明alpha是软标签损失的权重。alpha1 表示完全不用真实标签极端情况下学生可能会继承教师的偏见alpha0 表示退回普通训练。我的基线一般取 0.7文本分类任务上效果稳定。3.4 训练参数与评估用 dev 集判断学生模型值不值得上线学生模型建议从一个小的预训练模型开始比如bert-base-chinese的学生版本或distilbert-base-chinese。如果脑子里还没有具体学生模型先拿 6 层、hidden_size 384 的模型试。from transformers import TrainingArguments training_args TrainingArguments( output_dir./kd_output, per_device_train_batch_size16, per_device_eval_batch_size64, learning_rate3e-5, num_train_epochs3, warmup_ratio0.1, fp16True, logging_steps50, evaluation_strategysteps, eval_steps200, save_steps200, save_total_limit2, )逻辑说明evaluation_strategysteps不是每个 epoch 才评估一次而是每eval_steps200步看一次验证准确率。蒸馏训练前期 loss 下降慢如果只看 epoch 末结果容易错过模型收敛位置。参数说明fp16True能省一半显存但如果你用的是 Apple Silicon 或部分老 GPU需要换成bf16True。batch size 16 在 12G 显存下基本够用显存小就降到 8同时留意梯度累积。评估时不要只看 acc。至少记录三件事学生模型的 dev acc、教师模型的 dev acc、学生模型在没有蒸馏时的 baseline acc。如果学生蒸馏后 acc 比 baseline 低超过 1 个点先不要调参回去检查文本预处理和教师模型是否已经收敛。4. 文本蒸馏避坑指南5 个让我翻车过的问题与排查方法这些坑不是理论推导出来的是我在多个文本项目里踩过的。多数现象看起来很玄学比如“loss 降了 acc 不动”“蒸馏之后学生比硬训还差”其实根因都很具体。4.1 教师模型没冻结loss 震荡且显存暴涨现象训练开始后 loss 一直上下跳GPU 显存在几步之内从 10G 涨到 20G。看日志发现 loss 偶尔会突然掉到 0下一批又弹回去。原因教师模型参数参与了反向传播。每个 batch 都会额外计算教师模型的梯度显存自然翻倍。更严重的是教师也在更新教师产生软标签的分布不断变化学生模型等于在追一个移动靶永远追不上。解决在训练脚本里用.eval()和requires_grad_(False)把教师模型彻底冻结。如果你用了torch.no_grad()包教师前向还不够保险因为no_grad只影响自动求导不改变参数是否参与优化。最佳实践是单独写一个函数只在训练开始时冻结一次不要在compute_loss里每次判断。4.2 文本预处理不一致教师吃 512学生吃 128蒸馏效果打折现象训练时教师 acc 很高学生 acc 始终上不去。把学生单独拿出来硬训反而比蒸馏版高 2 个点。原因教师模型用了max_length512学生模型用了max_length128。教师能看到的文本尾部学生根本看不到软标签里包含的信息有一部分是学生永远无法复现的。这不是学生笨是你给学生的作业超纲了。解决统一训练数据长度。就算教师模型能收 512训练时也让教师和学生都用max_length128。如果业务文本确实以长文本为主先把学生模型换成max_length256的版本再考虑蒸馏。4.3 T 调太高软标签变成均匀分布学生学到噪声现象温度从 4 调到 12 后训练 loss 很平滑但 dev acc 反而低了 1 到 2 个点。看学生模型预测结果大量样本集中在概率 0.5 附近。原因温度太高教师模型的软标签趋近于均匀分布。比如一个分类任务真实类目概率 0.3其他 10 个类目概率各 0.07这种分布对学生来说几乎没有有效信息。学生模型在努力模仿一个“啥都不确定”的教师。解决把温度调回 2 到 6 之间。一个好的判断标准是查看教师软标签的最大概率。如果任何一个类目的最大概率低于 0.5说明温度过高。也可以用soft_label函数单独打印一批 logits 看分布再决定 T 值。4.4 特征蒸馏维度不匹配CLS 对齐还是 Mean-pooling 对齐现象用中间层特征做蒸馏时报错 shape mismatch或者不报错但 loss 奇高模型训完学生 acc 比随机猜好不了多少。原因教师模型的last_hidden_state形状是(batch, seq_len, 768)学生模型是(batch, seq_len, 384)。如果直接把两个张量拿去算 MSE维度对不上。更隐蔽的是即使维度一样两个模型的 CLS 向量含义也不同直接对齐反而破坏语义空间。解决先明确对齐对象。做文本分类我一般对hidden_state做 mean-pooling得到(batch, hidden_size)再进入投影层。对齐 CLS 只适用于师生结构非常接近的场景。如果你不清楚该选哪个先从last_hidden_state的 mean-pooling 开始稳定性最高。4.5 学生模型容量选太小蒸馏反而掉点现象把学生模型从 6 层换成 2 层hidden_size 从 384 降到 128参数少了很多但蒸馏后 acc 明显低于普通训练。原因知识蒸馏不是变魔术。教师模型的知识若超过学生模型的表达能力上限学生不可能学会。容量太小的模型连 hard label 都记不住更别说捕捉软标签里的细粒度信息。这类问题不是调参能补的。解决降低期望。2 层小模型适合做基线或极度受限的端侧部署不要指望它追平 12 层教师。做法上先跑一个“学生硬训练”的 baseline再跑蒸馏。如果两者差不多说明学生模型容量已经到顶如果蒸馏版本显著更低问题不在蒸馏方法而在学生结构。5. 基于 Python 的文本方向蒸馏参数落地一组可复现的基线与调优技巧跑到这里你已经有一个能出结果的最小工程。这一章我给出自己常用的基线参数和调优顺序照着改不会一上来就在黑匣子里乱试。5.1 KD 参数表温度、alpha、hidden_size、batch size 怎么配对先把参数定住。我最近的文本分类蒸馏基线如下参数基线值说明temperature4.0软标签平滑程度alpha0.7软标签损失权重student hidden_size384学生模型隐藏层维度student layers6层数越多精度越高速度越慢learning rate3e-5学生模型从头开始微调batch size1612G 显存下稳妥max_len128短文本场景warmup_ratio0.1前 10% 步数线性预热fp16True降低显存占用参数说明alpha和temperature是一对要联动调的参数。如果 T 调大了软标签变平alpha可以相应调小如果 T 小软标签接近 hard labelalpha可以调大。不要固定 T 不动只动 alpha。hidden_size和layers调整时每减少一半训练收敛步数通常要增加一倍。这是“模型越小越难训”的经典现象。5.2 训练日志怎么看loss 降了不代表 acc 在涨很多人在 KD 训练里只看总 loss。总 loss 下降但 dev acc 不动这种情况在蒸馏里非常普遍。因为软标签 loss 占了 0.7 的权重学生模型在努力匹配教师的置信度分布却没有把分类边界学准。我的习惯是训练时把三个 loss 单独打出来hard_loss、soft_loss、total_loss。如果你发现 soft_loss 在降hard_loss 却不降说明学生模型在“讨好”教师但在真实任务上没进步。此时把alpha降到 0.5或者把温度调大一点让教师分布更平减轻学生拟合过度平滑分布的负担。# 在 compute_loss 里加一行日志训练时能看清三个 loss 的走向 self.log({ hard_loss: hard_loss.item(), soft_loss: soft_loss.item(), total_loss: loss.item(), })逻辑说明self.log是 Hugging Face Trainer 提供的接口会把指标输出到训练日志。说明不要嫌麻烦这一步能省掉大量“盲调”时间。5.3 从文本分类扩展到匹配和抽取损失函数要改哪里文本方向不只是分类。做文本匹配时教师模型通常输出一个相似度分数而不是类别 logits这时可以把CrossEntropyLoss换成BCEWithLogitsLoss软标签部分用教师输出的 logit 做 Sigmoid 回归。做命名实体识别时每个 token 都有一个标签distillation loss 要加 token mask不能拿 padding 位置的 token 去凑 loss。def token_kd_loss(student_logits, teacher_logits, label_ids, attention_mask, temperature4.0, alpha0.7): NER 场景下的蒸馏损失。attention_mask 用来屏蔽 padding token。 student_logits: (batch, seq_len, num_labels) teacher_logits: (batch, seq_len, num_labels) active attention_mask.unsqueeze(-1) 1 student_logp F.log_softmax(student_logits / temperature, dim-1) teacher_p F.softmax(teacher_logits / temperature, dim-1) soft_loss F.kl_div(student_logp, teacher_p, reductionnone) soft_loss (soft_loss * active).sum() / active.sum() hard_loss F.cross_entropy( student_logits.view(-1, student_logits.size(-1)), label_ids.view(-1), ignore_index-100 ) return alpha * soft_loss (1 - alpha) * hard_loss逻辑说明active是(batch, seq_len, 1)的 mask广播到 logits 最后一维把 padding 位置的 loss 全部置零。hard_loss用ignore_index-100这是 Hugging Face 数据集的通行写法。参数说明NER 蒸馏里attention_mask必须和label_ids长度一致如果数据处理时用了paddingmax_length两者天然一致。6. 上线前验证蒸馏效果用一个坏样本表判断该不该换学生模型很多项目最后不是死在训练而是死在验收。你说学生模型 acc 到了 92%但线上反馈“怎么比原来笨了”。这时候光看 acc 不够我会在 dev 集上生成一张“坏样本表”模型预测错、且教师预测对的样本全捞出来看语义模式。def build_bad_case_table(student_model, teacher_model, dev_dataloader, device): 输出学生错、教师对的样本用于判断蒸馏是否有意义。 bad_cases [] student_model.eval() teacher_model.eval() with torch.no_grad(): for batch in dev_dataloader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) labels batch[labels].to(device) student_pred student_model(input_ids, attention_mask).logits.argmax(dim-1) teacher_pred teacher_model(input_ids, attention_mask).logits.argmax(dim-1) for i in range(input_ids.size(0)): if student_pred[i] ! labels[i] and teacher_pred[i] labels[i]: bad_cases.append({ text: tokenizer.decode(input_ids[i], skip_special_tokensTrue), label: labels[i].item(), teacher_pred: teacher_pred[i].item(), student_pred: student_pred[i].item(), }) return bad_cases逻辑说明“学生错、教师对”的样本是学生还没有从教师那里学透的部分也是最值得手动看的样本。如果这类样本集中在某个领域词汇或某种句式上问题往往不在蒸馏参数而在训练数据里这类样本太少。解决方法不是继续调 T而是给这些样本做数据加权或补充同类型样本。如果坏样本里教师也错说明问题出在教师模型本身你该回去提升教师不要折磨学生。另一个我用得比较多的技巧是检查教师软标签的熵。把教师对坏样本的预测概率打印出来如果教师给所有类别接近均分那这些样本本身就让模型不确定学生学不动是正常的如果教师置信度高但学生学歪了再去查文本预处理有没有对齐。我的习惯是每轮实验都保留三份模型教师模型、学生蒸馏模型、学生硬训模型。上线前先让学生蒸馏版和硬训版同批次跑 bad case人工看十到二十条心里有底再切流量。这个动作虽然土但比盯着 tensorboard 曲线猜来猜去靠谱得多。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑