From 1e9b6cd7e35bf2c12e4ebccff97451a9c863c42f Mon Sep 17 00:00:00 2001 From: iRAFEEK <182501111+iRAFEEK@users.noreply.github.com> Date: Mon, 3 Aug 2026 16:46:21 -0700 Subject: [PATCH 1/9] feat: add ReplaceSliceCopyWithSlicePass with contiguity detection (#10917) Slice analog of ReplaceViewCopyWithViewPass. Detects contiguous (outermost-dim, unit-step) slice_copy nodes eligible to be re-inplaced as zero-copy slices. Rewrite is gated behind offset-based sub-buffer aliasing support in memory planning (pending design discussion), so the pass currently runs as a safe no-op. --- .../replace_slice_copy_with_slice_pass.py | 123 ++++++++++++++++++ 1 file changed, 123 insertions(+) create mode 100644 exir/passes/replace_slice_copy_with_slice_pass.py 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..d0c53089a99 --- /dev/null +++ b/exir/passes/replace_slice_copy_with_slice_pass.py @@ -0,0 +1,123 @@ +# 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 + +"""Re-inplace contiguous ``slice_copy`` nodes as lightweight slices (#10917). + +This is the slice analog of :class:`ReplaceViewCopyWithViewPass`. A +``slice_copy`` addresses a sub-region of its input's storage; when that +sub-region is *contiguous* it can, in principle, be re-inplaced as a zero-copy +view into the base buffer instead of emitting a full-copy ``slice_copy`` kernel. + +Scope note (see #10917): + ``ReplaceViewCopyWithViewPass`` can reuse ``memory.view`` because a view + aliases the *entire* base buffer -- same ``nbytes`` and offset ``0`` (the + ``_ViewSpec`` guards ``nbytes == base.nbytes``). A slice aliases only a + *sub-region* at a non-zero byte offset with fewer bytes than the base, and + ExecuTorch has no offset-based aliasing mechanism in memory planning today. + Fully eliminating the copy therefore requires (a) memory-planning support + for offset sub-buffer aliasing and (b) a lightweight runtime op + (``et_slice``) mirroring ``et_view``. That runtime design is under + discussion with the maintainer. + + This pass implements the piece that is well-defined regardless of that + design decision: correctly identifying which ``slice_copy`` nodes are + *eligible* (contiguous) for re-inplacing. The rewrite is gated behind the + offset-aliasing support and is a no-op until it lands. +""" + +import logging + +import torch +from executorch.exir.dialects._ops import ops +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_contiguous_slice_copy(node: torch.fx.Node) -> bool: + """Return True if ``node`` is a ``slice_copy`` whose result is a contiguous + sub-region of a contiguous input, and is therefore eligible to be + re-inplaced as a zero-copy slice. + + A slice ``self[start:end:step]`` along ``dim`` is a contiguous sub-buffer of + a contiguous input only when it is taken along the outermost (first) storage + dimension with unit step. Slicing an inner dimension, or using ``step > 1``, + produces a strided (non-contiguous) result that cannot alias the base buffer + without a copy. + + Signature: ``slice_copy.Tensor(self, dim=0, start=None, end=None, step=1)``. + """ + if not _is_slice_copy(node): + return False + + args = node.args + self_arg = args[0] + + # dim defaults to 0; normalize negatives against the input rank. + dim = args[1] if len(args) > 1 else 0 + step = args[4] if len(args) > 4 else 1 + + if step != 1: + return False + + rank = None + self_val = self_arg.meta.get("val") if isinstance(self_arg, torch.fx.Node) else None + if self_val is not None and hasattr(self_val, "dim"): + rank = self_val.dim() + + if isinstance(dim, int) and dim < 0: + if rank is None: + # Cannot resolve a negative dim to the outermost dim without rank. + return False + dim = dim + rank + + # Only an outermost-dim, unit-step slice of a contiguous input is a + # contiguous sub-buffer that could alias the base storage. + return dim == 0 + + +class ReplaceSliceCopyWithSlicePass(PassBase): + """Re-inplace eligible (contiguous) ``slice_copy`` nodes as lightweight + slices. + + Until offset-based sub-buffer aliasing lands in memory planning (see the + module docstring and #10917), this pass only *identifies* eligible nodes and + performs no graph mutation, so it is safe to run in the pipeline. + """ + + def __init__(self) -> None: + super().__init__() + + def call(self, graph_module: torch.fx.GraphModule) -> PassResult: + n_eligible = 0 + for module in graph_module.modules(): + if not isinstance(module, torch.fx.GraphModule): + continue + for node in module.graph.nodes: + # A slice feeding the graph output can have its pointer modified + # at runtime, mirroring the view_copy pass's output guard. + if is_contiguous_slice_copy(node) and all( + u.op != "output" for u in node.users + ): + n_eligible += 1 + + logger.debug( + "ReplaceSliceCopyWithSlicePass: %d contiguous slice_copy node(s) " + "eligible for re-inplacing (rewrite pending offset-aliasing support, " + "#10917).", + n_eligible, + ) + # No mutation yet -> report unchanged. + return PassResult(graph_module, False) From 77f098bbbf72e36d2b86cad1b3701c20f8d319e9 Mon Sep 17 00:00:00 2001 From: iRAFEEK <182501111+iRAFEEK@users.noreply.github.com> Date: Mon, 3 Aug 2026 16:46:21 -0700 Subject: [PATCH 2/9] test: add unit tests for slice_copy contiguity detection (#10917) Covers outermost-dim/unit-step eligibility, negative-dim resolution, strided/inner-dim rejection, and that the pass is a safe no-op until the offset-aliasing rewrite lands. --- ...test_replace_slice_copy_with_slice_pass.py | 87 +++++++++++++++++++ 1 file changed, 87 insertions(+) create mode 100644 exir/tests/test_replace_slice_copy_with_slice_pass.py 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..4e2e468a5d7 --- /dev/null +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -0,0 +1,87 @@ +# 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 unittest + +import torch +from executorch.exir import to_edge +from executorch.exir.passes.replace_slice_copy_with_slice_pass import ( + _is_slice_copy, + is_contiguous_slice_copy, + ReplaceSliceCopyWithSlicePass, +) +from torch.export import export + + +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 test_pass_is_safe_noop_until_offset_aliasing_lands(self) -> None: + """The pass must run cleanly and not mutate the graph while the + offset-aliasing rewrite is still gated (see #10917).""" + + class M(torch.nn.Module): + def forward(self, x): + return x[0:2] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + before = gm.code + result = ReplaceSliceCopyWithSlicePass()(gm) + self.assertIsNotNone(result) + self.assertFalse(result.modified) + self.assertEqual(before, result.graph_module.code) + + 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() From 730cb5631e078e5956eddfdb6470811f8e771570 Mon Sep 17 00:00:00 2001 From: iRAFEEK <182501111+iRAFEEK@users.noreply.github.com> Date: Wed, 5 Aug 2026 09:44:08 -0700 Subject: [PATCH 3/9] style: collapse single-line assertEqual to satisfy lintrunner --- exir/tests/test_replace_slice_copy_with_slice_pass.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/exir/tests/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py index 4e2e468a5d7..2721206ef1e 100644 --- a/exir/tests/test_replace_slice_copy_with_slice_pass.py +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -78,9 +78,7 @@ 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)], [] - ) + self.assertEqual([n for n in gm.graph.nodes if is_contiguous_slice_copy(n)], []) if __name__ == "__main__": From 3b8b8bc45c2001fc05cec21ceb2fee2064a41f9d Mon Sep 17 00:00:00 2001 From: iRAFEEK <182501111+iRAFEEK@users.noreply.github.com> Date: Wed, 5 Aug 2026 11:42:30 -0700 Subject: [PATCH 4/9] style: format slice copy contiguity test --- exir/tests/test_replace_slice_copy_with_slice_pass.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/exir/tests/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py index 2721206ef1e..95751eea6ca 100644 --- a/exir/tests/test_replace_slice_copy_with_slice_pass.py +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -52,9 +52,7 @@ def forward(self, x): 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) - ] + eligible = [n for n in gm.graph.nodes if is_contiguous_slice_copy(n)] self.assertEqual(len(eligible), 1) def test_pass_is_safe_noop_until_offset_aliasing_lands(self) -> None: From 178458fd435e6b1c2531f1d181659d233a312fb8 Mon Sep 17 00:00:00 2001 From: iRAFEEK <182501111+iRAFEEK@users.noreply.github.com> Date: Mon, 10 Aug 2026 13:21:10 -0700 Subject: [PATCH 5/9] feat: implement zero-copy sub-buffer aliasing for contiguous slice_copy Replaces eligible contiguous slice_copy nodes with a memory.slice alias so the emitted program does not pay for a full tensor copy. _SliceSpec shares the base's mem_id and computes mem_offset = base.mem_offset + start * base.stride[0] * elem_size. The .pte format already carries (memory_id, memory_offset) via AllocationDetails, so no schema change is required. Memory planning handles memory.slice like memory.view -- the base spec is returned from get_node_tensor_specs, which extends the base's lifetime over the slice's consumers so the buffer is not reused while the alias is live. Emission mirrors _emit_view's elide path, needing no runtime kernel. Eligibility is gated to dim-0, unit-step slices with a non-negative start on a base that has the default dim order and its own allocation. Non-default layouts would otherwise be silently reinterpreted by the contiguous output stride, and an aliasing base (slice-of-slice or slice-of-view) has no concrete allocation to offset from. Everything outside those gates falls back to slice_copy unchanged. Also declares inplace_base on _SliceSpec, which the greedy memory planning algorithm reads. Verified locally against the executorch wheel runtime: - contiguous slices emit no slice_copy kernel (only aten::add) - outputs match eager for offset/lifetime/chained/3-D cases - ineligible slices still fall back to copy and stay correct - no regressions: exir/tests, exir/emit, exir/backend failure sets are identical to a pristine baseline --- exir/emit/_emitter.py | 17 + exir/memory.py | 15 +- exir/memory_planning.py | 5 + exir/pass_base.py | 2 +- exir/passes/__init__.py | 1 + .../replace_slice_copy_with_slice_pass.py | 292 ++++++++++++++---- exir/program/_program.py | 4 + exir/serde/export_serialize.py | 1 + exir/serde/serialize.py | 1 + ...test_replace_slice_copy_with_slice_pass.py | 135 +++++++- 10 files changed, 408 insertions(+), 65 deletions(-) diff --git a/exir/emit/_emitter.py b/exir/emit/_emitter.py index ec0d97df4b0..c92a653441a 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 ``_SliceSpec`` 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..9d2cf2d977e 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, @@ -914,6 +915,10 @@ def get_node_tensor_specs( base = node.args[0] assert isinstance(base, torch.fx.Node) specs = base.meta.get("spec") + elif node.target == memory.slice: + base = node.args[0] + assert isinstance(base, torch.fx.Node) + specs = base.meta.get("spec") 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/__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 index d0c53089a99..4799e98b3d2 100644 --- a/exir/passes/replace_slice_copy_with_slice_pass.py +++ b/exir/passes/replace_slice_copy_with_slice_pass.py @@ -8,36 +8,31 @@ """Re-inplace contiguous ``slice_copy`` nodes as lightweight slices (#10917). -This is the slice analog of :class:`ReplaceViewCopyWithViewPass`. A -``slice_copy`` addresses a sub-region of its input's storage; when that -sub-region is *contiguous* it can, in principle, be re-inplaced as a zero-copy -view into the base buffer instead of emitting a full-copy ``slice_copy`` kernel. - -Scope note (see #10917): - ``ReplaceViewCopyWithViewPass`` can reuse ``memory.view`` because a view - aliases the *entire* base buffer -- same ``nbytes`` and offset ``0`` (the - ``_ViewSpec`` guards ``nbytes == base.nbytes``). A slice aliases only a - *sub-region* at a non-zero byte offset with fewer bytes than the base, and - ExecuTorch has no offset-based aliasing mechanism in memory planning today. - Fully eliminating the copy therefore requires (a) memory-planning support - for offset sub-buffer aliasing and (b) a lightweight runtime op - (``et_slice``) mirroring ``et_view``. That runtime design is under - discussion with the maintainer. - - This pass implements the piece that is well-defined regardless of that - design decision: correctly identifying which ``slice_copy`` nodes are - *eligible* (contiguous) for re-inplacing. The rewrite is gated behind the - offset-aliasing support and is a no-op until it lands. +Slice analog of :class:`ReplaceViewCopyWithViewPass`. Contiguous slices taken +along the outermost dimension with unit step alias a sub-region of the base +buffer and can be represented with :class:`_SliceSpec`, which shares the +base's ``mem_id`` but uses a computed byte ``mem_offset``. """ import logging +from typing import Any, List, Optional import torch +from executorch.exir import memory from executorch.exir.dialects._ops import ops +from executorch.exir.sym_util import eval_shape +from executorch.exir.tensor import ( + contiguous_stride_from_shape, + determine_tensor_dynanism, + dim_order_from_stride, + TensorSpec, +) from torch.fx.passes.infra.pass_base import PassBase, PassResult logger: logging.Logger = logging.getLogger(__name__) +_SLICE_OP = memory.slice + def _is_slice_copy(node: torch.fx.Node) -> bool: return node.op == "call_function" and node.target in ( @@ -46,26 +41,21 @@ def _is_slice_copy(node: torch.fx.Node) -> bool: ) -def is_contiguous_slice_copy(node: torch.fx.Node) -> bool: - """Return True if ``node`` is a ``slice_copy`` whose result is a contiguous - sub-region of a contiguous input, and is therefore eligible to be - re-inplaced as a zero-copy slice. +def _normalize_dim(dim: int, rank: Optional[int]) -> Optional[int]: + if isinstance(dim, int) and dim < 0: + if rank is None: + return None + dim = dim + rank + return dim - A slice ``self[start:end:step]`` along ``dim`` is a contiguous sub-buffer of - a contiguous input only when it is taken along the outermost (first) storage - dimension with unit step. Slicing an inner dimension, or using ``step > 1``, - produces a strided (non-contiguous) result that cannot alias the base buffer - without a copy. - Signature: ``slice_copy.Tensor(self, dim=0, start=None, end=None, step=1)``. - """ +def is_contiguous_slice_copy(node: torch.fx.Node) -> bool: + """True if ``node`` is an outermost-dim, unit-step ``slice_copy``.""" if not _is_slice_copy(node): return False args = node.args self_arg = args[0] - - # dim defaults to 0; normalize negatives against the input rank. dim = args[1] if len(args) > 1 else 0 step = args[4] if len(args) > 4 else 1 @@ -77,47 +67,235 @@ def is_contiguous_slice_copy(node: torch.fx.Node) -> bool: if self_val is not None and hasattr(self_val, "dim"): rank = self_val.dim() - if isinstance(dim, int) and dim < 0: - if rank is None: - # Cannot resolve a negative dim to the outermost dim without rank. - return False - dim = dim + rank - - # Only an outermost-dim, unit-step slice of a contiguous input is a - # contiguous sub-buffer that could alias the base storage. + dim = _normalize_dim(dim, rank) return dim == 0 -class ReplaceSliceCopyWithSlicePass(PassBase): - """Re-inplace eligible (contiguous) ``slice_copy`` nodes as lightweight - slices. +def _slice_start_as_int(start: Any) -> int: + if start is None: + return 0 + if isinstance(start, int): + return start + if isinstance(start, torch.SymInt): + return int(eval_shape([start])[0]) + return int(start) + + +def _compute_slice_byte_offset(base: TensorSpec, dim: int, start: Any) -> int: + start_int = _slice_start_as_int(start) + if start_int < 0: + raise ValueError("memory.slice does not support negative slice starts.") + elem_size = torch._utils._element_size(base.dtype) + return start_int * base.stride[dim] * elem_size + + +def _is_aliasing_base(base: torch.fx.Node) -> bool: + """Whether ``base`` is itself an alias rather than a real allocation. + + ``memory.slice`` and ``memory.view`` nodes do not own storage, so a slice + taken from one has no concrete ``mem_offset`` to build on until the chain is + normalized down to the first real allocation. + """ + return base.op == "call_function" and base.target in (memory.slice, memory.view) + + +def _has_default_dim_order(spec: TensorSpec) -> bool: + """Whether ``spec`` has the standard contiguous dimension ordering. - Until offset-based sub-buffer aliasing lands in memory planning (see the - module docstring and #10917), this pass only *identifies* eligible nodes and - performs no graph mutation, so it is safe to run in the pipeline. + ``_SliceSpec`` computes a contiguous output stride. That is only a valid + alias for a dim-0 slice when the base itself has the default dim order. """ + return spec.dim_order == dim_order_from_stride( + contiguous_stride_from_shape(torch.Size(spec.shape)) + ) + + +class _SliceSpec(TensorSpec): + """TensorSpec for a zero-copy slice into a contiguous base buffer.""" + + def __init__( + self, + base: TensorSpec, + shape: List[int], + dim: int, + start: Any, + ) -> None: + if base.is_sparse: + raise Exception( + "_SliceSpec can only be created from non-sparse TensorSpec." + ) + if base.layout != torch.strided: + raise Exception(f"_SliceSpec requires strided layout, got {base.layout}.") + + self._base = base + self._byte_offset = _compute_slice_byte_offset(base, dim, start) + self._unguarded_access = False + + self._self_fields = [ + "debug", + "__repr__", + "shape", + "stride", + "dim_order", + "shape_dynamism", + "nbytes", + "allocated_memory", + "is_dynamic_shape_tensor", + "is_static_shape_tensor", + "is_upper_bound_tensor", + "is_dynamic_unbound_tensor", + "mem_offset", + ] + self._base_fields = [ + "scalar_type", + "const", + "alignment", + "storage", + "requires_grad", + "layout", + "is_sparse", + "init_mem_planning_fields", + "realign", + "from_tensor", + "lifetime", + "mem_id", + "mem_obj_id", + "dtype", + "extra_tensor_info", + "device", + "device_index", + # Read by the memory planning algorithms (e.g. ``greedy``). A slice + # is never itself an in-place target, so it defers to its base. + "inplace_base", + ] + + self.shape = list(shape) + self.stride = contiguous_stride_from_shape(torch.Size(self.shape)) + self.dim_order = dim_order_from_stride(self.stride) + self.shape_dynamism = determine_tensor_dynanism(torch.Size(self.shape)) + + if self.shape_dynamism != base.shape_dynamism: + raise Exception( + f"_SliceSpec shape_dynamism {self.shape_dynamism} != base {base.shape_dynamism}" + ) + if self.dtype != base.dtype: + raise Exception(f"_SliceSpec dtype {self.dtype} != base {base.dtype}") + + def __getattribute__(self, name: str): # pyre-ignore + if name in [ + "_base", + "_self_fields", + "_base_fields", + "_byte_offset", + "_unguarded_access", + ]: + return object.__getattribute__(self, name) + + self_fields = object.__getattribute__(self, "_self_fields") + base_fields = object.__getattribute__(self, "_base_fields") + + if name == "mem_offset": + base = object.__getattribute__(self, "_base") + base_offset = base.mem_offset + if base_offset is None: + return None + byte_offset = object.__getattribute__(self, "_byte_offset") + return base_offset + byte_offset + + if name in self_fields: + if name in ("nbytes", "allocated_memory"): + return TensorSpec.__getattribute__(self, name) + return object.__getattribute__(self, name) + + if name in base_fields: + base = object.__getattribute__(self, "_base") + return object.__getattribute__(base, name) + + return object.__getattribute__(self, name) + + def __setattr__(self, name: str, val) -> None: # pyre-ignore + if name in [ + "_base", + "_self_fields", + "_base_fields", + "_byte_offset", + "_unguarded_access", + ]: + object.__setattr__(self, name, val) + return + + if hasattr(self, "_self_fields") and name in self._self_fields: + if name == "mem_offset": + raise Exception("_SliceSpec.mem_offset is computed from the base.") + object.__setattr__(self, name, val) + return + + if hasattr(self, "_base_fields") and name in self._base_fields: + object.__setattr__(self._base, name, val) + return + + object.__setattr__(self, name, val) + + +class ReplaceSliceCopyWithSlicePass(PassBase): + """Replace eligible contiguous ``slice_copy`` nodes with ``memory.slice``.""" def __init__(self) -> None: super().__init__() def call(self, graph_module: torch.fx.GraphModule) -> PassResult: - n_eligible = 0 + n_replaced = 0 for module in graph_module.modules(): if not isinstance(module, torch.fx.GraphModule): continue for node in module.graph.nodes: - # A slice feeding the graph output can have its pointer modified - # at runtime, mirroring the view_copy pass's output guard. if is_contiguous_slice_copy(node) and all( u.op != "output" for u in node.users ): - n_eligible += 1 + base = node.args[0] + if ( + not isinstance(base, torch.fx.Node) + or "spec" not in base.meta + or not base.meta["spec"].is_static_shape_tensor + or not _has_default_dim_order(base.meta["spec"]) + ): + # Specs are populated by the lowering pipeline before this + # pass. Skip bare FX graphs so the pass remains safe to use + # in isolation as well. + continue + if _is_aliasing_base(base): + # The base is itself an alias (a slice or a view), so it + # has no allocation of its own to offset from. Chaining + # offsets through it would require normalizing to the + # first real allocation first, so leave this as a copy. + continue + dim = node.args[1] if len(node.args) > 1 else 0 + start = node.args[2] if len(node.args) > 2 else None + if _slice_start_as_int(start) < 0: + # Negative starts are relative to the end of the + # dimension. They cannot be expressed as a static + # offset without normalizing against the base shape. + continue + node.target = _SLICE_OP + shape = node.meta["val"].shape + node.meta["spec"] = _SliceSpec( + base.meta["spec"], list(shape), dim, start + ) + n_replaced += 1 + + module.recompile() logger.debug( - "ReplaceSliceCopyWithSlicePass: %d contiguous slice_copy node(s) " - "eligible for re-inplacing (rewrite pending offset-aliasing support, " - "#10917).", - n_eligible, + "ReplaceSliceCopyWithSlicePass: replaced %d slice_copy node(s) with %s.", + n_replaced, + _SLICE_OP, ) - # No mutation yet -> report unchanged. - return PassResult(graph_module, False) + return PassResult(graph_module, n_replaced > 0) + + def ensures(self, graph_module: torch.fx.GraphModule) -> None: + for module in graph_module.modules(): + if not isinstance(module, torch.fx.GraphModule): + continue + for node in module.graph.nodes: + if node.op == "call_function" and node.target == _SLICE_OP: + assert isinstance(node.meta["spec"], _SliceSpec) 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/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py index 95751eea6ca..77f0425f250 100644 --- a/exir/tests/test_replace_slice_copy_with_slice_pass.py +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -7,15 +7,22 @@ # pyre-strict import unittest +from typing import List import torch -from executorch.exir import to_edge +from executorch.exir import memory, to_edge from executorch.exir.passes.replace_slice_copy_with_slice_pass import ( + _compute_slice_byte_offset, _is_slice_copy, is_contiguous_slice_copy, ReplaceSliceCopyWithSlicePass, ) +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): @@ -55,20 +62,136 @@ def forward(self, x): eligible = [n for n in gm.graph.nodes if is_contiguous_slice_copy(n)] self.assertEqual(len(eligible), 1) - def test_pass_is_safe_noop_until_offset_aliasing_lands(self) -> None: - """The pass must run cleanly and not mutate the graph while the - offset-aliasing rewrite is still gated (see #10917).""" + 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 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): return x[0:2] + 1.0 gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) - before = gm.code + self._annotate_input_spec(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): + return x[0:2] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self._annotate_input_spec(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_skips_negative_start(self) -> None: + """Negative starts need shape-dependent normalization, so keep copying.""" + + class M(torch.nn.Module): + def forward(self, x): + return x[-2:] + 1.0 + + gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) + self._annotate_input_spec(gm) + + result = ReplaceSliceCopyWithSlicePass()(gm) self.assertFalse(result.modified) - self.assertEqual(before, result.graph_module.code) + with self.assertRaises(ValueError): + _compute_slice_byte_offset( + next(n for n in gm.graph.nodes if n.op == "placeholder").meta["spec"], + 0, + -2, + ) + + 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): + return x[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): + sliced = x[1:3] + 1.0 + # ``x`` is consumed *after* the slice, so the planner has to keep + # the base alive across the alias's lifetime. + return sliced.sum() + x.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): From 39b81ca816e5311c88fd58e8fe6f7da816b776bd Mon Sep 17 00:00:00 2001 From: iRAFEEK Date: Sat, 26 Sep 2026 05:13:19 +0900 Subject: [PATCH 6/9] fix: restrict zero-copy slices to planned static buffers --- exir/passes/BUCK | 14 +++ .../replace_slice_copy_with_slice_pass.py | 46 +++++++-- exir/program/BUCK | 1 + exir/tests/targets.bzl | 15 +++ exir/tests/test_quant_fusion_pass.py | 4 +- ...test_replace_slice_copy_with_slice_pass.py | 96 +++++++++++++++++-- 6 files changed, 156 insertions(+), 20 deletions(-) diff --git a/exir/passes/BUCK b/exir/passes/BUCK index 0f1ea2a7cdd..63f4fb605a6 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,19 @@ 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 = [ + "//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/replace_slice_copy_with_slice_pass.py b/exir/passes/replace_slice_copy_with_slice_pass.py index 4799e98b3d2..06be14fa25c 100644 --- a/exir/passes/replace_slice_copy_with_slice_pass.py +++ b/exir/passes/replace_slice_copy_with_slice_pass.py @@ -20,7 +20,6 @@ import torch from executorch.exir import memory from executorch.exir.dialects._ops import ops -from executorch.exir.sym_util import eval_shape from executorch.exir.tensor import ( contiguous_stride_from_shape, determine_tensor_dynanism, @@ -71,17 +70,25 @@ def is_contiguous_slice_copy(node: torch.fx.Node) -> bool: return dim == 0 -def _slice_start_as_int(start: Any) -> int: +def _slice_start_as_int(start: Optional[int]) -> int: if start is None: return 0 - if isinstance(start, int): - return start - if isinstance(start, torch.SymInt): - return int(eval_shape([start])[0]) - return int(start) + return start -def _compute_slice_byte_offset(base: TensorSpec, dim: int, start: Any) -> int: +def _is_static_slice_argument(value: Any) -> bool: + """Whether a slice parameter is known when the program is lowered. + + ``memory.slice`` encodes a byte offset in the emitted tensor metadata, so + it cannot represent an FX node or a symbolic value that is resolved only + at runtime. Keep those slices as ``aten.slice_copy``. + """ + return value is None or isinstance(value, int) + + +def _compute_slice_byte_offset( + base: TensorSpec, dim: int, start: Optional[int] +) -> int: start_int = _slice_start_as_int(start) if start_int < 0: raise ValueError("memory.slice does not support negative slice starts.") @@ -263,6 +270,14 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult: # pass. Skip bare FX graphs so the pass remains safe to use # in isolation as well. continue + output_spec = node.meta.get("spec") + if ( + not isinstance(output_spec, TensorSpec) + or not output_spec.is_static_shape_tensor + ): + # A static base can still produce a dynamic slice when + # bounds originate in the runtime graph. + continue if _is_aliasing_base(base): # The base is itself an alias (a slice or a view), so it # has no allocation of its own to offset from. Chaining @@ -271,6 +286,21 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult: continue dim = node.args[1] if len(node.args) > 1 else 0 start = node.args[2] if len(node.args) > 2 else None + end = node.args[3] if len(node.args) > 3 else None + step = node.args[4] if len(node.args) > 4 else 1 + if not all( + _is_static_slice_argument(arg) + for arg in (dim, start, end, step) + ): + # Do not turn symbolic bounds into a fixed byte offset. + # The regular slice_copy kernel handles runtime values. + continue + if base.op == "placeholder": + # Placeholders are externally owned buffers rather than + # intermediate allocations. Unlike memory.view, this + # pass has no runtime-kernel fallback, so only alias + # storage that memory planning owns. + continue if _slice_start_as_int(start) < 0: # Negative starts are relative to the end of the # dimension. They cannot be expressed as a static 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/tests/targets.bzl b/exir/tests/targets.bzl index f68ee23cc16..733c61dcd3e 100644 --- a/exir/tests/targets.bzl +++ b/exir/tests/targets.bzl @@ -458,6 +458,21 @@ 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:replace_slice_copy_with_slice_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..aeb9d10e212 100644 --- a/exir/tests/test_quant_fusion_pass.py +++ b/exir/tests/test_quant_fusion_pass.py @@ -183,9 +183,9 @@ def forward(self, x, y): ) m = m.to_executorch() - # check that we are using out variant of add and slice_copy + # The static dim-0 slice is now represented as a zero-copy memory alias. FileCheck().check("torch.ops.quantized_decomposed.add.out").check( - "torch.ops.aten.slice_copy.Tensor_out" + "executorch.exir.memory.slice" ).run(m.exported_program().graph_module.code) def test_cat(self) -> None: diff --git a/exir/tests/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py index 77f0425f250..7482f94447a 100644 --- a/exir/tests/test_replace_slice_copy_with_slice_pass.py +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -14,9 +14,11 @@ 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.schema import TensorShapeDynamism from executorch.exir.tensor import TensorSpec from executorch.extension.pybindings.portable_lib import ( _load_for_executorch_from_buffer, @@ -76,15 +78,23 @@ def _annotate_input_spec(self, gm: torch.fx.GraphModule) -> None: 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): - return x[0:2] + 1.0 + base = x + 0.0 + return base[0:2] + 1.0 gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) - self._annotate_input_spec(gm) + self._annotate_tensor_specs(gm) result = ReplaceSliceCopyWithSlicePass()(gm) self.assertIsNotNone(result) self.assertTrue(result.modified) @@ -104,10 +114,11 @@ def test_pass_skips_nondefault_base_dim_order(self) -> None: class M(torch.nn.Module): def forward(self, x): - return x[0:2] + 1.0 + base = x + 0.0 + return base[0:2] + 1.0 gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) - self._annotate_input_spec(gm) + 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) @@ -120,10 +131,11 @@ def test_pass_skips_negative_start(self) -> None: class M(torch.nn.Module): def forward(self, x): - return x[-2:] + 1.0 + base = x + 0.0 + return base[-2:] + 1.0 gm = self._edge_graph_module(M(), (torch.randn(4, 8),)) - self._annotate_input_spec(gm) + self._annotate_tensor_specs(gm) result = ReplaceSliceCopyWithSlicePass()(gm) self.assertFalse(result.modified) @@ -134,6 +146,68 @@ def forward(self, x): -2, ) + 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 @@ -144,7 +218,8 @@ def test_lowered_program_matches_eager_output(self) -> None: class M(torch.nn.Module): def forward(self, x): - return x[1:3] + 1.0 + base = x + 0.0 + return base[1:3] + 1.0 model = M().eval() example_input = torch.arange(32, dtype=torch.float32).reshape(4, 8) @@ -165,10 +240,11 @@ def test_base_outlives_slice_when_reused(self) -> None: class M(torch.nn.Module): def forward(self, x): - sliced = x[1:3] + 1.0 - # ``x`` is consumed *after* the slice, so the planner has to keep + 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() + x.sum() + return sliced.sum() + base.sum() model = M().eval() example_input = torch.arange(32, dtype=torch.float32).reshape(4, 8) From 59407f88add090be0eb736b0c9b6dc6ea449f820 Mon Sep 17 00:00:00 2001 From: iRAFEEK Date: Tue, 29 Sep 2026 03:48:26 +0900 Subject: [PATCH 7/9] style: format slice-copy pass and tests --- exir/passes/replace_slice_copy_with_slice_pass.py | 4 +--- exir/tests/test_replace_slice_copy_with_slice_pass.py | 4 +--- 2 files changed, 2 insertions(+), 6 deletions(-) diff --git a/exir/passes/replace_slice_copy_with_slice_pass.py b/exir/passes/replace_slice_copy_with_slice_pass.py index 06be14fa25c..6dda0d8a431 100644 --- a/exir/passes/replace_slice_copy_with_slice_pass.py +++ b/exir/passes/replace_slice_copy_with_slice_pass.py @@ -86,9 +86,7 @@ def _is_static_slice_argument(value: Any) -> bool: return value is None or isinstance(value, int) -def _compute_slice_byte_offset( - base: TensorSpec, dim: int, start: Optional[int] -) -> int: +def _compute_slice_byte_offset(base: TensorSpec, dim: int, start: Optional[int]) -> int: start_int = _slice_start_as_int(start) if start_int < 0: raise ValueError("memory.slice does not support negative slice starts.") diff --git a/exir/tests/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py index 7482f94447a..ff89f6d56e6 100644 --- a/exir/tests/test_replace_slice_copy_with_slice_pass.py +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -193,9 +193,7 @@ 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) - ) + 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) From 97870a95401825474c279421bab33606832f4f26 Mon Sep 17 00:00:00 2001 From: iRAFEEK Date: Wed, 30 Sep 2026 04:21:03 +0900 Subject: [PATCH 8/9] test: preserve input slice copy expectation --- exir/tests/test_quant_fusion_pass.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/exir/tests/test_quant_fusion_pass.py b/exir/tests/test_quant_fusion_pass.py index aeb9d10e212..a06f72a45f5 100644 --- a/exir/tests/test_quant_fusion_pass.py +++ b/exir/tests/test_quant_fusion_pass.py @@ -183,9 +183,10 @@ def forward(self, x, y): ) m = m.to_executorch() - # The static dim-0 slice is now represented as a zero-copy memory alias. + # 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( - "executorch.exir.memory.slice" + "torch.ops.aten.slice_copy.Tensor_out" ).run(m.exported_program().graph_module.code) def test_cat(self) -> None: From 45f93f4686580776a1ab85354ee793d9569362fd Mon Sep 17 00:00:00 2001 From: iRAFEEK Date: Fri, 2 Oct 2026 20:06:07 +0900 Subject: [PATCH 9/9] fix: make slice aliases safe with reinplacement --- exir/emit/_emitter.py | 2 +- exir/memory_planning.py | 8 +- exir/passes/BUCK | 1 + .../replace_slice_copy_with_slice_pass.py | 370 +++++------------- .../replace_view_copy_with_view_pass.py | 37 +- exir/tests/targets.bzl | 4 + ...test_replace_slice_copy_with_slice_pass.py | 193 ++++++++- 7 files changed, 319 insertions(+), 296 deletions(-) diff --git a/exir/emit/_emitter.py b/exir/emit/_emitter.py index c92a653441a..e56b4a1662f 100644 --- a/exir/emit/_emitter.py +++ b/exir/emit/_emitter.py @@ -1303,7 +1303,7 @@ def _emit_view(self, args: Tuple[_Argument, ...]) -> _EmitterValue: def _emit_slice(self, args: Tuple[_Argument, ...]) -> _EmitterValue: """Emit a statically memory-planned slice as a sub-buffer alias. - ``ReplaceSliceCopyWithSlicePass`` creates ``_SliceSpec`` values whose + ``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. diff --git a/exir/memory_planning.py b/exir/memory_planning.py index 9d2cf2d977e..5ac02ab42d3 100644 --- a/exir/memory_planning.py +++ b/exir/memory_planning.py @@ -911,14 +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") - elif node.target == 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/passes/BUCK b/exir/passes/BUCK index 63f4fb605a6..134f6360f3d 100644 --- a/exir/passes/BUCK +++ b/exir/passes/BUCK @@ -450,6 +450,7 @@ fbcode_target(_kind = runtime.python_library, "replace_slice_copy_with_slice_pass.py", ], deps = [ + ":replace_view_copy_with_view_pass", "//caffe2:torch", "//executorch/exir:memory", "//executorch/exir:tensor", diff --git a/exir/passes/replace_slice_copy_with_slice_pass.py b/exir/passes/replace_slice_copy_with_slice_pass.py index 6dda0d8a431..bfe419e9255 100644 --- a/exir/passes/replace_slice_copy_with_slice_pass.py +++ b/exir/passes/replace_slice_copy_with_slice_pass.py @@ -6,23 +6,20 @@ # pyre-strict -"""Re-inplace contiguous ``slice_copy`` nodes as lightweight slices (#10917). - -Slice analog of :class:`ReplaceViewCopyWithViewPass`. Contiguous slices taken -along the outermost dimension with unit step alias a sub-region of the base -buffer and can be represented with :class:`_SliceSpec`, which shares the -base's ``mem_id`` but uses a computed byte ``mem_offset``. -""" +"""Replace safe, static, contiguous slice copies with sub-buffer aliases.""" import logging -from typing import Any, List, Optional +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, - determine_tensor_dynanism, dim_order_from_stride, TensorSpec, ) @@ -30,8 +27,6 @@ logger: logging.Logger = logging.getLogger(__name__) -_SLICE_OP = memory.slice - def _is_slice_copy(node: torch.fx.Node) -> bool: return node.op == "call_function" and node.target in ( @@ -40,290 +35,121 @@ def _is_slice_copy(node: torch.fx.Node) -> bool: ) -def _normalize_dim(dim: int, rank: Optional[int]) -> Optional[int]: - if isinstance(dim, int) and dim < 0: - if rank is None: - return None - dim = dim + rank - return dim +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 if ``node`` is an outermost-dim, unit-step ``slice_copy``.""" - if not _is_slice_copy(node): + """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 - - args = node.args - self_arg = args[0] - dim = args[1] if len(args) > 1 else 0 - step = args[4] if len(args) > 4 else 1 - - if step != 1: + 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 - - rank = None - self_val = self_arg.meta.get("val") if isinstance(self_arg, torch.fx.Node) else None - if self_val is not None and hasattr(self_val, "dim"): - rank = self_val.dim() - - dim = _normalize_dim(dim, rank) + if dim < 0 and isinstance(val, torch.Tensor): + dim += val.dim() return dim == 0 -def _slice_start_as_int(start: Optional[int]) -> int: - if start is None: - return 0 - return start - - -def _is_static_slice_argument(value: Any) -> bool: - """Whether a slice parameter is known when the program is lowered. - - ``memory.slice`` encodes a byte offset in the emitted tensor metadata, so - it cannot represent an FX node or a symbolic value that is resolved only - at runtime. Keep those slices as ``aten.slice_copy``. - """ - return value is None or isinstance(value, int) - - def _compute_slice_byte_offset(base: TensorSpec, dim: int, start: Optional[int]) -> int: - start_int = _slice_start_as_int(start) - if start_int < 0: - raise ValueError("memory.slice does not support negative slice starts.") - elem_size = torch._utils._element_size(base.dtype) - return start_int * base.stride[dim] * elem_size - - -def _is_aliasing_base(base: torch.fx.Node) -> bool: - """Whether ``base`` is itself an alias rather than a real allocation. - - ``memory.slice`` and ``memory.view`` nodes do not own storage, so a slice - taken from one has no concrete ``mem_offset`` to build on until the chain is - normalized down to the first real allocation. - """ - return base.op == "call_function" and base.target in (memory.slice, memory.view) - - -def _has_default_dim_order(spec: TensorSpec) -> bool: - """Whether ``spec`` has the standard contiguous dimension ordering. - - ``_SliceSpec`` computes a contiguous output stride. That is only a valid - alias for a dim-0 slice when the base itself has the default dim order. - """ - return spec.dim_order == dim_order_from_stride( - contiguous_stride_from_shape(torch.Size(spec.shape)) - ) - - -class _SliceSpec(TensorSpec): - """TensorSpec for a zero-copy slice into a contiguous base buffer.""" - - def __init__( - self, - base: TensorSpec, - shape: List[int], - dim: int, - start: Any, - ) -> None: - if base.is_sparse: - raise Exception( - "_SliceSpec can only be created from non-sparse TensorSpec." - ) - if base.layout != torch.strided: - raise Exception(f"_SliceSpec requires strided layout, got {base.layout}.") - - self._base = base - self._byte_offset = _compute_slice_byte_offset(base, dim, start) - self._unguarded_access = False - - self._self_fields = [ - "debug", - "__repr__", - "shape", - "stride", - "dim_order", - "shape_dynamism", - "nbytes", - "allocated_memory", - "is_dynamic_shape_tensor", - "is_static_shape_tensor", - "is_upper_bound_tensor", - "is_dynamic_unbound_tensor", - "mem_offset", - ] - self._base_fields = [ - "scalar_type", - "const", - "alignment", - "storage", - "requires_grad", - "layout", - "is_sparse", - "init_mem_planning_fields", - "realign", - "from_tensor", - "lifetime", - "mem_id", - "mem_obj_id", - "dtype", - "extra_tensor_info", - "device", - "device_index", - # Read by the memory planning algorithms (e.g. ``greedy``). A slice - # is never itself an in-place target, so it defers to its base. - "inplace_base", - ] - - self.shape = list(shape) - self.stride = contiguous_stride_from_shape(torch.Size(self.shape)) - self.dim_order = dim_order_from_stride(self.stride) - self.shape_dynamism = determine_tensor_dynanism(torch.Size(self.shape)) - - if self.shape_dynamism != base.shape_dynamism: - raise Exception( - f"_SliceSpec shape_dynamism {self.shape_dynamism} != base {base.shape_dynamism}" - ) - if self.dtype != base.dtype: - raise Exception(f"_SliceSpec dtype {self.dtype} != base {base.dtype}") - - def __getattribute__(self, name: str): # pyre-ignore - if name in [ - "_base", - "_self_fields", - "_base_fields", - "_byte_offset", - "_unguarded_access", - ]: - return object.__getattribute__(self, name) - - self_fields = object.__getattribute__(self, "_self_fields") - base_fields = object.__getattribute__(self, "_base_fields") - - if name == "mem_offset": - base = object.__getattribute__(self, "_base") - base_offset = base.mem_offset - if base_offset is None: - return None - byte_offset = object.__getattribute__(self, "_byte_offset") - return base_offset + byte_offset - - if name in self_fields: - if name in ("nbytes", "allocated_memory"): - return TensorSpec.__getattribute__(self, name) - return object.__getattribute__(self, name) - - if name in base_fields: - base = object.__getattribute__(self, "_base") - return object.__getattribute__(base, name) - - return object.__getattribute__(self, name) - - def __setattr__(self, name: str, val) -> None: # pyre-ignore - if name in [ - "_base", - "_self_fields", - "_base_fields", - "_byte_offset", - "_unguarded_access", - ]: - object.__setattr__(self, name, val) - return - - if hasattr(self, "_self_fields") and name in self._self_fields: - if name == "mem_offset": - raise Exception("_SliceSpec.mem_offset is computed from the base.") - object.__setattr__(self, name, val) - return - - if hasattr(self, "_base_fields") and name in self._base_fields: - object.__setattr__(self._base, name, val) - return - - object.__setattr__(self, name, val) + 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 contiguous ``slice_copy`` nodes with ``memory.slice``.""" - - def __init__(self) -> None: - super().__init__() + """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 - for node in module.graph.nodes: - if is_contiguous_slice_copy(node) and all( - u.op != "output" for u in node.users + 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 ): - base = node.args[0] - if ( - not isinstance(base, torch.fx.Node) - or "spec" not in base.meta - or not base.meta["spec"].is_static_shape_tensor - or not _has_default_dim_order(base.meta["spec"]) - ): - # Specs are populated by the lowering pipeline before this - # pass. Skip bare FX graphs so the pass remains safe to use - # in isolation as well. - continue - output_spec = node.meta.get("spec") - if ( - not isinstance(output_spec, TensorSpec) - or not output_spec.is_static_shape_tensor - ): - # A static base can still produce a dynamic slice when - # bounds originate in the runtime graph. - continue - if _is_aliasing_base(base): - # The base is itself an alias (a slice or a view), so it - # has no allocation of its own to offset from. Chaining - # offsets through it would require normalizing to the - # first real allocation first, so leave this as a copy. - continue - dim = node.args[1] if len(node.args) > 1 else 0 - start = node.args[2] if len(node.args) > 2 else None - end = node.args[3] if len(node.args) > 3 else None - step = node.args[4] if len(node.args) > 4 else 1 - if not all( - _is_static_slice_argument(arg) - for arg in (dim, start, end, step) - ): - # Do not turn symbolic bounds into a fixed byte offset. - # The regular slice_copy kernel handles runtime values. - continue - if base.op == "placeholder": - # Placeholders are externally owned buffers rather than - # intermediate allocations. Unlike memory.view, this - # pass has no runtime-kernel fallback, so only alias - # storage that memory planning owns. - continue - if _slice_start_as_int(start) < 0: - # Negative starts are relative to the end of the - # dimension. They cannot be expressed as a static - # offset without normalizing against the base shape. - continue - node.target = _SLICE_OP - shape = node.meta["val"].shape - node.meta["spec"] = _SliceSpec( - base.meta["spec"], list(shape), dim, start + 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], ) - n_replaced += 1 - + 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( - "ReplaceSliceCopyWithSlicePass: replaced %d slice_copy node(s) with %s.", - n_replaced, - _SLICE_OP, - ) + 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 not isinstance(module, torch.fx.GraphModule): - continue - for node in module.graph.nodes: - if node.op == "call_function" and node.target == _SLICE_OP: - assert isinstance(node.meta["spec"], _SliceSpec) + 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/tests/targets.bzl b/exir/tests/targets.bzl index 733c61dcd3e..9ecd23e9565 100644 --- a/exir/tests/targets.bzl +++ b/exir/tests/targets.bzl @@ -468,7 +468,11 @@ def define_common_targets(is_fbcode = False): "//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 ], ) diff --git a/exir/tests/test_replace_slice_copy_with_slice_pass.py b/exir/tests/test_replace_slice_copy_with_slice_pass.py index ff89f6d56e6..f3c777b8d42 100644 --- a/exir/tests/test_replace_slice_copy_with_slice_pass.py +++ b/exir/tests/test_replace_slice_copy_with_slice_pass.py @@ -6,11 +6,16 @@ # 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, @@ -18,6 +23,11 @@ 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 ( @@ -126,8 +136,7 @@ def forward(self, x): result = ReplaceSliceCopyWithSlicePass()(gm) self.assertFalse(result.modified) - def test_pass_skips_negative_start(self) -> None: - """Negative starts need shape-dependent normalization, so keep copying.""" + def test_pass_normalizes_negative_start(self) -> None: class M(torch.nn.Module): def forward(self, x): @@ -138,13 +147,179 @@ def forward(self, x): self._annotate_tensor_specs(gm) result = ReplaceSliceCopyWithSlicePass()(gm) - self.assertFalse(result.modified) - with self.assertRaises(ValueError): - _compute_slice_byte_offset( - next(n for n in gm.graph.nodes if n.op == "placeholder").meta["spec"], - 0, - -2, - ) + 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."""