ARTICLE DETAIL

资讯详情

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

DASH:散度自适应监督视野的推理模型自蒸馏方法

DASH:散度自适应监督视野的推理模型自蒸馏方法 DASHDivergence-Adaptive Supervision Horizons for On-Policy Self-Distillation of Reasoning Models这个标题拆开看每个词都指向一个实际训练中逃不掉的问题推理模型Reasoning Models的推理轨迹是模型自己生成的训练信号也来自模型自身这就构成了自蒸馏Self-Distillation因为数据必须由当前策略实时采样这又天然落在 on-policy 的框架里而 Supervision Horizons 则回答一个更细的问题——当一条推理轨迹有几十甚至上百步时我们究竟应该对多少步给予监督。DASH 给出的思路是引入散度Divergence作为调节信号当前策略与参考策略偏离小的时候可以信任过程监督把监督视野拉长偏离大的时候过程步骤可信度下降就把视野收短。这篇文章会沿着概念、机制、实现、验证、排错这条线把这套设计讲透。1. 先理解推理模型自蒸馏为什么依赖 on-policy1.1 推理模型的自训练闭环推理模型指的是会先生成一段可见的思维链chain-of-thought再给出最终答案的模型。这类模型的训练常见两种信号来源一是人工标注二是模型自己的采样。后者就是 self-distillation用一个已经具备一定能力的模型去采样推理轨迹再筛选其中结果正确的轨迹作为训练数据让模型继续变强。在这个闭环里老师和学生是同一个或相近的模型。这里容易出现一个误区如果只是拿自己采样的数据反复训练很容易退化成自我模仿。模型会把训练集中常见但错误的推理模式固化下来表现为生成内容越来越长但答案准确率没有提升。解决思路通常有二一是用奖励模型或规则校验器筛选高质量轨迹二是做 on-policy 训练保证采样分布和当前策略始终保持一致。两者配合才是推理模型持续变强的基础。1.2 为什么 PPO 是 on-policy旧数据为什么不能无限复用PPO 是 on-policy 算法的典型代表。on-policy 的含义是更新策略时使用的经验必须来自当前策略的采样。PPO 每次迭代都重新用当前策略去环境中 rollout然后用这些最新数据更新一轮或若干轮。相比之下DQN 使用经验回放池旧经验可以被反复使用这是 off-policy。为什么推理模型训练特别在乎 on-policy因为推理模型的采样分布会随训练快速漂移。前一轮 checkpoint 采样的推理轨迹在新一轮 checkpoint 看来可能已经不是它自己会输出的内容。用旧数据训练等同于用别人的行为来更新自己策略会越训越偏。这也是 R1 风格训练中每一步都要重新采样、而不是复用离线数据的原因。1.3 On-policy 自蒸馏真正要解决的问题在 on-policy 自蒸馏框架下训练闭环是当前策略采样奖励评估策略更新再采样。真正的问题不是是否要 on-policy而是监督信号给到什么粒度。如果只给最终答案奖励中间几十步推理完全没有梯度信号模型不知道自己是哪一步想错的如果给每一步奖励错误步骤也会被同步放大。这两个极端之间需要一种机制来动态决定监督的粒度。DASH 中 Supervision Horizons 要处理的就是这个核心矛盾。2. 监督视野结果监督、过程监督与两者之间2.1 三种常见监督粒度监督视野可以直接对标准强学习里的 return horizon。下面表格列出了三种常见粒度。监督方式监督范围优点风险outcome-only只监督最终答案信号干净不容易强化错误中间步稀疏长链推理难学process / step-wise每一步都监督每一步都有学习信号中间步骤有噪声可能被错误 verifier 带偏adaptive horizon动态控制监督步数在探索期和收敛期使用不同粒度需要额外机制决定何时监督多少步实际项目中outcome-only 最稳定但学得慢process-step 学得快但对校验器要求极高。adaptive horizon 想做的就是两边的折中。2.2 固定视野的偏差固定使用结果监督时长推理链会出现 credit assignment 困难一条轨迹 50 步前 20 步思路正确第 21 步开始偏离最终答案错误。outcome-only 会把这 50 步全部判负模型无法定位问题出在哪一步只能靠大量采样去盲试。固定使用过程监督时问题变成谁来判断中间步骤是否正确如果用一个规则校验器或小型奖励模型对每一步打分打分本身的误差就会被放大。校验器认为某一步有价值实际是废话或错误推导模型就学会了写这种步骤能拿高分。所以固定视野不是不好而是无法同时应对探索期和收敛期两种状态。2.3 为什么视野要跟着散度走DASH 的关键假设是当前策略与参考策略之间的散度可以反映当前轨迹分布的稳定性和可信度。如果当前策略与参考策略几乎一致散度小说明模型的推理模式还在可信区间内中间步骤大概率是可学的此时可以扩大监督视野利用过程信号加速收敛。如果散度快速增大说明模型正在探索或漂移中间步骤的可信度下降此时只有最终答案或很短的后段轨迹才适合作为监督信号。用通俗的话说模型还没走偏的时候我们可以看它每一步走得好不好模型已经大改路线的时候我们只能看它最后有没有到达目标。这和人类带新人的方式很像——分歧小的时候可以拆步骤纠正分歧大的时候先保证大方向正确。3. DASH 机制拆解散度计算、视野映射与训练伪代码3.1 散度指标KL 散度怎么算散度计算通常使用当前策略与参考策略在每个 token 上的 KL 散度。参考策略可以是训练初期保存的 SFT 模型也可以是最近一次更新的 checkpoint或者用 EMA 维护的滑动版本。计算时要注意三点第一必须使用完整的 next-token 分布不能只取采样到的 token 的 logprob第二padding token 要屏蔽否则会稀释散度第三计算要在 log 域完成避免数值溢出。下面代码用于说明思路实际项目要结合自己的模型接口和分词器调整。import torch import torch.nn.functional as F torch.no_grad() def compute_kl_divergence(policy, ref_policy, trajectories, pad_token_id0): total_kl 0.0 valid_count 0 for input_ids, output_ids in trajectories: # output_ids 包含 reasoning tokens 和最终答案 logits_cur policy(input_ids, output_ids) logits_ref ref_policy(input_ids, output_ids) logp_cur F.log_softmax(logits_cur, dim-1) logp_ref F.log_softmax(logits_ref, dim-1) # 以参考策略为基准的 KL(pi_ref || pi_cur) per_token_kl torch.exp(logp_ref) * (logp_ref - logp_cur) per_token_kl per_token_kl.sum(dim-1) # 屏蔽 padding token mask (output_ids ! pad_token_id).float() per_token_kl per_token_kl * mask total_kl per_token_kl.sum().item() valid_count mask.sum().item() return total_kl / valid_count需要注意的是 KL 方向KL(pi_ref || pi_cur)和KL(pi_cur || pi_ref)数值不同含义也不同。前者衡量参考策略期望的 token 分布当前策略偏离了多少更贴近分布漂移的语义后者衡量当前采样点在参考分布下的意外程度。在一个项目里只能选一种并保持一致否则后面阈值全部失效。3.2 视野映射函数拿到散度后需要一个映射规则把它变成监督视野比例。常见做法是线性插值。def adapt_horizon_ratio(kl_value, kl_low1.0, kl_high4.0, h_min0.1, h_max1.0): if kl_value kl_low: return h_max elif kl_value kl_high: return h_min else: # 线性插值也可以换成 log 或分段函数 t (kl_value - kl_low) / (kl_high - kl_low) return h_min (h_max - h_min) * (1.0 - t)这里h_max1.0表示对完整推理轨迹做过程监督h_min0.1表示只监督最后一小段推理0则表示只监督最终答案。散度小视野比例高散度大视野比例低。阈值kl_low和kl_high不是固定值需要根据实际任务中 KL 的分布来标定。3.3 完整训练循环伪代码policy load_policy(sft_model) ref_policy copy.deepcopy(policy) for step in range(max_steps): prompts sample_prompts(batch_size) # on-policy 采样必须使用当前策略 trajectories policy.generate( prompts, temperature0.7, max_new_tokens2048 ) # 计算散度并自适应监督视野 kl compute_kl_divergence(policy, ref_policy, trajectories) horizon_ratio adapt_horizon_ratio(kl) # 根据视野构建训练目标 targets build_supervision_targets(trajectories, horizon_ratio, verifier) # 策略更新可以是 RL 目标也可以是一般的蒸馏/偏好目标 loss compute_loss(policy, trajectories, targets) loss.backward() optimizer.step() # 参考策略需要按节奏更新不能永远停在初始模型 if step % ref_update_steps 0: ref_policy copy.deepcopy(policy)整个循环的关键点有两个采样必须用当前策略参考策略要定期更新。如果参考策略一直不动KL 会随着训练单调增大视野一路被压到最短DASH 就退化成 outcome-only如果参考策略每一轮都紧跟自己KL 永远很小视野一直最长DASH 又退化成 process-full。4. 训练流程与关键配置4.1 Rollout 配置on-policy 自蒸馏的 rollout 阶段决定了数据质量配置上要专门对待。配置项常见选择说明temperature0.6 - 0.8采样阶段需要一定探索不能太低max_new_tokens512 - 4096按任务推理长度设置过长浪费时间num_samples per prompt4 - 16用于计算通过率或多数投票奖励verifier规则校验器 / 奖励模型数学代码场景可用规则开放场景要训练 PRM评估 temperature与 rollout 一致或更低不一致会导致评估无法反映训练分布一个常见错误是 rollout 用 temperature 0.8评估用 temperature 0.2然后发现评估分数和训练曲线对不上。建议先统一温度再观察模型行为。4.2 参考策略更新节奏参考策略的更新节奏直接影响 KL 的稳定性固定不变适合做 KL 正则约束防止模型偏离 SFT 太远但 DASH 的视野调节能力会被削弱。定期复制当前策略能反映相对漂移但周期性突变会让 KL 曲线出现锯齿。EMA 滑动更新平滑性最好推荐在探索和收敛之间做平衡。节奏的选择原则是让 KL 既有变化信号又不至于剧烈震荡。如果日志里 KL 曲线是单调递增的尖峰状多半是参考策略更新太慢。4.3 需要重点关注的超参数参数含义常见范围调大影响调小影响kl_low触发最大视野的散度阈值0.5 - 2.0更少进入长视野更容易进入长视野kl_high触发最小视野的散度阈值3.0 - 6.0更少进入短视野更容易进入短视野h_max / h_min最长/最短视野比例0.8 - 1.0 / 0.0 - 0.2更多过程监督更接近结果监督ref_update_steps参考策略更新频率100 - 500KL 更稳但滞后KL 波动大num_samples每个 prompt 采样数4 - 16奖励估计更准训练更慢上面数值只是起步参考换任务和模型规模后必须重新标定。正确做法是固定其他参数单独扫描kl_low和kl_high画散度分布直方图后选两个分位数作为初值。5. 运行验证与基线对比5.1 训练过程需要看的曲线跑 DASH 训练时下面几条曲线必须记录KL 散度曲线观察是否稳定参考策略更新节奏是否合理。视野比例曲线确认视野没有被某个极端值长时间锁死。验证集准确率或奖励曲线判断模型是否真正变强。平均轨迹长度过程监督下模型容易变啰嗦长度需要监控。一条健康的学习曲线应该是训练初期散度上升、视野收短模型主要靠结果监督完成粗调中期散度回落、视野逐渐拉长过程监督开始发挥作用后期视野和散度都稳定在小范围内。5.2 对比固定视野基线要判断自适应机制是否有效至少要和三种固定视野基线对比outcome-only、process-full、固定比例 0.5。评估指标至少包括 pass1、passk 和平均轨迹长度。方案pass1passk平均轨迹长度备注outcome-only待测待测待测信号稀疏长链任务提升慢process-full待测待测待测依赖 verifier 质量fixed 0.5待测待测待测折中但无法自适应DASH待测待测待测期望在稳定性和上限上取平衡这里表格里的数字需要你自己跑实验填充不能直接照搬。DASH 的价值通常不是每个 checkpoint 都最高而是训练过程更稳定不会在中途出现灾难性遗忘或奖励崩塌。5.3 如何合理解读结果如果 DASH 最终分数和 fixed 0.5 差不多但训练过程少了很多次大幅回退这也是收益。推理模型训练的成本主要在采样减少无效迭代就是节省成本。同时要做阈值敏感性分析把kl_low和kl_high上下调整 30%看最终指标是否稳定。如果小幅度调整就导致结果大幅波动说明自适应机制依赖的散度信号不稳定问题大概率出在散度计算或参考策略更新节奏上而不是视野映射本身。6. 常见问题与排查路径6.1 现象loss 正常但验证分数不涨可能原因过程监督把错误步骤当成了正样本。此时模型 loss 会下降因为它学会了输出符合 verifier 偏好的步骤而不是输出正确推理。排查方式随机抽 50 条训练轨迹人工检查 verifier 对每步的评分是否合理。如果评分明显错判优先修 verifier或把h_max调低。6.2 现象KL 一直很大视野长期锁在最短可能原因参考策略更新太慢模型已经偏离参考分布很远。排查方式打印 KL 曲线看是否单调上升。如果是缩短ref_update_steps或改 EMA。还要确认散度计算是否屏蔽了 padding以及是否统一了 KL 方向。6.3 现象视野比例来回跳训练不稳定可能原因kl_low和kl_high区间太窄KL 的每步波动就能跨越整个区间。排查方式先看 KL 分布直方图把阈值放到分布的两端分位数上而不是拍脑袋。也可以对视野比例做指数滑动平均抑制单步抖动。6.4 现象模型输出变长但正确答案变少可能原因过程监督奖励了无意义的长步骤模型学会写长废话来获得高分。排查方式检查短轨迹和长轨迹的奖励分布。如果长轨迹普遍高分建议加长度惩罚或在监督目标里强制保留 outcome-only 分量。问题现象常见原因检查方式处理建议loss 正常但验证分数不涨过程 verifier 评分错误人工抽查 50 条轨迹的逐步评分修 verifier 或降低 h_maxKL 一直很大视野锁最短参考策略更新过慢打印 KL 曲线缩短 ref_update_steps 或改 EMA视野比例来回跳阈值区间过窄看 KL 分布直方图用分位数标定阈值加滑动平均输出变长但正确率下降过程监督奖励长废话对比长/短轨迹奖励分布加长度惩罚保留 outcome-only 分量7. 最佳实践与扩展方向7.1 落地建议清单在实际项目里引入 DASH 思路前建议先对照这份清单检查环境散度计算统一在 log 域显式屏蔽 padding token。全程只使用一种 KL 方向不混用。参考策略不静止不动也不要跟得太紧。起步阶段先用 outcome-only warmup 若干步再开启自适应视野。训练日志必须记录 KL、视野比例、奖励、轨迹长度四个指标。每 N 步保存一次 checkpoint便于训练回退时定位是哪一步引入的回归。过程 verifier 上线前做小规模人工校验确认逐步评分质量。固定评估温度与 rollout 温度保持一致或单独说明差异。这套清单不只适用于 DASH任何涉及过程监督和策略自蒸馏的训练都能复用。7.2 与 ReAct 这类 reasoning acting 路线的联系推理模型的轨迹不只有纯文本思维链还可以包含工具调用和外部反馈。ReAct 的思路就是把 reasoning 和 acting 交替进行模型思考下一步调用工具观察结果再继续推理。这类轨迹天然自带更强的监督信号——工具返回的观察结果本身就是对中间步骤的校验。DASH 的监督视野概念在 ReAct 场景同样适用但要额外考虑动作结果的可验证性。比如模型调用计算器后得到明确结果这个中间步骤是可信的应该纳入监督模型自己生成的一段推测性推理可信度则要看散度。可以把动作结果是否可验证作为一个附加条件和散度一起决定视野这是比纯文本推理更自然的扩展方向。7.3 下一步可以尝试的改进如果实验已经验证了散度驱动的视野调节有效下一个阶段可以尝试把结果正确性、verifier 置信度作为第二信号与散度联合决定视野。做渐进式视野调度训练早期强制结果监督中期逐步放长视野后期再交给散度自动调节。在过程奖励不可靠时使用 outcome 优先的混合目标让最终答案监督始终占一定权重。与不同 on-policy 更新器如 PPO、GRPO结合观察自适应视野对策略更新稳定性的影响。推理模型训练里最容易被忽视的一点是采样、监督、更新三者之间的分布一致性。DASH 的监督视野设计本质上就是在维护这种一致性——散度小时信任中间步骤散度大时回归结果监督。理解这层逻辑后再去看自适应阈值、参考策略更新和 verifier 质量就不会只停留在调参层面。对新手来说最值得做的练习不是直接复现完整实验而是先跑通一条带 KL 日志的 on-policy 自蒸馏基线再逐步加入视野自适应逻辑每一步都保留固定视野对照。这样你才能知道自适应机制到底解决了什么问题又是靠什么信号在调节。
返回列表