DASH机制详解:发散度自适应监督时域,优化推理模型自蒸馏训练 📅 发布时间:2026/8/30 4:24:17 👁 浏览次数: 最近在做推理模型Reasoning Models的强化学习训练时我一直在思考一个问题当模型自己生成的长思维链Chain-of-Thought被当作训练信号时到底应该如何选择“监督范围”监督太短模型学不到长程推理监督太长训练又容易震荡甚至崩溃。这个问题的核心就是今天要聊的 DASH——Divergence-Adaptive Supervision Horizons for On-Policy Self-Distillation of Reasoning Models。如果你正在接触大模型推理优化、自蒸馏、PPO 策略训练或者想搞清楚 on-policy 与 self-distillation 之间到底是什么关系这篇文章值得收藏起来慢慢看。本文将围绕 DASH 的核心机制展开先讲清楚背景和基础概念再拆解发散度Divergence与监督时域Supervision Horizon的关系然后用伪代码实战演示 DASH 的思路最后分享训练调参过程中的注意事项与常见问题。无论你是刚入门 RLHF 的新手还是已经在调训练框架的算法工程师都可以从中找到可复用的经验。1. 背景与核心概念1.1 推理模型训练中的自蒸馏需求以 OpenAI o1 系、DeepSeek-R1 等为代表的一批推理模型让“让模型学会长思维链”成了大模型训练的关键目标。这类模型在推理时会在内部生成一段很长的思考过程经过反思、验证、回溯最终给出答案。训练这类模型时我们经常会碰到一个尴尬情况人工标注高质量思维链的成本极高而且人工思维链未必能覆盖模型的真实探索过程。于是自蒸馏Self-Distillation就成了一个重要思路——让模型自己生成推理轨迹然后从这些轨迹中提取监督信号再训练模型自身。简单说就是“用自己教自己”。在 Self-Distillation 的场景里教师模型和学生模型往往是同一个模型或者策略本身在不断更新。模型采样出一批推理轨迹我们挑出其中能在验证集上得到正确答案的轨迹作为训练目标再通过策略梯度或监督微调更新参数。1.2 什么是 On-Policy为什么 PPO 是 On-Policy 算法自蒸馏训练如果走策略优化路线就需要先搞清楚 on-policy 和 off-policy 的区别。很多初学者在学 PPO 时都会问为什么说 PPO 是 on-policy 的on-policy 的含义是训练时使用的数据必须来自当前策略也就是“正在学习的这版参数”最新采样出来的交互数据。PPO 的目标函数是$$L^{CLIP}(\theta) \mathbb{E}_t[\min(r_t(\theta)\hat{A}_t, \operatorname{clip}(r_t(\theta), 1-\epsilon, 1\epsilon)\hat{A}_t)]$$其中 $r_t(\theta) \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{old}}(a_t|s_t)}$$\pi_{\theta_{old}}$ 是采样时使用的旧策略。虽然 PPO 允许用重要性采样比例修正新旧策略差异但它本质上依赖当前策略采样的数据来优化策略所以 PPO 被归类为 on-policy 算法。相比之下DQN、SAC 这类算法可以复用历史经验池里的数据属于 off-policy。为什么这个区分很关键因为在自蒸馏训练中如果模型已经更新了几步旧的推理轨迹就不再是当前策略的产物。继续拿旧轨迹训练相当于在用不匹配的分布做监督容易导致训练偏差。1.3 Supervision Horizon监督时域是什么Supervision Horizon监督时域指的是在一次训练中我们决定用生成轨迹的哪一段作为监督信号。比如模型生成了 5000 个 token 的推理轨迹我们可以只监督最后 500 个 token靠近答案的部分监督从某个中间反思点之后的所有 token监督整条轨迹从第一个推理 token 到最后一个答案 token。这个“时域”的选择会直接影响模型的学习效果。时域太短模型相当于只学了一个局部片段缺少完整的推理逻辑时域太长模型需要在很长的范围内持续保证推理正确性优化难度大幅提升而且发散度较高的时候长程监督特别容易导致梯度噪声过大。DASH 的关键创新点就是把“监督时域”从一个固定值改成一个可以根据策略发散度Divergence动态调整的自适应值。2. DASH 机制拆解发散度如何决定监督时域2.1 核心思想发散度高的地方少教发散度低的地方多教DASH 的核心逻辑可以概括成一句话当前策略的分布发散度越高就缩短监督时域发散度越低就延长监督时域。直观理解是策略发散度高说明模型还在剧烈探索此时生成的轨迹不够稳定如果强行用整条长轨迹来监督会放大噪声训练容易震荡。这种情况下应该“少教一点”只监督相对可靠的局部片段。策略发散度低说明模型已经比较稳定轨迹质量较高此时可以“多教一点”把监督窗口扩大充分利用长轨迹中的推理结构。这种策略很像课程学习Curriculum Learning先让模型在短窗口内学会稳定表达再逐步扩展到长窗口、长推理链条。2.2 Divergence如何度量策略发散度在 DASH 的语境中发散度通常指的是当前策略与某个参考策略或旧策略之间的分布差异。常见的度量方式包括KL 散度$KL(\pi_\theta || \pi_{ref})$Token-level probability ratio当前策略生成某个 token 的概率与旧策略概率的比值动作分布的方差或者是策略熵的倒数。实际工程中最简单的方法是对轨迹中小批量 token 的 KL 散度取均值作为“发散度指标”。然后引入一个滑动窗口统计比如用指数移动平均EMA平滑发散度波动。这里给出一个简单的发散度计算思路import torch import torch.nn.functional as F def compute_kl_divergence(log_probs_current, log_probs_ref): 计算当前策略与参考策略的 KL 散度。 # log_probs_current: [batch_size, seq_len] # log_probs_ref: [batch_size, seq_len] return (log_probs_ref.exp() * (log_probs_ref - log_probs_current)).mean()需要注意的是实际训练中通常直接在 token 级别计算损失发散度可以按轨迹片段聚合也可以按全局 batch 聚合。DASH 论文中更关注“逐片段”的调节粒度因为推理轨迹里不同段落的风险差异很大。2.3 监督时域的自适应调节方式DASH 并不会每次重新生成一个随机时域而是设计了一个根据发散度平滑调整监督长度的机制。其大致流程是当前策略采样一段完整推理轨迹。将轨迹按语义或位置切分成多个监督片段。计算每个片段或整个轨迹的发散度。设定一个发散度容忍阈值如果发散度超过阈值监督时域缩短如果低于阈值监督时域延长。只对“被选中的时域窗口”内的 token 计算监督损失。用一个简化的 Python 伪代码来表达def adaptive_supervision_horizon(divergence, horizon, threshold, max_horizon, step1): 根据发散度动态调整监督时域。 divergence: 当前策略发散度 horizon: 当前监督时域长度 threshold: 发散度阈值 max_horizon: 最大时域长度 if divergence threshold: horizon max(horizon - step, 1) else: horizon min(horizon step, max_horizon) return horizon这只是最朴素的实现思路。更复杂的方案可以为发散度设置上下两个阈值形成滞回区间避免时域频繁抖动。也可以将发散度归一化后作为比例系数直接映射到[min_horizon, max_horizon]例如def linear_adaptive_horizon(divergence, min_h, max_h, d_min, d_max): ratio (divergence - d_min) / (d_max - d_min 1e-8) ratio min(max(ratio, 0.0), 1.0) horizon int(min_h (max_h - min_h) * (1.0 - ratio)) return horizon发散度越大时域越短发散度越小时域越长。2.4 DASH 与固定时域的关键差异传统固定时域训练Fixed Horizon只会设置一个常量完整轨迹从头到尾都用同样的监督强度。这会导致两个问题模型探索初期策略不稳定固定长时域容易把坏 token 的梯度放大训练不稳定模型收敛后期轨迹质量已经稳定固定短时域会浪费大量高质量长轨迹信息。DASH 本质上是把“时域选择”从一个超参数变成了一种自适应控制信号让训练过程动态匹配当前策略的探索状态。这也是“发散度自适应监督时域”这个名字的由来监督时域不是固定长度而是由发散度驱动的动态值。3. 为什么长思维链任务特别适合 DASH3.1 长轨迹中的“远端监督”信号稀疏在生成式推理任务中奖励往往只在最终答案处给出。整个思维链从第一步到最后一步可能跨越几千个 token但真正能直接反映正确性的信号只有最后那一个答案标签。这种“远端监督”remote supervision天然就有稀疏性的问题。如果只用最终答案做强化学习奖励中间的推理步骤很难得到有效指导如果整条轨迹都用自回归最大似然训练又容易让模型记住探索期的大量错误路径。DASH 通过动态时域控制相当于把远端监督“切”成多个近端监督片段先保证局部片段可靠再逐步扩展到远端。3.2 推理轨迹中不同段落的风险不同一条完整的推理轨迹往往包含多个阶段初始理解阶段模型解读问题提取关键条件。中间推导阶段模型执行逻辑运算、方程求解、代码模拟。反思修正阶段模型自查发现错误并重试。最终答案阶段模型汇总结果并输出。不同阶段对错误容忍度完全不同。中间推导一旦错了一步后面再长也可能全错最终答案阶段则是决定成败的关键。如果固定时域很多训练 signal 会被早期错误污染。DASH 的潜在优势在于它的发散度信号能捕捉“模型是否已经充分探索/稳定”从而决定要不要把监督焦点往后移动。3.3 与 Self-Distillation 的结合Self-Distillation 的一个难点是教师信号来自模型自身如果模型自身还在剧烈变化教师的输出也会剧烈变化学生学起来就很分裂。DASH 在一定程度上缓解了这个问题当策略发散度很高时监督时域自动缩短模型只从相对可信的片段中学习当策略收敛、发散度降低时监督时域自动延长让学生接触更完整的推理链。这正是“on-policy self-distillation”的组合逻辑数据来自当前策略监督范围也由当前策略的稳定性动态决定。4. 完整实战简化版 DASH 训练循环考虑到直接用大模型跑完整训练不现实下面我用 PyTorch 风格的代码展示一个简化版的 DASH 训练循环。这段代码不是完整可上线的训练脚本而是把 DASH 的核心机制单独抽出来方便理解。4.1 定义采样轨迹的数据结构import torch import torch.nn.functional as F from dataclasses import dataclass from typing import List, Optional dataclass class Trajectory: 一次采样得到的一条完整推理轨迹。 tokens: torch.Tensor # [seq_len] log_probs: torch.Tensor # [seq_len]当前策略对每个 token 的 log 概率 ref_log_probs: torch.Tensor # [seq_len]参考策略或旧策略的 log 概率 reward: float 0.0 # 最终奖励例如答案是否正确在实际训练中采样与优化通常交替进行当前策略采样 N 条轨迹然后利用这些轨迹计算损失并更新之后再次采样。4.2 计算轨迹发散度def compute_trajectory_divergence(traj: Trajectory) - float: 基于 token 级 KL 散度计算轨迹发散度。 这里简单使用平均 KL实际可加权或分片段计算。 log_p traj.log_probs log_ref traj.ref_log_probs kl (log_ref.exp() * (log_ref - log_p)).sum(dim-1) return kl.item()这里用逐 token 平均 KL 作为发散度指标。实际工程中可以换用F.kl_div但要注意输入输出顺序参考概率作为 weight 会更稳。4.3 实现 DASH 自适应时域控制器class DASHHorizonController: 基于发散度调整监督时域的自适应控制器。 def __init__(self, init_horizon: int, max_horizon: int, threshold: float): self.horizon init_horizon self.max_horizon max_horizon self.threshold threshold self.ema_divergence None self.alpha 0.9 # EMA 平滑系数 def update(self, divergence: float) - int: # 对发散度做指数滑动平均防止单条轨迹噪声过大 if self.ema_divergence is None: self.ema_divergence divergence else: self.ema_divergence ( self.alpha * self.ema_divergence (1 - self.alpha) * divergence ) # 自适应调节发散度高于阈值监督时域缩短反之延长 if self.ema_divergence self.threshold: self.horizon max(self.horizon - 1, 1) else: self.horizon min(self.horizon 1, self.max_horizon) return self.horizon这套控制器逻辑简单适合作为理解 DASH 的起点。如果你想在真实项目中体现更平滑的调节可以把step1改成基于发散度偏差的连续更新例如step int(k * (divergence - self.threshold)) self.horizon min(max(self.horizon - step, 1), self.max_horizon)4.4 DASH 训练主循环下面给出一个简化版训练主循环重点展示“采样on-policy→ 计算发散度 → 调节监督时域 → 计算监督损失 → 更新策略”的流程。def train_dash_step( policy_model, ref_model, tokenizer, prompt_ids, controller: DASHHorizonController, optimizer, device, max_new_tokens: int 512, ): 简化版 DASH 训练步骤。 真实场景还需处理 reward model、critic、clip 等细节。 # 1. 使用当前策略采样轨迹 with torch.no_grad(): outputs policy_model.generate( prompt_ids, max_new_tokensmax_new_tokens, return_dict_in_generateTrue, output_scoresTrue, temperature0.8, ) generated_ids outputs.sequences logprobs compute_logprobs(policy_model, prompt_ids, generated_ids) ref_logprobs compute_logprobs(ref_model, prompt_ids, generated_ids) traj Trajectory( tokensgenerated_ids[0], log_probslogprobs, ref_log_probsref_logprobs, reward0.0, # 这里应通过真实任务校验答案得到 ) # 2. 计算发散度调节监督时域 divergence compute_trajectory_divergence(traj) horizon controller.update(divergence) # 3. 只取最后一小段作为监督窗口 # 以中间为界监督窗口取轨迹末尾的 horizon 个 token supervision_start max(0, traj.tokens.shape[-1] - horizon) # 4. 构造监督损失这里用简单的正例监督取奖励为正的轨迹 if traj.reward 0: # 仅对监督窗口内的 token 计算交叉熵 logits policy_model(traj.tokens[:-1]).logits shift_logits logits[supervision_start - 1:-1] shift_labels traj.tokens[supervision_start:] loss F.cross_entropy( shift_logits.reshape(-1, shift_logits.size(-1)), shift_labels.reshape(-1), ignore_index-100, ) optimizer.zero_grad() loss.backward() optimizer.step() else: # 负样本轨迹可以跳过或使用策略梯度 / 负例抑制损失 loss None return loss, divergence, horizon注意上面这段代码省略了compute_logprobs的完整实现因为不同框架的实现方式差异较大。核心是理解这个循环先用当前策略采样再根据发散度调节监督窗口最后只对窗口内 token 监督。4.5 在强化学习框架中的位置如果你使用 PPO 训练推理模型DASH 可以集成到 loss 中。常见的做法是把 Mask 机制作用在 PPO 的 advantage 或 policy loss 上。比如# 假设 masks 形状为 [batch, seq_len]1 表示该 token 参与监督0 表示忽略 policy_loss -torch.minimum( ratio * advantage, torch.clamp(ratio, 1.0 - eps, 1.0 eps) * advantage ) masked_policy_loss (policy_loss * masks).mean()DASH 中的supervision horizon决定的就是masks中哪些位置是 1哪些位置是 0。5. DASH 与 ReAct 等“推理行动”范式的联系如果你经常关注大模型 Agent 方向可能会看到热词react: synergizing reasoning and acting in language models。ReAct 的核心思想是让语言模型交替进行推理Reasoning和行动Acting每一步思考之后可能调用工具、查询知识库再继续推理。ReAct 的典型流程是Thought思考模型分析当前状态决定下一步。Action行动调用搜索、代码执行、API 等工具。Observation观察获得工具返回的结果继续思考。自蒸馏 监督时域调整的思路同样可以应用到 ReAct 风格模型上。当模型在长 Agent 轨迹中探索时行动步骤之间的发散度往往更高我们可以用 DASH 的思想对“行动前后的短片段”做更精细的监督。这也解释了为什么reasoning models和ReAct这两个热词会和 DASH 出现在同一波技术视野里。例如在 Agent 训练中可以把一条 ReAct 轨迹看成由多个“Thought-Action-Observation”片段拼接而成。DASH 的动态时域调节可以理解为“监督不应该均匀覆盖所有步骤而应该聚焦在策略还不稳定的局部区间”。6. 工程调参与最佳实践6.1 发散度指标的选择如果只用 KL 散度一个标量可能丢失位置信息。建议按轨迹片段分别计算发散度而不是全轨迹平均。Token 级别概率比值发散度对模型更新非常敏感适合早停判断但需要做平滑。熵Entropy可以作为辅助指标策略熵过高说明探索过度熵过低说明过早收敛都可以纳入调节逻辑。6.2 时域调节策略阈值设计不宜过紧。发散度阈值太小时时域会频繁缩到最小相当于退化成只监督答案片段阈值太大时DASH 退化成固定长时域。建议对发散度做标准化用滑动窗口的均值和标准差将发散度归一化到[0,1]再映射到时域。不要在训练刚开始就设置一个很大的 max_horizon最好让模型先适应短窗口再逐步放开。6.3 训练稳定性梯度裁剪依然是必需品。DASH 改变了监督窗口梯度范数可能剧烈变化建议对整体梯度做 clip。KL 惩罚项可以加在 loss 中防止策略与参考模型偏离过远。发散度指标本身来自 KL但 loss 里的 KL 惩罚是另一回事。如果监督时域在训练中频繁抖动可以引入滞后比较机制只有连续 k 步发散度都超标才缩短时域避免单批次极端值造成误判。6.4 项目落地视角在实际项目中我会建议按下面的方式逐步验证 DASH先在一个小规模数据集上跑固定时域 baseline。添加发散度统计日志画出训练过程中发散度的变化曲线。固定发散度阈值观察动态时域是否与人工预期一致例如早阶段时域短、后期时域长。再逐步引入 DASH 的自适应调节对比收敛速度和最终奖励。7. 常见问题与排查思路问题现象常见原因解决思路监督时域一直缩到最小发散度阈值过小或策略持续大幅更新调大阈值或对发散度做 EMA 平滑监督时域始终不变发散度阈值设置过大DASH 没有触发观察发散度分布重新标定阈值训练 loss 震荡严重时域变化过快监督窗口不稳定引入滞回比较连续多次触发再调整模型过度依赖短窗口max_horizon 太小长轨迹监督不足逐步增大 max_horizon或者分段升温KL 发散度虚高新策略与旧策略差异过大关闭或降低 KL 惩罚上限检查学习率采样轨迹奖励普遍为负任务太复杂模型随机探索成功率低先用 SFT 微调模型再上 RL 类训练与 PPO 集成后 PPO loss 异常masks 维度不匹配或时域截取错误打印 masks shape逐 token 检查对齐还有一个容易被忽略的问题supervision horizon和max_new_tokens不要混为一谈。max_new_tokens控制生成长度是采样时的停止条件supervision horizon控制训练时哪些 token 参与损失计算。两者不是一回事。8. 总结与学习路线DASH 的核心价值我认为不在于它发明了某个全新的损失函数而在于它把“监督时域”这个原本需要人工反复尝试的超参数转变为一个可以由策略发散度动态驱动的自适应变量。这种思路对长思维链推理模型的训练特别有意义因为它直接解决了“远端监督稀疏”和“长轨迹不稳定”这两大痛点。如果你准备深入学习这个方向建议按下面的路线推进先搞清楚自回归生成模型的 token-level 监督原理能手动写出一个最小化的语言模型训练循环。动手实现一个简单的 on-policy 采样与优化流程理解为什么 PPO 是 on-policy 算法为什么策略更新后旧数据不再适用。在训练中加入 KL 发散度统计画出发散度曲线这是理解 DASH 的基础。复现一个基础的 DASH controller并在小规模推理任务上对比固定时域与自适应时域的效果。再逐步扩展把 DASH 与 ReAct、Self-Distillation、工具调用等场景结合。如果你已经在尝试训练推理模型我的建议是不要直接把 DASH 视为“银弹”。它更适合作为训练框架中的一个可调节模块与策略梯度、KL 惩罚、奖励归一化等方法配合使用。训练过程中把发散度、监督时域、奖励值三条曲线放在同一张监控面板上观察往往比盯着单个 loss 数字更能发现问题。希望这篇文章能帮你理解 DASH 这个前沿方向的底层逻辑。如果对你有帮助可以收藏备用后面实践过程中遇到训练不稳定或者时域调节失效的问题欢迎在评论区一起交流。