Files
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

163 lines
5.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 math
from dataclasses import dataclass, field
from enum import Enum
from pathlib import Path
from typing import Any
import torch
from models_utils import MODEL_REGISTRY, ModelType
class DataType(str, Enum):
"""Supported data types for model loading."""
HALF = "Half"
BFLOAT16 = "BFloat16"
FLOAT = "Float"
@property
def torch_dtype(self) -> torch.dtype:
return self._dtype_map[self.value]
DataType._dtype_map = {
DataType.HALF: torch.float16,
DataType.BFLOAT16: torch.bfloat16,
DataType.FLOAT: torch.float32,
}
class QuantFormat(str, Enum):
"""Supported quantization formats."""
INT8 = "int8"
FP8 = "fp8"
FP4 = "fp4"
class QuantAlgo(str, Enum):
"""Supported quantization algorithms."""
MAX = "max"
SVDQUANT = "svdquant"
SMOOTHQUANT = "smoothquant"
class CollectMethod(str, Enum):
"""Calibration collection methods."""
GLOBAL_MIN = "global_min"
MIN_MAX = "min-max"
MIN_MEAN = "min-mean"
MEAN_MAX = "mean-max"
DEFAULT = "default"
@dataclass
class QuantizationConfig:
"""Configuration for model quantization."""
format: QuantFormat = QuantFormat.INT8
algo: QuantAlgo = QuantAlgo.MAX
percentile: float = 1.0
collect_method: CollectMethod = CollectMethod.DEFAULT
alpha: float = 1.0 # SmoothQuant alpha
lowrank: int = 32 # SVDQuant lowrank
quantize_mha: bool = False
compress: bool = False
block_size: int = 16 # NVFP4 block size
def validate(self) -> None:
"""Validate configuration consistency."""
if self.format == QuantFormat.FP8 and self.collect_method != CollectMethod.DEFAULT:
raise NotImplementedError("Only 'default' collect method is implemented for FP8.")
if self.quantize_mha and self.format == QuantFormat.INT8:
raise ValueError("MHA quantization is only supported for FP8, not INT8.")
if self.compress and self.format == QuantFormat.INT8:
raise ValueError("Compression is only supported for FP8 and FP4, not INT8.")
@dataclass
class CalibrationConfig:
"""Configuration for calibration process."""
prompts_dataset: dict | Path
batch_size: int = 2
calib_size: int = 128
n_steps: int = 30
def validate(self) -> None:
"""Validate calibration configuration."""
if self.batch_size <= 0:
raise ValueError("Batch size must be positive.")
if self.calib_size <= 0:
raise ValueError("Calibration size must be positive.")
if self.n_steps <= 0:
raise ValueError("Number of steps must be positive.")
@property
def num_batches(self) -> int:
"""Calculate number of calibration batches."""
return math.ceil(self.calib_size / self.batch_size)
@dataclass
class ModelConfig:
"""Configuration for model loading and inference."""
model_type: ModelType = ModelType.FLUX_DEV
model_dtype: dict[str, torch.dtype] = field(default_factory=lambda: {"default": torch.float16})
backbone: list[str] = field(default_factory=list)
trt_high_precision_dtype: DataType = DataType.HALF
override_model_path: Path | None = None
cpu_offloading: bool = False
ltx_skip_upsampler: bool = False # Skip upsampler for LTX-Video (faster calibration)
extra_params: dict[str, Any] = field(default_factory=dict)
@property
def model_path(self) -> str:
"""Get the model path (override or default)."""
if self.override_model_path:
return str(self.override_model_path)
return MODEL_REGISTRY[self.model_type]
@dataclass
class ExportConfig:
"""Configuration for model export."""
quantized_torch_ckpt_path: Path | None = None
onnx_dir: Path | None = None
hf_ckpt_dir: Path | None = None
restore_from: Path | None = None
def validate(self) -> None:
"""Validate export configuration."""
if self.restore_from and not self.restore_from.exists():
raise FileNotFoundError(f"Restore checkpoint not found: {self.restore_from}")
if self.quantized_torch_ckpt_path:
parent_dir = self.quantized_torch_ckpt_path.parent
if not parent_dir.exists():
parent_dir.mkdir(parents=True, exist_ok=True)
if self.onnx_dir and not self.onnx_dir.exists():
self.onnx_dir.mkdir(parents=True, exist_ok=True)
if self.hf_ckpt_dir and not self.hf_ckpt_dir.exists():
self.hf_ckpt_dir.mkdir(parents=True, exist_ok=True)