已开启
[Bug]: [inductor][triton_experimental] torch.compile 生成的融合归约 kernel 读错地址:多级融合轴索引拆分不完备,layer_norm 输出错误数值(mobilevit_s accuracy fail #4307
AllenGuan创建于 21 天前
21 天前 添加了label:triage-review
TorchNPU-Bot
21 天前 评论:
21 天前 评论:
issue待分派,添加triage-review标签


21 天前 添加了label:bug
在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。
⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:
在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用
<TOKEN>等占位符替代原有内容。环境信息
🐛 问题描述
torch.compile(npu_backend=triton_experimental)编译"多步 reshape/transpose 视图链 + 逐元素加 + layer_norm"的融合归约 kernel 时,第三遍(normalize)load 的地址表达式残留在编译期未拆净的///%复合形式,且该表达式与正确地址不等价——融合轴把 batch 数位错误地混入空间数位,导致绝大多数元素读到错误地址、layer_norm 输出数值错误。复现步骤(自包含最小片段,任意目录直接运行):
import torch import torch_npu # noqa: F401 import torch.nn.functional as F torch.manual_seed(0) DEV = "npu" B, C, H, W = 128, 144, 32, 32 feat = torch.randn(B, C, H, W, device=DEV) # conv 输出(NCHW 连续) proj2d = torch.randn(B * 4 * 256, C, device=DEV) # mm 输出(行主) bias = torch.randn(C, device=DEV) ln_w = torch.randn(C, device=DEV) ln_b = torch.randn(C, device=DEV) def target(feat, proj2d, bias, ln_w, ln_b): # unfold 视图链(timm mobilevit 的 reshape/transpose 原样形态) x = feat.reshape(B * C * 16, 2, 16, 2).transpose(1, 2) x = x.reshape(B, C, 256, 4).transpose(1, 3) x = x.reshape(B * 4, 256, C) t = x + (bias + proj2d).view(B * 4, 256, C) return F.layer_norm(t, (C,), ln_w, ln_b) with torch.no_grad(): ref = target(feat, proj2d, bias, ln_w, ln_b) out = torch.compile(target, options={"npu_backend": "triton_experimental"})( feat, proj2d, bias, ln_w, ln_b ) torch.npu.synchronize() d = (out - ref).abs() print(f"max|d|={d.max().item():.3e} mean|d|={d.mean().item():.3e}") assert torch.allclose(out, ref, atol=2e-4, rtol=2e-4), "miscompile"模型级印证:timm
mobilevit_s(bs=128,inference)经同一后端编译后,社区 accuracy 流程判fail_accuracy(fp64 金标,tol=0.001,实测 max_abs=0.2739)。期望行为:compiled 与 eager 数值一致(fp32 量级,本例 max|d| ≈ 1.9e-06)。
实际行为:上述最小片段
max|d|=1.105e+01;断言失败。关键证据(
TORCH_COMPILE_DEBUG=1取得的 output_code.py,同名融合 kernel)——三遍 load 同一块输入,前两遍是干净的系数多项式,第三遍带未化简残式:# pass 1/2(正确): tmp0 = tl.load(in_ptr0 + (x8 + 2*x6 + 32*x9 + 64*x7 + 1024*r0_2 + 147456*x5), ...) # pass 3(错误): tmp20 = tl.load(in_ptr0 + (2*((((x1 + 4*x0) // 4) % 16)) + 32*x9 + 64*((x1 + 4*x0) // 64) + 1024*r0_2 + 147456*x5 + ((x1 % 2))), ...)kernel 内轴定义:
x0 = x6 + 16*x7、x4 = x8 + 2*x9、x1 = x4 + 4*x5(x1 是长 512 的融合轴,内含 batch 数位 x5)。全域枚举验证:残式 − 干净式= 2*x5 + 32*((x5+x6)//16),即 x5>0 的全部迭代点地址皆错(99.2%,仅 x5=0 板块幸免);例如 (x5=1, 其余=0) 时干净式地址 147456、残式 147458。第三遍 normalize 读错行,μ/σ 与被归一化数据不配套,输出错误。初步定位(供维护者参考):
0. 记号与两个算法
记法:
(a//b)%c即代码里的ModularIndexing(a,b,c)。轴的混合进制分解(全部是定义):读法(声明:非代码符号,仅助读):x5=batch(128)、x6/x7=W/H 高数位、x8/x9=W/H 低数位、
r0_2=通道(144)。正确地址(NCHW 单价表的必然读数):
算法 A(拆位,
_split_fused_axes_round):输入一个索引表达式;目标是找到「裸符号做底数的
a//c或(a//1)%c」,据此认定 a 是被压扁的融合轴(a = inner + c·outer),给两个数位各发一个新变量,替换
a//c → outer、a%c → inner、a → c·outer + inner。目的:把地址变成独立小轴的线性组合。每条替换都是恒等式(新变量的定义就是 a 的数位)。
算法 B(换绑,
_simplify_compound_indexing):输入一个还残留///%的表达式;目标是找到
s//d(s 是树上节点,节点 s 的 divisor 记为 δ),在同一棵树里找 divisor 恰为δ·d的既有节点 o,替换
s//d → o(Mod 形态(s//d)%m → o%m)。数学依据是整除链恒等式:隐含前提:s 必须真的是 T//δ(从 0 起覆盖全幅)。树上节点的元数据只有 (divisor, length)
两个数——"全幅起点"与"某次拆位拆出的内腔"在元数据里不可区分。
案发链条:A 只执行一轮就停 → 产物里留下一根"内腔轴"(T 的余数段,重锚在某个 divisor
槽位上)→ B 对它套恒等式时前提不成立 → 换绑不等价 → 错地址 → 数值错,且全程无编译错。
1. 坏 pass:pass 3 的 normalize load(master 单轮版)
1.1 算法 A 执行(仅一轮)
【本轮输入】【实测】
【算法目标】:找到
x3//c/(x3//1)%c,把 x3 的数位拆成新变量,实现"地址→小轴线性组合"。【证据收集】:x3 身上有三个模数的证据——
x3//1024、(x3//1)%256、x3//512。缺陷①(setdefault):每个符号只登记第一个模数——本轮登记 c=1024,256/512 的证据被丢弃
(x3 的数位分解是三层的
x3 = x0 + 256·x4 + 1024·x5,一轮只拆得动一层)。【替换】(按定义 x3 = 1024·x5 + x10,x10 := x3%1024):
(x3//1024)x5((x3//1)%256)((x10//1)%256)((x3//256)%4)(复合底数内)(x10//256)((x3//256)%4)%2(独立项)(x10//256)%2((x3//512)%2)(x10//512)【产物】【实测,P1】
【等价性】✓ 每条替换都是数位定义或 mod 恒等式,P1 ≡ P0。
【遗留】★
x10 = x0 + 256·x4(长 1024)自身仍是融合轴,压着两层旧数位,且裸轴证据齐全(
(x10//512)、(x10//256)、(x10//1)%256)——再执行一轮算法 A 即可拆掉。缺陷②(单轮):master 到此返回。半拆解的 P1 交给算法 B。
1.2 算法 B 执行(换绑)
【本轮输入】P1。【算法目标】:找到
s//d,用树上 divisor 对得上的既有节点 o 替换s//d → o。【前提】(T//δ)//d = T//(δ·d)要求 s = T//δ(全幅)。【逐原子替换】(x10 节点 divisor=1,δ=1):
(x10//512)→ 找 divisor=512 的节点 x9(前遍拆 x4 时铸)→x9。核对:
x10 = x0 + 256·x8 + 512·x9且x0+256·x8 < 512⇒x10//512 = x9精确成立。碰巧对。
(x10//256)→ 找 divisor=256 的节点 x1(原生 512 轴)→x1。★ 这里不等价。真值
x10//256 = (x0+256·x4)//256 = x4(值域 [0,4));写入x1 = x4 + 4·x5(值域 [0,512))。多出 4·x5。根源:恒等式要求 x10 = T//1 = 全幅,实际 x10 = T%1024(T = 1024·x5 + x10 ⇒
T//256 = 4·x5 + x4 = x1,而 x10//256 = x4——差恰为 4·x5)。divisor 字段无法区分。
(x10//256)%2→ 同一匹配 →x1%2。形式上同一条错误替换,数值侥幸无损:x1%2 = (4·x5+x4)%2 = x4%2,4·x5 在 mod 2 下湮灭。(x10//256) + 4·((x10//1)%256):底数非裸符号,整除提取不动;但第 2 条的结果已传导进底数;其中
((x10//1)%256)由 B 内置第二趟 mini 拆位正确替换为x0(x10%256 = x0,恒等)。底数成为
4·x0 + x1。★ 悲剧时序:第二趟本登记了 x10 的正确分解(x10 = 256·x4 + x0),但
(x10//256)证据原子已被第 2 条吃掉——同一根轴:%256 用途走正确的 x0,//256 用途走错误的 x1。
(4·x0+x1)//64(x1 不透明、长 512)——残式定形。【产物】【实测,P3】
【哪里不等价、差多少】:全域枚举
P3 − P0=2·x5 + 32·⌊(x5+x6)/16⌋——x5>0(batch≠0)的一切迭代点地址皆错,全域 99.2% 不等价。反例:x5=1、其余=0 时正确地址
147456、P3 给 147458。
1.3 不等价如何变成数值错误
tmp20 = tl.load(in_ptr0 + (2*(((x1+4*x0)//4)%16) + 64*((x1+4*x0)//64) + …)。替换产物是合法表达式,无断言失败——编译不报错,安静落纸。
x1 = x4 + 4*x5——多出的 4·x5 从此在设备上真实执行,每个 b>0 的迭代点读错行。
配套 ⇒ min_repro
max|d|=1.105e+01⇒ 模型fail_accuracy(0.2739)。欢迎加入社区,感谢您对社区的贡献 🎉!