diff --git a/backends/arm/test/models/test_nss.py b/backends/arm/test/models/test_nss.py index d738ac73f15..1eb2805b086 100644 --- a/backends/arm/test/models/test_nss.py +++ b/backends/arm/test/models/test_nss.py @@ -23,7 +23,6 @@ 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, @@ -65,32 +64,6 @@ 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.""" @@ -227,21 +200,6 @@ 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()