Skip to content

Commit 41d29de

Browse files
committed
Arm backend: Remove RESIZE that don't change shape
Sometimes, we get a TOSA RESIZE op stemming from upsample_bilinear2d or upsample_nearest2d that doesn't change the input & output resolution, it just passes from int8 to int32 and back to int8. Removing these resizes provides good uplift at runtime, especially if they operate on large tensors. Signed-off-by: George Gekov <george.gekov@arm.com> Change-Id: If0215efe423228f05c092795cbb28d13b61bcfc6
1 parent a7e7e99 commit 41d29de

4 files changed

Lines changed: 103 additions & 15 deletions

File tree

‎backends/arm/_passes/remove_noop_pass.py‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,8 @@ class RemoveNoopPass(ArmOpTargetedPass):
3232
exir_ops.edge.aten.copy.default,
3333
exir_ops.edge.aten.detach_copy.default,
3434
*_single_input_concat_ops,
35+
exir_ops.edge.aten.upsample_bilinear2d.vec,
36+
exir_ops.edge.aten.upsample_nearest2d.vec,
3537
exir_ops.backend.tosa.PAD.default,
3638
exir_ops.backend.tosa.SLICE.default,
3739
)
@@ -72,6 +74,24 @@ def call_operator(self, op, args, kwargs, meta, updated=False):
7274
return inputs[0]
7375
return super().call_operator(op, args, kwargs, meta, updated)
7476

77+
if op in (
78+
exir_ops.edge.aten.upsample_bilinear2d.vec,
79+
exir_ops.edge.aten.upsample_nearest2d.vec,
80+
):
81+
scale_factors = (
82+
args[3] if op == exir_ops.edge.aten.upsample_bilinear2d.vec else args[2]
83+
)
84+
if (
85+
isinstance(scale_factors, (list, tuple))
86+
and len(scale_factors) == 2
87+
and all(
88+
type(scale) in (int, float) and scale == 1.0
89+
for scale in scale_factors
90+
)
91+
):
92+
return args[0]
93+
return super().call_operator(op, args, kwargs, meta, updated)
94+
7595
if op == exir_ops.backend.tosa.PAD.default:
7696
padding = self._get_static_shape(args[1])
7797
# PAD is an identity when every before/after padding value is zero.

‎backends/arm/test/ops/test_upsample_bilinear2d.py‎

