Files
Model-Optimizer/examples/diffusers/quantization/calibration.py
T
jingyu-ml 26ae8da517 [2/3] Implicit Gemm NVFP4 (#1227)
### What does this PR do?

Type of change: new feature <!-- Use one of the following: Bug fix, new
feature, new example, new tests, documentation. -->

- Add Conv3D implicit GEMM kernel with BF16 WMMA tensor cores and fused
NVFP4 activation quantization for video diffusion VAE layers
- Integrate into _QuantConv3d via QuantModuleRegistry — automatically
dispatched when NVFP4 quantization is applied to nn.Conv3d
- Move kernel from `experimental/conv/ to modelopt/torch/kernels/conv/`;
move tests to `tests/gpu/torch/quantization/kernels/`

### Testing
<!-- Mention how have you tested your change if applicable. -->

- Added test cases to measure the difference between cuDNN and our CUDA
implicit GEMM kernel
- Added an NVFP4 fake quantization test using CUDA code

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ✅ <!--- If ❌, explain why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ <!---
Mandatory -->
- Did you write any new necessary tests?: ✅ <!--- Mandatory for new
features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ <!--- Only for new features, API changes, critical bug fixes or
backward incompatible changes. -->

### Additional Information
<!-- E.g. related issue. -->


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Per-backbone quantization/export in a single run with per-backbone
checkpoints and backbone-aware quant filters
* Configurable NVFP4 block-size via CLI/config; improved NVFP4 Conv3D
inference path and Wan 2.2 quantization support
* **Bug Fixes**
* Video-model calibration now respects extra params and forces video
decoding during calibration
* **Documentation**
* Added comprehensive Conv3D implicit‑GEMM kernel documentation; removed
experimental Conv3D prototype docs/benchmark
* **Tests**
* New Wan 2.2 quantization/export tests and expanded Conv3D/FP4 kernel
test coverage
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Jingyu Xin <jingyux@nvidia.com>
2026-04-19 12:20:14 +05:30

215 lines
9.1 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import logging
import warnings
from pathlib import Path
from typing import Any
from models_utils import MODEL_DEFAULTS, ModelType
from pipeline_manager import PipelineManager
from quantize_config import CalibrationConfig
from tqdm import tqdm
from utils import load_calib_prompts
class Calibrator:
"""Handles model calibration for quantization."""
def __init__(
self,
pipeline_manager: PipelineManager,
config: CalibrationConfig,
model_type: ModelType,
logger: logging.Logger,
):
"""
Initialize calibrator.
Args:
pipeline_manager: Pipeline manager with main and upsampler pipelines
config: Calibration configuration
model_type: Type of model being calibrated
logger: Logger instance
"""
self.pipeline_manager = pipeline_manager
self.pipe = pipeline_manager.pipe
self.pipe_upsample = pipeline_manager.pipe_upsample
self.config = config
self.model_type = model_type
self.logger = logger
def load_and_batch_prompts(self) -> list[list[str]]:
"""
Load calibration prompts from file.
Returns:
List of batched calibration prompts
"""
self.logger.info(f"Loading calibration prompts from {self.config.prompts_dataset}")
if isinstance(self.config.prompts_dataset, Path):
return load_calib_prompts(
self.config.batch_size,
self.config.prompts_dataset,
)
return load_calib_prompts(
self.config.batch_size,
self.config.prompts_dataset["name"],
self.config.prompts_dataset["split"],
self.config.prompts_dataset["column"],
)
def run_calibration(self, batched_prompts: list[list[str]]) -> None:
"""
Run calibration steps on the pipeline.
Args:
batched_prompts: List of batched calibration prompts
"""
self.logger.info(f"Starting calibration with {self.config.num_batches} batches")
extra_args = MODEL_DEFAULTS.get(self.model_type, {}).get("inference_extra_args", {})
with tqdm(total=self.config.num_batches, desc="Calibration", unit="batch") as pbar:
for i, prompt_batch in enumerate(batched_prompts):
if i >= self.config.num_batches:
break
if self.model_type == ModelType.LTX2:
self._run_ltx2_calibration(prompt_batch, extra_args)
elif self.model_type == ModelType.LTX_VIDEO_DEV:
# Special handling for LTX-Video
self._run_ltx_video_calibration(prompt_batch, extra_args)
elif self.model_type in [ModelType.WAN22_T2V_14b, ModelType.WAN22_T2V_5b]:
# Special handling for WAN video models
self._run_wan_video_calibration(prompt_batch, extra_args)
else:
common_args = {
"prompt": prompt_batch,
"num_inference_steps": self.config.n_steps,
}
self.pipe(**common_args, **extra_args).images
pbar.update(1)
self.logger.debug(f"Completed calibration batch {i + 1}/{self.config.num_batches}")
self.logger.info("Calibration completed successfully")
def _run_wan_video_calibration(
self, prompt_batch: list[str], extra_args: dict[str, Any]
) -> None:
extra_params = self.pipeline_manager.config.extra_params
kwargs = {}
kwargs["negative_prompt"] = extra_args["negative_prompt"]
kwargs["height"] = extra_params.get("height", extra_args["height"])
kwargs["width"] = extra_params.get("width", extra_args["width"])
kwargs["num_frames"] = extra_params.get("num_frames", extra_args["num_frames"])
kwargs["guidance_scale"] = extra_args["guidance_scale"]
if "guidance_scale_2" in extra_args:
kwargs["guidance_scale_2"] = extra_args["guidance_scale_2"]
kwargs["num_inference_steps"] = self.config.n_steps
self.pipe(prompt=prompt_batch, **kwargs).frames
def _run_ltx2_calibration(self, prompt_batch: list[str], extra_args: dict[str, Any]) -> None:
warnings.warn(
"LTX-2 packages (ltx-core, ltx-pipelines, ltx-trainer) are provided by Lightricks and are NOT "
"covered by the Apache 2.0 license governing NVIDIA Model Optimizer. You MUST comply "
"with the LTX Community License Agreement when installing and using LTX-2 with NVIDIA "
"Model Optimizer. Any derivative models or fine-tuned weights from LTX-2 remain "
"subject to the LTX Community License Agreement, not Apache 2.0. "
"See: https://github.com/Lightricks/LTX-2/blob/main/LICENSE",
UserWarning,
stacklevel=2,
)
from ltx_core.model.video_vae import TilingConfig
from ltx_pipelines.utils.constants import (
DEFAULT_AUDIO_GUIDER_PARAMS,
DEFAULT_VIDEO_GUIDER_PARAMS,
)
prompt = prompt_batch[0]
extra_params = self.pipeline_manager.config.extra_params
kwargs = {
"negative_prompt": extra_args.get(
"negative_prompt", "worst quality, inconsistent motion, blurry, jittery, distorted"
),
"seed": extra_params.get("seed", 0),
"height": extra_params.get("height", extra_args.get("height", 1024)),
"width": extra_params.get("width", extra_args.get("width", 1536)),
"num_frames": extra_params.get("num_frames", extra_args.get("num_frames", 121)),
"frame_rate": extra_params.get("frame_rate", extra_args.get("frame_rate", 24.0)),
"num_inference_steps": self.config.n_steps,
"video_guider_params": DEFAULT_VIDEO_GUIDER_PARAMS,
"audio_guider_params": DEFAULT_AUDIO_GUIDER_PARAMS,
"images": extra_params.get("images", []),
"tiling_config": extra_params.get("tiling_config", TilingConfig.default()),
}
decoded_video, decoded_audio = self.pipe(prompt=prompt, **kwargs)
# vae_decode_video returns a lazy generator — consume it so the
# video decoder's forward() actually runs during calibration.
for _ in decoded_video:
pass
def _run_ltx_video_calibration(
self, prompt_batch: list[str], extra_args: dict[str, Any]
) -> None:
"""
Run calibration for LTX-Video model using the full multi-stage pipeline.
Args:
prompt_batch: Batch of prompts
extra_args: Model-specific arguments
"""
# Extract specific args for LTX-Video
expected_height = extra_args.get("height", 512)
expected_width = extra_args.get("width", 704)
num_frames = extra_args.get("num_frames", 121)
negative_prompt = extra_args.get(
"negative_prompt", "worst quality, inconsistent motion, blurry, jittery, distorted"
)
def round_to_nearest_resolution_acceptable_by_vae(height, width):
height = height - (height % self.pipe.vae_spatial_compression_ratio)
width = width - (width % self.pipe.vae_spatial_compression_ratio)
return height, width
downscale_factor = 2 / 3
# Part 1: Generate video at smaller resolution
downscaled_height, downscaled_width = (
int(expected_height * downscale_factor),
int(expected_width * downscale_factor),
)
downscaled_height, downscaled_width = round_to_nearest_resolution_acceptable_by_vae(
downscaled_height, downscaled_width
)
# Generate initial latents at lower resolution
latents = self.pipe(
conditions=None,
prompt=prompt_batch,
negative_prompt=negative_prompt,
width=downscaled_width,
height=downscaled_height,
num_frames=num_frames,
num_inference_steps=self.config.n_steps,
output_type="latent",
).frames
# Part 2: Upscale generated video using latent upsampler (if available)
if self.pipe_upsample is not None:
_ = self.pipe_upsample(latents=latents, output_type="latent").frames
# Part 3: Denoise the upscaled video with few steps to improve texture
# However, in this example code, we will omit the upscale step since its optional.