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: Bug fix
Adds a compact SDXL and SDXL-Turbo mixed-precision FP4 recipe:
- block-16 NVFP4 for non-QKV Linear/GEMM layers;
- FP8 for Conv2d layers;
- high-precision Q/K/V projection Linears to preserve TensorRT
horizontal fusion;
- optional FP8 MHA quantization.
For SDXL FP4 export, Conv2d quantizers export directly through the
shared FP8 custom-op path. The previous `generate_fp8_scales` plus
`convert_zp_fp8` INT8 zero-point workaround is removed. The graph then
uses the existing FP8 Q/DQ normalization and `NVFP4QuantExporter`
lowering, with opset 23 for FLOAT4 support. Flux FP8 export also saves
the graph returned by its RoPE weight conversion.
This PR also changes shared exporter behavior:
- `_fp8_quantize` refreshes ONNX shape/type inference after applying the
custom FP8 operator's uint8 output metadata, affecting all FP8 ONNX
exports through this symbolic.
- `_quantized_sdpa` derives `disable_fp8_mha` from the live Q/K/V
quantizer state instead of a restored private module flag.
Other model recipe configurations remain unchanged.
### Usage
```bash
python quantize.py \
--model sdxl-1.0 \
--model-dtype Half \
--trt-high-precision-dtype Half \
--format fp4 \
--block-size 16 \
--batch-size 2 \
--calib-size 128 \
--n-steps 20 \
--quantized-torch-ckpt-save-path ./sdxl-fp4 \
--onnx-dir ./onnx-sdxl-fp4
```
### Testing
- CPU-only focused and generic NVFP4 exporter tests: 44 passed in 4.35
seconds.
- Focused Flux returned-graph save test: 1 passed.
- Required Linux unit CI at `034fe23ec` passed with the `all` dependency
set, including `tests/unit/examples/test_diffusers_fp4.py`.
- Latest changed-file pre-commit checks: all passed.
- TensorRT 10.14 on a B200 GPU:
- 302 native block-scaled NVFP4 GEMM tactics;
- 38 native FP8 Conv tactics;
- no FP4 Q/K/V projections;
- all 11 FP16 Q/K/V projection-fusion groups preserved;
- three alternating batch-2 profiles measured 18.614 ms FP4 versus
20.028 ms FP16 median UNet latency, a 7.06% reduction.
- FP8 SDXL/SD3 ONNX-to-TensorRT end-to-end runs were not executed
because they require explicit approval. The existing end-to-end test
matrix now includes SD3 FP8 alongside SDXL FP8.
### 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?: ✅ — no public API or CLI flags
change; the shared changes preserve the intended FP8 export and
attention behavior.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ — the shared NVFP4 opset, FP8 shape-inference, and Diffusers
attention-policy changes are recorded under bug fixes.
- Did you get Claude approval on this PR?: N/A
### Additional Information
Tracking: [5565357]
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
- **New Features**
- Added SDXL support for mixed NVFP4/FP8 quantization, including
convolution and softmax handling.
- Added an SDXL quantization preset for streamlined post-training
quantization workflows.
- Expanded FP4 ONNX export support to Flux and SDXL, with improved
FP4/FP8 graph processing and export reliability.
- Added automatic quantization policy and format restoration from
checkpoints.
- **Documentation**
- Documented SDXL layer behavior, optional FP8 attention quantization,
and Blackwell/TensorRT requirements for FP4 and FP8 deployment.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
> 🤖 _Generated by Codex (AI agent)._
---------
Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com>
Co-authored-by: Codex <codex@openai.com>
761 lines
28 KiB
Python
761 lines
28 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 argparse
|
|
import copy
|
|
import logging
|
|
import sys
|
|
import time as time
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import torch
|
|
from calibration import Calibrator
|
|
from config import (
|
|
FP8_DEFAULT_CONFIG,
|
|
INT8_DEFAULT_CONFIG,
|
|
NVFP4_DEFAULT_CONFIG,
|
|
NVFP4_FP8_CONV_CONFIG,
|
|
NVFP4_FP8_MHA_CONFIG,
|
|
reset_set_int8_config,
|
|
set_quant_config_attr,
|
|
)
|
|
from diffusers import DiffusionPipeline
|
|
from models_utils import (
|
|
MODEL_DEFAULTS,
|
|
ModelType,
|
|
build_block_range_quant_cfg,
|
|
get_model_filter_func,
|
|
parse_extra_params,
|
|
)
|
|
from pipeline_manager import PipelineManager
|
|
from quantize_config import (
|
|
CalibrationConfig,
|
|
CollectMethod,
|
|
DataType,
|
|
ExportConfig,
|
|
ModelConfig,
|
|
QuantAlgo,
|
|
QuantFormat,
|
|
QuantizationConfig,
|
|
)
|
|
from utils import check_conv_and_mha, check_lora
|
|
|
|
import modelopt.torch.opt as mto
|
|
import modelopt.torch.quantization as mtq
|
|
from modelopt.torch.export import export_hf_checkpoint
|
|
from modelopt.torch.quantization.nn import TensorQuantizer
|
|
|
|
_SDXL_MODEL_TYPES = (ModelType.SDXL_BASE, ModelType.SDXL_TURBO)
|
|
|
|
|
|
def setup_logging(verbose: bool = False) -> logging.Logger:
|
|
"""
|
|
Set up logging configuration.
|
|
|
|
Args:
|
|
verbose: Enable verbose logging
|
|
|
|
Returns:
|
|
Configured logger instance
|
|
"""
|
|
log_level = logging.DEBUG if verbose else logging.INFO
|
|
|
|
# Create custom formatter
|
|
formatter = logging.Formatter(
|
|
fmt="%(asctime)s | %(levelname)-8s | %(name)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S"
|
|
)
|
|
|
|
# Set up console handler
|
|
console_handler = logging.StreamHandler(sys.stdout)
|
|
console_handler.setFormatter(formatter)
|
|
|
|
# Configure root logger
|
|
logger = logging.getLogger(__name__)
|
|
logger.setLevel(log_level)
|
|
logger.addHandler(console_handler)
|
|
|
|
# Optionally reduce noise from other libraries
|
|
logging.getLogger("diffusers").setLevel(logging.WARNING)
|
|
logging.getLogger("transformers").setLevel(logging.WARNING)
|
|
|
|
return logger
|
|
|
|
|
|
class Quantizer:
|
|
"""Handles model quantization operations."""
|
|
|
|
def __init__(
|
|
self, config: QuantizationConfig, model_config: ModelConfig, logger: logging.Logger
|
|
):
|
|
"""
|
|
Initialize quantizer.
|
|
|
|
Args:
|
|
config: Quantization configuration
|
|
model_config: Model configuration
|
|
logger: Logger instance
|
|
"""
|
|
self.config = config
|
|
self.model_config = model_config
|
|
self.logger = logger
|
|
|
|
def get_quant_config(self, n_steps: int, backbone: torch.nn.Module) -> Any:
|
|
"""
|
|
Build quantization configuration based on format.
|
|
|
|
Args:
|
|
n_steps: Number of denoising steps
|
|
|
|
Returns:
|
|
Quantization configuration object
|
|
"""
|
|
self.logger.info(f"Building quantization config for {self.config.format.value}")
|
|
|
|
apply_int8_percentile_calibrator = False
|
|
if self.config.format == QuantFormat.INT8:
|
|
if self.config.algo == QuantAlgo.SMOOTHQUANT:
|
|
base_cfg = mtq.INT8_SMOOTHQUANT_CFG
|
|
else:
|
|
base_cfg = INT8_DEFAULT_CONFIG
|
|
apply_int8_percentile_calibrator = self.config.collect_method != CollectMethod.DEFAULT
|
|
elif self.config.format == QuantFormat.FP8:
|
|
base_cfg = FP8_DEFAULT_CONFIG
|
|
elif self.config.format == QuantFormat.FP4:
|
|
if self.model_config.model_type in _SDXL_MODEL_TYPES:
|
|
base_cfg = NVFP4_FP8_CONV_CONFIG
|
|
elif self.model_config.model_type.value.startswith("flux"):
|
|
base_cfg = NVFP4_FP8_MHA_CONFIG
|
|
else:
|
|
base_cfg = NVFP4_DEFAULT_CONFIG
|
|
else:
|
|
raise NotImplementedError(f"Unknown format {self.config.format}")
|
|
|
|
# Build a fresh config dict so runtime overrides never mutate the global constants.
|
|
base_cfg = copy.deepcopy(base_cfg)
|
|
|
|
if apply_int8_percentile_calibrator:
|
|
reset_set_int8_config(
|
|
base_cfg,
|
|
self.config.percentile,
|
|
n_steps,
|
|
collect_method=self.config.collect_method.value,
|
|
backbone=backbone,
|
|
)
|
|
|
|
quant_cfg_list = list(base_cfg["quant_cfg"])
|
|
|
|
if self.config.format == QuantFormat.FP4:
|
|
for i, entry in enumerate(quant_cfg_list):
|
|
if isinstance(entry, dict) and "block_sizes" in entry.get("cfg", {}):
|
|
new_block_sizes = {**entry["cfg"]["block_sizes"], -1: self.config.block_size}
|
|
quant_cfg_list[i] = {
|
|
**entry,
|
|
"cfg": {**entry["cfg"], "block_sizes": new_block_sizes},
|
|
}
|
|
|
|
if self.config.quantize_mha:
|
|
quant_cfg_list.append(
|
|
{
|
|
"quantizer_name": "*[qkv]_bmm_quantizer",
|
|
"cfg": {"num_bits": (4, 3), "axis": None},
|
|
}
|
|
)
|
|
|
|
# Apply the transformer-block-range recipe (e.g. Qwen-Image) BEFORE
|
|
# calibration. This restricts quantization to `transformer_blocks` and
|
|
# excludes the first/last N blocks. It must run before calibration so that
|
|
# SVDQuant does not mutate the weights of the excluded blocks. The recipe
|
|
# is format-agnostic (applies to FP8/NVFP4/SVDQuant alike).
|
|
block_range = MODEL_DEFAULTS.get(self.model_config.model_type, {}).get("block_range")
|
|
if block_range is not None:
|
|
recipe_rules = build_block_range_quant_cfg(
|
|
backbone,
|
|
exclude_first_n=block_range.get("exclude_first_n", 2),
|
|
exclude_last_n=block_range.get("exclude_last_n", 2),
|
|
block_module=block_range.get("block_module", "transformer_blocks"),
|
|
)
|
|
self.logger.info(
|
|
f"Applying block-range recipe ({len(recipe_rules)} rules) for "
|
|
f"{self.model_config.model_type.value}: quantize only "
|
|
f"'{block_range.get('block_module', 'transformer_blocks')}' excluding "
|
|
f"first {block_range.get('exclude_first_n', 2)} / last "
|
|
f"{block_range.get('exclude_last_n', 2)} blocks."
|
|
)
|
|
quant_cfg_list.extend(recipe_rules)
|
|
|
|
# Per-model SVDQuant exclusions (e.g. Qwen-Image's text-stream linears):
|
|
# matching layers skip the SVDQuant low-rank branch and AWQ smoothing but
|
|
# stay quantized with plain max calibration.
|
|
svdquant_skip_layers = None
|
|
if self.config.algo == QuantAlgo.SVDQUANT:
|
|
svdquant_skip_layers = MODEL_DEFAULTS.get(self.model_config.model_type, {}).get(
|
|
"svdquant_skip_layers"
|
|
)
|
|
if svdquant_skip_layers:
|
|
self.logger.info(
|
|
f"SVDQuant skip patterns for {self.model_config.model_type.value} "
|
|
f"(plain quantization): {svdquant_skip_layers}"
|
|
)
|
|
|
|
quant_config = {**base_cfg, "quant_cfg": quant_cfg_list}
|
|
set_quant_config_attr(
|
|
quant_config,
|
|
self.model_config.trt_high_precision_dtype.value,
|
|
self.config.algo.value,
|
|
alpha=self.config.alpha,
|
|
lowrank=self.config.lowrank,
|
|
skip_layers=svdquant_skip_layers,
|
|
)
|
|
self.logger.info(f"Quant config {quant_config}")
|
|
return quant_config
|
|
|
|
def quantize_model(
|
|
self,
|
|
backbone: torch.nn.Module,
|
|
quant_config: Any,
|
|
forward_loop: callable, # type: ignore[valid-type]
|
|
backbone_name: str = "transformer",
|
|
) -> torch.nn.Module:
|
|
"""
|
|
Apply quantization to the model.
|
|
|
|
Args:
|
|
backbone: Model backbone to quantize
|
|
quant_config: Quantization configuration
|
|
forward_loop: Forward pass function for calibration
|
|
backbone_name: Name of the backbone being quantized
|
|
"""
|
|
self.logger.info("Checking for LoRA layers...")
|
|
check_lora(backbone)
|
|
|
|
self.logger.info(f"Starting model quantization for {backbone_name}...")
|
|
mtq.quantize(backbone, quant_config, forward_loop)
|
|
# Get model-specific filter function
|
|
model_filter_func = get_model_filter_func(self.model_config.model_type, backbone_name)
|
|
self.logger.info(
|
|
f"Using filter function for {self.model_config.model_type.value}/{backbone_name}"
|
|
)
|
|
|
|
self.logger.info("Disabling specific quantizers...")
|
|
mtq.disable_quantizer(backbone, model_filter_func)
|
|
|
|
self.logger.info("Quantization completed successfully")
|
|
return backbone
|
|
|
|
|
|
class ExportManager:
|
|
"""Handles model export operations."""
|
|
|
|
def __init__(
|
|
self,
|
|
config: ExportConfig,
|
|
logger: logging.Logger,
|
|
pipeline_manager: PipelineManager | None = None,
|
|
):
|
|
"""
|
|
Initialize export manager.
|
|
|
|
Args:
|
|
config: Export configuration
|
|
logger: Logger instance
|
|
pipeline_manager: Pipeline manager for per-backbone IO
|
|
"""
|
|
self.config = config
|
|
self.logger = logger
|
|
self.pipeline_manager = pipeline_manager
|
|
|
|
def save_checkpoint(
|
|
self,
|
|
backbone: torch.nn.Module,
|
|
backbone_name: str | None = None,
|
|
) -> None:
|
|
"""
|
|
Save quantized model checkpoint.
|
|
|
|
Args:
|
|
backbone: The quantized backbone module to save (must be the same instance
|
|
that was passed to mtq.quantize, as it carries the _modelopt_state).
|
|
backbone_name: Optional name for the backbone file (defaults to "backbone").
|
|
"""
|
|
if not self.config.quantized_torch_ckpt_path:
|
|
return
|
|
|
|
ckpt_path = self.config.quantized_torch_ckpt_path
|
|
ckpt_path.mkdir(parents=True, exist_ok=True)
|
|
filename = f"{backbone_name}.pt" if backbone_name else "backbone.pt"
|
|
target_path = ckpt_path / filename
|
|
|
|
self.logger.info(f"Saving backbone to {target_path}")
|
|
mto.save(backbone, str(target_path))
|
|
|
|
self.logger.info("Checkpoint saved successfully")
|
|
|
|
def export_onnx(
|
|
self,
|
|
pipe: DiffusionPipeline,
|
|
backbone: torch.nn.Module,
|
|
model_type: ModelType,
|
|
quant_format: QuantFormat,
|
|
) -> None:
|
|
"""
|
|
Export model to ONNX format.
|
|
|
|
Args:
|
|
pipe: Diffusion pipeline
|
|
backbone: Model backbone
|
|
model_type: Type of model
|
|
quant_format: Quantization format
|
|
"""
|
|
if not self.config.onnx_dir:
|
|
return
|
|
|
|
# Deferred: the ONNX stack (onnx, onnx_graphsurgeon, ...) is only needed
|
|
# for --onnx-dir exports; HF-checkpoint-only runs must not require it.
|
|
from onnx_utils.export import modelopt_export_sd
|
|
|
|
self.logger.info(f"Starting ONNX export to {self.config.onnx_dir}")
|
|
|
|
self.logger.info("Preparing models for export...")
|
|
pipe.to("cpu")
|
|
torch.cuda.empty_cache()
|
|
backbone.to("cuda")
|
|
# Export to ONNX
|
|
backbone.eval()
|
|
with torch.no_grad():
|
|
self.logger.info("Exporting to ONNX...")
|
|
modelopt_export_sd(
|
|
backbone, str(self.config.onnx_dir), model_type.value, quant_format.value
|
|
)
|
|
|
|
self.logger.info("ONNX export completed successfully")
|
|
|
|
def restore_checkpoint(self) -> None:
|
|
"""
|
|
Restore a previously quantized model.
|
|
|
|
"""
|
|
if not self.config.restore_from:
|
|
return
|
|
|
|
restore_path = self.config.restore_from
|
|
if self.pipeline_manager is None:
|
|
raise RuntimeError("Pipeline manager is required for per-backbone checkpoints.")
|
|
|
|
if not restore_path.exists() or not restore_path.is_dir():
|
|
raise FileNotFoundError(f"Checkpoint directory not found: {restore_path}")
|
|
|
|
for backbone_name, backbone in self.pipeline_manager.iter_backbones():
|
|
source_path = restore_path / f"{backbone_name}.pt"
|
|
if not source_path.exists():
|
|
raise FileNotFoundError(
|
|
f"Checkpoint not found for '{backbone_name}' in {restore_path}"
|
|
)
|
|
self.logger.info(f"Restoring {backbone_name} from {source_path}")
|
|
mto.restore(backbone, str(source_path))
|
|
|
|
self.logger.info("Checkpoints restored successfully")
|
|
|
|
# TODO: should not do the any data type
|
|
def export_hf_ckpt(self, pipe: Any, model_config: ModelConfig | None = None) -> None:
|
|
"""
|
|
Export quantized model to HuggingFace checkpoint format.
|
|
|
|
Args:
|
|
pipe: Diffusion pipeline containing the quantized model
|
|
model_config: Model configuration (used to pass model-specific export kwargs)
|
|
"""
|
|
if not self.config.hf_ckpt_dir:
|
|
return
|
|
|
|
self.logger.info(f"Exporting HuggingFace checkpoint to {self.config.hf_ckpt_dir}")
|
|
kwargs: dict[str, Any] = {}
|
|
if model_config and model_config.model_type == ModelType.LTX2:
|
|
merged_path = model_config.extra_params.get("merged_base_safetensor_path")
|
|
if merged_path:
|
|
self.logger.info(f"Merging base safetensors from {merged_path} for LTX2 export")
|
|
kwargs["merged_base_safetensor_path"] = merged_path
|
|
if model_config:
|
|
for key in ("enable_swizzle_layout", "enable_layerwise_quant_metadata"):
|
|
val = model_config.extra_params.get(key)
|
|
if val is not None:
|
|
normalized = str(val).strip().lower()
|
|
if normalized in ("true", "1", "yes"):
|
|
kwargs[key] = True
|
|
elif normalized in ("false", "0", "no"):
|
|
kwargs[key] = False
|
|
else:
|
|
raise ValueError(
|
|
f"Invalid value for {key}: {val!r}. "
|
|
"Expected true/false, 1/0, or yes/no."
|
|
)
|
|
padding = model_config.extra_params.get("padding_strategy")
|
|
if padding is not None:
|
|
padding = str(padding).strip().lower()
|
|
if padding not in ("row", "row_col"):
|
|
raise ValueError(
|
|
f"Invalid padding_strategy: {padding!r}. Expected 'row' or 'row_col'."
|
|
)
|
|
kwargs["padding_strategy"] = padding
|
|
export_hf_checkpoint(pipe, export_dir=self.config.hf_ckpt_dir, **kwargs)
|
|
self.logger.info("HuggingFace checkpoint export completed successfully")
|
|
|
|
|
|
def create_argument_parser() -> argparse.ArgumentParser:
|
|
"""
|
|
Create and configure argument parser.
|
|
|
|
Returns:
|
|
Configured argument parser
|
|
"""
|
|
parser = argparse.ArgumentParser(
|
|
description="Enhanced Diffusion Model Quantization Tool",
|
|
formatter_class=argparse.RawDescriptionHelpFormatter,
|
|
epilog="""
|
|
Examples:
|
|
# Basic INT8 quantization with SmoothQuant
|
|
%(prog)s --model flux-dev --format int8 --quant-algo smoothquant --collect-method global_min
|
|
|
|
# FP8 quantization with ONNX export
|
|
%(prog)s --model sd3-medium --format fp8 --onnx-dir ./onnx_models/
|
|
|
|
# FP8 quantization with weight compression (reduces memory footprint)
|
|
%(prog)s --model flux-dev --format fp8 --compress
|
|
|
|
# Quantize LTX-Video model with full multi-stage pipeline
|
|
%(prog)s --model ltx-video-dev --format fp8 --batch-size 1 --calib-size 32
|
|
|
|
# Faster LTX-Video quantization (skip upsampler)
|
|
%(prog)s --model ltx-video-dev --format fp8 --batch-size 1 --calib-size 32 --ltx-skip-upsampler
|
|
|
|
# Restore and export a previously quantized model
|
|
%(prog)s --model flux-schnell --restore-from ./checkpoints/ --onnx-dir ./exports/
|
|
""",
|
|
)
|
|
model_group = parser.add_argument_group("Model Configuration")
|
|
model_group.add_argument(
|
|
"--model",
|
|
type=str,
|
|
default="flux-dev",
|
|
choices=[m.value for m in ModelType],
|
|
help="Model to load and quantize",
|
|
)
|
|
model_group.add_argument(
|
|
"--backbone",
|
|
nargs="+",
|
|
default=None,
|
|
help=(
|
|
"Model backbone(s) in the DiffusionPipeline to quantize. "
|
|
"Provide one or more names (e.g., 'transformer', 'video_decoder'). "
|
|
"If not provided, uses default based on model type."
|
|
),
|
|
)
|
|
model_group.add_argument(
|
|
"--model-dtype",
|
|
type=str,
|
|
default="Half",
|
|
choices=[d.value for d in DataType],
|
|
help="Precision for loading the pipeline. If you want different dtypes for separate components, "
|
|
"please specify using --component-dtype",
|
|
)
|
|
model_group.add_argument(
|
|
"--component-dtype",
|
|
action="append",
|
|
default=[],
|
|
help="Precision for loading each component of the model by format of name:dtype. "
|
|
"You can specify multiple components. "
|
|
"Example: --component-dtype vae:Half --component-dtype transformer:BFloat16",
|
|
)
|
|
model_group.add_argument(
|
|
"--override-model-path", type=str, help="Custom path to model (overrides default)"
|
|
)
|
|
model_group.add_argument(
|
|
"--cpu-offloading", action="store_true", help="Enable CPU offloading for limited VRAM"
|
|
)
|
|
model_group.add_argument(
|
|
"--ltx-skip-upsampler",
|
|
action="store_true",
|
|
help="Skip upsampler pipeline for LTX-Video (faster calibration, only quantizes main transformer)",
|
|
)
|
|
model_group.add_argument(
|
|
"--extra-param",
|
|
action="append",
|
|
default=[],
|
|
metavar="KEY=VALUE",
|
|
help=(
|
|
"Extra model-specific parameters in KEY=VALUE form. Can be provided multiple times. "
|
|
"These override model-specific CLI arguments when present."
|
|
),
|
|
)
|
|
quant_group = parser.add_argument_group("Quantization Configuration")
|
|
quant_group.add_argument(
|
|
"--format",
|
|
type=str,
|
|
default="int8",
|
|
choices=[f.value for f in QuantFormat],
|
|
help="Quantization format",
|
|
)
|
|
quant_group.add_argument(
|
|
"--quant-algo",
|
|
type=str,
|
|
default="max",
|
|
choices=[a.value for a in QuantAlgo],
|
|
help="Quantization algorithm",
|
|
)
|
|
quant_group.add_argument(
|
|
"--percentile",
|
|
type=float,
|
|
default=1.0,
|
|
help="Percentile for calibration, works for INT8, not including smoothquant",
|
|
)
|
|
quant_group.add_argument(
|
|
"--collect-method",
|
|
type=str,
|
|
default="default",
|
|
choices=[c.value for c in CollectMethod],
|
|
help="Calibration collection method, works for INT8, not including smoothquant",
|
|
)
|
|
quant_group.add_argument("--alpha", type=float, default=1.0, help="SmoothQuant alpha parameter")
|
|
quant_group.add_argument("--lowrank", type=int, default=32, help="SVDQuant lowrank parameter")
|
|
quant_group.add_argument(
|
|
"--quantize-mha", action="store_true", help="Quantizing MHA into FP8 if its True"
|
|
)
|
|
quant_group.add_argument(
|
|
"--compress",
|
|
action="store_true",
|
|
help="Compress quantized weights to reduce memory footprint (FP8/FP4 only)",
|
|
)
|
|
quant_group.add_argument(
|
|
"--block-size",
|
|
type=int,
|
|
default=16,
|
|
help="Block size for NVFP4 quantization (default: 16)",
|
|
)
|
|
|
|
calib_group = parser.add_argument_group("Calibration Configuration")
|
|
calib_group.add_argument("--batch-size", type=int, default=2, help="Batch size for calibration")
|
|
calib_group.add_argument(
|
|
"--calib-size", type=int, default=128, help="Total number of calibration samples"
|
|
)
|
|
calib_group.add_argument("--n-steps", type=int, default=30, help="Number of denoising steps")
|
|
calib_group.add_argument(
|
|
"--prompts-file",
|
|
type=str,
|
|
default=None,
|
|
help="Calibrate using prompts in the file instead of the default dataset.",
|
|
)
|
|
|
|
export_group = parser.add_argument_group("Export Configuration")
|
|
export_group.add_argument(
|
|
"--quantized-torch-ckpt-save-path",
|
|
type=str,
|
|
help="Path to save quantized PyTorch checkpoint",
|
|
)
|
|
export_group.add_argument("--onnx-dir", type=str, help="Directory for ONNX export")
|
|
export_group.add_argument(
|
|
"--hf-ckpt-dir",
|
|
type=str,
|
|
help="Directory for HuggingFace checkpoint export",
|
|
)
|
|
export_group.add_argument(
|
|
"--restore-from",
|
|
type=str,
|
|
help="Checkpoint directory; quantization format and MHA policy are restored automatically",
|
|
)
|
|
export_group.add_argument(
|
|
"--trt-high-precision-dtype",
|
|
type=str,
|
|
default="Half",
|
|
choices=[d.value for d in DataType],
|
|
help="Precision for TensorRT high-precision layers",
|
|
)
|
|
parser.add_argument("--verbose", action="store_true", help="Enable verbose logging")
|
|
|
|
return parser
|
|
|
|
|
|
def _infer_restored_quantization_format(
|
|
backbones: list[tuple[str, torch.nn.Module]],
|
|
) -> QuantFormat:
|
|
has_nvfp4 = False
|
|
has_fp8 = False
|
|
|
|
for _, backbone in backbones:
|
|
for module in backbone.modules():
|
|
if isinstance(module, TensorQuantizer) and module.is_enabled:
|
|
has_nvfp4 |= module.is_nvfp4_dynamic or module.is_nvfp4_static
|
|
has_fp8 |= module.is_fp8
|
|
|
|
if has_nvfp4:
|
|
return QuantFormat.FP4
|
|
if has_fp8:
|
|
return QuantFormat.FP8
|
|
return QuantFormat.INT8
|
|
|
|
|
|
def main() -> None:
|
|
from diffusers.models.normalization import RMSNorm as DiffuserRMSNorm
|
|
|
|
torch.nn.RMSNorm = DiffuserRMSNorm
|
|
torch.nn.modules.normalization.RMSNorm = DiffuserRMSNorm
|
|
|
|
parser = create_argument_parser()
|
|
args, unknown_args = parser.parse_known_args()
|
|
|
|
model_type = ModelType(args.model)
|
|
if args.backbone is None:
|
|
args.backbone = [MODEL_DEFAULTS[model_type]["backbone"]]
|
|
s = time.time()
|
|
|
|
model_dtype = {"default": DataType(args.model_dtype).torch_dtype}
|
|
for component_dtype in args.component_dtype:
|
|
component, dtype = component_dtype.split(":")
|
|
model_dtype[component] = DataType(dtype).torch_dtype
|
|
|
|
logger = setup_logging(args.verbose)
|
|
logger.info("Starting Enhanced Diffusion Model Quantization")
|
|
|
|
try:
|
|
extra_params = parse_extra_params(args.extra_param, unknown_args, logger)
|
|
model_config = ModelConfig(
|
|
model_type=model_type,
|
|
model_dtype=model_dtype,
|
|
backbone=args.backbone,
|
|
trt_high_precision_dtype=DataType(args.trt_high_precision_dtype),
|
|
override_model_path=Path(args.override_model_path)
|
|
if args.override_model_path
|
|
else None,
|
|
cpu_offloading=args.cpu_offloading,
|
|
ltx_skip_upsampler=args.ltx_skip_upsampler,
|
|
extra_params=extra_params,
|
|
)
|
|
|
|
quant_config = QuantizationConfig(
|
|
format=QuantFormat(args.format),
|
|
algo=QuantAlgo(args.quant_algo),
|
|
percentile=args.percentile,
|
|
collect_method=CollectMethod(args.collect_method),
|
|
alpha=args.alpha,
|
|
lowrank=args.lowrank,
|
|
quantize_mha=args.quantize_mha,
|
|
compress=args.compress,
|
|
block_size=args.block_size,
|
|
)
|
|
|
|
if args.prompts_file is not None:
|
|
prompts_file = Path(args.prompts_file)
|
|
assert prompts_file.exists(), (
|
|
f"User specified prompts file {prompts_file} does not exist."
|
|
)
|
|
prompts_dataset = prompts_file
|
|
else:
|
|
prompts_dataset = MODEL_DEFAULTS[model_type]["dataset"]
|
|
calib_config = CalibrationConfig(
|
|
prompts_dataset=prompts_dataset,
|
|
batch_size=args.batch_size,
|
|
calib_size=args.calib_size,
|
|
n_steps=args.n_steps,
|
|
)
|
|
|
|
export_config = ExportConfig(
|
|
quantized_torch_ckpt_path=Path(args.quantized_torch_ckpt_save_path)
|
|
if args.quantized_torch_ckpt_save_path
|
|
else None,
|
|
onnx_dir=Path(args.onnx_dir) if args.onnx_dir else None,
|
|
hf_ckpt_dir=Path(args.hf_ckpt_dir) if args.hf_ckpt_dir else None,
|
|
restore_from=Path(args.restore_from) if args.restore_from else None,
|
|
)
|
|
|
|
logger.info("Validating configurations...")
|
|
export_config.validate()
|
|
if not export_config.restore_from:
|
|
quant_config.validate()
|
|
calib_config.validate()
|
|
|
|
pipeline_manager = PipelineManager(model_config, logger)
|
|
pipe = pipeline_manager.create_pipeline()
|
|
pipeline_manager.setup_device()
|
|
|
|
export_manager = ExportManager(export_config, logger, pipeline_manager)
|
|
|
|
if export_config.restore_from:
|
|
export_manager.restore_checkpoint()
|
|
quant_config.format = _infer_restored_quantization_format(
|
|
list(pipeline_manager.iter_backbones())
|
|
)
|
|
logger.info(f"Detected restored quantization format: {quant_config.format.value}")
|
|
|
|
else:
|
|
logger.info("Initializing calibration...")
|
|
calibrator = Calibrator(pipeline_manager, calib_config, model_config.model_type, logger)
|
|
batched_prompts = calibrator.load_and_batch_prompts()
|
|
quantizer = Quantizer(quant_config, model_config, logger)
|
|
|
|
for backbone_name, backbone in pipeline_manager.iter_backbones():
|
|
logger.info(f"Quantizing backbone: {backbone_name}")
|
|
backbone_quant_config = quantizer.get_quant_config(calib_config.n_steps, backbone)
|
|
|
|
# Calibration runs the full pipeline (not just `mod`), so the
|
|
# closure intentionally ignores the backbone argument.
|
|
def forward_loop(mod):
|
|
calibrator.run_calibration(batched_prompts)
|
|
|
|
quantizer.quantize_model(
|
|
backbone,
|
|
backbone_quant_config,
|
|
forward_loop,
|
|
backbone_name=backbone_name,
|
|
)
|
|
|
|
# Compress model weights if requested (only for FP8/FP4)
|
|
if quant_config.compress:
|
|
logger.info(f"Compressing {backbone_name} weights...")
|
|
mtq.compress(backbone)
|
|
logger.info(f"{backbone_name} compression completed")
|
|
|
|
if backbone_name not in ("video_decoder", "vae"):
|
|
check_conv_and_mha(
|
|
backbone,
|
|
quant_config.format == QuantFormat.FP4
|
|
and model_config.model_type not in _SDXL_MODEL_TYPES,
|
|
quant_config.quantize_mha,
|
|
)
|
|
|
|
export_manager.save_checkpoint(backbone, backbone_name)
|
|
|
|
pipeline_manager.print_quant_summary()
|
|
|
|
for backbone_name, backbone in pipeline_manager.iter_backbones():
|
|
export_manager.export_onnx(
|
|
pipe,
|
|
backbone,
|
|
model_config.model_type,
|
|
quant_config.format,
|
|
)
|
|
|
|
export_manager.export_hf_ckpt(pipe, model_config)
|
|
|
|
logger.info(
|
|
f"Quantization process completed successfully! Time taken = {time.time() - s} seconds"
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Quantization failed: {e}", exc_info=True)
|
|
sys.exit(1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|