Files
jingyu-mlandClaude Opus 4.8 6b4ad85849 Qwen-Image diffusers PTQ: FP8 / NVFP4 / NVFP4-SVDQuant HF checkpoints (#1706)
### 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>
2026-07-07 17:21:31 -07:00

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