Skip to content
Open
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
1 change: 1 addition & 0 deletions exir/passes/BUCK
Original file line number Diff line number Diff line change
Expand Up @@ -453,6 +453,7 @@ fbcode_target(_kind = runtime.python_library,
"//executorch/exir:memory",
"//executorch/exir:tensor",
"//executorch/exir/dialects:lib",
"//executorch/exir/operator:convert",
],
)

Expand Down
178 changes: 174 additions & 4 deletions exir/passes/replace_view_copy_with_view_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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]
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down
104 changes: 104 additions & 0 deletions exir/tests/test_remove_view_copy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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]))
Loading