From c57b7cf008072f6aead9e7946c6373abfd32ca08 Mon Sep 17 00:00:00 2001 From: Yufeng Shi Date: Thu, 24 Sep 2026 10:45:05 +0100 Subject: [PATCH] Arm backend: Add static floating-point TopK lowering Delegate static rank-2 FP16/FP32 TopK with constant K=1..4 along the last dimension. Support largest=True and sorted=True, using FP for K=1 and FP+INT for K>1. Previously, TopK remained outside the Arm delegate. Existing range analysis narrowed only index paths proven safe for int32: scores -> TopK |-- values (FP) `-- indices (int64) |-- model output / gather / int64 arithmetic `-- cast(int32) -> proven-safe consumers Route all supported TopK index paths through int32 before partitioning. Restore int64 separately for consumers that need it, allowing delegated gathers to coexist with portable consumers. Keep TopK, tuple extraction and narrowing together in one delegate. Decompose TopK into repeated ARGMAX with cumulative masking, unrolling K selections at compile time. Mask each selected position with -inf while preserving previous masks. Gather values from the original scores using the selected indices. After decomposition (simplified): scores -> decomposed TopK |-- values (FP) `-- indices (int32) |-- proven-safe int32 consumers `-- cast(int64) per remaining consumer `-- model output / portable gather / int64 arithmetic Require finite scores without checking finiteness at runtime. Select equal scores in increasing index order, which may differ from PyTorch's tie ordering. Change-Id: I2b9a0b9e4e6a11d13a84c473161b0ad330e20b14 Signed-off-by: Yufeng Shi --- backends/arm/README.md | 22 ++ backends/arm/_passes/__init__.py | 1 + backends/arm/_passes/arm_pass_manager.py | 6 +- .../convert_int64_output_ops_to_int32.py | 149 +++++++- backends/arm/_passes/decompose_topk_pass.py | 320 ++++++++++++++++ backends/arm/operator_support/__init__.py | 1 + backends/arm/operator_support/topk_support.py | 40 ++ .../tosa_profile_supported_op_lists.py | 1 + .../tosa_supported_operators.py | 23 ++ .../test/misc/test_tosa_operator_support.py | 73 ++++ backends/arm/test/ops/test_topk.py | 341 ++++++++++++++++++ .../test_convert_int64_output_ops_to_int32.py | 221 ++++++++++++ .../test/passes/test_decompose_topk_pass.py | 63 ++++ backends/arm/test/targets.bzl | 2 + backends/arm/tosa/partitioner.py | 32 ++ .../source/backends/arm-vgf/VGF_op_support.md | 3 +- 16 files changed, 1278 insertions(+), 20 deletions(-) create mode 100644 backends/arm/_passes/decompose_topk_pass.py create mode 100644 backends/arm/operator_support/topk_support.py create mode 100644 backends/arm/test/ops/test_topk.py create mode 100644 backends/arm/test/passes/test_decompose_topk_pass.py diff --git a/backends/arm/README.md b/backends/arm/README.md index 19221145824..bbb8a8bb752 100644 --- a/backends/arm/README.md +++ b/backends/arm/README.md @@ -199,6 +199,24 @@ compilation. Reach for the step-by-step flow above when a recipe does not fit -- a custom quantization scheme, extra passes, or a compile spec the recipe does not expose. +#### TopK support + +The `to_edge_transform_and_lower` flow handles supported `torch.topk` calls +automatically. Supported configurations are: + +- FP16 or FP32 input with positive static shape `[T, E]`, where + `E <= 2^31 - 1`. +- Constant `1 <= K <= min(4, E)`, `dim=-1` or `dim=1`, `largest=True`, and + `sorted=True`. +- The TOSA FP profile for K=1, or FP+INT for K>1. The scores remain + floating-point in both cases. + +**Input scores must be finite.** This condition is not checked at runtime; +non-finite scores can produce incorrect results without triggering automatic +fallback. Equal scores are selected in increasing index order, so tied indices +may differ from PyTorch's results. Dynamic shapes or K, BF16, and quantized TopK +are unsupported. + ### Direct Drive (experimental, Ethos-U85 on Linux) workflow Direct Drive enables execution on Ethos-U85 via the Linux driver stack. @@ -419,6 +437,10 @@ List of model specific and optional passes: - Inserts int64 boundary casts where converted paths reach unsafe consumers or model outputs. - Keeps gather indices int64 so an undelegated gather remains valid. + - Prepares supported static TopK for delegation with int32 indices, + preserving int64 model outputs and consumers that require int64. + - For TopK configurations that cannot be delegated, downstream index + operations can still use int32 where range analysis proves it safe. - Supported Ops: - torch.ops.aten.topk.default - exir_ops.edge.aten.topk.default diff --git a/backends/arm/_passes/__init__.py b/backends/arm/_passes/__init__.py index 1801c320526..075677097bf 100644 --- a/backends/arm/_passes/__init__.py +++ b/backends/arm/_passes/__init__.py @@ -118,6 +118,7 @@ from .decompose_strided_slice_copy_pass import DecomposeStridedSliceCopyPass # noqa from .decompose_sum_pass import DecomposeSumPass # noqa from .decompose_tan_pass import DecomposeTanPass # noqa +from .decompose_topk_pass import DecomposeTopKPass # noqa from .decompose_tosa_unsupported_clamp_pass import ( # noqa DecomposeTOSAUnsupportedClampPass, ) diff --git a/backends/arm/_passes/arm_pass_manager.py b/backends/arm/_passes/arm_pass_manager.py index f82212cd5a2..e1585290170 100644 --- a/backends/arm/_passes/arm_pass_manager.py +++ b/backends/arm/_passes/arm_pass_manager.py @@ -107,6 +107,7 @@ DecomposeStridedSliceCopyPass, DecomposeSumPass, DecomposeTanPass, + DecomposeTopKPass, DecomposeTOSAUnsupportedClampPass, DecomposeTrilPass, DecomposeUnfoldToGatherPass, @@ -480,7 +481,9 @@ def transform_for_pre_decomposition_pipeline( if config.sdpa_safe_softmax_guard is SDPASafeSoftmaxGuardPolicy.AUTO: passes.append(DecomposeSDPAWithRegularSoftmaxPass()) - convert_pass = ConvertInt64OutputOpsToInt32Pass(convert_cast_ops=False) + convert_pass = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, tosa_spec=self.tosa_spec + ) if convert_pass.should_run(exported_program.graph_module): passes.append(convert_pass) @@ -540,6 +543,7 @@ def _tosa_pipeline( NormalizeDelegateIOLayoutPass(exported_program), FuseQuantizedActivationPass(), RewriteBoolToFp32CastViaInt8Pass(), + DecomposeTopKPass(self.tosa_spec), PrepareGatherIndicesPass(self.tosa_spec), CanonicalizeGatherPass(), ConvertToClampPass(), diff --git a/backends/arm/_passes/convert_int64_output_ops_to_int32.py b/backends/arm/_passes/convert_int64_output_ops_to_int32.py index 7e7a3bd3bea..3c14651ebfd 100644 --- a/backends/arm/_passes/convert_int64_output_ops_to_int32.py +++ b/backends/arm/_passes/convert_int64_output_ops_to_int32.py @@ -15,10 +15,16 @@ get_first_fake_tensor, set_node_arg, ) +from executorch.backends.arm._passes.decompose_topk_pass import ( + get_static_topk_config, + is_topk_indices_int32_cast, + topk_indices_only_feed_int32_casts, +) from executorch.backends.arm._passes.int32_range_analysis import ( Int32RangeAnalysis, ValueRange, ) +from executorch.backends.arm.tosa.specification import TosaSpecification from executorch.exir.dialects._ops import ops as exir_ops from executorch.exir.pass_base import ExportPass, PassResult @@ -43,6 +49,10 @@ class ConvertInt64OutputOpsToInt32Pass(ArmPass): apply the same bounded-index handling to ``getitem(topk, 1)`` while leaving the values output, ``getitem(topk, 0)``, unchanged. + With a TOSA specification, also prepare supported static TopK indices + for delegation. All immediate index users narrow to int32, with separate + int64 restoration casts for remaining consumers and graph outputs. + Argmax, argmin and extracted TopK indices are the bounded-index sources. Range propagation from those sources recognizes a separate allowlist of safe shape and arithmetic operations. @@ -60,6 +70,10 @@ class ConvertInt64OutputOpsToInt32Pass(ArmPass): ``"raise"`` (default) raises a ``RuntimeError`` at compile time. ``"warn"`` logs a warning and skips the cast for that node. ``"skip"`` silently skips the cast for that node. + tosa_spec (TosaSpecification | None): Enable target-aware TopK + boundary preparation when provided. Requires + ``convert_cast_ops=False`` to preserve int64 restoration casts. + Defaults to None, retaining generic range conversion only. """ @@ -72,6 +86,7 @@ def __init__( *args, convert_cast_ops: bool = True, on_overflow: Literal["raise", "warn", "skip"] = "raise", + tosa_spec: TosaSpecification | None = None, **kwargs, ) -> None: super().__init__(*args, **kwargs) @@ -81,6 +96,11 @@ def __init__( ) self.convert_cast_ops = convert_cast_ops self.on_overflow = on_overflow + if tosa_spec is not None and convert_cast_ops: + raise ValueError( + "TopK boundary preparation requires convert_cast_ops=False." + ) + self.tosa_spec = tosa_spec aten_cast_ops = Int32RangeAnalysis.aten_cast_ops edge_cast_ops = Int32RangeAnalysis.edge_cast_ops @@ -149,6 +169,40 @@ def _insert_int64_boundary( ) return boundaries[node] + def _get_or_create_int32_cast( + self, + graph: torch.fx.Graph, + source: torch.fx.Node, + to_copy_op, + int32_source: torch.fx.Node | None = None, + ) -> torch.fx.Node: + """Create or reuse an int32 cast immediately after its source. + + Args: + graph (torch.fx.Graph): Graph containing the source and cast. + source (torch.fx.Node): Bounded int64 index source. + to_copy_op (Any): Dialect-specific operator used for casts. + int32_source (torch.fx.Node | None): Existing narrowing cast to + reuse, or None to create one. + + Returns: + torch.fx.Node: Int32 cast positioned immediately after source. + + """ + if int32_source is None: + with graph.inserting_after(source): + int32_source = create_node( + graph, + to_copy_op, + args=(source,), + kwargs={"dtype": torch.int32}, + ) + self._safe_index_casts += 1 + else: + # Move the reused cast before any consumers redirected to it. + source.append(int32_source) + return int32_source + def _cast_safe_scalar_constants_to_int32( self, graph_module: torch.fx.GraphModule, @@ -192,6 +246,7 @@ def _cast_safe_index_paths_to_int32( source: torch.fx.Node, source_range: ValueRange, to_copy_op, + int32_source: torch.fx.Node | None = None, ) -> bool: """Convert proven-safe paths from a bounded index source to int32. @@ -206,6 +261,8 @@ def _cast_safe_index_paths_to_int32( source (torch.fx.Node): Int64 node with a statically known range. source_range (ValueRange): Inclusive minimum and maximum. to_copy_op (Any): Dialect-specific operator used for casts. + int32_source (torch.fx.Node | None): Existing narrowing cast to + reuse for safe index paths. Returns: bool: True when at least one path is converted to int32. @@ -219,13 +276,9 @@ def _cast_safe_index_paths_to_int32( graph = graph_module.graph original_users = {node: list(node.users) for node in ranges} - with graph.inserting_after(source): - cast_to_int32 = create_node( - graph, - to_copy_op, - args=(source,), - kwargs={"dtype": torch.int32}, - ) + int32_source = self._get_or_create_int32_cast( + graph, source, to_copy_op, int32_source + ) self._cast_safe_scalar_constants_to_int32( graph_module, analysis, safe_consumers, ranges, to_copy_op @@ -235,15 +288,51 @@ def _cast_safe_index_paths_to_int32( for node, users in original_users.items(): for user in users: if user in safe_consumers: - if node is source: - user.replace_input_with(source, cast_to_int32) + if node is source and user is not int32_source: + user.replace_input_with(source, int32_source) elif node is not source: boundary = self._insert_int64_boundary( graph, node, to_copy_op, boundaries ) user.replace_input_with(node, boundary) - self._safe_index_casts += 1 + return True + + def _prepare_topk_index_boundary( + self, + graph: torch.fx.Graph, + indices: torch.fx.Node, + to_copy_op, + ) -> bool: + """Restore int64 separately for remaining users of bounded indices. + + Args: + graph (torch.fx.Graph): Graph containing the index extraction. + indices (torch.fx.Node): Supported TopK's int64 index output. + to_copy_op (Any): Dialect-specific operator used for casts. + + Returns: + bool: True when remaining consumers received restoration casts. + + """ + remaining_users = [ + user for user in indices.users if not is_topk_indices_int32_cast(user) + ] + if not remaining_users: + return False + int32_indices = next( + (user for user in indices.users if is_topk_indices_int32_cast(user)), None + ) + int32_indices = self._get_or_create_int32_cast( + graph, indices, to_copy_op, int32_indices + ) + # Separate casts let gathers delegate independently of other consumers. + for user in remaining_users: + with graph.inserting_before(user): + int64_indices = create_node( + graph, to_copy_op, (int32_indices,), {"dtype": torch.int64} + ) + user.replace_input_with(indices, int64_indices) return True def _log_summary(self) -> None: @@ -263,7 +352,7 @@ def _convert_topk_indices( analysis: Int32RangeAnalysis, topk: torch.fx.Node, ) -> bool: - """Convert safe paths from extracted TopK indices. + """Convert safe TopK paths and optionally prepare delegate boundaries. Args: graph_module (torch.fx.GraphModule): Graph containing TopK. @@ -293,16 +382,39 @@ def _convert_topk_indices( logger.warning(msg) return False + # All extracted indices already feed int32 casts, so no further + # index conversion or boundary preparation is needed. + if topk_indices_only_feed_int32_casts(topk): + return False + + supports_topk_lowering = ( + self.tosa_spec is not None + and get_static_topk_config(topk, self.tosa_spec)[0] is not None + ) modified = False to_copy_op = self._get_decomposition(topk.target) + # Convert safe index paths to int32 regardless of TopK lowering support. + # For supported TopK, ensure every direct index user is an int32 cast, + # restoring int64 separately for consumers that still need it. for indices in index_getitems: - if get_first_fake_tensor(indices).dtype == torch.int64: - modified |= self._cast_safe_index_paths_to_int32( - graph_module, - analysis, - indices, - index_range, - to_copy_op, + if get_first_fake_tensor(indices).dtype != torch.int64: + continue + + existing_int32_cast = next( + (user for user in indices.users if is_topk_indices_int32_cast(user)), + None, + ) + modified |= self._cast_safe_index_paths_to_int32( + graph_module, + analysis, + indices, + index_range, + to_copy_op, + existing_int32_cast, + ) + if supports_topk_lowering: + modified |= self._prepare_topk_index_boundary( + graph_module.graph, indices, to_copy_op ) return modified @@ -391,6 +503,7 @@ def call(self, graph_module: torch.fx.GraphModule): if modified: graph_module.graph.eliminate_dead_code() + graph_module.graph.lint() graph_module.recompile() graph_module = super().call(graph_module).graph_module diff --git a/backends/arm/_passes/decompose_topk_pass.py b/backends/arm/_passes/decompose_topk_pass.py new file mode 100644 index 00000000000..421e52ac0eb --- /dev/null +++ b/backends/arm/_passes/decompose_topk_pass.py @@ -0,0 +1,320 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. +"""Share the static TopK contract and lower it using ARGMAX and masking.""" + +import operator +from dataclasses import dataclass +from typing import cast + +import torch +from executorch.backends.arm._passes.arm_pass import ArmPass +from executorch.backends.arm._passes.arm_pass_utils import ( + create_node, + get_first_fake_tensor, +) +from executorch.backends.arm._passes.canonicalize_gather_pass import ( + CanonicalizeGatherPass, +) +from executorch.backends.arm._passes.prepare_gather_indices_pass import ( + PrepareGatherIndicesPass, +) +from executorch.backends.arm.tosa.specification import TosaSpecification +from executorch.exir.dialects._ops import ops as exir_ops +from executorch.exir.pass_base import ExportPass, PassResult + +TOPK_OPS = (torch.ops.aten.topk.default, exir_ops.edge.aten.topk.default) +_CAST_OPS = ( + torch.ops.dim_order_ops._to_dim_order_copy.default, + exir_ops.edge.dim_order_ops._to_dim_order_copy.default, +) + + +@dataclass(frozen=True) +class StaticTopKConfig: + """Describe a supported last-dimension TopK on a static matrix.""" + + tokens: int + experts: int + k: int + + +def _get_score_shape( + scores: torch.fx.Node, +) -> tuple[tuple[int, int] | None, str | None]: + value = get_first_fake_tensor(scores) + if value.ndim != 2: + return None, "TopK requires rank-2 input." + if value.dtype not in (torch.float16, torch.float32): + return None, "TopK requires FP16 or FP32 scores." + if any(type(size) is not int or size <= 0 for size in value.shape): + return None, "TopK requires positive static T and E." + tokens, experts = value.shape + if experts > torch.iinfo(torch.int32).max: + return None, "TopK expert dimension exceeds the int32 index-size limit." + return (tokens, experts), None + + +def get_static_topk_config( + node: torch.fx.Node, tosa_spec: TosaSpecification +) -> tuple[StaticTopKConfig | None, str | None]: + """Check metadata and capabilities, without checking runtime finiteness. + + Args: + node (torch.fx.Node): ATen or Edge TopK node with valid arguments. + tosa_spec (TosaSpecification): Target capabilities. + + Returns: + tuple: ``(config, None)`` for supported TopK, or ``(None, reason)`` + explaining why TopK is unsupported. + + """ + names = ("self", "k", "dim", "largest", "sorted") + arguments = dict(zip(names, node.args)) + arguments.update(node.kwargs) + shape, reason = _get_score_shape(cast(torch.fx.Node, arguments["self"])) + if shape is None: + return None, reason + tokens, experts = shape + k = arguments["k"] + if type(k) is not int or not 1 <= k <= min(4, experts): + return None, "TopK requires constant 1 <= K <= min(4, E)." + dim = arguments.get("dim", -1) + if type(dim) is not int or dim not in (-1, 1): + return None, "TopK requires dim=-1 or dim=1." + if arguments.get("largest", True) is not True: + return None, "TopK requires largest=True." + if arguments.get("sorted", True) is not True: + return None, "TopK requires sorted=True." + if not tosa_spec.support_float(): + return None, "TopK requires the FP profile." + if k > 1 and not tosa_spec.support_integer(): + return None, "TopK with K>1 requires the FP and INT profiles." + if any( + user.target is not operator.getitem + or len(user.args) != 2 + or type(user.args[1]) is not int + or user.args[1] not in (0, 1) + for user in node.users + ): + return None, "TopK requires canonical values/indices tuple extraction." + return StaticTopKConfig(tokens, experts, k), None + + +def is_topk_indices_getitem(node: torch.fx.Node) -> bool: + """Return whether node extracts the indices output of ATen or Edge TopK. + + Args: + node (torch.fx.Node): Candidate tuple extraction node. + + Returns: + bool: True for ``operator.getitem(topk, 1)``, selecting the indices + from TopK's ``(values, indices)`` result. + + """ + return ( + node.target is operator.getitem + and len(node.args) == 2 + and type(node.args[1]) is int + and node.args[1] == 1 + and isinstance(node.args[0], torch.fx.Node) + and node.args[0].target in TOPK_OPS + ) + + +def is_topk_indices_int32_cast(node: torch.fx.Node) -> bool: + """Return whether node casts extracted TopK indices to int32. + + The checked node is the final cast in this pattern:: + + topk(scores, K) -> getitem(1) -> _to_dim_order_copy(dtype=int32) + + Args: + node (torch.fx.Node): Candidate cast node. + + Returns: + bool: True for an ATen or Edge ``_to_dim_order_copy`` to int32 + whose input is a TopK indices extraction. + + """ + return ( + node.target in _CAST_OPS + and len(node.args) == 1 + and node.kwargs.get("dtype") is torch.int32 + and isinstance(node.args[0], torch.fx.Node) + and is_topk_indices_getitem(node.args[0]) + ) + + +def topk_indices_only_feed_int32_casts(node: torch.fx.Node) -> bool: + """Return whether extracted TopK indices feed only int32 casts. + + Expected pattern:: + + topk -> getitem(1) -> _to_dim_order_copy(dtype=int32) + + Values consumers are ignored. If the indices output is unused, no + int32 cast is required. + + Args: + node (torch.fx.Node): TopK producer with canonical ``getitem`` users. + + Returns: + bool: True when every immediate indices consumer matches + ``is_topk_indices_int32_cast``, or the indices are unused. + + """ + for output in node.users: + if not is_topk_indices_getitem(output): + continue + for consumer in output.users: + if not is_topk_indices_int32_cast(consumer): + return False + return True + + +class DecomposeTopKPass(ArmPass): + """Lower finite-score TopK using repeated ARGMAX and index masking. + + Runs after Arm partitioning. Every used TopK index extraction must feed + only explicit int32 casts, matching this pattern:: + + topk(scores, K) + |-- getitem(0) --> value consumers + `-- getitem(1) --> _to_dim_order_copy(dtype=int32) --> consumers + + Either output may be unused. Before partitioning, + ``ConvertInt64OutputOpsToInt32Pass`` prepares restoration casts after + the int32 casts for consumers requiring int64. + + For scores of shape ``[T, E]``, the decomposition is:: + + remaining = scores + selected_indices = [] + for step in range(K): + index = argmax(remaining, dim=1).reshape(T, 1) + selected_indices.append(index) + if step + 1 < K: + remaining = where(arange(E) == index, -inf, remaining) + indices = concat(selected_indices, dim=1) # int32, shape [T, K] + values = gather(scores, dim=1, index=indices) + + Args: + tosa_spec (TosaSpecification): Target capabilities. + + """ + + _passes_required_after: set[type[ExportPass]] = { + PrepareGatherIndicesPass, + CanonicalizeGatherPass, + } + + def __init__(self, tosa_spec: TosaSpecification) -> None: + super().__init__() + self.tosa_spec = tosa_spec + + @staticmethod + def _build_topk_indices( + graph: torch.fx.Graph, + scores: torch.fx.Node, + config: StaticTopKConfig, + ) -> torch.fx.Node: + """Build int32 indices using repeated ARGMAX and cumulative masking. + + Args: + graph (torch.fx.Graph): Graph with the insertion point set. + scores (torch.fx.Node): Original score tensor of shape [T, E]. + config (StaticTopKConfig): Validated static TopK parameters. + + Returns: + torch.fx.Node: Int32 indices with shape [T, K]. + + """ + expert_ids: torch.fx.Node | None = None + minus_inf: torch.fx.Node | None = None + if config.k > 1: + scores_tensor = scores.meta["val"] + expert_ids = create_node( + graph, + exir_ops.edge.aten.arange.start_step, + (0, config.experts, 1), + {"dtype": torch.int32, "device": scores_tensor.device}, + ) + expert_ids = create_node( + graph, + exir_ops.edge.aten.view_copy.default, + (expert_ids, [1, config.experts]), + ) + minus_inf = create_node( + graph, + exir_ops.edge.aten.full.default, + ([1, 1], float("-inf")), + {"dtype": scores_tensor.dtype, "device": scores_tensor.device}, + ) + remaining = scores + selected_indices: list[torch.fx.Node] = [] + for step in range(config.k): + # ARGMAX reduces [T, E] to int32 indices [T]. + index = create_node( + graph, exir_ops.backend.tosa.ARGMAX.default, (remaining, 1) + ) + index = create_node( + graph, + exir_ops.edge.aten.view_copy.default, + (index, [config.tokens, 1]), + ) + selected_indices.append(index) + if step + 1 < config.k: + assert expert_ids is not None and minus_inf is not None + # Compare expert IDs [1, E] with selected indices [T, 1] + # to produce a mask [T, E]. + mask = create_node( + graph, exir_ops.edge.aten.eq.Tensor, (expert_ids, index) + ) + # Mask selected scores in remaining [T, E] with -inf. + remaining = create_node( + graph, + exir_ops.edge.aten.where.self, + (mask, minus_inf, remaining), + ) + if config.k == 1: + indices = selected_indices[0] + else: + # Concatenate K selections of shape [T, 1] into [T, K]. + indices = create_node( + graph, exir_ops.edge.aten.cat.default, (selected_indices, 1) + ) + return indices + + def call(self, graph_module: torch.fx.GraphModule) -> PassResult: + """Unroll selection and replace the TopK tuple's live extractions.""" + graph = graph_module.graph + modified = False + for topk in list(graph.nodes): + if topk.target not in TOPK_OPS: + continue + config, reason = get_static_topk_config(topk, self.tosa_spec) + if config is None: + raise RuntimeError(f"Unsupported delegated TopK: {reason}") + if not topk_indices_only_feed_int32_casts(topk): + raise RuntimeError("Delegated TopK requires prepared int32 indices.") + scores = topk.args[0] if topk.args else topk.kwargs["self"] + with graph.inserting_before(topk): + indices = self._build_topk_indices(graph, scores, config) + # Gather from scores [T, E] using indices [T, K] + # to produce values [T, K]. + values = create_node( + graph, exir_ops.edge.aten.gather.default, (scores, 1, indices) + ) + # Redirect TopK getitem consumers to the new values and indices. + for extraction in list(topk.users): + replacement = values if extraction.args[1] == 0 else indices + extraction.replace_all_uses_with(replacement) + modified = True + if modified: + graph.eliminate_dead_code() + graph.lint() + graph_module.recompile() + graph_module = super().call(graph_module).graph_module + return PassResult(graph_module, modified) diff --git a/backends/arm/operator_support/__init__.py b/backends/arm/operator_support/__init__.py index 2784d0db346..1472b327cc8 100644 --- a/backends/arm/operator_support/__init__.py +++ b/backends/arm/operator_support/__init__.py @@ -24,6 +24,7 @@ slice_copy_support, symint_arithmetic_support, to_dim_order_copy_support, + topk_support, tosa_supported_operators, unfold_copy_support, upsample_support, diff --git a/backends/arm/operator_support/topk_support.py b/backends/arm/operator_support/topk_support.py new file mode 100644 index 00000000000..03da595842b --- /dev/null +++ b/backends/arm/operator_support/topk_support.py @@ -0,0 +1,40 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. +"""Gate finite-score static TopK lowering and its index interfaces.""" + +import torch.fx +from executorch.backends.arm._passes.decompose_topk_pass import ( + get_static_topk_config, + topk_indices_only_feed_int32_casts, + TOPK_OPS, +) +from executorch.backends.arm.operator_support.tosa_supported_operators import ( + register_tosa_support_check, + SupportedTOSAOperatorCheck, +) +from executorch.backends.arm.tosa.specification import TosaSpecification + + +@register_tosa_support_check +class TopKSupported(SupportedTOSAOperatorCheck): + """Accept the static decomposition under its finite-score precondition.""" + + tosa_specs = TosaSpecification.all_versions_for_profile("FP") + targets = list(TOPK_OPS) + + def is_node_tosa_supported( + self, node: torch.fx.Node, tosa_spec: TosaSpecification + ) -> bool: + """Check the supported metadata and prepared index boundary.""" + config, reason = get_static_topk_config(node, tosa_spec) + if config is None: + self.reporter.report_reject(node, f"Unsupported TopK: {reason}") + return False + if not topk_indices_only_feed_int32_casts(node): + self.reporter.report_reject( + node, "TopK index users require an int32 narrowing boundary." + ) + return False + return True diff --git a/backends/arm/operator_support/tosa_profile_supported_op_lists.py b/backends/arm/operator_support/tosa_profile_supported_op_lists.py index 6c5ed465de1..cd4384dd708 100644 --- a/backends/arm/operator_support/tosa_profile_supported_op_lists.py +++ b/backends/arm/operator_support/tosa_profile_supported_op_lists.py @@ -270,6 +270,7 @@ exir_ops.edge.aten.pad.default, exir_ops.edge.aten.constant_pad_nd.default, exir_ops.edge.aten.argmax.default, + exir_ops.edge.aten.topk.default, exir_ops.edge.aten.amax.default, exir_ops.edge.aten.amin.default, exir_ops.edge.aten.eye.default, diff --git a/backends/arm/operator_support/tosa_supported_operators.py b/backends/arm/operator_support/tosa_supported_operators.py index 2b3c860b6d9..de845d9f25c 100644 --- a/backends/arm/operator_support/tosa_supported_operators.py +++ b/backends/arm/operator_support/tosa_supported_operators.py @@ -24,6 +24,13 @@ get_first_fake_tensor, is_submodule_node, ) +from executorch.backends.arm._passes.decompose_topk_pass import ( + get_static_topk_config, + is_topk_indices_getitem, + is_topk_indices_int32_cast, + topk_indices_only_feed_int32_casts, + TOPK_OPS, +) from executorch.backends.arm._passes.fuse_constant_ops_pass import ( ComputeConstantOpsAOTPass, ) @@ -1230,6 +1237,8 @@ def has_rejected_int64_output( return False if node.target in _ARGMAX_OPS: return not self._is_tosa_argmax_supported(node) + if self._is_prepared_topk_index(node): + return False return any( tensor.dtype == torch.int64 @@ -1237,6 +1246,16 @@ def has_rejected_int64_output( if isinstance(tensor, FakeTensor) ) + def _is_prepared_topk_index(self, node: torch.fx.Node) -> bool: + if node.target in TOPK_OPS: + source = node + elif is_topk_indices_getitem(node): + source = typing.cast(torch.fx.Node, node.args[0]) + else: + return False + config, _ = get_static_topk_config(source, self.tosa_spec) + return config is not None and topk_indices_only_feed_int32_casts(source) + def _is_argmax_int32_cast( self, node: torch.fx.Node, @@ -1366,6 +1385,10 @@ def _check_int64_input_nodes(self, node: torch.fx.Node) -> bool: # can be placed in the same delegate. if self._is_argmax_int32_cast(node, input_node): continue + if is_topk_indices_int32_cast(node) and self._is_prepared_topk_index( + input_node + ): + continue # Constant placeholder if ( diff --git a/backends/arm/test/misc/test_tosa_operator_support.py b/backends/arm/test/misc/test_tosa_operator_support.py index 91de62f89ae..74567a9fb21 100644 --- a/backends/arm/test/misc/test_tosa_operator_support.py +++ b/backends/arm/test/misc/test_tosa_operator_support.py @@ -3,11 +3,15 @@ # This source code is licensed under the BSD-style license found in the # LICENSE file in the root directory of this source tree. +import operator + import pytest import torch +from executorch.backends.arm._passes.decompose_topk_pass import get_static_topk_config from executorch.backends.arm.operator_support.index_tensor_support import ( IndexTensorSupported, ) +from executorch.backends.arm.operator_support.topk_support import TopKSupported from executorch.backends.arm.operator_support.tosa_supported_operators import ( CheckFPComparisonInputs, CheckKnownUnsupportedTOSASemantics, @@ -38,6 +42,14 @@ def _fp_comparison_checker() -> CheckFPComparisonInputs: return CheckFPComparisonInputs(WhyNoPartitionReporter()) +def _topk_node(shape=(2, 8), dtype=torch.float32, args=(2,), kwargs=None): + graph = torch.fx.Graph() + scores = _placeholder(graph, "scores", shape, dtype) + return graph.call_function( + exir_ops.edge.aten.topk.default, (scores, *args), kwargs or {} + ) + + @pytest.mark.parametrize( "target", ( @@ -183,3 +195,64 @@ def test_rejects_index_tensor_data_dependent_shapes( assert not checker.is_node_supported({}, node) assert "Symbolic value or index shapes" in reporter.get_table_report() assert not shape_env.guards + + +@pytest.mark.parametrize( + "shape,dtype,args,kwargs,reason", + [ + ((8,), torch.float32, (1,), {}, "rank-2"), + ((0, 8), torch.float32, (1,), {}, "positive static"), + ((2, 0), torch.float32, (1,), {}, "positive static"), + ((1, 2147483648), torch.float32, (1,), {}, "int32"), + ((2, 8), torch.float32, (0,), {}, "constant"), + ((2, 2), torch.float32, (3,), {}, "constant"), + ((2, 8), torch.float32, (True,), {}, "constant"), + ((2, 8), torch.float32, (2, 0), {}, "dim"), + ((2, 8), torch.float32, (2, 3), {}, "dim"), + ((2, 8), torch.float32, (2, True), {}, "dim"), + ((2, 8), torch.float32, (2, -1, False), {}, "largest"), + ((2, 8), torch.float32, (2,), {"sorted": False}, "sorted"), + ((2, 8), torch.bfloat16, (1,), {}, "FP16 or FP32"), + ((2, 8), torch.float64, (1,), {}, "FP16 or FP32"), + ((2, 8), torch.int32, (1,), {}, "FP16 or FP32"), + ], +) +def test_topk_unsupported_metadata(shape, dtype, args, kwargs, reason) -> None: + node = _topk_node(shape, dtype, args, kwargs) + spec = TosaSpecification.create_from_string("TOSA-1.0+FP+INT+bf16") + reporter = WhyNoPartitionReporter() + assert not TopKSupported(spec, reporter).is_node_supported({}, node) + assert reason in reporter.get_table_report() + + +def test_topk_rejects_runtime_k() -> None: + node = _topk_node() + with node.graph.inserting_before(node): + k = node.graph.placeholder("k") + node.args = (node.args[0], k) + config, reason = get_static_topk_config( + node, TosaSpecification.create_from_string("TOSA-1.0+FP+INT") + ) + assert config is None + assert reason is not None and "constant" in reason + + +@pytest.mark.parametrize("index", [-2, -1], ids=["values", "indices"]) +def test_topk_rejects_negative_tuple_indices(index) -> None: + topk = _topk_node() + extraction = topk.graph.call_function(operator.getitem, (topk, index)) + topk.graph.output(extraction) + spec = TosaSpecification.create_from_string("TOSA-1.0+FP+INT") + reporter = WhyNoPartitionReporter() + assert not TopKSupported(spec, reporter).is_node_supported({}, topk) + assert "canonical" in reporter.get_table_report() + + +@pytest.mark.parametrize("dim", [-1, 1]) +def test_topk_keyword_arguments(dim) -> None: + node = _topk_node(args=(), kwargs={"k": 4, "dim": dim}) + config, reason = get_static_topk_config( + node, TosaSpecification.create_from_string("TOSA-1.0+FP+INT") + ) + assert config is not None and config.k == 4 + assert reason is None diff --git a/backends/arm/test/ops/test_topk.py b/backends/arm/test/ops/test_topk.py new file mode 100644 index 00000000000..7d773ba82b8 --- /dev/null +++ b/backends/arm/test/ops/test_topk.py @@ -0,0 +1,341 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import operator + +import torch +from executorch.backends.arm._passes import ConvertInt64OutputOpsToInt32Pass +from executorch.backends.arm.test import common +from executorch.backends.arm.test.runner_utils import TosaReferenceModelDispatch +from executorch.backends.arm.test.tester.test_pipeline import ( + TosaPipelineFP, + VgfPipeline, +) +from executorch.backends.arm.tosa.compile_spec import TosaCompileSpec +from executorch.backends.arm.tosa.partitioner import TOSAPartitioner +from executorch.backends.test.harness.stages import StageType +from executorch.exir import EdgeCompileConfig, to_edge, to_edge_transform_and_lower +from executorch.exir.backend.operator_support import DontPartition +from executorch.exir.dialects._ops import ops as exir_ops +from executorch.exir.memory import alloc +from torch.fx.passes.operator_support import OperatorSupportBase + +aten_op = "torch.ops.aten.topk.default" +exir_op = "executorch_exir_dialects_edge__ops_aten_topk_default" +_CAST = exir_ops.edge.dim_order_ops._to_dim_order_copy.default +_CAST_OUT = torch.ops.dim_order_ops._to_dim_order_copy.out +_DELEGATE = torch.ops.higher_order.executorch_call_delegate +input_t = tuple[torch.Tensor] + + +class TopK(torch.nn.Module): + def __init__(self, k, dim=-1, output="both", largest=True, sorted=True): + super().__init__() + self.k = k + self.dim = dim + self.output = output + self.largest = largest + self.sorted = sorted + + def forward(self, scores): + values, indices = torch.topk( + scores, self.k, self.dim, self.largest, self.sorted + ) + if self.output == "values": + return values + if self.output == "indices": + return indices + if self.output == "int32": + return values, indices.int() + return values, indices + + +def _residual_ops(manager): + return [ + node + for node in manager.exported_program().graph.nodes + if node.op == "call_function" + and node.target not in (_DELEGATE, operator.getitem, alloc) + ] + + +def _lower(module, inputs, spec="TOSA-1.0+FP+INT", checks=None, dynamic_shapes=None): + return to_edge_transform_and_lower( + torch.export.export(module, inputs, dynamic_shapes=dynamic_shapes), + partitioner=[TOSAPartitioner(TosaCompileSpec(spec), additional_checks=checks)], + compile_config=EdgeCompileConfig(_check_ir_validity=False), + ) + + +def _execute(manager, inputs): + with TosaReferenceModelDispatch(): + return manager.exported_program().module()(*inputs) + + +test_data = { + f"{dtype}_k{k}_{shape}": (dtype, k, shape) + for dtype in (torch.float16, torch.float32) + for k in (1, 2, 3, 4) + for shape in ("equal", "larger") +} + + +@common.parametrize("test_data", test_data) +def test_topk_tosa_FP(test_data): + dtype, k, shape = test_data + # E: Number of experts; K: Number of activated experts. + # "equal" means E == K; "larger" means E > K. + experts = k if shape == "equal" else 8 + scores = torch.arange(experts, dtype=dtype).unsqueeze(0) + scores = torch.cat((scores.roll(2, 1), -scores.roll(3, 1)), dim=0) + module = TopK(k, dim=1 if shape == "equal" else -1) + pipeline = TosaPipelineFP[input_t]( + module, + (scores,), + aten_op, + exir_op, + tosa_extensions=[] if k == 1 else ["INT"], + atol=0, + rtol=0, + ) + pipeline.count_tosa_ops({"ARGMAX": k, "SELECT": k - 1, "GATHER": 1}) + pipeline.run() + + +@common.parametrize("dtype", {"fp16": torch.float16, "fp32": torch.float32}) +@common.parametrize("k", {"k1": 1, "k4": 4}) +@common.SkipIfNoModelConverter +def test_topk_vgf_no_quant(dtype, k): + scores = torch.tensor([[3.0, -1.0, 4.0, 0.0, 2.0, -2.0, 1.0, -3.0]], dtype=dtype) + pipeline = VgfPipeline[input_t]( + TopK(k), + (scores,), + aten_op, + exir_op, + quantize=False, + tosa_spec="TOSA-1.0+FP+INT", + atol=0, + rtol=0, + ) + pipeline.run() + + +@common.parametrize("k", {"k1": 1, "k4": 4}) +@common.parametrize("output", {key: key for key in ("values", "indices", "int32")}) +def test_topk_output_interfaces_tosa_FP(k, output): + scores = torch.tensor([[2.0, -1.0, 4.0, 0.0]], dtype=torch.float32) + model = TopK(k, output=output) + pipeline = TosaPipelineFP[input_t]( + model, + (scores,), + aten_op, + exir_op, + tosa_extensions=[] if k == 1 else ["INT"], + atol=0, + rtol=0, + ) + pipeline.run() + manager = pipeline.tester.get_artifact(StageType.TO_EDGE_TRANSFORM_AND_LOWER) + residual = _residual_ops(manager) + if output in ("values", "int32"): + assert not residual + else: + assert len(residual) == 1 + assert residual[0].target in (_CAST, _CAST_OUT) + assert residual[0].meta["val"].dtype == torch.int64 + + +@common.parametrize("dtype", {"fp16": torch.float16, "fp32": torch.float32}) +def test_topk_ties_and_extremes_tosa_FP(dtype): + k = 4 + minimum, maximum = torch.finfo(dtype).min, torch.finfo(dtype).max + scores = torch.tensor( + [ + [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], + [minimum] * 8, + [maximum] * 8, + [3.0, 2.0, 2.0, 2.0, 2.0, 1.0, -1.0, -2.0], + [-4.0, -4.0, -2.0, -3.0, -2.0, -5.0, -1.0, -1.0], + [minimum, maximum, 0.0, -0.0, 1.0, -1.0, minimum, maximum], + ], + dtype=dtype, + ) + pipeline = TosaPipelineFP[input_t]( + TopK(k), + (scores,), + aten_op, + exir_op, + tosa_extensions=["INT"], + run_on_tosa_ref_model=False, + ) + pipeline.run() + manager = pipeline.tester.get_artifact(StageType.TO_EDGE_TRANSFORM_AND_LOWER) + values, indices = _execute(manager, (scores,)) + assert values.shape == indices.shape == (scores.shape[0], k) + assert values.dtype == scores.dtype + assert indices.dtype == torch.int64 + assert torch.all((indices >= 0) & (indices < scores.shape[1])) + assert all( + row.unique().numel() == k for row in indices + ), "TopK must select K distinct indices per row" + assert torch.all(values[:, :-1] >= values[:, 1:]) + torch.testing.assert_close(values, scores.gather(1, indices), atol=0, rtol=0) + torch.testing.assert_close(values, scores.topk(k, dim=1).values, atol=0, rtol=0) + expected_indices = torch.argsort(scores, dim=1, descending=True, stable=True)[:, :k] + torch.testing.assert_close(indices, expected_indices, atol=0, rtol=0) + + +unsupported = { + "rank3": (TopK(1), torch.randn(2, 3, 8), "TOSA-1.0+FP+INT"), + "k5": (TopK(5), torch.randn(2, 8), "TOSA-1.0+FP+INT"), + "missing_int": (TopK(2), torch.randn(2, 8), "TOSA-1.0+FP"), + "missing_fp": (TopK(1), torch.randn(2, 8), "TOSA-1.0+INT"), +} + + +@common.parametrize("test_data", unsupported) +def test_topk_unsupported_tosa_FP(test_data): + model, scores, spec = test_data + manager = _lower(model, (scores,), spec) + assert any( + n.target == exir_ops.edge.aten.topk.default for n in _residual_ops(manager) + ) + assert not any( + n.target == _DELEGATE for n in manager.exported_program().graph.nodes + ) + + +@common.parametrize("axis", {"tokens": 0, "experts": 1}) +def test_topk_dynamic_shapes_not_delegated_tosa_FP(axis): + scores = torch.randn(3, 8) + shapes = ({axis: torch.export.Dim("varying", min=3, max=16)},) + manager = _lower(TopK(2), (scores,), dynamic_shapes=shapes) + assert any( + n.target == exir_ops.edge.aten.topk.default for n in _residual_ops(manager) + ) + + +class TopKGather(torch.nn.Module): + def forward(self, scores, features): + values, indices = torch.topk(scores, 3) + return values, indices, torch.gather(features, 1, indices) + + +@common.parametrize("portable", {"delegated": False, "portable": True}) +def test_topk_gather_boundary_tosa_FP(portable): + scores = torch.tensor([[3.0, 1.0, 5.0, 0.0, 2.0, 4.0]]) + features = torch.tensor([[10.0, 20.0, 30.0, 40.0, 50.0, 60.0]]) + checks = [DontPartition(exir_ops.edge.aten.gather.default)] if portable else None + manager = _lower(TopKGather(), (scores, features), checks=checks) + residual = _residual_ops(manager) + gathers = [n for n in residual if n.target == exir_ops.edge.aten.gather.default] + assert len(gathers) == int(portable) + if portable: + assert gathers[0].args[2].meta["val"].dtype == torch.int64 + assert not any(n.target == exir_ops.edge.aten.topk.default for n in residual) + assert manager.to_executorch().buffer + for current in (scores, scores.flip(1)): + actual = _execute(manager, (current, features)) + torch.testing.assert_close( + actual, TopKGather()(current, features), atol=0, rtol=0 + ) + + +def test_topk_mixed_gather_consumers_tosa_FP(): + class Model(torch.nn.Module): + def forward(self, scores, first, second): + indices = torch.topk(scores, 2).indices + return ( + torch.gather(first, 1, indices), + torch.gather(second, 1, indices), + indices, + ) + + class RejectSecondGather(OperatorSupportBase): + def is_node_supported(self, submodules, node): + return not ( + node.target == exir_ops.edge.aten.gather.default + and node.args[0].name == "second" + ) + + scores = torch.tensor([[2.0, 4.0, 1.0, 3.0]]) + inputs = (scores, scores + 10, scores + 20) + manager = _lower(Model(), inputs, checks=[RejectSecondGather()]) + residual = _residual_ops(manager) + assert not any(n.target == exir_ops.edge.aten.topk.default for n in residual) + gathers = [n for n in residual if n.target == exir_ops.edge.aten.gather.default] + assert len(gathers) == 1 + assert gathers[0].args[0].name == "second" + assert gathers[0].args[2].meta["val"].dtype == torch.int64 + assert manager.to_executorch().buffer + torch.testing.assert_close( + _execute(manager, inputs), Model()(*inputs), atol=0, rtol=0 + ) + + +@common.parametrize( + "target", + { + "topk": exir_ops.edge.aten.topk.default, + "getitem": operator.getitem, + "cast": _CAST, + }, +) +def test_topk_incomplete_index_chain_not_delegated_tosa_FP(target): + scores = torch.tensor([[3.0, 1.0, 2.0, 0.0]]) + manager = _lower(TopK(2), (scores,), checks=[DontPartition(target)]) + assert any( + n.target == exir_ops.edge.aten.topk.default for n in _residual_ops(manager) + ) + assert not any( + n.target == _DELEGATE for n in manager.exported_program().graph.nodes + ) + torch.testing.assert_close( + manager.exported_program().module()(scores), TopK(2)(scores) + ) + + +def test_topk_unsafe_index_arithmetic_tosa_FP(): + class Model(torch.nn.Module): + def forward(self, scores): + indices = torch.topk(scores, 1).indices + return indices * indices + + scores = torch.zeros(1, 50001) + scores[0, -1] = 1 + manager = _lower(Model(), (scores,), "TOSA-1.0+FP") + actual = _execute(manager, (scores,)) + assert actual.dtype == torch.int64 + assert actual.item() == 2500000000 + assert manager.to_executorch().buffer + + +@common.parametrize("prepared", {"prepared": True, "unprepared": False}) +def test_topk_legacy_partitioning_tosa_FP(prepared): + scores = torch.tensor([[3.0, 1.0, 0.0, 2.0]]) + compile_spec = TosaCompileSpec("TOSA-1.0+FP+INT") + manager = to_edge( + torch.export.export(TopK(2), (scores,)), + compile_config=EdgeCompileConfig(_check_ir_validity=False), + ) + if prepared: + manager = manager.transform( + [ + ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, tosa_spec=compile_spec.tosa_spec + ) + ] + ) + manager = manager.to_backend(TOSAPartitioner(compile_spec)) + has_topk = any( + n.target == exir_ops.edge.aten.topk.default for n in _residual_ops(manager) + ) + assert has_topk != prepared + if prepared: + assert manager.to_executorch().buffer + actual = _execute(manager, (scores,)) + else: + actual = manager.exported_program().module()(scores) + torch.testing.assert_close(actual, TopK(2)(scores), atol=0, rtol=0) diff --git a/backends/arm/test/passes/test_convert_int64_output_ops_to_int32.py b/backends/arm/test/passes/test_convert_int64_output_ops_to_int32.py index 23af5ad40a2..995556a1d77 100644 --- a/backends/arm/test/passes/test_convert_int64_output_ops_to_int32.py +++ b/backends/arm/test/passes/test_convert_int64_output_ops_to_int32.py @@ -3,6 +3,7 @@ # This source code is licensed under the BSD-style license found in the # LICENSE file in the root directory of this source tree. +import copy import operator from typing import Callable, Dict, Tuple @@ -14,6 +15,12 @@ ConvertInt64OutputOpsToInt32Pass, PrepareGatherIndicesPass, ) +from executorch.backends.arm._passes.decompose_topk_pass import ( + is_topk_indices_getitem, + is_topk_indices_int32_cast, + topk_indices_only_feed_int32_casts, + TOPK_OPS, +) from executorch.backends.arm._passes.prepare_gather_indices_pass import ( is_safe_int32_to_int64_gather_boundary, ) @@ -28,6 +35,7 @@ from executorch.exir.backend.operator_support import DontPartition, DontPartitionName from executorch.exir.dialects._ops import ops as exir_ops from torch.fx import Graph, GraphModule +from torch.utils import _pytree as pytree input_t1 = Tuple[torch.Tensor] # Input x @@ -1279,3 +1287,216 @@ def test_on_overflow_skip(): def test_on_overflow_invalid(): with pytest.raises(ValueError, match="on_overflow must be"): ConvertInt64OutputOpsToInt32Pass(on_overflow="blah") + + +class TopKPreparedConsumers(torch.nn.Module): + def __init__(self, output): + super().__init__() + self.output = output + + def forward(self, scores): + values, indices = torch.topk(scores, 3) + if self.output == "values": + return values + if self.output == "indices": + return indices + if self.output == "both": + return values, indices + if self.output == "int32": + return indices.int() + if self.output == "arithmetic": + return indices, indices * indices + return indices, torch.gather(scores, 1, indices), indices.unsqueeze(-1) + + +@pytest.mark.parametrize("edge", [False, True]) +@pytest.mark.parametrize( + "output", ["values", "indices", "both", "int32", "arithmetic", "gather"] +) +def test_topk_target_preparation_preserves_interfaces_and_is_idempotent(edge, output): + model = TopKPreparedConsumers(output) + scores = torch.tensor([[3.0, 1.0, 7.0, 2.0, 4.0, 6.0, 0.0, 5.0]]) + ep = torch.export.export(model, (scores,)) + if edge: + ep = to_edge( + ep, compile_config=EdgeCompileConfig(_check_ir_validity=False) + ).exported_program() + prepare = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, + tosa_spec=TosaSpecification.create_from_string("TOSA-1.0+FP+INT"), + ) + result = prepare(ep.graph_module) + actual = result.graph_module(scores) + expected = pytree.tree_leaves(model(scores)) + for a, b in zip(actual, expected, strict=True): + torch.testing.assert_close(a, b, atol=0, rtol=0) + topk = next(n for n in result.graph_module.graph.nodes if n.target in TOPK_OPS) + assert topk_indices_only_feed_int32_casts(topk) + if output == "gather": + graph = result.graph_module.graph + gather = next( + n + for n in graph.nodes + if n.target + in (torch.ops.aten.gather.default, exir_ops.edge.aten.gather.default) + ) + boundary = gather.args[2] + assert set(boundary.users) == {gather} + assert boundary is not graph.output_node().args[0][0] + assert boundary.meta["val"].dtype == torch.int64 + assert boundary.args[0].meta["val"].dtype == torch.int32 + graph_before = str(result.graph_module.graph) + repeated = prepare(result.graph_module) + assert not repeated.modified + assert str(repeated.graph_module.graph) == graph_before + + +@pytest.mark.parametrize( + "tosa_spec,convert_cast_ops", + [(None, False), ("TOSA-1.0+FP", False), (None, True)], +) +def test_topk_existing_int32_indices_need_no_conversion(tosa_spec, convert_cast_ops): + model = TopKPreparedConsumers("int32") + scores = torch.tensor([[3.0, 1.0, 0.0, 2.0]]) + ep = to_edge( + torch.export.export(model, (scores,)), + compile_config=EdgeCompileConfig(_check_ir_validity=False), + ).exported_program() + before = str(ep.graph) + result = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=convert_cast_ops, + tosa_spec=( + TosaSpecification.create_from_string(tosa_spec) if tosa_spec else None + ), + )(ep.graph_module) + assert not result.modified + assert str(result.graph_module.graph) == before + torch.testing.assert_close(result.graph_module(scores)[0], model(scores)) + + +def test_topk_target_preparation_does_not_narrow_unsafe_arithmetic(): + scores = torch.zeros(1, 50001) + scores[0, -3:] = torch.tensor([1.0, 2.0, 3.0]) + ep = torch.export.export(TopKPreparedConsumers("arithmetic"), (scores,)) + result = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, + tosa_spec=TosaSpecification.create_from_string("TOSA-1.0+FP+INT"), + )(ep.graph_module) + indices, squares = result.graph_module(scores) + assert indices.dtype == squares.dtype == torch.int64 + assert squares[0, 0].item() == 2500000000 + + +def test_topk_target_preparation_leaves_unsupported_target_unchanged(): + ep = torch.export.export(TopKPreparedConsumers("indices"), (torch.randn(2, 8),)) + before = str(ep.graph) + result = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, + tosa_spec=TosaSpecification.create_from_string("TOSA-1.0+FP"), + )(ep.graph_module) + assert not result.modified + assert str(result.graph_module.graph) == before + + +def test_topk_target_preparation_repeated_index_extractions(): + scores = torch.tensor([[3.0, 1.0, 0.0, 2.0]]) + ep = torch.export.export(TopKPreparedConsumers("indices"), (scores,)) + graph = ep.graph_module.graph + indices = next(n for n in graph.nodes if is_topk_indices_getitem(n)) + output = graph.output_node() + with graph.inserting_before(output): + duplicate = graph.call_function(operator.getitem, (indices.args[0], 1)) + duplicate.meta = indices.meta.copy() + output.args = ((indices, duplicate),) + spec = TosaSpecification.create_from_string("TOSA-1.0+FP+INT") + result = ConvertInt64OutputOpsToInt32Pass(convert_cast_ops=False, tosa_spec=spec)( + ep.graph_module + ) + expected = scores.topk(3).indices + first, second = result.graph_module(scores) + torch.testing.assert_close(first, expected) + torch.testing.assert_close(second, expected) + assert topk_indices_only_feed_int32_casts( + next(n for n in result.graph_module.graph.nodes if n.target in TOPK_OPS) + ) + + +@pytest.mark.parametrize("edge", [False, True]) +@pytest.mark.parametrize("tosa_spec", [None, "TOSA-1.0+FP", "TOSA-1.0+FP+INT"]) +def test_topk_reuses_existing_int32_cast(edge, tosa_spec): + scores = torch.tensor([[3.0, 1.0, 0.0, 2.0]]) + model = TopKPreparedConsumers("gather") + ep = torch.export.export(model, (scores,)) + if edge: + ep = to_edge( + ep, compile_config=EdgeCompileConfig(_check_ir_validity=False) + ).exported_program() + partial = ConvertInt64OutputOpsToInt32Pass(convert_cast_ops=False)( + ep.graph_module + ).graph_module + assert sum(is_topk_indices_int32_cast(n) for n in partial.graph.nodes) == 1 + assert not topk_indices_only_feed_int32_casts( + next(n for n in partial.graph.nodes if n.target in TOPK_OPS) + ) + result = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, + tosa_spec=( + TosaSpecification.create_from_string(tosa_spec) if tosa_spec else None + ), + )(partial) + int32_casts = [ + n for n in _cast_nodes(result.graph_module) if n.kwargs["dtype"] == torch.int32 + ] + assert len(int32_casts) == 1 + assert is_topk_indices_int32_cast(int32_casts[0]) + topk = next(n for n in result.graph_module.graph.nodes if n.target in TOPK_OPS) + assert topk_indices_only_feed_int32_casts(topk) == (tosa_spec == "TOSA-1.0+FP+INT") + for actual, expected in zip( + result.graph_module(scores), pytree.tree_leaves(model(scores)), strict=True + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + + +@pytest.mark.parametrize("edge", [False, True]) +@pytest.mark.parametrize("case", ["missing_int", "large_k", "symbolic"]) +def test_topk_target_preparation_retains_generic_conversion(edge, case): + class Model(torch.nn.Module): + def forward(self, scores): + indices = torch.topk(scores, 5 if case == "large_k" else 3).indices + return indices, indices.unsqueeze(-1) + + scores = torch.randn(2, 8) + shapes = ( + ({1: torch.export.Dim("experts", min=5, max=16)},) + if case == "symbolic" + else None + ) + ep = torch.export.export(Model(), (scores,), dynamic_shapes=shapes) + if edge: + ep = to_edge( + ep, compile_config=EdgeCompileConfig(_check_ir_validity=False) + ).exported_program() + generic = ConvertInt64OutputOpsToInt32Pass(convert_cast_ops=False)( + copy.deepcopy(ep.graph_module) + ) + result = ConvertInt64OutputOpsToInt32Pass( + convert_cast_ops=False, + tosa_spec=TosaSpecification.create_from_string( + "TOSA-1.0+FP" if case == "missing_int" else "TOSA-1.0+FP+INT" + ), + )(ep.graph_module) + assert generic.modified and result.modified + assert str(result.graph_module.graph) == str(generic.graph_module.graph) + assert not topk_indices_only_feed_int32_casts( + next(n for n in result.graph_module.graph.nodes if n.target in TOPK_OPS) + ) + torch.testing.assert_close( + result.graph_module(scores), Model()(scores), atol=0, rtol=0 + ) + + +def test_topk_target_preparation_requires_preserved_casts(): + with pytest.raises(ValueError, match="requires convert_cast_ops=False"): + ConvertInt64OutputOpsToInt32Pass( + tosa_spec=TosaSpecification.create_from_string("TOSA-1.0+FP+INT") + ) diff --git a/backends/arm/test/passes/test_decompose_topk_pass.py b/backends/arm/test/passes/test_decompose_topk_pass.py new file mode 100644 index 00000000000..2211a0ea9e5 --- /dev/null +++ b/backends/arm/test/passes/test_decompose_topk_pass.py @@ -0,0 +1,63 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import pytest +import torch +from executorch.backends.arm._passes import ( + ConvertInt64OutputOpsToInt32Pass, + DecomposeTopKPass, +) +from executorch.backends.arm._passes.decompose_topk_pass import TOPK_OPS +from executorch.backends.arm.tosa.specification import ( + TosaLoweringContext, + TosaSpecification, +) +from executorch.exir import EdgeCompileConfig, to_edge +from executorch.exir.dialects._ops import ops as exir_ops + + +@pytest.mark.parametrize("k", [1, 4]) +def test_decompose_topk_static_shapes_and_cumulative_masking(k): + dtype = torch.float32 + + class Model(torch.nn.Module): + def forward(self, scores): + values, indices = torch.topk(scores, k) + return values, indices.int() + + ep = to_edge( + torch.export.export(Model(), (torch.randn(2, 8, dtype=dtype),)), + compile_config=EdgeCompileConfig(_check_ir_validity=False), + ).exported_program() + spec = TosaSpecification.create_from_string("TOSA-1.0+FP+INT") + gm = ConvertInt64OutputOpsToInt32Pass(convert_cast_ops=False, tosa_spec=spec)( + ep.graph_module + ).graph_module + with TosaLoweringContext(spec): + result = DecomposeTopKPass(spec)(gm) + repeated = DecomposeTopKPass(spec)(result.graph_module) + assert result.modified and not repeated.modified + nodes = list(result.graph_module.graph.nodes) + assert not any(n.target in TOPK_OPS for n in nodes) + argmaxes = [n for n in nodes if n.target == exir_ops.backend.tosa.ARGMAX.default] + masks = [n for n in nodes if n.target == exir_ops.edge.aten.where.self] + assert len(argmaxes) == k + assert len(masks) == k - 1 + scores = next(n for n in nodes if n.op == "placeholder") + for step, argmax in enumerate(argmaxes): + assert argmax.meta["val"].dtype == torch.int32 + assert tuple(argmax.meta["val"].shape) == (2,) + assert argmax.args[0] is (scores if step == 0 else masks[step - 1]) + if step < k - 1: + assert masks[step].args[2] is argmax.args[0] + gather = next(n for n in nodes if n.target == exir_ops.edge.aten.gather.default) + assert gather.args[0] is scores + assert gather.args[2].meta["val"].dtype == torch.int32 + assert tuple(gather.meta["val"].shape) == (2, k) + assert gather.meta["val"].dtype == dtype + if k > 1: + sentinel = next(n for n in nodes if n.target == exir_ops.edge.aten.full.default) + assert sentinel.args[1] == float("-inf") + assert sentinel.kwargs["dtype"] == dtype diff --git a/backends/arm/test/targets.bzl b/backends/arm/test/targets.bzl index d45b8e86790..c81cf5113b4 100644 --- a/backends/arm/test/targets.bzl +++ b/backends/arm/test/targets.bzl @@ -38,6 +38,7 @@ def define_arm_tests(): "ops/test_view.py", "ops/test_cos.py", "ops/test_to_copy.py", + "ops/test_topk.py", "ops/test_exp.py", "ops/test_fft.py", "ops/test_flip.py", @@ -70,6 +71,7 @@ def define_arm_tests(): # "misc/test_evaluate_model.py", "misc/test_pass_pipeline_config.py", "misc/test_tosa_constant_pool.py", + "misc/test_tosa_operator_support.py", "misc/tosa_dialect/test_tosa_dialect_cast_to_block_scaled.py", "misc/tosa_dialect/test_tosa_dialect_mxfp_conv2d.py", "misc/tosa_dialect/test_tosa_dialect_mxfp_linear.py", diff --git a/backends/arm/tosa/partitioner.py b/backends/arm/tosa/partitioner.py index e35f64a7f62..db215d601b9 100644 --- a/backends/arm/tosa/partitioner.py +++ b/backends/arm/tosa/partitioner.py @@ -27,6 +27,11 @@ from executorch.backends.arm._passes.convert_expand_copy_to_repeat import ( calculate_multiples, ) +from executorch.backends.arm._passes.decompose_topk_pass import ( + is_topk_indices_getitem, + is_topk_indices_int32_cast, + TOPK_OPS, +) from executorch.backends.arm._passes.prepare_gather_indices_pass import ( is_safe_int32_to_int64_gather_boundary, ) @@ -310,6 +315,32 @@ def _find_connected_components(nodes: set[torch.fx.Node]) -> list[set[torch.fx.N return components +def _detag_incomplete_topk_chains(nodes: Iterable[torch.fx.Node]) -> set[str]: + """Keep TopK, tuple extraction, and index narrowing in one delegate.""" + affected_tags: set[str] = set() + for topk in nodes: + if topk.target not in TOPK_OPS: + continue + chain = {topk, *topk.users} + for extraction in topk.users: + if is_topk_indices_getitem(extraction): + chain.update( + user + for user in extraction.users + if is_topk_indices_int32_cast(user) + ) + tag = topk.meta.get("delegation_tag") + if tag is not None and all( + node.meta.get("delegation_tag") == tag for node in chain + ): + continue + for node in chain: + node_tag = node.meta.pop("delegation_tag", None) + if node_tag is not None: + affected_tags.add(node_tag) + return affected_tags + + def _detag_mixed_delegate_gather_boundaries( nodes: Iterable[torch.fx.Node], ) -> set[str]: @@ -739,6 +770,7 @@ def _tag_module( # noqa if active_tag in tags: tags.remove(active_tag) affected_tags = _detag_mixed_delegate_gather_boundaries(module.graph.nodes) + affected_tags.update(_detag_incomplete_topk_chains(module.graph.nodes)) if affected_tags: _retag_affected_partitions( module.graph.nodes, diff --git a/docs/source/backends/arm-vgf/VGF_op_support.md b/docs/source/backends/arm-vgf/VGF_op_support.md index 14776284ae9..4a2bc8f8816 100644 --- a/docs/source/backends/arm-vgf/VGF_op_support.md +++ b/docs/source/backends/arm-vgf/VGF_op_support.md @@ -6,7 +6,7 @@ This page lists VGF-supported PyTorch APIs and the dtype and quantization modes `8x8` means 8-bit activations and 8-bit weights. `16x8` means 16-bit activations and 8-bit weights. `8x4` means 8-bit activations and 4-bit weights. -Total supported PyTorch APIs: **157**. +Total supported PyTorch APIs: **158**. | PyTorch API | Support profile | DType | Quantization mode | | --- | --- | --- | --- | @@ -158,6 +158,7 @@ Total supported PyTorch APIs: **157**. | `torch.Tensor.repeat` | FP, INT | `FP16`, `BF16`, `INT8`, `INT16`, `BOOL` | 8x8, 16x8 | | `torch.Tensor.unfold` | FP, INT | `FP32`, `FP16`, `BF16`, `INT8`, `BOOL` | 8x8 | | `torch.Tensor.view` | FP, INT | `FP16`, `INT8` | 8x8 | +| `torch.topk` | FP | `FP32`, `FP16` | - | | `torch.transpose` / `torch.Tensor.transpose` | FP, INT | `FP32`, `INT8` | 8x8 | | `torch.tril` | FP | `FP32` | - | | `torch.unbind` | FP, INT | `FP32`, `INT8` | 8x8 |