Prhub

#48429 [BugFix] Restore full tokens for Qwen MTP When MoE SP

原始 PR 作者 gcanlin 合并时间 2026-07-13 13:29 文件变更 2 提交数 1 评论 1 代码增减 +20 / -2

执行摘要

修复 Qwen MTP 在 MoE SP 下的 token 损坏

PR #47006 优化了 Qwen MoE 通信,将 all-reduce 替换为 reduce-scatter,但导致 Qwen MTP 在 MoE SP 开启时 CUDA assert 崩溃。原因是 MTP 直接调用 Qwen 的 DecoderLayer(而非完整模型),而 #47006 将 gather 移到 attention 前,使得 MTP 调用时共享了错误的 token 视图。

建议精读本 PR,理解 _all_gather_hidden_and_residual 的修复模式。该模式可作为未来 PR #47006 类似变更影响 MTP 路径时的参考样板。同时建议为 MoE SP + MTP 组合添加自动化回归测试。

讨论亮点

PR 无直接 review 评论。自动审核机器人 claude[bot] 仅为 fork PR 留了一条自动化提示。维护者 ZJY0516 直接批准。没有公开的设计辩论。

实现拆解

  1. 导入新函数:在 qwen3_5_mtp.pyqwen3_next_mtp.py 中,从 vllm.model_executor.models.qwen3_next 导入 _all_gather_hidden_and_residual
  2. 缓存当前层引用:将 self.layers[current_step_idx] 赋值给 mtp_layer 变量,避免多次索引并便于后续访问层属性。
  3. 条件性 all-gather:在 norm 之前,检查 mtp_layer.use_attn_reduce_scatter_for_moe,若为 True(即 MoE SP 已启用),则调用 _all_gather_hidden_and_residual 恢复完整的 hidden_statesresidual 张量。_all_gather_hidden_and_residual 函数负责在 tensor-parallel 组内 all-gather 缩减的 token 维度,还原成完整序列视图。
文件 模块 状态 重要度
vllm/model_executor/models/qwen3_5_mtp.py 模型层 modified 6.44
vllm/model_executor/models/qwen3_next_mtp.py 模型层 modified 6.44

关键符号

Qwen3_5MultiTokenPredictor.forward Qwen3NextMultiTokenPredictor.forward

关键源码片段

vllm/model_executor/models/qwen3_5_mtp.py data-contract

Qwen3.5 MTP 模型文件,修复核心改动:在 norm 前调用 `_all_gather_hidden_and_residual` 恢复完整 token 视图。

# qwen3_5_mtp.py - 在 MTP forward 中根据层属性恢复完整 token 视图
# 新增导入
from vllm.model_executor.models.qwen3_next import (
    QwenNextMixtureOfExperts,
    _all_gather_hidden_and_residual, # 新增:用于 gather 缩减的 token 维度
    _is_shared_expert_fse_compatible,
)# forward 方法中的新增逻辑(位于 norm 之前)
if mtp_layer.use_attn_reduce_scatter_for_moe:
    # 当 MoE 使用 reduce-scatter 时,MTP 层输出中的 hidden_states 和 residual
    # 仅在部分 token 维度上有效(经过 scatter)。此调用在 tensor-parallel 组内
    # all-gather 完整的 token 视图,确保后续 norm 得到正确的全局表示。
    hidden_states, residual = _all_gather_hidden_and_residual(
        hidden_states,
        residual,
        positions.shape[-1], # 完整序列长度
        self.config.hidden_size, # 完整的 hidden size
    )
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
vllm/model_executor/models/qwen3_next_mtp.py data-contract

Qwen3Next MTP 模型文件,与 qwen3_5_mtp.py 完全对等的修复。

# qwen3_next_mtp.py - 与 qwen3_5_mtp.py 完全相同的修复模式
# 新增导入
from vllm.model_executor.models.qwen3_next import (
    Qwen3NextDecoderLayer,
    Qwen3NextModel,
    Qwen3NextRMSNorm,
    QwenNextMixtureOfExperts,
    _all_gather_hidden_and_residual, # 新增
)# forward 方法中的新增逻辑
if mtp_layer.use_attn_reduce_scatter_for_moe:
    hidden_states, residual = _all_gather_hidden_and_residual(
        hidden_states,
        residual,
        positions.shape[-1],
        self.config.hidden_size,
    )
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states

评论区精华

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

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

风险与影响

  • 回归风险低:改动范围小(2 个文件,共 +20/-2),仅在 MoE SP 条件分支下增加 all-gather 调用,对非 SP 路径无影响。
  • 缺失测试覆盖:PR 未附带对应单元测试或集成测试,仅提供了手动验证命令。未来若修改 _all_gather_hidden_and_residual 签名或行为,无自动化测试保障。
  • 依赖紧耦合use_attn_reduce_scatter_for_moe 属性依赖 DecoderLayer 的实现细节,若该属性重命名或语义变更,本修复将静默失效。
  • 用户:修复 Qwen3.5 和 Qwen3Next 使用 MTP 投机解码 + MoE SP 时的崩溃,使该组合正常工作。lm_eval gsm8k 测试显示准确率恢复至 ~85.8%。
  • 系统:对非 MTP 或非 MoE SP 配置无影响。
  • 团队:确认了架构层设计决策(gather 前置)对 MTP 路径的副作用,需在后续类似优化中同步考虑。
核心路径变更 缺少测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论