Skip to content

Commit 8f5a4ce

Browse files
committed
feat: added support for aten.rsqrt using new Neutron MLIR flow
1 parent 6f8194f commit 8f5a4ce

8 files changed

Lines changed: 159 additions & 0 deletions

File tree

backends/nxp/backend/edge_program_converter.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@
5656
exir_ops.edge.aten.permute_copy.default: PermuteCopyConverter, # noqa F405
5757
exir_ops.edge.aten.prelu.default: PReLUConverter, # noqa F405
5858
exir_ops.edge.aten.relu.default: ReLUConverter, # noqa F405
59+
exir_ops.edge.aten.rsqrt.default: RsqrtConverter, # noqa F405
5960
exir_ops.edge.aten.sigmoid.default: SigmoidConverter, # noqa F405
6061
exir_ops.edge.aten.slice_copy.Tensor: SliceTensorConverter, # noqa F405
6162
exir_ops.edge.aten._softmax.default: SoftmaxConverter, # 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
@@ -89,6 +89,9 @@
8989
from executorch.backends.nxp.backend.ir.converter.node_converters.ops_converters.relu_converter import (
9090
ReLUConverter,
9191
)
92+
from executorch.backends.nxp.backend.ir.converter.node_converters.ops_converters.rsqrt_converter import (
93+
RsqrtConverter,
94+
)
9295
from executorch.backends.nxp.backend.ir.converter.node_converters.ops_converters.sigmoid_converter import (
9396
SigmoidConverter,
9497
)
@@ -150,6 +153,7 @@
150153
"QDQPerTensorDequantizeConverter",
151154
"QDQQuantizeConverter",
152155
"ReLUConverter",
156+
"RsqrtConverter",
153157
"SigmoidConverter",
154158
"SliceTensorConverter",
155159
"SoftmaxConverter",
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
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.node_converter import (
9+
CustomDelegationOptions,
10+
NodeConverter,
11+
)
12+
from executorch.backends.nxp.backend.ir.lib.tflite.BuiltinOperator import (
13+
BuiltinOperator,
14+
)
15+
16+
from executorch.backends.nxp.backend.neutron_target_spec import NeutronTargetSpec
17+
from torch.fx import Node
18+
from torch.nn import Parameter
19+
20+
21+
class RsqrtConverter(NodeConverter):
22+
23+
@staticmethod
24+
def _is_supported_in_IR(
25+
node: Node,
26+
parameters_mapping: dict[str, Parameter],
27+
custom_delegation_options: CustomDelegationOptions,
28+
) -> bool:
29+
return True
30+
31+
@staticmethod
32+
def _is_supported_on_target(
33+
node: Node,
34+
neutron_target_spec: NeutronTargetSpec,
35+
parameters_mapping: dict[str, Parameter],
36+
custom_delegation_options: CustomDelegationOptions,
37+
) -> bool:
38+
if not NodeConverter.uses_quantization_type_for_io(
39+
node,
40+
supported_types=[torch.int8, torch.uint8],
41+
input_indices=[0],
42+
output_indices=[0],
43+
):
44+
return False
45+
46+
return True
47+
48+
def convert(self, node: Node):
49+
"""Convert the `aten.rsqrt.default` node to NeutronIR `RSQRT` operator.
50+
The ExecuTorch schema is:
51+
rsqrt(
52+
Tensor self
53+
) -> Tensor
54+
"""
55+
self.assert_convertible(node)
56+
57+
t_op = self._create_tflite_op_with_io_tensors(node)
58+
t_op.opcode_index = self.builder.op_code_index_for_op_type(
59+
BuiltinOperator.RSQRT
60+
)
61+
62+
self.builder.append_operators([t_op])

backends/nxp/neutron_partitioner.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -229,6 +229,7 @@ def tag_qdq_clusters(self, nodes: list[torch.fx.Node]):
229229
exir_ops.edge.aten.permute_copy.default: PermuteCopyConverter, # noqa F405
230230
exir_ops.edge.aten.prelu.default: PReLUConverter, # noqa F405
231231
exir_ops.edge.aten.relu.default: ReLUConverter, # noqa F405
232+
exir_ops.edge.aten.rsqrt.default: RsqrtConverter, # noqa F405
232233
exir_ops.edge.aten.sigmoid.default: SigmoidConverter, # noqa F405
233234
exir_ops.edge.aten.slice_copy.Tensor: SliceTensorConverter, # noqa F405
234235
exir_ops.edge.aten._softmax.default: SoftmaxConverter, # noqa F405

