Skip to content
Draft
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 @@ -28,7 +28,9 @@ def __init__(self, exported_program: ExportedProgram) -> None:

def _is_constant(self, node: torch.fx.Node) -> bool:
# Override fragile string match check with exported program check
return super()._is_constant(node) or is_param_node(self.exported_program, node)
exported_program = self.exported_program
assert exported_program is not None
return super()._is_constant(node) or is_param_node(exported_program, node)

def permute_subgraph(self, subgraph) -> bool:
# TABLE lookup inputs are already tied to the table layout.
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