记一次Triton-Ascend融合算子把variance算成0的排查
前言
事情的起点很简单,Qwen3.5-4B 在 Ascend 950PR 上可以正常启动,接口也能返回内容,但是输出完全是胡言乱语。
最麻烦的地方是服务看起来没有坏: 模型加载成功,显存正常,HTTP 也是 200,没有任何一个日志能看出来是某个融合算子算错了。只能从最终输出向前一层层比较中间值。
最后定位到的是 vLLM-Ascend 中的 split_qkv_rmsnorm_mrope 融合算子:K 的 RMSNorm variance 本来应该在 1 左右,实际却变成了四个 0。
1 | 正确结果: |
测试环境
1 | Device: Ascend 950PR |
最后发现这个问题应该和triton的版本有关
复现代码在
https://github.com/vllm-project/vllm-ascend/issues/12642
为什么四个0会让整个模型坏掉
当时模型的关键形状是:
1 | num_q_heads = 16 |
K RMSNorm 的核心计算并不复杂:
1 | squares = in_k_tensor * in_k_tensor |
如果 variance 被算成 0,在 eps=1e-6 时就会得到:
1 | 1 / sqrt(0 + 1e-6) = 1000 |
也就是说 K 会在归一化阶段被放大约 1000 倍,后面的 RoPE 和 Attention 即使全部正确,拿到的也已经是坏数据了。
先确认到底是哪一步错了
一开始不能直接假设 tl.sum 有问题。我把表达式拆开,并分别把中间结果写回设备内存:
1 | squares = in_k_tensor * in_k_tensor |
结果很奇怪:
1 | squares:正确 |
四个平方和是:
1 | [226.8115, 255.1978, 266.9960, 268.8226] |
除以 256 后应该得到约 0.89 到 1.05,不存在精度下溢的可能。
继续替换缩放操作后,现象更明显:
| 表达式 | 结果 |
|---|---|
square_sums * 1.0 |
正确 |
square_sums * 2.0 |
全0 |
square_sums * 0.5 |
全0 |
square_sums / 2 |
全0 |
square_sums / 256 |
全0 |
* 1.0 很可能被编译器当成恒等操作直接消掉了,只要真的生成一次缩放,结果就会错。
错误居然由后面的计算触发
更反直觉的是,单独运行上面的 RMSNorm,结果又完全正确。
只有后面继续执行这段 partial RoPE 计算时,前面的 variance 才会变成 0:
1 | orig_qk = extract_slice( |
逐步删除 kernel 中的计算,得到下面的边界:
| 保留的计算 | K variance |
|---|---|
| 只计算 RMSNorm | 正确 |
extract_slice(k_normalized) |
正确 |
直接输出 orig_qk |
正确 |
cat_y * sin_tensor |
正确 |
orig_qk * cos_tensor |
全0 |
| 完整 RoPE | 全0 |
后面的乘法当然不会在运行时穿越回去修改前面的变量。真正发生的是 Triton 会编译整个 kernel,后续消费者改变了前面中间值的布局和代码生成方式。
当下面这条数据流同时出现时,Ascend 后端会生成错误代码:
1 | [4,256] reduction |
继续改变 shape 后还发现:num_kv_heads=2 正常,num_kv_heads=4 错误。Q 有 16 个 head,所以同一套数学公式里 Q 正常、K 错误。
这一步很重要,因为它把问题从“RoPE 公式可能写错了”缩小到了“特定 shape 和布局组合触发了错误代码生成”。
最小规避方案
在编译器修复前,需要先让模型能够正确推理。最后采用的规避方式是在 reduction 后增加一次临时 store:
1 | squares = in_k_tensor * in_k_tensor |
这次写入会强迫编译器把 reduction 结果真正落到内存,切断错误的融合路径。
为了不增加 kernel 参数和临时 Tensor,我复用了当前 token 的 out_k 区域。这里只会临时覆盖前四个元素,kernel 末尾会把完整 K 写回,因此接口和最终输出都不变。
fix: 更新triton到nightly即可解决, 未知bug
记一次Triton-Ascend融合算子把variance算成0的排查
https://blog.novashen.top/2026/07/22/tech/AI Infra/记一次Triton-Ascend融合算子把variance算成0的排查/