资讯详情

资讯详情

大模型知识蒸馏实战:TinyBERT、KL散度与TRL库全解析

简介这份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 作者实验发现 T1 表现最好但一般蒸馏场景 T 大于 1 效果更好。温度系数的作用是平滑概率分布T 越大softmax 输出的分布越平缓每个类别的概率值更接近从而暴露出教师模型对类别之间相似性的判断。这部分知识是 hard target 给不了的——hard target 只告诉学生“正确答案是哪个”而 soft target 告诉学生“哪些类别容易混淆教师模型认为它们的关联有多强”。2.3 蒸馏配置代码逐行解读TinyBERT 开源的蒸馏配置可以直接复用我拆过这段代码逐行看下来收获很大distill_config DistillationConfig( # 温度系数tiny-bert 作者用 1 表现最好一般大于 1 比较好 temperatureself.temperature, # hard label 损失的权重 hard_label_weightself.hard_label_weight, # 预测层蒸馏 losssoft label 损失用交叉熵并稍微放大其权重 kd_loss_typeself.kd_loss_type, kd_loss_weightself.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 rateoptimizer AdamW(self.student_model.parameters(), lrself.lr) train_config TrainingConfig( output_dirself.student_model_dir, deviceself.student_trainer.device, data_parallelself.enable_parallel, ckpt_frequencyself.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_configtrain_config, distill_configdistill_config, model_Tself.teacher_model, model_Sself.student_model, adaptor_Tsimple_adaptor, adaptor_Ssimple_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 是蒸馏TRLTransformer Reinforcement Learning库是 HuggingFace 出品的后训练工具库覆盖 SFT、PPO、DPO 等训练范式。这里只看两个 trainerSFTTrainer 和 GKDTrainer。SFTTrainer 是有监督微调训练器利用输入输出对数据通过最小化模型输出与真实标签之间的损失让模型适配到特定下游任务。它的损失函数通常就是交叉熵衡量模型预测和实际标注的差异。GKDTrainer 则是知识蒸馏训练器核心差异在损失计算上——它计算学生模型和教师模型输出之间的散度如 JSD、KLD让学生模型学习教师模型的输出分布。两个 Trainer 的继承关系很简单GKDTrainer 继承自 SFTTrainerSFTTrainer 继承自 transformers 的 Trainer。这意味着 GKDTrainer 天然拥有 SFTTrainer 的全部能力只是在 compute_loss 上做了重写。4.2 SFTTrainer 调用与 transformers Trainer 的损失函数逻辑SFTTrainer 的调用非常简单trl 的 readme 直接给了 demofrom trl import SFTConfig, SFTTrainer from datasets import load_dataset dataset load_dataset(trl-lib/Capybara, splittrain) training_args SFTConfig(output_dirQwen/Qwen2.5-0.5B-SFT) trainer SFTTrainer( argstraining_args, modelQwen/Qwen2.5-0.5B, train_datasetdataset, ) 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, labelsNone, beta0.5, temperature1.0, reductionbatchmean, ): # 温度缩放 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, dtypestudent_log_probs.dtype) mixture_log_probs torch.logsumexp( torch.stack([ student_log_probs torch.log(beta), teacher_log_probs torch.log(1 - beta) ]), dim0, ) # 分别计算两个方向的 KL kl_teacher F.kl_div(mixture_log_probs, teacher_log_probs, reductionnone, log_targetTrue) kl_student F.kl_div(mixture_log_probs, student_log_probs, reductionnone, log_targetTrue) # 广义 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再加权求和。beta0.5 时就是标准的 JSD。这样做的好处是提供了对称性不会出现前向 KL 那种“学生模型被迫去覆盖教师模型低概率区域”的问题。compute_loss 的整体流程是先让学生模型前向传播拿到 logits再让教师模型在 eval 模式、torch.no_grad() 下前向传播拿到教师 logits然后用 prompts 的长度做 logits 切片对齐最后调用 generalized_jsd_loss。def compute_loss(self, model, inputs, return_outputsFalse, num_items_in_batchNone): # 学生模型前向 outputs_student model( input_idsinputs[input_ids], attention_maskinputs[attention_mask], ) # 教师模型 eval 模式不计算梯度 self.teacher_model.eval() with torch.no_grad(): outputs_teacher self.teacher_model( input_idsinputs[input_ids], attention_maskinputs[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_logitsshifted_student_logits, teacher_logitsshifted_teacher_logits, labelsshifted_labels, betaself.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 蒸馏时我通常先试 T2.0看 loss 曲线再调。坑五教师模型和学生模型词表不一致。现象forward 时报 shape mismatch。原因两个模型用的 tokenizer 不同或者词表大小不一样logits 的最后一维对不上。解决统一 tokenizer或者做 logits 映射。sparse 词表映射可以在蒸馏之前先做一次 tokenizer 对齐验证assert teacher_tokenizer.vocab_size student_tokenizer.vocab_size, vocab size mismatch4.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( modelmodel_config.model_name_or_path, teacher_modeltraining_args.teacher_model_name_or_path, argstraining_args, train_datasetdataset[args.dataset_train_split], eval_datasettest_data, processing_classtokenizer, peft_configget_peft_config(model_config), ) completions_callback LogCompletionsCallback( trainer, trainer.generation_config, num_prompts8 ) 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 切片错位每个都让我多花了至少一个晚上的调试时间。希望这篇指南能帮你把这些坑提前避开。本文还有配套的精品资源点击获取
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →