Prhub

#1926 Fix weight sync selector for frozen speculative drafts

原始 PR 作者 XinyuJiangCMU 合并时间 2026-07-30 08:34 文件变更 5 提交数 2 评论 1 代码增减 +40 / -9

执行摘要

修复冻结 draft 权重同步选择器

PR body 指出:raw conversion 模式下 trainer 可启用 speculative decoding 而不训练 MTP 层,此时 draft 冻结不应接收权重更新;但 Miles 之前未告知 SGLang 权重更新覆盖哪些 runner,默认 selector 'all' 让 draft 在每个训练步骤都被纳入更新会话,尽管 loader 过滤了所有 tensor,SGLang 仍对其执行后处理,且错误是静默的——job 不崩溃、loss 正常,只有 rollout 侧 spec_accept_length 从 2.71 掉到 1.00。根因是 target 收到未 shuffle 的新权重,而 draft 保留的是已 shuffle 的旧权重,SGLang 对 draft 二次应用非幂等的 AITER shuffle,产生错误排列。

值得精读。该 PR 展示了如何用一个单一决策函数统一权重更新覆盖范围,并通过端到端指标(spec_accept_length)验证修复效果;同时揭示了非幂等变换(AITER shuffle)在权重同步中可能引发的静默损坏,对处理分布式训练推理一致性有启发。建议关注其条件判断的扩展性,并补充针对 selector 分支的单元测试。

讨论亮点

Review 无实质讨论,yueming-yuanguapisolo 直接 APPROVED,guapisolo 仅回复 "LGTM"。PR body 中作者自证了问题的严重性和修复有效性:给出 AITER shuffle_weight 非幂等的最小复现(连续两次 shuffle 得到不同排列),并用对照组数据说明修复前 spec_accept_length 跌到 1.00、修复后回到 2.75。

实现拆解

  1. 新增 selector 决策函数:在 miles/backends/megatron_utils/update_weight/common.py 中新增 weight_update_selector(args),当 sglang_speculative_algorithm 启用、mtp_num_layers 缺失或为 0、且 megatron_to_hf_mode != "bridge" 时返回 "target",否则返回 "all"。这是判断 draft 是否冻结的唯一逻辑入口。
  2. 扩展更新会话入口begin_weight_update 增加 selector: str = "all" 参数,并通过 engine.begin_weight_update.remote(selector=selector) 传给每个 rollout 引擎,确保会话开始时就知道覆盖范围。
  3. 透传 selector 到所有 weight payload:在 miles/backends/sglang_utils/sglang_engine.py 中为 update_weights_from_tensorupdate_weights_from_distributedbegin_weight_update 增加 selector 参数并写入 HTTP payload。
  4. 接入两条更新路径:在 update_weight_from_tensor.py(colocate + tensor 路径)中,update_weights 调用 begin_weight_update 时传入 weight_update_selector(self.args),并在 _send_base_params_send_lora_params_send_to_colocated_engine 中透传 selector;在 update_weight_from_distributed/mixin.py_pause_and_prepare_engines 中缓存 self._weight_update_selector,供 _update_weight_implementation 调用 update_weights_from_distributed 时使用;broadcast.py 中的 update_weights_from_distributed 函数也增加 selector 参数并透传。
  5. 格式整理:第二个 commit 用 isort 将 common.py 的 import 折叠为单行,无逻辑变化。
  6. 测试与验证:本 PR 未新增单元测试,但 PR body 提供了 4 节点端到端验证数据(frozen-draft reference 2.7056,修复前 1.00,修复后 2.7490),以及 rollout 时间降低约 34% 的对比。
文件 模块 状态 重要度
miles/backends/megatron_utils/update_weight/common.py 权重同步 modified 6.83
miles/backends/sglang_utils/sglang_engine.py 引擎客户端 modified 6.5
miles/backends/megatron_utils/update_weight/update_weight_from_tensor.py 权重同步 modified 5.48
miles/backends/megatron_utils/update_weight/update_weight_from_distributed/mixin.py 权重同步 modified 5.07
miles/backends/megatron_utils/update_weight/update_weight_from_distributed/broadcast.py 权重同步 modified 4.59

关键符号

weight_update_selector begin_weight_update update_weights_from_tensor update_weights_from_distributed _pause_and_prepare_engines

关键源码片段

miles/backends/megatron_utils/update_weight/common.py core-logic

新增 weight_update_selector 并改造 begin_weight_update,是本 PR 的核心决策点与入口。

def weight_update_selector(args) -> str:
    """决定本次权重同步覆盖哪些 SGLang runner。    只有能证明 draft 冻结时才排除 draft,返回 "target";
    否则保守返回 "all"(target + draft 都覆盖)。
    """
    if (
        getattr(args, "sglang_speculative_algorithm", None) # 启用了 speculative decoding
        and not getattr(args, "mtp_num_layers", None) # 没有训练 MTP 层
        and getattr(args, "megatron_to_hf_mode", "raw") != "bridge" # 非 bridge 转换路径
    ):
        return "target"
    return "all"
​
​
def begin_weight_update(rollout_engines: Sequence[ActorHandle], selector: str = "all"):
    """在选中的 rollout 引擎上开启权重更新会话。    selector 必须与后续 weight payload 保持一致,
    避免更新会话覆盖范围与实际加载范围不一致。
    """
    ray.get([engine.begin_weight_update.remote(selector=selector) for engine in rollout_engines])

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  1. 依赖 SGLang 版本支持 selector 参数:如果 rollout 引擎的 SGLang 版本未实现该字段,可能忽略或报错;_check_weight_sync_results 已能捕获显式失败,但静默忽略仍可能保留旧行为。
  2. selector 判定条件脆弱weight_update_selector 依赖 mtp_num_layers 是否存在以及 megatron_to_hf_mode 是否等于 bridge,若训练配置中 MTP 层通过其他方式传递(如 mtp_num_layers=0 但实际有 MTP),会误返回 "target" 导致 draft 更新缺失;反之亦然。
  3. 缺少单元测试:该逻辑涉及多条配置组合,当前无自动化测试覆盖,回归风险完全依靠人工端到端验证。
  4. 影响所有权重同步路径:本 PR 同时改动了 colocate 与 distributed 两条路径,若任一透传遗漏(如 LoRA 路径),可能出现 selector 不一致(会话与 payload 范围不匹配),导致新的静默错误。

影响范围集中在使用 raw 转换模式 + speculative decoding 且不训练 MTP 层的训练场景(如 DeepSeek-V4-Flash FP8 on gfx950)。修复后这些配置下 spec_accept_length 从 1.00 恢复到 2.75 附近,rollout 时间降低约 34%,训练吞吐可预期提升。对无 speculative decoding 或 bridge 模式用户无行为变化(selector 仍为 all)。团队需确保升级的 SGLang 版本支持 selector 字段,否则需同步更新 SGLang。

静默错误修复 依赖 SGLang 版本支持 selector 缺少单元测试 逻辑依赖参数配置

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论