已合并
fix: iterate _maybe_split_fused_axes to fixpoint (multi-level digit extraction) #45161
AllenGuan创建于 14 天前
fix: iterate _maybe_split_fused_axes to fixpoint (multi-level digit extraction) #45161
已合并
Pull Request已成功合入, 合并人@ascend-robot
(感谢 AllenGuan 的贡献)atomgit-bot
14 天前 评论:
14 天前 评论:
变更摘要
本 PR 修复 torch_npu/_inductor/triton_experimental/codegen/triton.py 中 _maybe_split_fused_axes 单轮拆分不完备导致的 kernel 访存地址公式错误问题。原实现每个融合符号在一轮内只登记第一个模数且不再重跑,当迭代空间为多层进位结构(如 131072 = 128*4*256)时,一轮只拆最外层,新暴露的融合轴残留在复合底数(如 ModularIndexing(4*x0 + x1, ...))中,下游 _simplify_compound_indexing 对不透明融合轴按 range tree 换绑后产生错误地址直达生成的 kernel(实测融合 LN kernel 第三遍 load 地址全域 99.2% 不等价)。本次修改将原单轮实现改名为 _split_fused_axes_round 并原样保留其语义,同时把 _maybe_split_fused_axes 改为有界 fixpoint 包装器,迭代拆分至不动点,使索引收敛为叶轴的纯仿射表达式,仅改动工序⑤ codegen 的索引准备段,不涉及 IR、发射结构、接口与资料变更。
主要改动
- 新增有界 fixpoint 迭代包装器
_maybe_split_fused_axes: 迭代调用_split_fused_axes_round,每轮对变换后的表达式重新扫描证据,使上一轮新暴露的融合符号(如 1024-inner)在后续轮次继续被拆分,直至索引收敛为叶轴纯线性表达式,并新增类常量_SPLIT_FIXPOINT_MAX_ROUNDS = 6作为轮次上限(131072-flat 轴实测 ≤4 轮,6 留余量防活锁)。 - 双终止条件: 以
new_index == index(收敛)或new_index == prev(替换振荡兜底)任一命中即退出循环,避免无效迭代与振荡。 - 原单轮实现改名
_split_fused_axes_round: 函数体零改动,保留一轮拆分的全部原有语义,作为 fixpoint 迭代的单步原语被新包装器调用。


ascend-robot
14 天前 评论:
14 天前 评论:
atomgit-bot
14 天前 评论:
14 天前 评论:
14 天前 添加了label:ascend-cla/yes
此处折叠了166条消息 查看更多
8 天前 添加了label:approvedlgtm
8 天前 合入了pull request
AtlasAccount
8 天前 评论:
8 天前 评论:
流水线 pytorch_gitcode_PR_multiVersion#14692 [ commitID:38ec8e04 ] 已完成


