Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 119 additions & 0 deletions backends/arm/test/models/test_nss.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -26,13 +30,24 @@
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.nss.model.model_blocks_v1 import ( # type: ignore[import-not-found,import-untyped]
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

Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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():
Expand Down Expand Up @@ -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()
Loading