Skip to content

Commit 674a7ea

Browse files
Prevent unsafe view-copy replacement around mutations (pytorch#23279)
### Summary `ReplaceViewCopyWithViewPass` currently runs after reinplacement. A reinplaced mutation of the copied view can therefore become a mutation of the base after `view_copy` is changed to `memory.view`, even though the original copy kept the base unchanged. The inverse interaction is possible when the base is mutated while the copied view remains live. Add a reusable `is_copy_to_view_safe` check that follows existing view aliases, schema-declared mutation aliases, and reinplace allocation-sharing annotations. It rejects a substitution only when a write to either prospective storage family precedes a read from the other family. Mutations after the other value's last read remain eligible for view replacement. The helper accepts the set of existing aliasing operators so other copy-to-view passes, including slice-copy lowering, can share the same safety analysis. Authored with Codex. ### Test plan - Added regression coverage for mutation through a view while its base remains live. - Added regression coverage for mutation of the base while the view remains live. - Added coverage showing both directions still replace the copy when the other value is dead. - Ran the focused `test_remove_view_copy` cases and the existing `test_replace_view_copy_with_view_pass` case. - Ran Black, Flake8, `compileall`, and `git diff --check` on the changed Python files.
1 parent c2c2932 commit 674a7ea

3 files changed

Lines changed: 279 additions & 4 deletions

File tree

‎exir/passes/BUCK‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -453,6 +453,7 @@ fbcode_target(_kind = runtime.python_library,
453453
"//executorch/exir:memory",
454454
"//executorch/exir:tensor",
455455
"//executorch/exir/dialects:lib",
456+
"//executorch/exir/operator:convert",
456457
],
457458
)
458459

‎exir/passes/replace_view_copy_with_view_pass.py‎

Lines changed: 174 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -9,12 +9,17 @@
99

1010
import copy
1111
import logging
12-
from typing import Any, List, Tuple
12+
import operator
13+
from typing import Any, Dict, FrozenSet, Iterable, List, Optional, Set, Tuple
1314

1415
import torch
1516
from executorch.exir import memory
1617

