Prhub

#45307 [Bugfix] Fix trtllm fused allreduce+rms_norm for transformers backend

原始 PR 作者 tdoublep 合并时间 2026-06-16 16:34 文件变更 1 提交数 2 评论 3 代码增减 +7 / -0

执行摘要

修复 TRT-LLM 融合 allreduce + RMS Norm 的维度假设问题

PR body 提供了明确的 crash 场景:在 TP=2、使用 Transformers 后端加载 Qwen3-1.7B 时,allreduce_in.shape 解包为 (num_tokens, hidden_size) 会抛出 ValueError: too many values to unpack。该问题的根本原因是 Transformers 后端保留了 batch 维度,导致输入 shape 为 3D,与静态 2D 假设冲突。

值得合入:这是一个小而精准的兼容性修复,reviewer 明确 approve,且作者已根据反馈补充注释。建议直接合入 main 分支,无需精读详细设计。

讨论亮点

review 中 ZJY0516 建议为这一修复添加注释以说明这是针对 transformers 后端的处理,作者 tdoublep 响应“done”并在下一个 commit 中添加了注释。该讨论无争议,结论是补充的注释有效提升了代码可读性与上下文关联性。

实现拆解

  1. vllm/compilation/passes/fusion/allreduce_rms_fusion.pycall_trtllm_fused_allreduce_norm 函数中,于原有 2D 解包语句前插入形状兼容逻辑。
  2. 检测 allreduce_in.dim() != 2:如果不为 2D,则取最后一维作为 hidden,并对 allreduce_inresidual、以及可能的 norm_out 执行 .view(-1, hidden) 展平,保留 token 总数与 hidden size 的语义。
  3. 通过 review 添加了注释 # handle transformers backend passing outer batch dim.,明确该补丁的动机和适用场景。
  4. 无测试文件变更,但 PR body 提供了手动复现脚本作为验证。
文件 模块 状态 重要度
vllm/compilation/passes/fusion/allreduce_rms_fusion.py 编译融合 modified 5.4

关键符号

call_trtllm_fused_allreduce_norm

关键源码片段

vllm/compilation/passes/fusion/allreduce_rms_fusion.py core-logic

唯一修改的源文件,在 `call_trtllm_fused_allreduce_norm` 函数中添加了输入维度兼容性处理。

def call_trtllm_fused_allreduce_norm(
    allreduce_in: torch.Tensor,
    residual: torch.Tensor,
    rms_gamma: torch.Tensor,
    rms_eps: float,
    world_size: int,
    launch_with_pdl: bool,
    fp32_acc: bool,
    max_token_num: int,
    pattern_code: int,
    norm_out: torch.Tensor | None = None,
    quant_out: torch.Tensor | None = None,
    scale_out: torch.Tensor | None = None,
    scale_factor: torch.Tensor | None = None,
    weight_bias: float = 0.0,
) -> None:
    # handle transformers backend passing outer batch dim.
    # Transformers 后端传递的 residual 可能包含 batch 维度(shape `[batch, tokens, hidden]`)
    # 而此函数预期 2D 输入 `(tokens, hidden)`,需要将其展平。
    if allreduce_in.dim() != 2:
        hidden = allreduce_in.shape[-1]
        allreduce_in = allreduce_in.view(-1, hidden)
        residual = residual.view(-1, hidden)
        if norm_out is not None:
            norm_out = norm_out.view(-1, hidden)
​
    num_tokens, hidden_size = allreduce_in.shape
    # ... 后续逻辑不变

评论区精华

添加注释说明 Transformers 后端场景 documentation

ZJY0516 建议在维度检测处添加注释说明这是为 Transformers 后端准备的。

结论:作者在下一个 commit 中补充了注释,该 thread 已解决。 · 已解决

风险与影响

风险较低。变更只检查输入维度并在非 2D 时展平,不改变 2D 输入的原有路径,且 shape 调整基于 -1 推断,不会改变元素总数或跨行连续性。但需注意:若后续有 4D 或更高维度输入且语义与 (batch, tokens, hidden) 不同(例如 (batch, heads, tokens, hidden)),展平可能错误合并维度。当前仅服务 transformers 后端场景,风险可控。

  • 直接修复了 Transformers 后端在 TP>1 时使用 TRT-LLM 融合 allreduce+rms_norm 的崩溃问题。
  • 影响范围限于开启了 model_impl='transformers' 且使用兼容 GPU(具备 FlashInfer 融合能力)的用户。
  • 不会影响 vLLM 原生后端或其他融合模式。变更量为 7 行,侵入性极小。
核心路径变更

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论