Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions backends/arm/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions backends/arm/_passes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
6 changes: 5 additions & 1 deletion backends/arm/_passes/arm_pass_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@
DecomposeStridedSliceCopyPass,
DecomposeSumPass,
DecomposeTanPass,
DecomposeTopKPass,
DecomposeTOSAUnsupportedClampPass,
DecomposeTrilPass,
DecomposeUnfoldToGatherPass,
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -540,6 +543,7 @@ def _tosa_pipeline(
NormalizeDelegateIOLayoutPass(exported_program),
FuseQuantizedActivationPass(),
RewriteBoolToFp32CastViaInt8Pass(),
DecomposeTopKPass(self.tosa_spec),
PrepareGatherIndicesPass(self.tosa_spec),
CanonicalizeGatherPass(),
ConvertToClampPass(),
Expand Down
149 changes: 131 additions & 18 deletions backends/arm/_passes/convert_int64_output_ops_to_int32.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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.
Expand All @@ -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.
"""

Expand All @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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.
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down
Loading
Loading