Skip to content

Commit 83c4dc1

Browse files
authored
Revert pytorch#23005: Prevent edge_transform_passes from running twice (pytorch#23119)
Reverts pytorch#23005. The change in pytorch#23005 moved edge_transform_passes out of EdgeProgramManagerTransformStage, but this causes recipe-based lowering to lose the stage behavior expected by some existing callers. Restore the stage wiring and its regression test. Test plan: - `git diff --check upstream/main...HEAD` - Not run: `python -m pytest export/tests/test_export_session.py` (pytest is not installed in the current environment).
1 parent b0f0f1c commit 83c4dc1

3 files changed

Lines changed: 38 additions & 30 deletions

File tree

‎export/recipe.py‎

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -207,12 +207,9 @@ class LoweringRecipe:
207207
per-method partitioner lists. Use the dict form when
208208
backends need per-method compile specs.
209209
edge_transform_passes: Optional list of callables that take (method_name: str, exported_program: ExportedProgram)
210-
and return either List[PassType] or PassManager. Applied during
211-
TO_EDGE_TRANSFORM_AND_LOWER as per-ExportedProgram graph-module passes.
210+
and return either List[PassType] or PassManager to be applied during edge lowering.
212211
edge_manager_transform_passes: Optional list of callables that take EdgeProgramManager as argument
213-
and return passes to be applied. Applied sequentially by
214-
EDGE_PROGRAM_MANAGER_TRANSFORM, which runs after
215-
TO_EDGE_TRANSFORM_AND_LOWER in the default pipeline.
212+
and return passes to be applied. Applied sequentially after TO_EDGE stage.
216213
edge_compile_config: Optional edge compilation configuration
217214
pre_partitioning_callback: Optional callable invoked just before partitioning with
218215
`(partitioners, programs)` arguments.

‎export/stages.py‎

Lines changed: 36 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -734,6 +734,14 @@ class EdgeProgramManagerTransformStage(Stage):
734734

735735
def __init__(
736736
self,
737+
edge_transform_passes: (
738+
None
739+
| List[
740+
Callable[
741+
[str, ExportedProgram], List[PassType] | GraphModulePassManager
742+
]
743+
]
744+
) = None,
737745
edge_manager_transform_passes: (
738746
None
739747
| List[
@@ -751,6 +759,7 @@ def __init__(
751759
backends to control pass ordering and dependencies.
752760
"""
753761
super().__init__()
762+
self._edge_transform_passes = edge_transform_passes or []
754763
self._edge_manager_transform_passes = edge_manager_transform_passes or []
755764

756765
@classmethod
@@ -761,6 +770,7 @@ def from_recipe(
761770
return cls()
762771

763772
return cls(
773+
edge_transform_passes=lowering_recipe.edge_transform_passes,
764774
edge_manager_transform_passes=lowering_recipe.edge_manager_transform_passes,
765775
)
766776

@@ -793,10 +803,35 @@ def run(self, artifact: PipelineArtifact) -> None:
793803
f"Expected EdgeProgramManager but got {type(edge_program_manager)}"
794804
)
795805

796-
if not self._edge_manager_transform_passes:
806+
if not self._edge_transform_passes and not self._edge_manager_transform_passes:
797807
self._artifact = artifact
798808
return
799809

810+
# Detect if any callable returns PassManager
811+
pass_manager = None
812+
transform_passes = defaultdict(list)
813+
for method_name in edge_program_manager.methods:
814+
# Resolve transform passes if it's a callable
815+
ep = edge_program_manager.exported_program(method_name)
816+
for pass_callable in self._edge_transform_passes or []:
817+
if not callable(pass_callable):
818+
raise ValueError(
819+
"Transform passes must be a callable that resolves to passes"
820+
)
821+
passes = pass_callable(method_name, ep)
822+
if isinstance(passes, GraphModulePassManager):
823+
pass_manager = passes
824+
break
825+
else:
826+
transform_passes[method_name].extend(passes)
827+
if pass_manager:
828+
break
829+
830+
# See EdgeTransformAndLowerStage.run.
831+
final_passes = pass_manager or _drop_empty(transform_passes) or None
832+
if final_passes is not None:
833+
edge_program_manager = edge_program_manager.transform(final_passes)
834+
800835
# Run edge manager transform passes
801836
for pass_callable in self._edge_manager_transform_passes:
802837
passes = pass_callable(edge_program_manager)

‎export/tests/test_export_session.py‎

Lines changed: 0 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -757,30 +757,6 @@ def test_no_transform_passes_means_no_stage(self) -> None:
757757
StageType.EDGE_PROGRAM_MANAGER_TRANSFORM, session._get_default_pipeline()
758758
)
759759

760-
def test_edge_transform_passes_not_duplicated_in_default_pipeline(self) -> None:
761-
# Before the fix, from_recipe() gave edge_transform_passes to
762-
# EDGE_PROGRAM_MANAGER_TRANSFORM as well as TO_EDGE_TRANSFORM_AND_LOWER,
763-
# so they ran twice. Verify the stage no longer holds them at all.
764-
from executorch.export.stages import EdgeProgramManagerTransformStage
765-
766-
edge_pass = Mock()
767-
epm_pass = Mock()
768-
769-
stage = EdgeProgramManagerTransformStage.from_recipe(
770-
LoweringRecipe(
771-
edge_transform_passes=[edge_pass],
772-
edge_manager_transform_passes=[epm_pass],
773-
)
774-
)
775-
776-
# The stage must only know about edge_manager_transform_passes.
777-
self.assertEqual(stage._edge_manager_transform_passes, [epm_pass])
778-
self.assertFalse(
779-
hasattr(stage, "_edge_transform_passes"),
780-
"EdgeProgramManagerTransformStage must not hold edge_transform_passes "
781-
"because TO_EDGE_TRANSFORM_AND_LOWER already applies them.",
782-
)
783-
784760
def test_example_inputs_required_for_nn_module(self) -> None:
785761
"""Test that example_inputs are required for nn.Module."""
786762
with self.assertRaises(ValueError) as cm:

0 commit comments

Comments
 (0)