Skip to content

Commit 06dc587

Browse files
Arm backend: Use FP64 conv in NSS quantized ref
Temporarily evaluate reference convolutions in FP64 for the random-data INT test, casting each result back to FP32 to reduce sensitivity to host-dependent bias accumulation. Keep calibration, export, and qtol unchanged. Assert that the reference graph contains all 14 expected convolutions so the FP64 override cannot silently become ineffective. Preserve the existing quantization-stage settings when replacing the stage. Validation: NSS random-data INT test and file-specific lint pass on Ubuntu. Temporary workaround for MLETORCH-2609. Authored with assistance from OpenAI Codex. Signed-off-by: Sangwon Ha <sangwon.ha@arm.com> Change-Id: Ibaed52de6562b0e990a4372d6390349183a7ee51
1 parent fe3edc7 commit 06dc587

1 file changed

Lines changed: 42 additions & 0 deletions

File tree

‎backends/arm/test/models/test_nss.py‎

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
REAL_AND_RANDOM_DATA,
2424
skip_if_frozen_release,
2525
)
26+
from executorch.backends.arm.test.tester.quantize import ArmQuantize
2627
from executorch.backends.arm.test.tester.test_pipeline import (
2728
EthosU55PipelineINT,
2829
EthosU85PipelineINT,
@@ -64,6 +65,32 @@ def __init__(self, *args, **kwargs):
6465
self.auto_encoder = AutoEncoderV1()
6566

6667

68+
class _Fp64ConvReference(torch.fx.Interpreter):
69+
def call_function(self, target, args, kwargs):
70+
if target == torch.ops.aten.conv2d.default:
71+
x, weight, bias, *options = args
72+
return target(
73+
x.double(),
74+
weight.double(),
75+
bias.double() if bias is not None else None,
76+
*options,
77+
**kwargs,
78+
).to(x.dtype)
79+
return super().call_function(target, args, kwargs)
80+
81+
82+
class _NssFp64ReferenceQuantize(ArmQuantize):
83+
# TODO(MLETORCH-2609): FP32 bias accumulation changes quantization decisions
84+
# across hosts. Use FP64 only for the quantized reference's convolutions.
85+
def run_artifact(self, inputs):
86+
conv_count = sum(
87+
node.op == "call_function" and node.target == torch.ops.aten.conv2d.default
88+
for node in self.artifact.graph.nodes
89+
)
90+
assert conv_count == 14, f"Expected 14 NSS conv2d nodes, found {conv_count}"
91+
return _Fp64ConvReference(self.artifact).run(*inputs)
92+
93+
6794
def nss() -> AutoEncoderV1:
6895
"""Get an instance of NSS with weights loaded."""
6996

@@ -200,6 +227,21 @@ def test_nss_tosa_INT(use_real_data, is_qat):
200227
)
201228
if use_real_data:
202229
_set_nss_calibration_samples(pipeline)
230+
elif not is_qat:
231+
quantize_stage = pipeline._stages[pipeline.find_pos("quantize")].args[0]
232+
pipeline.change_args(
233+
"quantize",
234+
_NssFp64ReferenceQuantize(
235+
quantizer=quantize_stage.quantizer,
236+
quantization_config=quantize_stage.quantization_config,
237+
calibrate=quantize_stage.calibrate,
238+
calibration_samples=quantize_stage.calibration_samples,
239+
is_qat=quantize_stage.is_qat,
240+
set_global=False,
241+
fold_quantize=quantize_stage.fold_quantize,
242+
dynamic_shapes=quantize_stage.dynamic_shapes,
243+
),
244+
)
203245
pipeline.run()
204246

205247

0 commit comments

Comments
 (0)