
Soup 蒸馏训练数值稳定性修复在 FP32 对数空间计算 KL 散度损失与梯度【免费下载链接】SoupFine-tune LLMs from one YAML. Layer streaming trains an 8B model on a 4 GB laptop GPU.项目地址: https://gitcode.com/GitHub_Trending/soup12/Soup本篇文章深入解析 Soup 开源仓库中一项关键的蒸馏Knowledge Distillation训练稳定性修复——将 forward-KL、reverse-KL 与 Jensen-Shannon 三种散度统一迁移到 FP32 对数空间计算对应 changelog.d/0.75.0/736.fixed.md 记录的 #719 修复。你会理解低精度蒸馏损失为何会出现有限损失、非有限梯度的隐蔽故障掌握_compute_distill_term内核的 log-space 实现原理、非有限 logits 如何配合 AMP 步跳、以及 FP32 计算带来的时间与显存代价并能依据 docs/training.md 的完整配置范例在task: distill下复现这一行为。问题背景低精度蒸馏中的有限损失、非有限梯度知识蒸馏让学生模型student模仿冻结教师模型teacher的输出分布。Soup 的蒸馏训练器src/soup_cli/trainer/distill.py以 Hinton et al. 2015 的方式将训练损失组合为两部分的加权和loss 0.5 × CE(student) 0.5 × T² × D(teacher_logits / T || student_logits / T)其中 CE 是学生自身 logits 上的标准交叉熵T 为蒸馏温度D 为 token 级散度。仓库以 50/50 的固定比例混合两项源码中的_CE_WEIGHT 0.5、_DISTILL_WEIGHT 0.5。#719 修复针对的是一个隐蔽的数值稳定性缺陷在 FP16/BF16 低精度下低温度的有效 logits 很容易把概率下溢为 0而torch.kl_div在目标侧target side的导数在概率为 0 时会变成非有限值NaN/Inf。结果表现为归约后的 loss 是有限的但梯度是非有限的——训练看似正常推进实际参数更新已被污染。传统做法是加入逐 step 的有限值守卫finite-value guard来检测并修正但这种守卫需要在每个 step 同步设备既带来性能损耗也会掩盖问题的根源。修复方案三种散度统一迁移到 FP32 对数空间修复的核心思想对应 变更片段三种散度forward-KL、reverse-KL、JS一律在 FP32 中计算并对 FP16/BF16 输入返回 FP32 标量非有限 logits 直接传播用于 AMP自动混合精度的 GradScaler 步跳而不是用逐 step 守卫同步设备强行修正。官方文档 docs/training.md 明确解释了设计取舍The log-space reverse-KL and Jensen-Shannon formulas prevent finite losses with non-finite gradients; the upcast also preserves small losses that would underflow to zero. Non-finite logits propagate into the loss/gradients, allowing AMPs GradScaler to skip overflowed steps rather than aborting training with a validation exception.对数空间的 reverse-KL 与 JS 公式避免了有限损失伴随非有限梯度升精度还保住了会下溢为 0 的小损失非有限 logits 进入损失/梯度后AMP 的 GradScaler 可以跳过溢出的 step而不是让训练因校验异常而中止。源码级解析_compute_distill_term内核核心内核位于 src/soup_cli/trainer/distill.py#L50-L195。_compute_distill_term(student_logits, teacher_logits, divergence, temperature, labels, attention_mask, chunk_size, use_checkpoint)是一个纯张量内核输入为(batch, seq, vocab)的 logits输出为受限在训练 token 上的标量平均散度。1. 对齐KD 项与 CE 项测量相同的位置因果语言模型中位置 i 的 logits 预测 token i1因此 CE 项需要平移(logits[:, :-1] vs labels[:, 1:])。KD 项必须做同样的平移——否则训练 token 掩码labels ! -100会错位一个位置既丢掉每个 assistant 片段的首个预测 token又泄漏片段前的边界 token。源码在 L112-L118 统一完成对齐。2. 掩码只在训练 token 上度量散度掩码优先级为labels ! -100排除 padding 与 promptattention_mask排除 padding 全部位置。这与 CE 项的 ignore_index-100 语义保持一致L128-L144。若掩码全为假则直接返回 0 损失避免除零。3. FP32 升精度与 log-space 公式temp float(temperature) s student_logits.float() / temp t teacher_logits.float() / temp先在 FP32 中完成温度缩放再进入 log-space 计算L124-L126。三种散度都基于log_softmax与logaddexp推导全程不显式构造概率张量这也是 #736 相比初版内核avoid unused probability tensors的改进点forward_klp_t * (log_t - log_s)其中p_t log_t.exp()reverse_klp_s * (log_s - log_t)p_s log_s.exp()js利用log_m logaddexp(log_s, log_t) - log(2)在 log 空间直接构造混合分布避免先求概率再取对数带来的两次舍入然后0.5 * (kl(p_s||m) kl(p_t||m))。最终损失乘以temp * temp即 T² 温度缩放并转回 logits 的原始 dtype 返回L194-L195。4. Chunking 与激活检查点大词表如 Qwen 2.5 的 150k 词表加长序列会产生巨大的 logits 与概率中间张量。内核通过chunk_size按活跃响应 token 分块评估通过use_checkpoint使用torch.utils.checkpoint.checkpoint(..., use_reentrantFalse)非重入激活检查点在前向丢弃中间log_softmax张量、反向时重算L146-L192配套变更见 changelog.d/0.75.0/743.added.md 的 #722 工作。5. 非有限 logits 传播与 AMP 步跳与守卫修正相反修复后的内核不拦截NaN/Inf非有限值正常流入损失与梯度。由于关闭了逐 step 设备同步的守卫AMP 的 GradScaler 可以在溢出 step 后按既定机制缩小 scale 并跳过该 step训练得以继续而不是崩溃。这正是 tests/test_issue719_stable_distill_divergence.py 中test_non_finite_logits_propagate_to_amp的测试目标对三种散度、学生/教师任一侧注入 NaN 或 Inf 后断言 loss 非有限、梯度非有限——验证传播而非被修正。FP32 计算的代价时间、显存与硬件相关测量FP32 中间张量并非免费。docs/training.md 如实记录了维护者在 RTX 5070 上的测量BF16B1/S512/V32000前向加反向初版含守卫内核耗时是 main 分支的1.71 倍峰值显存是1.24 倍移除守卫后耗时从10.91 ms 降到 9.40 msmain 分支为 6.39 ms峰值显存从 376.5 MiB 降到 303.6 MiB两份 FP32 logits 副本本身在B4/S2048/V128256下就占用约7.8 GiB还需为额外中间张量预留预算。该文档特别注明这些测量与硬件强相关且未针对本次修订版重新运行。因此文章在引用时应如实说明这是初版守卫内核与 main 的对比而不是对修订版实现的基准测试。这提示大词表 长序列 大 batch 的蒸馏任务应显式配置 chunking 与检查点来压住峰值显存。配置指南如何在task: distill下使用完整配置示例摘自 docs/training.md 知识蒸馏章节base: HuggingFaceTB/SmolLM2-135M task: distill modality: text backend: transformers data: train: ./data/chat.jsonl max_length: 2048 chat_template: chatml training: teacher_model: meta-llama/Llama-3.1-8B distill_divergence: forward_kl # kl | forward_kl | reverse_kl | js distill_temperature: 2.0 distill_chunk_size: 256 # token chunk size for divergence evaluation distill_checkpoint: true # non-reentrant activation checkpointing epochs: 3 lr: 5e-5 quantization: 4bit # quantizes student only各字段语义与校验边界配置字段的 schema 定义与校验在 src/soup_cli/config/schema.py 及其引用的 src/soup_cli/utils/distill.py字段可选值 / 默认值说明teacher_modelHF repo id 或本地路径必填缺失时 schema 直接拒绝taskdistilldistill_divergencekl/forward_kl/reverse_kl/jskl是forward_kl的别名规范化后即标准蒸馏reverse_kl为 mode-seekingjs为对称散度distill_temperature默认2.0边界[0.05, 100.0]只接受有限值拒绝 bool/NaN/±infdistill_chunk_size正整数示例 256按活跃响应 tokenlabels ! -100分块评估散度中等值64~256兼顾峰值显存与 kernel 启动开销distill_checkpointbool默认false非重入激活检查点开启但未设 chunk_size 时全部活跃 token 单块处理保留字节下降但瞬时峰值不受限distill_modetoken默认/sequencesequence为跨 tokenizer 友好的硬标签 KD教师生成续写、学生纯 CE与uld_strategy互斥setup 阶段拒绝组合校验器层面的防护还包括distill_divergence拒绝 bool/空串/null 字节/超长16 字符上限并做大小写规范化teacher_model有 512 字符上限backendmlx不支持taskdistill。此外 src/soup_cli/config/schema.py 的模型级校验会保证这些distill_*字段只在taskdistill时被接受避免误配置静默失效。教师模型生命周期教师只加载一次在 setup 中冻结requires_grad_(False).eval()并在torch.no_grad()下前向以控制显存上界学生按标准 PEFT 流程套 LoRA。教师与学生可能落在不同设备如学生被 Trainer 移到 CUDA、教师留在 CPU训练器会自动桥接设备把输入搬到教师设备、把教师 logits 搬回学生设备distill.py#L740-L756。若教师与学生的词表大小不一致taskdistill会报错需要配置跨 tokenizer 的uld_strategy或改用distill_mode: sequence。回归测试如何锁定这一修复tests/test_issue719_stable_distill_divergence.py 从四个维度为本次修复提供可验证的证据低温有限性test_low_temperature_loss_and_gradients_stay_finite对float32/bfloat16/float16×reverse_kl/js×labels/attention_mask掩码做参数化在温度 0.05 下断言 loss 与梯度全部有限且被掩码的位置梯度严格为 0forward-KL 与概率空间参考一致test_forward_kl_matches_probability_space_reference以F.kl_div(..., reductionbatchmean) × T²为独立参考校验内核输出的数值等价性非有限传播test_non_finite_logits_propagate_to_amp三种散度 × 学生/教师两侧 × NaN/Inf断言 loss 与梯度保持非有限——这是 AMP 步跳机制的前提双精度参考test_divergence_matches_double_precision_reference用独立的双精度概率空间算术构造期望值断言内核输出 dtype 为float32且与参考在rtol1e-5, atol2e-7内一致。总结#719 修复把 Soup 蒸馏训练器从守卫修正转向数值空间根治三种散度在 FP32 对数空间计算既消除了低精度下有限损失、非有限梯度的隐蔽故障又保住了会下溢的小损失同时让非有限 logits 通过 AMP 步跳自然处理去掉了逐 step 的设备同步开销。从 变更片段、训练文档、内核实现 到 回归测试仓库给出了从设计决策、性能代价到行为契约的完整闭环。若你正在使用task: distill进行大词表蒸馏记得结合distill_chunk_size与distill_checkpoint控制 FP32 中间张量带来的显存开销。【免费下载链接】SoupFine-tune LLMs from one YAML. Layer streaming trains an 8B model on a 4 GB laptop GPU.项目地址: https://gitcode.com/GitHub_Trending/soup12/Soup创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考