十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

DASH:动态调整监督范围的推理模型自蒸馏策略

DASH:动态调整监督范围的推理模型自蒸馏策略 训练推理模型的时候很多人其实踩过同一个隐形的坑让小模型学习大模型的思考过程并不是把长思维链的 token 全部拿来做监督训练就有效。监督范围supervision horizon设得太短学生模型学不到完整的推理骨架设得太长教师模型早期的错误步骤又会被学生一字不差地复制进自己的分布。更麻烦的是这个“合适的长度”并不是一个固定值它随着训练的推进不断变化和当前模型的策略、数据分布、任务难度都有关。DASHDivergence-Adaptive Supervision Horizons正是为了解决这个矛盾提出的。它不是给所有样本设置一个固定的监督长度而是根据当前学生策略与目标分布之间的发散程度动态调整监督范围。这个思路听起来只是“调一下超参”但它背后的理念变了监督强度不再是一个需要人工反复试验的常数而是训练过程中一个可在线计算的动态控制信号。读完这篇博客你会理解三件事为什么 supervision horizon 是推理模型自蒸馏里容易被忽略的关键超参为什么固定监督长度会带来两种典型失败模式以及如何在自己的 on-policy self-distillation 训练循环中加入发散自适应监督范围。最后我会给出一套可运行的概念验证代码方便你直接改造成自己的实验脚本。1. 这篇博客真正要解决的问题推理模型reasoning model和普通大语言模型的区别不只是在最终答案前多生成一段文字。以目前常见的 o1 类、R1 类模型为例它们的训练目标之一就是让模型学会“先思考、再回答”。这种思考过程通常被称为 reasoning trace 或 chain-of-thought模型需要在给出结论之前自己生成假设、验证、回退、修正最后才收敛到答案。当你想把这种能力迁移到更小的模型上时最简单的做法是蒸馏让教师模型生成 reasoning trace然后用这些 trace 去监督学生模型训练。项目实践里大家很快会发现直接蒸馏并不总是有效。原因集中在两个层面。第一个层面是数据质量。教师模型生成的长 trace 并不等于每一步都是正确的其中包含大量试错、冗余甚至错误中间结论。如果学生被要求逐 token 复现整段 trace它学到的不只是推理能力还有教师模型自言自语时的噪声。第二个层面是监督强度。训练时我们到底应该监督到第几个 token如果只监督前几十个 token学生根本接触不到完整的推理过程如果监督到最后一步学生又可能被教师早期那些“不完美但实际有用”的思路锁死。更麻烦的是在训练的不同阶段模型的分布一直在变一个固定的监督长度不可能同时适配训练初期和后期。DASH 的核心判断是监督长度不应该由人手动固定而应该由学生策略和教师/目标策略之间的分布差异自动决定。发散了就多给监督收敛了就减少监督把探索空间还给模型。这就是它和普通蒸馏方案最本质的差异。2. 背景从监督微调到 on-policy 自蒸馏2.1 推理模型训练范式的变化传统的指令微调SFT用的是人工标注或规则构造的输入输出对目标是把模型从“会续写文本”变成“会遵循指令完成任务”。但推理模型需要的能力不是“会写答案”而是“会写过程”。过程性的监督数据很难通过标注员稳定获得因为不同人的推理风格差异很大而且很多推理步骤本身没有唯一正确标准。因此推理模型训练开始大量依赖模型自身生成的数据。这就推动了两个方向的流行一个是蒸馏用更大的教师模型生成高质量推理轨迹另一个是强化学习和自蒸馏利用模型自身采样的轨迹做训练信号。2.2 蒸馏、自蒸馏和 on-policy 的区别这里先把几个容易混淆的术语拆开看。蒸馏通常指学生模型学习教师模型的输出分布或生成结果。自蒸馏则是教师和学生来自同族模型例如使用同一个模型更强的 checkpoint 或者更大参数量的同系列模型。on-policy 强调的是“数据由当前策略实时生成”而不是使用一个预先准备好的静态数据集。在推理模型场景里on-policy self-distillation 可以粗略理解为学生模型在当前参数状态下不断采样新轨迹同时参考一个相对更强的模型或更成熟的 checkpoint 来提供监督。每一步训练都在线生成数据模型每更新一次采样分布就变一次因此监督信号始终贴合当前策略。2.3 为什么 on-policy 对推理模型这么重要如果使用离线数据集蒸馏模型更新后数据分布就和当前策略脱节了。尤其是在长思维链场景中训练初期学生只能生成很短的思考过程中期可能突然开始生成大量重复 token后期又可能出现过度压缩推理步骤。离线数据无法覆盖这些不同阶段的问题。on-policy 采样让训练过程随时能看到“模型现在到底是怎么想的”。但这也带来一个新的工程问题模型实时生成的轨迹质量波动很大如果监督策略不随之调整训练极易在“过早锁死”和“过度漂移”之间摇摆。这正是 DASH 试图用发散度来制衡的原因。3. DASH 要解决的问题监督范围为什么不能固定3.1 什么是 Supervision Horizon在序列到序列训练中supervision horizon 指的是训练损失覆盖生成序列前多少个 token。简单说如果一条推理轨迹有 1000 个 token训练时只对前 300 个 token 计算蒸馏损失那么 horizon 就是 300。很多实现里这个值是一个固定超参。大家可能觉得整条轨迹都监督信息不是最多吗问题在于蒸馏损失对模型的影响不是“越多越好”。它本质上是在告诉模型你的每一步输出都应该尽量接近参考分布。但参考分布并不保证每一步都是最优的尤其在推理早期。3.2 固定监督范围的两种失败模式第一种失败模式是监督范围太短。学生只学到推理开头后面的思考过程完全靠自己自由发挥。对能力较强的学生来说这给了它探索空间但对能力较弱的学生它根本没有见过完整的推理结构生成的后半段很容易退化成重复和空话。第二种失败模式是监督范围太长。学生从头到尾被要求对齐教师分布早期推理中的试探性内容、错误分支、无意义重复都会被模型当成“标准答案”吸收。训练后期我们会看到一个普遍现象学生生成的结果和教师高度相似但一遇到分布外的难题就崩掉因为它的内在推理能力并没有真正建立起来。更关键的是这两种失败不是静态的。训练初期学生发散很大需要长监督来稳定训练后期学生已经接近目标分布过长的监督反而会压制探索。固定 horizon 本质上是在用平均主义处理一个动态过程自然处处不合适。3.3 KL 散度的不对称让问题更复杂既然要动态调整监督范围就需要一个信号来度量“学生和教师现在差多远”。最常用的是 KL 散度。但 KL 散度不是对称的前向 KL 和后向 KL 对分布差异的敏感点完全不同。前向 KL 更关注“教师模型中高概率但学生模型中低概率”的位置它会促使学生覆盖教师的全部高概率区域后向 KL 则更关注“学生模型中高概率但教师模型中低概率”的位置它容易让学生只锁定一个教师高概率的模式。在长序列推理场景如果不用正确方向的 KL 作为控制信号模型训练很容易走入“过于保守”或“过于激进”两个极端。DASH 把这样的发散信号当作在线“误差计”根据误差大小决定接下来给多少监督。可以说它把监督范围从一个静态超参变成了一个由分布距离驱动的闭环控制量。4. DASH 的核心原理发散自适应的监督范围4.1 一句话概括 DASHDASH 的做法可以概括为在 on-policy self-distillation 训练过程中实时估计学生策略与参考策略的分布发散度然后根据发散度选择本次训练步应该监督到第几个 token。发散度大说明学生距离参考策略还很远需要延长监督发散度小说明学生已经基本靠近参考策略可以缩短监督把后续生成交给学生自己探索。这个思路和课程学习有相似之处但方向不同。课程学习一般按训练步数或样本难度来编排课程而 DASH 是按模型实际分布与目标分布的距离来编排监督强度。它不关心我们“训练到第几步”它关心的是“模型现在到底有多接近目标”。4.2 用学车来类比 DASH想象一个教练教学生开车。刚开始学生对路况完全不熟教练会频繁接管方向盘甚至每一步都给出明确指令。这时候如果教练完全放手车很容易冲出去。但等学生开了一段路操作越来越稳教练如果还一直抢方向盘学生就永远练不出自己的判断。一个好的教练会根据学生的实际表现调整接管程度。弯道多、车速快的时候多接管直线、路况好、学生稳定的时候少接管。DASH 的监督范围就是这个“接管程度”。发散度衡量的就是“学生现在开得稳不稳”KL 大说明偏离路线远需要多监督KL 小说明已经比较稳可以松手。4.3 DASH 的动态调节策略在实现层面DASH 的调节逻辑可以抽象成一个带阈值的控制策略。通常不会只看单个 token 的 KL因为噪声太大一般会先对一段窗口内的 KL 做平滑。调节规则可以设计为如果最近一段窗口的平均发散度低于下限阈值说明学生分布已经贴近参考分布horizon 可以缩短。如果平均发散度高于上限阈值说明学生偏离较大horizon 需要延长。如果发散度落在两个阈值之间则保持当前 horizon 不变。为了避免 horizon 在训练中高频抖动一般还会加入滞后带或者对 horizon 本身做平滑。horizon 的最小值和最大值也需要限定避免出现“0 监督”或者“整条轨迹完全复制教师”的极端情况。5. DASH 与 ReAct 的关系及适用边界5.1 不要把 DASH 和 ReAct 混为一谈搜索 DASH 时很容易看到另一个热词ReActReasoning and Acting in Language Models。ReAct 强调的是在推理过程中交替进行思考和环境操作通过行动获取外部信息再基于新信息继续推理。它解决的是模型“只会在脑子里想不会动手查”的问题。DASH 关注的是训练阶段的监督信号分配重点在自蒸馏过程中怎么确定监督范围。两者的层级不同。ReAct 改变的是推理过程的交互协议DASH 改变的是训练过程的监督策略。如果非要联系可以这样理解用 DASH 训练出的模型可以在部署时配合 ReAct 使用先让模型学会结构化推理再让它学会与外部工具交互。5.2 DASH 适合什么场景DASH 最适合的场景是你有一个较强的推理模型或成熟 checkpoint 作为教师想训练一个同族或小规模学生模型并且训练数据主要来自模型在线采样。这种场景下学生分布变化快固定监督长度很难调DASH 的自适应机制正好能派上用场。对于已经在做 on-policy self-distillation 的团队DASH 的改造成本也相对可控。你不需要更换全部训练框架只需要在计算蒸馏损失之前多算一个发散度信号然后动态更新 horizon。5.3 DASH 不适合什么场景如果你使用的是完全离线、人工清洗过的固定数据集且教师生成的轨迹已经做过严格筛选和去噪那么固定 horizon 已经够用DASH 带来的收益可能不明显。如果你的教师模型和学生模型能力差距非常大比如用 70B 模型蒸馏 0.5B 模型那么发散度在很长一段时间内都会处于高位DASH 的动态调节区间会被顶到最大值效果上接近于固定长监督。这时候更值得提升的是训练数据质量和模型容量而不是监督策略。如果你的推理轨迹都特别短比如只有几十个 token那么 horizon 的调整空间太小DASH 能发挥的作用有限。这个问题更适合用增强学习或偏好优化来解决。6. 概念验证实现一个 DASH 风格训练循环这一部分给出一个概念验证实现目的是演示如何在 on-policy self-distillation 中加入发散自适应监督范围。代码不是论文官方实现而是一个可运行的最小框架你可以在自己的训练任务上调整。6.1 环境准备建议使用 Python 3.10 以上版本配合 PyTorch 2.x 和 Hugging Face Transformers。模型选择以你实际训练环境为准。为了演示我会用 Qwen2.5 系列这种同系列不同参数量的模型作为教师和学生的例子但代码本身不绑定特定模型。需要安装的基础依赖pip install torch transformers datasets accelerate如果是在多卡环境训练还建议安装 deepspeed 或 peft但本文的核心逻辑不依赖这些库。6.2 定义 DASH 训练配置先用 dataclass 把 DASH 的核心超参集中管理。这个配置文件包含教师模型、学生模型、基础训练参数以及发散度的阈值和 horizon 范围。# 文件路径dash_config.py from dataclasses import dataclass dataclass class DashConfig: # 模型 base_model: str Qwen/Qwen2.5-1.5B-Instruct teacher_model: str Qwen/Qwen2.5-7B-Instruct # 数据与采样 max_length: int 2048 batch_size: int 4 max_new_tokens: int 1024 # 发散度阈值 kl_lower: float 0.3 kl_upper: float 1.2 # horizon 范围与步长 min_horizon: int 256 max_horizon: int 1536 horizon_step: int 64 # horizon 平滑 use_ema: bool True ema_alpha: float 0.2 # 训练 lr: float 1e-5 total_steps: int 2000 grad_accumulation: int 8这里的 kl_lower 和 kl_upper 对训练效果影响很大。不同模型、不同分词的分布尺度不同建议先固定 horizon 训练几百步观察 KL 的数值分布再设置阈值。6.3 计算发散度并选择监督范围这一步是 DASH 的核心。首先我们需要在训练时拿到教师和学生模型对当前生成轨迹每个 token 的 log-probabilities。这里用前向 KL 来度量发散度。所谓前向 KL就是以学生分布为基准计算学生分布与教师分布的差异。然后我们把每个位置的 KL 按窗口平均得到一个平滑的发散度信号。最后根据这个信号动态调整 horizon。# 文件路径dash_core.py import torch import torch.nn.functional as F def compute_kl_trace( student_logprobs: torch.Tensor, teacher_logprobs: torch.Tensor, attention_mask: torch.Tensor, ) - torch.Tensor: 计算每个 token 位置的前向 KL 散度。 参数 student_logprobs: [batch, seq_len, vocab] 学生模型的 log_softmax 输出 teacher_logprobs: [batch, seq_len, vocab] 教师模型的 log_softmax 输出 attention_mask: [batch, seq_len] 1 表示有效 token0 表示 padding 返回 [seq_len] 每个 token 位置的 KL 均值 # 前向 KL sum(p_s * (log p_s - log p_t)) log_ratio student_logprobs - teacher_logprobs student_probs student_logprobs.exp() per_token_kl (student_probs * log_ratio).sum(dim-1) per_token_kl per_token_kl * attention_mask.float() # 只统计有效 token避免 padding 干扰 token_counts attention_mask.float().sum(dim0) token_counts token_counts.clamp(min1) return per_token_kl.sum(dim0) / token_counts def smooth_kl_trace(kl_trace: torch.Tensor, window_size: int 32) - torch.Tensor: 对 KL 序列做滑动窗口平均减少单点噪声。 if kl_trace.numel() window_size: return kl_trace.mean() kernel torch.ones(window_size) / window_size # 这里用 1D 卷积做平滑方便批量处理 smoothed F.conv1d( kl_trace.view(1, 1, -1), kernel.view(1, 1, -1), paddingwindow_size // 2, ) return smoothed.view(-1) def select_horizon( smoothed_kl: torch.Tensor, prev_horizon: int, cfg, ) - int: 根据发散度动态选择新的监督范围。 avg_kl smoothed_kl.mean().item() new_horizon prev_horizon if avg_kl cfg.kl_lower: # 学生已经贴近教师分布可以减少监督 new_horizon min(cfg.max_horizon, prev_horizon - cfg.horizon_step) new_horizon max(cfg.min_horizon, new_horizon) elif avg_kl cfg.kl_upper: # 学生偏离过大需要延长监督 new_horizon min(cfg.max_horizon, prev_horizon cfg.horizon_step) new_horizon max(cfg.min_horizon, new_horizon) # 落在阈值之间则保持不变 if cfg.use_ema: # 对 horizon 做指数平滑防止抖动 new_horizon int(cfg.ema_alpha * new_horizon (1 - cfg.ema_alpha) * prev_horizon) return max(cfg.min_horizon, min(cfg.max_horizon, new_horizon))这段代码中最核心的是 select_horizon 函数。它不直接依赖训练步数而是依赖当前学生分布和教师分布的 KL 信号。如果 KL 一直在阈值区间内波动horizon 就会稳定在某个值附近而不是固定不变。6.4 训练循环把发散度接入损失计算训练循环主要分为四个阶段学生模型按当前策略采样生成推理轨迹。用教师模型和学生模型分别计算 log-probabilities。计算发散度并决定当前步的 horizon。对前 horizon 个 token 计算蒸馏损失反向传播更新学生模型。下面是一个核心训练步的代码。为了保持清晰我假设 teacher 和 student 都是 Hugging Face 的 CausalLM 模型并且已经通过 accelerate 或普通 PyTorch 封装好。# 文件路径train_step.py import torch import torch.nn.functional as F from dash_config import DashConfig from dash_core import compute_kl_trace, select_horizon def compute_distill_loss(student_logits, teacher_logits, horizon, attention_mask): 只对前 horizon 个 token 计算蒸馏损失。 蒸馏损失使用 KL 散度目标是让学生分布逼近教师分布。 seq_len student_logits.size(1) student_logprobs F.log_softmax(student_logits, dim-1) teacher_logprobs F.log_softmax(teacher_logits, dim-1) # 每个 token 的 KL per_token_kl F.kl_div( student_logprobs, teacher_logprobs, reductionnone, log_targetTrue, ).sum(dim-1) # 只保留前 horizon 个 token horizon_mask torch.arange(seq_len, devicestudent_logits.device) horizon horizon_mask horizon_mask.float() per_token_kl per_token_kl * horizon_mask.unsqueeze(0) # 过滤 padding loss_mask horizon_mask.unsqueeze(0) * attention_mask.float() denom loss_mask.sum().clamp(min1.0) return per_token_kl.sum() / denom def train_step(student, teacher, tokenizer, batch, cfg, optimizer, prev_horizon): prompts batch[prompt] student_inputs tokenizer(prompts, return_tensorspt, paddingTrue).to(student.device) # 1. 学生模型 on-policy 采样 with torch.no_grad(): gen_outputs student.generate( **student_inputs, max_new_tokenscfg.max_new_tokens, return_dict_in_generateTrue, output_scoresTrue, ) generated_sequences gen_outputs.sequences seq_len generated_sequences.size(1) # 2. 计算学生和教师对同一段生成轨迹的 logits student_logits student(generated_sequences).logits with torch.no_grad(): teacher_logits teacher(generated_sequences).logits # 这里简化处理 attention mask实际训练中需要构造生成序列的 mask attention_mask torch.ones_like(generated_sequences) # 3. 计算发散度并选择 horizon teacher_logprobs F.log_softmax(teacher_logits, dim-1) student_logprobs F.log_softmax(student_logits, dim-1) kl_trace compute_kl_trace( student_logprobs[:, :cfg.max_length, :], teacher_logprobs[:, :cfg.max_length, :], attention_mask[:, :cfg.max_length], ) horizon select_horizon(kl_trace, prev_horizon, cfg) # 4. 计算蒸馏损失并更新 loss compute_distill_loss( student_logits[:, :cfg.max_length, :], teacher_logits[:, :cfg.max_length, :], horizonmin(horizon, cfg.max_length), attention_maskattention_mask[:, :cfg.max_length], ) loss.backward() optimizer.step() optimizer.zero_grad() return loss.item(), horizon这个 train_step 是简化版重点展示了 DASH 的接入方式。实际工程里你还需要处理 padding mask、梯度裁剪、混合精度、EMA 等问题但核心逻辑已经完整先根据 KL 动态算出 horizon再按 horizon 截断蒸馏损失。这里我特别想提醒一点student_logits 是用学生当前策略生成的序列重新前向计算出来的。也就是说我们先用学生模型采样轨迹再用学生模型当前参数对这些轨迹重新打分。这是 on-policy 训练的标准做法否则梯度会穿过采样过程导致训练不稳定。6.5 运行与验证如果你的环境和模型准备就绪可以把上面的模块拼成一个最小训练脚本。训练时建议每隔固定步数打印 loss 和 horizon。python train_step.py训练日志可能类似于下面这种格式具体数值取决于你的模型和数据这里只展示结构step100 loss2.345 horizon768 avg_kl0.68 step200 loss2.102 horizon832 avg_kl0.91 step300 loss1.864 horizon1024 avg_kl1.15 step400 loss1.635 horizon768 avg_kl0.55 step500 loss1.421 horizon640 avg_kl0.31如果 horizon 一直顶在最大值说明学生的发散度持续超过阈值这时候 DASH 没有起太多调节作用你需要检查 KL 计算是否准确或者教师和学生模型差距是否过大。7. 如何评估 DASH 训练是否生效7.1 离线指标判断训练过程中比较重要的指标有三个loss、horizon、发散度。loss 下降说明训练在收敛但不代表模型推理能力变强horizon 变化说明自适应机制在起效发散度下降说明学生分布正在接近教师分布。但更可靠的判断方式是评估集上的推理准确率和生成质量。建议固定一批数学题、逻辑推理题或代码题在每个 checkpoint 上做一次评估看准确率是否随训练提升以及模型生成的 reasoning trace 是否出现更多有效步骤。7.2 动态行为观察DASH 是否生效最直观的观察点是 horizon 的变化趋势。理想的 DASH 训练会呈现出“初期 horizon 较高中后期逐渐下降最后稳定在一个合理区间”的模式。如果你的训练日志里 horizon 几乎没有变化说明阈值设置不合理。建议先跑一个固定 horizon 的 baseline收集 KL 的分布再把 kl_lower 和 kl_upper 分别设置在分位数附近。7.3 对生成结果做结构化检查除了数值指标还可以人工检查学生模型生成的推理轨迹。重点看三件事第一学生是否真的在“思考”还是只是在复制教师开头的几个 token第二学生的推理中间步骤是否比固定监督训练更长或者更有阶段性第三在难题上学生是否出现“开头正确但结尾崩坏”的情况。如果这些问题出现通常说明监督范围的控制时机不太合适需要调整阈值或 horizon 的限制范围。8. 常见问题与排查思路问题现象可能原因排查方式解决方案horizon 一直顶在最大值kl_lower 设置过低或教师学生差距过大输出 KL 分布直方图观察均值和中位数提高 kl_upper 和 kl_lower适当减小模型差距horizon 频繁抖动阈值区间过窄或 KL 噪声太大检查平滑窗口和 EMA alpha增大 KL 平滑窗口让阈值区间更宽训练 loss 异常升高KL 计算包含 padding 或 prompt 位置检查 attention mask 是否覆盖生成序列过滤 padding只计算 reasoning trace 部分学生重复生成同一句话监督范围过短或 min_horizon 过低查看生成序列的重复比例提高 min_horizon或加入重复惩罚学生完全复刻教师风格max_horizon 过大监督过强比较学生和教师的生成文本相似度降低 max_horizon或调低 kl_upper显存不足同时计算教师和学生 logits查看显存占用和 batch size使用梯度检查点、减少 batch、教师模型走推理模式不反传9. 最佳实践与工程建议9.1 先确认固定 horizon 的 baseline很多团队一上来就调 DASH 阈值结果发现训练不稳定。更稳妥的做法是先用固定 horizon 跑通整个流程记录训练中的 KL 分布和 loss 曲线。这个 baseline 是后续判断 DASH 是否有效的基础。没有 baseline 就调 DASH等于没有参照物。9.2 发散度信号要做平滑不要裸用单个 token 的 KL 波动非常大直接用来决定 horizon 会让监督范围频繁跳变。建议至少对 KL 序列做一个窗口平滑再对 horizon 做 EMA。这样训练曲线更稳定代码改动也不大。9.3 horizon 必须有上下界无论发散度信号多么好horizon 都要限制在合理范围内。如果 horizon 太小学生可能完全没有教师监督如果 horizon 太大整个训练退化成逐 token 复刻。设置 min_horizon 和 max_horizon本质上是给自适应机制加一个安全边界。9.4 KL 计算方向要与训练目标一致如果你希望学生尽量覆盖教师的所有高概率行为使用前向 KL 比较合适如果你希望学生专注在自身高概率区域、避免去拟合教师低频噪声可以考虑反向 KL。DASH 的原始理念核心是将发散度作为控制信号具体 KL 方向要根据你的任务目标选择。9.5 训练数据与生成内容的安全过滤推理模型训练中教师模型可能生成有害、偏见或虚构的内容。无论蒸馏还是自蒸馏都必须对训练数据和最终模型的输出做安全评估。不要因为追求特殊 token 的相似度而忽略内容安全。往训练管线中加一层过滤和审计在工程上是必要的。9.6 日志记录保持完整训练中至少记录 step、loss、horizon、avg_kl、当前学习率。如果之后发现某个 checkpoint 效果特别好或特别差这些日志能帮你快速归因。horizon 的动态曲线本身也会成为非常有效的诊断工具。10. 总结与下一步推理模型的自蒸馏并不只是把大模型的输出复制给一个小模型那么简单。真正的难点在于模型的推理能力分布时刻在变固定监督长度无法同时满足“初期稳”和“后期放”两个要求。DASH 给出的方案是把 supervision horizon 从人工超参变成一个由发散度驱动的在线控制信号让监督强度跟随学生策略与目标分布的距离自动调整。这篇博客从监督范围的根本矛盾讲起解释了 DASH 的核心设计也给出了一套概念性的 PyTorch 实现。你可以在自己的 on-policy self-distillation 实验基础上把 select_horizon 和动态损失计算加进去先跑一个固定 horizon 的 baseline再对比 DASH 带来的变化。下一步值得继续深入的方向有三个第一把 DASH 和强化学习策略梯度混合看看动态监督范围能否降低策略训练的方差第二在更长推理轨迹和更大模型规模上验证阈值规律第三把发散度信号进一步细化到 token 级别而不仅仅是序列级平均这样监督范围的控制可以更精细。希望这篇博客能给你一个足够清晰的出发点。
返回列表