Prhub

#46987 [XPU] [RMSNorm] revert weightless change on xpu

原始 PR 作者 zufangzhu 合并时间 2026-06-30 03:03 文件变更 1 提交数 2 评论 0 代码增减 +4 / -11

执行摘要

XPU RMSNorm 回退 weightless 优化

在 XPU 上,weightless RMSNorm(weight=None)会回退到 native 实现,分解成大量小 aten 算子(pow、mean、rsqrt、mul、add),小 batch 时每个 kernel 都很小,模型受限于 host/dispatch,显著增加 CPU dispatch 的 kernel 数量,损害 TPOT。详见 PR body 描述。

建议精读了解 XPU 平台上的 kernel dispatch 优化策略。回退逻辑清晰,评审充分,可安全合入。

讨论亮点

本 PR 没有 reviewer 评论或讨论。两位 reviewers(xinyu-intel、jikunshang)均直接批准,未提出异议。

实现拆解

修改 vllm/kernels/xpu_ops.pyrms_normfused_add_rms_norm 两个函数的 weightless 分支:

  1. rms_norm 函数(+4/-11):当 weight is None 时,不再调用 native 实现,而是创建与 x 的最后一个维度相同形状的全 1 张量作为 weight,然后直接调用 torch.ops._C.rms_norm 完成计算。
  2. fused_add_rms_norm 函数(+4/-11):类似地,当 weight is None 时,不再调用 native 实现并逐元素复制,而是创建全 1 weight 张量,直接调用 torch.ops._C.fused_add_rms_norm 完成计算。

该变更仅涉及一个文件,无新增测试,由代码逻辑显式保证:全 1 weight 在 RMSNorm 中相当于无权重缩放,数学上等价于 weightless 情况。

文件 模块 状态 重要度
vllm/kernels/xpu_ops.py XPU 内核 modified 5.79

关键符号

rms_norm fused_add_rms_norm

关键源码片段

vllm/kernels/xpu_ops.py core-logic

核心变更文件,修改了 RMSNorm 和 fused_add_rms_norm 两个函数的 weightless 分支,用构造全 1 weight 代替 native fallback。

    def rms_norm(
        x: Tensor, weight: Tensor | None, epsilon: float, variance_size: int | None = None
    ) -> Tensor:
        assert variance_size is None
        if weight is None:
            # 在 XPU 上,原生 _C kernel 不接受 weight=None。
            # 之前 fallback 到 native 实现会分解为大量小算子,
            # 导致 CPU dispatch 开销大。这里构造全 1 的 weight 张量,
            # 数学上等价于 weightless(因为 RMSNorm 中 weight 为 1 相当于无缩放),
            # 且能直接调用 _C kernel,避免多次 kernel launch。
            weight = torch.ones(x.shape[-1], device=x.device, dtype=x.dtype)
        output = torch.empty(x.shape, device=x.device, dtype=x.dtype)
        torch.ops._C.rms_norm(output, x, weight, epsilon)
        return output

评论区精华

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

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

风险与影响

  • 正确性风险:使用全 1 weight 替代 weight=None 在数学上等价(因为 RMSNorm 公式中 weight 为 1 相当于无缩放),但需确认 _C.rms_norm kernel 对全 1 weight 的处理是否与 weightless 完全一致。
  • 性能风险:构造 torch.ones 张量有微小开销,但远小于 native 实现的多次 kernel launch,整体收益明确。
  • 兼容性风险:仅影响 XPU 平台,不涉及其他后端或量化路径。
  • 影响范围:仅 XPU 平台,修改了两个函数:rms_normfused_add_rms_norm,影响所有使用 RMSNorm 且 weight 为 None 的场景(通常在某些模型 head 或实验性结构中)。
  • 性能影响:小 batch 场景下 TPOT 显著改善,因减少了 CPU dispatch 次数。
  • 代码维护:代码量减少(-11 行),逻辑更简洁,不再依赖 native fallback。
仅影响 XPU 平台

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论