Skip to content

Commit 96e87e3

Browse files
committed
Arm backend: Add PReLU decomposition pass
Decompose aten.prelu.default into clamp, mul, and add so the Arm backend can lower it through existing TOSA-supported primitives. Scalar PReLU weights are used directly; 1D per-channel weights are reshaped for broadcasting over the channel dimension. This intentionally uses the same lowering for U55 and U85 instead of a U85-specific where/select path. The pass only supports scalar or 1D PReLU weights, assumes per-channel weights apply to dim 1, and relies on the existing clamp/mul/add/view_copy backend and quantizer support for dtype and shape coverage. Signed-off-by: Per Held <per.held@arm.com> Change-Id: Ib3889c27eb51b7aa21c99439b8edadffc57e0e13
1 parent 8134bb2 commit 96e87e3

5 files changed

Lines changed: 209 additions & 0 deletions

File tree

backends/arm/_passes/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,7 @@
8181
from .decompose_meandim_pass import DecomposeMeanDimPass # noqa
8282
from .decompose_ne_pass import DecomposeNotEqualPass # noqa
8383
from .decompose_permute_for_u55_pass import DecomposePermuteForU55Pass # noqa
84+
from .decompose_prelu_pass import DecomposePReLUPass # noqa
8485
from .decompose_quant_nodes import DecomposeQuantNodesPass # noqa
8586
from .decompose_remainder_pass import DecomposeRemainderPass # noqa
8687
from .decompose_rnn_pass import DecomposeRnnPass # noqa

backends/arm/_passes/arm_pass_manager.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,7 @@
8181
DecomposeMeanDimPass,
8282
DecomposeNotEqualPass,
8383
DecomposePermuteForU55Pass,
84+
DecomposePReLUPass,
8485
DecomposeQuantNodesPass,
8586
DecomposeRemainderPass,
8687
DecomposeRnnPass,
@@ -579,6 +580,7 @@ def _tosa_pipeline(
579580
ReplaceScalarWithTensorByProfilePass(),
580581
RewriteLeLtToGeGtPass(),
581582
DecomposeLeakyReLUPass(), # Emits full_like so before ConvertFullLikeToFullPass
583+
DecomposePReLUPass(),
582584
ConvertFullLikeToFullPass(),
583585
MatchArgDtypePass(),
584586
UnsqueezeScalarPlaceholdersPass(exported_program),
@@ -731,6 +733,7 @@ def transform_for_annotation_pipeline(self, graph_module: GraphModule):
731733
DecomposeMeanDimPass(graph_module, self.tosa_spec, tfa_pass=True),
732734
DecomposeAdaptiveAvgPool2dPass(tfa_pass=True),
733735
DecomposeAvgPool2dPass(tfa_pass=True),
736+
DecomposePReLUPass(tfa_pass=True),
734737
]
735738
)
736739

Lines changed: 107 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,107 @@
1+
# Copyright 2026 Arm Limited and/or its affiliates.
2+
#
3+
# This source code is licensed under the BSD-style license found in the
4+
# LICENSE file in the root directory of this source tree.
5+
6+
from typing import Set, Type
7+
8+
import torch
9+
from executorch.backends.arm._passes import ArmOpTargetedPass
10+
from executorch.exir.dialects._ops import ops as exir_ops
11+
from executorch.exir.pass_base import ExportPass
12+
13+
edge_ops = (exir_ops.edge.aten.prelu.default,)
14+
torch_ops = (torch.ops.aten.prelu.default,)
15+
16+
17+
def _get_prelu_ops(op) -> tuple:
18+
if op in edge_ops:
19+
return (
20+
exir_ops.edge.aten.clamp.default,
21+
exir_ops.edge.aten.mul.Tensor,
22+
exir_ops.edge.aten.add.Tensor,
23+
exir_ops.edge.aten.view_copy.default,
24+
)
25+
if op in torch_ops:
26+
return (
27+
torch.ops.aten.clamp.default,
28+
torch.ops.aten.mul.Tensor,
29+
torch.ops.aten.add.Tensor,
30+
torch.ops.aten.view_copy.default,
31+
)
32+
raise RuntimeError(f"Can't get decomposition ops for op {op}")
33+
34+
35+
def _weight_shape(input_rank: int, weight_shape: torch.Size) -> tuple[int, ...] | None:
36+
weight_dims = tuple(int(dim) for dim in weight_shape)
37+
if len(weight_dims) == 0 or weight_dims == (1,):
38+
return None
39+
if len(weight_dims) != 1:
40+
raise RuntimeError(f"Unsupported PReLU weight shape: {weight_dims}")
41+
if input_rank < 2:
42+
raise RuntimeError(
43+
f"Per-channel PReLU weight requires input rank >= 2, got {input_rank}"
44+
)
45+
return (1, weight_dims[0], *([1] * (input_rank - 2)))
46+
47+
48+
class DecomposePReLUPass(ArmOpTargetedPass):
49+
"""Decompose PReLU into primitive TOSA-supported operations.
50+
51+
PReLU(x, weight) = max(0, x) + weight * min(0, x)
52+
53+
Example:
54+
%op1 = clamp(x,0,None) (equivalent to max(0,x))
55+
%op2 = clamp(x,None,0) (equivalent to min(0,x))
56+
%op3 = weight
57+
%op4 = mul(%op3,%op2)
58+
%op5 = add(%op1,%op4)
59+
60+
"""
61+
62+
_passes_required_after: Set[Type[ExportPass]] = set()
63+
target_ops = edge_ops + torch_ops
64+
check_allowed_to_transform = True
65+
66+
def call_operator(self, op, args, kwargs, meta):
67+
if (
68+
op not in self.target_ops
69+
or not self.allowed_to_transform(meta)
70+
or self._is_quantized_meta(meta)
71+
):
72+
return super().call_operator(op, args, kwargs, meta)
73+
74+
x, weight = args
75+
clamp, mul, add, view = _get_prelu_ops(op)
76+
77+
positive = super().call_operator(
78+
op=clamp, args=(x, 0, None), kwargs=kwargs, meta=meta, updated=True
79+
)
80+
negative = super().call_operator(
81+
op=clamp, args=(x, None, 0), kwargs=kwargs, meta=meta, updated=True
82+
)
83+
84+
input_rank = len(x.data.shape)
85+
reshape_shape = _weight_shape(input_rank, weight.data.shape)
86+
if reshape_shape is not None:
87+
weight = super().call_operator(
88+
op=view,
89+
args=(weight, reshape_shape),
90+
kwargs={},
91+
meta=meta,
92+
)
93+
94+
scaled_negative = super().call_operator(
95+
op=mul,
96+
args=(negative, weight),
97+
kwargs=kwargs,
98+
meta=meta,
99+
updated=True,
100+
)
101+
return super().call_operator(
102+
op=add,
103+
args=(positive, scaled_negative),
104+
kwargs=kwargs,
105+
meta=meta,
106+
updated=True,
107+
)

