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