知识蒸馏降本复盘:用 70B 教师模型训练 7B 学生模型的全流程

知识蒸馏降本复盘:用 70B 教师模型训练 7B 学生模型的全流程 知识蒸馏降本复盘用 70B 教师模型训练 7B 学生模型的全流程一、算力诅咒的解法70B 模型的效果7B 模型的成本业务侧使用 Llama-3-70B 构建的智能客服系统在回答质量上获得了业务方的高度认可人工评估满意度 91%。但每月 85 万的推理成本让预算难以持续。降本诉求明确将推理成本降到每月 15 万以内同时保持回答质量不低于当前的 85% 水平。一个直接的思路是切换到 7B 或 13B 模型。但直接部署开源的 7B 模型在业务方的内部测试集上满意度骤降至 62%根本不可用。知识蒸馏Knowledge Distillation提供了另一条路径利用 70B 模型的输出作为软标签来训练 7B 模型使得小模型学会模仿大模型的输出分布。二、蒸馏数据构建与损失函数设计蒸馏的质量高度依赖教师模型生成的训练数据质量。简单地用 70B 模型批量生成 QA 对是不够的——教师模型自身的错误会被传递到学生模型。引入了多轮过滤机制# 知识蒸馏数据生成 —— 多轮过滤确保训练数据质量 class DistillationDatasetBuilder: def __init__(self, teacher_model, diversity_threshold: float 0.85): self.teacher teacher_model self.diversity_threshold diversity_threshold self.seen_embeddings [] # 已生成样本的嵌入用于多样性过滤 def generate_qa_pair(self, seed_question: str) - dict: 生成一条蒸馏训练样本 1. 教师模型生成答案 2. 教师模型自评分一致性检查 3. 多样性过滤 4. 只保留高评分 高多样性的样本 # Step 1: 教师模型生成 Top-K 候选答案 candidates self.teacher.generate( seed_question, num_return_sequences5, # 生成 5 个候选 temperature0.8, # 适中温度兼顾多样性和质量 ) # Step 2: 教师模型自评一致性 —— 用同一个问题问两次看答案是否一致 answer1 self.teacher.generate(seed_question, temperature0.3) answer2 self.teacher.generate(seed_question, temperature0.3) consistency_score compute_semantic_similarity(answer1, answer2) # Step 3: 多样性检查 —— 避免蒸馏数据过于单一 answer_emb self._embed(candidates[0].text) if self._is_too_similar(answer_emb): return None # 丢弃避免学生模型过拟合到重复模式 # Step 4: 保留高评分样本 if consistency_score 0.8: return None # 教师模型自己都不确定不纳入训练 return { question: seed_question, teacher_answer: candidates[0].text, teacher_logits: candidates[0].logits, # 软标签 consistency: consistency_score }蒸馏的核心损失函数KL 散度 任务损失# 蒸馏损失函数 —— KL 散度软标签 交叉熵硬标签 def distillation_loss( student_logits: torch.Tensor, # 学生模型的输出 logits teacher_logits: torch.Tensor, # 教师模型的输出 logits软标签 labels: torch.Tensor, # 真实标签硬标签 temperature: float 4.0, # 蒸馏温度越高分布越平滑 alpha: float 0.7, # 软标签权重 ) - torch.Tensor: 总损失 α × KL_div(软标签) (1-α) × CE(硬标签) 温度参数 T 的作用 - T 越高教师输出的概率分布越平滑类间差异变小 - 平滑分布能传递更多的类间关系知识 - 实验发现 T4 在 70B→7B 蒸馏中效果最佳 # KL 散度损失让学生输出的平滑分布模仿教师的平滑分布 soft_student F.log_softmax(student_logits / temperature, dim-1) soft_teacher F.softmax(teacher_logits / temperature, dim-1) kl_loss F.kl_div(soft_student, soft_teacher, reductionbatchmean) # 恢复温度对梯度的影响梯度缩放因子 kl_loss kl_loss * (temperature ** 2) # 硬标签损失经典交叉熵防止学生模型完全偏离正确答案 ce_loss F.cross_entropy(student_logits, labels) return alpha * kl_loss (1 - alpha) * ce_loss三、实验数据与蒸馏效果蒸馏训练配置50 万条蒸馏数据LoRArank64微调单张 A100 训练 18 小时。评估维度70B 教师7B 基座7B 蒸馏蒸馏提升MMLU71.262.166.84.7GSM8K62.538.254.316.1业务满意度91%62%87%25%月推理成本85 万3 万3 万-96.5%首 Token 延迟2.1s0.4s0.4s-81%最显著的提升出现在数学推理16.1和业务满意度25%上。这说明教师模型传递的不仅仅是正确答案是什么更多的是从问题到答案的推理路径——这正是 KL 散度损失捕获的内容。四、蒸馏的适用边界蒸馏方案不是通用银弹有几个明确的边界条件教师模型必须足够强如果教师模型本身表现一般80% 满意度蒸馏的效果边际递减蒸馏数据量有甜点从 10 万条提升到 50 万条业务满意度从 78% 提升到 87%再加到 100 万条仅提升到 87.5%。投入产出比在 50 万条附近收敛创意生成类任务蒸馏困难对于创意写作、诗歌等开放式任务教师模型的输出多样性本身就不高蒸馏会进一步压缩多样性。五、总结知识蒸馏降本的核心经验蒸馏效果增量远大于直接微调7B 底座模型的满意度 62%SFT 微调可达 72%蒸馏可达 87%。蒸馏的相对收益25%是 SFT 的两倍以上教师数据质量 蒸馏数据量多轮过滤自评一致性 多样性检查丢失了约 40% 的候选数据但提升了蒸馏效果 8~12 个百分点。宁可少用高质量数据也不要用海量低质量数据蒸馏是最便宜的降本手段相较模型量化精度有损、投机解码架构改造大知识蒸馏只需训练成本投入不需要修改推理引擎上线风险最低α0.7 是经验上的最优软硬标签平衡点α 过低0.5蒸馏失去了软标签的类间知识传递效果α 过高0.9可能忽视硬标签的正确引导。适用场景推荐在 70B→13B/7B 的蒸馏路径上使用。70B→1B 的蒸馏会导致精度退化不可接受满意度降至 65% 以下。