From 07bfe01c98625feefe20fc7ed75e01cb7affb1a6 Mon Sep 17 00:00:00 2001 From: Sangwon Ha Date: Wed, 23 Sep 2026 17:41:05 +0100 Subject: [PATCH] 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 Change-Id: Ibaed52de6562b0e990a4372d6390349183a7ee51 --- backends/arm/test/models/test_nss.py | 42 ++++++++++++++++++++++++++++ 1 file changed, 42 insertions(+) diff --git a/backends/arm/test/models/test_nss.py b/backends/arm/test/models/test_nss.py index 1eb2805b086..d738ac73f15 100644 --- a/backends/arm/test/models/test_nss.py +++ b/backends/arm/test/models/test_nss.py @@ -23,6 +23,7 @@ REAL_AND_RANDOM_DATA, skip_if_frozen_release, ) +from executorch.backends.arm.test.tester.quantize import ArmQuantize from executorch.backends.arm.test.tester.test_pipeline import ( EthosU55PipelineINT, EthosU85PipelineINT, @@ -64,6 +65,32 @@ def __init__(self, *args, **kwargs): self.auto_encoder = AutoEncoderV1() +class _Fp64ConvReference(torch.fx.Interpreter): + def call_function(self, target, args, kwargs): + if target == torch.ops.aten.conv2d.default: + x, weight, bias, *options = args + return target( + x.double(), + weight.double(), + bias.double() if bias is not None else None, + *options, + **kwargs, + ).to(x.dtype) + return super().call_function(target, args, kwargs) + + +class _NssFp64ReferenceQuantize(ArmQuantize): + # TODO(MLETORCH-2609): FP32 bias accumulation changes quantization decisions + # across hosts. Use FP64 only for the quantized reference's convolutions. + def run_artifact(self, inputs): + conv_count = sum( + node.op == "call_function" and node.target == torch.ops.aten.conv2d.default + for node in self.artifact.graph.nodes + ) + assert conv_count == 14, f"Expected 14 NSS conv2d nodes, found {conv_count}" + return _Fp64ConvReference(self.artifact).run(*inputs) + + def nss() -> AutoEncoderV1: """Get an instance of NSS with weights loaded.""" @@ -200,6 +227,21 @@ def test_nss_tosa_INT(use_real_data, is_qat): ) if use_real_data: _set_nss_calibration_samples(pipeline) + elif not is_qat: + quantize_stage = pipeline._stages[pipeline.find_pos("quantize")].args[0] + pipeline.change_args( + "quantize", + _NssFp64ReferenceQuantize( + quantizer=quantize_stage.quantizer, + quantization_config=quantize_stage.quantization_config, + calibrate=quantize_stage.calibrate, + calibration_samples=quantize_stage.calibration_samples, + is_qat=quantize_stage.is_qat, + set_global=False, + fold_quantize=quantize_stage.fold_quantize, + dynamic_shapes=quantize_stage.dynamic_shapes, + ), + ) pipeline.run()