
简介面向深度学习开发者与研究工作者、聚焦使用Python语言在文本分类与语义理解等任务中完成教师模型到学生模型知识蒸馏的实战工程包完整覆盖从数据预处理、教师模型与学生模型搭建、学生模型拟合教师软标签概率分布的硬标签交叉熵与软标签KL散度联合损失定义到训练与评估的闭环。工程基于Hugging Face的Transformers库实践BERT、XLNet等预训练模型的压缩帮助读者快速搭建并运行完整的蒸馏训练流程压缩包整体仅926KB共包含32个文件以9个Python源码脚本为核心其中涵盖教师模型、学生模型与蒸馏主程序并辅以预训练模型文件、词表与模型配置、JSON与TXT数据文件、XML工程配置以及Markdown说明文档目录结构清晰便于按需查阅。目前已有470人学习并下载使用对照源码与配置文件可清楚理解蒸馏训练各环节的具体实现细节适合希望入手模型压缩、降低文本模型推理资源占用的中高级Python开发者也可作为高校相关项目或竞赛训练的参考。1. 用 Python 给文本模型做知识蒸馏先把“大模型的能力”压缩到能上线的大小在文本分类、实体识别、语义匹配这些场景里一个常见矛盾是教师模型效果够好但太重学生模型能上线但效果差一大截。用 Python 做知识蒸馏Knowledge Distillation思路是先让大模型在文本数据上输出“软标签”再把这份判断能力教给一个小模型从而让线上模型在参数量缩小十几倍之后尽量保住精度。它不是去压缩文件、也不是蒸馏词向量而是把模型在训练数据上学到的内在规律迁移到更小的结构里。适合手里已有训练好的文本大模型、但部署预算有限或者上线延迟要求很苛刻的团队。这篇文章会讲清楚原理、最小可运行流程、超参设置以及我在实际文本项目里踩过的坑。2. 知识蒸馏的原理与文本任务的三个关键选择从软标签到温度参数在动手写代码之前先要理解知识蒸馏在文本任务里到底迁移的是什么。很多初次接触的人把它理解为“让小学生抄大模型的正确答案”但如果只是抄答案那和用大模型标注数据再训练一个小模型没有本质区别。蒸馏真正有价值的地方在于让学生模型学习大模型在输出里暴露出来的“犹豫过程”。2.1 教师模型、学生模型与软标签先分清“知识”到底指什么对于文本分类来说一个训练好的教师模型在输入“等了一个小时还没上菜差评”之后输出的 logits 并不只是“负面高、正面低”这么简单。原始 logits 经过 softmax 后可能得到负面 0.7、中性 0.2、正面 0.1。这 0.2 和 0.1 在传统硬标签里会被直接抹掉因为在打标数据里这条样本的标签只有“负面”。但恰恰是这些非最高类别的概率包含了教师模型从大量语料中学到的语义相似性这家餐厅虽然很差但“差评”这种表达在语料中偶尔也会出现在中性或正面语境里。所以文本方向的知识蒸馏最核心的知识是教师模型的 logits 分布或者更准确地说是“软化后的概率分布”。对于序列标注任务知识是每个 token 在实体类别上的分布对于句向量和语义匹配任务知识是相似度分数或者排序关系对于生成式文本任务知识是解码器每个时间步对下一个 token 的预测分布。当你把教师模型换成部署时真正要用的小模型这些分布就是最值得教的“经验”。选教师和学生模型时我常用的做法是教师模型直接使用已经微调好的文本模型比如中文场景下的bert-base-chinese、roberta-base或更重的chinese-roberta-wwm-ext学生模型则优先选预训练小模型例如bert-tiny、albert-tiny或者直接按教师结构调整层数、hidden_size。需要特别注意的是学生模型的 tokenizer 尽量和教师保持一致。文本任务不像图像输入要经过 token 化换个 tokenizer 会导致序列切分方式不一致教师和学生看到的是同一句话的不同切法蒸馏效果自然跑偏。2.2 温度 T 怎么设分类、匹配、生成任务的不同区间温度参数是知识蒸馏里最像“玄学”但又有规律可循的一个超参。标准做法是在 softmax 之前把 logits 除以温度 Tsoft_prob_i exp(logit_i / T) / sum(exp(logit_j / T))T 越高分布越平滑模型的“犹豫”越明显T 越低分布越接近 one-hot学生能学到的知识越少。Hinton 那篇经典论文里给出的建议是从 3 附近开始试但对于文本方向任务这个值并不是固定的。以我自己的经验文本分类任务一般先设 T4然后按 1、2、4、8 扫一遍。分类数据的类别数通常在几十以内如果教师模型已经很自信logits 差距很大T4 能有效拉开中低概率类别之间的差异但如果 T 太大比如 10 甚至更高所有类别概率都被抹平学生等于在看一个“没有什么信息量”的均匀分布训练会非常慢甚至不收敛。反过来T1 时蒸馏退化成普通交叉熵学生只看到 argmax 标签学不到软知识。对于 NER 和序列标注任务情况略有不同。因为每个 token 都要输出一个实体类别分布类别数少且多数样本的标签是 O非实体此时模型输出很容易偏向 O 类。我一般会把 T 稍微调小到 2 到 4 之间避免把 O 类的噪声扩散给太多 token。对于文本匹配任务如果你在蒸馏一个 cross-encoder 到 bi-encoder温度更多作用在相似度得分上先对教师输出的相似度矩阵除以 T再做 softmax让每个样本在 batch 内的相对排序变得平滑。生成式文本任务则要谨慎解码器每一步都在预测词语分布词表往往很大T 过高会让学生生成一些语法正常但语义漂移的句子。常见的做法是先固定 T3但不直接比较之后的第一步生成结果而是观察学生模型在验证集上的困惑度变化。还有一个工程习惯保存教师 logits 时不要只保存已经除以温度后的概率。更稳妥的做法是保存原始 logits 的 npy 或 pt 文件之后扫描多个 T 时直接重放。这样一次教师推理可以反复使用不用因为调 T 重新跑一遍大模型。2.3 损失函数KL 散度与交叉熵的结合方式文本方向蒸馏最常见的损失函数是把两个部分加在一起学生模型对教师软分布的 KL 散度以及学生模型对原始硬标签的交叉熵。前者负责学知识后者负责兜底防止教师模型在某些样本上犯错时把学生带偏。一个可以减少实现对计算误差的小函数需要把温度缩放写进去import torch import torch.nn.functional as F def soft_cross_entropy(student_logits, teacher_logits, temperature4.0): # 学生 logits 先除以温度再做 log_softmax student_log_probs F.log_softmax(student_logits / temperature, dim-1) # 教师 logits 同样除以温度得到软目标分布 teacher_probs F.softmax(teacher_logits / temperature, dim-1) # batchmean 表示对 batch 维求平均 loss F.kl_div(student_log_probs, teacher_probs, reductionbatchmean) # 乘以温度的平方让 loss 量级和普通交叉熵接近方便调学习率 return loss * (temperature ** 2)这段代码里有两个关键点。第一学生侧必须用log_softmax而不是先softmax再取 logKL 散度的输入要求是 log 概率和概率这样数值更稳定梯度也能正确回传。第二temperature ** 2这个缩放因子非常关键。如果不乘回去温度一旦提高KL 的梯度会按比例变小相当于隐式降低了学习率乘回去之后不同温度下的梯度量级才能保持可比你调整学习率才有意义。如果你手头有标注数据就可以把两个损失组合起来def distillation_loss(student_logits, teacher_logits, hard_labelsNone, temperature4.0, alpha0.8): kd_loss soft_cross_entropy(student_logits, teacher_logits, temperature) if hard_labels is not None: ce_loss F.cross_entropy(student_logits, hard_labels) # alpha 越大越偏向软标签 return alpha * kd_loss (1 - alpha) * ce_loss # 完全没有标注数据时只学教师分布 return kd_lossalpha的经验区间在 0.5 到 0.9 之间。如果教师模型很强、标注数据质量也高可以把 alpha 放低一点让硬标签起更大作用如果蒸馏的大部分数据是未标注伪标签alpha 可以直接设到 0.9 以上甚至纯用软标签。你可能会在论文里看到一些自适应调 alpha 的方案但在实际文本项目里固定 0.8 起步往往已经够用。3. 用 Python 跑通一个文本蒸馏最小例子以文本分类为主线理论说得再多不如先把最小可复现的流程跑通。我以文本二分类为例演示从加载模型到完成一个 epoch 的完整代码。这里不会依赖某个特定平台只使用 PyTorch 和 Hugging Face Transformers这两者在文本方向已经成为事实标准。3.1 选型教师、学生和 tokenizer 怎么搭第一步是确定教师和学生模型。对于中文文本一个很常见的搭配是教师用bert-base-chinese学生用一个小型的bert-tiny或albert-tiny。但如果你的学生模型是从零随机初始化的需要大量数据和训练步数才能追平教师。更稳妥的做法是从一个已经预训练过的小模型开始再在上面蒸馏。具体到实现我会先用AutoConfig读教师配置然后按比例裁剪出一个较小的学生结构from transformers import AutoConfig, AutoTokenizer, AutoModelForSequenceClassification teacher_model_name bert-base-chinese student_config AutoConfig.from_pretrained(teacher_model_name) # 这里只做一个示例把隐藏层数从 12 层减到 4 层 student_config.num_hidden_layers 4 student_config.num_labels 2 # 学生模型的结构由 config 决定这里是随机初始化 student_model AutoModelForSequenceClassification.from_config(student_config) tokenizer AutoTokenizer.from_pretrained(teacher_model_name)这里我做了两件事。一是从教师配置里复制了词表、hidden_size、attention 头数等参数只改了层数和标签数这样学生的输入输出和教师完全对齐二是复用了教师的 tokenizer保证同一句话在教师和学生中切出的 token 序列完全一致。需要注意num_hidden_layers从 12 减到 4 之后隐藏层和注意力权重是随机初始化的所以后期训练步数不能太少。如果不想从随机初始化开始另一条路线是直接从 Hugging Face Hub 上拉一个已经预训练好的小模型只要它的词表覆盖你的语言即可。这样学生模型一上来就具备基本语言能力蒸馏更像是在做“领域适配和压缩”。3.2 构造蒸馏数据标注数据和未标注伪标签一起用数据是文本蒸馏的另一个大头。我一般会构造两条数据流一条是有标注的数据直接参与交叉熵损失另一条是更大量的未标注文本让教师模型先跑一遍产生软标签后喂给学生。常见做法是把未标注文本存成一个 UTF-8 的 CSV一列是文本内容后面可选一列 label如果没有 label也没关系教师模型输出的 logits 就是学生要学的目标。下面的代码演示如何准备数据集from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(bert-base-chinese) texts [ 这家餐厅服务很好菜品味道不错, 等了一个小时还没上菜差评, 性价比还可以就是位置有点偏, ] # 如果你有真实 label就放到这个列表里没有就全部填 None labels [1, 0, None] encodings tokenizer( texts, max_length128, paddingTrue, truncationTrue, return_tensorspt, ) print(encodings.keys()) # dict_keys([input_ids, token_type_ids, attention_mask])这段代码会输出一个 batch 的input_ids、token_type_ids和attention_mask。对中文 BERT 来说这三个字段都要用其中token_type_ids在句子对任务里尤其重要。padding 时我默认选max_length128如果你的文本很长建议先统计长度分布不要无脑往 512 填否则训练会慢很多。蒸馏数据的核心原则是教师模型见到的文本分布要尽量贴近线上真实输入。如果你之后要在客服场景上线却拿一堆新闻语料做蒸馏学生学到的软标签就带偏了。就算只有几千条未标注客服文本也好过几十万条不相关文本。3.3 训练脚本骨架先缓存教师 Logits再训练学生教师模型推理一次的成本通常很高在训练学生时不建议每步都重新跑教师。我的做法是先把教师的输出 logits 缓存下来再进入学生训练循环。这也能保证同一个 batch 的软标签是固定的不会因为教师模型被切到 eval 模式时产生随机 dropout 而抖动。先看教师 logits 缓存这一段import torch from torch.utils.data import TensorDataset, DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) teacher_model AutoModelForSequenceClassification.from_pretrained( bert-base-chinese, num_labels2 ) teacher_model.to(device) teacher_model.eval() # 把 tokenizer 结果搬到设备上 input_ids encodings[input_ids].to(device) attention_mask encodings[attention_mask].to(device) with torch.no_grad(): teacher_logits teacher_model( input_idsinput_ids, attention_maskattention_mask, ).logits.detach().cpu() # 把教师输出和学生输入打包成一个可以反复迭代的数据集 dataset TensorDataset( encodings[input_ids], encodings[attention_mask], teacher_logits, ) dataloader DataLoader(dataset, batch_size8, shuffleTrue)这段代码的关键在于teacher_model.eval()和torch.no_grad()。如果漏了eval()教师在推理时仍然会开 dropout同一句话每次跑出来的 logits 都不一样学生就要去拟合一个不断抖动的目标训练很难稳定。detach().cpu()的作用是把 logits 从计算图里拆出来放到 CPU 上否则即使是no_grad也会一直占着显存。接着是真正的学生训练循环from transformers import AdamW student_model.to(device) student_model.train() optimizer AdamW(student_model.parameters(), lr2e-5) for epoch in range(3): total_loss 0.0 for step, batch in enumerate(dataloader): input_ids batch[0].to(device) attention_mask batch[1].to(device) teacher_logits batch[2].to(device) student_logits student_model( input_idsinput_ids, attention_maskattention_mask, ).logits # teacher_logits 是已经缓存的原始 logits没有过温度 loss soft_cross_entropy( student_logits, teacher_logits, temperature4.0, ) loss.backward() optimizer.step() optimizer.zero_grad() total_loss loss.item() print(fepoch {epoch} loss {total_loss / len(dataloader):.4f})训练循环里只有一个损失函数因为当前数据集没有硬标签。如果后续加入带标签数据只需要把distillation_loss里的hard_labels传进去并把alpha调到 0.8 左右即可。学生模型的forward结果和教师一样都是SequenceClassifierOutput所以两者可以直接对齐。3.4 参数一览温度、alpha、batch_size 和学习率怎么联动最后给出一个我常用的参数表和取值逻辑。参数常用区间说明temperature1 到 8文本分类从 4 开始NER 用 2 到 4生成式任务建议从 3 开始alpha0.5 到 0.9有硬标签时 0.7 到 0.8纯无监督时直接 1.0batch_size16 到 32BERT 类模型按 8 到 16 起步显存不够就用梯度累积学习率2e-5 到 5e-5学生模型如果已经从预训练权重加载用 2e-5 更稳最大长度128 到 512以线上文本长度分布为准越长越慢训练轮数3 到 10小模型收敛快但要监控验证集避免过拟合学习率是文本蒸馏里最容易翻车的一个点。如果学生模型是完全随机初始化的直接用 2e-5 会让它很难收敛可以尝试 1e-4 到 3e-4 这种稍大的学习率如果学生模型是加载过的预训练小模型再用大学习率反而会把原来的语言能力冲掉。我的建议是先用较小的训练轮数跑通观察 loss 下降曲线再考虑扫学习率。4. 文本方向蒸馏的进阶用法从分类到 NER、句向量与生成式文本分类只是最简单的起点。在实际业务里你会遇到序列标注、文本匹配、生成式摘要或翻译等任务这些任务的蒸馏方法和分类有共性但各自有必须处理的细节。4.1 序列标注怎么蒸馏对齐 token 级 logits 与 ignore_indexNER 和文本分类最大的区别是教师和学生输出的 logits 维度变成[batch_size, sequence_length, num_labels]。每个 token 都要算一个 KL 散度但对 padding 位置必须屏蔽。import torch.nn.functional as F # 假设 teacher_logits 和 student_logits 都是 [B, L, num_labels] # labels 是 [B, L]padding 部分已由 tokenizer 填为 -100 def token_kd_loss(student_logits, teacher_logits, labels, temperature4.0): student_log_probs F.log_softmax(student_logits / temperature, dim-1) teacher_probs F.softmax(teacher_logits / temperature, dim-1) # 在类别维度上计算 KLreductionnone 方便手动 mask kl_loss F.kl_div( student_log_probs, teacher_probs, reductionnone, ) # kl_loss 现在是 [B, L, num_labels]对类别维求和 kl_loss kl_loss.sum(dim-1) # labels ! -100 的位置是有效 token mask labels ! -100 kl_loss kl_loss * mask return kl_loss.sum() / mask.sum()这里最容易踩坑的是labels的值。Hugging Face 的分词器在 padding 时不会自动把标签设为 -100你要自己做一次对齐。常见做法是先用tokenizer把文本 token 化再把原有的实体标签按照 offset mapping 映射到 token 级padding 后把标签列表补成 -100。如果漏了这个步骤模型会对 padding token 也计算实体损失损失值会被稀释学生最后倾向于把所有 token 都预测为 O。另外教师在 NER 上的软分布往往非常偏O 类概率占 90% 以上。这时就算做了 KL 散度学生也可能只学到“什么都不预测就是最安全的”。我建议在 NER 蒸馏时对教师 logits 先做一次温度缩放并且不要用太小的 T必要时可以把教师的非 O 类 logits 乘一个放大系数。4.2 句向量与文本匹配从 cross-encoder 蒸馏到 bi-encoder语义匹配场景中线上常用的做法是 bi-encoder把句子分别编码成句向量再做余弦相似度或内积。但在训练阶段效果更好的是 cross-encoder把两个句子拼成一个输入通过交互层判断相似度。cross-encoder 效果强但推理慢bi-encoder 快但效果弱。蒸馏在这里的迁移路径很明确把 cross-encoder 的知识教给 bi-encoder。做法是先对一批句子对构造样本包括正例、负例和 hard negative教师模型输出每个句子对的相似度 logits。学生模型分别编码句子得到句向量计算相似度。蒸馏损失可以有两种选择一种是对齐相似度分数直接算 MSE另一种是对 batch 内相似度矩阵做 KL 散度让学生学会相对关系import torch.nn.functional as F # teacher_sims: [batch, batch]表示句对之间的相似度分数 # student_sims: [batch, batch]由句向量内积得到 def sim_kd_loss(student_sims, teacher_sims, temperature2.0): # 先对相似度矩阵按行做 softmax得到“哪个样本最相似”的分布 teacher_prob F.softmax(teacher_sims / temperature, dim-1) student_prob F.log_softmax(student_sims / temperature, dim-1) return F.kl_div(student_prob, teacher_prob, reductionbatchmean) * (temperature ** 2)这段代码假设一个 batch 内每个句子和该 batch 内其他句子的相似度都能由教师计算。实际工程中教师是 cross-encoder需要在 batch 内两两拼接后前向显存开销是 O(N^2)。所以 batch 通常不能太大。还有一种更省显存的做法是只让教师给每组三元组打一个分数用 margin ranking loss 作为蒸馏目标的替代这种方式训练起来更快但软信息量变少了学生更容易丢掉细微的排序差异。4.3 生成式文本蒸馏解码分布和 teacher forcing 的坑生成式任务比如摘要、翻译、对话回复是做蒸馏时最复杂的一类。因为在生成模型里知识不仅存在于 encoder 输出也存在于解码器每一步对词表的预测分布。如果只比较最终生成的文本等于又把蒸馏退化成了论文里的“hard label”。一个最简化的生成式蒸馏训练流程如下先用教师模型对训练文本做一次 beam search 或 greedy 解码生成目标文本。然后把教师模型切换到 teacher forcing 模式输入原始文本和教师生成的结果得到每一步解码的 logits并缓存。学生模型用同样的输入重新前向与缓存的教师 logits 计算 KL 散度。# 伪代码示例假设 teacher 和 student 都是 Seq2Seq 模型 teacher_outputs teacher( input_idsinput_ids, decoder_input_idsteacher_decoded_ids, # 教师自己解码出来的 token ) teacher_decoder_logits teacher_outputs.logits student_outputs student( input_idsinput_ids, decoder_input_idsteacher_decoded_ids, ) loss soft_cross_entropy( student_outputs.logits, teacher_decoder_logits, temperature3.0, )这里最关键的细节是decoder_input_ids必须来自教师解码的结果而不是学生自己的输出。如果让学生用自己的前一步预测来生成下一步输入seq2seq 训练里的 exposure bias 会把偏差一路放大最终学生生成的文本可能迅速偏向重复词和退化解。很多新手在这里翻车以为把两个模型放到同一个训练循环里就行结果学生模型学了两步就输出空白。还有一个实际取舍生成式蒸馏的代价非常高因为教师每一步的 logits 都要缓存磁盘占用会很大。一个 30k 词表的 T5 模型每缓存一个 token 的 logits 就是 30k 个 float。所以工程上更常见的做法是先对少量高质量数据做蒸馏再配合强化学习或额外的交叉熵微调来修正学生的生成风格。4.4 隐藏层特征对齐当 logits 蒸馏不够时的下一级手段如果你试了温度、试了 alpha学生模型的结果还是差一截可以考虑让学生去模仿教师的中间层表示。这一步在分类任务里提升没有 logits 蒸馏明显但在句子对和 NER 任务里中间层特征往往包含更丰富的上下文信息。常见的做法是加一个回归损失让学生某一层的输出经过线性投影后去逼近教师对应层的输出import torch.nn as nn projection nn.Linear(student_hidden_size, teacher_hidden_size) student_hidden student_model.extract_hidden_state(text_batch) teacher_hidden teacher_model.extract_hidden_state(text_batch).detach() # 只对有效 token 对齐 mask attention_mask.unsqueeze(-1).float() mse_loss ((projection(student_hidden) - teacher_hidden) * mask).pow(2).sum() mse_loss mse_loss / mask.sum()需要说明的是教师和学生的层数往往不一致所以你要指定对应的层索引比如教师第 8 层对学生第 3 层。这个操作需要手动验证层语义是否对齐否则等于让模型去拟合两个分布差异很大的向量反而干扰分类头的学习。隐藏层对齐的开销比纯 logits 蒸馏大得多我一般只在学生模型退化到“明显说不出话”的时候才加。5. 文本蒸馏常见踩坑与排查效果不好时先查这四个环节这部分是我在维护多套蒸馏 pipeline 后总结出的血泪经验。下面每条都按照现象、原因、解决的顺序展开可以直接对照排查。5.1 学生比教师差太多先查温度、alpha 和训练步数现象学生模型在验证集上的准确率比教师低 10 个点以上甚至和随机初始化模型差不太多。 原因最常见的是温度设太高教师软标签被磨得很平学生看什么都是“五五开”其次是 alpha 太小交叉熵主导学生本质还在训硬标签没有学到软知识再就是训练步数不够小模型没有充分拟合教师的软分布。 解决把温度先降到 3 或 4alpha 提到 0.8训练轮数增加到 5 轮以上。先在小验证集上跑通一条曲线看 loss 是否持续下降。如果 loss 在下降但验证集不涨再考虑换更大的学生结构。5.2 教师模型自己都没训好蒸馏等于传错现象学生蒸馏结束后错误模式明显和教师一致教师判断错的样本学生也错在同一个地方。 原因不是学生的问题而是教师本身有系统偏差。比如教师模型是在旧的预处理逻辑下训练的而当前数据用的新分词器或者教师只在一个很小的数据集上微调过泛化不够。 解决蒸馏之前先评估教师模型的独立指标确保教师在目标语料上的表现达到可接受水平。如果教师有多个历史 checkpoint可以把它们的 logits 做一个简单平均减少单模型过拟合带来的偏置。我遇到最典型的案例是教师模型在训练时用了增强后的文本但无缝进了原始文本结果所有样本都被预测成同一个类。这种情况再调温度都没有用。5.3 有标注数据上蒸馏导致过拟合现象训练集 loss 很低验证集 loss 在第二个 epoch 就开始回升学生的验证集准确率反而不如不蒸馏直接训练的小模型。 原因有标注数据量太少教师给出软标签时也包含了自身对这几百条数据的过拟合噪声学生模型容量小但把这种噪声背下来了。 解决把部分标注数据换成大量未标注数据用教师软标签做自训练在损失函数里把 alpha 调高降低 hard label 的权重。同时给加入 weight decay、在 dropout 层之后多做一些随机 Mask 的数据增强。还有一个办法是只保留教师预测置信度较高的伪标签比如教师对某条样本最高类别概率小于 0.6就直接扔掉不要让学生去学模糊样本。5.4 类别不平衡被软标签放大现象学生在多数类上表现不错少数类基本全错验证集的整体准确率被多数类拉高看起来好像还可以。 原因教师模型在训练时已经受到类别不平衡影响多数类的 logits 绝对值偏高softmax 之后把少数类概率压得更低。学生再通过 KL 散度去对齐这个分布等于把不平衡又学了一遍。 解决训练时用带权重的采样器或者对少数类样本在蒸馏损失上乘一个大于 1 的系数。也可以把教师 logits 先做一次类别维度的 z-score 归一化再算软标签去掉类别本身的数值偏差。更简单直接的做法是把大多数类别样本的软标签温度调高一点让少数类的概率可见性更强。5.5 NER 和句子对里的 mask 不一致现象NER 模型蒸馏后所有 token 都预测成 O句子对模型的相似度在验证集上不升反降。 原因标签对齐漏了 padding token句子对蒸馏时没有把token_type_ids传给学生导致学生把两个句子的位置信息混在一起。还有一个常见原因是教师和学生的 tokenizer 不一致学生把“中”切成[中, 国]教师切成了[中国]logits 的位置根本对不上。 解决先用一小批数据打印教师和学生的input_ids与attention_mask人工检查对齐关系。训练脚本里 NER 专门用labels ! -100做 mask句子对任务确保学生 forward 时传入token_type_ids。如果师生的 tokenizer 不一致最好的办法是直接用教师 tokenizer 初始化学生的 embedding 层。5.6 显存不够导致教师推理经常中断现象教师模型单个 batch 推理就显存溢出训练进程中断Cache 文件没写完整。 原因教师模型隐藏层维度大logits 缓存又多放在 GPU 上持续不释放。 解决教师推理时把 batch 调小到 4 或 8同时用torch.no_grad()包住每次拿到 logits 后立刻detach().cpu()不要留在 GPU 上。如果数据量很大直接把 logits 分片写盘训练学生时按索引顺序读取。这个流程虽然麻烦但能让你在单卡环境中也能跑较大的教师模型。6. 蒸馏效果的验证方法与一个常用的提分技巧先压到中型再微调蒸馏代码跑通只是第一步验证“学生是否真的学到了”才是工程上线前最关键的环节。我的经验是不能只看训练集 loss也不能只看验证集准确率至少要做三件事。首先在同一个验证集上同时评估教师、学生、以及“不蒸馏直接微调学生”三份结果。这样才能公平判断蒸馏带来的增量。如果蒸馏学生和直接训练学生几乎一样说明教师软标签没有提供更多信息问题多半出在温度或数据构造上。其次要关注学生的预测置信度分布。文本分类中蒸馏学生往往会比直接训练的学生给出更平滑的预测分数这在低误报场景里是好事但在需要决策阈值时要重新校准阈值。最后如果是上线后的文本系统最好做线上 AB 对比重点看学生模型在长尾样本上的表现不要只看平均指标。我自己常用的一个提分技巧是“两阶段蒸馏”第一次先把大模型蒸馏到一个中等尺寸模型比如从 12 层压到 8 层第二次把 8 层模型蒸馏到真正上线的 4 层小模型。这个做法的好处是教师在每一步都离学生的容量更近学生学到的知识不会被“压得太碎”。实际操作时第一阶段的 alpha 可以设 0.6第二阶段再提高到 0.9。以下是我在验证阶段常用的一个评估函数def evaluate(model, dataloader): model.eval() correct 0 total 0 with torch.no_grad(): for batch in dataloader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) labels batch[label].to(device) logits model(input_idsinput_ids, attention_maskattention_mask).logits preds logits.argmax(dim-1) correct (preds labels).sum().item() total labels.size(0) return correct / total这个函数看起来简单但它能暴露一个很隐蔽的问题如果你把model.eval()漏了学生的 dropout 没关验证结果会上下跳动。我踩过这种坑折腾了半天调参数最后发现只是评估阶段忘了切换模式。我现在的习惯是把蒸馏脚本拆成三个独立阶段缓存教师 logits、训练学生、评估与导出。每个阶段之间通过磁盘文件连接这样任意一个阶段出问题都可以单独重跑不用每次从头开始。如果你正要开始做文本方向的知识蒸馏建议你先用几百条数据把这条链路跑通感受一下温度和 alpha 对结果的真实影响再决定要不要投入更多资源。希望这篇笔记能帮到你。本文还有配套的精品资源点击获取