已合并
[inductor]support pointwise group #42716
[inductor]support pointwise group #42716
已合并
rain-666创建于 7月25日
共 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 paths7 - Indirect memory ops (gather/scatter/index_select) with SIMT/SIMD paths
8 - View/reshape/cat + compute fusion with dynamic shapes8 - 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 elements13 - nonzero produces data-dependent output shape[0] = count of non-zero elements
12 - Unbacked SymInt guarded by runtime assertions under dynamo14 - Unbacked SymInt guarded by runtime assertions under dynamo
13 15 
@@ -17,7 +19,10 @@ Reference patterns from:
17"""19"""
18 20 
19import unittest21import unittest
22+from types import SimpleNamespace
23+from unittest.mock import patch
20 24 
25+import sympy
21import torch26import torch
22import torch.nn.functional as F27import torch.nn.functional as F
23from torch._dynamo.testing import CompileCounterWithBackend28from torch._dynamo.testing import CompileCounterWithBackend
@@ -31,6 +36,10 @@ from torch.testing._internal.common_utils import (
31 36 
32import torch_npu37import torch_npu
33import torch_npu._inductor38import 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 
35torch._dynamo.config.cache_size_limit = 12844torch._dynamo.config.cache_size_limit = 128
36 45 
@@ -40,6 +49,16 @@ if not torch.npu.is_available():
40device = "npu"49device = "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 helpers63# 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 Tests652+# 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 
636class TestNonzeroDynamicShapes(DynamicShapeTestMixin, TestCase):795class TestNonzeroDynamicShapes(DynamicShapeTestMixin, TestCase):
@@ -854,6 +1013,7 @@ class TestNonzeroDynamicShapes(DynamicShapeTestMixin, TestCase):
854 1013 
855instantiate_parametrized_tests(TestFusionDynamicShapes)1014instantiate_parametrized_tests(TestFusionDynamicShapes)
856instantiate_parametrized_tests(TestSymbolicGroupElementwise)1015instantiate_parametrized_tests(TestSymbolicGroupElementwise)
1016+instantiate_parametrized_tests(TestPointwiseSymbolicGrouping)
857instantiate_parametrized_tests(TestNonzeroDynamicShapes)1017instantiate_parametrized_tests(TestNonzeroDynamicShapes)
858 1018 
859 1019 
@@ -315,6 +315,80 @@ class SplitTiling:
315 return None315 return None
316 return axis316 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_axis643 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 its646+ if template in ("persistent_reduction", "reduction"):
554- # runtime size (grid==1 over it). Grid parallelism, if any, comes from647+ primary_axis = self._dynamic_reduction_tiling_axis()
555- # the STATIC non-reduction split axes; their blocks are passed so the648+ else:
556- # grid is computed from them (build_grouped_launch_policy treats a649+ 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 None652 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_axis656 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 None662 return None
569 return GroupedKernelMeta(663 return GroupedKernelMeta(