Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,9 @@


class RemovePermutesAroundElementwiseTosaOps(RemovePermutesAroundElementwiseOps):
# The base takes an optional program; this pass always has one.
exported_program: ExportedProgram

def __init__(self, exported_program: ExportedProgram) -> None:
super().__init__(
extra_permutable_ops={
Expand Down
1 change: 1 addition & 0 deletions backends/transforms/channels_last_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,7 @@ def _permute_copy(input, dims):
lib.impl("max_pool2d", _max_pool2d, "CompositeExplicitAutograd")
register_fake("channels_last::max_pool2d", _max_pool2d, lib=lib)


lib.define(
"grid_sampler_2d(Tensor input, Tensor grid, int interpolation_mode, "
"int padding_mode, bool align_corners) -> Tensor"
Expand Down
12 changes: 8 additions & 4 deletions backends/transforms/decompose_channels_last_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,10 @@
exir_ops.edge.channels_last.grid_sampler_2d.default: exir_ops.edge.aten.grid_sampler_2d.default,
}

_DIRECT_DECOMPOSITIONS = {
exir_ops.edge.channels_last.permute_copy.default: exir_ops.edge.aten.permute_copy.default,
}


class DecomposeChannelsLastPass(ExportPass):
"""Decompose channels_last dialect ops into permute + aten op + permute.
Expand All @@ -39,6 +43,10 @@ class DecomposeChannelsLastPass(ExportPass):
"""

def call_operator(self, op, args, kwargs, meta):
direct_op = _DIRECT_DECOMPOSITIONS.get(op)
if direct_op is not None:
return super().call_operator(direct_op, args, kwargs, meta)

aten_op = _DECOMPOSITIONS.get(op)
if aten_op is not None:
nchw_in = super().call_operator(
Expand Down Expand Up @@ -90,8 +98,4 @@ def call_operator(self, op, args, kwargs, meta):
meta,
)
return values, indices
if op == exir_ops.edge.channels_last.permute_copy.default:
return super().call_operator(
exir_ops.edge.aten.permute_copy.default, args, kwargs, meta
)
return super().call_operator(op, args, kwargs, meta)
37 changes: 37 additions & 0 deletions backends/transforms/fuse_transpose_or_permute_op_pairs_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

# pyre-unsafe

from collections import deque
from typing import Any, Callable, cast

import torch
Expand Down Expand Up @@ -41,6 +42,42 @@ class FuseTransposeOrPermuteOpPairsPass(FuseOpPairsAcrossBranchesPass):
exir_ops.edge.quantized_decomposed.dequantize_per_tensor.default,
}

def __init__(
self,
can_propagate: Callable[[torch.fx.Node], bool] | None = None,
) -> None:
super().__init__()
self.can_propagate = can_propagate

def get_fuse_candidates(
self,
producer: torch.fx.Node,
consumer_op_packets: set[EdgeOpOverloadPacket],
bypass_ops: set[EdgeOpOverload],
) -> list[torch.fx.Node]:
if self.can_propagate is None:
return super().get_fuse_candidates(
producer, consumer_op_packets, bypass_ops
)

users = deque(producer.users)
visited: set[torch.fx.Node] = set()
removal_candidates = []
while users:
user = users.popleft()
if user in visited:
continue
visited.add(user)
if user.target in bypass_ops:
if not self.can_propagate(user):
return []
users.extend(user.users)
elif self.can_fuse_for_chain(producer, user, consumer_op_packets):
removal_candidates.append(user)
else:
return []
return removal_candidates

def can_fuse_for_chain(
self,
producer: torch.fx.Node,
Expand Down
Loading
Loading