From 14f33bed130ca11d5f8f0d36e43424797cb56fdf Mon Sep 17 00:00:00 2001 From: Oscar Andersson Date: Thu, 9 Jul 2026 15:31:40 +0200 Subject: [PATCH] Arm backend: Add fused i8 warp+downsample example Adds example where int8 warp+downsample is fused into one custom shader. This pattern is present in rife model and such a fusion can reduce memory bandwidth significantly as it avoids processing pixels that are not used by downsampling. Signed-off-by: Oscar Andersson Change-Id: I0bec6fb6d3fe879900d744cb92f6c4aab16238ff --- examples/arm/QAT_example/qat_loop.py | 9 +- examples/arm/QAT_example/rife_vgf/__init__.py | 6 + .../arm/QAT_example/rife_vgf/extension.py | 189 ++++++ .../QAT_example/rife_vgf/passes/__init__.py | 8 + .../rewrite_warp_downsample_to_tosa_custom.py | 548 ++++++++++++++++++ ..._rewrite_warp_downsample_to_tosa_custom.py | 471 +++++++++++++++ .../QAT_example/rife_vgf/shaders/__init__.py | 241 ++++++++ ...ownsample2_sampler_int8_align_corners.glsl | 70 +++ ...mple2_sampler_int8_align_corners.spirv.b64 | 1 + ...ownsample4_sampler_int8_align_corners.glsl | 70 +++ ...mple4_sampler_int8_align_corners.spirv.b64 | 1 + ...ownsample8_sampler_int8_align_corners.glsl | 70 +++ ...mple8_sampler_int8_align_corners.spirv.b64 | 1 + 13 files changed, 1684 insertions(+), 1 deletion(-) create mode 100644 examples/arm/QAT_example/rife_vgf/__init__.py create mode 100644 examples/arm/QAT_example/rife_vgf/extension.py create mode 100644 examples/arm/QAT_example/rife_vgf/passes/__init__.py create mode 100644 examples/arm/QAT_example/rife_vgf/passes/rewrite_warp_downsample_to_tosa_custom.py create mode 100644 examples/arm/QAT_example/rife_vgf/passes/test_rewrite_warp_downsample_to_tosa_custom.py create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/__init__.py create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.glsl create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.spirv.b64 create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.glsl create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.spirv.b64 create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.glsl create mode 100644 examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.spirv.b64 diff --git a/examples/arm/QAT_example/qat_loop.py b/examples/arm/QAT_example/qat_loop.py index 04a3d6728df..80aa5a982fb 100644 --- a/examples/arm/QAT_example/qat_loop.py +++ b/examples/arm/QAT_example/qat_loop.py @@ -35,7 +35,7 @@ import numpy as np import torch import torch.nn.functional as F -from PIL import Image +from PIL import Image # type: ignore[import-untyped] from torch.export import export from torchao.quantization.pt2e import ( move_exported_model_to_eval, @@ -50,6 +50,10 @@ EXECUTORCH_ROOT = Path(__file__).resolve().parents[3] sys.path.insert(0, str(EXECUTORCH_ROOT)) +from examples.arm.QAT_example.rife_vgf import ( # noqa: E402 + configure_rife_vgf, + configure_rife_vgf_quantizer, +) from executorch.backends.arm.quantizer import ( # noqa: E402 get_uint8_io_quantization_config, get_vgf_snorm_quantization_config, @@ -665,6 +669,7 @@ def make_quantizer( quantizer = VgfQuantizer(compile_spec, use_composable_quantizer=True) global_config = get_vgf_snorm_quantization_config(is_qat=is_qat) quantizer.set_global(global_config) + configure_rife_vgf_quantizer(quantizer) if io_quantization == "int8": quantizer.set_io(global_config) @@ -801,6 +806,7 @@ def export_vgf_artifact( artifact_dir.mkdir(parents=True, exist_ok=True) compile_spec = make_vgf_compile_spec(artifact_dir) partitioner = VgfPartitioner(compile_spec) + configure_rife_vgf(partitioner) cast(Any, model).to(memory_format=torch.channels_last) inputs = ( inputs[0].contiguous(memory_format=torch.channels_last), @@ -1791,6 +1797,7 @@ def main() -> int: print(f"Data source: {source}") model = load_rife_model(args.model_root, args.checkpoint) + configure_rife_vgf() eager_channels_last_model = clone_model(model, channels_last=True) ptq_model, qat_model, quantization_reports = build_requested_quantized_models( args, diff --git a/examples/arm/QAT_example/rife_vgf/__init__.py b/examples/arm/QAT_example/rife_vgf/__init__.py new file mode 100644 index 00000000000..dc35a7af5e7 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/__init__.py @@ -0,0 +1,6 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from .extension import configure_rife_vgf, configure_rife_vgf_quantizer # noqa: F401 diff --git a/examples/arm/QAT_example/rife_vgf/extension.py b/examples/arm/QAT_example/rife_vgf/extension.py new file mode 100644 index 00000000000..9dab929f91c --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/extension.py @@ -0,0 +1,189 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import torch + +from executorch.backends.arm._passes import RewriteConvPass +from executorch.backends.arm._passes.arm_pass_manager import ( + _registered_pass_insertions, + register_pass_insertions_before, +) +from executorch.backends.arm.common.annotation_meta import ArmAnnotationInfo +from executorch.backends.arm.vgf import VgfPartitioner +from executorch.backends.cortex_m.quantizer_reporter import ( + QuantizerInfo, + QuantizerReporterUser, +) +from executorch.exir.pass_base import ExportPass +from torch._ops import OpOverload +from torchao.quantization.pt2e.quantizer import ( + FixedQParamsQuantizationSpec, + QuantizationAnnotation, + Quantizer, +) +from torchao.quantization.pt2e.quantizer.quantizer import Q_ANNOTATION_KEY + +from .passes import RewriteWarpDownsampleToTosaCustomPass + +_RIFE_LIBRARY: torch.library.Library | None = None +_RIFE_FAKE_IMPLS_REGISTERED = False + + +def _warp_downsample_fake( + image: torch.Tensor, flow: torch.Tensor, scale: int +) -> torch.Tensor: + del flow + return torch.empty( + ( + image.shape[0], + image.shape[1], + image.shape[2] // scale, + image.shape[3] // scale, + ), + dtype=image.dtype, + device=image.device, + ) + + +def _has_warp_downsample_op(scale: int) -> bool: + try: + getattr(torch.ops.rife, f"warp_downsample{scale}").default + except AttributeError: + return False + return True + + +def _register_warp_downsample_fake_impls() -> None: + global _RIFE_FAKE_IMPLS_REGISTERED + if _RIFE_FAKE_IMPLS_REGISTERED: + return + + for scale in (2, 4, 8): + + def _warp_downsample_fake_impl( + image: torch.Tensor, + flow: torch.Tensor, + scale: int = scale, + ) -> torch.Tensor: + return _warp_downsample_fake(image, flow, scale) + + try: + torch.library.register_fake(f"rife::warp_downsample{scale}")( + _warp_downsample_fake_impl + ) + except RuntimeError as error: + message = str(error) + if "already" not in message and "CompositeImplicitAutograd" not in message: + raise + + _RIFE_FAKE_IMPLS_REGISTERED = True + + +def _ensure_warp_downsample_ops_defined() -> None: + global _RIFE_LIBRARY + if _RIFE_LIBRARY is None: + _RIFE_LIBRARY = torch.library.Library("rife", "FRAGMENT") + missing_scales = [ + scale for scale in (2, 4, 8) if not _has_warp_downsample_op(scale) + ] + for scale in missing_scales: + _RIFE_LIBRARY.define( + f"warp_downsample{scale}(Tensor image, Tensor flow) -> Tensor" + ) + + +def _ensure_warp_downsample_ops_registered() -> None: + _ensure_warp_downsample_ops_defined() + _register_warp_downsample_fake_impls() + + +def _warp_downsample_target(scale: int) -> OpOverload: + _ensure_warp_downsample_ops_registered() + return getattr(torch.ops.rife, f"warp_downsample{scale}").default + + +def _warp_downsample_targets() -> tuple[OpOverload, ...]: + return tuple(_warp_downsample_target(scale) for scale in (2, 4, 8)) + + +def _warp_downsample_snorm_qspec() -> FixedQParamsQuantizationSpec: + return FixedQParamsQuantizationSpec( + dtype=torch.int8, + scale=1.0 / 127.0, + zero_point=0, + quant_min=-127, + quant_max=127, + qscheme=torch.per_tensor_symmetric, + is_dynamic=False, + ) + + +class _WarpDownsampleQuantizer(Quantizer, QuantizerReporterUser): + def __init__(self) -> None: + super().__init__() + QuantizerReporterUser.__init__(self) + self.targets = set(_warp_downsample_targets()) + self.snorm_qspec = _warp_downsample_snorm_qspec() + + def get_quantizer_info(self) -> QuantizerInfo: + return QuantizerInfo( + self.__class__.__name__, + "rife.warp_downsample{2,4,8}", + "rife_warp_downsample_snorm", + "examples.arm.QAT_example.rife_vgf", + ) + + def annotate(self, model: torch.fx.GraphModule) -> torch.fx.GraphModule: + for node in model.graph.nodes: + if ( + node.op != "call_function" + or node.target not in self.targets + or len(node.args) != 2 + ): + continue + image = node.args[0] + if not isinstance(image, torch.fx.Node): + continue + node.meta[Q_ANNOTATION_KEY] = QuantizationAnnotation( + input_qspec_map={image: self.snorm_qspec}, + output_qspec=self.snorm_qspec, + _annotated=True, + ) + meta_custom = node.meta.get("custom", {}) + meta_custom[ArmAnnotationInfo.CUSTOM_META_KEY] = ArmAnnotationInfo( + quantized=True + ) + node.meta["custom"] = meta_custom + self.report_accept([node]) + return model + + def validate(self, model: torch.fx.GraphModule) -> None: + return None + + +def _register_pass_before(target_pass_type: type, pass_: ExportPass) -> None: + existing_insertions = _registered_pass_insertions.get(target_pass_type) + if existing_insertions is not None and any( + isinstance(existing_pass, type(pass_)) + for existing_pass in existing_insertions.before_passes + ): + return + register_pass_insertions_before(target_pass_type, [pass_]) + + +def configure_rife_vgf(partitioner: VgfPartitioner | None = None) -> None: + """Enable the Practical-RIFE warp-downsample VGF extension.""" + _ensure_warp_downsample_ops_registered() + if partitioner is not None: + for target in _warp_downsample_targets(): + partitioner.register_custom_partition_op(target) + + _register_pass_before(RewriteConvPass, RewriteWarpDownsampleToTosaCustomPass()) + + +def configure_rife_vgf_quantizer(quantizer) -> None: + """Quantize RIFE warp-downsample image input/output as int8 SNORM.""" + _ensure_warp_downsample_ops_registered() + quantizer.add_quantizer(_WarpDownsampleQuantizer()) diff --git a/examples/arm/QAT_example/rife_vgf/passes/__init__.py b/examples/arm/QAT_example/rife_vgf/passes/__init__.py new file mode 100644 index 00000000000..8f60fd420b8 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/passes/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from .rewrite_warp_downsample_to_tosa_custom import ( # noqa: F401 + RewriteWarpDownsampleToTosaCustomPass, +) diff --git a/examples/arm/QAT_example/rife_vgf/passes/rewrite_warp_downsample_to_tosa_custom.py b/examples/arm/QAT_example/rife_vgf/passes/rewrite_warp_downsample_to_tosa_custom.py new file mode 100644 index 00000000000..4e92b5b2c42 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/passes/rewrite_warp_downsample_to_tosa_custom.py @@ -0,0 +1,548 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import math +import operator +from typing import Set, Type + +import torch +from examples.arm.QAT_example.rife_vgf.shaders import ( + build_warp_downsample_payload, + warp_downsample_operator_name, +) +from executorch.backends.arm._passes import ArmPass +from executorch.backends.arm._passes.arm_pass_utils import create_node +from executorch.backends.arm._passes.fold_qdq_with_annotated_qparams_pass import ( + get_input_qparams, + get_output_qparams, +) +from executorch.backends.arm._passes.quant_args import QuantArgs +from executorch.backends.arm.constants import NHWC_INVERSE_ORDER, NHWC_ORDER +from executorch.backends.arm.tosa.dialect.ops.custom import register_fake_tosa +from executorch.backends.arm.vgf.shaders.grid_sampler import ( + CUSTOM_SHADER_DOMAIN_NAME, + encode_payload, +) +from executorch.exir.dialects._ops import ops as exir_ops +from executorch.exir.pass_base import ExportPass, PassResult +from torch.fx.passes.shape_prop import _extract_tensor_metadata + + +def _target_name(target: object) -> str | None: + schema = getattr(target, "_schema", None) + schema_name = getattr(schema, "name", None) + overload_name = getattr(schema, "overload_name", None) + if isinstance(schema_name, str) and schema_name.startswith("rife::"): + name = schema_name.replace("::", ".", 1) + if overload_name in (None, "", "default"): + return f"{name}.default" + if isinstance(overload_name, str): + return f"{name}.{overload_name}" + + target_name = getattr(target, "__name__", None) + if isinstance(target_name, str): + if target_name.startswith("rife."): + return target_name + namespace = getattr(target, "namespace", None) + if isinstance(namespace, str): + return f"{namespace}.{target_name}" + return None + + +def _target_scale(target: object) -> int | None: + target_name = _target_name(target) + for scale in (2, 4, 8): + if target_name == f"rife.warp_downsample{scale}.default": + return scale + return None + + +def _warp_downsample_custom_fake_impl( + inputs, operator_name, domain_name, implementation_attrs +) -> list[torch.Tensor]: + assert domain_name == CUSTOM_SHADER_DOMAIN_NAME + _ = implementation_attrs + input_tensor, flow = inputs + _ = flow + scale = None + for candidate in (2, 4, 8): + if operator_name == warp_downsample_operator_name(candidate): + scale = candidate + break + assert scale is not None + return [ + torch.empty( + ( + input_tensor.shape[0], + input_tensor.shape[1] // scale, + input_tensor.shape[2] // scale, + input_tensor.shape[-1], + ), + dtype=input_tensor.dtype, + device=input_tensor.device, + ) + ] + + +for _scale in (2, 4, 8): + register_fake_tosa(warp_downsample_operator_name(_scale))( + _warp_downsample_custom_fake_impl + ) + + +def _set_fake_tensor_meta(node: torch.fx.Node, value) -> None: + node.meta["val"] = value + if isinstance(value, list): + if value: + node.meta["tensor_meta"] = _extract_tensor_metadata(value[0]) + else: + node.meta["tensor_meta"] = _extract_tensor_metadata(value) + + +def _permute_to_nhwc( + graph: torch.fx.Graph, + tensor: torch.fx.Node, + from_node: torch.fx.Node, +) -> torch.fx.Node: + nhwc_tensor = create_node( + graph, + op_target=exir_ops.edge.aten.permute_copy.default, + args=(tensor, list(NHWC_ORDER)), + from_node=from_node, + ) + _set_fake_tensor_meta( + nhwc_tensor, + exir_ops.edge.aten.permute_copy.default(tensor.meta["val"], list(NHWC_ORDER)), + ) + return nhwc_tensor + + +def _node_precedes(node: torch.fx.Node, reference: torch.fx.Node) -> bool: + if node.graph is not reference.graph: + return False + for candidate in node.graph.nodes: + if candidate is node: + return True + if candidate is reference: + return False + return False + + +def _is_supported_input(input_tensor: torch.fx.Node) -> bool: + value = input_tensor.meta.get("val") + return ( + isinstance(value, torch.Tensor) + and len(value.shape) == 4 + and int(value.shape[0]) == 1 + and int(value.shape[1]) in (3, 4) + and value.dtype == torch.int8 + ) + + +def _pad_c3_to_c4( + graph: torch.fx.Graph, + input_tensor: torch.fx.Node, + from_node: torch.fx.Node, +) -> torch.fx.Node: + input_val = input_tensor.meta["val"] + for user in input_tensor.users: + if ( + user.op == "call_function" + and user.target == exir_ops.edge.aten.constant_pad_nd.default + and len(user.args) >= 3 + and user.args[0] is input_tensor + and isinstance(user.args[1], (list, tuple)) + and list(user.args[1]) == [0, 0, 0, 0, 0, 1] + and user.args[2] == 0 + and _node_precedes(user, from_node) + ): + user_val = user.meta.get("val") + if ( + isinstance(user_val, torch.Tensor) + and len(user_val.shape) == 4 + and tuple(user_val.shape[:1]) == tuple(input_val.shape[:1]) + and int(user_val.shape[1]) == 4 + and tuple(user_val.shape[2:]) == tuple(input_val.shape[2:]) + ): + return user + + padded = create_node( + graph, + op_target=exir_ops.edge.aten.constant_pad_nd.default, + args=(input_tensor, [0, 0, 0, 0, 0, 1], 0), + from_node=from_node, + inherit_qparams=True, + ) + _set_fake_tensor_meta( + padded, + exir_ops.edge.aten.constant_pad_nd.default( + input_tensor.meta["val"], [0, 0, 0, 0, 0, 1], 0 + ), + ) + return padded + + +def _slice_c4_to_c3( + graph: torch.fx.Graph, + output_tensor: torch.fx.Node, + from_node: torch.fx.Node, +) -> torch.fx.Node: + sliced = create_node( + graph, + op_target=exir_ops.edge.aten.slice_copy.Tensor, + args=(output_tensor, 1, 0, 3), + from_node=from_node, + inherit_qparams=True, + ) + _set_fake_tensor_meta( + sliced, + exir_ops.edge.aten.slice_copy.Tensor(output_tensor.meta["val"], 1, 0, 3), + ) + return sliced + + +def _flow_arg_and_channel_offset( + flow: torch.fx.Node, +) -> tuple[torch.fx.Node, int]: + if ( + flow.op != "call_function" + or flow.target != exir_ops.edge.aten.slice_copy.Tensor + or len(flow.args) < 4 + ): + return flow, 0 + + flow_arg, dim, start, end, *step = flow.args + if ( + not isinstance(flow_arg, torch.fx.Node) + or dim != 1 + or not isinstance(start, int) + or not isinstance(end, int) + or (start, end) not in ((0, 2), (2, 4)) + or step not in ([], [1]) + or flow.kwargs.get("step", 1) != 1 + ): + return flow, 0 + + flow_arg_value = flow_arg.meta.get("val") + flow_value = flow.meta.get("val") + if ( + isinstance(flow_arg_value, torch.Tensor) + and isinstance(flow_value, torch.Tensor) + and len(flow_arg_value.shape) == 4 + and len(flow_value.shape) == 4 + and tuple(flow_arg_value.shape[:2]) == (1, 4) + and tuple(flow_value.shape[:2]) == (1, 2) + ): + return flow_arg, int(start) + return flow, 0 + + +def _restore_normalized_boundary_input(input_tensor: torch.fx.Node) -> torch.fx.Node: + """Undo delegate IO layout normalization for sampler-backed C4 inputs. + + NormalizeDelegateIOLayoutPass exposes channels-last tensors at delegate + boundaries by changing the placeholder shape to NHWC and inserting an + inverse permute back to NCHW. For warp_downsample, the custom shader rewrite + then inserts an NCHW->NHWC permute, and later cleanup can cancel the pair, + leaving a public COMBINED_IMAGE_SAMPLER input. Keep the public boundary as a + tensor by using the original placeholder as NCHW and letting the rewrite add + the shader-local NHWC permute. + + """ + if ( + input_tensor.op != "call_function" + or input_tensor.target != exir_ops.edge.aten.permute_copy.default + or len(input_tensor.args) < 2 + or len(input_tensor.users) != 1 + ): + return input_tensor + permutation = input_tensor.args[1] + if not isinstance(permutation, (list, tuple)) or list(permutation) != list( + NHWC_INVERSE_ORDER + ): + return input_tensor + + source = input_tensor.args[0] + if ( + not isinstance(source, torch.fx.Node) + or source.op != "placeholder" + or len(source.users) != 1 + ): + return input_tensor + + input_value = input_tensor.meta.get("val") + source_value = source.meta.get("val") + if ( + not isinstance(input_value, torch.Tensor) + or not isinstance(source_value, torch.Tensor) + or tuple(source_value.shape) + != tuple(input_value.shape[axis] for axis in NHWC_ORDER) + ): + return input_tensor + + source.meta["val"] = input_value + source.meta["tensor_meta"] = _extract_tensor_metadata(input_value) + for key in ("input_qparams", "output_qparams"): + if key in input_tensor.meta: + source.meta[key] = input_tensor.meta[key] + return source + + +def _is_supported_flow(flow: torch.fx.Node, flow_channel_offset: int) -> bool: + value = flow.meta.get("val") + return ( + isinstance(value, torch.Tensor) + and len(value.shape) == 4 + and int(value.shape[0]) == 1 + and int(value.shape[1]) >= flow_channel_offset + 2 + and int(value.shape[1]) in (2, 4) + and value.dtype == torch.int8 + ) + + +def _has_supported_shape_contract( + input_tensor: torch.fx.Node, + flow: torch.fx.Node, + output: torch.fx.Node, + scale: int, + flow_channel_offset: int, +) -> bool: + input_value = input_tensor.meta.get("val") + flow_value = flow.meta.get("val") + output_value = output.meta.get("val") + if not ( + isinstance(input_value, torch.Tensor) + and isinstance(flow_value, torch.Tensor) + and isinstance(output_value, torch.Tensor) + and len(input_value.shape) == 4 + and len(flow_value.shape) == 4 + and len(output_value.shape) == 4 + ): + return False + + input_n, input_c, input_h, input_w = (int(dim) for dim in input_value.shape) + flow_n, flow_c, flow_h, flow_w = (int(dim) for dim in flow_value.shape) + output_n, output_c, output_h, output_w = (int(dim) for dim in output_value.shape) + return ( + input_n == flow_n == output_n == 1 + and input_c == output_c + and flow_c >= flow_channel_offset + 2 + and flow_h == input_h + and flow_w == input_w + and input_h % scale == 0 + and input_w % scale == 0 + and output_h == input_h // scale + and output_w == input_w // scale + ) + + +def _uses_int8_snorm_qparams(qparams: QuantArgs) -> bool: + return ( + not qparams.per_channel + and math.isclose( + qparams.get_scale_per_tensor(), 1.0 / 127.0, rel_tol=1e-6, abs_tol=1e-9 + ) + and qparams.get_zp_per_tensor() == 0 + and qparams.qmin == -127 + and qparams.qmax == 127 + and qparams.dtype == torch.int8 + ) + + +def _uses_warp_downsample_int8_snorm_metadata(node: torch.fx.Node) -> bool: + try: + input_qparams = get_input_qparams(node) + output_qparams = get_output_qparams(node) + except ValueError: + return False + image_qparams = input_qparams.get(0) + if image_qparams is None or not output_qparams: + return False + return _uses_int8_snorm_qparams(image_qparams) and _uses_int8_snorm_qparams( + next(iter(output_qparams.values())) + ) + + +class RewriteWarpDownsampleToTosaCustomPass(ArmPass): + """Rewrite ``rife.warp_downsample{2,4,8}`` nodes to ``tosa.CUSTOM``.""" + + _passes_required_after: Set[Type[ExportPass]] = set() + + @staticmethod + def _encode_payload( + scale: int, + input_tensor: torch.fx.Node, + flow_tensor: torch.fx.Node, + output_tensor: torch.fx.Node, + output_shape: tuple[int, ...] | None = None, + output_dtype: torch.dtype | None = None, + flow_qparams: QuantArgs | None = None, + flow_channel_offset: int = 0, + ) -> list[int]: + input_val = input_tensor.meta.get("val") + flow_val = flow_tensor.meta.get("val") + output_val = output_tensor.meta.get("val") + if input_val is None or flow_val is None or output_val is None: + raise RuntimeError("warp_downsample node is missing tensor metadata") + if flow_qparams is None: + raise RuntimeError("int8 warp_downsample flow is missing input qparams") + payload = build_warp_downsample_payload( + scale=scale, + input_shape=tuple(input_val.shape), + output_shape=( + output_shape if output_shape is not None else tuple(output_val.shape) + ), + input_dtype=input_val.dtype, + output_dtype=output_dtype if output_dtype is not None else output_val.dtype, + flow_dtype=flow_val.dtype, + flow_scale=( + flow_qparams.get_scale_per_tensor() + if flow_qparams is not None + else None + ), + flow_zero_point=( + flow_qparams.get_zp_per_tensor() if flow_qparams is not None else None + ), + flow_channel_offset=flow_channel_offset, + ) + return encode_payload(payload) + + def call(self, graph_module): # noqa: C901 + modified = False + for node in list(graph_module.graph.nodes): + if node.op != "call_function": + continue + scale = _target_scale(node.target) + if scale is None: + continue + + if len(node.args) != 2: + raise RuntimeError("warp_downsample VGF rewrite requires two inputs") + input_tensor, flow = node.args + flow, flow_channel_offset = _flow_arg_and_channel_offset(flow) + try: + input_qparams = get_input_qparams(node) + except ValueError: + input_qparams = {} + flow_qparams = input_qparams.get(1) + use_quantized_image_payload = _uses_warp_downsample_int8_snorm_metadata( + node + ) + if not use_quantized_image_payload: + raise RuntimeError( + "warp_downsample int8 VGF rewrite requires SNORM qparams " + "scale=1/127, zp=0, qmin=-127, qmax=127 on input/output" + ) + if flow_qparams is None: + raise RuntimeError("int8 warp_downsample flow is missing input qparams") + output_dtype = torch.int8 + if not _is_supported_input(input_tensor): + raise RuntimeError( + "warp_downsample VGF rewrite requires int8 NCHW " + "input [1, 3 or 4, H, W]" + ) + if not _is_supported_flow(flow, flow_channel_offset): + raise RuntimeError( + "warp_downsample VGF rewrite requires int8 NCHW flow " + "[1, 2 or 4, H, W] with enough channels for flow_channel_offset" + ) + if not _has_supported_shape_contract( + input_tensor, flow, node, scale, flow_channel_offset + ): + raise RuntimeError( + "warp_downsample VGF rewrite requires input [1, C, H, W], " + "flow [1, 2 or 4, H, W], and output " + "[1, C, H / scale, W / scale]" + ) + + operator_name = warp_downsample_operator_name(scale) + output_channel_count = int(node.meta["val"].shape[1]) + with graph_module.graph.inserting_before(node): + input_tensor = _restore_normalized_boundary_input(input_tensor) + custom_input = input_tensor + if int(custom_input.meta["val"].shape[1]) == 3: + custom_input = _pad_c3_to_c4(graph_module.graph, custom_input, node) + custom_output_shape = tuple(node.meta["val"].shape) + if output_channel_count == 3: + custom_output_shape = ( + int(node.meta["val"].shape[0]), + 4, + int(node.meta["val"].shape[2]), + int(node.meta["val"].shape[3]), + ) + implementation_attrs = self._encode_payload( + scale, + custom_input, + flow, + node, + output_shape=custom_output_shape, + output_dtype=output_dtype, + flow_qparams=flow_qparams, + flow_channel_offset=flow_channel_offset, + ) + nhwc_input = _permute_to_nhwc( + graph_module.graph, + custom_input, + custom_input, + ) + custom_node = create_node( + graph_module.graph, + op_target=exir_ops.backend.tosa.CUSTOM.default, + args=([nhwc_input, flow],), + kwargs={ + "operator_name": operator_name, + "domain_name": CUSTOM_SHADER_DOMAIN_NAME, + "implementation_attrs": implementation_attrs, + }, + from_node=node, + inherit_qparams=True, + ) + with graph_module.graph.inserting_after(custom_node): + getitem_node = graph_module.graph.create_node( + "call_function", + operator.getitem, + args=(custom_node, 0), + kwargs={}, + ) + custom_output = _warp_downsample_custom_fake_impl( + [nhwc_input.meta["val"], flow.meta["val"]], + operator_name, + CUSTOM_SHADER_DOMAIN_NAME, + implementation_attrs, + )[0] + _set_fake_tensor_meta(custom_node, [custom_output]) + getitem_node.meta = dict(node.meta) + _set_fake_tensor_meta(getitem_node, custom_output) + + with graph_module.graph.inserting_after(getitem_node): + output = create_node( + graph_module.graph, + op_target=exir_ops.edge.aten.permute_copy.default, + args=(getitem_node, list(NHWC_INVERSE_ORDER)), + from_node=node, + ) + output.meta = dict(node.meta) + nchw_custom_output = exir_ops.edge.aten.permute_copy.default( + custom_output, list(NHWC_INVERSE_ORDER) + ) + _set_fake_tensor_meta(output, nchw_custom_output) + + if output_channel_count == 3: + with graph_module.graph.inserting_after(output): + output = _slice_c4_to_c3(graph_module.graph, output, node) + output.meta = dict(node.meta) + _set_fake_tensor_meta(output, node.meta["val"]) + + node.replace_all_uses_with(output) + graph_module.graph.erase_node(node) + modified = True + + if modified: + graph_module.graph.eliminate_dead_code() + graph_module.graph.lint() + graph_module.recompile() + graph_module = super().call(graph_module).graph_module + + return PassResult(graph_module, modified) diff --git a/examples/arm/QAT_example/rife_vgf/passes/test_rewrite_warp_downsample_to_tosa_custom.py b/examples/arm/QAT_example/rife_vgf/passes/test_rewrite_warp_downsample_to_tosa_custom.py new file mode 100644 index 00000000000..fdfe770f3ce --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/passes/test_rewrite_warp_downsample_to_tosa_custom.py @@ -0,0 +1,471 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import json + +import pytest +import torch +from examples.arm.QAT_example.rife_vgf import shaders as warp_downsample_shaders +from examples.arm.QAT_example.rife_vgf.extension import ( + _ensure_warp_downsample_ops_registered, +) +from examples.arm.QAT_example.rife_vgf.passes.rewrite_warp_downsample_to_tosa_custom import ( + _restore_normalized_boundary_input, + _target_scale, + RewriteWarpDownsampleToTosaCustomPass, +) +from executorch.backends.arm._passes.quant_args import QuantArgs +from executorch.backends.arm.constants import NHWC_INVERSE_ORDER +from executorch.backends.arm.tosa.specification import ( + TosaLoweringContext, + TosaSpecification, +) +from executorch.exir.dialects._ops import ops as exir_ops +from executorch.exir.pass_base import PassResult + + +class _Schema: + name = "rife::warp_downsample8" + overload_name = "" + + +class _EdgeTarget: + __name__ = "rife.warp_downsample8.default" + _schema = _Schema() + + +class _UnsupportedEdgeTarget: + __name__ = "rife.warp_downsample20.default" + + +@pytest.fixture(autouse=True) +def _stub_warp_downsample_shader_compile(monkeypatch) -> None: + monkeypatch.setattr( + warp_downsample_shaders, + "_compile_shader_source", + lambda source: "compiled_shader", + ) + + +def test_target_scale_matches_edge_op_name_exactly() -> None: + assert _target_scale(_EdgeTarget()) == 8 + assert _target_scale(_UnsupportedEdgeTarget()) is None + + +def _warp_downsample_graph( + *, + image_shape: tuple[int, int, int, int] = (1, 4, 16, 16), + flow_shape: tuple[int, int, int, int] = (1, 2, 16, 16), + output_shape: tuple[int, int, int, int] = (1, 4, 8, 8), + dtype: torch.dtype = torch.int8, + with_qparams: bool = True, +) -> torch.fx.GraphModule: + _ensure_warp_downsample_ops_registered() + graph = torch.fx.Graph() + image = graph.placeholder("image") + flow = graph.placeholder("flow") + image.meta["val"] = torch.empty(*image_shape, dtype=dtype) + flow.meta["val"] = torch.empty(*flow_shape, dtype=dtype) + warp = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image, flow), + ) + warp.meta["val"] = torch.empty(*output_shape, dtype=dtype) + if with_qparams: + snorm_qparams = QuantArgs(1.0 / 127.0, 0, -127, 127, torch.int8) + flow_qparams = QuantArgs(0.25, -3, -128, 127, torch.int8) + warp.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp.meta["output_qparams"] = {0: snorm_qparams} + graph.output(warp) + return torch.fx.GraphModule(torch.nn.Module(), graph) + + +def _rewrite_warp_downsample(graph_module: torch.fx.GraphModule) -> PassResult: + result = RewriteWarpDownsampleToTosaCustomPass()(graph_module) + assert result is not None + return result + + +def test_rewrite_warp_downsample_rejects_flow_spatial_mismatch() -> None: + graph_module = _warp_downsample_graph(flow_shape=(1, 2, 8, 16)) + + with pytest.raises(RuntimeError, match=r"flow \[1, 2 or 4, H, W\]"): + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + RewriteWarpDownsampleToTosaCustomPass()(graph_module) + + +def test_rewrite_warp_downsample_rejects_output_shape_mismatch() -> None: + graph_module = _warp_downsample_graph(output_shape=(1, 4, 7, 8)) + + with pytest.raises(RuntimeError, match="H / scale"): + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + RewriteWarpDownsampleToTosaCustomPass()(graph_module) + + +def test_rewrite_warp_downsample_rejects_fp32() -> None: + graph_module = _warp_downsample_graph(dtype=torch.float32) + + with pytest.raises(RuntimeError, match="requires int8 NCHW input"): + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + RewriteWarpDownsampleToTosaCustomPass()(graph_module) + + +def test_rewrite_warp_downsample_pads_c3_input_and_slices_c3_output() -> None: + graph_module = _warp_downsample_graph( + image_shape=(1, 3, 16, 16), + output_shape=(1, 3, 8, 8), + ) + + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + result = _rewrite_warp_downsample(graph_module) + + assert result.modified + custom_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.CUSTOM.default + ] + pad_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.edge.aten.constant_pad_nd.default + ] + slice_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.edge.aten.slice_copy.Tensor + ] + assert len(custom_nodes) == 1 + assert len(pad_nodes) == 1 + assert len(slice_nodes) == 1 + assert custom_nodes[0].meta["val"][0].shape == torch.Size((1, 8, 8, 4)) + assert slice_nodes[0].meta["val"].shape == torch.Size((1, 3, 8, 8)) + + +def test_rewrite_warp_downsample_allows_shared_flow() -> None: + _ensure_warp_downsample_ops_registered() + graph = torch.fx.Graph() + image0 = graph.placeholder("image0") + image1 = graph.placeholder("image1") + flow = graph.placeholder("flow") + image0.meta["val"] = torch.empty(1, 4, 16, 16, dtype=torch.int8) + image1.meta["val"] = torch.empty(1, 4, 16, 16, dtype=torch.int8) + flow.meta["val"] = torch.empty(1, 2, 16, 16, dtype=torch.int8) + warp0 = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image0, flow), + ) + warp0.meta["val"] = torch.empty(1, 4, 8, 8, dtype=torch.int8) + warp1 = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image1, flow), + ) + warp1.meta["val"] = torch.empty(1, 4, 8, 8, dtype=torch.int8) + snorm_qparams = QuantArgs(1.0 / 127.0, 0, -127, 127, torch.int8) + flow_qparams = QuantArgs(0.25, -3, -128, 127, torch.int8) + warp0.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp0.meta["output_qparams"] = {0: snorm_qparams} + warp1.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp1.meta["output_qparams"] = {0: snorm_qparams} + graph.output((warp0, warp1)) + graph_module = torch.fx.GraphModule(torch.nn.Module(), graph) + + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + result = _rewrite_warp_downsample(graph_module) + + assert result.modified + custom_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.CUSTOM.default + ] + assert len(custom_nodes) == 2 + + +def test_rewrite_warp_downsample_reuses_shared_c3_padding() -> None: + _ensure_warp_downsample_ops_registered() + graph = torch.fx.Graph() + image = graph.placeholder("image") + flow0 = graph.placeholder("flow0") + flow1 = graph.placeholder("flow1") + image.meta["val"] = torch.empty(1, 3, 16, 16, dtype=torch.int8) + flow0.meta["val"] = torch.empty(1, 2, 16, 16, dtype=torch.int8) + flow1.meta["val"] = torch.empty(1, 2, 16, 16, dtype=torch.int8) + warp0 = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image, flow0), + ) + warp0.meta["val"] = torch.empty(1, 3, 8, 8, dtype=torch.int8) + warp1 = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image, flow1), + ) + warp1.meta["val"] = torch.empty(1, 3, 8, 8, dtype=torch.int8) + snorm_qparams = QuantArgs(1.0 / 127.0, 0, -127, 127, torch.int8) + flow_qparams = QuantArgs(0.25, -3, -128, 127, torch.int8) + warp0.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp0.meta["output_qparams"] = {0: snorm_qparams} + warp1.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp1.meta["output_qparams"] = {0: snorm_qparams} + graph.output((warp0, warp1)) + graph_module = torch.fx.GraphModule(torch.nn.Module(), graph) + + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + result = _rewrite_warp_downsample(graph_module) + + pad_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.edge.aten.constant_pad_nd.default + ] + custom_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.CUSTOM.default + ] + assert len(pad_nodes) == 1 + assert len(custom_nodes) == 2 + + +def test_rewrite_warp_downsample_does_not_reuse_later_padding() -> None: + _ensure_warp_downsample_ops_registered() + graph = torch.fx.Graph() + image = graph.placeholder("image") + flow = graph.placeholder("flow") + image.meta["val"] = torch.empty(1, 3, 16, 16, dtype=torch.int8) + flow.meta["val"] = torch.empty(1, 2, 16, 16, dtype=torch.int8) + warp = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image, flow), + ) + warp.meta["val"] = torch.empty(1, 3, 8, 8, dtype=torch.int8) + snorm_qparams = QuantArgs(1.0 / 127.0, 0, -127, 127, torch.int8) + flow_qparams = QuantArgs(0.25, -3, -128, 127, torch.int8) + warp.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp.meta["output_qparams"] = {0: snorm_qparams} + later_pad = graph.call_function( + exir_ops.edge.aten.constant_pad_nd.default, + (image, [0, 0, 0, 0, 0, 1], 0), + ) + later_pad.meta["val"] = torch.empty(1, 4, 16, 16, dtype=torch.int8) + graph.output((warp, later_pad)) + graph_module = torch.fx.GraphModule(torch.nn.Module(), graph) + + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + result = _rewrite_warp_downsample(graph_module) + + pad_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.edge.aten.constant_pad_nd.default + ] + assert len(pad_nodes) == 2 + result.graph_module.graph.lint() + + +def test_rewrite_warp_downsample_folds_flow_slice_into_channel_offset( + monkeypatch, +) -> None: + captured_source = {} + + def _capture_shader(source: str) -> str: + captured_source["source"] = source + return "compiled_shader" + + monkeypatch.setattr( + warp_downsample_shaders, + "_compile_shader_source", + _capture_shader, + ) + _ensure_warp_downsample_ops_registered() + graph = torch.fx.Graph() + image = graph.placeholder("image") + flow = graph.placeholder("flow") + image.meta["val"] = torch.empty(1, 4, 16, 16, dtype=torch.int8) + flow.meta["val"] = torch.empty(1, 4, 16, 16, dtype=torch.int8) + sliced_flow = graph.call_function( + exir_ops.edge.aten.slice_copy.Tensor, + (flow, 1, 2, 4), + ) + sliced_flow.meta["val"] = torch.empty(1, 2, 16, 16, dtype=torch.int8) + warp = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image, sliced_flow), + ) + warp.meta["val"] = torch.empty(1, 4, 8, 8, dtype=torch.int8) + snorm_qparams = QuantArgs(1.0 / 127.0, 0, -127, 127, torch.int8) + flow_qparams = QuantArgs(0.25, -3, -128, 127, torch.int8) + warp.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp.meta["output_qparams"] = {0: snorm_qparams} + graph.output(warp) + graph_module = torch.fx.GraphModule(torch.nn.Module(), graph) + + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + result = _rewrite_warp_downsample(graph_module) + + custom_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.CUSTOM.default + ] + assert len(custom_nodes) == 1 + custom_node = custom_nodes[0] + _, custom_flow = custom_node.args[0] + assert custom_flow.name == "flow" + assert all( + node.target != exir_ops.edge.aten.slice_copy.Tensor + for node in result.graph_module.graph.nodes + ) + assert "kFlowChannelOffset = 2u" in captured_source["source"] + assert "uint[](0u, channel, uint(p.y), uint(p.x))" in captured_source["source"] + payload = json.loads(bytes(custom_node.kwargs["implementation_attrs"]).decode()) + assert payload["input_1_vkformat"] == "VK_FORMAT_R8_SINT" + + +def test_rewrite_warp_downsample_uses_int8_flow_payload(monkeypatch) -> None: + monkeypatch.setattr( + warp_downsample_shaders, + "_compile_shader_source", + lambda source: "compiled_shader", + ) + _ensure_warp_downsample_ops_registered() + graph = torch.fx.Graph() + image = graph.placeholder("image") + flow = graph.placeholder("flow") + image.meta["val"] = torch.empty(1, 4, 16, 16, dtype=torch.int8) + flow.meta["val"] = torch.empty(1, 2, 16, 16, dtype=torch.int8) + warp = graph.call_function( + torch.ops.rife.warp_downsample2.default, + (image, flow), + ) + warp.meta["val"] = torch.empty(1, 4, 8, 8, dtype=torch.int8) + snorm_qparams = QuantArgs(1.0 / 127.0, 0, -127, 127, torch.int8) + flow_qparams = QuantArgs(0.25, -3, -128, 127, torch.int8) + warp.meta["input_qparams"] = {0: snorm_qparams, 1: flow_qparams} + warp.meta["output_qparams"] = {0: snorm_qparams} + graph.output(warp) + graph_module = torch.fx.GraphModule(torch.nn.Module(), graph) + + with TosaLoweringContext(TosaSpecification.create_from_string("TOSA-1.0+INT")): + result = _rewrite_warp_downsample(graph_module) + + custom_nodes = [ + node + for node in result.graph_module.graph.nodes + if node.target == exir_ops.backend.tosa.CUSTOM.default + ] + assert len(custom_nodes) == 1 + custom_node = custom_nodes[0] + payload = json.loads(bytes(custom_node.kwargs["implementation_attrs"]).decode()) + assert payload["input_1_vkformat"] == "VK_FORMAT_R8_SINT" + assert payload["shader_code"] == "compiled_shader" + assert set(custom_node.meta["input_qparams"]) == {0, 1} + + +def _center_2x2_reference(image: torch.Tensor, scale: int) -> torch.Tensor: + offset0 = scale // 2 - 1 + offset1 = scale // 2 + output_h = image.shape[2] // scale + output_w = image.shape[3] // scale + result = torch.empty(image.shape[0], image.shape[1], output_h, output_w) + for y in range(output_h): + for x in range(output_w): + base_y = y * scale + base_x = x * scale + result[:, :, y, x] = 0.25 * ( + image[:, :, base_y + offset0, base_x + offset0] + + image[:, :, base_y + offset0, base_x + offset1] + + image[:, :, base_y + offset1, base_x + offset0] + + image[:, :, base_y + offset1, base_x + offset1] + ) + return result + + +def _full_window_avg_reference(image: torch.Tensor, scale: int) -> torch.Tensor: + return torch.nn.functional.avg_pool2d(image, kernel_size=scale, stride=scale) + + +@pytest.mark.parametrize("scale", (4, 8)) +def test_warp_downsample_reference_uses_center_2x2_not_full_window_average( + scale: int, +) -> None: + side = scale * 2 + image = torch.arange(float(side * side)).square().reshape(1, 1, side, side) + + reference = _center_2x2_reference(image, scale) + + offset0 = scale // 2 - 1 + offset1 = scale // 2 + expected = torch.empty_like(reference) + for y in range(2): + for x in range(2): + base_y = y * scale + base_x = x * scale + expected[:, :, y, x] = torch.tensor( + [ + [ + 0.25 + * ( + image[0, 0, base_y + offset0, base_x + offset0] + + image[0, 0, base_y + offset0, base_x + offset1] + + image[0, 0, base_y + offset1, base_x + offset0] + + image[0, 0, base_y + offset1, base_x + offset1] + ) + ] + ] + ) + + assert torch.equal(reference, expected) + assert not torch.equal(reference, _full_window_avg_reference(image, scale)) + + +def _normalized_boundary_graph( + *, shared_source: bool = False, shared_permute: bool = False +) -> tuple[torch.fx.GraphModule, torch.fx.Node, torch.fx.Node]: + graph = torch.fx.Graph() + source = graph.placeholder("source") + source.meta["val"] = torch.empty(1, 16, 16, 4) + inverse = graph.call_function( + exir_ops.edge.aten.permute_copy.default, + (source, list(NHWC_INVERSE_ORDER)), + ) + inverse.meta["val"] = torch.empty(1, 4, 16, 16) + if shared_source: + extra = graph.call_function( + exir_ops.edge.aten.alias_copy.default, + (source,), + ) + extra.meta["val"] = source.meta["val"] + if shared_permute: + extra = graph.call_function( + exir_ops.edge.aten.alias_copy.default, + (inverse,), + ) + extra.meta["val"] = inverse.meta["val"] + graph.output(inverse) + return torch.fx.GraphModule(torch.nn.Module(), graph), source, inverse + + +def test_restore_normalized_boundary_input_keeps_exclusive_boundary() -> None: + _, source, inverse = _normalized_boundary_graph() + + restored = _restore_normalized_boundary_input(inverse) + + assert restored is source + assert tuple(source.meta["val"].shape) == (1, 4, 16, 16) + + +@pytest.mark.parametrize("shared_source,shared_permute", ((True, False), (False, True))) +def test_restore_normalized_boundary_input_skips_shared_boundary( + shared_source: bool, shared_permute: bool +) -> None: + _, source, inverse = _normalized_boundary_graph( + shared_source=shared_source, shared_permute=shared_permute + ) + + restored = _restore_normalized_boundary_input(inverse) + + assert restored is inverse + assert tuple(source.meta["val"].shape) == (1, 16, 16, 4) diff --git a/examples/arm/QAT_example/rife_vgf/shaders/__init__.py b/examples/arm/QAT_example/rife_vgf/shaders/__init__.py new file mode 100644 index 00000000000..b8a9f1e7755 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/__init__.py @@ -0,0 +1,241 @@ +# Copyright 2026 Arm Limited and/or its affiliates. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import base64 +import shutil +import subprocess # nosec B404 +import tempfile +from pathlib import Path +from typing import Any + +from executorch.backends.arm.vgf.shaders.grid_sampler import ( + GRID_SAMPLER_2D_QUANTIZED_GRID_VK_FORMAT, + GRID_SAMPLER_2D_SAMPLER_INT8_VK_FORMAT, + GRID_SAMPLER_2D_SHADER_ENTRY_POINT, + GRID_SAMPLER_2D_SHADER_LANGUAGE, +) + +WARP_DOWNSAMPLE_OPERATOR_PREFIX = "rife.warp_downsample" +WARP_DOWNSAMPLE_WORKGROUP_SIZES = [8, 8, 1] +SUPPORTED_WARP_DOWNSAMPLE_SCALES = (2, 4, 8) + + +def warp_downsample_operator_name(scale: int) -> str: + if scale not in SUPPORTED_WARP_DOWNSAMPLE_SCALES: + raise ValueError(f"Unsupported warp_downsample scale {scale}") + return f"{WARP_DOWNSAMPLE_OPERATOR_PREFIX}{scale}" + + +def _dispatch_shape_for_output_shape(output_shape: tuple[int, ...]) -> list[int]: + if len(output_shape) != 4: + raise ValueError( + "warp_downsample output_shape must be rank 4 NCHW, " + f"got shape {output_shape}" + ) + output_batch = int(output_shape[0]) + output_height = int(output_shape[2]) + output_width = int(output_shape[3]) + group_x, group_y, group_z = WARP_DOWNSAMPLE_WORKGROUP_SIZES + return [ + (output_width + group_x - 1) // group_x, + (output_height + group_y - 1) // group_y, + (output_batch + group_z - 1) // group_z, + ] + + +def _sampler_config() -> dict[str, str]: + return { + "min_filter": "VK_FILTER_LINEAR", + "mag_filter": "VK_FILTER_LINEAR", + "address_mode_u": "VK_SAMPLER_ADDRESS_MODE_CLAMP_TO_EDGE", + "address_mode_v": "VK_SAMPLER_ADDRESS_MODE_CLAMP_TO_EDGE", + "border_color": "VK_BORDER_COLOR_FLOAT_TRANSPARENT_BLACK", + } + + +def _format_float(value: float) -> str: + return format(float(value), ".9g") + + +def _flow_shader_source( + *, + scale: int, + flow_dtype: Any | None, + flow_scale: float | None, + flow_zero_point: int | None, + flow_channel_offset: int, +) -> str: + offset0 = scale // 2 - 1 + offset1 = scale // 2 + if str(flow_dtype) != "torch.int8": + raise ValueError("warp_downsample flow payload supports only int8") + if flow_scale is None or flow_zero_point is None: + raise ValueError("int8 flow requires flow_scale and flow_zero_point") + read_flow_value = f""" int8_t value[1]; + tensorReadARM(flow, coords, value); + return (float(value[0]) - float({int(flow_zero_point)})) * {_format_float(flow_scale)};""" + return f"""// Copyright 2026 Arm Limited and/or its affiliates. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#version 450 +#extension GL_ARM_tensors : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +layout(set = 0, binding = 0) uniform sampler2D inputImage; +layout(set = 0, binding = 1) uniform tensorARM flow; +layout(set = 0, binding = 2, rgba8_snorm) uniform writeonly image2D outImage; + +layout(local_size_x = 8, local_size_y = 8, local_size_z = 1) in; + +const int kScale = {scale}; +const int kOffset0 = {offset0}; +const int kOffset1 = {offset1}; +const uint kFlowChannelOffset = {int(flow_channel_offset)}u; + +float readFlowChannel(ivec2 p, uint channel) {{ + uint coords[4] = uint[](0u, channel, uint(p.y), uint(p.x)); +{read_flow_value} +}} + +vec2 readFlowXY(ivec2 p) {{ + return vec2( + readFlowChannel(p, kFlowChannelOffset), + readFlowChannel(p, kFlowChannelOffset + 1u)); +}} + +vec2 alignCornersUv(vec2 gridXY) {{ + vec2 inputSize = vec2(textureSize(inputImage, 0)); + vec2 texel = (gridXY + vec2(1.0)) * vec2(0.5) * (inputSize - vec2(1.0)); + return (texel + vec2(0.5)) / inputSize; +}} + +vec2 baseGridXY(ivec2 p, ivec2 fullSize) {{ + return vec2( + (2.0 * float(p.x)) / float(fullSize.x - 1) - 1.0, + (2.0 * float(p.y)) / float(fullSize.y - 1) - 1.0); +}} + +vec4 sampleWarped(ivec2 p, ivec2 fullSize) {{ + vec2 flowXY = readFlowXY(p); + vec2 gridXY = baseGridXY(p, fullSize) + flowXY; + return texture(inputImage, alignCornersUv(gridXY)); +}} + +void main() {{ + ivec2 outSize = imageSize(outImage); + ivec2 gid = ivec2(gl_GlobalInvocationID.xy); + if (gid.x >= outSize.x || gid.y >= outSize.y) {{ + return; + }} + + ivec2 fullSize = textureSize(inputImage, 0); + ivec2 base = gid * kScale; + ivec2 p00 = base + ivec2(kOffset0, kOffset0); + ivec2 p01 = base + ivec2(kOffset1, kOffset0); + ivec2 p10 = base + ivec2(kOffset0, kOffset1); + ivec2 p11 = base + ivec2(kOffset1, kOffset1); + + vec4 value = 0.25 * ( + sampleWarped(p00, fullSize) + + sampleWarped(p01, fullSize) + + sampleWarped(p10, fullSize) + + sampleWarped(p11, fullSize)); + imageStore(outImage, gid, value); +}} +""" + + +def _compile_shader_source(source: str) -> str: + glslc = shutil.which("glslc") + if glslc is None: + raise RuntimeError("glslc is required to compile RIFE VGF custom shaders") + with tempfile.TemporaryDirectory() as tmpdir: + source_path = Path(tmpdir) / "warp_downsample.glsl" + spirv_path = Path(tmpdir) / "warp_downsample.spv" + source_path.write_text(source, encoding="utf-8") + subprocess.run( # nosec B603 + [glslc, "-fshader-stage=compute", str(source_path), "-o", str(spirv_path)], + check=True, + ) + return base64.b64encode(spirv_path.read_bytes()).decode("ascii") + + +def _shader_code( + *, + scale: int, + flow_dtype: Any | None, + flow_scale: float | None, + flow_zero_point: int | None, + flow_channel_offset: int, +) -> str: + return _compile_shader_source( + _flow_shader_source( + scale=scale, + flow_dtype=flow_dtype, + flow_scale=flow_scale, + flow_zero_point=flow_zero_point, + flow_channel_offset=flow_channel_offset, + ) + ) + + +def build_warp_downsample_payload( + scale: int, + input_shape: tuple[int, ...], + output_shape: tuple[int, ...], + input_dtype: Any, + output_dtype: Any | None = None, + flow_dtype: Any | None = None, + flow_scale: float | None = None, + flow_zero_point: int | None = None, + flow_channel_offset: int = 0, +) -> dict[str, Any]: + if scale not in SUPPORTED_WARP_DOWNSAMPLE_SCALES: + raise ValueError(f"Unsupported warp_downsample scale {scale}") + if output_dtype is None: + output_dtype = input_dtype + if str(input_dtype) != "torch.int8" or str(output_dtype) != "torch.int8": + raise ValueError( + "warp_downsample supports only matching int8 RGBA image payloads" + ) + if len(input_shape) != 4 or int(input_shape[0]) != 1 or int(input_shape[1]) != 4: + raise ValueError( + "warp_downsample currently requires NCHW input shape [1, 4, H, W]" + ) + if str(flow_dtype) != "torch.int8": + raise ValueError("warp_downsample supports only int8 flow payloads") + shader_code = _shader_code( + scale=scale, + flow_dtype=flow_dtype, + flow_scale=flow_scale, + flow_zero_point=flow_zero_point, + flow_channel_offset=flow_channel_offset, + ) + return { + "entry_point": GRID_SAMPLER_2D_SHADER_ENTRY_POINT, + # Current runtime consumes this field as dispatch counts, not local + # shader workgroup size. The shader uses an 8x8 output-space work + # volume per workgroup. + "workgroup_sizes": _dispatch_shape_for_output_shape(output_shape), + "shader_language": GRID_SAMPLER_2D_SHADER_LANGUAGE, + "shader_code": shader_code, + "input_0_binding": 0, + "input_0_descriptorset": 0, + "input_0_type": "Image", + "input_0_vkformat": GRID_SAMPLER_2D_SAMPLER_INT8_VK_FORMAT, + "input_0_vkdescriptortype": "VK_DESCRIPTOR_TYPE_COMBINED_IMAGE_SAMPLER", + "input_0_sampler": _sampler_config(), + "input_1_binding": 1, + "input_1_descriptorset": 0, + "input_1_type": "Tensor", + "input_1_vkformat": GRID_SAMPLER_2D_QUANTIZED_GRID_VK_FORMAT, + "input_1_vkdescriptortype": "VK_DESCRIPTOR_TYPE_TENSOR_ARM", + "output_0_binding": 2, + "output_0_descriptorset": 0, + "output_0_type": "Image", + "output_0_vkformat": GRID_SAMPLER_2D_SAMPLER_INT8_VK_FORMAT, + "output_0_vkdescriptortype": "VK_DESCRIPTOR_TYPE_STORAGE_IMAGE", + } diff --git a/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.glsl b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.glsl new file mode 100644 index 00000000000..bde54186286 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.glsl @@ -0,0 +1,70 @@ +// Copyright 2026 Arm Limited and/or its affiliates. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#version 450 +#extension GL_ARM_tensors : require + +layout(set = 0, binding = 0) uniform sampler2D inputImage; +layout(set = 0, binding = 1) uniform tensorARM flow; +layout(set = 0, binding = 2, rgba8_snorm) uniform writeonly image2D outImage; + +layout(local_size_x = 8, local_size_y = 8, local_size_z = 1) in; + +// This shader approximates downsampling: each output pixel samples the central +// 2x2 texels in the corresponding full-resolution scale x scale window. +// It is not an area/average-pool downsample over the whole window. +const int kScale = 2; +const int kOffset0 = 0; +const int kOffset1 = 1; + +vec2 readFlowXY(ivec2 p) { + uint xCoords[4] = uint[](0u, 0u, uint(p.y), uint(p.x)); + uint yCoords[4] = uint[](0u, 1u, uint(p.y), uint(p.x)); + float xVal[1]; + float yVal[1]; + tensorReadARM(flow, xCoords, xVal); + tensorReadARM(flow, yCoords, yVal); + return vec2(xVal[0], yVal[0]); +} + +vec2 alignCornersUv(vec2 gridXY) { + vec2 inputSize = vec2(textureSize(inputImage, 0)); + vec2 texel = (gridXY + vec2(1.0)) * vec2(0.5) * (inputSize - vec2(1.0)); + return (texel + vec2(0.5)) / inputSize; +} + +vec2 baseGridXY(ivec2 p, ivec2 fullSize) { + return vec2( + (2.0 * float(p.x)) / float(fullSize.x - 1) - 1.0, + (2.0 * float(p.y)) / float(fullSize.y - 1) - 1.0); +} + +vec4 sampleWarped(ivec2 p, ivec2 fullSize) { + vec2 flowXY = readFlowXY(p); + vec2 gridXY = baseGridXY(p, fullSize) + flowXY; + return texture(inputImage, alignCornersUv(gridXY)); +} + +void main() { + ivec2 outSize = imageSize(outImage); + ivec2 gid = ivec2(gl_GlobalInvocationID.xy); + if (gid.x >= outSize.x || gid.y >= outSize.y) { + return; + } + + ivec2 fullSize = textureSize(inputImage, 0); + ivec2 base = gid * kScale; + ivec2 p00 = base + ivec2(kOffset0, kOffset0); + ivec2 p01 = base + ivec2(kOffset1, kOffset0); + ivec2 p10 = base + ivec2(kOffset0, kOffset1); + ivec2 p11 = base + ivec2(kOffset1, kOffset1); + + vec4 value = 0.25 * ( + sampleWarped(p00, fullSize) + + sampleWarped(p01, fullSize) + + sampleWarped(p10, fullSize) + + sampleWarped(p11, fullSize)); + imageStore(outImage, gid, value); +} diff --git a/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.spirv.b64 b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.spirv.b64 new file mode 100644 index 00000000000..3fa3ebbba8f --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample2_sampler_int8_align_corners.spirv.b64 @@ -0,0 +1 @@ +AwIjBwAAAQALAA0A7AAAAAAAAAARAAIAAQAAABEAAgAyAAAAEQACAE4QAAAKAAUAU1BWX0FSTV90ZW5zb3JzAAsABgABAAAAR0xTTC5zdGQuNDUwAAAAAA4AAwAAAAAAAQAAAA8ABgAFAAAABAAAAG1haW4AAAAAnAAAABAABgAEAAAAEQAAAAgAAAAIAAAAAQAAAAMAAwACAAAAwgEAAAQABQBHTF9BUk1fdGVuc29ycwAABAAKAEdMX0dPT0dMRV9jcHBfc3R5bGVfbGluZV9kaXJlY3RpdmUAAAQACABHTF9HT09HTEVfaW5jbHVkZV9kaXJlY3RpdmUABQAEAAQAAABtYWluAAAAAAUABgANAAAAcmVhZEZsb3dYWSh2aTI7AAUAAwAMAAAAcAAAAAUABwASAAAAYWxpZ25Db3JuZXJzVXYodmYyOwAFAAQAEQAAAGdyaWRYWQAABQAHABcAAABiYXNlR3JpZFhZKHZpMjt2aTI7AAUAAwAVAAAAcAAAAAUABQAWAAAAZnVsbFNpemUAAAAABQAIAB0AAABzYW1wbGVXYXJwZWQodmkyO3ZpMjsAAAAFAAMAGwAAAHAAAAAFAAUAHAAAAGZ1bGxTaXplAAAAAAUABAAjAAAAeENvb3JkcwAFAAQALgAAAHlDb29yZHMABQAEADgAAABmbG93AAAAAAUABAA9AAAAeFZhbAAAAAAFAAQAQQAAAHlWYWwAAAAABQAFAEwAAABpbnB1dFNpemUAAAAFAAUAUAAAAGlucHV0SW1hZ2UAAAUABABVAAAAdGV4ZWwAAAAFAAQAfwAAAGZsb3dYWQAABQAEAIAAAABwYXJhbQAAAAUABACDAAAAZ3JpZFhZAAAFAAQAhAAAAHBhcmFtAAAABQAEAIYAAABwYXJhbQAAAAUABACMAAAAcGFyYW0AAAAFAAQAkwAAAG91dFNpemUABQAFAJYAAABvdXRJbWFnZQAAAAAFAAMAmQAAAGdpZAAFAAgAnAAAAGdsX0dsb2JhbEludm9jYXRpb25JRAAAAAUABQCzAAAAZnVsbFNpemUAAAAABQAEALcAAABiYXNlAAAAAAUAAwC8AAAAcDAwAAUAAwDAAAAAcDAxAAUAAwDEAAAAcDEwAAUAAwDIAAAAcDExAAUABADNAAAAdmFsdWUAAAAFAAQAzwAAAHBhcmFtAAAABQAEANEAAABwYXJhbQAAAAUABADUAAAAcGFyYW0AAAAFAAQA1gAAAHBhcmFtAAAABQAEANoAAABwYXJhbQAAAAUABADcAAAAcGFyYW0AAAAFAAQA4AAAAHBhcmFtAAAABQAEAOIAAABwYXJhbQAAAEcABAA4AAAAIQAAAAEAAABHAAQAOAAAACIAAAAAAAAARwAEAFAAAAAhAAAAAAAAAEcABABQAAAAIgAAAAAAAABHAAMAlgAAABkAAABHAAQAlgAAACEAAAACAAAARwAEAJYAAAAiAAAAAAAAAEcABACcAAAACwAAABwAAABHAAQA6wAAAAsAAAAZAAAAEwACAAIAAAAhAAMAAwAAAAIAAAAVAAQABgAAACAAAAABAAAAFwAEAAcAAAAGAAAAAgAAACAABAAIAAAABwAAAAcAAAAWAAMACQAAACAAAAAXAAQACgAAAAkAAAACAAAAIQAEAAsAAAAKAAAACAAAACAABAAPAAAABwAAAAoAAAAhAAQAEAAAAAoAAAAPAAAAIQAFABQAAAAKAAAACAAAAAgAAAAXAAQAGQAAAAkAAAAEAAAAIQAFABoAAAAZAAAACAAAAAgAAAAVAAQAHwAAACAAAAAAAAAAKwAEAB8AAAAgAAAABAAAABwABAAhAAAAHwAAACAAAAAgAAQAIgAAAAcAAAAhAAAAKwAEAB8AAAAkAAAAAAAAACsABAAfAAAAJQAAAAEAAAAgAAQAJgAAAAcAAAAGAAAAQxAEADYAAAAJAAAAIAAAACAABAA3AAAAAAAAADYAAAA7AAQANwAAADgAAAAAAAAAHAAEADsAAAAJAAAAJQAAACAABAA8AAAABwAAADsAAAArAAQABgAAAEMAAAAAAAAAIAAEAEQAAAAHAAAACQAAABkACQBNAAAACQAAAAEAAAAAAAAAAAAAAAAAAAABAAAAAAAAABsAAwBOAAAATQAAACAABABPAAAAAAAAAE4AAAA7AAQATwAAAFAAAAAAAAAAKwAEAAkAAABXAAAAAACAPywABQAKAAAAWAAAAFcAAABXAAAAKwAEAAkAAABaAAAAAAAAPywABQAKAAAAWwAAAFoAAABaAAAAKwAEAAkAAABmAAAAAAAAQCsABAAGAAAAbQAAAAEAAAArAAQACQAAAI8AAAAAAAAAGQAJAJQAAAAJAAAAAQAAAAAAAAAAAAAAAAAAAAIAAAAFAAAAIAAEAJUAAAAAAAAAlAAAADsABACVAAAAlgAAAAAAAAAXAAQAmgAAAB8AAAADAAAAIAAEAJsAAAABAAAAmgAAADsABACbAAAAnAAAAAEAAAAXAAQAnQAAAB8AAAACAAAAFAACAKEAAAArAAQABgAAALkAAAACAAAALAAFAAcAAAC+AAAAQwAAAEMAAAAsAAUABwAAAMIAAABtAAAAQwAAACwABQAHAAAAxgAAAEMAAABtAAAALAAFAAcAAADKAAAAbQAAAG0AAAAgAAQAzAAAAAcAAAAZAAAAKwAEAAkAAADOAAAAAACAPisABAAfAAAA6gAAAAgAAAAsAAYAmgAAAOsAAADqAAAA6gAAACUAAAA2AAUAAgAAAAQAAAAAAAAAAwAAAPgAAgAFAAAAOwAEAAgAAACTAAAABwAAADsABAAIAAAAmQAAAAcAAAA7AAQACAAAALMAAAAHAAAAOwAEAAgAAAC3AAAABwAAADsABAAIAAAAvAAAAAcAAAA7AAQACAAAAMAAAAAHAAAAOwAEAAgAAADEAAAABwAAADsABAAIAAAAyAAAAAcAAAA7AAQAzAAAAM0AAAAHAAAAOwAEAAgAAADPAAAABwAAADsABAAIAAAA0QAAAAcAAAA7AAQACAAAANQAAAAHAAAAOwAEAAgAAADWAAAABwAAADsABAAIAAAA2gAAAAcAAAA7AAQACAAAANwAAAAHAAAAOwAEAAgAAADgAAAABwAAADsABAAIAAAA4gAAAAcAAAA9AAQAlAAAAJcAAACWAAAAaAAEAAcAAACYAAAAlwAAAD4AAwCTAAAAmAAAAD0ABACaAAAAngAAAJwAAABPAAcAnQAAAJ8AAACeAAAAngAAAAAAAAABAAAAfAAEAAcAAACgAAAAnwAAAD4AAwCZAAAAoAAAAEEABQAmAAAAogAAAJkAAAAkAAAAPQAEAAYAAACjAAAAogAAAEEABQAmAAAApAAAAJMAAAAkAAAAPQAEAAYAAAClAAAApAAAAK8ABQChAAAApgAAAKMAAAClAAAAqAAEAKEAAACnAAAApgAAAPcAAwCpAAAAAAAAAPoABACnAAAAqAAAAKkAAAD4AAIAqAAAAEEABQAmAAAAqgAAAJkAAAAlAAAAPQAEAAYAAACrAAAAqgAAAEEABQAmAAAArAAAAJMAAAAlAAAAPQAEAAYAAACtAAAArAAAAK8ABQChAAAArgAAAKsAAACtAAAA+QACAKkAAAD4AAIAqQAAAPUABwChAAAArwAAAKYAAAAFAAAArgAAAKgAAAD3AAMAsQAAAAAAAAD6AAQArwAAALAAAACxAAAA+AACALAAAAD9AAEA+AACALEAAAA9AAQATgAAALQAAABQAAAAZAAEAE0AAAC1AAAAtAAAAGcABQAHAAAAtgAAALUAAABDAAAAPgADALMAAAC2AAAAPQAEAAcAAAC4AAAAmQAAAFAABQAHAAAAugAAALkAAAC5AAAAhAAFAAcAAAC7AAAAuAAAALoAAAA+AAMAtwAAALsAAAA9AAQABwAAAL0AAAC3AAAAgAAFAAcAAAC/AAAAvQAAAL4AAAA+AAMAvAAAAL8AAAA9AAQABwAAAMEAAAC3AAAAgAAFAAcAAADDAAAAwQAAAMIAAAA+AAMAwAAAAMMAAAA9AAQABwAAAMUAAAC3AAAAgAAFAAcAAADHAAAAxQAAAMYAAAA+AAMAxAAAAMcAAAA9AAQABwAAAMkAAAC3AAAAgAAFAAcAAADLAAAAyQAAAMoAAAA+AAMAyAAAAMsAAAA9AAQABwAAANAAAAC8AAAAPgADAM8AAADQAAAAPQAEAAcAAADSAAAAswAAAD4AAwDRAAAA0gAAADkABgAZAAAA0wAAAB0AAADPAAAA0QAAAD0ABAAHAAAA1QAAAMAAAAA+AAMA1AAAANUAAAA9AAQABwAAANcAAACzAAAAPgADANYAAADXAAAAOQAGABkAAADYAAAAHQAAANQAAADWAAAAgQAFABkAAADZAAAA0wAAANgAAAA9AAQABwAAANsAAADEAAAAPgADANoAAADbAAAAPQAEAAcAAADdAAAAswAAAD4AAwDcAAAA3QAAADkABgAZAAAA3gAAAB0AAADaAAAA3AAAAIEABQAZAAAA3wAAANkAAADeAAAAPQAEAAcAAADhAAAAyAAAAD4AAwDgAAAA4QAAAD0ABAAHAAAA4wAAALMAAAA+AAMA4gAAAOMAAAA5AAYAGQAAAOQAAAAdAAAA4AAAAOIAAACBAAUAGQAAAOUAAADfAAAA5AAAAI4ABQAZAAAA5gAAAOUAAADOAAAAPgADAM0AAADmAAAAPQAEAJQAAADnAAAAlgAAAD0ABAAHAAAA6AAAAJkAAAA9AAQAGQAAAOkAAADNAAAAYwAEAOcAAADoAAAA6QAAAP0AAQA4AAEANgAFAAoAAAANAAAAAAAAAAsAAAA3AAMACAAAAAwAAAD4AAIADgAAADsABAAiAAAAIwAAAAcAAAA7AAQAIgAAAC4AAAAHAAAAOwAEADwAAAA9AAAABwAAADsABAA8AAAAQQAAAAcAAABBAAUAJgAAACcAAAAMAAAAJQAAAD0ABAAGAAAAKAAAACcAAAB8AAQAHwAAACkAAAAoAAAAQQAFACYAAAAqAAAADAAAACQAAAA9AAQABgAAACsAAAAqAAAAfAAEAB8AAAAsAAAAKwAAAFAABwAhAAAALQAAACQAAAAkAAAAKQAAACwAAAA+AAMAIwAAAC0AAABBAAUAJgAAAC8AAAAMAAAAJQAAAD0ABAAGAAAAMAAAAC8AAAB8AAQAHwAAADEAAAAwAAAAQQAFACYAAAAyAAAADAAAACQAAAA9AAQABgAAADMAAAAyAAAAfAAEAB8AAAA0AAAAMwAAAFAABwAhAAAANQAAACQAAAAlAAAAMQAAADQAAAA+AAMALgAAADUAAAA9AAQANgAAADkAAAA4AAAAPQAEACEAAAA6AAAAIwAAAEQQBQA7AAAAPgAAADkAAAA6AAAAPgADAD0AAAA+AAAAPQAEADYAAAA/AAAAOAAAAD0ABAAhAAAAQAAAAC4AAABEEAUAOwAAAEIAAAA/AAAAQAAAAD4AAwBBAAAAQgAAAEEABQBEAAAARQAAAD0AAABDAAAAPQAEAAkAAABGAAAARQAAAEEABQBEAAAARwAAAEEAAABDAAAAPQAEAAkAAABIAAAARwAAAFAABQAKAAAASQAAAEYAAABIAAAA/gACAEkAAAA4AAEANgAFAAoAAAASAAAAAAAAABAAAAA3AAMADwAAABEAAAD4AAIAEwAAADsABAAPAAAATAAAAAcAAAA7AAQADwAAAFUAAAAHAAAAPQAEAE4AAABRAAAAUAAAAGQABABNAAAAUgAAAFEAAABnAAUABwAAAFMAAABSAAAAQwAAAG8ABAAKAAAAVAAAAFMAAAA+AAMATAAAAFQAAAA9AAQACgAAAFYAAAARAAAAgQAFAAoAAABZAAAAVgAAAFgAAACFAAUACgAAAFwAAABZAAAAWwAAAD0ABAAKAAAAXQAAAEwAAACDAAUACgAAAF4AAABdAAAAWAAAAIUABQAKAAAAXwAAAFwAAABeAAAAPgADAFUAAABfAAAAPQAEAAoAAABgAAAAVQAAAIEABQAKAAAAYQAAAGAAAABbAAAAPQAEAAoAAABiAAAATAAAAIgABQAKAAAAYwAAAGEAAABiAAAA/gACAGMAAAA4AAEANgAFAAoAAAAXAAAAAAAAABQAAAA3AAMACAAAABUAAAA3AAMACAAAABYAAAD4AAIAGAAAAEEABQAmAAAAZwAAABUAAAAkAAAAPQAEAAYAAABoAAAAZwAAAG8ABAAJAAAAaQAAAGgAAACFAAUACQAAAGoAAABmAAAAaQAAAEEABQAmAAAAawAAABYAAAAkAAAAPQAEAAYAAABsAAAAawAAAIIABQAGAAAAbgAAAGwAAABtAAAAbwAEAAkAAABvAAAAbgAAAIgABQAJAAAAcAAAAGoAAABvAAAAgwAFAAkAAABxAAAAcAAAAFcAAABBAAUAJgAAAHIAAAAVAAAAJQAAAD0ABAAGAAAAcwAAAHIAAABvAAQACQAAAHQAAABzAAAAhQAFAAkAAAB1AAAAZgAAAHQAAABBAAUAJgAAAHYAAAAWAAAAJQAAAD0ABAAGAAAAdwAAAHYAAACCAAUABgAAAHgAAAB3AAAAbQAAAG8ABAAJAAAAeQAAAHgAAACIAAUACQAAAHoAAAB1AAAAeQAAAIMABQAJAAAAewAAAHoAAABXAAAAUAAFAAoAAAB8AAAAcQAAAHsAAAD+AAIAfAAAADgAAQA2AAUAGQAAAB0AAAAAAAAAGgAAADcAAwAIAAAAGwAAADcAAwAIAAAAHAAAAPgAAgAeAAAAOwAEAA8AAAB/AAAABwAAADsABAAIAAAAgAAAAAcAAAA7AAQADwAAAIMAAAAHAAAAOwAEAAgAAACEAAAABwAAADsABAAIAAAAhgAAAAcAAAA7AAQADwAAAIwAAAAHAAAAPQAEAAcAAACBAAAAGwAAAD4AAwCAAAAAgQAAADkABQAKAAAAggAAAA0AAACAAAAAPgADAH8AAACCAAAAPQAEAAcAAACFAAAAGwAAAD4AAwCEAAAAhQAAAD0ABAAHAAAAhwAAABwAAAA+AAMAhgAAAIcAAAA5AAYACgAAAIgAAAAXAAAAhAAAAIYAAAA9AAQACgAAAIkAAAB/AAAAgQAFAAoAAACKAAAAiAAAAIkAAAA+AAMAgwAAAIoAAAA9AAQATgAAAIsAAABQAAAAPQAEAAoAAACNAAAAgwAAAD4AAwCMAAAAjQAAADkABQAKAAAAjgAAABIAAACMAAAAWAAHABkAAACQAAAAiwAAAI4AAAACAAAAjwAAAP4AAgCQAAAAOAABAA== \ No newline at end of file diff --git a/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.glsl b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.glsl new file mode 100644 index 00000000000..5f70a30f15b --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.glsl @@ -0,0 +1,70 @@ +// Copyright 2026 Arm Limited and/or its affiliates. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#version 450 +#extension GL_ARM_tensors : require + +layout(set = 0, binding = 0) uniform sampler2D inputImage; +layout(set = 0, binding = 1) uniform tensorARM flow; +layout(set = 0, binding = 2, rgba8_snorm) uniform writeonly image2D outImage; + +layout(local_size_x = 8, local_size_y = 8, local_size_z = 1) in; + +// This shader approximates downsampling: each output pixel samples the central +// 2x2 texels in the corresponding full-resolution scale x scale window. +// It is not an area/average-pool downsample over the whole window. +const int kScale = 4; +const int kOffset0 = 1; +const int kOffset1 = 2; + +vec2 readFlowXY(ivec2 p) { + uint xCoords[4] = uint[](0u, 0u, uint(p.y), uint(p.x)); + uint yCoords[4] = uint[](0u, 1u, uint(p.y), uint(p.x)); + float xVal[1]; + float yVal[1]; + tensorReadARM(flow, xCoords, xVal); + tensorReadARM(flow, yCoords, yVal); + return vec2(xVal[0], yVal[0]); +} + +vec2 alignCornersUv(vec2 gridXY) { + vec2 inputSize = vec2(textureSize(inputImage, 0)); + vec2 texel = (gridXY + vec2(1.0)) * vec2(0.5) * (inputSize - vec2(1.0)); + return (texel + vec2(0.5)) / inputSize; +} + +vec2 baseGridXY(ivec2 p, ivec2 fullSize) { + return vec2( + (2.0 * float(p.x)) / float(fullSize.x - 1) - 1.0, + (2.0 * float(p.y)) / float(fullSize.y - 1) - 1.0); +} + +vec4 sampleWarped(ivec2 p, ivec2 fullSize) { + vec2 flowXY = readFlowXY(p); + vec2 gridXY = baseGridXY(p, fullSize) + flowXY; + return texture(inputImage, alignCornersUv(gridXY)); +} + +void main() { + ivec2 outSize = imageSize(outImage); + ivec2 gid = ivec2(gl_GlobalInvocationID.xy); + if (gid.x >= outSize.x || gid.y >= outSize.y) { + return; + } + + ivec2 fullSize = textureSize(inputImage, 0); + ivec2 base = gid * kScale; + ivec2 p00 = base + ivec2(kOffset0, kOffset0); + ivec2 p01 = base + ivec2(kOffset1, kOffset0); + ivec2 p10 = base + ivec2(kOffset0, kOffset1); + ivec2 p11 = base + ivec2(kOffset1, kOffset1); + + vec4 value = 0.25 * ( + sampleWarped(p00, fullSize) + + sampleWarped(p01, fullSize) + + sampleWarped(p10, fullSize) + + sampleWarped(p11, fullSize)); + imageStore(outImage, gid, value); +} diff --git a/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.spirv.b64 b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.spirv.b64 new file mode 100644 index 00000000000..7db4d41f427 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample4_sampler_int8_align_corners.spirv.b64 @@ -0,0 +1 @@ +AwIjBwAAAQALAA0A7QAAAAAAAAARAAIAAQAAABEAAgAyAAAAEQACAE4QAAAKAAUAU1BWX0FSTV90ZW5zb3JzAAsABgABAAAAR0xTTC5zdGQuNDUwAAAAAA4AAwAAAAAAAQAAAA8ABgAFAAAABAAAAG1haW4AAAAAnAAAABAABgAEAAAAEQAAAAgAAAAIAAAAAQAAAAMAAwACAAAAwgEAAAQABQBHTF9BUk1fdGVuc29ycwAABAAKAEdMX0dPT0dMRV9jcHBfc3R5bGVfbGluZV9kaXJlY3RpdmUAAAQACABHTF9HT09HTEVfaW5jbHVkZV9kaXJlY3RpdmUABQAEAAQAAABtYWluAAAAAAUABgANAAAAcmVhZEZsb3dYWSh2aTI7AAUAAwAMAAAAcAAAAAUABwASAAAAYWxpZ25Db3JuZXJzVXYodmYyOwAFAAQAEQAAAGdyaWRYWQAABQAHABcAAABiYXNlR3JpZFhZKHZpMjt2aTI7AAUAAwAVAAAAcAAAAAUABQAWAAAAZnVsbFNpemUAAAAABQAIAB0AAABzYW1wbGVXYXJwZWQodmkyO3ZpMjsAAAAFAAMAGwAAAHAAAAAFAAUAHAAAAGZ1bGxTaXplAAAAAAUABAAjAAAAeENvb3JkcwAFAAQALgAAAHlDb29yZHMABQAEADgAAABmbG93AAAAAAUABAA9AAAAeFZhbAAAAAAFAAQAQQAAAHlWYWwAAAAABQAFAEwAAABpbnB1dFNpemUAAAAFAAUAUAAAAGlucHV0SW1hZ2UAAAUABABVAAAAdGV4ZWwAAAAFAAQAfwAAAGZsb3dYWQAABQAEAIAAAABwYXJhbQAAAAUABACDAAAAZ3JpZFhZAAAFAAQAhAAAAHBhcmFtAAAABQAEAIYAAABwYXJhbQAAAAUABACMAAAAcGFyYW0AAAAFAAQAkwAAAG91dFNpemUABQAFAJYAAABvdXRJbWFnZQAAAAAFAAMAmQAAAGdpZAAFAAgAnAAAAGdsX0dsb2JhbEludm9jYXRpb25JRAAAAAUABQCzAAAAZnVsbFNpemUAAAAABQAEALcAAABiYXNlAAAAAAUAAwC8AAAAcDAwAAUAAwDAAAAAcDAxAAUAAwDFAAAAcDEwAAUAAwDJAAAAcDExAAUABADOAAAAdmFsdWUAAAAFAAQA0AAAAHBhcmFtAAAABQAEANIAAABwYXJhbQAAAAUABADVAAAAcGFyYW0AAAAFAAQA1wAAAHBhcmFtAAAABQAEANsAAABwYXJhbQAAAAUABADdAAAAcGFyYW0AAAAFAAQA4QAAAHBhcmFtAAAABQAEAOMAAABwYXJhbQAAAEcABAA4AAAAIQAAAAEAAABHAAQAOAAAACIAAAAAAAAARwAEAFAAAAAhAAAAAAAAAEcABABQAAAAIgAAAAAAAABHAAMAlgAAABkAAABHAAQAlgAAACEAAAACAAAARwAEAJYAAAAiAAAAAAAAAEcABACcAAAACwAAABwAAABHAAQA7AAAAAsAAAAZAAAAEwACAAIAAAAhAAMAAwAAAAIAAAAVAAQABgAAACAAAAABAAAAFwAEAAcAAAAGAAAAAgAAACAABAAIAAAABwAAAAcAAAAWAAMACQAAACAAAAAXAAQACgAAAAkAAAACAAAAIQAEAAsAAAAKAAAACAAAACAABAAPAAAABwAAAAoAAAAhAAQAEAAAAAoAAAAPAAAAIQAFABQAAAAKAAAACAAAAAgAAAAXAAQAGQAAAAkAAAAEAAAAIQAFABoAAAAZAAAACAAAAAgAAAAVAAQAHwAAACAAAAAAAAAAKwAEAB8AAAAgAAAABAAAABwABAAhAAAAHwAAACAAAAAgAAQAIgAAAAcAAAAhAAAAKwAEAB8AAAAkAAAAAAAAACsABAAfAAAAJQAAAAEAAAAgAAQAJgAAAAcAAAAGAAAAQxAEADYAAAAJAAAAIAAAACAABAA3AAAAAAAAADYAAAA7AAQANwAAADgAAAAAAAAAHAAEADsAAAAJAAAAJQAAACAABAA8AAAABwAAADsAAAArAAQABgAAAEMAAAAAAAAAIAAEAEQAAAAHAAAACQAAABkACQBNAAAACQAAAAEAAAAAAAAAAAAAAAAAAAABAAAAAAAAABsAAwBOAAAATQAAACAABABPAAAAAAAAAE4AAAA7AAQATwAAAFAAAAAAAAAAKwAEAAkAAABXAAAAAACAPywABQAKAAAAWAAAAFcAAABXAAAAKwAEAAkAAABaAAAAAAAAPywABQAKAAAAWwAAAFoAAABaAAAAKwAEAAkAAABmAAAAAAAAQCsABAAGAAAAbQAAAAEAAAArAAQACQAAAI8AAAAAAAAAGQAJAJQAAAAJAAAAAQAAAAAAAAAAAAAAAAAAAAIAAAAFAAAAIAAEAJUAAAAAAAAAlAAAADsABACVAAAAlgAAAAAAAAAXAAQAmgAAAB8AAAADAAAAIAAEAJsAAAABAAAAmgAAADsABACbAAAAnAAAAAEAAAAXAAQAnQAAAB8AAAACAAAAFAACAKEAAAArAAQABgAAALkAAAAEAAAALAAFAAcAAAC+AAAAbQAAAG0AAAArAAQABgAAAMIAAAACAAAALAAFAAcAAADDAAAAwgAAAG0AAAAsAAUABwAAAMcAAABtAAAAwgAAACwABQAHAAAAywAAAMIAAADCAAAAIAAEAM0AAAAHAAAAGQAAACsABAAJAAAAzwAAAAAAgD4rAAQAHwAAAOsAAAAIAAAALAAGAJoAAADsAAAA6wAAAOsAAAAlAAAANgAFAAIAAAAEAAAAAAAAAAMAAAD4AAIABQAAADsABAAIAAAAkwAAAAcAAAA7AAQACAAAAJkAAAAHAAAAOwAEAAgAAACzAAAABwAAADsABAAIAAAAtwAAAAcAAAA7AAQACAAAALwAAAAHAAAAOwAEAAgAAADAAAAABwAAADsABAAIAAAAxQAAAAcAAAA7AAQACAAAAMkAAAAHAAAAOwAEAM0AAADOAAAABwAAADsABAAIAAAA0AAAAAcAAAA7AAQACAAAANIAAAAHAAAAOwAEAAgAAADVAAAABwAAADsABAAIAAAA1wAAAAcAAAA7AAQACAAAANsAAAAHAAAAOwAEAAgAAADdAAAABwAAADsABAAIAAAA4QAAAAcAAAA7AAQACAAAAOMAAAAHAAAAPQAEAJQAAACXAAAAlgAAAGgABAAHAAAAmAAAAJcAAAA+AAMAkwAAAJgAAAA9AAQAmgAAAJ4AAACcAAAATwAHAJ0AAACfAAAAngAAAJ4AAAAAAAAAAQAAAHwABAAHAAAAoAAAAJ8AAAA+AAMAmQAAAKAAAABBAAUAJgAAAKIAAACZAAAAJAAAAD0ABAAGAAAAowAAAKIAAABBAAUAJgAAAKQAAACTAAAAJAAAAD0ABAAGAAAApQAAAKQAAACvAAUAoQAAAKYAAACjAAAApQAAAKgABAChAAAApwAAAKYAAAD3AAMAqQAAAAAAAAD6AAQApwAAAKgAAACpAAAA+AACAKgAAABBAAUAJgAAAKoAAACZAAAAJQAAAD0ABAAGAAAAqwAAAKoAAABBAAUAJgAAAKwAAACTAAAAJQAAAD0ABAAGAAAArQAAAKwAAACvAAUAoQAAAK4AAACrAAAArQAAAPkAAgCpAAAA+AACAKkAAAD1AAcAoQAAAK8AAACmAAAABQAAAK4AAACoAAAA9wADALEAAAAAAAAA+gAEAK8AAACwAAAAsQAAAPgAAgCwAAAA/QABAPgAAgCxAAAAPQAEAE4AAAC0AAAAUAAAAGQABABNAAAAtQAAALQAAABnAAUABwAAALYAAAC1AAAAQwAAAD4AAwCzAAAAtgAAAD0ABAAHAAAAuAAAAJkAAABQAAUABwAAALoAAAC5AAAAuQAAAIQABQAHAAAAuwAAALgAAAC6AAAAPgADALcAAAC7AAAAPQAEAAcAAAC9AAAAtwAAAIAABQAHAAAAvwAAAL0AAAC+AAAAPgADALwAAAC/AAAAPQAEAAcAAADBAAAAtwAAAIAABQAHAAAAxAAAAMEAAADDAAAAPgADAMAAAADEAAAAPQAEAAcAAADGAAAAtwAAAIAABQAHAAAAyAAAAMYAAADHAAAAPgADAMUAAADIAAAAPQAEAAcAAADKAAAAtwAAAIAABQAHAAAAzAAAAMoAAADLAAAAPgADAMkAAADMAAAAPQAEAAcAAADRAAAAvAAAAD4AAwDQAAAA0QAAAD0ABAAHAAAA0wAAALMAAAA+AAMA0gAAANMAAAA5AAYAGQAAANQAAAAdAAAA0AAAANIAAAA9AAQABwAAANYAAADAAAAAPgADANUAAADWAAAAPQAEAAcAAADYAAAAswAAAD4AAwDXAAAA2AAAADkABgAZAAAA2QAAAB0AAADVAAAA1wAAAIEABQAZAAAA2gAAANQAAADZAAAAPQAEAAcAAADcAAAAxQAAAD4AAwDbAAAA3AAAAD0ABAAHAAAA3gAAALMAAAA+AAMA3QAAAN4AAAA5AAYAGQAAAN8AAAAdAAAA2wAAAN0AAACBAAUAGQAAAOAAAADaAAAA3wAAAD0ABAAHAAAA4gAAAMkAAAA+AAMA4QAAAOIAAAA9AAQABwAAAOQAAACzAAAAPgADAOMAAADkAAAAOQAGABkAAADlAAAAHQAAAOEAAADjAAAAgQAFABkAAADmAAAA4AAAAOUAAACOAAUAGQAAAOcAAADmAAAAzwAAAD4AAwDOAAAA5wAAAD0ABACUAAAA6AAAAJYAAAA9AAQABwAAAOkAAACZAAAAPQAEABkAAADqAAAAzgAAAGMABADoAAAA6QAAAOoAAAD9AAEAOAABADYABQAKAAAADQAAAAAAAAALAAAANwADAAgAAAAMAAAA+AACAA4AAAA7AAQAIgAAACMAAAAHAAAAOwAEACIAAAAuAAAABwAAADsABAA8AAAAPQAAAAcAAAA7AAQAPAAAAEEAAAAHAAAAQQAFACYAAAAnAAAADAAAACUAAAA9AAQABgAAACgAAAAnAAAAfAAEAB8AAAApAAAAKAAAAEEABQAmAAAAKgAAAAwAAAAkAAAAPQAEAAYAAAArAAAAKgAAAHwABAAfAAAALAAAACsAAABQAAcAIQAAAC0AAAAkAAAAJAAAACkAAAAsAAAAPgADACMAAAAtAAAAQQAFACYAAAAvAAAADAAAACUAAAA9AAQABgAAADAAAAAvAAAAfAAEAB8AAAAxAAAAMAAAAEEABQAmAAAAMgAAAAwAAAAkAAAAPQAEAAYAAAAzAAAAMgAAAHwABAAfAAAANAAAADMAAABQAAcAIQAAADUAAAAkAAAAJQAAADEAAAA0AAAAPgADAC4AAAA1AAAAPQAEADYAAAA5AAAAOAAAAD0ABAAhAAAAOgAAACMAAABEEAUAOwAAAD4AAAA5AAAAOgAAAD4AAwA9AAAAPgAAAD0ABAA2AAAAPwAAADgAAAA9AAQAIQAAAEAAAAAuAAAARBAFADsAAABCAAAAPwAAAEAAAAA+AAMAQQAAAEIAAABBAAUARAAAAEUAAAA9AAAAQwAAAD0ABAAJAAAARgAAAEUAAABBAAUARAAAAEcAAABBAAAAQwAAAD0ABAAJAAAASAAAAEcAAABQAAUACgAAAEkAAABGAAAASAAAAP4AAgBJAAAAOAABADYABQAKAAAAEgAAAAAAAAAQAAAANwADAA8AAAARAAAA+AACABMAAAA7AAQADwAAAEwAAAAHAAAAOwAEAA8AAABVAAAABwAAAD0ABABOAAAAUQAAAFAAAABkAAQATQAAAFIAAABRAAAAZwAFAAcAAABTAAAAUgAAAEMAAABvAAQACgAAAFQAAABTAAAAPgADAEwAAABUAAAAPQAEAAoAAABWAAAAEQAAAIEABQAKAAAAWQAAAFYAAABYAAAAhQAFAAoAAABcAAAAWQAAAFsAAAA9AAQACgAAAF0AAABMAAAAgwAFAAoAAABeAAAAXQAAAFgAAACFAAUACgAAAF8AAABcAAAAXgAAAD4AAwBVAAAAXwAAAD0ABAAKAAAAYAAAAFUAAACBAAUACgAAAGEAAABgAAAAWwAAAD0ABAAKAAAAYgAAAEwAAACIAAUACgAAAGMAAABhAAAAYgAAAP4AAgBjAAAAOAABADYABQAKAAAAFwAAAAAAAAAUAAAANwADAAgAAAAVAAAANwADAAgAAAAWAAAA+AACABgAAABBAAUAJgAAAGcAAAAVAAAAJAAAAD0ABAAGAAAAaAAAAGcAAABvAAQACQAAAGkAAABoAAAAhQAFAAkAAABqAAAAZgAAAGkAAABBAAUAJgAAAGsAAAAWAAAAJAAAAD0ABAAGAAAAbAAAAGsAAACCAAUABgAAAG4AAABsAAAAbQAAAG8ABAAJAAAAbwAAAG4AAACIAAUACQAAAHAAAABqAAAAbwAAAIMABQAJAAAAcQAAAHAAAABXAAAAQQAFACYAAAByAAAAFQAAACUAAAA9AAQABgAAAHMAAAByAAAAbwAEAAkAAAB0AAAAcwAAAIUABQAJAAAAdQAAAGYAAAB0AAAAQQAFACYAAAB2AAAAFgAAACUAAAA9AAQABgAAAHcAAAB2AAAAggAFAAYAAAB4AAAAdwAAAG0AAABvAAQACQAAAHkAAAB4AAAAiAAFAAkAAAB6AAAAdQAAAHkAAACDAAUACQAAAHsAAAB6AAAAVwAAAFAABQAKAAAAfAAAAHEAAAB7AAAA/gACAHwAAAA4AAEANgAFABkAAAAdAAAAAAAAABoAAAA3AAMACAAAABsAAAA3AAMACAAAABwAAAD4AAIAHgAAADsABAAPAAAAfwAAAAcAAAA7AAQACAAAAIAAAAAHAAAAOwAEAA8AAACDAAAABwAAADsABAAIAAAAhAAAAAcAAAA7AAQACAAAAIYAAAAHAAAAOwAEAA8AAACMAAAABwAAAD0ABAAHAAAAgQAAABsAAAA+AAMAgAAAAIEAAAA5AAUACgAAAIIAAAANAAAAgAAAAD4AAwB/AAAAggAAAD0ABAAHAAAAhQAAABsAAAA+AAMAhAAAAIUAAAA9AAQABwAAAIcAAAAcAAAAPgADAIYAAACHAAAAOQAGAAoAAACIAAAAFwAAAIQAAACGAAAAPQAEAAoAAACJAAAAfwAAAIEABQAKAAAAigAAAIgAAACJAAAAPgADAIMAAACKAAAAPQAEAE4AAACLAAAAUAAAAD0ABAAKAAAAjQAAAIMAAAA+AAMAjAAAAI0AAAA5AAUACgAAAI4AAAASAAAAjAAAAFgABwAZAAAAkAAAAIsAAACOAAAAAgAAAI8AAAD+AAIAkAAAADgAAQA= \ No newline at end of file diff --git a/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.glsl b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.glsl new file mode 100644 index 00000000000..3830acb6750 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.glsl @@ -0,0 +1,70 @@ +// Copyright 2026 Arm Limited and/or its affiliates. +// +// This source code is licensed under the BSD-style license found in the +// LICENSE file in the root directory of this source tree. + +#version 450 +#extension GL_ARM_tensors : require + +layout(set = 0, binding = 0) uniform sampler2D inputImage; +layout(set = 0, binding = 1) uniform tensorARM flow; +layout(set = 0, binding = 2, rgba8_snorm) uniform writeonly image2D outImage; + +layout(local_size_x = 8, local_size_y = 8, local_size_z = 1) in; + +// This shader approximates downsampling: each output pixel samples the central +// 2x2 texels in the corresponding full-resolution scale x scale window. +// It is not an area/average-pool downsample over the whole window. +const int kScale = 8; +const int kOffset0 = 3; +const int kOffset1 = 4; + +vec2 readFlowXY(ivec2 p) { + uint xCoords[4] = uint[](0u, 0u, uint(p.y), uint(p.x)); + uint yCoords[4] = uint[](0u, 1u, uint(p.y), uint(p.x)); + float xVal[1]; + float yVal[1]; + tensorReadARM(flow, xCoords, xVal); + tensorReadARM(flow, yCoords, yVal); + return vec2(xVal[0], yVal[0]); +} + +vec2 alignCornersUv(vec2 gridXY) { + vec2 inputSize = vec2(textureSize(inputImage, 0)); + vec2 texel = (gridXY + vec2(1.0)) * vec2(0.5) * (inputSize - vec2(1.0)); + return (texel + vec2(0.5)) / inputSize; +} + +vec2 baseGridXY(ivec2 p, ivec2 fullSize) { + return vec2( + (2.0 * float(p.x)) / float(fullSize.x - 1) - 1.0, + (2.0 * float(p.y)) / float(fullSize.y - 1) - 1.0); +} + +vec4 sampleWarped(ivec2 p, ivec2 fullSize) { + vec2 flowXY = readFlowXY(p); + vec2 gridXY = baseGridXY(p, fullSize) + flowXY; + return texture(inputImage, alignCornersUv(gridXY)); +} + +void main() { + ivec2 outSize = imageSize(outImage); + ivec2 gid = ivec2(gl_GlobalInvocationID.xy); + if (gid.x >= outSize.x || gid.y >= outSize.y) { + return; + } + + ivec2 fullSize = textureSize(inputImage, 0); + ivec2 base = gid * kScale; + ivec2 p00 = base + ivec2(kOffset0, kOffset0); + ivec2 p01 = base + ivec2(kOffset1, kOffset0); + ivec2 p10 = base + ivec2(kOffset0, kOffset1); + ivec2 p11 = base + ivec2(kOffset1, kOffset1); + + vec4 value = 0.25 * ( + sampleWarped(p00, fullSize) + + sampleWarped(p01, fullSize) + + sampleWarped(p10, fullSize) + + sampleWarped(p11, fullSize)); + imageStore(outImage, gid, value); +} diff --git a/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.spirv.b64 b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.spirv.b64 new file mode 100644 index 00000000000..845c4e1ea84 --- /dev/null +++ b/examples/arm/QAT_example/rife_vgf/shaders/warp_downsample8_sampler_int8_align_corners.spirv.b64 @@ -0,0 +1 @@ +AwIjBwAAAQALAA0A7gAAAAAAAAARAAIAAQAAABEAAgAyAAAAEQACAE4QAAAKAAUAU1BWX0FSTV90ZW5zb3JzAAsABgABAAAAR0xTTC5zdGQuNDUwAAAAAA4AAwAAAAAAAQAAAA8ABgAFAAAABAAAAG1haW4AAAAAnAAAABAABgAEAAAAEQAAAAgAAAAIAAAAAQAAAAMAAwACAAAAwgEAAAQABQBHTF9BUk1fdGVuc29ycwAABAAKAEdMX0dPT0dMRV9jcHBfc3R5bGVfbGluZV9kaXJlY3RpdmUAAAQACABHTF9HT09HTEVfaW5jbHVkZV9kaXJlY3RpdmUABQAEAAQAAABtYWluAAAAAAUABgANAAAAcmVhZEZsb3dYWSh2aTI7AAUAAwAMAAAAcAAAAAUABwASAAAAYWxpZ25Db3JuZXJzVXYodmYyOwAFAAQAEQAAAGdyaWRYWQAABQAHABcAAABiYXNlR3JpZFhZKHZpMjt2aTI7AAUAAwAVAAAAcAAAAAUABQAWAAAAZnVsbFNpemUAAAAABQAIAB0AAABzYW1wbGVXYXJwZWQodmkyO3ZpMjsAAAAFAAMAGwAAAHAAAAAFAAUAHAAAAGZ1bGxTaXplAAAAAAUABAAjAAAAeENvb3JkcwAFAAQALgAAAHlDb29yZHMABQAEADgAAABmbG93AAAAAAUABAA9AAAAeFZhbAAAAAAFAAQAQQAAAHlWYWwAAAAABQAFAEwAAABpbnB1dFNpemUAAAAFAAUAUAAAAGlucHV0SW1hZ2UAAAUABABVAAAAdGV4ZWwAAAAFAAQAfwAAAGZsb3dYWQAABQAEAIAAAABwYXJhbQAAAAUABACDAAAAZ3JpZFhZAAAFAAQAhAAAAHBhcmFtAAAABQAEAIYAAABwYXJhbQAAAAUABACMAAAAcGFyYW0AAAAFAAQAkwAAAG91dFNpemUABQAFAJYAAABvdXRJbWFnZQAAAAAFAAMAmQAAAGdpZAAFAAgAnAAAAGdsX0dsb2JhbEludm9jYXRpb25JRAAAAAUABQCzAAAAZnVsbFNpemUAAAAABQAEALcAAABiYXNlAAAAAAUAAwC8AAAAcDAwAAUAAwDBAAAAcDAxAAUAAwDGAAAAcDEwAAUAAwDKAAAAcDExAAUABADPAAAAdmFsdWUAAAAFAAQA0QAAAHBhcmFtAAAABQAEANMAAABwYXJhbQAAAAUABADWAAAAcGFyYW0AAAAFAAQA2AAAAHBhcmFtAAAABQAEANwAAABwYXJhbQAAAAUABADeAAAAcGFyYW0AAAAFAAQA4gAAAHBhcmFtAAAABQAEAOQAAABwYXJhbQAAAEcABAA4AAAAIQAAAAEAAABHAAQAOAAAACIAAAAAAAAARwAEAFAAAAAhAAAAAAAAAEcABABQAAAAIgAAAAAAAABHAAMAlgAAABkAAABHAAQAlgAAACEAAAACAAAARwAEAJYAAAAiAAAAAAAAAEcABACcAAAACwAAABwAAABHAAQA7QAAAAsAAAAZAAAAEwACAAIAAAAhAAMAAwAAAAIAAAAVAAQABgAAACAAAAABAAAAFwAEAAcAAAAGAAAAAgAAACAABAAIAAAABwAAAAcAAAAWAAMACQAAACAAAAAXAAQACgAAAAkAAAACAAAAIQAEAAsAAAAKAAAACAAAACAABAAPAAAABwAAAAoAAAAhAAQAEAAAAAoAAAAPAAAAIQAFABQAAAAKAAAACAAAAAgAAAAXAAQAGQAAAAkAAAAEAAAAIQAFABoAAAAZAAAACAAAAAgAAAAVAAQAHwAAACAAAAAAAAAAKwAEAB8AAAAgAAAABAAAABwABAAhAAAAHwAAACAAAAAgAAQAIgAAAAcAAAAhAAAAKwAEAB8AAAAkAAAAAAAAACsABAAfAAAAJQAAAAEAAAAgAAQAJgAAAAcAAAAGAAAAQxAEADYAAAAJAAAAIAAAACAABAA3AAAAAAAAADYAAAA7AAQANwAAADgAAAAAAAAAHAAEADsAAAAJAAAAJQAAACAABAA8AAAABwAAADsAAAArAAQABgAAAEMAAAAAAAAAIAAEAEQAAAAHAAAACQAAABkACQBNAAAACQAAAAEAAAAAAAAAAAAAAAAAAAABAAAAAAAAABsAAwBOAAAATQAAACAABABPAAAAAAAAAE4AAAA7AAQATwAAAFAAAAAAAAAAKwAEAAkAAABXAAAAAACAPywABQAKAAAAWAAAAFcAAABXAAAAKwAEAAkAAABaAAAAAAAAPywABQAKAAAAWwAAAFoAAABaAAAAKwAEAAkAAABmAAAAAAAAQCsABAAGAAAAbQAAAAEAAAArAAQACQAAAI8AAAAAAAAAGQAJAJQAAAAJAAAAAQAAAAAAAAAAAAAAAAAAAAIAAAAFAAAAIAAEAJUAAAAAAAAAlAAAADsABACVAAAAlgAAAAAAAAAXAAQAmgAAAB8AAAADAAAAIAAEAJsAAAABAAAAmgAAADsABACbAAAAnAAAAAEAAAAXAAQAnQAAAB8AAAACAAAAFAACAKEAAAArAAQABgAAALkAAAAIAAAAKwAEAAYAAAC+AAAAAwAAACwABQAHAAAAvwAAAL4AAAC+AAAAKwAEAAYAAADDAAAABAAAACwABQAHAAAAxAAAAMMAAAC+AAAALAAFAAcAAADIAAAAvgAAAMMAAAAsAAUABwAAAMwAAADDAAAAwwAAACAABADOAAAABwAAABkAAAArAAQACQAAANAAAAAAAIA+KwAEAB8AAADsAAAACAAAACwABgCaAAAA7QAAAOwAAADsAAAAJQAAADYABQACAAAABAAAAAAAAAADAAAA+AACAAUAAAA7AAQACAAAAJMAAAAHAAAAOwAEAAgAAACZAAAABwAAADsABAAIAAAAswAAAAcAAAA7AAQACAAAALcAAAAHAAAAOwAEAAgAAAC8AAAABwAAADsABAAIAAAAwQAAAAcAAAA7AAQACAAAAMYAAAAHAAAAOwAEAAgAAADKAAAABwAAADsABADOAAAAzwAAAAcAAAA7AAQACAAAANEAAAAHAAAAOwAEAAgAAADTAAAABwAAADsABAAIAAAA1gAAAAcAAAA7AAQACAAAANgAAAAHAAAAOwAEAAgAAADcAAAABwAAADsABAAIAAAA3gAAAAcAAAA7AAQACAAAAOIAAAAHAAAAOwAEAAgAAADkAAAABwAAAD0ABACUAAAAlwAAAJYAAABoAAQABwAAAJgAAACXAAAAPgADAJMAAACYAAAAPQAEAJoAAACeAAAAnAAAAE8ABwCdAAAAnwAAAJ4AAACeAAAAAAAAAAEAAAB8AAQABwAAAKAAAACfAAAAPgADAJkAAACgAAAAQQAFACYAAACiAAAAmQAAACQAAAA9AAQABgAAAKMAAACiAAAAQQAFACYAAACkAAAAkwAAACQAAAA9AAQABgAAAKUAAACkAAAArwAFAKEAAACmAAAAowAAAKUAAACoAAQAoQAAAKcAAACmAAAA9wADAKkAAAAAAAAA+gAEAKcAAACoAAAAqQAAAPgAAgCoAAAAQQAFACYAAACqAAAAmQAAACUAAAA9AAQABgAAAKsAAACqAAAAQQAFACYAAACsAAAAkwAAACUAAAA9AAQABgAAAK0AAACsAAAArwAFAKEAAACuAAAAqwAAAK0AAAD5AAIAqQAAAPgAAgCpAAAA9QAHAKEAAACvAAAApgAAAAUAAACuAAAAqAAAAPcAAwCxAAAAAAAAAPoABACvAAAAsAAAALEAAAD4AAIAsAAAAP0AAQD4AAIAsQAAAD0ABABOAAAAtAAAAFAAAABkAAQATQAAALUAAAC0AAAAZwAFAAcAAAC2AAAAtQAAAEMAAAA+AAMAswAAALYAAAA9AAQABwAAALgAAACZAAAAUAAFAAcAAAC6AAAAuQAAALkAAACEAAUABwAAALsAAAC4AAAAugAAAD4AAwC3AAAAuwAAAD0ABAAHAAAAvQAAALcAAACAAAUABwAAAMAAAAC9AAAAvwAAAD4AAwC8AAAAwAAAAD0ABAAHAAAAwgAAALcAAACAAAUABwAAAMUAAADCAAAAxAAAAD4AAwDBAAAAxQAAAD0ABAAHAAAAxwAAALcAAACAAAUABwAAAMkAAADHAAAAyAAAAD4AAwDGAAAAyQAAAD0ABAAHAAAAywAAALcAAACAAAUABwAAAM0AAADLAAAAzAAAAD4AAwDKAAAAzQAAAD0ABAAHAAAA0gAAALwAAAA+AAMA0QAAANIAAAA9AAQABwAAANQAAACzAAAAPgADANMAAADUAAAAOQAGABkAAADVAAAAHQAAANEAAADTAAAAPQAEAAcAAADXAAAAwQAAAD4AAwDWAAAA1wAAAD0ABAAHAAAA2QAAALMAAAA+AAMA2AAAANkAAAA5AAYAGQAAANoAAAAdAAAA1gAAANgAAACBAAUAGQAAANsAAADVAAAA2gAAAD0ABAAHAAAA3QAAAMYAAAA+AAMA3AAAAN0AAAA9AAQABwAAAN8AAACzAAAAPgADAN4AAADfAAAAOQAGABkAAADgAAAAHQAAANwAAADeAAAAgQAFABkAAADhAAAA2wAAAOAAAAA9AAQABwAAAOMAAADKAAAAPgADAOIAAADjAAAAPQAEAAcAAADlAAAAswAAAD4AAwDkAAAA5QAAADkABgAZAAAA5gAAAB0AAADiAAAA5AAAAIEABQAZAAAA5wAAAOEAAADmAAAAjgAFABkAAADoAAAA5wAAANAAAAA+AAMAzwAAAOgAAAA9AAQAlAAAAOkAAACWAAAAPQAEAAcAAADqAAAAmQAAAD0ABAAZAAAA6wAAAM8AAABjAAQA6QAAAOoAAADrAAAA/QABADgAAQA2AAUACgAAAA0AAAAAAAAACwAAADcAAwAIAAAADAAAAPgAAgAOAAAAOwAEACIAAAAjAAAABwAAADsABAAiAAAALgAAAAcAAAA7AAQAPAAAAD0AAAAHAAAAOwAEADwAAABBAAAABwAAAEEABQAmAAAAJwAAAAwAAAAlAAAAPQAEAAYAAAAoAAAAJwAAAHwABAAfAAAAKQAAACgAAABBAAUAJgAAACoAAAAMAAAAJAAAAD0ABAAGAAAAKwAAACoAAAB8AAQAHwAAACwAAAArAAAAUAAHACEAAAAtAAAAJAAAACQAAAApAAAALAAAAD4AAwAjAAAALQAAAEEABQAmAAAALwAAAAwAAAAlAAAAPQAEAAYAAAAwAAAALwAAAHwABAAfAAAAMQAAADAAAABBAAUAJgAAADIAAAAMAAAAJAAAAD0ABAAGAAAAMwAAADIAAAB8AAQAHwAAADQAAAAzAAAAUAAHACEAAAA1AAAAJAAAACUAAAAxAAAANAAAAD4AAwAuAAAANQAAAD0ABAA2AAAAOQAAADgAAAA9AAQAIQAAADoAAAAjAAAARBAFADsAAAA+AAAAOQAAADoAAAA+AAMAPQAAAD4AAAA9AAQANgAAAD8AAAA4AAAAPQAEACEAAABAAAAALgAAAEQQBQA7AAAAQgAAAD8AAABAAAAAPgADAEEAAABCAAAAQQAFAEQAAABFAAAAPQAAAEMAAAA9AAQACQAAAEYAAABFAAAAQQAFAEQAAABHAAAAQQAAAEMAAAA9AAQACQAAAEgAAABHAAAAUAAFAAoAAABJAAAARgAAAEgAAAD+AAIASQAAADgAAQA2AAUACgAAABIAAAAAAAAAEAAAADcAAwAPAAAAEQAAAPgAAgATAAAAOwAEAA8AAABMAAAABwAAADsABAAPAAAAVQAAAAcAAAA9AAQATgAAAFEAAABQAAAAZAAEAE0AAABSAAAAUQAAAGcABQAHAAAAUwAAAFIAAABDAAAAbwAEAAoAAABUAAAAUwAAAD4AAwBMAAAAVAAAAD0ABAAKAAAAVgAAABEAAACBAAUACgAAAFkAAABWAAAAWAAAAIUABQAKAAAAXAAAAFkAAABbAAAAPQAEAAoAAABdAAAATAAAAIMABQAKAAAAXgAAAF0AAABYAAAAhQAFAAoAAABfAAAAXAAAAF4AAAA+AAMAVQAAAF8AAAA9AAQACgAAAGAAAABVAAAAgQAFAAoAAABhAAAAYAAAAFsAAAA9AAQACgAAAGIAAABMAAAAiAAFAAoAAABjAAAAYQAAAGIAAAD+AAIAYwAAADgAAQA2AAUACgAAABcAAAAAAAAAFAAAADcAAwAIAAAAFQAAADcAAwAIAAAAFgAAAPgAAgAYAAAAQQAFACYAAABnAAAAFQAAACQAAAA9AAQABgAAAGgAAABnAAAAbwAEAAkAAABpAAAAaAAAAIUABQAJAAAAagAAAGYAAABpAAAAQQAFACYAAABrAAAAFgAAACQAAAA9AAQABgAAAGwAAABrAAAAggAFAAYAAABuAAAAbAAAAG0AAABvAAQACQAAAG8AAABuAAAAiAAFAAkAAABwAAAAagAAAG8AAACDAAUACQAAAHEAAABwAAAAVwAAAEEABQAmAAAAcgAAABUAAAAlAAAAPQAEAAYAAABzAAAAcgAAAG8ABAAJAAAAdAAAAHMAAACFAAUACQAAAHUAAABmAAAAdAAAAEEABQAmAAAAdgAAABYAAAAlAAAAPQAEAAYAAAB3AAAAdgAAAIIABQAGAAAAeAAAAHcAAABtAAAAbwAEAAkAAAB5AAAAeAAAAIgABQAJAAAAegAAAHUAAAB5AAAAgwAFAAkAAAB7AAAAegAAAFcAAABQAAUACgAAAHwAAABxAAAAewAAAP4AAgB8AAAAOAABADYABQAZAAAAHQAAAAAAAAAaAAAANwADAAgAAAAbAAAANwADAAgAAAAcAAAA+AACAB4AAAA7AAQADwAAAH8AAAAHAAAAOwAEAAgAAACAAAAABwAAADsABAAPAAAAgwAAAAcAAAA7AAQACAAAAIQAAAAHAAAAOwAEAAgAAACGAAAABwAAADsABAAPAAAAjAAAAAcAAAA9AAQABwAAAIEAAAAbAAAAPgADAIAAAACBAAAAOQAFAAoAAACCAAAADQAAAIAAAAA+AAMAfwAAAIIAAAA9AAQABwAAAIUAAAAbAAAAPgADAIQAAACFAAAAPQAEAAcAAACHAAAAHAAAAD4AAwCGAAAAhwAAADkABgAKAAAAiAAAABcAAACEAAAAhgAAAD0ABAAKAAAAiQAAAH8AAACBAAUACgAAAIoAAACIAAAAiQAAAD4AAwCDAAAAigAAAD0ABABOAAAAiwAAAFAAAAA9AAQACgAAAI0AAACDAAAAPgADAIwAAACNAAAAOQAFAAoAAACOAAAAEgAAAIwAAABYAAcAGQAAAJAAAACLAAAAjgAAAAIAAACPAAAA/gACAJAAAAA4AAEA \ No newline at end of file