Skip to content

Commit c8d5189

Browse files
Arm backend: Add pre-decomposition partitioner pipeline (pytorch#22514)
Route the EXIR pre-decomposition hook through an Arm pass pipeline. Store compile specs on concrete partitioners so the pipeline can be configured consistently for TOSA, Ethos-U, and VGF. Keep the pipeline a no-op until backend-specific passes are registered, and update generated partitioner docs and public API manifests. Assisted by Codex. Change-Id: I1963bad577714907e5ba22eed765333d9816308c Signed-off-by: Yufeng Shi <yufeng.shi@arm.com>
1 parent 8ea353d commit c8d5189

7 files changed

Lines changed: 74 additions & 0 deletions

File tree

‎backends/arm/_passes/arm_pass_manager.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -458,6 +458,12 @@ def _tosa_context(self, graph_module: GraphModule) -> TosaLoweringContext:
458458
shape_env = _get_shape_env_from_gm(graph_module)
459459
return TosaLoweringContext(self.tosa_spec, shape_env)
460460

461+
def transform_for_pre_decomposition_pipeline(
462+
self, exported_program: ExportedProgram
463+
) -> ExportedProgram:
464+
"""Apply Arm passes before default ATen decompositions."""
465+
return exported_program
466+
461467
def _transform_graph_module(self, graph_module: GraphModule):
462468
# TFA and control-flow submodule paths operate on bare GraphModules
463469
# without a standalone ExportedProgram to keep in sync.

‎backends/arm/ethosu/partitioner.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ def __init__(
3434
self.delegation_spec = DelegationSpec(
3535
EthosUBackend.__name__, compile_spec._to_list()
3636
)
37+
self.compile_spec = compile_spec
3738
self.additional_checks = additional_checks
3839
self.tosa_spec = compile_spec.tosa_spec
3940
self._decomposable_resize_support = DecomposableResizeSupported(self.tosa_spec)

‎backends/arm/public_api_manifests/api_manifest_running.toml‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,10 @@ signature = "EthosUPartitioner.partition(self, exported_program: torch.export.ex
6060
kind = "function"
6161
signature = "EthosUPartitioner.register_custom_partition_op(self, op: torch._ops.OpOverload) -> None"
6262

63+
[python.EthosUPartitioner.transform_for_pre_decomposition]
64+
kind = "function"
65+
signature = "EthosUPartitioner.transform_for_pre_decomposition(self, exported_program: torch.export.exported_program.ExportedProgram) -> torch.export.exported_program.ExportedProgram"
66+
6367
[python.EthosUQuantizer]
6468
kind = "class"
6569
signature = "EthosUQuantizer(compile_spec: 'EthosUCompileSpec', use_composable_quantizer: 'bool' = True) -> 'None'"
@@ -180,6 +184,10 @@ signature = "VgfPartitioner.partition(self, exported_program: torch.export.expor
180184
kind = "function"
181185
signature = "VgfPartitioner.register_custom_partition_op(self, op: torch._ops.OpOverload) -> None"
182186

187+
[python.VgfPartitioner.transform_for_pre_decomposition]
188+
kind = "function"
189+
signature = "VgfPartitioner.transform_for_pre_decomposition(self, exported_program: torch.export.exported_program.ExportedProgram) -> torch.export.exported_program.ExportedProgram"
190+
183191
[python.VgfQuantizer]
184192
kind = "class"
185193
signature = "VgfQuantizer(compile_spec: 'VgfCompileSpec', use_composable_quantizer: 'bool' = True) -> 'None'"

‎backends/arm/tosa/partitioner.py‎

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
from typing import Callable, cast, List, Mapping, Optional, Sequence, Tuple
2222

2323
import torch
24+
from executorch.backends.arm._passes.arm_pass_manager import ArmPassManager
2425
from executorch.backends.arm._passes.arm_pass_utils import get_first_fake_tensor
2526
from executorch.backends.arm._passes.convert_expand_copy_to_repeat import (
2627
calculate_multiples,
@@ -32,6 +33,7 @@
3233
is_exact_tosa_boundary_bilinear_downscale,
3334
)
3435

36+
from executorch.backends.arm.common.arm_compile_spec import ArmCompileSpec
3537
from executorch.backends.arm.common.type import ensure_type
3638
from executorch.backends.arm.constants import DQ_OPS, Q_OPS
3739
from executorch.backends.arm.operator_support.tosa_supported_operators import (
@@ -361,6 +363,8 @@ class TOSAPartitioner(Partitioner):
361363
362364
"""
363365

366+
compile_spec: ArmCompileSpec
367+
364368
def __init__(
365369
self,
366370
compile_spec: TosaCompileSpec,
@@ -381,12 +385,34 @@ def __init__(
381385
self.delegation_spec = DelegationSpec(
382386
TOSABackend.__name__, compile_spec._to_list()
383387
)
388+
self.compile_spec = compile_spec
384389
self.tosa_spec = compile_spec.tosa_spec
385390
self.additional_checks = additional_checks
386391
self._decomposable_resize_support = DecomposableResizeSupported(self.tosa_spec)
387392
self._custom_partition_ops: set[torch._ops.OpOverload] = set()
388393
self.intermediate_path = compile_spec._get_intermediate_path()
389394

395+
def transform_for_pre_decomposition(
396+
self, exported_program: ExportedProgram
397+
) -> ExportedProgram:
398+
"""Apply required Arm passes before default ATen decompositions.
399+
400+
EXIR invokes this backend extension hook automatically through
401+
``to_edge_transform_and_lower``. Model export users should not call it
402+
directly.
403+
404+
Args:
405+
exported_program (ExportedProgram): The ATen-dialect program to
406+
transform.
407+
408+
Returns:
409+
ExportedProgram: The transformed ATen-dialect program.
410+
411+
"""
412+
return ArmPassManager(
413+
self.compile_spec
414+
).transform_for_pre_decomposition_pipeline(exported_program)
415+
390416
def register_custom_partition_op(self, op: torch._ops.OpOverload) -> None:
391417
"""Register a custom op to be considered supported."""
392418
self._custom_partition_ops.add(op)

‎backends/arm/vgf/partitioner.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@ def __init__(
3535
self.delegation_spec = DelegationSpec(
3636
VgfBackend.__name__, compile_spec._to_list()
3737
)
38+
self.compile_spec = compile_spec
3839
self.additional_checks = additional_checks
3940
self.tosa_spec = compile_spec.tosa_spec
4041
self._decomposable_resize_support = DecomposableResizeSupported(self.tosa_spec)

‎docs/source/backends/arm-ethos-u/arm-ethos-u-partitioner.md‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,3 +50,19 @@ Returns:
5050
def EthosUPartitioner.register_custom_partition_op(self, op: torch._ops.OpOverload) -> None:
5151
```
5252
Register a custom op to be considered supported.
53+
54+
```python
55+
def EthosUPartitioner.transform_for_pre_decomposition(self, exported_program: torch.export.exported_program.ExportedProgram) -> torch.export.exported_program.ExportedProgram:
56+
```
57+
Apply required Arm passes before default ATen decompositions.
58+
59+
EXIR invokes this backend extension hook automatically through
60+
``to_edge_transform_and_lower``. Model export users should not call it
61+
directly.
62+
63+
Args:
64+
- **exported_program (ExportedProgram)**: The ATen-dialect program to
65+
transform.
66+
67+
Returns:
68+
- **ExportedProgram**: The transformed ATen-dialect program.

‎docs/source/backends/arm-vgf/arm-vgf-partitioner.md‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,3 +50,19 @@ Returns:
5050
def VgfPartitioner.register_custom_partition_op(self, op: torch._ops.OpOverload) -> None:
5151
```
5252
Register a custom op to be considered supported.
53+
54+
```python
55+
def VgfPartitioner.transform_for_pre_decomposition(self, exported_program: torch.export.exported_program.ExportedProgram) -> torch.export.exported_program.ExportedProgram:
56+
```
57+
Apply required Arm passes before default ATen decompositions.
58+
59+
EXIR invokes this backend extension hook automatically through
60+
``to_edge_transform_and_lower``. Model export users should not call it
61+
directly.
62+
63+
Args:
64+
- **exported_program (ExportedProgram)**: The ATen-dialect program to
65+
transform.
66+
67+
Returns:
68+
- **ExportedProgram**: The transformed ATen-dialect program.

0 commit comments

Comments
 (0)