From 867406862b142b43449ce457b32b256c9ad17eb9 Mon Sep 17 00:00:00 2001 From: Jacob Szwejbka Date: Wed, 30 Sep 2026 11:10:42 -0700 Subject: [PATCH] Guard view-copy replacement across mutations Replace view_copy with a storage alias only when writes to the prospective base and view storage cannot be observed through the other value. Preserve the optimization when the other value is already dead. Authored with Codex. --- exir/passes/BUCK | 1 + .../replace_view_copy_with_view_pass.py | 178 +++++++++++++++++- exir/tests/test_remove_view_copy.py | 104 ++++++++++ 3 files changed, 279 insertions(+), 4 deletions(-) diff --git a/exir/passes/BUCK b/exir/passes/BUCK index 6e47151614f..0f1ea2a7cdd 100644 --- a/exir/passes/BUCK +++ b/exir/passes/BUCK @@ -453,6 +453,7 @@ fbcode_target(_kind = runtime.python_library, "//executorch/exir:memory", "//executorch/exir:tensor", "//executorch/exir/dialects:lib", + "//executorch/exir/operator:convert", ], ) diff --git a/exir/passes/replace_view_copy_with_view_pass.py b/exir/passes/replace_view_copy_with_view_pass.py index 947952d7692..166ad699b15 100644 --- a/exir/passes/replace_view_copy_with_view_pass.py +++ b/exir/passes/replace_view_copy_with_view_pass.py @@ -9,12 +9,17 @@ import copy import logging -from typing import Any, List, Tuple +import operator +from typing import Any, Dict, FrozenSet, Iterable, List, Optional, Set, Tuple import torch from executorch.exir import memory from executorch.exir.dialects._ops import ops +from executorch.exir.operator.convert import ( + output_to_aliased_input_map, + unwrap_op_overload, +) from executorch.exir.tensor import ( contiguous_stride_from_shape, determine_tensor_dynanism, @@ -37,6 +42,163 @@ def _is_view_copy(node: torch.fx.Node) -> bool: _VIEW_OP = memory.view +def _schema(node: torch.fx.Node) -> Optional[torch.FunctionSchema]: + if node.op != "call_function": + return None + try: + return unwrap_op_overload(node.target)._schema + except (AttributeError, TypeError): + return None + + +def _schema_arg(node: torch.fx.Node, schema: torch.FunctionSchema, index: int) -> Any: + if index < len(node.args): + return node.args[index] + return node.kwargs.get(schema.arguments[index].name) + + +def _mutated_inputs(node: torch.fx.Node) -> Set[torch.fx.Node]: + mutated: Set[torch.fx.Node] = set() + + share_idx = node.meta.get("_share_alloc_with_arg_idx") + if isinstance(share_idx, int) and share_idx < len(node.args): + arg = node.args[share_idx] + if isinstance(arg, torch.fx.Node): + mutated.add(arg) + + schema = _schema(node) + if schema is None: + return mutated + for index, argument in enumerate(schema.arguments): + alias_info = argument.alias_info + if alias_info is None or not alias_info.is_write: + continue + arg = _schema_arg(node, schema, index) + if isinstance(arg, torch.fx.Node): + mutated.add(arg) + return mutated + + +def _alias_source( + node: torch.fx.Node, aliasing_ops: FrozenSet[Any] +) -> Optional[torch.fx.Node]: + if node.op != "call_function": + return None + + if node.target in aliasing_ops and node.args: + base = node.args[0] + return base if isinstance(base, torch.fx.Node) else None + + share_idx = node.meta.get("_share_alloc_with_arg_idx") + if isinstance(share_idx, int) and share_idx < len(node.args): + base = node.args[share_idx] + return base if isinstance(base, torch.fx.Node) else None + + if node.target == operator.getitem and len(node.args) == 2: + container, output_index = node.args + if not isinstance(container, torch.fx.Node) or not isinstance( + output_index, int + ): + return None + schema = _schema(container) + if schema is None: + return None + input_index = output_to_aliased_input_map(schema).get(output_index) + if input_index is None: + return None + base = _schema_arg(container, schema, input_index) + return base if isinstance(base, torch.fx.Node) else None + + schema = _schema(node) + if schema is None or len(schema.returns) != 1: + return None + input_index = output_to_aliased_input_map(schema).get(0) + if input_index is None: + return None + base = _schema_arg(node, schema, input_index) + return base if isinstance(base, torch.fx.Node) else None + + +def _alias_root( + node: torch.fx.Node, + aliasing_ops: FrozenSet[Any], + roots: Dict[torch.fx.Node, torch.fx.Node], +) -> torch.fx.Node: + if node in roots: + return roots[node] + source = _alias_source(node, aliasing_ops) + root = ( + node + if source is None or source is node + else _alias_root(source, aliasing_ops, roots) + ) + roots[node] = root + return root + + +def _is_alias_only_node(node: torch.fx.Node, aliasing_ops: FrozenSet[Any]) -> bool: + if node.op != "call_function": + return False + if node.target in aliasing_ops: + return True + return ( + node.target == operator.getitem + and _alias_source(node, aliasing_ops) is not None + ) + + +def is_copy_to_view_safe( + node: torch.fx.Node, + aliasing_ops: Optional[Iterable[Any]] = None, +) -> bool: + """Return whether replacing a copy with an alias preserves mutation semantics. + + The replacement merges the storage of ``node`` and its first argument. A + mutation of either storage is safe only after the other storage's last read. + Existing aliases and outputs of in-place operations are included in each + storage group. + """ + if not node.args or not isinstance(node.args[0], torch.fx.Node): + return False + + aliases = ( + frozenset(aliasing_ops) if aliasing_ops is not None else frozenset({_VIEW_OP}) + ) + roots: Dict[torch.fx.Node, torch.fx.Node] = {} + base_root = _alias_root(node.args[0], aliases, roots) + copy_root = _alias_root(node, aliases, roots) + if base_root is copy_root: + return True + + nodes = list(node.graph.nodes) + copy_index = nodes.index(node) + last_base_read = copy_index + last_copy_read = copy_index + + for index, current in enumerate(nodes): + if _is_alias_only_node(current, aliases): + continue + input_roots = { + _alias_root(input_node, aliases, roots) + for input_node in current.all_input_nodes + } + if base_root in input_roots: + last_base_read = index + if copy_root in input_roots: + last_copy_read = index + + for index, current in enumerate(nodes[copy_index + 1 :], copy_index + 1): + mutated_roots = { + _alias_root(input_node, aliases, roots) + for input_node in _mutated_inputs(current) + } + if copy_root in mutated_roots and index <= last_base_read: + return False + if base_root in mutated_roots and index <= last_copy_read: + return False + return True + + class _Guard: def __init__( self, name: str, field_lambda, expected_val: Any # pyre-ignore[2] @@ -278,10 +440,16 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult: for module in graph_module.modules(): if not isinstance(module, torch.fx.GraphModule): continue - for node in module.graph.nodes: + # Process consumers before producers so nested view copies are + # analyzed with their final aliasing behavior. + for node in reversed(module.graph.nodes): # Note: We only replace view_copy nodes that are not output, since # the output pointer could be modified at runtime (T187925929) - if _is_view_copy(node) and all(u.op != "output" for u in node.users): + if ( + _is_view_copy(node) + and all(u.op != "output" for u in node.users) + and is_copy_to_view_safe(node) + ): base, _ = node.args node.target = _VIEW_OP @@ -309,7 +477,9 @@ def ensures(self, graph_module: torch.fx.GraphModule) -> None: # Note: We only replace view_copy nodes that are not output, since # the output pointer could be modified at runtime (T187925929) assert not ( - _is_view_copy(node) and all(u.op != "output" for u in node.users) + _is_view_copy(node) + and all(u.op != "output" for u in node.users) + and is_copy_to_view_safe(node) ) if node.op == "call_function" and node.target == _VIEW_OP: assert isinstance(node.meta["spec"], _ViewSpec) diff --git a/exir/tests/test_remove_view_copy.py b/exir/tests/test_remove_view_copy.py index bea7e3ff83c..fe2558dc34c 100644 --- a/exir/tests/test_remove_view_copy.py +++ b/exir/tests/test_remove_view_copy.py @@ -12,6 +12,14 @@ from executorch.exir import memory, to_edge from executorch.exir.capture._config import ExecutorchBackendConfig from executorch.exir.passes import MemoryPlanningPass +from executorch.exir.passes.normalize_view_copy_base_pass import ( + NormalizeViewCopyBasePass, +) +from executorch.exir.passes.reinplace import reinplace_pass +from executorch.exir.passes.replace_view_copy_with_view_pass import ( + ReplaceViewCopyWithViewPass, +) +from executorch.exir.passes.spec_prop_pass import SpecPropPass class TestModel1(nn.Module): @@ -42,6 +50,18 @@ def get_example_inputs(self): class TestRemoveViewCopy(unittest.TestCase): + def _run_view_and_reinplace_passes( + self, model: nn.Module, example_inputs: tuple + ) -> torch.fx.GraphModule: + ep = to_edge( + torch.export.export(model.eval(), example_inputs, strict=True) + ).exported_program() + reinplace_pass(ep) + graph_module = SpecPropPass()(ep.graph_module).graph_module + NormalizeViewCopyBasePass()(graph_module) + ReplaceViewCopyWithViewPass()(graph_module) + return graph_module + def test_disable(self) -> None: model = TestModel1() model.eval() @@ -234,3 +254,87 @@ def forward(self, x): plan = etpm.executorch_program.execution_plan[0] op_names = [op.name for op in plan.operators] self.assertTrue("executorch_prim::et_view" in op_names) + + def test_mutated_view_with_live_base_is_not_replaced(self) -> None: + class TestModel(nn.Module): + def forward(self, x, indices, values): + base = torch.relu(x) + viewed = base.view(4, 3) + changed = torch.ops.aten.index_put.default(viewed, [indices], values) + return base, changed + + inputs = ( + torch.arange(12, dtype=torch.float32).reshape(4, 3), + torch.tensor([0]), + torch.tensor([[100.0, 101.0, 102.0]]), + ) + expected = TestModel()(*copy.deepcopy(inputs)) + graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs) + actual = graph_module(*copy.deepcopy(inputs)) + + self.assertFalse(any(n.target == memory.view for n in graph_module.graph.nodes)) + self.assertTrue(torch.equal(expected[0], actual[0])) + self.assertTrue(torch.equal(expected[1], actual[1])) + + def test_mutated_view_with_dead_base_is_replaced(self) -> None: + class TestModel(nn.Module): + def forward(self, x, indices, values): + base = torch.relu(x) + viewed = base.view(4, 3) + return torch.ops.aten.index_put.default(viewed, [indices], values) + + inputs = ( + torch.arange(12, dtype=torch.float32).reshape(4, 3), + torch.tensor([0]), + torch.tensor([[100.0, 101.0, 102.0]]), + ) + expected = TestModel()(*copy.deepcopy(inputs)) + graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs) + actual = graph_module(*copy.deepcopy(inputs)) + + self.assertTrue(any(n.target == memory.view for n in graph_module.graph.nodes)) + self.assertTrue(torch.equal(expected, actual[0])) + + def test_base_mutation_after_last_view_read_allows_replacement(self) -> None: + class TestModel(nn.Module): + def forward(self, x, indices, values): + base = torch.relu(x) + viewed = base.view(4, 3) + observed = viewed.clone() + changed = torch.ops.aten.index_put.default(base, [indices], values) + return observed, changed + + inputs = ( + torch.arange(12, dtype=torch.float32).reshape(4, 3), + torch.tensor([0]), + torch.tensor([[100.0, 101.0, 102.0]]), + ) + expected = TestModel()(*copy.deepcopy(inputs)) + graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs) + actual = graph_module(*copy.deepcopy(inputs)) + + self.assertTrue(any(n.target == memory.view for n in graph_module.graph.nodes)) + self.assertTrue(torch.equal(expected[0], actual[0])) + self.assertTrue(torch.equal(expected[1], actual[1])) + + def test_base_mutation_before_view_read_prevents_replacement(self) -> None: + class TestModel(nn.Module): + def forward(self, x, indices, values): + base = torch.relu(x) + viewed = base.view(4, 3) + changed = torch.ops.aten.index_put.default(base, [indices], values) + observed = viewed.clone() + return changed, observed + + inputs = ( + torch.arange(12, dtype=torch.float32).reshape(4, 3), + torch.tensor([0]), + torch.tensor([[100.0, 101.0, 102.0]]), + ) + expected = TestModel()(*copy.deepcopy(inputs)) + graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs) + actual = graph_module(*copy.deepcopy(inputs)) + + self.assertFalse(any(n.target == memory.view for n in graph_module.graph.nodes)) + self.assertTrue(torch.equal(expected[0], actual[0])) + self.assertTrue(torch.equal(expected[1], actual[1]))