mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-03 11:49:47 +08:00
### What does this PR do? Type of change: New feature Adds **Qwen-Image** (`Qwen/Qwen-Image`, `QwenImageTransformer2DModel`) to the diffusers quantization example and exports HuggingFace checkpoints in three precisions — **FP8**, **NVFP4**, and **NVFP4 + SVDQuant** — through the unified HF export. - Registers `--model qwen-image` (lazy diffusers import; no `trust_remote_code`). - Transformer-block-range recipe: quantizes only the linears under `transformer_blocks`, keeping the **first 2 / last 2** blocks (and everything outside `transformer_blocks`) in original precision. Applied **before** calibration so SVDQuant never mutates the excluded blocks. Expressed with the top-level `enable` `QuantizerCfgEntry` field (disable-all → re-enable `transformer_blocks` → disable first/last-N). - SVDQuant export (AWQ-style): promotes quantizer-owned tensors to clean module-level safetensors keys at export time — `weight_quantizer.svdquant_lora_a/b → <module>.svdquant_lora_a/b` and `input_quantizer._pre_quant_scale → <module>.pre_quant_scale` — with a documented `NVFP4_SVD` `quantization_config` (`group_size`, `has_zero_point: false`, `pre_quant_scale: true`, `lora_rank`). **Core SVDQuant quantization code (`modelopt/torch/quantization`) is unchanged.** - Shared export-path change — **intentionally global** (applies to all diffusers exports — SDXL / Flux / Wan, not just Qwen; the full export suite was verified green on GB200): `hide_quantizers_from_state_dict` now strips quantizer state from *all* modules (not just quant-linears) so calibrated norm-layer input quantizers no longer leak `input_quantizer._amax`. (An earlier `max_shard_size` workaround was dropped after merging `main`: #1794 makes the ComfyUI layerwise-metadata post-processing a no-op unless explicitly opted in, so a default sharded export no longer hits the unsupported-sharded path.) ### Usage ```bash python examples/diffusers/quantization/quantize.py \ --model qwen-image --override-model-path <Qwen-Image> --model-dtype BFloat16 \ --format fp4 --quant-algo svdquant --lowrank 32 \ --calib-size 64 --n-steps 20 \ --hf-ckpt-dir <out> # FP8: --format fp8 --quant-algo max # NVFP4: --format fp4 --quant-algo max ``` ### Testing - Focused unit + example tests pass on GB200 (sm_100): block-range recipe, `NVFP4_SVD` config schema, SVDQuant forward/fold (LoRA stays on `weight_quantizer`), Qwen dummy-input / strict-QKV-fusion / promotion, pipeline loading, and the diffusers HF-export test for Qwen FP8 / NVFP4 / SVDQuant. - Full `tests/examples/diffusers/test_export_diffusers_hf_ckpt.py` is green (SDXL, Flux, Qwen, Wan2.2) — confirms the shared export changes do not regress other models. - End-to-end on the real `Qwen/Qwen-Image` (~20B): all three formats export valid HF checkpoints — only `transformer_blocks` 2..57 quantized, nothing outside, no quantizer/`_amax` leak, correct `weight_scale`(`_2`)/`input_scale`, promoted SVDQuant keys (rank-consistent shapes), and the expected `quantization_config`. ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ <!-- live-model LoRA storage unchanged; existing exports unaffected --> - 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?: ❌ <!-- new-example feature; add a CHANGELOG.rst entry if required --> - Did you get Claude approval on this PR?: ❌ <!-- run /claude review --> ### Additional Information All changes are confined to the diffusers example (`examples/diffusers/quantization`) plus the shared export path (`modelopt/torch/export`); the core quantization library is untouched. ### Follow-up (next step): fused-QKV SVDQuant for sglang / Nunchaku This export keeps attention `q/k/v` (and `add_q/k/v_proj`) as **separate** projections — the diffusers-native layout. That matches sglang's bf16 / FP8 / plain-NVFP4 paths (which also keep QKV separate) and ModelOpt/TRT-LLM consumers, so those load 1:1. sglang's **NVFP4-SVDQuant (Nunchaku)** path, however, builds a **fused** `to_qkv` with a *single* fused rank-r LoRA in Nunchaku-native format (`proj_down`/`proj_up`, `smooth_factor`, `wscales`/`wtscale`). Our per-projection tensors (`svdquant_lora_a/b` + `pre_quant_scale`; three independent rank-r decompositions) are not directly loadable there — and cannot be fused at load time, because the fp16 weight residual needed to derive a single fused rank-r is not preserved after export. **Planned next step:** an opt-in fused-QKV SVDQuant export mode that fuses q/k/v **before** SVDQuant calibration (yielding one rank-r over the fused weight) and emits a Nunchaku-compatible layout, enabling lower-latency fused-QKV inference in sglang. Tracked as a separate follow-up. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added Qwen-Image (`QWEN_IMAGE`) model quantization and Diffusers export support * Added NVFP4_SVD (SVDQuant) export configuration support * Added transformer block-range quantization recipes (exclude first/last blocks) * **Bug Fixes** * Improved missing-pipeline error messaging for Qwen-Image * Prevented quantizer-related tensor/buffer leakage by promoting and cleaning quantizer outputs during export * **Tests** * Added Qwen-Image HF checkpoint export tests and offline fixtures * Added unit coverage for SVDQuant promotion/clean state-dict keys <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Jingyu Xin <jingyux@nvidia.com> Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
389 lines
14 KiB
Python
389 lines
14 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
|
|
from collections.abc import Callable
|
|
from enum import Enum
|
|
from typing import Any
|
|
|
|
import torch
|
|
from diffusers import (
|
|
DiffusionPipeline,
|
|
FluxPipeline,
|
|
LTXConditionPipeline,
|
|
StableDiffusion3Pipeline,
|
|
WanPipeline,
|
|
)
|
|
|
|
try:
|
|
from diffusers import Flux2Pipeline
|
|
except ImportError:
|
|
Flux2Pipeline = None
|
|
|
|
# Qwen-Image classes were added in a recent diffusers release; import lazily so
|
|
# this example still imports on older diffusers versions.
|
|
try:
|
|
from diffusers import QwenImagePipeline
|
|
except ImportError:
|
|
QwenImagePipeline = None
|
|
from utils import (
|
|
filter_func_default,
|
|
filter_func_flux_dev,
|
|
filter_func_ltx2_vae,
|
|
filter_func_ltx_video,
|
|
filter_func_qwen_image,
|
|
filter_func_wan_vae,
|
|
filter_func_wan_video,
|
|
)
|
|
|
|
|
|
class ModelType(str, Enum):
|
|
"""Supported model types."""
|
|
|
|
SDXL_BASE = "sdxl-1.0"
|
|
SDXL_TURBO = "sdxl-turbo"
|
|
SD3_MEDIUM = "sd3-medium"
|
|
SD35_MEDIUM = "sd3.5-medium"
|
|
FLUX_DEV = "flux-dev"
|
|
FLUX_SCHNELL = "flux-schnell"
|
|
FLUX2_DEV = "flux2-dev"
|
|
LTX_VIDEO_DEV = "ltx-video-dev"
|
|
LTX2 = "ltx-2"
|
|
WAN22_T2V_14b = "wan2.2-t2v-14b"
|
|
WAN22_T2V_5b = "wan2.2-t2v-5b"
|
|
QWEN_IMAGE = "qwen-image"
|
|
|
|
|
|
_FILTER_FUNC_MAP: dict[ModelType, Callable[[str], bool]] = {
|
|
ModelType.FLUX_DEV: filter_func_flux_dev,
|
|
ModelType.FLUX2_DEV: filter_func_flux_dev,
|
|
ModelType.LTX_VIDEO_DEV: filter_func_ltx_video,
|
|
ModelType.LTX2: filter_func_ltx_video,
|
|
ModelType.WAN22_T2V_14b: filter_func_wan_video,
|
|
ModelType.WAN22_T2V_5b: filter_func_wan_video,
|
|
ModelType.QWEN_IMAGE: filter_func_qwen_image,
|
|
}
|
|
|
|
_VAE_FILTER_FUNC_MAP: dict[tuple[ModelType, str], Callable[[str], bool]] = {
|
|
(ModelType.LTX2, "video_decoder"): filter_func_ltx2_vae,
|
|
(ModelType.WAN22_T2V_14b, "vae"): filter_func_wan_vae,
|
|
(ModelType.WAN22_T2V_5b, "vae"): filter_func_wan_vae,
|
|
}
|
|
|
|
|
|
def get_model_filter_func(
|
|
model_type: ModelType, backbone_name: str = "transformer"
|
|
) -> Callable[[str], bool]:
|
|
"""Get the appropriate filter function for a given model type and backbone."""
|
|
vae_func = _VAE_FILTER_FUNC_MAP.get((model_type, backbone_name))
|
|
if vae_func is not None:
|
|
return vae_func
|
|
return _FILTER_FUNC_MAP.get(model_type, filter_func_default)
|
|
|
|
|
|
# Model registry with HuggingFace model IDs
|
|
MODEL_REGISTRY: dict[ModelType, str] = {
|
|
ModelType.SDXL_BASE: "stabilityai/stable-diffusion-xl-base-1.0",
|
|
ModelType.SDXL_TURBO: "stabilityai/sdxl-turbo",
|
|
ModelType.SD3_MEDIUM: "stabilityai/stable-diffusion-3-medium-diffusers",
|
|
ModelType.SD35_MEDIUM: "stabilityai/stable-diffusion-3.5-medium",
|
|
ModelType.FLUX_DEV: "black-forest-labs/FLUX.1-dev",
|
|
ModelType.FLUX_SCHNELL: "black-forest-labs/FLUX.1-schnell",
|
|
ModelType.FLUX2_DEV: "black-forest-labs/FLUX.2-dev",
|
|
ModelType.LTX_VIDEO_DEV: "Lightricks/LTX-Video-0.9.7-dev",
|
|
ModelType.LTX2: "Lightricks/LTX-2",
|
|
ModelType.WAN22_T2V_14b: "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
|
ModelType.WAN22_T2V_5b: "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
|
ModelType.QWEN_IMAGE: "Qwen/Qwen-Image",
|
|
}
|
|
|
|
MODEL_PIPELINE: dict[ModelType, type[DiffusionPipeline] | None] = {
|
|
ModelType.SDXL_BASE: DiffusionPipeline,
|
|
ModelType.SDXL_TURBO: DiffusionPipeline,
|
|
ModelType.SD3_MEDIUM: StableDiffusion3Pipeline,
|
|
ModelType.SD35_MEDIUM: StableDiffusion3Pipeline,
|
|
ModelType.FLUX_DEV: FluxPipeline,
|
|
ModelType.FLUX_SCHNELL: FluxPipeline,
|
|
ModelType.FLUX2_DEV: Flux2Pipeline,
|
|
ModelType.LTX_VIDEO_DEV: LTXConditionPipeline,
|
|
ModelType.LTX2: None,
|
|
ModelType.WAN22_T2V_14b: WanPipeline,
|
|
ModelType.WAN22_T2V_5b: WanPipeline,
|
|
ModelType.QWEN_IMAGE: QwenImagePipeline,
|
|
}
|
|
|
|
# Shared dataset configurations
|
|
_SD_PROMPTS_DATASET = {
|
|
"name": "Gustavosta/Stable-Diffusion-Prompts",
|
|
"split": "train",
|
|
"column": "Prompt",
|
|
}
|
|
|
|
_OPENVID_DATASET = {
|
|
"name": "nkp37/OpenVid-1M",
|
|
"split": "train",
|
|
"column": "caption",
|
|
}
|
|
|
|
# Model family base configurations
|
|
_SDXL_BASE_CONFIG: dict[str, Any] = {
|
|
"backbone": "unet",
|
|
"dataset": _SD_PROMPTS_DATASET,
|
|
}
|
|
|
|
_SD3_BASE_CONFIG: dict[str, Any] = {
|
|
"backbone": "transformer",
|
|
"dataset": _SD_PROMPTS_DATASET,
|
|
}
|
|
|
|
_FLUX_BASE_CONFIG: dict[str, Any] = {
|
|
"backbone": "transformer",
|
|
"dataset": _SD_PROMPTS_DATASET,
|
|
"inference_extra_args": {
|
|
"height": 1024,
|
|
"width": 1024,
|
|
"guidance_scale": 3.5,
|
|
"max_sequence_length": 512,
|
|
},
|
|
}
|
|
|
|
_WAN_BASE_CONFIG: dict[str, Any] = {
|
|
"backbone": "transformer",
|
|
"dataset": _OPENVID_DATASET,
|
|
}
|
|
|
|
# Model-specific default arguments for calibration
|
|
MODEL_DEFAULTS: dict[ModelType, dict[str, Any]] = {
|
|
ModelType.SDXL_BASE: _SDXL_BASE_CONFIG,
|
|
ModelType.SDXL_TURBO: _SDXL_BASE_CONFIG,
|
|
ModelType.SD3_MEDIUM: _SD3_BASE_CONFIG,
|
|
ModelType.SD35_MEDIUM: _SD3_BASE_CONFIG,
|
|
ModelType.FLUX_DEV: _FLUX_BASE_CONFIG,
|
|
ModelType.FLUX_SCHNELL: _FLUX_BASE_CONFIG,
|
|
ModelType.FLUX2_DEV: {
|
|
"backbone": "transformer",
|
|
"dataset": _SD_PROMPTS_DATASET,
|
|
"inference_extra_args": {
|
|
"height": 768,
|
|
"width": 1024,
|
|
"guidance_scale": 4.0,
|
|
},
|
|
},
|
|
ModelType.LTX_VIDEO_DEV: {
|
|
"backbone": "transformer",
|
|
"dataset": _OPENVID_DATASET,
|
|
"inference_extra_args": {
|
|
"height": 512,
|
|
"width": 704,
|
|
"num_frames": 121,
|
|
"negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted",
|
|
},
|
|
},
|
|
ModelType.LTX2: {
|
|
"backbone": "transformer",
|
|
"dataset": _OPENVID_DATASET,
|
|
"inference_extra_args": {
|
|
"height": 768,
|
|
"width": 1280,
|
|
"num_frames": 121,
|
|
"frame_rate": 24.0,
|
|
"negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted",
|
|
},
|
|
},
|
|
ModelType.WAN22_T2V_14b: {
|
|
**_WAN_BASE_CONFIG,
|
|
"from_pretrained_extra_args": {
|
|
"boundary_ratio": 0.875,
|
|
},
|
|
"inference_extra_args": {
|
|
"height": 720,
|
|
"width": 1280,
|
|
"num_frames": 81,
|
|
"fps": 16,
|
|
"guidance_scale": 4.0,
|
|
"guidance_scale_2": 3.0,
|
|
"negative_prompt": (
|
|
"vivid colors, overexposed, static, blurry details, subtitles, style, "
|
|
"work of art, painting, picture, still, overall grayish, worst quality, "
|
|
"low quality, JPEG artifacts, ugly, deformed, extra fingers, poorly drawn hands, "
|
|
"poorly drawn face, deformed, disfigured, deformed limbs, fused fingers, "
|
|
"static image, cluttered background, three legs, many people in the background, "
|
|
"walking backwards"
|
|
),
|
|
},
|
|
},
|
|
ModelType.WAN22_T2V_5b: {
|
|
**_WAN_BASE_CONFIG,
|
|
"inference_extra_args": {
|
|
"height": 512,
|
|
"width": 768,
|
|
"num_frames": 81,
|
|
"fps": 16,
|
|
"guidance_scale": 5.0,
|
|
"negative_prompt": (
|
|
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留" # noqa: RUF001
|
|
",丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体," # noqa: RUF001
|
|
"手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" # noqa: RUF001
|
|
),
|
|
},
|
|
},
|
|
ModelType.QWEN_IMAGE: {
|
|
"backbone": "transformer",
|
|
"dataset": _SD_PROMPTS_DATASET,
|
|
"inference_extra_args": {
|
|
"height": 1024,
|
|
"width": 1024,
|
|
},
|
|
# Quantize only ``transformer_blocks``; keep the first 2 and last 2 blocks
|
|
# (and everything outside ``transformer_blocks``) in original precision.
|
|
# Applied before calibration via ``build_block_range_quant_cfg`` so SVDQuant
|
|
# never mutates the excluded blocks' weights.
|
|
"block_range": {
|
|
"exclude_first_n": 2,
|
|
"exclude_last_n": 2,
|
|
"block_module": "transformer_blocks",
|
|
},
|
|
# The text-stream linears (joint-attention added-KV projections and the
|
|
# txt MLP) and the modulation linears cannot use the SVDQuant low-rank
|
|
# branch; they are exported as plain NVFP4 instead (no pre_quant_scale,
|
|
# no svdquant_lora_a/b). The remaining image-stream linears keep full
|
|
# SVDQuant.
|
|
"svdquant_skip_layers": [
|
|
"*.attn.add_q_proj",
|
|
"*.attn.add_k_proj",
|
|
"*.attn.add_v_proj",
|
|
"*.attn.to_add_out",
|
|
"*.txt_mlp.net.0.proj",
|
|
"*.txt_mlp.net.2",
|
|
"*.img_mod.1",
|
|
"*.txt_mod.1",
|
|
],
|
|
},
|
|
}
|
|
|
|
|
|
def _coerce_extra_param_value(value: str) -> Any:
|
|
lowered = value.lower()
|
|
if lowered in {"true", "false"}:
|
|
return lowered == "true"
|
|
try:
|
|
return int(value)
|
|
except ValueError:
|
|
pass
|
|
try:
|
|
return float(value)
|
|
except ValueError:
|
|
return value
|
|
|
|
|
|
def parse_extra_params(
|
|
kv_args: list[str], unknown_args: list[str], logger: logging.Logger
|
|
) -> dict[str, Any]:
|
|
extra_params: dict[str, Any] = {}
|
|
for item in kv_args:
|
|
if "=" not in item:
|
|
raise ValueError(f"Invalid --extra-param value: '{item}'. Expected KEY=VALUE.")
|
|
key, value = item.split("=", 1)
|
|
extra_params[key] = _coerce_extra_param_value(value)
|
|
|
|
i = 0
|
|
while i < len(unknown_args):
|
|
token = unknown_args[i]
|
|
if token.startswith("--extra_param."):
|
|
key = token[len("--extra_param.") :]
|
|
value = "true"
|
|
if i + 1 < len(unknown_args) and not unknown_args[i + 1].startswith("--"):
|
|
value = unknown_args[i + 1]
|
|
i += 1
|
|
extra_params[key] = _coerce_extra_param_value(value)
|
|
elif token.startswith("--extra_param"):
|
|
raise ValueError(
|
|
"Use --extra_param.KEY VALUE or --extra-param KEY=VALUE for extra parameters."
|
|
)
|
|
else:
|
|
logger.warning("Ignoring unknown argument: %s", token)
|
|
i += 1
|
|
|
|
return extra_params
|
|
|
|
|
|
def build_block_range_quant_cfg(
|
|
backbone: torch.nn.Module,
|
|
exclude_first_n: int,
|
|
exclude_last_n: int,
|
|
block_module: str = "transformer_blocks",
|
|
) -> list[dict[str, Any]]:
|
|
"""Build ordered ``quant_cfg`` rules for a transformer-block-only recipe.
|
|
|
|
The rules quantize only the linears under ``block_module`` while keeping the
|
|
first ``exclude_first_n`` and last ``exclude_last_n`` blocks -- and everything
|
|
outside ``block_module`` -- in original precision.
|
|
|
|
The rules are meant to be appended to the ``quant_cfg`` list consumed by
|
|
``mtq.quantize`` so the selection is applied BEFORE calibration. This is
|
|
required for SVDQuant, whose calibration subtracts a low-rank residual from
|
|
the weights of every *enabled* linear: disabling the excluded blocks only
|
|
after calibration would leave their weights mutated instead of bit-identical
|
|
to the original precision.
|
|
|
|
Rules are applied in order with later rules overriding earlier ones:
|
|
1. disable every linear weight/input quantizer,
|
|
2. re-enable only those under ``block_module`` (``enable`` is a top-level
|
|
QuantizerCfgEntry toggle; a ``None`` cfg keeps the base preset's quant params),
|
|
3. disable the first/last ``n`` blocks.
|
|
|
|
Raises:
|
|
ValueError: if the backbone has no ``block_module`` list, or it has fewer
|
|
than ``exclude_first_n + exclude_last_n + 2`` blocks (it requires at
|
|
least two quantized middle blocks).
|
|
"""
|
|
blocks = getattr(backbone, block_module, None)
|
|
if blocks is None or not hasattr(blocks, "__len__"):
|
|
raise ValueError(
|
|
f"Backbone {type(backbone).__name__} has no '{block_module}' module list; "
|
|
"cannot build the transformer-block-range recipe."
|
|
)
|
|
num_blocks = len(blocks)
|
|
# Require at least two quantized middle blocks so the recipe actually
|
|
# quantizes something (excluding first/last alone could otherwise leave 0-1
|
|
# quantized blocks). For the default 2+2 recipe this means n >= 6.
|
|
min_blocks = exclude_first_n + exclude_last_n + 2
|
|
if num_blocks < min_blocks:
|
|
raise ValueError(
|
|
f"'{block_module}' has only {num_blocks} block(s); excluding the first "
|
|
f"{exclude_first_n} and last {exclude_last_n} requires at least {min_blocks} blocks "
|
|
f"(at least 2 quantized middle blocks)."
|
|
)
|
|
|
|
excluded = sorted(
|
|
set(range(exclude_first_n)) | set(range(num_blocks - exclude_last_n, num_blocks))
|
|
)
|
|
# `enable` is a top-level QuantizerCfgEntry field (independent of `cfg`); a `None`
|
|
# cfg leaves the base preset's quant params untouched, so disabling then
|
|
# re-enabling restores the original (FP8/NVFP4/...) attributes. Putting `enable`
|
|
# under `cfg` is rejected by the QuantizerAttributeConfig validator.
|
|
rules: list[dict[str, Any]] = [
|
|
{"quantizer_name": "*weight_quantizer", "enable": False},
|
|
{"quantizer_name": "*input_quantizer", "enable": False},
|
|
{"quantizer_name": f"*{block_module}.*weight_quantizer", "enable": True},
|
|
{"quantizer_name": f"*{block_module}.*input_quantizer", "enable": True},
|
|
]
|
|
for idx in excluded:
|
|
rules.append(
|
|
{"quantizer_name": f"*{block_module}.{idx}.*weight_quantizer", "enable": False}
|
|
)
|
|
rules.append({"quantizer_name": f"*{block_module}.{idx}.*input_quantizer", "enable": False})
|
|
return rules
|