Lines changed: 0 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -34,8 +34,6 @@
3434
"rand_double_size": lambda: (torch.rand(2, 4, 8, 3), (16, 6), None, True),
3535
"rand_one_double_scale": lambda: (torch.rand(2, 4, 1, 1), None, 2.0, True),
3636
"rand_one_double_size": lambda: (torch.rand(2, 4, 1, 1), (2, 2), None, True),
37-
"rand_one_same_scale": lambda: (torch.rand(2, 4, 1, 1), None, 1.0, True),
38-
"rand_one_same_size": lambda: (torch.rand(2, 4, 1, 1), (1, 1), None, True),
3937
# Can't compare outputs as the rounding when selecting the nearest pixel is
4038
# different between PyTorch and TOSA. Just check the legalization went well.
4139
# TODO Improve the test infrastructure to support more in depth verification
@@ -86,18 +84,6 @@
8684
None,
8785
True,
8886
),
89-
"randn_one_same_scale_negative": lambda: (
90-
torch.randn(2, 4, 1, 1),
91-
None,
92-
1.0,
93-
True,
94-
),
95-
"randn_one_same_size_negative": lambda: (
96-
torch.randn(2, 4, 1, 1),
97-
(1, 1),
98-
None,
99-
True,
100-
),
10187
}
10288
test_data_suite_tosa_bf16 = {
10389
"randn_double_scale_bf16": lambda: (

‎backends/arm/test/ops/test_upsample_nearest2d.py‎

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,7 +41,6 @@
4141
"rand_double_size": lambda: (torch.rand(2, 4, 8, 3), (16, 6), None, True),
4242
"rand_one_double_scale": lambda: (torch.rand(2, 4, 1, 1), None, 2.0, True),
4343
"rand_one_double_size": lambda: (torch.rand(2, 4, 1, 1), (2, 2), None, True),
44-
"rand_one_same_scale": lambda: (torch.rand(2, 4, 1, 1), None, 1.0, True),
4544
"rand_one_same_size": lambda: (torch.rand(2, 4, 1, 1), (1, 1), None, True),
4645
# Can't compare outputs as the rounding when selecting the nearest pixel is
4746
# different between PyTorch and TOSA. Just check the legalization went well.

‎backends/arm/test/passes/test_remove_data_layout_noops.py‎

Lines changed: 83 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,89 @@ def test_keep_multi_input_concat():
118118
assert _count_target(result, exir_ops.edge.aten.cat.default) == 1
119119

120120

121+
def test_remove_identity_bilinear_upsample():
122+
graph = Graph()
123+
x = graph.placeholder("x")
124+
x.meta["val"] = torch.ones(1, 4, 768, 384)
125+
upsample = _call(
126+
graph,
127+
exir_ops.edge.aten.upsample_bilinear2d.vec,
128+
(x, None, False, [1.0, 1.0]),
129+
torch.ones(1, 4, 768, 384),
130+
)
131+
graph.output((upsample,))
132+
133+
result = _run_remove_noop(GraphModule(torch.nn.Module(), graph))
134+
135+
assert _count_target(result, exir_ops.edge.aten.upsample_bilinear2d.vec) == 0
136+
assert result.graph.output_node().args[0][0].op == "placeholder"
137+
138+
139+
def test_remove_identity_nearest_upsample():
140+
graph = Graph()
141+
x = graph.placeholder("x")
142+
x.meta["val"] = torch.ones(1, 4, 768, 384)
143+
upsample = _call(
144+
graph,
145+
exir_ops.edge.aten.upsample_nearest2d.vec,
146+
(x, None, [1.0, 1.0]),
147+
torch.ones(1, 4, 768, 384),
148+
)
149+
graph.output((upsample,))
150+
151+
result = _run_remove_noop(GraphModule(torch.nn.Module(), graph))
152+
153+
assert _count_target(result, exir_ops.edge.aten.upsample_nearest2d.vec) == 0
154+
assert result.graph.output_node().args[0][0].op == "placeholder"
155+
156+
157+
def test_keep_bilinear_upsample_with_rounded_identity_shape():
158+
graph = Graph()
159+
x = graph.placeholder("x")
160+
x.meta["val"] = torch.ones(1, 4, 3, 3)
161+
upsample = _call(
162+
graph,
163+
exir_ops.edge.aten.upsample_bilinear2d.vec,
164+
(x, None, False, [1.1, 1.1]),
165+
torch.ones(1, 4, 3, 3),
166+
)
167+
graph.output((upsample,))
168+
169+
result = _run_remove_noop(GraphModule(torch.nn.Module(), graph))
170+
171+
assert _count_target(result, exir_ops.edge.aten.upsample_bilinear2d.vec) == 1
172+
173+
174+
class _IdentityUpsampleModule(torch.nn.Module):
175+
def forward(self, x):
176+
return torch.nn.functional.interpolate(
177+
x, scale_factor=1.0, mode="bilinear", align_corners=False
178+
)
179+
180+
181+
def test_remove_identity_bilinear_upsample_backend_pipeline():
182+
exported_program = export(
183+
_IdentityUpsampleModule(), (torch.ones(1, 4, 768, 384),), strict=True
184+
)
185+
edge_program = to_edge(
186+
exported_program,
187+
compile_config=EdgeCompileConfig(_check_ir_validity=False),
188+
).exported_program()
189+
190+
assert (
191+
_count_target(
192+
edge_program.graph_module, exir_ops.edge.aten.upsample_bilinear2d.vec
193+
)
194+
== 1
195+
)
196+
197+
graph_module = ArmPassManager(
198+
TosaCompileSpec("TOSA-1.0+FP")
199+
).transform_to_backend_pipeline(edge_program, edge_program.graph_module)
200+
201+
assert _count_target(graph_module, exir_ops.backend.tosa.RESIZE.default) == 0
202+
203+
121204
def test_remove_full_slice_and_unused_shape_constants():
122205
graph_module = _tosa_data_layout_graph(
123206
exir_ops.backend.tosa.SLICE.default,

0 commit comments

Comments
 (0)