Skip to content

Commit 02096e1

Browse files
NXP backend: Enable Amax with new Neutron flow (#20628)
### Summary Add tests verifying correct support for amax by the Neutron backend using the new Neutron MLIR flow. ### Test plan Unit tests provided. cc @robert-kalmar
1 parent 10d3009 commit 02096e1

17 files changed

Lines changed: 610 additions & 11 deletions

backends/nxp/backend/edge_helper.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
from executorch.backends.nxp.tests.ops_aliases import (
1212
AddTensor,
13+
Amax,
1314
Amin,
1415
Cat,
1516
Clone,
@@ -47,6 +48,7 @@
4748
no_op_candidates = {
4849
AddTensor,
4950
Amin,
51+
Amax,
5052
MulTensor,
5153
PermuteCopy,
5254
SubTensor,

backends/nxp/backend/edge_program_converter.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
exir_ops.edge.aten._adaptive_avg_pool2d.default: AdaptiveAvgPool2dConverter, # noqa F405
3232
exir_ops.edge.aten.addmm.default: AddMMConverter, # noqa F405
3333
exir_ops.edge.aten.add.Tensor: AddTensorConverter, # noqa F405
34+
exir_ops.edge.aten.amax.default: AmaxConverter, # noqa F405
3435
exir_ops.edge.aten.amin.default: AminConverter, # noqa F405
3536
exir_ops.edge.aten.avg_pool2d.default: AvgPool2dConverter, # noqa F405
3637
exir_ops.edge.aten.bmm.default: BMMConverter, # noqa F405

backends/nxp/backend/ir/converter/node_converters/ops_converters/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,9 @@
1010
from executorch.backends.nxp.backend.ir.converter.node_converters.ops_converters.addmm_converter import (
1111
AddMMConverter,
1212
)
13+
from executorch.backends.nxp.backend.ir.converter.node_converters.ops_converters.amax_converter import (
14+
AmaxConverter,
15+
)
1316
from executorch.backends.nxp.backend.ir.converter.node_converters.ops_converters.amin_converter import (
1417
AminConverter,
1518
)
@@ -116,6 +119,7 @@
116119
"AdaptiveAvgPool2dConverter",
117120
"AddMMConverter",
118121
"AddTensorConverter",
122+
"AmaxConverter",
119123
"AminConverter",
120124
"AvgPool2dConverter",
121125
"BMMConverter",
Lines changed: 81 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,81 @@
1+
# Copyright 2026 NXP
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+
import torch
7+
8+
from executorch.backends.nxp.backend.ir.converter.conversion.common import OpsList
9+
from executorch.backends.nxp.backend.ir.converter.node_converter import (
10+
CustomDelegationOptions,
11+
NodeConverter,
12+
)
13+
from executorch.backends.nxp.backend.ir.converter.node_converters.shared.reduce_utils import (
14+
convert_axes_from_attribute,
15+
get_dim_and_handle_io_formats,
16+
get_reduce_node_attrs,
17+
)
18+
from executorch.backends.nxp.backend.ir.tflite_generator.builtin_options import (
19+
reduce_max_options,
20+
)
21+
from executorch.backends.nxp.backend.neutron_target_spec import NeutronTargetSpec
22+
from torch.fx import Node
23+
from torch.nn import Parameter
24+
25+
26+
class AmaxConverter(NodeConverter):
27+
28+
@staticmethod
29+
def _is_supported_on_target(
30+
node: Node,
31+
neutron_target_spec: NeutronTargetSpec,
32+
parameters_mapping: dict[str, Parameter],
33+
custom_delegation_options: CustomDelegationOptions,
34+
) -> bool:
35+
if not NodeConverter.uses_quantization_type_for_io(
36+
node,
37+
supported_types=[torch.int8, torch.uint8],
38+
input_indices=[0],
39+
output_indices=[0],
40+
):
41+
return False
42+
43+
return True
44+
45+
@staticmethod
46+
def _is_supported_in_IR(
47+
node: Node,
48+
parameters_mapping: dict[str, Parameter],
49+
custom_delegation_options: CustomDelegationOptions,
50+
) -> bool:
51+
if not NodeConverter._has_shared_q_params_if_quantized(node):
52+
return False
53+
54+
return True
55+
56+
def convert(self, node: Node):
57+
"""Convert the 'amax' operator to NeutronIR 'ReduceMax'.
58+
The ExecuTorch schema is:
59+
amax(
60+
Tensor self,
61+
int[1]? dim,
62+
bool keepdim=False,
63+
) -> Tensor
64+
"""
65+
self.assert_convertible(node)
66+
67+
dim, keepdim = get_reduce_node_attrs(node)
68+
69+
t_op = self._create_tflite_op_with_io_tensors(node)
70+
t_op.builtin_options = reduce_max_options.ReduceMax(keepdim)
71+
72+
ops = OpsList(middle_op=t_op)
73+
# dim default value is None, in that case no changes to dim or io_formats are needed and all dims are reduced
74+
dim = (
75+
get_dim_and_handle_io_formats(self.builder, ops, dim, keepdim)
76+
if dim is not None
77+
else None
78+
)
79+
80+
convert_axes_from_attribute(t_op, self.builder, dim)
81+
self.builder.append_operators(ops.flatten())

backends/nxp/backend/ir/converter/node_converters/ops_converters/amin_converter.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -70,7 +70,12 @@ def convert(self, node: Node):
7070
t_op.builtin_options = reduce_min_options.ReduceMin(keepdim)
7171

7272
ops = OpsList(middle_op=t_op)
73-
dim = get_dim_and_handle_io_formats(self.builder, ops, dim, keepdim)
73+
# dim default value is None, it that case no changes to dim or io_formats are needed and all dims are reduced
74+
dim = (
75+
get_dim_and_handle_io_formats(self.builder, ops, dim, keepdim)
76+
if dim is not None
77+
else None
78+
)
7479

7580
convert_axes_from_attribute(t_op, self.builder, dim)
7681
self.builder.append_operators(ops.flatten())

backends/nxp/backend/ir/converter/node_converters/ops_converters/sum_dim_int_list_converter.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -72,7 +72,12 @@ def convert(self, node: Node):
7272
t_op.builtin_options = sum_options.Sum(keepdim)
7373

7474
ops = OpsList(middle_op=t_op)
75-
dim = get_dim_and_handle_io_formats(self.builder, ops, dim, keepdim)
75+
# dim default value is None, it that case no changes to dim or io_formats are needed and all dims are reduced
76+
dim = (
77+
get_dim_and_handle_io_formats(self.builder, ops, dim, keepdim)
78+
if dim is not None and dim != []
79+
else None
80+
)
7681

7782
convert_axes_from_attribute(t_op, self.builder, dim)
7883
self.builder.append_operators(ops.flatten())

backends/nxp/backend/ir/converter/node_converters/shared/reduce_utils.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,7 @@ def _normalize_and_to_channel_last_dim(dim: list[int], rank: int) -> list[int]:
5959

6060

6161
def get_reduce_node_attrs(node: Node) -> tuple[list[int], bool]:
62-
dim = node.args[1]
62+
dim = node.args[1] if len(node.args) >= 2 else None
6363
keepdim = node.args[2] if len(node.args) >= 3 else False
6464
return dim, keepdim
6565

backends/nxp/backend/node_format_inference.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from executorch.backends.nxp.backend.edge_program_converter import functions_converters
1818
from executorch.backends.nxp.tests.ops_aliases import (
1919
AdaptiveAvgPool2D,
20+
Amax,
2021
Amin,
2122
AvgPool2D,
2223
Convolution,
@@ -67,6 +68,7 @@ class NodeFormatInference:
6768
ViewCopy,
6869
PermuteCopy,
6970
MeanDim,
71+
Amax,
7072
Amin,
7173
SumDimIntList,
7274
}
@@ -147,9 +149,10 @@ def _infer_format_of_nodes(self, node: Node):
147149
self._node_inputs[node][0], DataFormat.FORMATLESS
148150
)
149151

150-
elif op_type in [MeanDim, Amin, SumDimIntList]:
152+
elif op_type in [MeanDim, Amax, Amin, SumDimIntList]:
151153
# The operator schema is:
152-
# <reduce_op>(Tensor self, int[1]? dim, bool keepdim=False, *, ScalarType? dtype=None) -> Tensor
154+
# <reduce_op>(Tensor self, int[1]? dim, bool keepdim=False, *, ScalarType? dtype=None) -> Tensor or
155+
# <reduce_op>(Tensor self, int[1]? dim, bool keepdim=False) -> Tensor
153156
keep_dim = try_get_arg(node, 2) or False
154157
if keep_dim:
155158
# The operator preserves the rank, so we can handle it as an operator that can use any node format.

backends/nxp/neutron_partitioner.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -204,6 +204,7 @@ def tag_qdq_clusters(self, nodes: list[torch.fx.Node]):
204204
exir_ops.edge.aten._adaptive_avg_pool2d.default: AdaptiveAvgPool2dConverter, # noqa F405
205205
exir_ops.edge.aten.addmm.default: AddMMConverter, # noqa F405
206206
exir_ops.edge.aten.add.Tensor: AddTensorConverter, # noqa F405
207+
exir_ops.edge.aten.amax.default: AmaxConverter, # noqa F405
207208
exir_ops.edge.aten.amin.default: AminConverter, # noqa F405
208209
exir_ops.edge.aten.avg_pool2d.default: AvgPool2dConverter, # noqa F405
209210
exir_ops.edge.aten.bmm.default: BMMConverter, # noqa F405

backends/nxp/quantizer/neutron_quantizer.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
AdaptiveAvgPoolPattern,
1717
AddmmPattern,
1818
AddTensorPattern,
19+
AmaxPattern,
1920
AminPattern,
2021
AvgPool1DPattern,
2122
AvgPool2DPattern,
@@ -58,6 +59,7 @@
5859
SqueezePattern,
5960
SubTensorPattern,
6061
SumDimIntListPattern,
62+
SumPattern,
6163
TanhInPlacePattern,
6264
TanhPattern,
6365
TransposeIntPattern,
@@ -263,6 +265,7 @@ def __init__(self, neutron_target_spec: NeutronTargetSpec, is_qat: bool = False)
263265
OpQuantizer(AdaptiveAvgPoolPattern(is_qat=is_qat), static_qconfig),
264266
OpQuantizer(AddTensorPattern(is_qat=is_qat), static_qconfig),
265267
OpQuantizer(AddmmPattern(self, is_qat=is_qat), static_fc_qconfig),
268+
OpQuantizer(AmaxPattern(is_qat=is_qat), static_qconfig),
266269
OpQuantizer(AminPattern(is_qat=is_qat), static_qconfig),
267270
OpQuantizer(AvgPool1DPattern(is_qat=is_qat), static_qconfig),
268271
OpQuantizer(AvgPool2DPattern(is_qat=is_qat), static_qconfig),
@@ -304,6 +307,7 @@ def __init__(self, neutron_target_spec: NeutronTargetSpec, is_qat: bool = False)
304307
OpQuantizer(SqueezePattern(is_qat=is_qat), static_qconfig),
305308
OpQuantizer(SubTensorPattern(is_qat=is_qat), static_qconfig),
306309
OpQuantizer(SumDimIntListPattern(is_qat=is_qat), static_qconfig),
310+
OpQuantizer(SumPattern(is_qat=is_qat), static_qconfig),
307311
OpQuantizer(TanhPattern(is_qat=is_qat), static_qconfig),
308312
OpQuantizer(TanhInPlacePattern(is_qat=is_qat), static_qconfig),
309313
OpQuantizer(TransposeIntPattern(is_qat=is_qat), static_qconfig),

0 commit comments

Comments
 (0)