已开启
[Bug]: [inductor][triton_experimental] torch.compile 生成的融合归约 kernel 读错地址:多级融合轴索引拆分不完备,layer_norm 输出错误数值(mobilevit_s accuracy fail #4307
AllenGuan创建于  21 天前
AllenGuan成员
21 天前 创建

在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。

⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:

  • API 令牌或密钥
  • 密码或身份验证凭证
  • 私有网址或接口地址
  • 个人或机密数据
  • ...

在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用 <TOKEN> 等占位符替代原有内容。

环境信息

- 操作系统 Ubuntu 22.04.5 LTS (aarch64)
- 昇腾硬件信息 910B2
- CANN软件版本 9.1.0-beta.1
- 安装的对应软件版本 torch 2.13.0+cpu / torch_npu 2.13.0(源码 master 分支 c8ec201682 构建)/ triton-ascend 3.2.2

🐛 问题描述

torch.compilenpu_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*x7x4 = x8 + 2*x9x1 = 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)。轴的混合进制分解(全部是定义):

x3 = x0 + 256·x1        x1 = x4 + 4·x5        x0 = x6 + 16·x7        x4 = x8 + 2·x9

读法(声明:非代码符号,仅助读):x5=batch(128)、x6/x7=W/H 高数位、x8/x9=W/H 低数位、
r0_2=通道(144)。正确地址(NCHW 单价表的必然读数):

addr = 147456·x5 + 1024·r0_2 + 64·x7 + 32·x9 + 2·x6 + x8

算法 A(拆位,_split_fused_axes_round:输入一个索引表达式;目标是找到「裸符号做底数
a//c(a//1)%c」,据此认定 a 是被压扁的融合轴(a = inner + c·outer),给两个数位
各发一个新变量,替换 a//c → outera%c → innera → 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)。数学依据是整除链恒等式:

(T//δ)//d = T//(δ·d)

隐含前提:s 必须真的是 T//δ(从 0 起覆盖全幅)。树上节点的元数据只有 (divisor, length)
两个数——"全幅起点"与"某次拆位拆出的内腔"在元数据里不可区分

案发链条:A 只执行一轮就停 → 产物里留下一根"内腔轴"(T 的余数段,重锚在某个 divisor
槽位上)→ B 对它套恒等式时前提不成立 → 换绑不等价 → 错地址 → 数值错,且全程无编译错。


1. 坏 pass:pass 3 的 normalize load(master 单轮版)

1.1 算法 A 执行(仅一轮)

【本轮输入】【实测】

P0 = 1024*r0_2 + 147456*(x3//1024)
   + 64*(( 4*((x3//1)%256) + ((x3//256)%4) )//64)
   + 32*((x3//512)%2)
   + 2*(( 4*((x3//1)%256) + ((x3//256)%4) )//4)%16
   + ((x3//256)%4)%2

【算法目标】:找到 x3//c / (x3//1)%c,把 x3 的数位拆成新变量,实现"地址→小轴线性
组合"。【证据收集】:x3 身上有三个模数的证据——x3//1024(x3//1)%256x3//512
缺陷①(setdefault):每个符号只登记第一个模数——本轮登记 c=1024,256/512 的证据被丢弃
(x3 的数位分解是三层的 x3 = x0 + 256·x4 + 1024·x5,一轮只拆得动一层)。

【替换】(按定义 x3 = 1024·x5 + x10,x10 := x3%1024):

P0 中的子式 替换后 数学依据
(x3//1024) x5 定义
((x3//1)%256) ((x10//1)%256) 1024·x5 ≡ 0 (mod 256)
((x3//256)%4)(复合底数内) (x10//256) 折到 x10 数位后,x10//256 < 4 ⇒ %4 恒等消失
((x3//256)%4)%2(独立项) (x10//256)%2 同上,仅 %2 保留
((x3//512)%2) (x10//512) 2·x5 ≡ 0 (mod 2);x10//512 < 2 ⇒ %2 恒等消失

【产物】【实测,P1】

P1 = 1024*r0_2 + 147456*x5 + 32*(x10//512)
   + 64*(( (x10//256) + 4*((x10//1)%256) )//64)
   + (x10//256)%2
   + 2*(( (x10//256) + 4*((x10//1)%256) )//4)%16

【等价性】✓ 每条替换都是数位定义或 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):

  1. (x10//512) → 找 divisor=512 的节点 x9(前遍拆 x4 时铸)→ x9
    核对:x10 = x0 + 256·x8 + 512·x9x0+256·x8 < 512x10//512 = x9 精确成立。
    碰巧对
  2. (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 字段无法区分。
  3. (x10//256)%2 → 同一匹配 → x1%2。形式上同一条错误替换,数值侥幸无损:
    x1%2 = (4·x5+x4)%2 = x4%2,4·x5 在 mod 2 下湮灭。
  4. 复合底数 (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
  5. 终趟范围推理收不动 (4·x0+x1)//64(x1 不透明、长 512)——残式定形

【产物】【实测,P3】

P3 = 1024*r0_2 + 147456*x5 + 32*x9 + 64*((4*x0 + x1)//64) + (x1//1)%2 + 2*((4*x0 + x1)//4)%16

【哪里不等价、差多少】:全域枚举 P3 − P0 = 2·x5 + 32·⌊(x5+x6)/16⌋——x5>0
(batch≠0)的一切迭代点地址皆错,全域 99.2% 不等价。反例:x5=1、其余=0 时正确地址
147456、P3 给 147458。

1.3 不等价如何变成数值错误

  • P3 原样打印进 kernel:tmp20 = tl.load(in_ptr0 + (2*(((x1+4*x0)//4)%16) + 64*((x1+4*x0)//64) + …)
    替换产物是合法表达式,无断言失败——编译不报错,安静落纸
  • kernel 头部为被引用融合轴发射重建赋值 x1 = x4 + 4*x5——多出的 4·x5 从此在设备上
    真实执行
    ,每个 b>0 的迭代点读错行。
  • pass 3 是 normalize:μ/σ 来自地址正确的前两遍,被归一化数据从错行读来——分子分母不
    配套 ⇒ min_repro max|d|=1.105e+01 ⇒ 模型 fail_accuracy(0.2739)。

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
AAllenGuan成员
21 天前 关联了看板:FrameworkPTAdapter 版本issue看板
TorchNPU-BotTorchNPU-Bot成员
21 天前 添加了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
21 天前 评论:

issue待分派,添加triage-review标签

likedislike
ascend-robotascend-robot成员
21 天前 添加了label:bug
AAllenGuan成员
21 天前 修改了issue 的描述
AAllenGuan成员
21 天前 修改了issue 的描述
AAllenGuan成员
19 天前 修改了issue 的描述