Skip to content
Merged
Show file tree
Hide file tree
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
9 changes: 8 additions & 1 deletion examples/arm/QAT_example/qat_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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,
Expand Down
6 changes: 6 additions & 0 deletions examples/arm/QAT_example/rife_vgf/__init__.py
Original file line number Diff line number Diff line change
@@ -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
189 changes: 189 additions & 0 deletions examples/arm/QAT_example/rife_vgf/extension.py
Original file line number Diff line number Diff line change
@@ -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())
8 changes: 8 additions & 0 deletions examples/arm/QAT_example/rife_vgf/passes/__init__.py
Original file line number Diff line number Diff line change
@@ -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,
)
Loading
Loading