建议 PR 标题:
fix(inductor): iterate _maybe_split_fused_axes to fixpoint (multi-level digit extraction)【合入来源】
社区 issue 链接待创建后补充(bug 现象与复现最小片段见该 issue)。
【修改方案】
背景:triton_experimental 代码生成为了让 NPU 后端获得线性(可向量化证明)的访存地址,在
prepare_indexing流水线中用_maybe_split_fused_axes把"多根叶轴压扁而成的融合轴"按FloorDiv/ModularIndexing证据拆回独立子轴。现实现存在双重不完备:一轮里candidates.setdefault对每个融合符号只登记第一个模数,且一轮结束不再重跑。当迭代空间是多层进位结构(如131072 = 128*4*256)时,一轮只拆掉最外层,替换后新暴露的融合轴(及其数位)残留在复合底数(如(4*x0 + x1))里;下游_simplify_compound_indexing对不透明融合轴只能按 range tree 做结构换绑,会把"值域小于模数的商式"换绑成同名长轴——错误的地址公式直达生成的 kernel(实测:融合 LN kernel 第三遍 load 地址全域 99.2% 不等价,输出错误数值)。【改动】算法 A 的替换规则一字未改,只把"执行一次"改成"重复执行直到表达式不再变化"
(
new_index == index收敛退出;== prev防振荡;上限 6 轮防活锁)。【轮 1】输入 P0 → 与 §1.1 完全相同 → P1,x10 暴露。区别仅在:不返回。
【轮 2】输入 P1。重新扫证据——证据在 x10 身上。登记 c=512,替换
(x10 = 512·x9 + x11,x11 := x10%512,长 512):
(x10//512) → x9(恒等)(x10//256) → (x11//256) + 2·x9(512·x9 整除提取,恒等)(x10//1)%256 → (x11//1)%256(512·x9 ≡ 0 mod 256,恒等)新暴露 x11(长 512)仍压一层,证据在场。(实录拆位链:1024→512→256。)
【轮 3】输入轮 2 产物。拆 x11 @ c=256,替换(x11 = 256·x8 + x0):
(x11//256) → x8、((x11//1)%256) → x0(恒等)。★ 铸新变量走树缓存:
lookup(1,256)命中既有 x0、lookup(256,2)命中既有 x8——数位最终落回前两遍 load 已在用的规范叶轴(三遍多项式逐字一致的机械保证)。
**【轮 4】**总账重放代入早前登记的
x0 → 16·x7 + x6,范围推理收尾:(4·x0 + x4)//64 = x7(分子低位 4·x6+x4 ≤ 63,归零)((4·x0 + x4)//4)%16 = x0%16 = x6扫描无新证据 →
new_index == index→ 收敛退出。产物:纯线性六项,零///%:【每轮等价性】✓ 全程只用算法 A 原有的恒等替换——fixpoint 不改变替换的数学性质,只是
拆到没有融合轴为止。
**【算法 B 随后的行为】**输入 A' 里一个
///%原子都不剩 → B 入口卫兵if not index.has(...)直接返回——B 一条替换都没执行。不是它被修聪明了,是它的输入里再没有能让它犯错的原子。
组件交互:仅工序⑤ codegen 的索引准备段(
triton_experimental/codegen/triton.py)内部变化,不触碰其他 pass、不改 IR 与发射结构。取舍说明:不在下游_simplify_compound_indexing里对半拆解形态加容错——半拆解形态本身就是错误地址的产房,在拆位工序内收敛到纯仿射才是根治,下游无需感知。【资料变更】
不涉及
【接口变更】
不涉及
【功能验证】
F.layer_norm(迭代空间 131072×144,融合归约 kernel)——修复前 compiled vs eagermax|d|=1.105e+01;修复后max|d|=1.907e-06,且三遍 load 的地址多项式逐字一致(修复前第三遍残留2*((((x1 + 4*x0) // 4) % 16)) + 64*((x1 + 4*x0) // 64) + ((x1 % 2))残式,修复后与第一/二遍同为x8 + 2*x6 + 32*x9 + 64*x7 + 1024*r0_2 + 147456*x5)。tl.load/tl.store访存行,修复后含///%残式的行数为 0。mobilevit_s(bs=128,inference,无 fallback)——accuracy 对比 max_abs 0.2739 → 0.0546(残差为融合 kernel 与 aclnn 的合法算法差异经深网络的放大,非错误地址:NPU eager 自身距 fp64 金标 0.43,compiled 0.425,两者同距);性能上 LN 融合族 174.7 → 94.5 ms/fwd(-46%),榜首归约 kernelaiv_scalar_ratio0.421 → 0.0011(残式的逐元素标量寻址消失,转为纯访存受限)。test_triton_experimental_autotune.py66 passed、test_triton_experimental_enable.py13 passed(清缓存后全量)。pytorch/test/_inductor/test_triton_experimental_split_fused_axes.py,15 例全 PASS):不走端到端编译,直调三道 pass(_maybe_split_fused_axes/_split_fused_axes_round/_simplify_compound_indexing),配真IterationRangesRoot节点铸造与独立 sizevars:套件经三视角对抗审查(含对实现的变异分析);未覆盖的既有分支(simplify 的 MI 换绑臂、fold_fused_mod 体、split_root_modulo redirect、符号模数)已在文件头 Scope 注记为后续项。
【CheckList】