diff --git a/exir/emit/_emitter.py b/exir/emit/_emitter.py index ec0d97df4b0..e56b4a1662f 100644 --- a/exir/emit/_emitter.py +++ b/exir/emit/_emitter.py @@ -1300,6 +1300,20 @@ def _emit_view(self, args: Tuple[_Argument, ...]) -> _EmitterValue: self.chain.instructions.append(kernel) return out_arg + def _emit_slice(self, args: Tuple[_Argument, ...]) -> _EmitterValue: + """Emit a statically memory-planned slice as a sub-buffer alias. + + ``ReplaceSliceCopyWithSlicePass`` creates ``_ViewSpec`` values whose + ``mem_offset`` is the byte offset into their base allocation. No kernel + is needed for a static planned slice: the tensor value can be emitted + directly from that specification. + """ + assert 4 <= len(args) <= 5 + spec = self.node.meta["spec"] + assert spec.is_static_shape_tensor + assert spec.mem_id is not None and spec.mem_offset is not None + return self._emit_spec(spec) + def _add_debug_handle( self, emitter_id: int, @@ -1829,6 +1843,9 @@ def call_function( # pyre-fixme[14] elif target == memory.view: return self._emit_view(args) + elif target == memory.slice: + return self._emit_slice(args) + elif target == memory.free: assert len(args) == 1 # pyre-ignore diff --git a/exir/memory.py b/exir/memory.py index 36a244bc02f..8295218cae3 100644 --- a/exir/memory.py +++ b/exir/memory.py @@ -6,7 +6,7 @@ # pyre-strict -from typing import List, Tuple, Union +from typing import List, Optional, Tuple, Union import torch from executorch.exir.sym_util import eval_shape @@ -48,3 +48,16 @@ def view(base: torch.Tensor, size: List[int]) -> torch.Tensor: It is used to elide view_copy nodes. """ return base.view(size) + + +def slice( # noqa: A001 + base: torch.Tensor, + dim: int = 0, + start: Optional[int] = None, + end: Optional[int] = None, + step: int = 1, +) -> torch.Tensor: + """ + Mimics ``aten.slice.Tensor`` for eliding contiguous ``slice_copy`` nodes. + """ + return torch.ops.aten.slice.Tensor(base, dim, start, end, step) diff --git a/exir/memory_planning.py b/exir/memory_planning.py index 44b53585455..5ac02ab42d3 100644 --- a/exir/memory_planning.py +++ b/exir/memory_planning.py @@ -647,6 +647,7 @@ def collect_specs_from_nodes( # noqa: C901 in [ memory.alloc, memory.view, + memory.slice, operator.getitem, torch.ops.higher_order.cond, exir_while, @@ -910,10 +911,10 @@ def get_node_tensor_specs( has no tensor specs. """ # get tensor specs - if node.target == memory.view: + if node.target in (memory.view, memory.slice): base = node.args[0] assert isinstance(base, torch.fx.Node) - specs = base.meta.get("spec") + return get_node_tensor_specs(base) else: specs = node.meta.get("spec") diff --git a/exir/pass_base.py b/exir/pass_base.py index e7c2caa66a9..45c6f9500ac 100644 --- a/exir/pass_base.py +++ b/exir/pass_base.py @@ -1007,7 +1007,7 @@ def call_function( # TODO according to zhengxu ExportPassBase should not be aware of # memory.alloc. Check this comment: # https://www.internalfb.com/diff/D42758019?dst_version_fbid=5906016402813292&transaction_fbid=1104713900200176 - elif target == memory.alloc: + elif target in (memory.alloc, memory.slice): return self.callback._fx( "call_function", target, diff --git a/exir/passes/BUCK b/exir/passes/BUCK index 0f1ea2a7cdd..134f6360f3d 100644 --- a/exir/passes/BUCK +++ b/exir/passes/BUCK @@ -29,6 +29,7 @@ fbcode_target(_kind = runtime.python_library, ":replace_aten_with_edge_pass", ":replace_broken_ops_with_function_ops_pass", ":replace_edge_with_backend_pass", + ":replace_slice_copy_with_slice_pass", ":replace_sym_size_op_pass", ":scalar_to_tensor_pass", ":spec_prop_pass", @@ -443,6 +444,20 @@ fbcode_target(_kind = runtime.python_library, ], ) +fbcode_target(_kind = runtime.python_library, + name = "replace_slice_copy_with_slice_pass", + srcs = [ + "replace_slice_copy_with_slice_pass.py", + ], + deps = [ + ":replace_view_copy_with_view_pass", + "//caffe2:torch", + "//executorch/exir:memory", + "//executorch/exir:tensor", + "//executorch/exir/dialects:lib", + ], +) + fbcode_target(_kind = runtime.python_library, name = "replace_view_copy_with_view_pass", srcs = [ diff --git a/exir/passes/__init__.py b/exir/passes/__init__.py index 57c4c313112..fe48bfc6446 100644 --- a/exir/passes/__init__.py +++ b/exir/passes/__init__.py @@ -269,6 +269,7 @@ def callWithLoggerEnabled(self, graph_module: torch.fx.GraphModule) -> None: # it's retraced after running to_out_variant with the first trace. memory.alloc, memory.view, + memory.slice, executorch_call_delegate, } to_out_var_skiplist.update(_EXECUTORCH_SYM_OPS) diff --git a/exir/passes/replace_slice_copy_with_slice_pass.py b/exir/passes/replace_slice_copy_with_slice_pass.py new file mode 100644 index 00000000000..bfe419e9255 --- /dev/null +++ b/exir/passes/replace_slice_copy_with_slice_pass.py @@ -0,0 +1,155 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +# pyre-strict + +"""Replace safe, static, contiguous slice copies with sub-buffer aliases.""" + +import logging +from typing import Any, Optional + +import torch +from executorch.exir import memory +from executorch.exir.dialects._ops import ops +from executorch.exir.passes.replace_view_copy_with_view_pass import ( + _ViewSpec, + is_copy_to_view_safe, +) +from executorch.exir.tensor import ( + contiguous_stride_from_shape, + dim_order_from_stride, + TensorSpec, +) +from torch.fx.passes.infra.pass_base import PassBase, PassResult + +logger: logging.Logger = logging.getLogger(__name__) + + +def _is_slice_copy(node: torch.fx.Node) -> bool: + return node.op == "call_function" and node.target in ( + torch.ops.aten.slice_copy.Tensor, + ops.edge.aten.slice_copy.Tensor, + ) + + +def _is_static_slice_argument(value: Any) -> bool: + return value is None or isinstance(value, int) + + +def is_contiguous_slice_copy(node: torch.fx.Node) -> bool: + """True for static-argument, outermost-dimension, unit-step slices.""" + if ( + not _is_slice_copy(node) + or node.kwargs + or not all(_is_static_slice_argument(arg) for arg in node.args[1:]) + ): + return False + dim = node.args[1] if len(node.args) > 1 else 0 + step = node.args[4] if len(node.args) > 4 else 1 + base = node.args[0] + val = base.meta.get("val") if isinstance(base, torch.fx.Node) else None + if not isinstance(dim, int) or step != 1: + return False + if dim < 0 and isinstance(val, torch.Tensor): + dim += val.dim() + return dim == 0 + + +def _compute_slice_byte_offset(base: TensorSpec, dim: int, start: Optional[int]) -> int: + start, _, _ = slice(start, None, 1).indices(base.shape[dim]) + return start * base.stride[dim] * torch._utils._element_size(base.dtype) + + +def _static_slice_bounds(node: torch.fx.Node) -> Optional[tuple[int, int]]: + base = node.args[0] + if not isinstance(base, torch.fx.Node) or base.op == "placeholder": + return None + # Retain copies of aliases, including pending slice candidates. + if _is_slice_copy(base) or base.target in (memory.slice, memory.view): + return None + base_spec = base.meta.get("spec") + output_spec = node.meta.get("spec") + if ( + not isinstance(base_spec, TensorSpec) + or not isinstance(output_spec, TensorSpec) + or not base_spec.is_static_shape_tensor + or not output_spec.is_static_shape_tensor + or output_spec.dtype != base_spec.dtype + or base_spec.const + or base_spec.is_sparse + or base_spec.layout != torch.strided + or base_spec.dim_order + != dim_order_from_stride( + contiguous_stride_from_shape(torch.Size(base_spec.shape)) + ) + ): + return None + start = node.args[2] if len(node.args) > 2 else None + end = node.args[3] if len(node.args) > 3 else None + start, end, _ = slice(start, end, 1).indices(base_spec.shape[0]) + shape = [max(end - start, 0), *base_spec.shape[1:]] + # Empty results retain their kernel; do not emit an alias at + # the end of (or into an empty) allocation. + if 0 in shape or list(output_spec.shape) != shape: + return None + return start, end + + +class ReplaceSliceCopyWithSlicePass(PassBase): + """Replace eligible slice copies after view-copy replacement.""" + + def call(self, graph_module: torch.fx.GraphModule) -> PassResult: + n_replaced = 0 + for module in graph_module.modules(): + if not isinstance(module, torch.fx.GraphModule): + continue + replacements = {} + # Analyze consumers first, with already-decided aliasing behavior. + # Specs are rebuilt in forward order below so views of slices use + # the final base spec, including its byte offset. + for node in reversed(module.graph.nodes): + if not is_contiguous_slice_copy(node) or any( + u.op == "output" for u in node.users + ): + continue + bounds = _static_slice_bounds(node) + if bounds is None: + continue + start, end = bounds + base = node.args[0] + base_spec = base.meta["spec"] + if not is_copy_to_view_safe(node, (memory.view, memory.slice)): + continue + replacements[node] = _compute_slice_byte_offset(base_spec, 0, start) + node.target = memory.slice + node.args = (base, 0, start, end, 1) + n_replaced += 1 + + updated_specs = set() + for node in module.graph.nodes: + if node in replacements: + node.meta["spec"] = _ViewSpec( + node.args[0].meta["spec"], + list(node.meta["spec"].shape), + byte_offset=replacements[node], + ) + updated_specs.add(node) + elif node.target == memory.view and node.args[0] in updated_specs: + node.meta["spec"] = _ViewSpec( + node.args[0].meta["spec"], list(node.meta["spec"].shape) + ) + updated_specs.add(node) + module.recompile() + + logger.debug("Replaced %d slice_copy nodes with memory.slice", n_replaced) + return PassResult(graph_module, n_replaced > 0) + + def ensures(self, graph_module: torch.fx.GraphModule) -> None: + for module in graph_module.modules(): + if isinstance(module, torch.fx.GraphModule): + for node in module.graph.nodes: + if node.target == memory.slice: + assert isinstance(node.meta["spec"], _ViewSpec) diff --git a/exir/passes/replace_view_copy_with_view_pass.py b/exir/passes/replace_view_copy_with_view_pass.py index 166ad699b15..651a3d6a9fc 100644 --- a/exir/passes/replace_view_copy_with_view_pass.py +++ b/exir/passes/replace_view_copy_with_view_pass.py @@ -217,7 +217,9 @@ def __call__(self, view_spec) -> None: # pyre-ignore[2] class _ViewSpec(TensorSpec): - def __init__(self, base: TensorSpec, shape: List[int]) -> None: + def __init__( + self, base: TensorSpec, shape: List[int], *, byte_offset: int = 0 + ) -> None: """ A _ViewSpec is TensorSpec that shares non-size related fields with its base. The size-related fields are: shape, stride, dim_order, and shape_dynamism. @@ -228,7 +230,8 @@ def __init__(self, base: TensorSpec, shape: List[int]) -> None: A _ViewSpec can only be created from a non-sparse, strided TensorSpec. On creation, a _ViewSpec must be compatible with its base with respect to - shape_dynamism, dtype, and nbytes. + shape_dynamism and dtype. Its byte range must fit within the base. + byte_offset is relative to the base's memory-planned offset. A _ViewSpec contains _guards that are evaluated on every __getattribute__ call. The purpose of the guards is to make sure the _ViewSpec is still compatible @@ -239,6 +242,7 @@ def __init__(self, base: TensorSpec, shape: List[int]) -> None: # Any attribute that is not in _self_fields or _base_fields will # raise an Exception. If TensorSpec is extended with a new attribute, # we should explicitly decide how _ViewSpec will handle it. + self._byte_offset = byte_offset self._self_fields = [ # We need to get the debug method from self # so that the object id it prints is correct. @@ -348,11 +352,20 @@ def __init__(self, base: TensorSpec, shape: List[int]) -> None: _Guard("dtype", lambda view_spec: view_spec.dtype, base.dtype) ) - # We do not guard nbytes because dynamic symints are replaced by upper bounds. - # We do guard on rank, though - if self.nbytes() != base.nbytes(): + # Dynamic symints are replaced by upper bounds, so only retain a + # byte-range guard for static specs. + if byte_offset < 0 or byte_offset + self.nbytes() > base.nbytes(): raise Exception( - f"_ViewSpec is incompatible with its base on creation. It has nbytes={self.nbytes()}, but its base has nbytes={base.nbytes()}." + f"_ViewSpec byte range ({byte_offset}, {byte_offset + self.nbytes()}) exceeds base size {base.nbytes()}." + ) + if self.is_static_shape_tensor: + self._guards.append( + _Guard( + "byte_range", + lambda view_spec: view_spec._byte_offset + view_spec.nbytes() + <= view_spec._base.nbytes(), + True, + ) ) self._guards.append( _Guard("rank", lambda view_spec: len(view_spec.shape), len(shape)) @@ -376,6 +389,7 @@ def __getattribute__(self, name: str): # pyre-ignore "_guards", "_unguarded_access", "_run_guards", + "_byte_offset", ]: return object.__getattribute__(self, name) @@ -383,7 +397,9 @@ def __getattribute__(self, name: str): # pyre-ignore if name in self._self_fields: val = object.__getattribute__(self, name) elif name in self._base_fields: - val = object.__getattribute__(self._base, name) + val = self._base.__getattribute__(name) + if name == "mem_offset" and val is not None: + val += self._byte_offset else: if len(name) > 0 and name[0] != "_": logger.warning( @@ -404,6 +420,7 @@ def __setattr__(self, name: str, val) -> None: # pyre-ignore "_guards", "_unguarded_access", "_run_guards", + "_byte_offset", ]: object.__setattr__(self, name, val) return @@ -413,7 +430,11 @@ def __setattr__(self, name: str, val) -> None: # pyre-ignore return if name in self._base_fields: - object.__setattr__(self._base, name, val) + if name == "mem_offset" and self._byte_offset: + raise ValueError( + "Set the base allocation offset, not a sub-view offset." + ) + self._base.__setattr__(name, val) return if len(name) > 0 and name[0] != "_": diff --git a/exir/program/BUCK b/exir/program/BUCK index 8e7b59e0ba0..abdd4d6069d 100644 --- a/exir/program/BUCK +++ b/exir/program/BUCK @@ -46,6 +46,7 @@ fbcode_target(_kind = runtime.python_library, "//executorch/exir/passes:remove_graph_asserts_pass", "//executorch/exir/passes:remove_mixed_type_operators", "//executorch/exir/passes:replace_aten_with_edge_pass", + "//executorch/exir/passes:replace_slice_copy_with_slice_pass", "//executorch/exir/passes:replace_view_copy_with_view_pass", "//executorch/exir/passes:spec_prop_pass", "//executorch/exir/passes:weights_to_outputs_pass", diff --git a/exir/program/_program.py b/exir/program/_program.py index 94c5cad4786..b76827a597b 100644 --- a/exir/program/_program.py +++ b/exir/program/_program.py @@ -73,6 +73,9 @@ ) from executorch.exir.passes.remove_mixed_type_operators import RemoveMixedTypeOperators from executorch.exir.passes.replace_aten_with_edge_pass import aten_to_edge +from executorch.exir.passes.replace_slice_copy_with_slice_pass import ( + ReplaceSliceCopyWithSlicePass, +) from executorch.exir.passes.replace_view_copy_with_view_pass import ( ReplaceViewCopyWithViewPass, ) @@ -751,6 +754,7 @@ def pre_memory_planning_passes( NormalizeViewCopyBasePass(), dead_code_elimination_pass, ReplaceViewCopyWithViewPass(), + ReplaceSliceCopyWithSlicePass(), sym_shape_eval_pass, config.to_out_var_pass, ] diff --git a/exir/serde/export_serialize.py b/exir/serde/export_serialize.py index 670c244ca2f..3355c12806c 100644 --- a/exir/serde/export_serialize.py +++ b/exir/serde/export_serialize.py @@ -213,6 +213,7 @@ def _reverse_map(d: Dict[Any, Enum]): _KNOWN_FUNCTIONS = { exir.memory.view, + exir.memory.slice, } diff --git a/exir/serde/serialize.py b/exir/serde/serialize.py index a2eb2491067..849f7d425ad 100644 --- a/exir/serde/serialize.py +++ b/exir/serde/serialize.py @@ -376,6 +376,7 @@ def serialize( _KNOWN_FUNCTIONS_MAP = { "executorch.exir.memory.view": exir.memory.view, + "executorch.exir.memory.slice": exir.memory.slice, } diff --git a/exir/tests/targets.bzl b/exir/tests/targets.bzl index f68ee23cc16..9ecd23e9565 100644 --- a/exir/tests/targets.bzl +++ b/exir/tests/targets.bzl @@ -458,6 +458,25 @@ def define_common_targets(is_fbcode = False): ], ) + python_unittest( + name = "replace_slice_copy_with_slice_pass", + srcs = [ + "test_replace_slice_copy_with_slice_pass.py", + ], + deps = [ + "//caffe2:torch", + "//executorch/exir:lib", + "//executorch/exir:memory", + "//executorch/exir/passes:lib", + "//executorch/exir/passes:normalize_view_copy_base_pass", + "//executorch/exir/passes:reinplace_pass", + "//executorch/exir/passes:replace_slice_copy_with_slice_pass", + "//executorch/exir/passes:replace_view_copy_with_view_pass", + "//executorch/exir/passes:spec_prop_pass", + "//executorch/extension/pybindings:portable_lib", # @manual + ], + ) + python_unittest( name = "test_remove_view_copy", srcs = [ diff --git a/exir/tests/test_quant_fusion_pass.py b/exir/tests/test_quant_fusion_pass.py index 499419a1e10..a06f72a45f5 100644 --- a/exir/tests/test_quant_fusion_pass.py +++ b/exir/tests/test_quant_fusion_pass.py @@ -183,7 +183,8 @@ def forward(self, x, y): ) m = m.to_executorch() - # check that we are using out variant of add and slice_copy + # The slice aliases an input tensor, so it remains a slice_copy rather than + # a zero-copy memory alias. FileCheck().check("torch.ops.quantized_decomposed.add.out").check( "torch.ops.aten.slice_copy.Tensor_out" ).run(m.exported_program().graph_module.code) diff --git a/exir/tests/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py new file mode 100644 index 00000000000..f3c777b8d42 --- /dev/null +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -0,0 +1,455 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +# pyre-strict + +import copy +import unittest +from typing import List + +import torch +from executorch.exir import memory, to_edge +from executorch.exir.passes.normalize_view_copy_base_pass import ( + NormalizeViewCopyBasePass, +) +from executorch.exir.passes.reinplace import reinplace_pass +from executorch.exir.passes.replace_slice_copy_with_slice_pass import ( + _compute_slice_byte_offset, + _is_slice_copy, + _is_static_slice_argument, + is_contiguous_slice_copy, + ReplaceSliceCopyWithSlicePass, +) +from executorch.exir.passes.replace_view_copy_with_view_pass import ( + _ViewSpec, + ReplaceViewCopyWithViewPass, +) +from executorch.exir.passes.spec_prop_pass import SpecPropPass +from executorch.exir.schema import TensorShapeDynamism +from executorch.exir.tensor import TensorSpec +from executorch.extension.pybindings.portable_lib import ( + _load_for_executorch_from_buffer, +) +from torch.export import export +from torch.testing import assert_close + + +class TestReplaceSliceCopyWithSlicePass(unittest.TestCase): + def _edge_graph_module( + self, module: torch.nn.Module, inputs: tuple + ) -> torch.fx.GraphModule: + ep = export(module.eval(), inputs, strict=True) + return to_edge(ep).exported_program().graph_module + + def test_contiguity_classification(self) -> None: + """A unit-step slice along the outermost dim is contiguous (eligible); + inner-dim or strided slices are not.""" + + class M(torch.nn.Module): + def forward(self, x): + a = x[0:2] # dim 0, step 1 -> contiguous (eligible) + b = x[:, 1:3] # dim 1 -> strided (not eligible) + c = x[0:4:2] # dim 0, step 2 -> strided (not eligible) + return a.sum() + b.sum() + c.sum() + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + slice_nodes = [n for n in gm.graph.nodes if _is_slice_copy(n)] + eligible = [n for n in slice_nodes if is_contiguous_slice_copy(n)] + + self.assertEqual(len(slice_nodes), 3) + self.assertEqual(len(eligible), 1) + + def test_negative_outermost_dim_is_contiguous(self) -> None: + """A negative dim that resolves to the outermost dim is still eligible.""" + + class M(torch.nn.Module): + def forward(self, x): + # dim=-2 on a rank-2 tensor resolves to dim 0. + return torch.ops.aten.slice_copy.Tensor(x, -2, 0, 2).sum() + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + eligible = [n for n in gm.graph.nodes if is_contiguous_slice_copy(n)] + self.assertEqual(len(eligible), 1) + + def _annotate_input_spec(self, gm: torch.fx.GraphModule) -> None: + """Populate ``spec`` on every tensor placeholder. + + The lowering pipeline normally does this before the pass runs. Note + that ``to_edge`` lifts scalar constants to their own placeholders, so + annotating only the first placeholder would miss the real input. + """ + for node in gm.graph.nodes: + if node.op != "placeholder": + continue + val = node.meta.get("val") + if isinstance(val, torch.Tensor): + node.meta["spec"] = TensorSpec.from_tensor(val) + + def _annotate_tensor_specs(self, gm: torch.fx.GraphModule) -> None: + """Populate static specs for all tensor nodes in a small FX test graph.""" + for node in gm.graph.nodes: + val = node.meta.get("val") + if isinstance(val, torch.Tensor): + node.meta["spec"] = TensorSpec.from_tensor(val) + + def test_pass_replaces_annotated_contiguous_slice(self) -> None: + """A statically annotated dim-0 slice becomes a memory alias.""" + + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + return base[0:2] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self._annotate_tensor_specs(gm) + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertIsNotNone(result) + self.assertTrue(result.modified) + self.assertEqual( + len( + [ + n + for n in result.graph_module.graph.nodes + if n.op == "call_function" and n.target == memory.slice + ] + ), + 1, + ) + + def test_pass_skips_nondefault_base_dim_order(self) -> None: + """Avoid aliases that would reinterpret a non-contiguous base layout.""" + + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + return base[0:2] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self._annotate_tensor_specs(gm) + # Mutate the layout of the slice's own base, not just any placeholder. + slice_node = next(n for n in gm.graph.nodes if _is_slice_copy(n)) + slice_node.args[0].meta["spec"].dim_order = (1, 0) + + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertFalse(result.modified) + + def test_pass_normalizes_negative_start(self) -> None: + + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + return base[-2:] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self._annotate_tensor_specs(gm) + + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertTrue(result.modified) + sliced = next(n for n in gm.graph.nodes if n.target == memory.slice) + self.assertEqual(sliced.args[2:4], (2, 4)) + base_spec = sliced.args[0].meta["spec"] + base_spec.mem_offset = 128 + self.assertEqual(sliced.meta["spec"].mem_offset, 128 + 2 * 8 * 4) + self.assertEqual(_compute_slice_byte_offset(base_spec, 0, -2), 2 * 8 * 4) + + def test_clamped_bounds_and_empty_slices(self) -> None: + class M(torch.nn.Module): + def __init__(self, start, end): + super().__init__() + self.start = start + self.end = end + + def forward(self, x): + base = x + 0.0 + return base[self.start : self.end] + 1.0 + + x = torch.arange(32, dtype=torch.float32).reshape(4, 8) + for start, end, expected_alias in ( + (2, 99, True), + (-99, 2, True), + (-3, -1, True), + (None, None, True), + (99, 100, False), + (3, 1, False), + (0, -99, False), + ): + with self.subTest(start=start, end=end): + model = M(start, end) + ep = to_edge(export(model, (x,), strict=True)).exported_program() + gm = ep.graph_module + self._annotate_tensor_specs(gm) + result = ReplaceSliceCopyWithSlicePass()(gm) + # Export may remove a no-op full slice. + if start is not None or end is not None: + self.assertEqual(result.modified, expected_alias) + assert_close(ep.module()(x), model(x)) + for node in gm.graph.nodes: + if node.target == memory.slice: + base_spec = node.args[0].meta["spec"] + base_spec.mem_offset = 256 + spec = node.meta["spec"] + self.assertGreaterEqual(spec.mem_offset, 256) + self.assertLessEqual( + spec.mem_offset + spec.nbytes(), 256 + base_spec.nbytes() + ) + et = to_edge(export(model, (x,), strict=True)).to_executorch() + runtime = _load_for_executorch_from_buffer(et.buffer) + assert_close(runtime.forward((x,))[0], model(x)) + + def test_shared_view_spec_byte_range(self) -> None: + base = TensorSpec.from_tensor(torch.empty(4, 8)) + sliced = _ViewSpec(base, [2, 8], byte_offset=32) + viewed = _ViewSpec(sliced, [16]) + self.assertIsNone(viewed.mem_offset) + base.mem_offset = 128 + base.mem_id = 1 + self.assertEqual(sliced.mem_offset, 160) + self.assertEqual(viewed.mem_offset, 160) + self.assertEqual(viewed.mem_id, 1) + base.mem_offset = 256 + self.assertEqual(viewed.mem_offset, 288) + for offset in (-1, 100): + with self.subTest(offset=offset), self.assertRaises(Exception): + _ViewSpec(base, [2, 8], byte_offset=offset) + + def test_reinplace_mutation_safety(self) -> None: + class M(torch.nn.Module): + def __init__(self, mutate_slice, other_is_live): + super().__init__() + self.mutate_slice = mutate_slice + self.other_is_live = other_is_live + + def forward(self, x, indices, values): + base = torch.relu(x) + sliced = torch.ops.aten.slice_copy.Tensor(base, 0, 1, 3) + # Include a downstream view so alias families span both passes. + viewed = sliced.view(2, 8) + if self.mutate_slice: + changed = torch.ops.aten.index_put.default( + viewed, [indices], values + ) + return (base, changed) if self.other_is_live else changed + if self.other_is_live: + changed = torch.ops.aten.index_put.default(base, [indices], values) + return changed, viewed.clone() + observed = viewed.clone() + changed = torch.ops.aten.index_put.default(base, [indices], values) + return observed, changed + + inputs = ( + torch.arange(32, dtype=torch.float32).reshape(4, 8), + torch.tensor([1]), + torch.full((1, 8), 100.0), + ) + for mutate_slice in (False, True): + for other_is_live in (False, True): + with self.subTest(mutate_slice=mutate_slice, live=other_is_live): + model = M(mutate_slice, other_is_live) + expected = model(*copy.deepcopy(inputs)) + ep = to_edge(export(model, inputs, strict=True)).exported_program() + reinplace_pass(ep) + gm = SpecPropPass()(ep.graph_module).graph_module + NormalizeViewCopyBasePass()(gm) + ReplaceViewCopyWithViewPass()(gm) + ReplaceSliceCopyWithSlicePass()(gm) + aliases = [n for n in gm.graph.nodes if n.target == memory.slice] + self.assertEqual(bool(aliases), not other_is_live) + actual = gm(*copy.deepcopy(inputs)) + assert_close( + actual, expected if isinstance(expected, tuple) else (expected,) + ) + + def test_view_of_slice_uses_updated_spec(self) -> None: + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + return base[1:3].view(16) + 1.0 + + x = torch.arange(32, dtype=torch.float32).reshape(4, 8) + model = M() + gm = self._edge_graph_module(model, (x,)) + self._annotate_tensor_specs(gm) + NormalizeViewCopyBasePass()(gm) + ReplaceViewCopyWithViewPass()(gm) + ReplaceSliceCopyWithSlicePass()(gm) + sliced = next(n for n in gm.graph.nodes if n.target == memory.slice) + viewed = next(n for n in gm.graph.nodes if n.target == memory.view) + self.assertIs(viewed.meta["spec"]._base, sliced.meta["spec"]) + sliced.args[0].meta["spec"].mem_offset = 128 + self.assertEqual(viewed.meta["spec"].mem_offset, 160) + et = to_edge(export(model, (x,), strict=True)).to_executorch() + runtime = _load_for_executorch_from_buffer(et.buffer) + assert_close(runtime.forward((x,))[0], model(x)) + + def test_sibling_slice_aliases_preserve_mutation_safety(self) -> None: + for annotated in (False, True): + with self.subTest(allocation_sharing_annotation=annotated): + graph = torch.fx.Graph() + x = graph.placeholder("x") + base = graph.call_function(torch.ops.aten.relu.default, (x,)) + first = graph.call_function( + torch.ops.aten.slice_copy.Tensor, (base, 0, 1, 3) + ) + second = graph.call_function( + torch.ops.aten.slice_copy.Tensor, (base, 0, 1, 3) + ) + changed = graph.call_function( + ( + torch.ops.aten.add.Tensor + if annotated + else torch.ops.aten.add_.Tensor + ), + (first, 10), + ) + if annotated: + changed.meta["_share_alloc_with_arg_idx"] = 0 + observed = graph.call_function(torch.ops.aten.clone.default, (second,)) + graph.output((changed, observed)) + for node in (x, base): + node.meta["val"] = torch.empty(4, 8) + for node in (first, second, changed, observed): + node.meta["val"] = torch.empty(2, 8) + gm = torch.fx.GraphModule(torch.nn.Module(), graph) + self._annotate_tensor_specs(gm) + inputs = torch.arange(32, dtype=torch.float32).reshape(4, 8) + expected = gm(inputs) + ReplaceSliceCopyWithSlicePass()(gm) + self.assertTrue(_is_slice_copy(first)) + self.assertEqual(second.target, memory.slice) + assert_close(gm(inputs), expected) + + def test_static_slice_argument_check_rejects_runtime_nodes(self) -> None: + """Runtime graph values cannot be encoded as fixed alias offsets.""" + graph = torch.fx.Graph() + x = graph.placeholder("x") + start = graph.placeholder("start") + base = graph.call_function(torch.ops.aten.relu.default, (x,)) + sliced = graph.call_function( + torch.ops.aten.slice_copy.Tensor, (base, 0, start, 2) + ) + relu = graph.call_function(torch.ops.aten.relu.default, (sliced,)) + graph.output(relu) + + x.meta["val"] = torch.empty(4, 8) + x.meta["spec"] = TensorSpec.from_tensor(x.meta["val"]) + base.meta["val"] = torch.empty(4, 8) + base.meta["spec"] = TensorSpec.from_tensor(base.meta["val"]) + sliced.meta["val"] = torch.empty(2, 8) + sliced.meta["spec"] = TensorSpec.from_tensor(sliced.meta["val"]) + gm = torch.fx.GraphModule(torch.nn.Module(), graph) + + self.assertFalse(_is_static_slice_argument(start)) + self.assertTrue(_is_static_slice_argument(1)) + self.assertTrue(_is_static_slice_argument(None)) + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertFalse(result.modified) + self.assertEqual(sliced.target, torch.ops.aten.slice_copy.Tensor) + + def test_pass_skips_dynamic_output_shape(self) -> None: + """A dynamic slice must retain its copy kernel even with a static base.""" + + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + return base[0:2] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self._annotate_tensor_specs(gm) + slice_node = next(n for n in gm.graph.nodes if _is_slice_copy(n)) + slice_node.meta["spec"].shape_dynamism = TensorShapeDynamism.DYNAMIC_BOUND + + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertFalse(result.modified) + + def test_pass_skips_placeholder_base(self) -> None: + """External input tensors have no memory-planned allocation to alias.""" + graph = torch.fx.Graph() + x = graph.placeholder("x") + sliced = graph.call_function(torch.ops.aten.slice_copy.Tensor, (x, 0, 1, 3)) + relu = graph.call_function(torch.ops.aten.relu.default, (sliced,)) + graph.output(relu) + x.meta["val"] = torch.empty(4, 8) + x.meta["spec"] = TensorSpec.from_tensor(x.meta["val"]) + sliced.meta["val"] = torch.empty(2, 8) + sliced.meta["spec"] = TensorSpec.from_tensor(sliced.meta["val"]) + gm = torch.fx.GraphModule(torch.nn.Module(), graph) + + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertFalse(result.modified) + self.assertEqual(sliced.target, torch.ops.aten.slice_copy.Tensor) + + def _emitted_operators(self, program) -> List[str]: + return [ + str(op) for op in program.executorch_program.execution_plan[0].operators + ] + + def test_lowered_program_matches_eager_output(self) -> None: + """The emitted sub-buffer alias executes with the original semantics.""" + + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + return base[1:3] + 1.0 + + model = M().eval() + example_input = torch.arange(32, dtype=torch.float32).reshape(4, 8) + et = to_edge(export(model, (example_input,), strict=True)).to_executorch() + + # The slice must be aliased away, not merely produce the right answer -- + # falling back to a copy would also pass a numerical check alone. + self.assertFalse( + any("slice_copy" in op for op in self._emitted_operators(et)), + "expected the contiguous slice to be elided, but slice_copy was emitted", + ) + + runtime_module = _load_for_executorch_from_buffer(et.buffer) + assert_close(runtime_module.forward((example_input,))[0], model(example_input)) + + def test_base_outlives_slice_when_reused(self) -> None: + """The base buffer must not be reused while the alias is still live.""" + + class M(torch.nn.Module): + def forward(self, x): + base = x + 0.0 + sliced = base[1:3] + 1.0 + # ``base`` is consumed *after* the slice, so the planner has to keep + # the base alive across the alias's lifetime. + return sliced.sum() + base.sum() + + model = M().eval() + example_input = torch.arange(32, dtype=torch.float32).reshape(4, 8) + et = to_edge(export(model, (example_input,), strict=True)).to_executorch() + runtime_module = _load_for_executorch_from_buffer(et.buffer) + + assert_close(runtime_module.forward((example_input,))[0], model(example_input)) + + def test_chained_slice_falls_back_to_copy(self) -> None: + """A slice of a slice has no concrete base allocation to offset from.""" + + class M(torch.nn.Module): + def forward(self, x): + return x[0:3][1:2] + 1.0 + + model = M().eval() + example_input = torch.arange(32, dtype=torch.float32).reshape(4, 8) + # Must lower and execute correctly rather than tripping over an + # aliasing base during memory planning. + et = to_edge(export(model, (example_input,), strict=True)).to_executorch() + runtime_module = _load_for_executorch_from_buffer(et.buffer) + + assert_close(runtime_module.forward((example_input,))[0], model(example_input)) + + def test_non_slice_nodes_are_ignored(self) -> None: + class M(torch.nn.Module): + def forward(self, x): + return (x + 1.0).relu() + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self.assertEqual([n for n in gm.graph.nodes if is_contiguous_slice_copy(n)], []) + + +if __name__ == "__main__": + unittest.main()