diff --git a/backends/arm/test/models/test_nfru.py b/backends/arm/test/models/test_nfru.py index e00302fe971..fdc9a991d0d 100644 --- a/backends/arm/test/models/test_nfru.py +++ b/backends/arm/test/models/test_nfru.py @@ -9,6 +9,10 @@ os.environ.setdefault("HF_HUB_DISABLE_XET", "1") import torch +from executorch.backends.arm.quantizer import ( + get_symmetric_quantization_config, + TOSAQuantizer, +) from executorch.backends.arm.scripts.neural_graphics_test_data import ( _NFRU_INPUT_CHANNELS, iter_nfru_test_calibration_samples, @@ -26,10 +30,19 @@ 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 from ng_model_gym.usecases.nfru.model.nfru_v1_nn import ( # type: ignore[import-not-found,import-untyped] NFRUAutoEncoder, ) +from torchao.quantization.pt2e import ( + allow_exported_model_train_eval, + move_exported_model_to_eval, +) +from torchao.quantization.pt2e.quantize_pt2e import convert_pt2e, prepare_qat_pt2e input_t = Tuple[torch.Tensor] # Input x @@ -60,6 +73,57 @@ def nfru() -> NFRUAutoEncoder: return model +def prequantized_nfru(inputs: input_t) -> torch.fx.GraphModule: + weights = hf_hub_download( # nosec B615 + repo_id="Arm/neural-frame-rate-upscaling", + filename="nfru_v1_int8.pt", + revision="main", + ) + checkpoint = torch.load( + weights, + map_location=torch.device("cpu"), + weights_only=True, + )["model_state_dict"] + prefix = "network.auto_encoder." + 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(nfru().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=False, is_qat=True) + ) + prepared = prepare_qat_pt2e(exported, quantizer) + + parameter_keys = list(dict(prepared.named_parameters())) + lifted_parameter_keys = [ + key for key in state_dict if key.startswith("_param_constant") + ] + buffer_keys = [ + key + for key, _ in prepared.named_buffers() + if not key.startswith("activation_post_process_") + ] + lifted_buffer_keys = [ + key for key in state_dict if key.startswith("_tensor_constant") + ] + assert len(parameter_keys) == len(lifted_parameter_keys) + assert len(buffer_keys) == len(lifted_buffer_keys) + for lifted_key, parameter_key in zip( + lifted_parameter_keys, parameter_keys, strict=True + ): + state_dict[parameter_key] = state_dict.pop(lifted_key) + for lifted_key, buffer_key in zip(lifted_buffer_keys, buffer_keys, strict=True): + state_dict[buffer_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_nfru_verification_inputs() @@ -119,6 +183,28 @@ def test_nfru_tosa_INT(use_real_data, is_qat): pipeline.run() +@common.parametrize("use_real_data", input_test_data) +def test_nfru_prequantized_tosa_INT(use_real_data): + inputs = example_inputs() if use_real_data else random_inputs() + pipeline = TosaPipelineINT[input_t]( + prequantized_nfru(inputs), + inputs, + aten_op=[], + exir_op=[], + atol=0.2, + qtol=2 if use_real_data else 1, + ) + 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() + + @common.parametrize("use_real_data", input_test_data) def test_nfru_tosa_INT_a16w8(use_real_data): pipeline = TosaPipelineINT[input_t]( @@ -179,6 +265,31 @@ def test_nfru_vgf_quant(use_real_data, is_qat): pipeline.run() +@common.SkipIfNoModelConverter +@common.parametrize("use_real_data", input_test_data) +def test_nfru_prequantized_vgf_INT(use_real_data): + inputs = example_inputs() if use_real_data else random_inputs() + pipeline = VgfPipeline[input_t]( + prequantized_nfru(inputs), + inputs, + aten_op=[], + exir_op=[], + tosa_version="TOSA-1.0+INT", + quantize=True, + atol=0.2, + qtol=2 if use_real_data else 1, + ) + 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() + + @common.SkipIfNoModelConverter @common.parametrize("use_real_data", input_test_data) def test_nfru_vgf_quant_a16w8(use_real_data):