已合并
[inductor]support pointwise group #42716
rain-666创建于 7月25日
[inductor]support pointwise group #42716
已合并
共 2 个文件变更+264-10
| @@ -7,7 +7,9 @@ Part 1: Multi-operator fusion tests that exercise NPU's SplitTiling system: | |||
| 7 | - Indirect memory ops (gather/scatter/index_select) with SIMT/SIMD paths | 7 | - Indirect memory ops (gather/scatter/index_select) with SIMT/SIMD paths |
| 8 | - View/reshape/cat + compute fusion with dynamic shapes | 8 | - View/reshape/cat + compute fusion with dynamic shapes |
| 9 | 9 | ||
| 10 | -Part 2: Nonzero data-dependent dynamic shape tests: | 10 | +Part 2: Pointwise symbolic grouping policy tests. |
| 11 | + | ||
| 12 | +Part 3: Nonzero data-dependent dynamic shape tests: | ||
| 11 | - nonzero produces data-dependent output shape[0] = count of non-zero elements | 13 | - nonzero produces data-dependent output shape[0] = count of non-zero elements |
| 12 | - Unbacked SymInt guarded by runtime assertions under dynamo | 14 | - Unbacked SymInt guarded by runtime assertions under dynamo |
| 13 | 15 | ||
| @@ -17,7 +19,10 @@ Reference patterns from: | |||
| 17 | """ | 19 | """ |
| 18 | 20 | ||
| 19 | import unittest | 21 | import unittest |
| 22 | +from types import SimpleNamespace | ||
| 23 | +from unittest.mock import patch | ||
| 20 | 24 | ||
| 25 | +import sympy | ||
| 21 | import torch | 26 | import torch |
| 22 | import torch.nn.functional as F | 27 | import torch.nn.functional as F |
| 23 | from torch._dynamo.testing import CompileCounterWithBackend | 28 | from torch._dynamo.testing import CompileCounterWithBackend |
| @@ -31,6 +36,10 @@ from torch.testing._internal.common_utils import ( | |||
| 31 | 36 | ||
| 32 | import torch_npu | 37 | import torch_npu |
| 33 | import torch_npu._inductor | 38 | import torch_npu._inductor |
| 39 | +from torch_npu._inductor.codegen import split_tiling as split_tiling_module | ||
| 40 | +from torch_npu._inductor.codegen.split_tiling import SplitTiling | ||
| 41 | +from torch_npu._inductor.config import num_vector_core | ||
| 42 | +from torch_npu._inductor.runtime import triton_heuristics | ||
| 34 | 43 | ||
| 35 | torch._dynamo.config.cache_size_limit = 128 | 44 | torch._dynamo.config.cache_size_limit = 128 |
| 36 | 45 | ||
| @@ -40,6 +49,16 @@ if not torch.npu.is_available(): | |||
| 40 | device = "npu" | 49 | device = "npu" |
| 41 | 50 | ||
| 42 | 51 | ||
| 52 | +def make_axis(name, length): | ||
| 53 | + symbol = sympy.Symbol(name) | ||
| 54 | + return SimpleNamespace( | ||
| 55 | + name=name, | ||
| 56 | + prefix=name[0], | ||
| 57 | + length=length, | ||
| 58 | + symbol=lambda: symbol, | ||
| 59 | + ) | ||
| 60 | + | ||
| 61 | + | ||
| 43 | # ============================================================ | 62 | # ============================================================ |
| 44 | # Base class: provides compile counting + dynamic=True helpers | 63 | # Base class: provides compile counting + dynamic=True helpers |
| 45 | # ============================================================ | 64 | # ============================================================ |
| @@ -630,7 +649,147 @@ class TestSymbolicGroupElementwise(TestCase): | |||
| 630 | self._run_and_check(fn, inputs, next_inputs, None) | 649 | self._run_and_check(fn, inputs, next_inputs, None) |
| 631 | 650 | ||
| 632 | # ============================================================ | 651 | # ============================================================ |
| 633 | -# Part 2: Nonzero Data-Dependent Dynamic Shape Tests | 652 | +# Part 2: Pointwise Symbolic Grouping Policy Tests |
| 653 | +# ============================================================ | ||
| 654 | + | ||
| 655 | +class TestPointwiseSymbolicGrouping(TestCase): | ||
| 656 | + def test_static_split_dynamic_tiling_group(self): | ||
| 657 | + # Transpose and broadcast commonly leave the dynamic inner axis tiled | ||
| 658 | + # while a static outer axis supplies grid parallelism. | ||
| 659 | + split_axis = make_axis("x0", sympy.Integer(128)) | ||
| 660 | + tiling_axis = make_axis("x1", sympy.Symbol("s0", positive=True)) | ||
| 661 | + kernel = SimpleNamespace( | ||
| 662 | + persistent_reduction=False, | ||
| 663 | + inside_reduction=False, | ||
| 664 | + sorted_axis=[split_axis, tiling_axis], | ||
| 665 | + split_axis=[split_axis], | ||
| 666 | + tiling_axis=[tiling_axis], | ||
| 667 | + features=SimpleNamespace(scheduler_nodes=lambda: ()), | ||
| 668 | + ) | ||
| 669 | + split_tiling = object.__new__(SplitTiling) | ||
| 670 | + split_tiling.kernel = kernel | ||
| 671 | + x0, x1 = (axis.symbol() for axis in kernel.sorted_axis) | ||
| 672 | + dynamic_stride = sympy.Symbol("s1", positive=True) | ||
| 673 | + split_tiling.indexing = [x1 + 128 * x0, x0 + dynamic_stride * x1] | ||
| 674 | + | ||
| 675 | + sizevars = SimpleNamespace(size_hint=lambda expr: 64) | ||
| 676 | + virtualized = SimpleNamespace(graph=SimpleNamespace(sizevars=sizevars)) | ||
| 677 | + with patch.object(split_tiling_module, "V", virtualized): | ||
| 678 | + self.assertEqual(split_tiling._pointwise_layout_kind(), "transpose") | ||
| 679 | + meta = split_tiling._build_grouped_meta() | ||
| 680 | + | ||
| 681 | + self.assertIsNotNone(meta) | ||
| 682 | + self.assertEqual(meta.primary_group_axis, "x1") | ||
| 683 | + self.assertEqual(meta.static_split_axes, ("x0",)) | ||
| 684 | + self.assertEqual(meta.runtime_block_arg_names, ("X0BLOCK",)) | ||
| 685 | + self.assertEqual(meta.group_features[0].source, "axis") | ||
| 686 | + self.assertEqual(meta.group_features[0].axis_names, ("x1",)) | ||
| 687 | + self.assertEqual(meta.group_features[0].buckets, (64, 128, 256, 512)) | ||
| 688 | + | ||
| 689 | + def test_static_split_dynamic_tiling_broadcast_group(self): | ||
| 690 | + split_axis = make_axis("x0", sympy.Integer(128)) | ||
| 691 | + tiling_axis = make_axis("x1", sympy.Symbol("s0", positive=True)) | ||
| 692 | + kernel = SimpleNamespace( | ||
| 693 | + persistent_reduction=False, | ||
| 694 | + inside_reduction=False, | ||
| 695 | + sorted_axis=[split_axis, tiling_axis], | ||
| 696 | + split_axis=[split_axis], | ||
| 697 | + tiling_axis=[tiling_axis], | ||
| 698 | + features=SimpleNamespace(scheduler_nodes=lambda: ()), | ||
| 699 | + ) | ||
| 700 | + split_tiling = object.__new__(SplitTiling) | ||
| 701 | + split_tiling.kernel = kernel | ||
| 702 | + x0, x1 = (axis.symbol() for axis in kernel.sorted_axis) | ||
| 703 | + split_tiling.indexing = [x0, x0 + 128 * x1] | ||
| 704 | + | ||
| 705 | + self.assertEqual(split_tiling._pointwise_layout_kind(), "broadcast") | ||
| 706 | + meta = split_tiling._build_grouped_meta() | ||
| 707 | + | ||
| 708 | + self.assertIsNotNone(meta) | ||
| 709 | + self.assertEqual(meta.primary_group_axis, "x1") | ||
| 710 | + self.assertEqual(meta.group_features[0].name, "pointwise_broadcast_axis") | ||
| 711 | + self.assertEqual(meta.group_features[0].source, "axis") | ||
| 712 | + self.assertEqual(meta.group_features[0].axis_names, ("x1",)) | ||
| 713 | + self.assertEqual(meta.group_features[0].buckets, (16, 64, 256, 1024, 4096)) | ||
| 714 | + | ||
| 715 | + combined = split_tiling._build_group_features( | ||
| 716 | + None, tiling_axis, "transpose_broadcast" | ||
| 717 | + ) | ||
| 718 | + self.assertEqual(combined[0].name, "pointwise_broadcast_axis") | ||
| 719 | + self.assertEqual(combined[0].buckets, (16, 64, 256, 1024, 4096)) | ||
| 720 | + | ||
| 721 | + def test_existing_pointwise_group_feature_is_unchanged(self): | ||
| 722 | + primary_axis = make_axis("x0", sympy.Symbol("s0", positive=True)) | ||
| 723 | + kernel = SimpleNamespace( | ||
| 724 | + persistent_reduction=False, | ||
| 725 | + inside_reduction=False, | ||
| 726 | + sorted_axis=[primary_axis], | ||
| 727 | + ) | ||
| 728 | + split_tiling = object.__new__(SplitTiling) | ||
| 729 | + split_tiling.kernel = kernel | ||
| 730 | + | ||
| 731 | + features = split_tiling._build_group_features(None, primary_axis) | ||
| 732 | + | ||
| 733 | + self.assertEqual(len(features), 1) | ||
| 734 | + self.assertEqual(features[0].name, "pointwise") | ||
| 735 | + self.assertEqual(features[0].source, "outer_product") | ||
| 736 | + self.assertEqual(features[0].axis_names, ("x0",)) | ||
| 737 | + self.assertEqual(features[0].buckets, (num_vector_core * 4096,)) | ||
| 738 | + | ||
| 739 | + def test_plain_pointwise_does_not_use_tiling_fallback(self): | ||
| 740 | + split_axis = make_axis("x0", sympy.Integer(128)) | ||
| 741 | + tiling_axis = make_axis("x1", sympy.Symbol("s0", positive=True)) | ||
| 742 | + kernel = SimpleNamespace( | ||
| 743 | + persistent_reduction=False, | ||
| 744 | + inside_reduction=False, | ||
| 745 | + sorted_axis=[split_axis, tiling_axis], | ||
| 746 | + split_axis=[split_axis], | ||
| 747 | + tiling_axis=[tiling_axis], | ||
| 748 | + features=SimpleNamespace(scheduler_nodes=lambda: ()), | ||
| 749 | + ) | ||
| 750 | + split_tiling = object.__new__(SplitTiling) | ||
| 751 | + split_tiling.kernel = kernel | ||
| 752 | + x0, x1 = (axis.symbol() for axis in kernel.sorted_axis) | ||
| 753 | + split_tiling.indexing = [x1 + 128 * x0] | ||
| 754 | + | ||
| 755 | + self.assertIsNone(split_tiling._build_grouped_meta()) | ||
| 756 | + | ||
| 757 | + def test_tiling_primary_keeps_static_grid_block(self): | ||
| 758 | + cfg = {"kwargs": {"X0BLOCK": 32, "X1BLOCK_SUB": 64}} | ||
| 759 | + with patch.object(triton_heuristics, "config_to_dict", return_value=cfg): | ||
| 760 | + policy = triton_heuristics.build_grouped_launch_policy( | ||
| 761 | + group_id=0, | ||
| 762 | + cfg=object(), | ||
| 763 | + runtime_block_arg_names=("X0BLOCK",), | ||
| 764 | + group_features=(), | ||
| 765 | + primary_group_axis="x1", | ||
| 766 | + primary_feature_index=0, | ||
| 767 | + axis_env={"x0": 128, "x1": 64}, | ||
| 768 | + npu_num_vector_core=32, | ||
| 769 | + ) | ||
| 770 | + | ||
| 771 | + self.assertEqual(policy["static_blocks"], (("X0BLOCK", 32),)) | ||
| 772 | + self.assertEqual(policy["runtime_block_rules"], ()) | ||
| 773 | + self.assertEqual(policy["grid_target"], 1) | ||
| 774 | + | ||
| 775 | + def test_pointwise_tiling_fallback_requires_one_dynamic_axis(self): | ||
| 776 | + x0 = make_axis("x0", sympy.Symbol("s0", positive=True)) | ||
| 777 | + x1 = make_axis("x1", sympy.Symbol("s1", positive=True)) | ||
| 778 | + kernel = SimpleNamespace( | ||
| 779 | + persistent_reduction=False, | ||
| 780 | + inside_reduction=False, | ||
| 781 | + sorted_axis=[x0, x1], | ||
| 782 | + split_axis=[], | ||
| 783 | + tiling_axis=[x0, x1], | ||
| 784 | + ) | ||
| 785 | + split_tiling = object.__new__(SplitTiling) | ||
| 786 | + split_tiling.kernel = kernel | ||
| 787 | + | ||
| 788 | + self.assertIsNone(split_tiling._dynamic_pointwise_tiling_axis()) | ||
| 789 | + | ||
| 790 | + | ||
| 791 | +# ============================================================ | ||
| 792 | +# Part 3: Nonzero Data-Dependent Dynamic Shape Tests | ||
| 634 | # ============================================================ | 793 | # ============================================================ |
| 635 | 794 | ||
| 636 | class TestNonzeroDynamicShapes(DynamicShapeTestMixin, TestCase): | 795 | class TestNonzeroDynamicShapes(DynamicShapeTestMixin, TestCase): |
| @@ -854,6 +1013,7 @@ class TestNonzeroDynamicShapes(DynamicShapeTestMixin, TestCase): | |||
| 854 | 1013 | ||
| 855 | instantiate_parametrized_tests(TestFusionDynamicShapes) | 1014 | instantiate_parametrized_tests(TestFusionDynamicShapes) |
| 856 | instantiate_parametrized_tests(TestSymbolicGroupElementwise) | 1015 | instantiate_parametrized_tests(TestSymbolicGroupElementwise) |
| 1016 | +instantiate_parametrized_tests(TestPointwiseSymbolicGrouping) | ||
| 857 | instantiate_parametrized_tests(TestNonzeroDynamicShapes) | 1017 | instantiate_parametrized_tests(TestNonzeroDynamicShapes) |
| 858 | 1018 | ||
| 859 | 1019 | ||
| @@ -315,6 +315,80 @@ class SplitTiling: | |||
| 315 | return None | 315 | return None |
| 316 | return axis | 316 | return axis |
| 317 | 317 | ||
| 318 | + def _dynamic_pointwise_tiling_axis(self): | ||
| 319 | + """Return the sole dynamic pointwise axis when it is tiled, not split. | ||
| 320 | + | ||
| 321 | + Transpose and broadcast kernels commonly keep the dynamic inner axis in | ||
| 322 | + the tiling space while a static outer axis drives the grid. That axis can | ||
| 323 | + still select grouped compile variants even though it has no grid block. | ||
| 324 | + """ | ||
| 325 | + if self.kernel.persistent_reduction or self.kernel.inside_reduction: | ||
| 326 | + return None | ||
| 327 | + dynamic_axes = [ | ||
| 328 | + axis | ||
| 329 | + for axis in self.kernel.sorted_axis | ||
| 330 | + if not isinstance(axis.length, sympy.Integer) | ||
| 331 | + ] | ||
| 332 | + if len(dynamic_axes) != 1: | ||
| 333 | + return None | ||
| 334 | + axis = dynamic_axes[0] | ||
| 335 | + if axis in self.kernel.split_axis or axis not in self.kernel.tiling_axis: | ||
| 336 | + return None | ||
| 337 | + if self._pointwise_layout_kind() is None: | ||
| 338 | + return None | ||
| 339 | + return axis | ||
| 340 | + | ||
| 341 | + def _pointwise_layout_kind(self): | ||
| 342 | + """Classify transpose/broadcast from the kernel's memory indexings.""" | ||
| 343 | + axis_symbols = tuple(axis.symbol() for axis in self.kernel.sorted_axis) | ||
| 344 | + if not axis_symbols or not self.indexing: | ||
| 345 | + return None | ||
| 346 | + | ||
| 347 | + stride_orders = set() | ||
| 348 | + has_broadcast = False | ||
| 349 | + for index in self.indexing: | ||
| 350 | + axis_strides = [] | ||
| 351 | + unsupported_index = False | ||
| 352 | + for axis in axis_symbols: | ||
| 353 | + stride = sympy.expand(index).coeff(axis) | ||
| 354 | + if stride == 0 and axis in index.free_symbols: | ||
| 355 | + unsupported_index = True | ||
| 356 | + break | ||
| 357 | + if stride != 0: | ||
| 358 | + axis_strides.append((axis, stride)) | ||
| 359 | + if unsupported_index: | ||
| 360 | + continue | ||
| 361 | + | ||
| 362 | + present = [axis for axis, _ in axis_strides] | ||
| 363 | + if len(present) < len(axis_symbols): | ||
| 364 | + has_broadcast = True | ||
| 365 | + continue | ||
| 366 | + | ||
| 367 | + strides = [] | ||
| 368 | + for axis, stride in axis_strides: | ||
| 369 | + try: | ||
| 370 | + stride = int(stride) | ||
| 371 | + except (TypeError, ValueError): | ||
| 372 | + try: | ||
| 373 | + stride = int(V.graph.sizevars.size_hint(stride)) | ||
| 374 | + except (AttributeError, KeyError, TypeError, ValueError): | ||
| 375 | + strides = [] | ||
| 376 | + break | ||
| 377 | + strides.append((stride, axis)) | ||
| 378 | + if strides: | ||
| 379 | + stride_orders.add( | ||
| 380 | + tuple(axis for _, axis in sorted(strides, key=lambda x: x[0])) | ||
| 381 | + ) | ||
| 382 | + | ||
| 383 | + has_transpose = len(stride_orders) > 1 | ||
| 384 | + if has_transpose and has_broadcast: | ||
| 385 | + return "transpose_broadcast" | ||
| 386 | + if has_transpose: | ||
| 387 | + return "transpose" | ||
| 388 | + if has_broadcast: | ||
| 389 | + return "broadcast" | ||
| 390 | + return None | ||
| 391 | + | ||
| 318 | def non_reduction_axis_names(self): | 392 | def non_reduction_axis_names(self): |
| 319 | return tuple(axis.name for axis in self.kernel.sorted_axis if axis.prefix != "r") | 393 | return tuple(axis.name for axis in self.kernel.sorted_axis if axis.prefix != "r") |
| 320 | 394 | ||
| @@ -342,7 +416,7 @@ class SplitTiling: | |||
| 342 | _REDUCTION_BUCKETS = (8192,) | 416 | _REDUCTION_BUCKETS = (8192,) |
| 343 | _OUTER_BUCKETS = (256,) | 417 | _OUTER_BUCKETS = (256,) |
| 344 | 418 | ||
| 345 | - def _build_group_features(self, workload, primary_axis): | 419 | + def _build_group_features(self, workload, primary_axis, pointwise_layout=None): |
| 346 | if self.kernel.persistent_reduction or self.kernel.inside_reduction: | 420 | if self.kernel.persistent_reduction or self.kernel.inside_reduction: |
| 347 | outer_names = self.non_reduction_axis_names() | 421 | outer_names = self.non_reduction_axis_names() |
| 348 | reduction_names = self.reduction_axis_names() | 422 | reduction_names = self.reduction_axis_names() |
| @@ -376,6 +450,24 @@ class SplitTiling: | |||
| 376 | (num_vector_core * 4096,), | 450 | (num_vector_core * 4096,), |
| 377 | ), | 451 | ), |
| 378 | ) | 452 | ) |
| 453 | + if pointwise_layout in ("broadcast", "transpose_broadcast"): | ||
| 454 | + return ( | ||
| 455 | + GroupFeatureSpec( | ||
| 456 | + "pointwise_broadcast_axis", | ||
| 457 | + "axis", | ||
| 458 | + (primary_axis.name,), | ||
| 459 | + (16, 64, 256, 1024, 4096), | ||
| 460 | + ), | ||
| 461 | + ) | ||
| 462 | + if pointwise_layout == "transpose": | ||
| 463 | + return ( | ||
| 464 | + GroupFeatureSpec( | ||
| 465 | + "pointwise_transpose_axis", | ||
| 466 | + "axis", | ||
| 467 | + (primary_axis.name,), | ||
| 468 | + (64, 128, 256, 512), | ||
| 469 | + ), | ||
| 470 | + ) | ||
| 379 | return ( | 471 | return ( |
| 380 | GroupFeatureSpec( | 472 | GroupFeatureSpec( |
| 381 | "pointwise", | 473 | "pointwise", |
| @@ -538,6 +630,7 @@ class SplitTiling: | |||
| 538 | dynamic_split_axes, static_split_axes = self._classify_split_axes() | 630 | dynamic_split_axes, static_split_axes = self._classify_split_axes() |
| 539 | template = self._grouped_template_name() | 631 | template = self._grouped_template_name() |
| 540 | workload = self._classify_group_workload(template) | 632 | workload = self._classify_group_workload(template) |
| 633 | + pointwise_layout = None | ||
| 541 | if dynamic_split_axes: | 634 | if dynamic_split_axes: |
| 542 | primary_axis = self._select_primary_group_axis(dynamic_split_axes) | 635 | primary_axis = self._select_primary_group_axis(dynamic_split_axes) |
| 543 | if primary_axis is None: | 636 | if primary_axis is None: |
| @@ -550,12 +643,11 @@ class SplitTiling: | |||
| 550 | f"{axis.name.upper()}BLOCK" for axis in self.kernel.split_axis | 643 | f"{axis.name.upper()}BLOCK" for axis in self.kernel.split_axis |
| 551 | ) | 644 | ) |
| 552 | else: | 645 | else: |
| 553 | - # No dynamic grid split axis: bucket the reduction tiling axis by its | 646 | + if template in ("persistent_reduction", "reduction"): |
| 554 | - # runtime size (grid==1 over it). Grid parallelism, if any, comes from | 647 | + primary_axis = self._dynamic_reduction_tiling_axis() |
| 555 | - # the STATIC non-reduction split axes; their blocks are passed so the | 648 | + else: |
| 556 | - # grid is computed from them (build_grouped_launch_policy treats a | 649 | + primary_axis = self._dynamic_pointwise_tiling_axis() |
| 557 | - # reduction-tiling primary specially -- no runtime rule for it). | 650 | + pointwise_layout = self._pointwise_layout_kind() |
| 558 | - primary_axis = self._dynamic_reduction_tiling_axis() | ||
| 559 | if primary_axis is None: | 651 | if primary_axis is None: |
| 560 | return None | 652 | return None |
| 561 | static_names = tuple(axis.name for axis in static_split_axes) | 653 | static_names = tuple(axis.name for axis in static_split_axes) |
| @@ -563,7 +655,9 @@ class SplitTiling: | |||
| 563 | runtime_block_arg_names = tuple( | 655 | runtime_block_arg_names = tuple( |
| 564 | f"{axis.name.upper()}BLOCK" for axis in self.kernel.split_axis | 656 | f"{axis.name.upper()}BLOCK" for axis in self.kernel.split_axis |
| 565 | ) | 657 | ) |
| 566 | - feature_specs = self._build_group_features(workload, primary_axis) | 658 | + feature_specs = self._build_group_features( |
| 659 | + workload, primary_axis, pointwise_layout | ||
| 660 | + ) | ||
| 567 | if not feature_specs: | 661 | if not feature_specs: |
| 568 | return None | 662 | return None |
| 569 | return GroupedKernelMeta( | 663 | return GroupedKernelMeta( |