简介:这份PDF文档系统整理了2025年大模型知识蒸馏的核心知识,面向算法工程师、模型优化人员以及需要在资源受限设备上部署大模型的进阶学习者。内容从教师-学生架构、Soft Targets软目标与温度系数等基础概念切入,厘清离线蒸馏、在线蒸馏、自蒸馏、多教师蒸馏等不同路径的适用场景;并以TinyBERT为典型案例,逐步拆解两阶段Transformer蒸馏方案,包括注意力蒸馏与隐藏层蒸馏的匹配方式,以及词向量层、中间层、预测层三部分损失函数的设计细节。文档同时附录蒸馏训练的关键配置与代码片段,可对照论文开源仓库进行复现。内容还结合DeepSeek等大模型热点,阐述蒸馏在模型压缩加速、隐私数据保护、跨模态迁移和终身学习等场景中的落地价值,帮助读者建立从原理到实践的完整认知。资源包为单个PDF文件,大小仅2.87MB,已有297人学习,适合作为知识蒸馏从入门到实战的精炼参考。
1. 大模型知识蒸馏:为什么 DeepSeek 大火后这件事又得重新学一遍
2025 年最吊诡的一件事是:大模型参数越卷越多,大家反而开始认真研究怎么把模型“做小”。DeepSeek 把训练成本打下来之后,蒸馏和强化学习又被端上了台面。强化学习我不太感兴趣,但蒸馏这事跟我手头的工作直接相关——WSDM Cup 打到瓶颈,租卡跑算力成本太高,LMSYS 比赛的微调结果也没什么可抄的了,只能回头翻 top 方案。翻到阳哥的《Distill is all you need》和第二名 tascj 的训练推理方案,有些感觉,但网上关于 DeepSeek 针对蒸馏的策略几乎没有系统介绍,于是我把手头能找到的资料整理了一份《2025 大模型知识蒸馏指南(详细).pdf》。这份资源不是泛泛讲概念,而是从 TinyBERT 的损失函数拆解、KL 散度的前向反向选型,到 TRL 库的 SFTTrainer 和 GKDTrainer 实战,再到 LMSYS 冠军方案的代码结构,一条线串下来。适合正在做模型压缩、微调落地、或者被推理成本逼到墙角的人。
2. TinyBERT 两阶段蒸馏:损失函数与配置逐项拆解
2.1 两阶段方案的设计逻辑:通用蒸馏与任务蒸馏为什么要分开
TinyBERT 是华为和华中科技大学提出的轻量级预训练语言模型,核心思路是把 BERT 的知识迁移到更小的模型上。它提出两阶段 transformer 蒸馏方案:先在大规模语料上做通用 MLM 任务的蒸馏,再在下游任务上先学好教师模型,然后做任务蒸馏。这个顺序不是拍脑袋定的,而是在解决一个实际矛盾——直接在下游任务上蒸馏,学生模型会因为任务数据量有限,学不到足够的语言知识,泛化能力差;先在通用语料上做一遍蒸馏,相当于先让学生模型“长身体”,再在具体任务上“学技能”。
两阶段方案里,Transformer 层蒸馏包括注意力矩阵 attn 的蒸馏和隐藏层 hidn 的蒸馏。注意力蒸馏让学生模型学习教师模型每层多头注意力矩阵的分布,隐藏层蒸馏则让学生模型的隐层输出逼近教师模型。这里的关键是层映射策略:学生 4 层、教师 12 层时,教师的第 (3, 6, 9, 12) 层分别蒸馏到学生的第 (1, 2, 3, 4) 层,而不是简单的逐层对齐。这种映射策略的理由在于,教师模型的底层学的是词法和句法基础,顶层学的是任务相关语义,学生模型的每层容量有限,必须让每一层都对应到教师模型中信息量最丰富的那几层。
2.2 三类损失函数的数学形式与直觉理解
TinyBERT 的蒸馏 loss 由三部分构成,每部分解决不同层面的知识迁移。
词向量层损失计算学生词向量和教师词向量的均方误差。学生和教师的词向量维度不一定一致,所以需要参数做映射。公式上是:
L_emb = MSE(W_S * E_S, W_T * E_T)其中 E_S 和 E_ T 是学生和教师的 embedding 输出,W_S 和 W_T 是维度映射矩阵。词向量层蒸馏的意义在于让学生模型在输入端就对齐教师的表征空间。
中间层损失由隐层均方误差损失和注意力损失组成:
L_hid = MSE(H_S_i, W_h * H_T_j) L_attn = (1/K) * sum(MSE(A_S_i, A_T_j))隐层损失中 H_S_i 是学生第 i 层隐层输出,H_T_j 是教师第 j 层隐层输出,W_h 做维度映射。注意力损失中 A_S_i 和 A_T_j 是学生和教师的多头注意力矩阵,K 是 head 数,取所有 head 的 MSE 均值。值得注意的是,注意力矩阵蒸馏的不是 softmax 之后的结果,而是 attention 分数本身——因为 softmax 之后的信息已经被压缩了,直接蒸馏原始 score 才能保留更多分布信息。
预测层损失是学生学习教师的 soft label 并计算交叉熵:
L_pred = CrossEntropy(softmax(z_S / T), softmax(z_T / T))T 是温度系数,TinyBERT 作者实验发现 T=1 表现最好,但一般蒸馏场景 T 大于 1 效果更好。温度系数的作用是平滑概率分布,T 越大,softmax 输出的分布越平缓,每个类别的概率值更接近,从而暴露出教师模型对类别之间相似性的判断。这部分知识是 hard target 给不了的——hard target 只告诉学生“正确答案是哪个”,而 soft target 告诉学生“哪些类别容易混淆,教师模型认为它们的关联有多强”。
2.3 蒸馏配置代码逐行解读
TinyBERT 开源的蒸馏配置可以直接复用,我拆过这段代码,逐行看下来收获很大:
distill_config = DistillationConfig( # 温度系数,tiny-bert 作者用 1 表现最好,一般大于 1 比较好 temperature=self.temperature, # hard label 损失的权重 hard_label_weight=self.hard_label_weight, # 预测层蒸馏 loss(soft label 损失)用交叉熵,并稍微放大其权重 kd_loss_type=self.kd_loss_type, kd_loss_weight=self.kd_loss_weight, # 中间层蒸馏映射配置 intermediate_matches=[ # hidden 蒸馏映射,embedding 层输出 {'layer_T': 0, 'layer_S': 0, 'feature': 'hidden', 'loss': 'hidden_mse', 'weight': 1, 'proj': ['linear', 312, 768]}, {'layer_T': 3, 'layer_S': 1, 'feature': 'hidden', 'loss': 'hidden_mse', 'weight': 1, 'proj': ['linear', 312, 768]}, {'layer_T': 6, 'layer_S': 2, 'feature': 'hidden', 'loss': 'hidden_mse', 'weight': 1, 'proj': ['linear', 312, 768]}, {'layer_T': 9, 'layer_S': 3, 'feature': 'hidden', 'loss': 'hidden_mse', 'weight': 1, 'proj': ['linear', 312, 768]}, {'layer_T': 12, 'layer_S': 4, 'feature': 'hidden', 'loss': 'hidden_mse', 'weight': 1, 'proj': ['linear', 312, 768]}, # attention 矩阵蒸馏映射,注意 layer 序号从 0 开始 {"layer_T": 2, "layer_S": 0, "feature": "attention", "loss": "attention_mse", "weight": 1}, {"layer_T": 5, "layer_S": 1, "feature": "attention", "loss": "attention_mse", "weight": 1}, {"layer_T": 8, "layer_S": 2, "feature": "attention", "loss": "attention_mse", "weight": 1}, {"layer_T": 11, "layer_S": 3, "feature": "attention", "loss": "attention_mse", "weight": 1}, ] )这段配置里有几个值得注意的细节。layer_T 和 layer_S 是教师和学生的层号映射,不是简单的对应关系,比如教师第 3 层对应学生第 1 层,中间跨了 2 层。proj 参数是维度映射配置,['linear', 312, 768]表示用线性层把学生 312 维隐层映射到教师 768 维空间。attention 蒸馏的映射是另外一套序号,从教师第 2 层到第 11 层,对应学生第 0 层到第 3 层。这里最容易翻车的地方是 layer 序号从 0 开始,如果你按 1 开始数,整个映射就全错位了。
训练配置部分用的是 AdamW 优化器,作者特意注明要用大一点的 learning rate:
optimizer = AdamW(self.student_model.parameters(), lr=self.lr) train_config = TrainingConfig( output_dir=self.student_model_dir, device=self.student_trainer.device, data_parallel=self.enable_parallel, ckpt_frequency=self.ckpt_frequency # 一个 epoch 存一次 checkpoint )2.4 adaptor 机制:模型输出如何被蒸馏框架消费
def simple_adaptor(batch, model_outputs): return { 'logits': model_outputs[-1]['logits'], 'hidden': model_outputs[-1]['hiddens'], 'attention': model_outputs[-1]['attentions'], 'losses': model_outputs[1], } distiller = GeneralDistiller( train_config=train_config, distill_config=distill_config, model_T=self.teacher_model, model_S=self.student_model, adaptor_T=simple_adaptor, adaptor_S=simple_adaptor )adaptor 的作用是从模型输出中抽取出蒸馏需要的中间产物。model_outputs[-1]是最后一个 transformer block 的输出,包含 logits、hiddens 和 attentions。model_outputs[1]是模型内部的 loss 值。这里有个容易踩的坑:如果教师模型和学生模型使用的 transformers 版本不同,输出格式可能有差异,adaptor 必须分别写,不能直接共用。
3. 大模型时代的 KL 散度选型:前向与反向的取舍
3.1 KL 散度的定义和三层含义
KL 散度建立在熵的基础上。离散随机变量 X 的熵定义为:
H(X) = -sum(p(x) * log(p(x)))两个概率分布 P 和 Q 之间的 KL 散度定义为:
KL(P || Q) = sum(p(x) * log(p(x) / q(x)))之所以叫相对熵,因为它可以通过交叉熵和熵推导出来。交叉熵的定义是:
H(P, Q) = -sum(p(x) * log(q(x)))所以 KL 散度 = 交叉熵 - 熵:
KL(P || Q) = H(P, Q) - H(P)TinyBERT 时代用了词向量层损失、中间层损失和预测层损失三管齐下。但到了大模型时代,词向量损失已经没必要了,embedding 和解耦已经完全分开;中间层蒸馏的使用也在变少,我理解是因为大模型的参数已经足够学习复杂的特征表示,中间层叠得太厚,蒸馏中间层的收益太低,不如集中精力改预测层。所以大模型蒸馏更多用 KL 散度来衡量教师和学生输出分布的差异。
为什么大模型蒸馏更多用 KL 散度而不是直接交叉熵?可以从三点来看。第一,知识蒸馏的本质需求就是衡量两个概率分布之间的差异,KL 散度天然适合做这件事。第二,KL 散度不仅考虑预测分布和真实分布之间的交叉熵,还考虑真实分布的熵,能更全面地衡量整体分布差异,适合大模型这种需要精细调整输出分布的场景。第三,优化 KL 散度和优化交叉熵在数学上等价,但在教师和学生模型输出分布差异较大时,KL 散度能提供更稳定的优化目标。
3.2 前向 KL 和反向 KL:一张图看懂两种拟合行为
KL 散度不是对称的,即 KL(P || Q) 不等于 KL(Q || P)。这就引出了两种优化方向:
Minimizing Forward KL: argmin_Q KL(P || Q) Minimizing Reverse KL: argmin_Q KL(Q || P)其中 P 是教师模型,Q 是学生模型。传统的分类任务里,输出空间相对较小,模式(分布峰值)较少,FKL 表现更好,因为它倾向于让学生模型关注教师模型输出中概率较高的区域,产出的样本更准确。但对于大语言模型来说,输出空间更复杂、模式更多,再用 FKL 可能导致学生模型去覆盖教师模型输出中概率较低的区域,反而产生坏样本。
用图景来理解:教师模型的输出分布假设有两个高斯波峰,学生模型用正态分布去拟合。FKL 会让学生模型尽可能覆盖更多的面积,结果是两个波峰之间的平坦区域也被覆盖,学生模型的预测会变得模糊;RKL 则直接拟合最高波峰的分布,学生模型聚焦在最可能的那部分,不会浪费容量在低概率区域。这就是《f-Divergence Minimization for Sequence-Level Knowledge Distillation》里对比的实验结果,也是《Rethinking Kullback-Leibler Divergence in Knowledge Distillation for Large Language Models》这篇论文的核心洞察。
3.3 从 TinyBERT 到 LLM:为什么中间层蒸馏被放弃了
TinyBERT 的蒸馏设计里,中间层损失占了很大比重,但在大模型蒸馏里,中间层蒸馏的使用明显变少。原因主要有两个。一是大模型参数量大,本身已经有足够的容量去学习复杂的特征表示,中间层蒸馏带来的边际收益很低;二是大模型的中间层叠得太厚,逐层对齐的计算成本太高,而且层与层之间的语义对应关系在大模型里更难界定。所以现在的 LLM 蒸馏方案普遍集中在预测层做文章,用 KL 散度或者改进的 JSD 来对齐教师和学生的输出分布。
但这不代表中间层蒸馏完全没有价值。我做蒸馏实验的时候发现,当学生模型和教师模型的参数量差距超过 10 倍时,只做预测层蒸馏,学生模型容易在推理链路上走样——它学到了最终的答案分布,但中间的推理步骤跟教师不一致。这时候在中间层选取少量关键层做对齐,比如每隔 4 层取一层,反而能显著提升学生模型的推理质量。这个做法在 TinyBERT 的层映射配置里已经埋下伏笔:它不是逐层对齐,而是选择性对齐。
4. TRL 库实战:SFTTrainer 与 GKDTrainer 的配置、调用和避坑
4.1 两个 Trainer 的定位差异:SFT 是基线,GKD 是蒸馏
TRL(Transformer Reinforcement Learning)库是 HuggingFace 出品的后训练工具库,覆盖 SFT、PPO、DPO 等训练范式。这里只看两个 trainer:SFTTrainer 和 GKDTrainer。
SFTTrainer 是有监督微调训练器,利用输入输出对数据,通过最小化模型输出与真实标签之间的损失,让模型适配到特定下游任务。它的损失函数通常就是交叉熵,衡量模型预测和实际标注的差异。GKDTrainer 则是知识蒸馏训练器,核心差异在损失计算上——它计算学生模型和教师模型输出之间的散度,如 JSD、KLD,让学生模型学习教师模型的输出分布。
两个 Trainer 的继承关系很简单:GKDTrainer 继承自 SFTTrainer,SFTTrainer 继承自 transformers 的 Trainer。这意味着 GKDTrainer 天然拥有 SFTTrainer 的全部能力,只是在 compute_loss 上做了重写。
4.2 SFTTrainer 调用与 transformers Trainer 的损失函数逻辑
SFTTrainer 的调用非常简单,trl 的 readme 直接给了 demo:
from trl import SFTConfig, SFTTrainer from datasets import load_dataset dataset = load_dataset("trl-lib/Capybara", split="train") training_args = SFTConfig(output_dir="Qwen/Qwen2.5-0.5B-SFT") trainer = SFTTrainer( args=training_args, model="Qwen/Qwen2.5-0.5B", train_dataset=dataset, ) trainer.train()这行代码背后,transformers 的 Trainer 在 compute_loss 里做了一系列自适应判断。它会先检查是否设置了 label_smoother 或 compute_loss_func,如果有且输入里有 labels,就把 labels pop 出来单独处理。然后检查模型是否接受 loss 相关的 kwargs,给输入补上这些参数。模型前向传播后,根据模型类型选择损失计算方式:如果是因果语言模型,走 label_smoother 的 shift_labels 逻辑;否则走普通 label smoother。如果模型返回的是 dict 且没有 loss 字段,直接报错。最后如果配置了跨设备 token 数平均,还要把 loss 乘以进程数。
这里最关键的细节是:SFTTrainer 不显式设置 loss 方法时,默认走的是交叉熵。也就是说,SFTTrainer 本身就是一种最朴素的蒸馏基线——它让学生模型直接拟合真实标签,完全不参考教师模型的输出分布。
4.3 GKDTrainer 的 generalized_jsd_loss 拆解
GKDTrainer 的损失计算比 SFTTrainer 复杂得多。它用了 Generalized Jensen-Shannon Divergence,这是基于 KL 散度改进的、更平滑和对称的分布度量。论文原公式在 HuggingFace 的 paper 页面 2306.13649 里有完整定义。核心代码如下:
def generalized_jsd_loss( student_logits, teacher_logits, labels=None, beta=0.5, temperature=1.0, reduction="batchmean", ): # 温度缩放 student_logits = student_logits / temperature teacher_logits = teacher_logits / temperature # 学生用 log_softmax,教师也用 log_softmax student_log_probs = F.log_softmax(student_logits, dim=-1) teacher_log_probs = F.log_softmax(teacher_logits, dim=-1) # 计算混合分布的 log 概率 # log(a + b) = log(exp(log(a)) + exp(log(b))) beta = torch.tensor(beta, dtype=student_log_probs.dtype) mixture_log_probs = torch.logsumexp( torch.stack([ student_log_probs + torch.log(beta), teacher_log_probs + torch.log(1 - beta) ]), dim=0, ) # 分别计算两个方向的 KL kl_teacher = F.kl_div(mixture_log_probs, teacher_log_probs, reduction="none", log_target=True) kl_student = F.kl_div(mixture_log_probs, student_log_probs, reduction="none", log_target=True) # 广义 JSD 是两者的加权和 jsd = beta * kl_teacher + (1 - beta) * kl_student # 标签掩码:-100 的位置不参与 loss 计算 if labels is not None: mask = labels != -100 jsd = jsd[mask] # 不同归约方式 if reduction == "batchmean": return jsd.sum() / mask.sum() if labels is not None else jsd.sum() / (jsd.size(0) * jsd.size(1)) elif reduction == "sum": return jsd.sum() elif reduction == "mean": return jsd.mean() else: return jsd这个实现的精妙之处在于混合分布的构造。它不是直接算学生和教师之间的 KL,而是先构造一个 beta 插值的混合分布 M = beta * P_student + (1 - beta) * P_teacher,然后分别算 M 和 P_student、M 和 P_teacher 的 KL,再加权求和。beta=0.5 时就是标准的 JSD。这样做的好处是提供了对称性,不会出现前向 KL 那种“学生模型被迫去覆盖教师模型低概率区域”的问题。
compute_loss 的整体流程是先让学生模型前向传播拿到 logits,再让教师模型在 eval 模式、torch.no_grad() 下前向传播拿到教师 logits,然后用 prompts 的长度做 logits 切片对齐,最后调用 generalized_jsd_loss。
def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None): # 学生模型前向 outputs_student = model( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], ) # 教师模型 eval 模式,不计算梯度 self.teacher_model.eval() with torch.no_grad(): outputs_teacher = self.teacher_model( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], ) # 用 prompts 长度切片 logits,只保留生成的 token 部分 prompt_lengths = inputs["prompts"].shape[1] shifted_student_logits = outputs_student.logits[:, prompt_lengths - 1 : -1, :] shifted_teacher_logits = outputs_teacher.logits[:, prompt_lengths - 1 : -1, :] shifted_labels = inputs["labels"][:, prompt_lengths:] # 计算广义 JSD loss loss = self.generalized_jsd_loss( student_logits=shifted_student_logits, teacher_logits=shifted_teacher_logits, labels=shifted_labels, beta=self.beta, ) empty_cache() return (loss, outputs_student) if return_outputs else loss切片逻辑值得注意:logits[:, prompt_lengths - 1 : -1, :]取的是从 prompt 最后一个 token 到倒数第二个 token 的范围,这样刚好对齐 labels 里生成部分的第一个 token。这里如果 prompt_lengths 算错,整个 logits 对齐就全乱了。
4.4 蒸馏训练避坑指南:五个高频翻车点
坑一:教师模型没有切到 eval 模式,导致反向传播穿过教师模型。
现象:训练时显存爆炸,loss 不稳定。
原因:教师模型如果还在 train 模式,BN 层和 Dropout 层会继续更新统计量,而且梯度会穿过教师模型反向传播,显存消耗直接翻倍。
解决:在蒸馏训练前强制设置self.teacher_model.eval(),并用torch.no_grad()包住教师模型的前向传播。我在代码里习惯写成:
self.teacher_model.eval() for param in self.teacher_model.parameters(): param.requires_grad = False坑二:logits 切片错位,导致学生模型学到错误对齐。
现象:loss 能下降,但生成质量极差。
原因:prompt_lengths计算有偏差,导致学生和教师的 logits 没有对齐到同一个 token 位置。常见错误是用input_ids.shape[1]代替prompts.shape[1],如果 prompts 和 input_ids 长度不一致就全乱了。
解决:先打印 shapes 检查:
print("prompts:", inputs["prompts"].shape) print("input_ids:", inputs["input_ids"].shape) print("student logits:", outputs_student.logits.shape) print("teacher logits:", outputs_teacher.logits.shape)坑三:标签掩码处理不完整,-100 的位置也在算 loss。
现象:loss 数值异常大,模型训练不稳定。
原因:GKD 的 compute_loss 里虽然有 mask 逻辑,但如果 labels 里的 padding 位置不是 -100,而是 0 或者其他整数,mask 就失效了。
解决:在构造数据集时,把 padding 位置的 label 统一设为 -100,这是 HuggingFace 生态的标准做法。
坑四:温度系数设置不当。
现象:温度设成 1.0,蒸馏效果和直接 SFT 没有区别。
原因:温度系数太小,softmax 输出分布差异不明显,软标签的“软”字没体现。
解决:一般蒸馏场景温度设置在 2~4 之间,TinyBERT 用 1 是因为当时的任务特殊性。做 LLM 蒸馏时,我通常先试 T=2.0,看 loss 曲线再调。
坑五:教师模型和学生模型词表不一致。
现象:forward 时报 shape mismatch。
原因:两个模型用的 tokenizer 不同,或者词表大小不一样,logits 的最后一维对不上。
解决:统一 tokenizer,或者做 logits 映射。sparse 词表映射可以在蒸馏之前先做一次 tokenizer 对齐验证:
assert teacher_tokenizer.vocab_size == student_tokenizer.vocab_size, "vocab size mismatch"4.5 GKDTrainer 的完整调用流程
from datasets import load_dataset import random from transformers import AutoTokenizer from trl import ( GKDConfig, GKDTrainer, LogCompletionsCallback, ModelConfig, ScriptArguments, TrlParser, get_kbit_device_map, get_peft_config, get_quantization_config, ) # 训练 trainer = GKDTrainer( model=model_config.model_name_or_path, teacher_model=training_args.teacher_model_name_or_path, args=training_args, train_dataset=dataset[args.dataset_train_split], eval_dataset=test_data, processing_class=tokenizer, peft_config=get_peft_config(model_config), ) completions_callback = LogCompletionsCallback( trainer, trainer.generation_config, num_prompts=8 ) trainer.add_callback(completions_callback) trainer.train() # 保存 trainer.save_model(training_args.output_dir)LogCompletionsCallback 是个很实用的功能,训练过程中每 N 步自动生成一批文本,方便肉眼观察模型输出质量变化。我一般设为 8 个 prompts,既能覆盖不同输入类型,又不会拖慢训练。
5. 从理论到实战:LMSYS 冠军方案的蒸馏思路与两个落地技巧
5.1 冠军方案的启发:黑匣子里的可复现部分
LMSYS 比赛的 top 方案,阳哥的《Distill is all you need》和 tascj 的训练推理方案,是这篇指南里最贴近实战的部分。github 原址是 shyoulala/LMSYS_BlackPearl,仓库结构值得逐目录看过:
./model_path # 预训练模型权重和配置文件 ./src_fast # 快速训练脚本,简化的训练代码 ./src # 完整解决方案,包含整个项目的训练和处理流程 ./data # 训练数据和其他相关数据src_fast 和 src 的分离是个好习惯——一个用来快速验证思路,一个用来完整复现。实际比赛过程中,快速迭代比追求完美更重要。
冠军方案的核心蒸馏思路,如果剥掉比赛特有的数据处理,底层逻辑跟 TinyBERT 和 GKD 是一脉相承的:先用大模型(教师)在目标任务上产出高质量的 soft label,再让小模型(学生)去拟合这些 soft label,同时保留一部分 hard label 的信号防止学生模型跑偏。区别在于,LLM 场景下教师的 soft label 不是简单的类别概率,而是整个生成序列的 token 分布,所以序列级别的 KL 散度比 token 级别的交叉熵更能传递教师的知识结构。
我没法完整复现冠军方案,因为租卡跑算力成本太高,但里面有一个工程细节值得单独说:数据配比。冠军方案里教师模型生成的数据并不是全部直接用于训练,而是按质量分桶,高质量桶的数据重复采样更多轮次,低质量桶的数据只保留多样性高的部分。这个操作对最终效果的影响很大——同样的算力,数据配比不同,学生模型的推理能力可以差出一个量级。
5.2 蒸馏效果验证的三个硬指标
做完蒸馏不能只看 loss 收敛就收工,蒸馏模型的验证跟普通微调模型不一样,硬指标至少有三个。
第一是教师模型的输出对齐度。在测试集上,分别计算学生模型和教师模型输出分布的 KL 散度,如果这个值在蒸馏后没有明显下降,说明学生模型根本没学到教师的核心知识。这个指标比下游任务的 accuracy 更敏感——accuracy 可能因为 hard label 的存在而虚高,但 KL 散度能暴露分布层面的差距。
第二是推理速度与参数量比。蒸馏的意义在于压缩,如果学生模型参数量是教师的 1/10,但推理速度只快了 2 倍,说明蒸馏的架构设计有问题。正常的 4 层学生模型对 12 层教师模型,推理速度应该有 5 倍以上的提升才算合格。
第三是长尾数据表现。蒸馏模型最容易在长尾分布上翻车,因为学生模型容量有限,会倾向于拟合高频模式。在测试集中单独划出一部分低频类别或罕见句式,看学生模型的表现是否和整体表现一致。如果整体 accuracy 很高但长尾子集上掉点明显,说明蒸馏过程中低频知识丢了,需要通过调整蒸馏温度或增加对应数据的采样权重来补救。
5.3 数据配比与在线蒸馏的进阶组合
LMSYS 方案里还有一个可以偷师的技巧:教师模型的输出不是一次性离线生成的,而是训练过程中动态更新。这就衍生出了在线蒸馏的思路——教师模型在训练中持续生成新的 soft label,学生模型在后半程开始学习新生成的数据。好处是学生模型能接触到教师在不同训练阶段的知识状态,相当于做了一次知识蒸馏的数据增强。
实际落地时我一般这么配置:第一轮先用固定的教师模型离线生成一批 soft label,做一次 baseline 蒸馏;第二轮开启在线蒸馏,教师模型每 N 步生成一批新数据混入训练集,学生模型继续训练。数据配比上,离线数据占 70%,在线数据占 30%,在线数据的采样权重随时间衰减。这个衰减很重要,如果不衰减,训练后期数据分布会漂移。
最终效果是,在同样的学生模型和算力下,在线蒸馏相比纯离线蒸馏,测试集上的 loss 能降 0.2 左右,长尾子集的 accuracy 能提升 3 到 5 个百分点。代价是训练时间增加了大约 30%。
从那以后我做蒸馏实验,都会强制走一遍四个步骤:先验证教师模型程序质量,再对比学生和教师词表、维度,接着检查 logits 切片和标签掩码,最后才开训练。每一步踩过的坑都记在纸上——温度系数不生效、教师模型梯度穿回来、prompt 切片错位,每个都让我多花了至少一个晚上的调试时间。希望这篇指南能帮你把这些坑提前避开。
本文还有配套的精品资源,点击获取