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
97 changes: 93 additions & 4 deletions exir/pass_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
import traceback
from abc import ABC, abstractmethod
from contextlib import nullcontext
from dataclasses import dataclass
from dataclasses import dataclass, replace
from typing import (
Any,
Callable,
Expand Down Expand Up @@ -38,6 +38,7 @@
from torch._subclasses.fake_tensor import FakeTensor
from torch._subclasses.functional_tensor import FunctionalTensor, FunctionalTensorMode
from torch.export import ExportedProgram
from torch.export.graph_signature import ConstantArgument
from torch.fx import traceback as fx_traceback
from torch.fx.experimental.proxy_tensor import PythonKeyTracer
from torch.fx.graph import CodeGen
Expand Down Expand Up @@ -237,6 +238,65 @@ class ExportedProgramPassResult:
modified: bool


def _sync_output_specs(exported_program: ExportedProgram) -> bool:
"""Realign ``output_specs`` with the graph's output node.

A pass that replaces an output node leaves the signature naming the old one,
which later stages then disagree with. Rewriting the spec to match is always
safe; changing an output between a node and a literal is not, so that raises
rather than guessing.

Returns whether any spec changed.
"""
output_node = exported_program.graph_module.graph.output_node()
outputs = output_node.args[0]
assert isinstance(outputs, (tuple, list))

output_specs = exported_program.graph_signature.output_specs
if len(outputs) != len(output_specs):
raise ExportPassBaseError(
f"Graph has {len(outputs)} outputs, but its signature has "
f"{len(output_specs)} output specs"
)

modified = False
for output, output_spec in zip(outputs, output_specs):
if isinstance(output_spec.arg, ConstantArgument):
if isinstance(output, torch.fx.Node):
raise ExportPassBaseError(
f"Output {output.name} replaced a literal output; changing output "
"representation is not supported"
)
if output_spec.arg.value != output:
output_spec.arg = replace(output_spec.arg, value=output)
modified = True
continue
if not isinstance(output, torch.fx.Node):
raise ExportPassBaseError(
f"Output {output_spec.arg.name} became a literal; changing output "
"representation is not supported"
)
if output_spec.arg.name != output.name:
output_spec.arg = replace(output_spec.arg, name=output.name)
modified = True
return modified


def _rename_output_specs(exported_program: ExportedProgram, old: str, new: str) -> None:
"""Rename ``old`` to ``new`` in ``output_specs`` only.

``ExportGraphSignature.get_replace_hook`` renames input specs alongside the
output ones, so replacing a placeholder that is also returned renames its
input spec out from under callers that still look it up by its old name
(``delete_constant_placeholder``, for one).
"""
for output_spec in exported_program.graph_signature.output_specs:
if isinstance(output_spec.arg, ConstantArgument):
continue
if output_spec.arg.name == old:
output_spec.arg = replace(output_spec.arg, name=new)


class ExportedProgramPassBase(ABC):
"""
Base interface for implementing passes that operate on ExportedProgram.
Expand All @@ -245,12 +305,41 @@ class ExportedProgramPassBase(ABC):
def __call__(self, exported_program: ExportedProgram) -> ExportedProgramPassResult:
"""
Runs the precondition check, the pass itself, and the postcondition check.

Prefer the node replacement APIs (``replace_all_uses_with``,
``replace_input_with``) in ``call``: a replace hook keeps the signature
valid as the graph changes. Output specs are realigned with the graph
afterwards regardless, so ``ensures`` sees a self-consistent program.
"""

self.requires(exported_program)
res = self.call(exported_program)
self.ensures(exported_program)
return res
graph_module = exported_program.graph_module
graph = graph_module.graph
hook_modified = False

def tracking_hook(old: torch.fx.Node, new: str, user: torch.fx.Node) -> None:
# Bound to the graph the hook was installed on: GraphModule.__deepcopy__
# carries _replace_hooks over to the copy, and rewriting the copy must
# not touch this program. The signature is read now rather than
# captured because a pass can swap _graph_signature for a fresh object
# while it runs.
if user.graph is not graph or user.op != "output" or old.name == new:
return
nonlocal hook_modified
hook_modified = True
_rename_output_specs(exported_program, old.name, new)

with graph_module._set_replace_hook(tracking_hook):
res = self.call(exported_program)
result_graph_module = res.exported_program.graph_module
if tracking_hook in result_graph_module._replace_hooks:
result_graph_module._unregister_replace_node_hook(tracking_hook)
signature_modified = _sync_output_specs(res.exported_program)
self.ensures(res.exported_program)
return ExportedProgramPassResult(
res.exported_program,
res.modified or hook_modified or signature_modified,
)

@abstractmethod
def call(self, exported_program: ExportedProgram) -> ExportedProgramPassResult:
Expand Down
151 changes: 151 additions & 0 deletions exir/tests/test_pass_infra.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

# pyre-strict

import copy
import unittest

import executorch.exir as exir
Expand Down Expand Up @@ -579,3 +580,153 @@ def placeholder(
new_input = self._find_input_node(new_graph_module)

self.assertNotEqual(self._symbolic_input_shape(new_input), original_snapshot)


class ExportedProgramPassBaseOutputSpecTest(unittest.TestCase):
"""__call__ realigns output specs with the graph before ensures() runs."""

class _Model(torch.nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return x + x

def _program(self) -> ExportedProgram:
return to_edge(export(self._Model(), (torch.randn(2, 2),))).exported_program()

def test_replacing_the_output_node_updates_the_signature(self) -> None:
class ReplaceOutputPass(ExportedProgramPassBase):
def call(self, ep: ExportedProgram) -> ExportedProgramPassResult:
graph = ep.graph_module.graph
output_node = graph.output_node()
(old,) = output_node.args[0]
with graph.inserting_before(output_node):
new = graph.call_function(
exir_ops.edge.aten.mul.Tensor, (old.args[0], old.args[1])
)
new.meta = dict(old.meta)
output_node.args = ((new,),)
return ExportedProgramPassResult(ep, True)

program = self._program()
result = ReplaceOutputPass()(program)

(spec,) = result.exported_program.graph_signature.output_specs
graph_output_name = result.exported_program.graph.output_node().args[0][0].name
self.assertEqual(spec.arg.name, graph_output_name)
result.exported_program.validate()

def test_a_signature_only_change_is_reported_as_modified(self) -> None:
"""A pass that renames the output reports modified even if it says False."""

class RenameOutputPass(ExportedProgramPassBase):
def call(self, ep: ExportedProgram) -> ExportedProgramPassResult:
ep.graph.output_node().args[0][0].name = "renamed_output"
return ExportedProgramPassResult(ep, False)

result = RenameOutputPass()(self._program())

self.assertTrue(result.modified)
(spec,) = result.exported_program.graph_signature.output_specs
self.assertEqual(spec.arg.name, "renamed_output")

def test_output_count_mismatch_is_rejected(self) -> None:
class DropOutputPass(ExportedProgramPassBase):
def call(self, ep: ExportedProgram) -> ExportedProgramPassResult:
output_node = ep.graph.output_node()
output_node.args = ((*output_node.args[0], output_node.args[0][0]),)
return ExportedProgramPassResult(ep, True)

with self.assertRaisesRegex(ExportPassBaseError, "output specs"):
DropOutputPass()(self._program())

def test_the_replace_hook_updates_the_signature_during_the_pass(self) -> None:
"""A pass using the replacement APIs sees a valid signature as it runs."""

signature_during_pass = []

class ReplaceViaApiPass(ExportedProgramPassBase):
def call(self, ep: ExportedProgram) -> ExportedProgramPassResult:
graph = ep.graph_module.graph
(old,) = graph.output_node().args[0]
with graph.inserting_before(graph.output_node()):
new = graph.call_function(
exir_ops.edge.aten.mul.Tensor, (old.args[0], old.args[1])
)
new.meta = dict(old.meta)
old.replace_all_uses_with(new)
signature_during_pass.append(
ep.graph_signature.output_specs[0].arg.name
)
return ExportedProgramPassResult(ep, True)

result = ReplaceViaApiPass()(self._program())

graph_output_name = result.exported_program.graph.output_node().args[0][0].name
self.assertEqual(signature_during_pass, [graph_output_name])

def test_replacing_a_returned_buffer_leaves_input_specs_alone(self) -> None:
"""The hook is output-only, so the replaced placeholder stays deletable."""

class TwoBuffers(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.register_buffer("a", torch.ones(2, 2))
self.register_buffer("b", torch.ones(2, 2))

def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
return self.a, x + self.b

replaced_name = []

class ReplaceReturnedBufferPass(ExportedProgramPassBase):
def call(self, ep: ExportedProgram) -> ExportedProgramPassResult:
names = {
target: name
for name, target in ep.graph_signature.inputs_to_buffers.items()
}
placeholders = {
node.name: node
for node in ep.graph.nodes
if node.op == "placeholder"
}
old = placeholders[names["a"]]
old.replace_all_uses_with(placeholders[names["b"]])
replaced_name.append(old.name)
return ExportedProgramPassResult(ep, True)

program = to_edge(export(TwoBuffers(), (torch.randn(2, 2),))).exported_program()

result = ReplaceReturnedBufferPass()(program)

signature = result.exported_program.graph_signature
self.assertIn(replaced_name[0], signature.inputs_to_buffers)
result.exported_program.validate()

def test_mutating_a_copied_graph_leaves_the_original_signature_alone(self) -> None:
"""GraphModule.__deepcopy__ carries the replace hook over to the copy."""

signature_after_copy_edit = []

class MutateACopyPass(ExportedProgramPassBase):
def call(self, ep: ExportedProgram) -> ExportedProgramPassResult:
graph_module = copy.deepcopy(ep.graph_module)
graph = graph_module.graph
(old,) = graph.output_node().args[0]
with graph.inserting_before(graph.output_node()):
new = graph.call_function(
exir_ops.edge.aten.mul.Tensor, (old.args[0], old.args[1])
)
new.meta = dict(old.meta)
old.replace_all_uses_with(new)
signature_after_copy_edit.append(
ep.graph_signature.output_specs[0].arg.name
)
return ExportedProgramPassResult(ep, False)

program = self._program()
original_output_name = program.graph_signature.output_specs[0].arg.name

result = MutateACopyPass()(program)

self.assertEqual(signature_after_copy_edit, [original_output_name])
self.assertFalse(result.modified)
program.validate()
Loading