已合并
[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
已合并
共 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 | + | ||
| 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_kwargs | 121 | 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 | |||
| 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_kernel | 130 | # 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.indexing | 658 | 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 reduction | 662 | + |
| 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 = False | 664 | + 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 reduction | 669 | # 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.GridNpu | 550 | 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 这个地方你是不是有限制呀?reduction_dims 要是同时reduce 两维 这不就错了吗?如果只考虑一种场景,请添加注释 ![]() ![]() | |||
| 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 fixed | 716 | # remove this method once the original invocation is fixed |
| 554 | import hashlib | 717 | 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 | ||


增加Scan类Op的用例