Skip to content
17 changes: 17 additions & 0 deletions exir/emit/_emitter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
15 changes: 14 additions & 1 deletion exir/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
5 changes: 5 additions & 0 deletions exir/memory_planning.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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")

Expand Down
2 changes: 1 addition & 1 deletion exir/pass_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
14 changes: 14 additions & 0 deletions exir/passes/BUCK
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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 = [
Expand Down
1 change: 1 addition & 0 deletions exir/passes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading
Loading