已合并
[Fix] [inductor] Fix the compilation failure issue of Scan IR under NPU + Inductor environment #31313
zhudada0120创建于 3月3日
[Fix] [inductor] Fix the compilation failure issue of Scan IR under NPU + Inductor environment #31313
已合并
zhudada0120创建于 3月3日
共 3 个文件变更+285-20
@@ -0,0 +1,94 @@
1+import torch
2+from torch.testing._internal.common_utils import (
3+ run_tests,
4+ parametrize,
5+ instantiate_parametrized_tests,
6+)
7+from testutils import TestUtils
8+import torch_npu
9+ 
10+ 
11+def _make_fn_add_cumsum_sum(dim: int):
12+ def fn(x):
13+ a = x + 1
14+ b = torch.ops.aten.cumsum(a, dim=dim)
15+ c = torch.ops.aten.sum(b, dim=dim)
16+ return c
17+ 
18+ return fn
19+ 
20+ 
21+def _make_fn_add_cumsum_sum_diff_dim(dim: int):
22+ def fn(x):
23+ a = x + 1
24+ b = torch.ops.aten.cumsum(a, dim=dim)
25+ c = torch.ops.aten.sum(b, dim=-1)
26+ return c
27+ 
28+ return fn
29+ 
30+ 
31+def _make_fn_add_cumsum_add(dim: int):
32+ def fn(x):
33+ a = x + 1
34+ b = torch.ops.aten.cumsum(a, dim=dim)
35+ c = b + 1
36+ return c
37+ 
38+ return fn
39+ 
40+ 
41+def _build_test_cases():
42+ # fn spec: (name, fn_factory, input_dtype, rtol, atol)
43+ fn_cases = [
44+ ("add_cumsum_sum", _make_fn_add_cumsum_sum, torch.int32, 0, 0),
45+ ("add_cumsum_add", _make_fn_add_cumsum_add, torch.int32, 0, 0),
46+ ("add_cumsum_sum_diff_dim", _make_fn_add_cumsum_sum_diff_dim, torch.int32, 0, 0),
47+ ]
48+ 
49+ # shape/dim spec: ((shape), dim)
50+ shape_dim_cases = [
51+ ((3, 4, 5), 0),
52+ ((3, 4, 5), 1),
53+ ((3, 4, 5), -1),
54+ ((8192, 4, 5), 0),
55+ ((3, 8192, 5), 1),
56+ ((3, 4, 8192), -1),
57+ ((3, 4), 0),
58+ ((3, 4), -1),
59+ ]
60+ 
61+ test_cases = []
62+ for shape, dim in shape_dim_cases:
63+ for fn_name, fn_factory, input_dtype, rtol, atol in fn_cases:
64+ test_cases.append((shape, dim, fn_name, fn_factory, input_dtype, rtol, atol))
65+ return test_cases
66+ 
67+ 
68+TEST_CASES = _build_test_cases()
69+ 
70+ 
71+class TestScan(TestUtils):
72+ @parametrize(
73+ "shape, dim, fn_name, fn_factory, input_dtype, rtol, atol",
74+ TEST_CASES,
75+ )
76+ def test_scan_aten_op(self, shape, dim, fn_name, fn_factory, input_dtype, rtol, atol):
77+ 
78+ x = torch.ones(shape, device=torch.device("npu"), dtype=input_dtype)
79+ 
80+ fn = fn_factory(dim)
81+ compiled = torch.compile(fn, backend="inductor", dynamic=False)
82+ 
83+ out_inductor = compiled(x).to(torch.int32)
84+ out_eager = fn(x).to(torch.int32)
85+ max_abs_diff = (out_inductor - out_eager).abs().max().item()
86+ print(f"=== fn={fn_name} case={(shape, dim)} dtype={input_dtype} ===")
87+ print(f"allclose: {torch.allclose(out_inductor, out_eager, rtol=rtol, atol=atol)}")
88+ self.assertEqual(out_eager, out_inductor, rtol=rtol, atol=atol)
89+ 
90+ 
91+instantiate_parametrized_tests(TestScan)
92+ 
93+if __name__ == "__main__":
94+ run_tests()
@@ -121,12 +121,11 @@ class NPUTritonScheduling(TritonScheduling):
121 self, kernel_features: SIMDKernelFeatures, kernel_args, kernel_kwargs121 self, kernel_features: SIMDKernelFeatures, kernel_args, kernel_kwargs
122 ) -> List[SIMDKernel]:122 ) -> List[SIMDKernel]:
123 123 
124- return [124+ if kernel_features.contains_op("scan"):
W
Wweizhan43月4日

增加Scan类Op的用例

likedislike
zhudada0120
zhudada0120
3月9日 评论:
125- self.kernel_type(125+ kernel_kwargs = dict(kernel_kwargs)
126- *kernel_args,126+ kernel_kwargs["override_cooperative_reduction"] = False
127- **kernel_kwargs,127+ 
128- )128+ return [self.kernel_type(*kernel_args, **kernel_kwargs)]
129- ]
130 129 
131 # transform indexing before call codegen_node_schedule_with_kernel130 # transform indexing before call codegen_node_schedule_with_kernel
132 def codegen_node_schedule(self, kernel_features: SIMDKernelFeatures, nodes):131 def codegen_node_schedule(self, kernel_features: SIMDKernelFeatures, nodes):
@@ -657,12 +656,16 @@ class NPUTritonScheduling(TritonScheduling):
657 split_tiling = SplitTiling(kernel)656 split_tiling = SplitTiling(kernel)
658 split_tiling.select_split_tiling_axis()657 split_tiling.select_split_tiling_axis()
659 kernel.load_store_indexing = split_tiling.indexing658 kernel.load_store_indexing = split_tiling.indexing
660- # ReductionAnalysis depends on kernel.load_store_indexing 659+ # ReductionAnalysis depends on kernel.load_store_indexing.
661- if kernel.inside_reduction:660+ if kernel.inside_reduction and getattr(kernel, "find_reduction_node", None) is not None:
662- kernel.reduce_analysis = ReductionAnalysis(kernel)661+ from torch._inductor import ir
663- # pure_simt_kernel, high dim reduction don't use persitent reduction662+ 
664- if kernel.is_unified_simt_kernel() and kernel.reduction_dim() != len(kernel.golden_var_list) - 1:663+ reduction_node = kernel.find_reduction_node()
665- kernel.persistent_reduction = False664+ if reduction_node is not None and isinstance(reduction_node, ir.Reduction):
665+ kernel.reduce_analysis = ReductionAnalysis(kernel)
666+ # pure_simt_kernel, high dim reduction don't use persitent reduction
667+ if kernel.is_unified_simt_kernel() and kernel.reduction_dim() != len(kernel.golden_var_list) - 1:
668+ kernel.persistent_reduction = False
666 # no_loop_axis depends on persistent reduction669 # no_loop_axis depends on persistent reduction
667 split_tiling.select_no_loop_axis()670 split_tiling.select_no_loop_axis()
668 671 
@@ -549,6 +549,169 @@ class NPUIndexTritonKernel(TritonKernel):
549 def _get_grid_type(self) -> type[triton_heuristics.GridExpr]:549 def _get_grid_type(self) -> type[triton_heuristics.GridExpr]:
550 return npu_triton_heuristics.GridNpu550 return npu_triton_heuristics.GridNpu
551 551 
552+ 
553+ def scan(
554+ self,
555+ dtypes: tuple[torch.dtype, ...],
556+ combine_fn: Callable[
557+ [tuple[CSEVariable, ...], tuple[CSEVariable, ...]], tuple[CSEVariable, ...]
558+ ],
559+ values: tuple[CSEVariable, ...],
560+ ) -> tuple[CSEVariable, ...]:
561+ """NPU override for ops.scan codegen.
562+ 
563+ Upstream TritonKernel.scan assumes the reduction/scan dimension is the
564+ last dimension in the broadcasted tensor (layout [X, R]) and uses
565+ `dim = triton_tensor_ndim() - num_reduction_dims`.
566+ 
567+ NPU index codegen may produce either an R-first ([R, X...]) or
568+ R-last ([X..., R]) dense layout depending on how `golden_var_list` is
569+ derived for the current kernel. We must scan along the actual reduction
570+ ("r") axis position in the broadcasted dense tensor, instead of
571+ hard-coding an axis.
572+ """
573+ 
574+ assert self.inside_reduction
575+ assert not self.cooperative_reduction, "TODO"
576+ 
577+ masks = OrderedSet(f"{tree.prefix}mask" for tree in self.range_trees)
578+ self.filter_masks(masks)
579+ masks = sorted(masks)
580+ assert not self._load_mask, "ops.scan not supported inside ops.masked"
581+ 
582+ broadcasted_values: list[CSEVariable] = []
583+ accumulators: list[CSEVariable] = []
584+ 
585+ dtypes = tuple(upcast_compute_type(dtype) for dtype in dtypes)
586+ cse_compute = functools.partial(self.cse.generate, self.compute)
587+ combine_helper_fn = self._lift_helper(combine_fn, len(values), dtypes)
588+ 
589+ # Pick the scan dimension by locating the reduction axis in the dense
590+ # layout used by dense_size_list()/dense_size_str().
591+ #
592+ # dense_size_list() orders dims according to reversed(golden_var_list),
593+ # i.e. the i-th size corresponds to reversed(golden_var_list)[i].
594+ _ = self.dense_size_list()
595+ golden = list(self.golden_var_list or [])
596+ dense_vars = list(reversed(golden))
597+ reduction_dims = [
D
Ddezheng8893月5日

这个地方你是不是有限制呀?reduction_dims 要是同时reduce 两维 这不就错了吗?如果只考虑一种场景,请添加注释

likedislike
zhudada0120
zhudada0120
3月6日 评论:
598+ i
599+ for i, v in enumerate(dense_vars)
600+ if str(v).startswith("r")
601+ ]
602+ # Scan is single-axis and only fuses with reductions on the same r-axis.
603+ if reduction_dims:
604+ dim = reduction_dims[-1]
605+ else:
606+ # Fallback to upstream heuristic if we failed to infer the layout.
607+ dim = self.triton_tensor_ndim() - self.num_reduction_dims
608+ 
609+ # Derive per-axis reduction-loop symbol names from the scan dim.
610+ dense_ndim = len(self.dense_size_list())
611+ scan_axis_sym = dense_vars[dim] if 0 <= dim < len(dense_vars) else None
612+ scan_axis = (
613+ self.range_tree_nodes.get(scan_axis_sym) if scan_axis_sym in self.range_tree_nodes else None
614+ )
615+ scan_axis_name = (
616+ getattr(scan_axis, "name", None) or (str(scan_axis_sym) if scan_axis_sym is not None else "r")
617+ )
618+ rbase_sym = f"base_{scan_axis_name}"
619+ rblock_sym = f"{scan_axis_name.upper()}BLOCK_SUB"
620+ if getattr(scan_axis, "is_split_axis", False):
621+ roffset_expr = f"{scan_axis_name}_offset + (loop_{scan_axis_name} * {rblock_sym})"
622+ else:
623+ roffset_expr = f"(loop_{scan_axis_name} * {rblock_sym})"
624+ reshape_sizes = ["1"] * dense_ndim
625+ if 0 <= dim < dense_ndim:
626+ reshape_sizes[dim] = rblock_sym
627+ rbase_broadcast = (
628+ f"tl.broadcast_to({rbase_sym}.reshape({', '.join(reshape_sizes)}), {self.dense_size_str()})"
629+ if dense_ndim > 1
630+ else rbase_sym
631+ )
632+ 
633+ for value, dtype in zip(values, dtypes):
634+ value_dtype = self.cse.generate(
635+ self.compute,
636+ f"{value}.to({triton_compute_type(dtype)})",
637+ dtype=dtype,
638+ )
639+ value = self.cse.generate(
640+ self.compute,
641+ f"tl.broadcast_to({value_dtype}, {self.dense_size_str()})",
642+ dtype=dtype,
643+ )
644+ broadcasted_values.append(value)
645+ 
646+ acc_type = triton_acc_type(dtype)
647+ 
648+ if not self.persistent_reduction:
649+ accumulator = self.cse.newvar(dtype=dtype)
650+ reduced_size = self.dense_size_list()
651+ reduced_size[dim] = "1"
652+ reduced_size = f"[{', '.join(reduced_size)}]"
653+ 
654+ default = "float('nan')" if dtype.is_floating_point else "-1"
655+ self.body.writeline(
656+ f"{accumulator} = tl.full({reduced_size}, {default}, {acc_type})"
657+ )
658+ accumulators.append(accumulator)
659+ 
660+ def csv(vs):
661+ return " ".join(f"{v}," for v in vs)
662+ 
663+ def cse_multiple(line, in_values, in_masks, in_dtypes):
664+ n = len(in_values)
665+ cache_keys = [f"{line}, {i}, {in_masks}" for i in range(n)]
666+ if all(self.cse.contains(cache_key) for cache_key in cache_keys):
667+ return [self.cse.get(cache_key) for cache_key in cache_keys]
668+ result_vars = [self.cse.newvar(dtype=_dtype) for _dtype in in_dtypes]
669+ self.compute.writeline(f"{csv(result_vars)} = {line}")
670+ for result_var, cache_key in zip(result_vars, cache_keys):
671+ if in_masks:
672+ result_var.mask_vars = in_masks # type: ignore[attr-defined]
673+ self.cse.put(cache_key, result_var)
674+ return tuple(result_vars)
675+ 
676+ partial_scan_vars = cse_multiple(
677+ f"tl.associative_scan(({csv(broadcasted_values)}), {dim}, {combine_helper_fn})",
678+ values,
679+ masks,
680+ dtypes,
681+ )
682+ 
683+ if not self.persistent_reduction:
684+ partial_reduce_vars = [
685+ cse_compute(
686+ f"triton_helpers.select_one(({partial_scan_var}), ({rbase_broadcast}) == ({rblock_sym} - 1), dim={dim}, keep_dims=True)",
687+ dtype=upcast_compute_type(partial_scan_var.dtype),
688+ )
689+ for partial_scan_var in partial_scan_vars
690+ ]
691+ accs_next = combine_fn(tuple(accumulators), tuple(partial_reduce_vars))
692+ full_scan_vars = combine_fn(tuple(accumulators), partial_scan_vars)
693+ result_vars = [
694+ cse_compute(
695+ f"tl.where(({roffset_expr}) > 0, {full_scan}, {partial_scan})",
696+ dtype=partial_scan.dtype,
697+ )
698+ for full_scan, partial_scan in zip(full_scan_vars, partial_scan_vars)
699+ ]
700+ for acc_next, accumulator, partial_reduce in zip(
701+ accs_next, accumulators, partial_reduce_vars
702+ ):
703+ self.compute.writeline(
704+ f"{accumulator} = tl.where(({roffset_expr}) > 0, {acc_next}, {partial_reduce})"
705+ )
706+ else:
707+ result_vars = partial_scan_vars
708+ 
709+ for result_var in result_vars:
710+ assert isinstance(result_var, TritonCSEVariable)
711+ result_var.mask_vars = OrderedSet(masks)
712+ 
713+ return tuple(result_vars)
714+ 
552 def patch_triton_hash(self):715 def patch_triton_hash(self):
553 # remove this method once the original invocation is fixed716 # remove this method once the original invocation is fixed
554 import hashlib717 import hashlib
@@ -1322,11 +1485,12 @@ class NPUIndexTritonKernel(TritonKernel):
1322 1485 
1323 def dense_size_list(self) -> List[str]:1486 def dense_size_list(self) -> List[str]:
1324 if self.inside_reduction:1487 if self.inside_reduction:
1325- if not self.reduce_analysis:1488+ if self.find_reduction_node() is not None:
1326- self.reduce_analysis = ReductionAnalysis(self)1489+ if not self.reduce_analysis:
1327- if self.is_contiguous_reduction():1490+ self.reduce_analysis = ReductionAnalysis(self)
1328- return self.reduce_analysis.dense_post_reduction_list()1491+ if self.is_contiguous_reduction():
1329- return self.reduce_analysis.dense_size_list()1492+ return self.reduce_analysis.dense_post_reduction_list()
1493+ return self.reduce_analysis.dense_size_list()
1330 1494 
1331 if not self.golden_var_list:1495 if not self.golden_var_list:
1332 self.select_golden_varlist()1496 self.select_golden_varlist()
@@ -1362,9 +1526,13 @@ class NPUIndexTritonKernel(TritonKernel):
1362 1526 
1363 def dense_size_str(self):1527 def dense_size_str(self):
1364 if self.inside_reduction:1528 if self.inside_reduction:
1365- if not self.reduce_analysis:1529+ if self.find_reduction_node() is not None:
1366- self.reduce_analysis = ReductionAnalysis(self)1530+ if not self.reduce_analysis:
1367- return self.reduce_analysis.dense_size_str()1531+ self.reduce_analysis = ReductionAnalysis(self)
1532+ return self.reduce_analysis.dense_size_str()
1533+ # Scan fallback: generate dense shape directly.
1534+ sizes = self.dense_size_list()
1535+ return f"[{', '.join(sizes)}]"
1368 sizes = self.dense_size_list()1536 sizes = self.dense_size_list()
1369 return f"[{', '.join(sizes)}]"1537 return f"[{', '.join(sizes)}]"
1370 1538