Skip to content

Commit d0da5f8

Browse files
authored
Add support to PropagateSlice for custom unary/binary ops
Summary: As titled. Differential Revision: D119399660 Pull Request resolved: #22657
1 parent d1d5f40 commit d0da5f8

2 files changed

Lines changed: 78 additions & 4 deletions

File tree

backends/cadence/aot/reorder_ops.py

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1179,17 +1179,23 @@ class PropagateSlice(RemoveOrReplacePassInterface):
11791179
Handles any slice dim and any step size.
11801180
"""
11811181

1182-
def __init__(self) -> None:
1182+
def __init__(
1183+
self,
1184+
additional_unary_targets: Optional[list[EdgeOpOverload]] = None,
1185+
additional_binary_targets: Optional[list[EdgeOpOverload]] = None,
1186+
) -> None:
11831187
super().__init__()
1184-
elementwise_targets = [
1188+
unary_targets = [
11851189
exir_ops.edge.quantized_decomposed.quantize_per_tensor.default,
11861190
exir_ops.edge.cadence.quantize_per_tensor.default,
11871191
exir_ops.edge.quantized_decomposed.dequantize_per_tensor.default,
11881192
exir_ops.edge.cadence.dequantize_per_tensor.default,
1193+
*(additional_unary_targets or []),
11891194
]
11901195
binary_targets = [
11911196
exir_ops.edge.aten.add.Tensor,
11921197
exir_ops.edge.aten.mul.Tensor,
1198+
*(additional_binary_targets or []),
11931199
]
11941200
self._dispatch: dict[
11951201
EdgeOpOverload,
@@ -1198,7 +1204,7 @@ def __init__(self) -> None:
11981204
Callable[[torch.fx.Node, torch.fx.Node], bool],
11991205
],
12001206
] = {}
1201-
for t in elementwise_targets:
1207+
for t in unary_targets:
12021208
self._dispatch[t] = (
12031209
self._should_swap_elementwise,
12041210
self._swap_elementwise_slice,
@@ -1224,7 +1230,8 @@ def _should_swap_elementwise(
12241230
def _swap_elementwise_slice(
12251231
self, op_node: torch.fx.Node, slice_node: torch.fx.Node
12261232
) -> bool:
1227-
op_input = get_arg(op_node, "input", torch.fx.Node)
1233+
op_input = op_node.args[0]
1234+
assert isinstance(op_input, torch.fx.Node)
12281235
graph = slice_node.graph
12291236

12301237
slice_dim = get_arg(slice_node, "dim", int)

backends/cadence/aot/tests/test_reorder_ops_passes.py

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1358,6 +1358,73 @@ def test_unsupported_parent_not_swapped(self) -> None:
13581358

13591359
self.assertFalse(result.modified)
13601360

1361+
def test_swap_additional_unary_target(self) -> None:
1362+
x_data = torch.randn(4, 60, 1, 1)
1363+
builder = GraphBuilder()
1364+
x = builder.placeholder("x", x_data)
1365+
relu = builder.call_operator(exir_ops.edge.aten.relu.default, args=(x,))
1366+
sliced = builder.call_operator(
1367+
exir_ops.edge.aten.slice_copy.Tensor,
1368+
args=(relu, 0, 0, 4, 2),
1369+
)
1370+
builder.output([sliced])
1371+
gm = builder.get_graph_module()
1372+
1373+
result = transform_and_check_numerics(
1374+
gm,
1375+
(x_data,),
1376+
PropagateSlice(additional_unary_targets=[exir_ops.edge.aten.relu.default]),
1377+
)
1378+
1379+
self.assertTrue(result.modified)
1380+
slice_nodes = gm.graph.find_nodes(
1381+
op="call_function", target=exir_ops.edge.aten.slice_copy.Tensor
1382+
)
1383+
self.assertEqual(len(slice_nodes), 1)
1384+
relu_nodes = gm.graph.find_nodes(
1385+
op="call_function", target=exir_ops.edge.aten.relu.default
1386+
)
1387+
self.assertEqual(len(relu_nodes), 1)
1388+
self.assertIs(relu_nodes[0].args[0], slice_nodes[0])
1389+
self.assertEqual(list(relu_nodes[0].meta["val"].shape), [2, 60, 1, 1])
1390+
1391+
def test_swap_additional_binary_target(self) -> None:
1392+
lhs_data = torch.randn(1, 60, 1, 1)
1393+
rhs_data = torch.randn(4, 60, 1, 1)
1394+
builder = GraphBuilder()
1395+
lhs = builder.placeholder("lhs", lhs_data)
1396+
rhs = builder.placeholder("rhs", rhs_data)
1397+
sub = builder.call_operator(
1398+
exir_ops.edge.aten.sub.Tensor,
1399+
args=(lhs, rhs),
1400+
)
1401+
sliced = builder.call_operator(
1402+
exir_ops.edge.aten.slice_copy.Tensor,
1403+
args=(sub, 0, 0, 4, 2),
1404+
)
1405+
builder.output([sliced])
1406+
gm = builder.get_graph_module()
1407+
1408+
result = transform_and_check_numerics(
1409+
gm,
1410+
(lhs_data, rhs_data),
1411+
PropagateSlice(additional_binary_targets=[exir_ops.edge.aten.sub.Tensor]),
1412+
)
1413+
1414+
self.assertTrue(result.modified)
1415+
slice_nodes = gm.graph.find_nodes(
1416+
op="call_function", target=exir_ops.edge.aten.slice_copy.Tensor
1417+
)
1418+
self.assertEqual(len(slice_nodes), 1)
1419+
self.assertEqual(slice_nodes[0].args[0].name, "rhs")
1420+
sub_nodes = gm.graph.find_nodes(
1421+
op="call_function", target=exir_ops.edge.aten.sub.Tensor
1422+
)
1423+
self.assertEqual(len(sub_nodes), 1)
1424+
self.assertIs(sub_nodes[0].args[0], lhs.node)
1425+
self.assertIs(sub_nodes[0].args[1], slice_nodes[0])
1426+
self.assertEqual(list(sub_nodes[0].meta["val"].shape), [2, 60, 1, 1])
1427+
13611428
def test_swap_broadcast_mul_slice_on_broadcast_dim(self) -> None:
13621429
"""[1,60,1,1] * [4,1,1,1] → [4,60,1,1] → slice(dim=0, step=2)
13631430
Only the [4,1,1,1] input should be sliced."""

0 commit comments

Comments
 (0)