Prhub

#42379 [Bugfix] Fix RMSNorm kernels to multiply in weight's native dtype

原始 PR 作者 liulanze 合并时间 2026-05-30 14:16 文件变更 2 提交数 2 评论 17 代码增减 +10 / -23

执行摘要

修复 RMSNorm 内核中权重始终以 FP32 相乘的回归 bug

修复由PR #40860引入的回归:RMSNorm CUDA内核在最终乘法时始终将权重转换为float32,与vllm/ir/ops/layernorm.pyx = x.to(weight.dtype) * weight规范不符。导致每层约0.03的误差,在多层Q/K RMSNorm模型中(如DeepSeek V4)累积至完全错误(关联Issue #42325)。

这是一个精准的回归修复案例,值得阅读关键源码片段和讨论精华。对数值精度敏感的开发者可深入了解CUDA内核中dtype控制的权衡。推荐精读。

讨论亮点

@yewentao256 要求提供lm_eval指标:"Thanks for the work! Please add a lm_eval metrics to show we don't hurt in e2e acc"。作者补充TinyLlama结果(评分不变)后,yewentao256批准。

@zyongye 提出设计分歧:"Only check with python IR implementation is not ideal... ideally all non Tensor core operation shall be operate on FP32 to maximize precision",建议修改IR实现而非CUDA内核。作者以精度测试无回归为由坚持原方案,最终未采纳该建议,PR按原方式合并。

CI失败确认:@AndreasKaratzas 分析AMD MI300上的故障为硬件问题,与PR无关,并协助force merge。

实现拆解

  1. 修复 rms_norm_kernel 及其fused变体:在 csrc/libtorch_stable/layernorm_kernels.cu 的vectorized、FP16/BF16优化路径和通用scalar路径下,将 x * s_variance * w(w已转为float)改为 static_cast<scalar_t>(x * s_variance) * weight_elem,确保缩放到标量类型后再与原生dtype权重相乘。

  2. 修复 rms_norm_static_fp8_quant_kernel 及其fused变体:在 csrc/libtorch_stable/layernorm_quant_kernels.cu 的对应三处位置同步修改,维持fused路径与非fused复合路径的数值一致性,否则 test_fused_rms_norm_quant 会因精度边界分歧而失败。

  3. 验证:通过 tests/kernels/core/test_layernorm.pytests/kernels/ir/test_layernorm.py 回归测试(100%通过),并在TinyLlama上运行lm_eval确认端到端精度无变化。

文件 模块 状态 重要度
csrc/libtorch_stable/layernorm_kernels.cu CUDA 内核 modified 4.24
csrc/libtorch_stable/layernorm_quant_kernels.cu CUDA 内核 modified 4.06

关键符号

rms_norm_kernel fused_add_rms_norm_kernel rms_norm_static_fp8_quant_kernel fused_add_rms_norm_static_fp8_quant_kernel

关键源码片段

csrc/libtorch_stable/layernorm_kernels.cu core-logic

核心修复文件,包含 rms_norm_kernel 和 fused_add_rms_norm_kernel 的三个实现路径,是本次 bugfix 的主体。

// csrc/libtorch_stable/layernorm_kernels.cu - 修复后的 rms_norm_kernel (vectorized 路径 )
#pragma unroll
for (int j = 0; j < VEC_SIZE; j++) {
    float x = static_cast<float>(src1.val[j]);
    // 之前 (regressed) 的代码:
    // float w = static_cast<float>(src2.val[j]);
    // dst.val[j] = static_cast<scalar_t>(x * s_variance * w);
    // 修复后:先缩放至 scalar_t(如 bfloat16),再乘以权重原生 dtype 值
    // 这样可以保证乘法结果精确匹配 Python 参考实现 `x.to(weight.dtype) * weight`
    dst.val[j] = static_cast<scalar_t>(x * s_variance) * src2.val[j];
}

评论区精华

要求提供 lm_eval 端到端精度验证 测试

yewentao256 要求提供 lm_eval 指标以证明端到端精度无退化。作者补充了 TinyLlama 上的结果(评分完全一致)。

结论:作者补充数据后,yewentao256 批准。 · 已解决

关于在 FP32 还是原生 dtype 中乘法的设计分歧 设计

zyongye 认为所有非 Tnesor Core 操作应在 FP32 中最大化精度,建议修改 IR 而非 CUDA 实现。作者以测试无回归为由坚持当前修复方案。

结论:未达成共识,但 PR 按作者方案合并。后续是否调整为 IR 驱动尚未定论。 · 已解决

风险与影响

  • 回归风险:对依赖FP32乘法精度的极端情况(如自定义训练后权重)可能有微小行为变化,但主流路径和端到端测试均无影响。
  • 兼容性风险:FP8量化内核同步修改需保证fused与非fused路径一致,现有测试已覆盖。
  • 数值精度敏感:部分模型(如DeepSeek V4)的Q/K norm层可能间接依赖FP32的额外精度,但视觉上无退化。
  • 无性能影响:仅改变乘法精度顺序,计算量不变。
  • 用户影响:使用BF16/FP16权重RMSNorm的模型(DeepSeek V4、LLaDA-2等)将恢复正确的数值输出,输出与v0.19.1一致;对FP32权重模型无影响。
  • 系统影响:无,仅CUDA内核代码变更。
  • 团队影响:修复跨版本回归,提升数值稳定性,减少后续诊断成本。
核心路径变更 回归风险 数值精度敏感

关联 Issue

#40860 [Feat] DeepSeek V4 Rebased
#42325 [Bug]: RMSNorm kernel ignores weight dtype, always uses FP32 (regression in v0.20.0)

完整报告

参与讨论