已合并
动态shape场景下表达式中存在Max及除数是符号变量的单项式//和%运算适配 #42762
zzll创建于 7月25日
动态shape场景下表达式中存在Max及除数是符号变量的单项式//和%运算适配 #42762
已合并
zzll创建于 7月25日
3 个文件变更+99-13
@@ -9,7 +9,7 @@ from torch.testing._internal.common_utils import run_tests, parametrize, instant
9from torch._inductor import config9from torch._inductor import config
10from testutils import TestUtils10from testutils import TestUtils
11import torch_npu11import torch_npu
12-from unittest import skip12+ 
13 13 
14os.environ["INDUCTOR_ASCEND_DUMP_FX_GRAPH"] = "1"14os.environ["INDUCTOR_ASCEND_DUMP_FX_GRAPH"] = "1"
15os.environ["TORCH_COMPILE_DEBUG"] = "1"15os.environ["TORCH_COMPILE_DEBUG"] = "1"
@@ -68,7 +68,7 @@ class TestDebugMsg(TestUtils):
68 content68 content
69 )69 )
70 70 
71- @skip("dynamic linear skip")71+ 
72 @parametrize('shape_x', [(32, 8, 64)])72 @parametrize('shape_x', [(32, 8, 64)])
73 @parametrize('shape_y', [(32, 1, 64)])73 @parametrize('shape_y', [(32, 1, 64)])
74 @parametrize('dtype', ['float32'])74 @parametrize('dtype', ['float32'])
@@ -269,18 +269,7 @@ def rebuild_flattened_dims(indexing):
269 if find_index_in_substitute(index, kernel):269 if find_index_in_substitute(index, kernel):
270 new_index = sympy_subs(index, kernel.expr_substituted)270 new_index = sympy_subs(index, kernel.expr_substituted)
271 indexing[key] = new_index271 indexing[key] = new_index
272- # 删除kernel.expr_substituted中未使用的升维轴
273- for key, index in indexing.items():
274- index_symbols = index.free_symbols
275- remaining_del = [v for v in kernel.expr_substituted.values() if v not in index_symbols]
276- kernel.dim_up_temp.update(remaining_del)
277 272 
278- if len(kernel.dim_up_temp) > 0:
279- keys_to_remove = [k for k, v in kernel.expr_substituted.items() if v in kernel.dim_up_temp]
280- for expr in keys_to_remove:
281- del kernel.expr_substituted[expr]
282- for var in kernel.dim_up_temp:
283- del kernel.range_tree_nodes[var]
284 log.debug(273 log.debug(
285 "rebuild_flattened_dims: range_tree_nodes_substituted=%s, store_items=%s",274 "rebuild_flattened_dims: range_tree_nodes_substituted=%s, store_items=%s",
286 kernel.range_tree_nodes_substituted,275 kernel.range_tree_nodes_substituted,
@@ -1,3 +1,4 @@
1+import collections
1import contextlib2import contextlib
2import dataclasses3import dataclasses
3import functools4import functools
@@ -1505,10 +1506,106 @@ class NPUIndexTritonKernel(TritonKernel):
1505 with self:1506 with self:
1506 self._mark_store_index_keys()1507 self._mark_store_index_keys()
1507 self._transform_schedule_indexing()1508 self._transform_schedule_indexing()
1509+ self._remove_unused_dim_up_axes()
1508 self._record_store_unified_indexing()1510 self._record_store_unified_indexing()
1509 self._remove_substituted_dims_from_kernel()1511 self._remove_substituted_dims_from_kernel()
1510 self._finalize_kernel_codegen_dims()1512 self._finalize_kernel_codegen_dims()
1511 1513 
1514+ 
1515+ def _index_symbol_users(self):
1516+ """
1517+ Map every symbol the scheduled nodes' indices reference to the index keys
1518+ using it, or None when a node has no transformed indexing yet.
1519+ 
1520+ indexing_exprs are still expressed in loop-body vars instead of kernel
1521+ axes, so a node without transformed indexing carries no usable axis
1522+ information and a partial map would make the missing axes look unused.
1523+ """
1524+ axis_users = collections.defaultdict(list)
1525+ for node in self._iter_schedule_nodes(self.node_schedule):
1526+ indexing = node._body.indexing
1527+ if indexing is None:
1528+ log.warning("%s has no transformed indexing", node)
1529+ return None
1530+ for key, index in indexing.items():
1531+ for sym in getattr(index, "free_symbols", set()):
1532+ axis_users[sym].append(key)
1533+ return axis_users
1534+ 
1535+ 
1536+ def _drop_dead_substitution_candidates(self, removed_axes):
1537+ """
1538+ Forget parent expansions that rebuild a removed axis.
1539+ 
1540+ Their axes no longer exist in the kernel, so substituting such an
1541+ expansion into an index would reference a variable codegen never defines.
1542+ """
1543+ for var in list(self.range_tree_nodes_substituted):
1544+ candidates = self.range_tree_nodes_substituted[var]
1545+ alive = [
1546+ candidate
1547+ for candidate in candidates
1548+ if not getattr(candidate[1], "free_symbols", set()) & removed_axes
1549+ ]
1550+ if len(alive) == len(candidates):
1551+ continue
1552+ if alive:
1553+ self.range_tree_nodes_substituted[var] = alive
1554+ else:
1555+ del self.range_tree_nodes_substituted[var]
atomgit-bot
atomgit-botatomgit-bot7月25日

🟡 Medium Priority

_drop_dead_substitution_candidates 方法中,当 alive(未被移除的候选列表)既非空又不等于原始 candidates 时(即只有部分候选因涉及已移除的轴而失效),代码只执行了 log.info,却未将 self.range_tree_nodes_substituted[var] 更新为 alive。失效的候选表达式仍保留在字典中,后续 _remove_substituted_dims_from_kernel(line 1721)和 _should_preserve_substituted_var(line 1704)会读取这些过时的候选,导致:

  1. len(candidates) 被高估,可能错误保留不应保留的轴;
  2. 失效候选表达式引用了已从 range_tree_nodes 移除的符号变量,若后续被用于索引替换,codegen 会引用未定义的变量。

触发条件:range_tree_nodes_substituted 中某个 var 有多个候选(多个 (length, expr) 对),且其中至少一个候选的 free_symbols 与已移除轴相交,但不是全部。

修复方向:在 if not alive: 分支后面增加 else: self.range_tree_nodes_substituted[var] = alive

建议:在 if not alive: 分支后添加 else: 分支,将 self.range_tree_nodes_substituted[var] 赋值为过滤后的 alive 列表。

改动建议
1555
+ if len(alive) == len(candidates):
1556
+ continue
1557
+ if not alive:
1555
- del self.range_tree_nodes_substituted[var]
1558
+ del self.range_tree_nodes_substituted[var]
1559
+ else:
1560
+ self.range_tree_nodes_substituted[var] = alive
应用建议
likedislike
zzll
7月25日 评论:
1556+ log.info(
1557+ "Dropped expansion candidates of %s rebuilding removed axes: %s",
1558+ var,
1559+ [candidate for candidate in candidates if candidate not in alive],
1560+ )
atomgit-bot
atomgit-botatomgit-bot7月25日

🟡 Medium Priority

_drop_dead_substitution_candidates 方法(第 1536–1558 行)在清理引用已移除轴的候选展开式时存在逻辑缺陷:当 alive 列表既不为空也不等于 candidates 全部(即部分候选存活、部分已死),代码既不 continue 也不更新 self.range_tree_nodes_substituted[var],导致已死的候选条目残留在 range_tree_nodes_substituted 中。

具体路径(第 1544–1553 行): alive = [candidate for candidate in candidates if not getattr(candidate[1], "free_symbols", set()) & removed_axes] if len(alive) == len(candidates): continue # 全活 → 跳过 if not alive: del self.range_tree_nodes_substituted[var] # 全死 → 删除

BUG: 部分存活时,fall-through 到 log.info,但未执行

self.range_tree_nodes_substituted[var] = alive

后果:range_tree_nodes_substituted 中残留引用已移除符号的候选展开式,后续 _should_preserve_substituted_var 可能因 len(candidates) 包含死条目而做出错误判断;substituted_dims_in_indexing 若使用这些死候选,会引用代码生成中已不存在的变量。

建议:在 if not alive: 分支后增加 else: 分支,将 self.range_tree_nodes_substituted[var] = alive,确保部分存活时列表被正确更新为仅含活候选。

likedislike
zzll
7月25日 评论:
1561+ 
1562+ 
1563+ def _remove_unused_dim_up_axes(self):
1564+ """
1565+ Drop rebuilt (dim-up) axes that no index expression references any more.
1566+ 
1567+ An axis survives exactly when some index still references it, so codegen
1568+ always defines the axes the generated index expressions use, and nothing
1569+ else. Must run after substituted_dims_in_indexing has expanded the parent
1570+ axes: a rebuilt axis may only show up once its parent is expanded, e.g. a
1571+ broadcast dim that the Store index alone uses. Judging it any earlier
1572+ drops an axis that is still needed, and codegen then emits no definition
1573+ for a variable the Store index already references.
1574+ """
1575+ axis_users = self._index_symbol_users()
1576+ if axis_users is None:
1577+ log.warning("Keep all dim-up axes, transformed indexing is incomplete")
1578+ return
1579+ 
1580+ unused_dims = {
1581+ v for v in self.expr_substituted.values() if v not in axis_users
1582+ }
1583+ log.debug(
1584+ "dim-up axes: unused=%s, kept=%s",
1585+ sorted(unused_dims, key=str),
1586+ {
1587+ str(v): axis_users[v]
1588+ for v in self.expr_substituted.values()
1589+ if v in axis_users
1590+ },
1591+ )
1592+ if not unused_dims:
1593+ return
1594+ 
1595+ self.dim_up_temp.update(unused_dims)
1596+ keys_to_remove = [
1597+ k for k, v in self.expr_substituted.items() if v in unused_dims
1598+ ]
1599+ for expr in keys_to_remove:
1600+ del self.expr_substituted[expr]
1601+ for var in unused_dims:
1602+ node = self.range_tree_nodes.get(var)
1603+ if node is not None:
1604+ node.parent.remove_entry(var)
1605+ self._drop_dead_substitution_candidates(unused_dims)
1606+ log.info("Removed unused dim-up axes: %s", sorted(unused_dims, key=str))
1607+ 
1608+ 
1512 def _store_keeps_unified_anchor(self, var, index):1609 def _store_keeps_unified_anchor(self, var, index):
1513 """1610 """
1514 Check if a Store index keeps the unified axis as the only axis with the same prefix.1611 Check if a Store index keeps the unified axis as the only axis with the same prefix.