backends/nxp/quantizer/neutron_quantizer.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,7 @@
5252
ReluInPlacePattern,
5353
ReluPattern,
5454
ReshapePattern,
55+
RsqrtPattern,
5556
SharedSpecPattern,
5657
SigmoidPattern,
5758
SliceTensorPattern,
@@ -303,6 +304,7 @@ def __init__(self, neutron_target_spec: NeutronTargetSpec, is_qat: bool = False)
303304
OpQuantizer(ReluPattern(is_qat=is_qat), static_qconfig),
304305
OpQuantizer(ReluInPlacePattern(is_qat=is_qat), static_qconfig),
305306
OpQuantizer(ReshapePattern(is_qat=is_qat), static_qconfig),
307+
OpQuantizer(RsqrtPattern(is_qat=is_qat), static_qconfig),
306308
OpQuantizer(SigmoidPattern(is_qat=is_qat), static_qconfig),
307309
OpQuantizer(SliceTensorPattern(is_qat=is_qat), static_qconfig),
308310
OpQuantizer(SoftMaxPattern(is_qat=is_qat), static_qconfig),

backends/nxp/quantizer/patterns.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1089,6 +1089,13 @@ def partition_types(self):
10891089
return [torch.ops.aten.reshape.default]
10901090

10911091

1092+
class RsqrtPattern(SingleInputBasicPattern):
1093+
"""Quantizer for the `aten.rsqrt.default` operator."""
1094+
1095+
def partition_types(self):
1096+
return [torch.ops.aten.rsqrt.default]
1097+
1098+
10921099
class ViewPattern(SharedSpecPattern):
10931100
"""
10941101
Quantizer for View operator.
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 numpy as np
7+
8+
# noinspection PyUnusedImports
9+
import pytest
10+
import torch
11+
12+
from executorch.backends.nxp.tests.dataset_creator import RandomDatasetCreator
13+
from executorch.backends.nxp.tests.graph_verifier import DetailedGraphVerifier
14+
from executorch.backends.nxp.tests.model_output_comparator import (
15+
AllCloseOutputComparator,
16+
)
17+
from executorch.backends.nxp.tests.nsys_testing import lower_run_compare
18+
from executorch.backends.nxp.tests.ops_aliases import Rsqrt
19+
from executorch.backends.nxp.tests.use_qat import * # noqa F403
20+
21+
22+
@pytest.fixture(autouse=True)
23+
def reseed_model_per_test_run():
24+
torch.manual_seed(23)
25+
np.random.seed(23)
26+
27+
28+
class RsqrtModule(torch.nn.Module):
29+
def __init__(self):
30+
super().__init__()
31+
32+
def forward(self, x):
33+
return torch.rsqrt(x)
34+
35+
36+
class TestRsqrt:
37+
def assert_delegated(self, model, input_shape, mocker, request, use_qat=False):
38+
graph_verifier = DetailedGraphVerifier(
39+
mocker,
40+
expected_delegated_ops={Rsqrt: 1},
41+
expected_non_delegated_ops={},
42+
)
43+
44+
# Use positive-only values because rsqrt is only defined for x > 0.
45+
dataset_creator = RandomDatasetCreator(low=0.1, high=2.0)
46+
47+
# Allow a single quantization bit error in the output.
48+
comparator = AllCloseOutputComparator(atol=1)
49+
50+
lower_run_compare(
51+
model,
52+
input_shape,
53+
graph_verifier,
54+
request,
55+
dataset_creator,
56+
comparator,
57+
use_qat=use_qat,
58+
)
59+
60+
def test__basic_nsys_inference(self, mocker, request):
61+
input_shape = (2, 13, 7, 9)
62+
model = RsqrtModule()
63+
self.assert_delegated(model, input_shape, mocker, request)
64+
65+
def test__basic_nsys_inference__qat(self, mocker, request, use_qat):
66+
input_shape = (3, 5, 7, 11)
67+
model = RsqrtModule()
68+
self.assert_delegated(model, input_shape, mocker, request, use_qat=use_qat)
69+
70+
@pytest.mark.parametrize(
71+
"input_shape",
72+
[
73+
pytest.param((2,), id="1D"),
74+
pytest.param((2, 3), id="2D"),
75+
pytest.param((2, 3, 5), id="3D"),
76+
pytest.param((2, 3, 5, 7), id="4D"),
77+
],
78+
)
79+
def test__input_shapes(self, mocker, request, input_shape):
80+
model = RsqrtModule()
81+
self.assert_delegated(model, input_shape, mocker, request)

backends/nxp/tests/ops_aliases.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@
4848
QuantizePerChannel = exir_ops.edge.quantized_decomposed.quantize_per_channel.default
4949
QuantizePerTensor = exir_ops.edge.quantized_decomposed.quantize_per_tensor.default
5050
Relu = exir_ops.edge.aten.relu.default
51+
Rsqrt = exir_ops.edge.aten.rsqrt.default
5152
Sigmoid = exir_ops.edge.aten.sigmoid.default
5253
Slice = exir_ops.edge.aten.slice.Tensor
5354
SliceCopy = exir_ops.edge.aten.slice_copy.Tensor

0 commit comments

Comments
 (0)