1718
from executorch.exir.dialects._ops import ops
19+
from executorch.exir.operator.convert import (
20+
output_to_aliased_input_map,
21+
unwrap_op_overload,
22+
)
1823
from executorch.exir.tensor import (
1924
contiguous_stride_from_shape,
2025
determine_tensor_dynanism,
@@ -37,6 +42,163 @@ def _is_view_copy(node: torch.fx.Node) -> bool:
3742
_VIEW_OP = memory.view
3843

3944

45+
def _schema(node: torch.fx.Node) -> Optional[torch.FunctionSchema]:
46+
if node.op != "call_function":
47+
return None
48+
try:
49+
return unwrap_op_overload(node.target)._schema
50+
except (AttributeError, TypeError):
51+
return None
52+
53+
54+
def _schema_arg(node: torch.fx.Node, schema: torch.FunctionSchema, index: int) -> Any:
55+
if index < len(node.args):
56+
return node.args[index]
57+
return node.kwargs.get(schema.arguments[index].name)
58+
59+
60+
def _mutated_inputs(node: torch.fx.Node) -> Set[torch.fx.Node]:
61+
mutated: Set[torch.fx.Node] = set()
62+
63+
share_idx = node.meta.get("_share_alloc_with_arg_idx")
64+
if isinstance(share_idx, int) and share_idx < len(node.args):
65+
arg = node.args[share_idx]
66+
if isinstance(arg, torch.fx.Node):
67+
mutated.add(arg)
68+
69+
schema = _schema(node)
70+
if schema is None:
71+
return mutated
72+
for index, argument in enumerate(schema.arguments):
73+
alias_info = argument.alias_info
74+
if alias_info is None or not alias_info.is_write:
75+
continue
76+
arg = _schema_arg(node, schema, index)
77+
if isinstance(arg, torch.fx.Node):
78+
mutated.add(arg)
79+
return mutated
80+
81+
82+
def _alias_source(
83+
node: torch.fx.Node, aliasing_ops: FrozenSet[Any]
84+
) -> Optional[torch.fx.Node]:
85+
if node.op != "call_function":
86+
return None
87+
88+
if node.target in aliasing_ops and node.args:
89+
base = node.args[0]
90+
return base if isinstance(base, torch.fx.Node) else None
91+
92+
share_idx = node.meta.get("_share_alloc_with_arg_idx")
93+
if isinstance(share_idx, int) and share_idx < len(node.args):
94+
base = node.args[share_idx]
95+
return base if isinstance(base, torch.fx.Node) else None
96+
97+
if node.target == operator.getitem and len(node.args) == 2:
98+
container, output_index = node.args
99+
if not isinstance(container, torch.fx.Node) or not isinstance(
100+
output_index, int
101+
):
102+
return None
103+
schema = _schema(container)
104+
if schema is None:
105+
return None
106+
input_index = output_to_aliased_input_map(schema).get(output_index)
107+
if input_index is None:
108+
return None
109+
base = _schema_arg(container, schema, input_index)
110+
return base if isinstance(base, torch.fx.Node) else None
111+
112+
schema = _schema(node)
113+
if schema is None or len(schema.returns) != 1:
114+
return None
115+
input_index = output_to_aliased_input_map(schema).get(0)
116+
if input_index is None:
117+
return None
118+
base = _schema_arg(node, schema, input_index)
119+
return base if isinstance(base, torch.fx.Node) else None
120+
121+
122+
def _alias_root(
123+
node: torch.fx.Node,
124+
aliasing_ops: FrozenSet[Any],
125+
roots: Dict[torch.fx.Node, torch.fx.Node],
126+
) -> torch.fx.Node:
127+
if node in roots:
128+
return roots[node]
129+
source = _alias_source(node, aliasing_ops)
130+
root = (
131+
node
132+
if source is None or source is node
133+
else _alias_root(source, aliasing_ops, roots)
134+
)
135+
roots[node] = root
136+
return root
137+
138+
139+
def _is_alias_only_node(node: torch.fx.Node, aliasing_ops: FrozenSet[Any]) -> bool:
140+
if node.op != "call_function":
141+
return False
142+
if node.target in aliasing_ops:
143+
return True
144+
return (
145+
node.target == operator.getitem
146+
and _alias_source(node, aliasing_ops) is not None
147+
)
148+
149+
150+
def is_copy_to_view_safe(
151+
node: torch.fx.Node,
152+
aliasing_ops: Optional[Iterable[Any]] = None,
153+
) -> bool:
154+
"""Return whether replacing a copy with an alias preserves mutation semantics.
155+
156+
The replacement merges the storage of ``node`` and its first argument. A
157+
mutation of either storage is safe only after the other storage's last read.
158+
Existing aliases and outputs of in-place operations are included in each
159+
storage group.
160+
"""
161+
if not node.args or not isinstance(node.args[0], torch.fx.Node):
162+
return False
163+
164+
aliases = (
165+
frozenset(aliasing_ops) if aliasing_ops is not None else frozenset({_VIEW_OP})
166+
)
167+
roots: Dict[torch.fx.Node, torch.fx.Node] = {}
168+
base_root = _alias_root(node.args[0], aliases, roots)
169+
copy_root = _alias_root(node, aliases, roots)
170+
if base_root is copy_root:
171+
return True
172+
173+
nodes = list(node.graph.nodes)
174+
copy_index = nodes.index(node)
175+
last_base_read = copy_index
176+
last_copy_read = copy_index
177+
178+
for index, current in enumerate(nodes):
179+
if _is_alias_only_node(current, aliases):
180+
continue
181+
input_roots = {
182+
_alias_root(input_node, aliases, roots)
183+
for input_node in current.all_input_nodes
184+
}
185+
if base_root in input_roots:
186+
last_base_read = index
187+
if copy_root in input_roots:
188+
last_copy_read = index
189+
190+
for index, current in enumerate(nodes[copy_index + 1 :], copy_index + 1):
191+
mutated_roots = {
192+
_alias_root(input_node, aliases, roots)
193+
for input_node in _mutated_inputs(current)
194+
}
195+
if copy_root in mutated_roots and index <= last_base_read:
196+
return False
197+
if base_root in mutated_roots and index <= last_copy_read:
198+
return False
199+
return True
200+
201+
40202
class _Guard:
41203
def __init__(
42204
self, name: str, field_lambda, expected_val: Any # pyre-ignore[2]
@@ -278,10 +440,16 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult:
278440
for module in graph_module.modules():
279441
if not isinstance(module, torch.fx.GraphModule):
280442
continue
281-
for node in module.graph.nodes:
443+
# Process consumers before producers so nested view copies are
444+
# analyzed with their final aliasing behavior.
445+
for node in reversed(module.graph.nodes):
282446
# Note: We only replace view_copy nodes that are not output, since
283447
# the output pointer could be modified at runtime (T187925929)
284-
if _is_view_copy(node) and all(u.op != "output" for u in node.users):
448+
if (
449+
_is_view_copy(node)
450+
and all(u.op != "output" for u in node.users)
451+
and is_copy_to_view_safe(node)
452+
):
285453
base, _ = node.args
286454
node.target = _VIEW_OP
287455

@@ -309,7 +477,9 @@ def ensures(self, graph_module: torch.fx.GraphModule) -> None:
309477
# Note: We only replace view_copy nodes that are not output, since
310478
# the output pointer could be modified at runtime (T187925929)
311479
assert not (
312-
_is_view_copy(node) and all(u.op != "output" for u in node.users)
480+
_is_view_copy(node)
481+
and all(u.op != "output" for u in node.users)
482+
and is_copy_to_view_safe(node)
313483
)
314484
if node.op == "call_function" and node.target == _VIEW_OP:
315485
assert isinstance(node.meta["spec"], _ViewSpec)

‎exir/tests/test_remove_view_copy.py‎

Lines changed: 104 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,14 @@
1212
from executorch.exir import memory, to_edge
1313
from executorch.exir.capture._config import ExecutorchBackendConfig
1414
from executorch.exir.passes import MemoryPlanningPass
15+
from executorch.exir.passes.normalize_view_copy_base_pass import (
16+
NormalizeViewCopyBasePass,
17+
)
18+
from executorch.exir.passes.reinplace import reinplace_pass
19+
from executorch.exir.passes.replace_view_copy_with_view_pass import (
20+
ReplaceViewCopyWithViewPass,
21+
)
22+
from executorch.exir.passes.spec_prop_pass import SpecPropPass
1523

1624

1725
class TestModel1(nn.Module):
@@ -42,6 +50,18 @@ def get_example_inputs(self):
4250

4351

4452
class TestRemoveViewCopy(unittest.TestCase):
53+
def _run_view_and_reinplace_passes(
54+
self, model: nn.Module, example_inputs: tuple
55+
) -> torch.fx.GraphModule:
56+
ep = to_edge(
57+
torch.export.export(model.eval(), example_inputs, strict=True)
58+
).exported_program()
59+
reinplace_pass(ep)
60+
graph_module = SpecPropPass()(ep.graph_module).graph_module
61+
NormalizeViewCopyBasePass()(graph_module)
62+
ReplaceViewCopyWithViewPass()(graph_module)
63+
return graph_module
64+
4565
def test_disable(self) -> None:
4666
model = TestModel1()
4767
model.eval()
@@ -234,3 +254,87 @@ def forward(self, x):
234254
plan = etpm.executorch_program.execution_plan[0]
235255
op_names = [op.name for op in plan.operators]
236256
self.assertTrue("executorch_prim::et_view" in op_names)
257+
258+
def test_mutated_view_with_live_base_is_not_replaced(self) -> None:
259+
class TestModel(nn.Module):
260+
def forward(self, x, indices, values):
261+
base = torch.relu(x)
262+
viewed = base.view(4, 3)
263+
changed = torch.ops.aten.index_put.default(viewed, [indices], values)
264+
return base, changed
265+
266+
inputs = (
267+
torch.arange(12, dtype=torch.float32).reshape(4, 3),
268+
torch.tensor([0]),
269+
torch.tensor([[100.0, 101.0, 102.0]]),
270+
)
271+
expected = TestModel()(*copy.deepcopy(inputs))
272+
graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs)
273+
actual = graph_module(*copy.deepcopy(inputs))
274+
275+
self.assertFalse(any(n.target == memory.view for n in graph_module.graph.nodes))
276+
self.assertTrue(torch.equal(expected[0], actual[0]))
277+
self.assertTrue(torch.equal(expected[1], actual[1]))
278+
279+
def test_mutated_view_with_dead_base_is_replaced(self) -> None:
280+
class TestModel(nn.Module):
281+
def forward(self, x, indices, values):
282+
base = torch.relu(x)
283+
viewed = base.view(4, 3)
284+
return torch.ops.aten.index_put.default(viewed, [indices], values)
285+
286+
inputs = (
287+
torch.arange(12, dtype=torch.float32).reshape(4, 3),
288+
torch.tensor([0]),
289+
torch.tensor([[100.0, 101.0, 102.0]]),
290+
)
291+
expected = TestModel()(*copy.deepcopy(inputs))
292+
graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs)
293+
actual = graph_module(*copy.deepcopy(inputs))
294+
295+
self.assertTrue(any(n.target == memory.view for n in graph_module.graph.nodes))
296+
self.assertTrue(torch.equal(expected, actual[0]))
297+
298+
def test_base_mutation_after_last_view_read_allows_replacement(self) -> None:
299+
class TestModel(nn.Module):
300+
def forward(self, x, indices, values):
301+
base = torch.relu(x)
302+
viewed = base.view(4, 3)
303+
observed = viewed.clone()
304+
changed = torch.ops.aten.index_put.default(base, [indices], values)
305+
return observed, changed
306+
307+
inputs = (
308+
torch.arange(12, dtype=torch.float32).reshape(4, 3),
309+
torch.tensor([0]),
310+
torch.tensor([[100.0, 101.0, 102.0]]),
311+
)
312+
expected = TestModel()(*copy.deepcopy(inputs))
313+
graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs)
314+
actual = graph_module(*copy.deepcopy(inputs))
315+
316+
self.assertTrue(any(n.target == memory.view for n in graph_module.graph.nodes))
317+
self.assertTrue(torch.equal(expected[0], actual[0]))
318+
self.assertTrue(torch.equal(expected[1], actual[1]))
319+
320+
def test_base_mutation_before_view_read_prevents_replacement(self) -> None:
321+
class TestModel(nn.Module):
322+
def forward(self, x, indices, values):
323+
base = torch.relu(x)
324+
viewed = base.view(4, 3)
325+
changed = torch.ops.aten.index_put.default(base, [indices], values)
326+
observed = viewed.clone()
327+
return changed, observed
328+
329+
inputs = (
330+
torch.arange(12, dtype=torch.float32).reshape(4, 3),
331+
torch.tensor([0]),
332+
torch.tensor([[100.0, 101.0, 102.0]]),
333+
)
334+
expected = TestModel()(*copy.deepcopy(inputs))
335+
graph_module = self._run_view_and_reinplace_passes(TestModel(), inputs)
336+
actual = graph_module(*copy.deepcopy(inputs))
337+
338+
self.assertFalse(any(n.target == memory.view for n in graph_module.graph.nodes))
339+
self.assertTrue(torch.equal(expected[0], actual[0]))
340+
self.assertTrue(torch.equal(expected[1], actual[1]))

0 commit comments

Comments
 (0)