backends/arm/quantizer/quantizer_support.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -185,6 +185,7 @@ def check_pattern(cls, pattern):
185185
(torch.ops.aten.var.correction,),
186186
(torch.ops.aten.leaky_relu.default,),
187187
(torch.ops.aten.leaky_relu_.default,),
188+
(torch.ops.aten.prelu.default,),
188189
(torch.ops.aten.linalg_vector_norm.default,),
189190
(torch.ops.aten.log_softmax.int,),
190191
(torch.ops.aten.round.default,),
Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,97 @@
1+
# Copyright 2026 Arm Limited and/or its affiliates.
2+
#
3+
# This source code is licensed under the BSD-style license found in the
4+
# LICENSE file in the root directory of this source tree.
5+
6+
from typing import Tuple
7+
8+
import torch
9+
from executorch.backends.arm.test import common
10+
from executorch.backends.arm.test.tester.test_pipeline import (
11+
EthosU55PipelineINT,
12+
EthosU85PipelineINT,
13+
TosaPipelineFP,
14+
TosaPipelineINT,
15+
)
16+
17+
aten_op = "torch.ops.aten.prelu.default"
18+
exir_op = "executorch_exir_dialects_edge__ops_aten_prelu_default"
19+
input_t1 = Tuple[torch.Tensor]
20+
21+
22+
class PReLU(torch.nn.Module):
23+
def __init__(self, num_parameters: int = 1):
24+
super().__init__()
25+
self.activation = torch.nn.PReLU(num_parameters=num_parameters)
26+
27+
def forward(self, x: torch.Tensor):
28+
return self.activation(x)
29+
30+
test_data: dict[str, tuple[input_t1, int]] = {
31+
"scalar_2d": ((torch.randn(4, 5),), 1),
32+
"scalar_4d": ((torch.randn(1, 3, 8, 8),), 1),
33+
"per_channel_3d": ((torch.randn(2, 4, 5),), 4),
34+
"per_channel_4d": ((torch.randn(1, 3, 8, 8),), 3),
35+
}
36+
37+
38+
@common.parametrize("test_data", PReLU.test_data)
39+
def test_prelu_tosa_FP(test_data):
40+
data, num_parameters = test_data
41+
pipeline = TosaPipelineFP[input_t1](
42+
PReLU(num_parameters),
43+
data,
44+
[],
45+
use_to_edge_transform_and_lower=True,
46+
)
47+
pipeline.add_stage_after(
48+
"to_edge_transform_and_lower", pipeline.tester.check_not, [exir_op]
49+
)
50+
pipeline.run()
51+
52+
53+
@common.parametrize("test_data", PReLU.test_data)
54+
def test_prelu_tosa_INT(test_data):
55+
data, num_parameters = test_data
56+
pipeline = TosaPipelineINT[input_t1](
57+
PReLU(num_parameters),
58+
data,
59+
[],
60+
use_to_edge_transform_and_lower=True,
61+
)
62+
pipeline.add_stage_after(
63+
"to_edge_transform_and_lower", pipeline.tester.check_not, [exir_op]
64+
)
65+
pipeline.run()
66+
67+
68+
@common.parametrize("test_data", PReLU.test_data)
69+
@common.XfailIfNoCorstone300
70+
def test_prelu_u55_INT(test_data):
71+
data, num_parameters = test_data
72+
pipeline = EthosU55PipelineINT[input_t1](
73+
PReLU(num_parameters),
74+
data,
75+
[],
76+
use_to_edge_transform_and_lower=True,
77+
)
78+
pipeline.add_stage_after(
79+
"to_edge_transform_and_lower", pipeline.tester.check_not, [exir_op]
80+
)
81+
pipeline.run()
82+
83+
84+
@common.parametrize("test_data", PReLU.test_data)
85+
@common.XfailIfNoCorstone320
86+
def test_prelu_u85_INT(test_data):
87+
data, num_parameters = test_data
88+
pipeline = EthosU85PipelineINT[input_t1](
89+
PReLU(num_parameters),
90+
data,
91+
[],
92+
use_to_edge_transform_and_lower=True,
93+
)
94+
pipeline.add_stage_after(
95+
"to_edge_transform_and_lower", pipeline.tester.check_not, [exir_op]
96+
)
97+
pipeline.run()

0 commit comments

Comments
 (0)