记一次Triton-Ascend融合算子把variance算成0的排查

前言

事情的起点很简单,Qwen3.5-4B 在 Ascend 950PR 上可以正常启动,接口也能返回内容,但是输出完全是胡言乱语。

最麻烦的地方是服务看起来没有坏: 模型加载成功,显存正常,HTTP 也是 200,没有任何一个日志能看出来是某个融合算子算错了。只能从最终输出向前一层层比较中间值。

最后定位到的是 vLLM-Ascend 中的 split_qkv_rmsnorm_mrope 融合算子:K 的 RMSNorm variance 本来应该在 1 左右,实际却变成了四个 0。

1
2
3
4
5
正确结果:
[1.181689, 1.050795, 1.118657, 1.191830]

融合算子结果:
[0, 0, 0, 0]

测试环境

1
2
3
4
5
6
7
8
9
10
11
Device:          Ascend 950PR
Model: Qwen3.5-4B
Model dtype: BF16
vLLM: 0.21.0
vLLM-Ascend: 0.21.0rc1
triton-ascend: 3.2.1
num_q_heads: 16
num_kv_heads: 4
head_size: 256
rope_dim: 64
partial RoPE: enabled

最后发现这个问题应该和triton的版本有关

复现代码在
https://github.com/vllm-project/vllm-ascend/issues/12642

为什么四个0会让整个模型坏掉

当时模型的关键形状是:

1
2
3
4
num_q_heads  = 16
num_kv_heads = 4
head_size = 256
rope_dim = 64

K RMSNorm 的核心计算并不复杂:

1
2
3
4
squares = in_k_tensor * in_k_tensor
variances = tl.sum(squares, axis=1) / head_size
reciprocal_std = 1 / tl.sqrt(variances + eps)
k_normalized = in_k_tensor * reciprocal_std

如果 variance 被算成 0,在 eps=1e-6 时就会得到:

1
1 / sqrt(0 + 1e-6) = 1000

也就是说 K 会在归一化阶段被放大约 1000 倍,后面的 RoPE 和 Attention 即使全部正确,拿到的也已经是坏数据了。

先确认到底是哪一步错了

一开始不能直接假设 tl.sum 有问题。我把表达式拆开,并分别把中间结果写回设备内存:

1
2
3
squares = in_k_tensor * in_k_tensor
square_sums = tl.sum(squares, axis=1)
variances = square_sums / head_size

结果很奇怪:

1
2
3
squares:正确
square_sums:正确
variances:全0

四个平方和是:

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
2
3
4
5
6
7
orig_qk = extract_slice(
k_normalized,
offsets=(0, 0),
sizes=(num_kv_heads, rope_dim),
strides=(1, 1),
)
roped_k = orig_qk * cos_tensor

逐步删除 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
2
3
4
5
6
[4,256] reduction
→ [4] scalar scale
→ [4,1] broadcast
→ [4,256] normalized tensor
→ extract_slice [4,64]
→ 与 [1,64] 广播相乘

继续改变 shape 后还发现:num_kv_heads=2 正常,num_kv_heads=4 错误。Q 有 16 个 head,所以同一套数学公式里 Q 正常、K 错误。

这一步很重要,因为它把问题从“RoPE 公式可能写错了”缩小到了“特定 shape 和布局组合触发了错误代码生成”。

最小规避方案

在编译器修复前,需要先让模型能够正确推理。最后采用的规避方式是在 reduction 后增加一次临时 store:

1
2
3
4
5
6
7
8
9
10
11
squares = in_k_tensor * in_k_tensor
square_sums = tl.sum(squares, axis=1)

tl.store(
out_k_ptr
+ (block_offset + index) * kv_size
+ tl.arange(0, num_kv_heads),
square_sums,
)

variances = square_sums / head_size

这次写入会强迫编译器把 reduction 结果真正落到内存,切断错误的融合路径。

为了不增加 kernel 参数和临时 Tensor,我复用了当前 token 的 out_k 区域。这里只会临时覆盖前四个元素,kernel 末尾会把完整 K 写回,因此接口和最终输出都不变。

fix: 更新triton到nightly即可解决, 未知bug

作者

Noah Shen

发布于

2026-07-22

更新于

2026-07-22

许可协议

评论