mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### 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>
163 lines
5.1 KiB
Python
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)
|