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]))