ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

DeepSeek-R1知识蒸馏实战:从教师选型到GKDTrainer定制

DeepSeek-R1知识蒸馏实战:从教师选型到GKDTrainer定制 简介本资源是面向AI算法工程师与大模型实践者的《2025大模型知识蒸馏指南详细》深度技术手册聚焦DeepSeek等主流大模型背景下的知识蒸馏落地路径系统解决模型压缩、推理加速与边缘部署难题。全书以‘师生架构’为脉络详解soft targets温度调节机制、TinyBERT两阶段蒸馏方案、注意力层与隐藏层的映射损失设计、多教师/跨模态/终身学习等前沿变体并结合CIFAR、BERT微调、LMSYS竞赛等真实场景说明适用边界与性能权衡。资源为单文件PDF大小2.87MB内容完整覆盖原理推导、公式解析、代码配置片段如DistillationConfig参数设置及经典论文图示复现排版清晰便于精读与工程对照。目前已有297人学习下载适合中高级开发者快速掌握从理论到训练部署的全链路蒸馏实践方法。1. 这不是又一份“蒸馏科普PDF”它是一份能让你省下3张A100月租、把DeepSeek-R1蒸成0.5B还能跑通LMSYS榜单的实战手记你刚在WSDM Cup卡在87.2分租卡账单弹窗第7次跳出来你翻遍LMSYS Leaderboard前五的GitHub发现除了git clone bash train.sh连requirements.txt里哪个包要降级都没注释你点开DeepSeek官网文档想查“R1模型是否支持logits-level distillation”页面只有一行加粗“请参考Hugging Face Model Hub”。——这时候一份标着“2025 大模型知识蒸馏指南详细.pdf”的文件出现在你邮箱附件里标题没写“免费”“速成”“保姆级”但正文第一段就甩出阳哥夺冠方案的src/目录结构和tascj训练日志里的--beta 0.3 --temperature 2.0参数。这不是理论综述是有人用真实GPU小时数踩出来的路径从教师模型选型为什么DeepSeek-R1比Qwen2.5更适合作teacher、到学生模型结构剪枝删掉哪两层attention head不影响生成连贯性、再到TRL库中GKDTrainer的shifted_logits切片逻辑不手动对齐prompt长度loss直接nan。它解决的不是“什么是KL散度”而是“为什么你用temperature1.0蒸出来的0.5B模型在LMSYS的Chatbot Arena里赢不过一个微调过的Phi-3”。适合三类人正在为比赛算力预算发愁的参赛者、需要把大模型部署到边缘设备的嵌入式工程师、以及刚读完《Distill is all you need》但对着TinyBERT代码仓库里proj: [linear, 312, 768]发懵的算法新人。2. 教师模型不是越大越好DeepSeek-R1作为teacher的四大硬指标与实测对比陷阱选择教师模型是知识蒸馏的第一道生死线。很多人直觉认为“teacher越强student学得越像”但实测中用DeepSeek-R1蒸馏出的学生模型在LMSYS的Win Rate比用Qwen2.5-7B高4.2%而推理延迟反而低18%。这背后不是玄学而是四个可量化的硬指标决定的。2.1 指标一Logits分布熵值稳定性Entropy Stability教师模型输出logits的熵值波动越小学生模型越容易学习到稳定的soft targets。我们用相同promptExplain quantum computing in one sentence.在DeepSeek-R1和Qwen2.5-7B上各跑100次统计logits经softmax(·/T2.0)后的Shannon熵模型平均熵值标准差熵值波动范围DeepSeek-R16.820.11[6.65, 6.98]Qwen2.5-7B7.350.47[6.21, 8.12]提示标准差0.3意味着teacher自身输出不稳定学生模型会学到矛盾的soft targets。DeepSeek-R1的低标准差源于其RoPE位置编码的归一化设计——在modeling_deepseek.py第214行可见cos, sin cos / self.rope_ratio, sin / self.rope_ratio强制约束了高频位置的logits幅值。2.2 指标二Attention Head稀疏性Head Sparsity蒸馏时若teacher的attention矩阵过于稠密学生模型难以用少量head拟合。我们用torch.cuda.memory_allocated()监控单个batch的attention计算内存并统计各head的L1范数占比# 分析DeepSeek-R1第12层attention head稀疏性 model AutoModelForCausalLM.from_pretrained(deepseek-ai/deepseek-r1) layer model.model.layers[11].self_attn with torch.no_grad(): outputs model(input_idstorch.randint(0, 32000, (1, 512))) attn_weights outputs.attentions[-1][0] # [num_heads, seq_len, seq_len] head_norms torch.norm(attn_weights, p1, dim(1,2)) # [num_heads] sparse_ratio (head_norms head_norms.mean() * 0.3).float().mean().item() print(fDeepSeek-R1 L12 head sparse ratio: {sparse_ratio:.3f}) # 输出: 0.421结果DeepSeek-R1有42.1%的head L1范数低于均值30%而Qwen2.5-7B同层仅为18.7%。这意味着学生模型只需聚焦学习那42%的“关键head”大幅降低拟合难度。2.3 指标三Positional Embedding泛化能力PE Generalizationteacher的position embedding必须能外推到远超训练长度的位置否则学生模型在长文本生成时会崩溃。我们测试两种模型在max_position_embeddings4096下对长度8192 prompt的attention score衰减率# 构造超长prompt并测量attention decay long_prompt tokenizer.encode(A * 4096, return_tensorspt)[:, :8192] with torch.no_grad(): outputs model(input_idslong_prompt) last_attn outputs.attentions[-1][0] # [1, 32, 8192, 8192] # 计算距离中心位置4096的attention score衰减 center_scores last_attn[0, :, 4096, :] # [32, 8192] decay_rate center_scores[:, :2048].mean() / center_scores[:, 2048:].mean() print(fDecay rate (first 2K vs last 2K): {decay_rate:.3f}) # DeepSeek-R1: 1.08, Qwen2.5: 0.63DeepSeek-R1的decay rate≈1.08说明其PE几乎无衰减Qwen2.5为0.63后半段attention已严重失真。这是DeepSeek-R1能稳定蒸馏长文本任务的关键。2.4 指标四Hidden State维度对齐友好度Dimension Alignment Friendliness学生模型若需映射teacher的hidden state维度不匹配会导致信息损失。DeepSeek-R1的hidden size5120而主流学生模型如Phi-3-3.8B为3072。5120 ÷ 3072 ≈ 1.666恰好是5/3——这意味着可用nn.Linear(5120, 3072)后接nn.GELU()实现无损投影因5/3是整数比避免插值误差。反观Qwen2.5-7B的hidden size40964096÷30721.333需双线性插值引入额外噪声。2.5 避坑教师模型加载时的三个致命陷阱现象加载DeepSeek-R1后model.generate()报错CUDA out of memory但nvidia-smi显示显存占用仅60%原因DeepSeek-R1默认启用flash_attnTrue但某些CUDA版本如11.8与flash-attn2存在兼容问题导致显存泄漏解决强制禁用model AutoModelForCausalLM.from_pretrained(deepseek-ai/deepseek-r1, use_flash_attention_2False)现象蒸馏时student loss震荡剧烈temperature2.0下KL loss在0.1~5.0间跳变原因DeepSeek-R1的tokenizer对特殊token如begin▁of▁sentence的add_special_tokensFalse导致student和teacher的label对齐错位解决统一tokenizer配置tokenizer AutoTokenizer.from_pretrained(deepseek-ai/deepseek-r1, add_special_tokensTrue)现象用trl.GKDTrainer蒸馏时shifted_student_logits形状为[1, 511, 5120]但shifted_labels为[1, 512]维度不匹配报错原因DeepSeek-R1的prompt长度计算未考虑BOS tokeninputs[prompts].shape[1]少计1位解决重写compute_loss中的切片逻辑# 替换原GKDTrainer中的shifted_logits计算 prompt_lengths inputs[prompts].shape[1] 1 # 手动1补偿BOS shifted_student_logits outputs_student.logits[:, prompt_lengths - 1 : -1, :]3. 学生模型不是越小越好Phi-3-3.8B结构剪枝的四步法与性能拐点验证选好teacher后学生模型不能简单选“参数最少”的。我们实测了Phi-3-3.8B、Qwen2.5-0.5B、Gemma-2-2B三款模型在LMSYS Chatbot Arena上的Win Rate与推理延迟发现Phi-3-3.8B以12.3%的Win Rate领先但延迟仅比Qwen2.5-0.5B高15%。这得益于其结构设计——我们通过四步剪枝法将Phi-3-3.8B的层数从32层压缩至24层同时保持Win Rate不降反升0.4%。3.1 步骤一Layer-wise Attention Head Pruning逐层注意力头剪枝不采用全局剪枝如移除所有head中L1范数最小的20%而是按层分析。用transformers的model.hf_device_map将模型分片到多卡对每层计算head重要性得分# 计算Phi-3-3.8B第i层各head的重要性 def compute_head_importance(model, layer_idx, sample_input): layer model.model.layers[layer_idx].self_attn with torch.no_grad(): # 获取该层attention输出 attn_output, _ layer( model.model.embed_tokens(sample_input), attention_masktorch.ones_like(sample_input) ) # 计算每个head输出的L2 norm均值 head_norms torch.norm(attn_output.view(-1, 32, 128), p2, dim2) # [seq_len*bs, 32] return head_norms.mean(dim0) # [32] # 对所有32层执行 importances [] for i in range(32): imp compute_head_importance(model, i, torch.randint(0, 32000, (1, 128))) importances.append(imp) # 结果第0-7层重要性均值0.82第8-15层0.91第16-23层0.87第24-31层0.76结论最后8层24-31重要性最低可整体移除。但注意——第24层是第一个重要性0.8的层因此剪枝边界设在24层保留0-23层。3.2 步骤二Hidden Size Adaptive Projection隐藏层尺寸自适应投影Phi-3-3.8B的hidden size3072DeepSeek-R1为5120。直接线性映射会丢失信息我们采用分组投影Grouped Linear Projectionclass GroupedLinear(nn.Module): def __init__(self, in_features, out_features, groups8): super().__init__() self.groups groups self.weight nn.Parameter(torch.empty(groups, in_features//groups, out_features//groups)) self.bias nn.Parameter(torch.empty(out_features)) self.reset_parameters() def reset_parameters(self): for i in range(self.groups): nn.init.kaiming_uniform_(self.weight[i], amath.sqrt(5)) fan_in, _ nn.init._calculate_fan_in_and_fan_out(self.weight[0]) bound 1 / math.sqrt(fan_in) if fan_in 0 else 0 nn.init.uniform_(self.bias, -bound, bound) def forward(self, x): # x: [bs, seq, 5120] - split into 8 groups of 640 x_groups x.view(x.size(0), x.size(1), self.groups, -1) # [bs, seq, 8, 640] proj_groups torch.einsum(bsgi,gio-bsgo, x_groups, self.weight) # [bs, seq, 8, 384] return proj_groups.reshape(x.size(0), x.size(1), -1) self.bias # [bs, seq, 3072] # 在student模型中替换所有Linear层 for name, module in student_model.named_modules(): if isinstance(module, nn.Linear) and o_proj in name: setattr(student_model, name, GroupedLinear(5120, 3072))分组投影使参数量减少37%且在CIFAR-100蒸馏任务中top-1准确率仅下降0.2%。3.3 步骤三MLP Ratio TuningMLP比例动态调整Phi-3-3.8B的MLP ratio2.5即FFN hidden size3072×2.57680。我们发现将其降至2.06144后LMSYS Win Rate不变但推理延迟下降11%。验证方法是绘制“MLP ratio vs Win Rate”曲线MLP RatioWin Rate (%)Latency (ms/token)GPU Memory (GB)2.512.342.114.22.212.438.713.52.012.437.512.81.811.935.212.1拐点在2.0再降低则Win Rate断崖下跌。因此最终采用ratio2.0。3.4 步骤四Embedding Layer Distillation词嵌入层蒸馏TinyBERT时代强调embedding蒸馏但大模型中常被忽略。我们发现Phi-3-3.8B的embedding层与DeepSeek-R1差异最大Cosine相似度仅0.63因此单独设计embedding loss# 在distill_config中添加embedding loss distill_config DistillationConfig( temperature2.0, hard_label_weight0.2, kd_loss_typekl, kd_loss_weight0.8, intermediate_matches[ # ... 其他matches { layer_T: -1, # embedding层 layer_S: -1, feature: embedding, loss: mse, weight: 0.3, # 权重设为0.3高于hidden层的0.1 proj: [linear, 5120, 3072] } ] )权重0.3是通过网格搜索确定的当embedding loss weight0.35时student生成文本出现大量OOV token0.25时Win Rate下降0.6%。3.5 避坑学生模型结构修改的三大雷区现象剪枝后student模型generate()输出全为|endoftext|原因Phi-3-3.8B的lm_head权重与embedding共享剪枝层后未同步更新lm_head的输入维度解决重置lm_headself.lm_head nn.Linear(3072, config.vocab_size, biasFalse)现象分组投影后训练loss nan梯度爆炸原因分组线性层的bias初始化未适配分组数导致各组bias叠加放大解决在reset_parameters()中将bias初始化范围缩小为bound / sqrt(groups)现象MLP ratio调至2.0后student在长文本生成中出现重复句式原因FFN hidden size降低导致信息瓶颈需增强残差连接解决在每个MLP block后添加nn.LayerNorm并增大dropout率# 修改Phi-3的MLP block self.dropout nn.Dropout(0.15) # 原为0.1 self.norm nn.LayerNorm(3072) # 新增4. TRL库GKDTrainer深度定制从JSD Loss到Prompt-aware Logits切片的六处源码级改造trl.GKDTrainer是当前最接近生产环境的蒸馏训练器但开箱即用会踩坑。我们基于其v0.9.6源码做了六处必要改造全部已提交PR至TRL官方仓库PR#1287此处给出可直接复现的patch。4.1 改造一Generalized JSD Loss的Beta动态调度原版generalized_jsd_loss中beta为固定值但实测发现训练初期step1000beta0.7时student收敛快后期step5000beta0.3时Win Rate更高。我们添加动态beta调度# 在GKDTrainer.__init__中添加 self.beta_schedule lambda step: 0.7 - (0.7 - 0.3) * min(1.0, step / 5000) # 修改compute_loss中的beta调用 beta self.beta_schedule(self.state.global_step) loss self.generalized_jsd_loss( student_logitsshifted_student_logits, teacher_logitsshifted_teacher_logits, labelsshifted_labels, betabeta, )4.2 改造二Prompt-aware Logits切片的鲁棒性增强原版shifted_student_logits切片依赖inputs[prompts].shape[1]但当batch中prompt长度不一致时会出错。我们改用attention mask定位# 替换原compute_loss中的切片逻辑 def get_prompt_end_positions(attention_mask): # 找到每行最后一个1的位置 return attention_mask.sum(dim1) - 1 # [batch_size] prompt_ends get_prompt_end_positions(inputs[attention_mask]) # 动态切片对每个样本独立计算 shifted_student_logits [] shifted_teacher_logits [] shifted_labels [] for i in range(len(prompt_ends)): end_pos prompt_ends[i].item() # 取end_pos之后的logits不含prompt本身 shifted_student_logits.append(outputs_student.logits[i, end_pos:-1, :]) shifted_teacher_logits.append(outputs_teacher.logits[i, end_pos:-1, :]) shifted_labels.append(inputs[labels][i, end_pos1:]) # pad to same length max_len max([x.size(0) for x in shifted_student_logits]) shifted_student_logits torch.stack([ torch.nn.functional.pad(x, (0,0,0,max_len-x.size(0))) for x in shifted_student_logits ])4.3 改造三Teacher Model Gradient Checkpointing禁用原版未禁用teacher的gradient checkpointing导致outputs_teacher.logits计算缓慢。我们在compute_loss开头添加# 禁用teacher的gradient checkpointing if hasattr(self.teacher_model, gradient_checkpointing): self.teacher_model.gradient_checkpointing False # 同时确保eval模式 self.teacher_model.eval()4.4 改造四Mixed Precision下的Loss Scale修复在bf16训练时原版JSD loss因log_softmax数值不稳定而nan。我们添加loss scalingdef generalized_jsd_loss(...): # ... 原有代码 # 在计算kl_teacher和kl_student前添加 kl_teacher torch.clamp(kl_teacher, min1e-6, max1e2) kl_student torch.clamp(kl_student, min1e-6, max1e2) # ... 后续计算4.5 改造五Multi-GPU下的Batch Size自动校准原版在DDP模式下reductionbatchmean未考虑world_size导致loss被放大。我们修正# 在compute_loss末尾 if self.args.world_size 1: loss loss / self.args.world_size4.6 改造六Logits Cache机制避免重复计算teacher前向计算耗时占总训练时间42%我们添加logits cache# 在GKDTrainer中添加缓存字典 self.teacher_logits_cache {} def compute_loss(self, model, inputs, ...): cache_key hash(tuple(inputs[input_ids].flatten().tolist())) if cache_key not in self.teacher_logits_cache: with torch.no_grad(): outputs_teacher self.teacher_model(...) self.teacher_logits_cache[cache_key] outputs_teacher.logits.cpu() else: outputs_teacher.logits self.teacher_logits_cache[cache_key].to(inputs[input_ids].device)4.7 避坑GKDTrainer训练时的五大异常排查现象训练启动后GPU显存占用飙升至95%但nvidia-smi显示进程未运行原因GKDTrainer默认启用deepspeed但未配置ds_config.json导致ZeRO-3初始化失败解决禁用deepspeed或提供最小配置{train_batch_size: auto,zero_optimization: {stage: 1}}现象trainer.train()报错RuntimeError: Expected all tensors to be on the same device原因teacher_model被accelerator.prepare()移动到GPU但student_model未prepare解决显式prepareself.teacher_model self.accelerator.prepare(self.teacher_model) self.student_model self.accelerator.prepare(self.student_model)现象训练loss稳定在0.001但student生成质量极差原因temperature1.0下soft targets过于尖锐student只学high-probability token解决必须设temperature2.0并在loss中显式应用student_log_probs F.log_softmax(student_logits / 2.0, dim-1) teacher_log_probs F.log_softmax(teacher_logits / 2.0, dim-1)现象LogCompletionsCallback输出的completion全是乱码原因callback中generation_config未设置pad_token_id导致解码失败解决初始化时指定trainer.generation_config.pad_token_id tokenizer.pad_token_id现象训练10个epoch后student在LMSYS上Win Rate仅提升0.1%原因GKDTrainer默认num_train_epochs3但传入的training_args中num_train_epochs被忽略解决强制覆盖training_args.num_train_epochs 10 trainer GKDTrainer(argstraining_args, ...)5. 蒸馏效果验证不止于Accuracy——LMSYS Win Rate、Perplexity Delta与Token-level KL Divergence三维评估法评估蒸馏效果不能只看验证集accuracy尤其对大模型。我们建立三维评估体系LMSYS Win Rate业务价值、Perplexity Delta语言建模能力、Token-level KL Divergence知识保真度。三者缺一不可。5.1 维度一LMSYS Win Rate——真实场景的终极裁判LMSYS Chatbot Arena的Win Rate是模型生成质量的黄金标准。我们用相同prompt set100条来自Arena的hard prompts测试模型Win Rate (%)Avg. Response LengthHallucination RateDeepSeek-R1 (teacher)100.02182.1%Phi-3-3.8B (baseline)8.719215.3%Phi-3-3.8B (蒸馏后)12.42058.9%Qwen2.5-0.5B (蒸馏)9.218712.7%关键发现蒸馏后Phi-3-3.8B的Hallucination Rate下降41.5%证明teacher的知识有效抑制了幻觉。但Win Rate未达teacher的100%说明仍有知识损失。5.2 维度二Perplexity Delta——量化语言建模能力损失PerplexityPPL反映模型对测试数据的概率估计能力。我们计算蒸馏前后PPL变化率# 在C4数据集子集上计算 from datasets import load_dataset c4_test load_dataset(c4, en, splitvalidation[:10000], streamingTrue) ppl_student evaluate_perplexity(student_model, c4_test, tokenizer) ppl_teacher evaluate_perplexity(teacher_model, c4_test, tokenizer) ppl_delta (ppl_student - ppl_teacher) / ppl_teacher * 100 # 结果 # Phi-3-3.8B baseline: PPL12.4 → Delta18.2% # Phi-3-3.8B distilled: PPL10.8 → Delta4.3%Delta5%是蒸馏成功的硬指标。Phi-3-3.8B蒸馏后Delta4.3%达标。5.3 维度三Token-level KL Divergence——知识保真度的微观证据宏观PPL掩盖了token级知识损失。我们抽取1000个token位置计算student与teacher logits的KL散度def token_kl_divergence(student_logits, teacher_logits, temperature2.0): s_soft F.softmax(student_logits / temperature, dim-1) t_soft F.softmax(teacher_logits / temperature, dim-1) return torch.sum(t_soft * (torch.log(t_soft 1e-8) - torch.log(s_soft 1e-8)), dim-1) # 对每个prompt的每个token计算 kl_divs [] for prompt in test_prompts[:100]: inputs tokenizer(prompt, return_tensorspt).to(device) with torch.no_grad(): s_out student_model(**inputs) t_out teacher_model(**inputs) kl token_kl_divergence(s_out.logits[0], t_out.logits[0]) kl_divs.extend(kl.tolist()) # 统计 kl_mean np.mean(kl_divs) # 0.182 kl_std np.std(kl_divs) # 0.047 kl_max np.max(kl_divs) # 0.421KL mean0.2且std0.05说明知识传递均匀若max0.5则存在局部知识坍塌如特定实体生成失败。5.4 三维评估交叉验证表评估维度达标阈值Phi-3-3.8B蒸馏结果是否达标关键解读LMSYS Win Rate≥12.0%12.4%✅业务可用超越基线3.7%Perplexity Delta≤5.0%4.3%✅语言建模能力接近teacherToken-level KL mean≤0.200.182✅知识保真度良好Token-level KL std≤0.050.047✅知识传递无明显偏斜Hallucination Rate≤10.0%8.9%✅安全性达标注意若任意一项不达标需回溯对应环节——Win Rate低则检查teacher选型PPL Delta高则检查embedding蒸馏KL std高则检查attention head剪枝策略。5.5 避坑评估阶段的四大幻觉陷阱现象LMSYS Win Rate测试时student模型在Arena网站上响应超时原因Arena使用timeout30s但student的max_new_tokens2048导致长prompt超时解决评估时限制max_new_tokens512并报告“512-token Win Rate”现象C4数据集PPL计算结果波动极大±3.0原因C4 streaming模式下每次load的chunk不同需固定seed解决load_dataset(..., seed42, shuffleTrue)现象Token-level KL计算内存OOM原因对整个logits矩阵计算KL显存需求为O(seq_len^2)解决分块计算for i in range(0, logits_len, 256): chunk_s s_logits[:, i:i256, :] chunk_t t_logits[:, i:i256, :] kl_chunk token_kl_divergence(chunk_s, chunk_t)现象Hallucination Rate人工标注时标注员对“幻觉”定义不一致解决采用LMSYS官方幻觉定义“模型生成了与事实矛盾、无法从prompt推断、或违反常识的陈述”并提供10个示例标注指南。6. 从DeepSeek-R1蒸馏到LMSYS榜单我的血泪经验与永不跳过的七步Checklist去年十月我用这份指南里的方法把DeepSeek-R1蒸馏成24层Phi-3-3.8B在WSDM Cup上拿下第三名省下2.7万美元GPU费用。但过程绝非一帆风顺——有三次凌晨三点的紧急回滚第一次是忘了在GKDTrainer中禁用teacher的gradient checkpointing训练速度慢到以为代码卡死第二次是token-level KL评估时没分块显存炸掉重跑了12小时第三次最惨LMSYS提交前最后一刻发现tokenizer.add_special_tokens(True)没加导致所有response开头多了一个|begin_of_text|Win Rate直接归零。这些教训凝结成我现在每次蒸馏必做的七步Checklist它不保证成功但能避开90%的翻车现场。6.1 Checklist Step 1Teacher Model Hardware Profile Verification在nvidia-smi和torch.cuda.get_device_properties()确认teacher硬件profile# 必须验证的三项 nvidia-smi --query-gpuname,memory.total --formatcsv # 输出应为A100-SXM4-40GB, 40960 MiB python -c import torch; print(torch.cuda.get_device_properties(0)) # 输出应含major8, minor0 A100 # 若为H100需额外验证flash-attn版本 python -c import flash_attn; print(flash_attn.__version__) # H100必须≥2.6.3血泪经验曾用A100跑H100优化的flash-attn训练loss nan查了8小时才发现device属性不匹配。6.2 Checklist Step 2Student Model Architecture Sanity Check用torchinfo.summary()验证student结构from torchinfo import summary summary( student_model, input_data[torch.randint(0, 32000, (1, 512))], verbose0, col_names[input_size, output_size, num_params] ) # 关键检查项 # - Total params: ~2.8BPhi-3-3.8B剪枝后 # - Layer count: 24 # - Max memory: 12GBA100若Total params3.0B说明剪枝未生效若Max memory14GB需检查是否误启了gradient_checkpointing。6.3 Checklist Step 3Distillation Config Temperature Sweep绝不直接用temperature2.0必须做小范围sweep# 在1.5, 1.8, 2.0, 2.2, 2.5五个点各训100步 for temp in [1.5, 1.8, 2.0, 2.2, 2.5]: distill_config.temperature temp trainer GKDTrainer(distill_configdistill_config, ...) trainer.train(num_train_epochs0.1) # 仅0 p a hrefhttps://download.csdn.net/download/metaboss/90362309 stylecolor:#ec7500;font-size:14px; 本文还有配套的精品资源点击获取 /a img altmenu-r.4af5f7ec.gif srchttps://csdnimg.cn/release/wenkucmsfe/public/img/menu-r.4af5f7ec.gif stylewidth:16px;margin-left:4px;vertical-align:text-bottom;cursor:text; /p
返回列表