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()