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

273 lines
12 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
import warnings
from collections.abc import Iterator
from typing import Any
import torch
from diffusers import DiffusionPipeline, LTXLatentUpsamplePipeline
from models_utils import MODEL_DEFAULTS, MODEL_PIPELINE, MODEL_REGISTRY, ModelType
from quantize_config import ModelConfig
import modelopt.torch.quantization as mtq
class PipelineManager:
"""Manages diffusion pipeline creation and configuration."""
def __init__(self, config: ModelConfig, logger: logging.Logger):
"""
Initialize pipeline manager.
Args:
config: Model configuration
logger: Logger instance
"""
self.config = config
self.logger = logger
self.pipe: Any | None = None
self.pipe_upsample: LTXLatentUpsamplePipeline | None = None # For LTX-Video upsampling
self._transformer: torch.nn.Module | None = None
self._video_decoder: torch.nn.Module | None = None
@staticmethod
def create_pipeline_from(
model_type: ModelType,
torch_dtype: torch.dtype | dict[str, str | torch.dtype] = torch.bfloat16,
override_model_path: str | None = None,
) -> DiffusionPipeline:
"""
Create and return an appropriate pipeline based on configuration.
Returns:
Configured diffusion pipeline
Raises:
ValueError: If model type is unsupported
"""
pipeline_cls = MODEL_PIPELINE[model_type]
if pipeline_cls is None:
raise ValueError(
f"Model type {model_type.value} is not supported by the installed diffusers "
"version; upgrade diffusers to a release that provides its pipeline."
)
model_id = (
MODEL_REGISTRY[model_type] if override_model_path is None else override_model_path
)
pipe = pipeline_cls.from_pretrained(
model_id,
torch_dtype=torch_dtype,
use_safetensors=True,
**MODEL_DEFAULTS[model_type].get("from_pretrained_extra_args", {}),
)
pipe.set_progress_bar_config(disable=True)
return pipe
def create_pipeline(self) -> Any:
"""
Create and return an appropriate pipeline based on configuration.
Returns:
Configured diffusion pipeline
Raises:
ValueError: If model type is unsupported
"""
self.logger.info(f"Creating pipeline for {self.config.model_type.value}")
self.logger.info(f"Model path: {self.config.model_path}")
self.logger.info(f"Data type: {self.config.model_dtype}")
try:
if self.config.model_type == ModelType.LTX2:
from modelopt.torch.quantization.plugins.diffusion import ltx2 as ltx2_plugin
ltx2_plugin.register_ltx2_quant_linear()
self.pipe = self._create_ltx2_pipeline()
self.logger.info("LTX-2 pipeline created successfully")
return self.pipe
pipeline_cls = MODEL_PIPELINE[self.config.model_type]
if pipeline_cls is None:
raise ValueError(
f"Model type {self.config.model_type.value} is not supported by the "
"installed diffusers version; upgrade diffusers to a release that "
"provides its pipeline."
)
self.pipe = pipeline_cls.from_pretrained(
self.config.model_path,
torch_dtype=self.config.model_dtype,
use_safetensors=True,
**MODEL_DEFAULTS[self.config.model_type].get("from_pretrained_extra_args", {}),
)
if self.config.model_type == ModelType.LTX_VIDEO_DEV:
# Optionally load the upsampler pipeline for LTX-Video
if not self.config.ltx_skip_upsampler:
self.logger.info("Loading LTX-Video upsampler pipeline...")
self.pipe_upsample = LTXLatentUpsamplePipeline.from_pretrained(
"Lightricks/ltxv-spatial-upscaler-0.9.7",
vae=self.pipe.vae,
torch_dtype=self.config.model_dtype,
)
self.pipe_upsample.set_progress_bar_config(disable=True)
else:
self.logger.info("Skipping upsampler pipeline for faster calibration")
self.pipe.set_progress_bar_config(disable=True)
self.logger.info("Pipeline created successfully")
return self.pipe
except Exception as e:
self.logger.error(f"Failed to create pipeline: {e}")
raise
def setup_device(self) -> None:
"""Configure pipeline device placement."""
if not self.pipe:
raise RuntimeError("Pipeline not created. Call create_pipeline() first.")
if self.config.model_type == ModelType.LTX2:
self.logger.info("Skipping device setup for LTX-2 pipeline (handled internally)")
return
if self.config.cpu_offloading:
self.logger.info("Enabling CPU offloading for memory efficiency")
self.pipe.enable_model_cpu_offload()
if self.pipe_upsample:
self.pipe_upsample.enable_model_cpu_offload()
else:
self.logger.info("Moving pipeline to CUDA")
self.pipe.to("cuda")
if self.pipe_upsample:
self.logger.info("Moving upsampler pipeline to CUDA")
self.pipe_upsample.to("cuda")
# Enable VAE tiling for LTX-Video to save memory
if self.config.model_type == ModelType.LTX_VIDEO_DEV:
if hasattr(self.pipe, "vae") and hasattr(self.pipe.vae, "enable_tiling"):
self.logger.info("Enabling VAE tiling for LTX-Video")
self.pipe.vae.enable_tiling()
def iter_backbones(self) -> Iterator[tuple[str, torch.nn.Module]]:
"""
Yield (backbone_name, module) pairs.
"""
if not self.pipe:
raise RuntimeError("Pipeline not created. Call create_pipeline() first.")
names = list(self.config.backbone)
if not names:
raise RuntimeError("No backbone names provided.")
if self.config.model_type == ModelType.LTX2:
for name in names:
if name == "video_decoder":
self._ensure_ltx2_video_decoder_cached()
yield name, self._video_decoder
elif name == "transformer":
self._ensure_ltx2_transformer_cached()
yield name, self._transformer
else:
raise ValueError(
f"Unsupported LTX-2 backbone name '{name}'. "
"Expected 'transformer' or 'video_decoder'."
)
return
for name in names:
module = getattr(self.pipe, name, None)
if module is None:
raise RuntimeError(f"Pipeline missing backbone module '{name}'.")
yield name, module
def _ensure_ltx2_transformer_cached(self) -> None:
if not self.pipe:
raise RuntimeError("Pipeline not created. Call create_pipeline() first.")
if self._transformer is None:
transformer = self.pipe.stage_1_model_ledger.transformer()
self.pipe.stage_1_model_ledger.transformer = lambda: transformer
self._transformer = transformer
def _ensure_ltx2_video_decoder_cached(self) -> None:
if not self.pipe:
raise RuntimeError("Pipeline not created. Call create_pipeline() first.")
if self._video_decoder is None:
video_decoder = self.pipe.stage_1_model_ledger.video_decoder()
# Cache it so subsequent calls return the same (quantized) instance
self.pipe.stage_1_model_ledger.video_decoder = lambda: video_decoder
self.pipe.stage_2_model_ledger.video_decoder = lambda: video_decoder
self._video_decoder = video_decoder
def _create_ltx2_pipeline(self) -> Any:
params = dict(self.config.extra_params)
checkpoint_path = params.pop("checkpoint_path", None)
distilled_lora_path = params.pop("distilled_lora_path", None)
distilled_lora_strength = params.pop("distilled_lora_strength", 0.8)
spatial_upsampler_path = params.pop("spatial_upsampler_path", None)
gemma_root = params.pop("gemma_root", None)
fp8_quantization = params.pop("fp8_quantization", None) or params.pop(
"fp8transformer", False
)
params.pop("merged_base_safetensor_path", None)
params.pop("enable_swizzle_layout", None)
params.pop("padding_strategy", None)
params.pop("enable_layerwise_quant_metadata", None)
if not checkpoint_path:
raise ValueError("Missing required extra_param: checkpoint_path.")
if not distilled_lora_path:
raise ValueError("Missing required extra_param: distilled_lora_path.")
if not spatial_upsampler_path:
raise ValueError("Missing required extra_param: spatial_upsampler_path.")
if not gemma_root:
raise ValueError("Missing required extra_param: gemma_root.")
warnings.warn(
"LTX-2 packages (ltx-core, ltx-pipelines, ltx-trainer) are provided by Lightricks and are NOT "
"covered by the Apache 2.0 license governing NVIDIA Model Optimizer. You MUST comply "
"with the LTX Community License Agreement when installing and using LTX-2 with NVIDIA "
"Model Optimizer. Any derivative models or fine-tuned weights from LTX-2 remain "
"subject to the LTX Community License Agreement, not Apache 2.0. "
"See: https://github.com/Lightricks/LTX-2/blob/main/LICENSE",
UserWarning,
stacklevel=2,
)
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
from ltx_core.quantization import QuantizationPolicy
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
distilled_lora = [
LoraPathStrengthAndSDOps(
str(distilled_lora_path),
float(distilled_lora_strength),
LTXV_LORA_COMFY_RENAMING_MAP,
)
]
pipeline_kwargs = {
"checkpoint_path": str(checkpoint_path),
"distilled_lora": distilled_lora,
"spatial_upsampler_path": str(spatial_upsampler_path),
"gemma_root": str(gemma_root),
"loras": [],
"quantization": QuantizationPolicy.fp8_cast() if fp8_quantization else None,
}
pipeline_kwargs.update(params)
return TI2VidTwoStagesPipeline(**pipeline_kwargs)
def print_quant_summary(self):
for name, backbone in self.iter_backbones():
self.logger.info(f"{name} quantization info:")
mtq.print_quant_summary(backbone)