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
4 changes: 1 addition & 3 deletions .claude/skills/cortex-m/SKILL.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,10 +33,8 @@ exported = export(quantized, example_inputs)
edge = to_edge_transform_and_lower(
exported,
compile_config=cortex_m_edge_compile_config(),
transform_passes=CortexMPassManager(),
)
edge._edge_programs["forward"] = CortexMPassManager(
edge.exported_program(), CortexMPassManager.pass_list
).transform()
et_program = edge.to_executorch()
```

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -247,10 +247,8 @@ def _export_cortex_m(
_core_aten_ops_exception_list=[torch.ops.aten.max_pool2d.default],
),
constant_methods=metadata,
transform_passes=CortexMPassManager(target_config=target_config),
)
edge._edge_programs["forward"] = CortexMPassManager(
edge.exported_program(), target_config=target_config
).transform()
return edge.to_executorch()


Expand Down
11 changes: 4 additions & 7 deletions backends/arm/scripts/aot_arm_compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -983,15 +983,12 @@ def _to_channels_last(x):
edge = to_edge_transform_and_lower(
exported_program,
compile_config=cortex_m_edge_compile_config(),
transform_passes=CortexMPassManager(
target_config=target_config,
use_explicit_layout=args.cortex_m_explicit_layout,
),
)

pass_manager = CortexMPassManager(
edge.exported_program(),
target_config=target_config,
use_explicit_layout=args.cortex_m_explicit_layout,
)
edge._edge_programs["forward"] = pass_manager.transform()

return model_quant, edge, example_inputs


Expand Down
2 changes: 1 addition & 1 deletion backends/cortex_m/passes/cortex_m_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ class CortexMPass(ExportPass):
"""Base class for passes that need the Cortex-M target config.

Passes that subclass this declare `exported_program` and `target_config`
in their `__init__`; `CortexMPassManager.transform()` injects both
in their `__init__`; `CortexMPassManager` injects both
automatically when running the pass list.
"""

Expand Down
85 changes: 52 additions & 33 deletions backends/cortex_m/passes/cortex_m_pass_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,12 @@
from executorch.backends.transforms.replace_squeeze_unsqueeze_with_view import (
ReplaceSqueezeAndUnsqueezeWithViewPass,
)
from executorch.exir.pass_base import ExportPass
from executorch.exir.pass_manager import PassManager
from executorch.exir.pass_base import (
ExportedProgramPassBase,
ExportedProgramPassResult,
ExportPass,
)
from executorch.exir.pass_manager import ExportedProgramPassManager, PassType
from executorch.exir.program._program import _transform, lift_constant_tensor_pass
from torch.export import ExportedProgram

Expand All @@ -50,7 +54,36 @@
PassClass = Type[ExportPass]


class CortexMPassManager(PassManager):
class _CortexMLoweringPass(ExportedProgramPassBase):
def __init__(
self, pass_classes: list[PassClass], target_config: CortexMTargetConfig
) -> None:
self.pass_classes = pass_classes
self.target_config = target_config

def call(self, exported_program: ExportedProgram) -> ExportedProgramPassResult:
modified = False
for pass_cls in self.pass_classes:
signature = inspect.signature(pass_cls)
kwargs: dict[str, Any] = {}
if "exported_program" in signature.parameters:
kwargs["exported_program"] = exported_program
if "target_config" in signature.parameters:
kwargs["target_config"] = self.target_config

transform_pass = pass_cls(**kwargs)
transformed = _transform(exported_program, transform_pass)
modified |= transformed is not exported_program
exported_program = transformed

# Passes can introduce tensor attributes that must become program inputs.
buffer_count = len(exported_program.graph_signature.buffers)
exported_program = lift_constant_tensor_pass(exported_program)
modified |= len(exported_program.graph_signature.buffers) != buffer_count
return ExportedProgramPassResult(exported_program, modified)


class CortexMPassManager(ExportedProgramPassManager):
legacy_pass_list: list[PassClass] = [
# Run before folding so qparams attach to max_pool2d values, not tuple + getitem.
RemoveGetItemPass,
Expand Down Expand Up @@ -98,17 +131,16 @@ class CortexMPassManager(PassManager):

def __init__(
self,
exported_program: ExportedProgram | None,
exported_program: ExportedProgram | None = None,
passes: Optional[list[PassClass]] = None,
target_config: Optional[CortexMTargetConfig] = None,
use_explicit_layout: bool = False,
) -> None:
"""Initialize the Cortex-M pass manager.

Args:
exported_program: The exported program to transform. Required
before calling ``transform()``; may be ``None`` for callers
that only use ``transform_for_annotation()``.
exported_program: Optional program for the legacy ``transform()``
entry point. Omit when using ``edge.transform(pass_manager)``.
passes: Optional override of the pass list. Defaults to
the legacy or explicit-layout pass list selected by
``use_explicit_layout``.
Expand All @@ -119,20 +151,26 @@ def __init__(
use_explicit_layout: Select the experimental explicit-layout pass
sequence. Legacy lowering remains the default.
"""
super().__init__(passes=[])
self.exported_program = exported_program
# PassManager.passes is typed as callables; this manager stores pass classes which are initialized at transform time with the exported_program.
default_passes = (
self.explicit_layout_pass_list
if use_explicit_layout
else self.legacy_pass_list
)
self.passes: list[PassClass] = ( # type: ignore[assignment]
passes if passes is not None else default_passes # type: ignore[assignment]
)
pass_classes = passes if passes is not None else default_passes
for pass_cls in pass_classes:
if not isinstance(pass_cls, type):
raise ValueError(
f"{type(self).__name__} expects pass classes, not instances; "
f"got {pass_cls!r}"
)
self.target_config: CortexMTargetConfig = target_config or CortexMTargetConfig(
cpu=CortexM.M55
)
lowering_passes: list[PassType] = [
_CortexMLoweringPass(pass_classes, self.target_config)
]
super().__init__(lowering_passes)

def transform_for_annotation(self, model):
passes = self.pass_list_transform_for_annotation
Expand All @@ -148,24 +186,5 @@ def transform(self) -> ExportedProgram:
f"got {exported_program!r}"
)

for pass_cls in self.passes:
if not isinstance(pass_cls, type):
raise ValueError(
f"{type(self).__name__} expects pass classes, not instances; "
f"got {pass_cls!r}"
)

signature = inspect.signature(pass_cls)
kwargs: dict[str, Any] = {}
if "exported_program" in signature.parameters:
kwargs["exported_program"] = exported_program
if "target_config" in signature.parameters:
kwargs["target_config"] = self.target_config

transform_pass = pass_cls(**kwargs)
exported_program = _transform(exported_program, transform_pass)

# All constant tensors should be lifted to buffers at this point, re-run
# lift_constant_tensor_pass in case new ones have been introduced.
exported_program = lift_constant_tensor_pass(exported_program)
return exported_program
result = self(exported_program)
return result.exported_program if result.modified else exported_program
2 changes: 1 addition & 1 deletion backends/cortex_m/quantizer/quantizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,5 +117,5 @@ def validate(self, model: GraphModule) -> None:
return None

def transform_for_annotation(self, model: GraphModule) -> GraphModule:
pass_manager = CortexMPassManager(None)
pass_manager = CortexMPassManager()
return pass_manager.transform_for_annotation(model)
4 changes: 2 additions & 2 deletions backends/cortex_m/test/misc/test_target_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,7 @@ def test_default_target_config_is_m55(self):
CortexMPassManager,
)

pm = CortexMPassManager(exported_program=None)
pm = CortexMPassManager()
assert pm.target_config.cpu == CortexM.M55
assert pm.target_config.backend == cmsis_nn.Backend.MVE

Expand All @@ -112,6 +112,6 @@ def test_explicit_target_config_threaded(self):
)

target_config = CortexMTargetConfig(cpu=CortexM.M33)
pm = CortexMPassManager(exported_program=None, target_config=target_config)
pm = CortexMPassManager(target_config=target_config)
assert pm.target_config.cpu == CortexM.M33
assert pm.target_config.backend == cmsis_nn.Backend.DSP
21 changes: 21 additions & 0 deletions backends/cortex_m/test/targets.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ def define_common_targets(is_fbcode = False):
"//executorch/backends/cortex_m/passes:cortex_passes",
"//executorch/backends/cortex_m/quantizer:quantizer",
"//executorch/backends/test/harness:tester",
"//executorch/exir:lib",
],
)

Expand Down Expand Up @@ -101,6 +102,26 @@ def define_common_targets(is_fbcode = False):
],
)

python_pytest(
name = "test_pass_manager",
srcs = ["test_pass_manager.py"],
compile = "with-source",
typing = False,
deps = [
"//caffe2:torch",
"//pytorch/ao:torchao", # @manual
"//executorch/backends/cortex_m:edge_compile_config",
"//executorch/backends/cortex_m:target_config",
"//executorch/backends/cortex_m/passes:cortex_passes",
"//executorch/backends/cortex_m/quantizer:quantizer",
"//executorch/exir:lib",
"//executorch/exir:pass_base",
"//executorch/exir/_serialize:lib",
"//executorch/exir/dialects:lib",
"fbsource//third-party/pypi/pytest:pytest",
],
)



python_pytest(
Expand Down
23 changes: 9 additions & 14 deletions backends/cortex_m/test/test_explicit_layout_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,15 +4,12 @@
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from functools import partial

import pytest
import torch
from executorch.backends.cortex_m.passes.cortex_m_pass_manager import CortexMPassManager
from executorch.backends.cortex_m.quantizer.quantizer import CortexMQuantizer
from executorch.backends.cortex_m.target_config import CortexM, CortexMTargetConfig
from executorch.backends.cortex_m.test.tester import CortexMTester
from executorch.backends.test.harness.stages import Quantize, RunPasses, StageType
from executorch.backends.cortex_m.test.tester import CortexMRunPasses, CortexMTester
from executorch.backends.test.harness.stages import Quantize, StageType
from executorch.exir.dialects._ops import ops as exir_ops
from torch.fx import Node

Expand Down Expand Up @@ -75,13 +72,9 @@ def _count(exported_program, target) -> int:
def _run_explicit_layout_pass_manager(tester: CortexMTester) -> CortexMTester:
target_config = CortexMTargetConfig(cpu=CortexM.M55)
tester.run_passes(
RunPasses(
partial(
CortexMPassManager,
target_config=target_config,
use_explicit_layout=True,
), # type: ignore[arg-type]
CortexMPassManager.explicit_layout_pass_list, # type: ignore[arg-type]
CortexMRunPasses(
target_config=target_config,
use_explicit_layout=True,
)
)
return tester
Expand Down Expand Up @@ -216,5 +209,7 @@ def test_explicit_layout_rejects_unsupported_spatial_operator():
with pytest.raises(Exception) as caught:
_run_explicit_layout_passes(tester)

assert caught.value.__cause__ is not None
assert "NHWC-eligible" in str(caught.value.__cause__)
error = caught.value
while error.__cause__ is not None:
error = error.__cause__
assert "NHWC-eligible" in str(error)
Loading
Loading