From 022cab3d399f80ddc3e117bf10183492c211fca5 Mon Sep 17 00:00:00 2001 From: Michiel Olieslagers Date: Fri, 11 Sep 2026 11:27:30 +0000 Subject: [PATCH] Arm backend: Test prequantized NSS weights Load the published NSS int8 checkpoint and restore its legacy PT2E state. Convert it without calibration or requantization, then exercise the resulting graph through the TOSA and VGF test pipelines. Authored with assistance from Codex. Signed-off-by: Michiel Olieslagers Change-Id: I871bbca4ac82c866a196dc43329125122e42b989 --- backends/arm/test/models/test_nss.py | 119 +++++++++++++++++++++++++++ 1 file changed, 119 insertions(+) diff --git a/backends/arm/test/models/test_nss.py b/backends/arm/test/models/test_nss.py index 8a5a5f0d5a6..1eb2805b086 100644 --- a/backends/arm/test/models/test_nss.py +++ b/backends/arm/test/models/test_nss.py @@ -7,6 +7,10 @@ import pytest import torch +from executorch.backends.arm.quantizer import ( + get_symmetric_quantization_config, + TOSAQuantizer, +) from executorch.backends.arm.scripts.neural_graphics_test_data import ( _NSS_INPUT_CHANNELS, iter_nss_test_calibration_samples, @@ -26,6 +30,10 @@ TosaPipelineINT, VgfPipeline, ) +from executorch.backends.arm.tosa import TosaSpecification +from executorch.backends.transforms.duplicate_dynamic_quant_chain import ( + DuplicateDynamicQuantChainPass, +) from huggingface_hub import hf_hub_download @@ -33,6 +41,13 @@ AutoEncoderV1, ) from torch.export import Dim +from torchao.quantization.pt2e import ( + allow_exported_model_train_eval, + FixedQParamsFakeQuantize, + FixedQParamsObserver, + move_exported_model_to_eval, +) +from torchao.quantization.pt2e.quantize_pt2e import convert_pt2e, prepare_qat_pt2e input_t = Tuple[torch.Tensor] # Input x @@ -71,6 +86,62 @@ def nss() -> AutoEncoderV1: return nss_model.auto_encoder +def prequantized_nss(inputs: input_t) -> torch.fx.GraphModule: + weights = hf_hub_download( # nosec B615 + repo_id="Arm/neural-super-sampling", + filename="nss_v1_0_1_high_int8.pt", + revision="main", + ) + checkpoint = torch.load( + weights, map_location=torch.device("cpu"), weights_only=True + )["model_state_dict"] + prefix = "autoencoder." + assert all(key.startswith(prefix) for key in checkpoint) + state_dict = {key.removeprefix(prefix): value for key, value in checkpoint.items()} + + exported = torch.export.export(nss().eval(), inputs, strict=True).module() + quantizer = TOSAQuantizer(TosaSpecification.create_from_string("TOSA-1.0+INT")) + quantizer.set_global( + get_symmetric_quantization_config(is_per_channel=True, is_qat=True) + ) + prepared = prepare_qat_pt2e(exported, quantizer) + + fixed_observers = { + key.removesuffix(".activation_post_process.scale") + for key in state_dict + if key.endswith(".activation_post_process.scale") + } + for name in fixed_observers: + scale = state_dict[f"{name}.activation_post_process.scale"].item() + zero_point = state_dict[f"{name}.activation_post_process.zero_point"].item() + observer = FixedQParamsObserver.with_args( + scale=scale, + zero_point=zero_point, + dtype=torch.int8, + qscheme=torch.per_tensor_affine, + quant_min=-127, + quant_max=127, + ) + prepared.set_submodule(name, FixedQParamsFakeQuantize(observer=observer)) + + parameter_keys = list(dict(prepared.named_parameters())) + lifted_parameter_keys = [ + key for key in state_dict if key.startswith("_param_constant") + ] + assert len(parameter_keys) == len(lifted_parameter_keys) + for lifted_key, parameter_key in zip( + lifted_parameter_keys, parameter_keys, strict=True + ): + state_dict[parameter_key] = state_dict.pop(lifted_key) + + prepared.load_state_dict(state_dict, strict=True) + move_exported_model_to_eval(prepared) + converted = convert_pt2e(prepared) + DuplicateDynamicQuantChainPass()(converted) + allow_exported_model_train_eval(converted) + return converted + + def example_inputs(): return load_nss_verification_inputs() @@ -132,6 +203,28 @@ def test_nss_tosa_INT(use_real_data, is_qat): pipeline.run() +@common.parametrize("use_real_data", input_test_data) +def test_nss_prequantized_tosa_INT(use_real_data): + inputs = example_inputs() if use_real_data else random_inputs() + pipeline = TosaPipelineINT[input_t]( + prequantized_nss(inputs), + inputs, + aten_op=[], + exir_op=[], + use_to_edge_transform_and_lower=True, + qtol=12 if use_real_data else 8, + ) + pipeline.pop_stage("quantize") + pipeline.pop_stage("check.quant_nodes") + pipeline.add_stage_after( + "export", + pipeline.tester.check, + ["torch.ops.quantized_decomposed.dequantize_per_tensor.default"], + suffix="prequant_nodes", + ) + pipeline.run() + + @pytest.mark.skip(reason="No support for aten_upsample_nearest2d_vec on U55") @common.XfailIfNoCorstone300 def test_nss_u55_INT(): @@ -202,3 +295,29 @@ def test_nss_vgf_INT(use_real_data, is_qat): if use_real_data: _set_nss_calibration_samples(pipeline) pipeline.run() + + +@common.SkipIfNoModelConverter +@common.parametrize("use_real_data", input_test_data) +def test_nss_prequantized_vgf_INT(use_real_data): + inputs = example_inputs() if use_real_data else random_inputs() + pipeline = VgfPipeline[input_t]( + prequantized_nss(inputs), + inputs, + aten_op=[], + exir_op=[], + use_to_edge_transform_and_lower=True, + run_on_vulkan_runtime=True, + quantize=True, + tosa_version="TOSA-1.0+INT", + qtol=12 if use_real_data else 8, + ) + pipeline.pop_stage("quantize") + pipeline.pop_stage("check.quant_nodes") + pipeline.add_stage_after( + "export", + pipeline.tester.check, + ["torch.ops.quantized_decomposed.dequantize_per_tensor.default"], + suffix="prequant_nodes", + ) + pipeline.run()