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 |