[1/3] Diffusion ckpt export for NVFP4 & FP8 (#781)

## What does this PR do?

**Type of change:** New feature <!-- Use one of the following: Bug fix,
new feature, new example, new tests, documentation. -->

**Overview:** 

This PR adds support for exporting quantized diffusers models (DiT,
Flux, SD3, UNet, etc.) to HuggingFace checkpoint format, enabling
deployment to inference frameworks like SGLang, vLLM, and TensorRT-LLM.

**Changes**

New file: `diffusers_utils.py`
- Dummy input generation for various diffusion models
- Pipeline component extraction helpers
- QKV projection detection and grouping
- `hide_quantizers_from_state_dict()` context manager for clean saves

Refactored: `unified_export_hf.py`
- New `_fuse_qkv_linears_diffusion()` for QKV amax fusion
- `_export_diffusers_checkpoint()` to export full pipelines (models +
tokenizers + schedulers etc.)

Plans

- [x] [1/3] Add the basic functionalities to support limited image
models with NVFP4 + FP8, with some refactoring on the previous LLM code
and the diffusers example. PIC: @jingyu-ml
- [ ] [2/3] Add support to more video gen modelsPIC: @jingyu-ml 
- [ ] [3/3] Add test cases, refactor on the doc, and all related README.
PIC: @jingyu-ml

## Usage
<!-- You can potentially add a usage example below. -->
```
mtq.quantize(pipe, quant_config, forward_call)
export_hf_checkpoint(pipe, export_dir=hf_ckpt_dir)
```

## Testing
<!-- Mention how have you tested your change if applicable. -->

## Before your PR is "*Ready for review*"
<!-- If you haven't finished some of the above items you can still open
`Draft` PR. -->

- **Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)**
and your commits are signed.
- **Is this change backward compatible?**: Yes <!--- If No, explain why.
-->
- **Did you write any new necessary tests?**:No
- **Did you add or update any necessary documentation?**:No
- **Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**:No
<!--- Only for new features, API changes, critical bug fixes or bw
breaking changes. -->

## Additional Information
<!-- E.g. related issue. -->


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

## New Features
* Added HuggingFace checkpoint export support for quantized diffusion
models with configurable output directory
* Introduced new `--hf-ckpt-dir` CLI argument for specifying checkpoint
export destination
* Extended export functionality to support selective component exports
from diffusion pipelines
* Enhanced quantized model export with improved component handling and
multi-stage checkpoint generation

<sub>✏️ Tip: You can customize this high-level summary in your review
settings.</sub>
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Jingyu Xin <jingyux@nvidia.com>
This commit is contained in:
jingyu-ml
2026-01-21 23:11:09 +00:00
committed by GitHub
parent 563a1e09c6
commit 668b8a19e8
8 changed files with 1080 additions and 202 deletions
@@ -0,0 +1,194 @@
# 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.
from collections.abc import Callable
from enum import Enum
from typing import Any
from diffusers import (
DiffusionPipeline,
FluxPipeline,
LTXConditionPipeline,
StableDiffusion3Pipeline,
WanPipeline,
)
from utils import (
filter_func_default,
filter_func_flux_dev,
filter_func_ltx_video,
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"
LTX_VIDEO_DEV = "ltx-video-dev"
WAN22_T2V_14b = "wan2.2-t2v-14b"
WAN22_T2V_5b = "wan2.2-t2v-5b"
def get_model_filter_func(model_type: ModelType) -> Callable[[str], bool]:
"""
Get the appropriate filter function for a given model type.
Args:
model_type: The model type enum
Returns:
A filter function appropriate for the model type
"""
filter_func_map = {
ModelType.FLUX_DEV: filter_func_flux_dev,
ModelType.FLUX_SCHNELL: filter_func_default,
ModelType.SDXL_BASE: filter_func_default,
ModelType.SDXL_TURBO: filter_func_default,
ModelType.SD3_MEDIUM: filter_func_default,
ModelType.SD35_MEDIUM: filter_func_default,
ModelType.LTX_VIDEO_DEV: filter_func_ltx_video,
ModelType.WAN22_T2V_14b: filter_func_wan_video,
ModelType.WAN22_T2V_5b: filter_func_wan_video,
}
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.LTX_VIDEO_DEV: "Lightricks/LTX-Video-0.9.7-dev",
ModelType.WAN22_T2V_14b: "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
ModelType.WAN22_T2V_5b: "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
}
MODEL_PIPELINE: dict[ModelType, type[DiffusionPipeline]] = {
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.LTX_VIDEO_DEV: LTXConditionPipeline,
ModelType.WAN22_T2V_14b: WanPipeline,
ModelType.WAN22_T2V_5b: WanPipeline,
}
# 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.LTX_VIDEO_DEV: {
"backbone": "transformer",
"dataset": _SD_PROMPTS_DATASET,
"inference_extra_args": {
"height": 512,
"width": 704,
"num_frames": 121,
"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
),
},
},
}
+20 -196
View File
@@ -17,7 +17,6 @@ import argparse
import logging
import sys
import time as time
from collections.abc import Callable
from dataclasses import dataclass, field
from enum import Enum
from pathlib import Path
@@ -45,43 +44,23 @@ if __name__ == "__main__":
torch.nn.RMSNorm = DiffuserRMSNorm
torch.nn.modules.normalization.RMSNorm = DiffuserRMSNorm
from diffusers import (
DiffusionPipeline,
FluxPipeline,
LTXConditionPipeline,
LTXLatentUpsamplePipeline,
StableDiffusion3Pipeline,
WanPipeline,
from diffusers import DiffusionPipeline, LTXLatentUpsamplePipeline
from models_utils import (
MODEL_DEFAULTS,
MODEL_PIPELINE,
MODEL_REGISTRY,
ModelType,
get_model_filter_func,
)
from onnx_utils.export import generate_fp8_scales, modelopt_export_sd
from tqdm import tqdm
from utils import (
check_conv_and_mha,
check_lora,
filter_func_default,
filter_func_ltx_video,
filter_func_wan_video,
load_calib_prompts,
)
from utils import check_conv_and_mha, check_lora, load_calib_prompts
import modelopt.torch.opt as mto
import modelopt.torch.quantization as mtq
from modelopt.torch.export import export_hf_checkpoint
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"
LTX_VIDEO_DEV = "ltx-video-dev"
WAN22_T2V = "wan2.2-t2v-14b"
class DataType(str, Enum):
"""Supported data types for model loading."""
@@ -127,155 +106,6 @@ class CollectMethod(str, Enum):
DEFAULT = "default"
def get_model_filter_func(model_type: ModelType) -> Callable[[str], bool]:
"""
Get the appropriate filter function for a given model type.
Args:
model_type: The model type enum
Returns:
A filter function appropriate for the model type
"""
filter_func_map = {
ModelType.FLUX_DEV: filter_func_default,
ModelType.FLUX_SCHNELL: filter_func_default,
ModelType.SDXL_BASE: filter_func_default,
ModelType.SDXL_TURBO: filter_func_default,
ModelType.SD3_MEDIUM: filter_func_default,
ModelType.SD35_MEDIUM: filter_func_default,
ModelType.LTX_VIDEO_DEV: filter_func_ltx_video,
ModelType.WAN22_T2V: filter_func_wan_video,
}
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.LTX_VIDEO_DEV: "Lightricks/LTX-Video-0.9.7-dev",
ModelType.WAN22_T2V: "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
}
MODEL_PIPELINE: dict[ModelType, type[DiffusionPipeline]] = {
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.LTX_VIDEO_DEV: LTXConditionPipeline,
ModelType.WAN22_T2V: WanPipeline,
}
# Model-specific default arguments for calibration
MODEL_DEFAULTS: dict[ModelType, dict[str, Any]] = {
ModelType.SDXL_BASE: {
"backbone": "unet",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
},
ModelType.SDXL_TURBO: {
"backbone": "unet",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
},
ModelType.SD3_MEDIUM: {
"backbone": "transformer",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
},
ModelType.SD35_MEDIUM: {
"backbone": "transformer",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
},
ModelType.FLUX_DEV: {
"backbone": "transformer",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
"inference_extra_args": {
"height": 1024,
"width": 1024,
"guidance_scale": 3.5,
"max_sequence_length": 512,
},
},
ModelType.FLUX_SCHNELL: {
"backbone": "transformer",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
"inference_extra_args": {
"height": 1024,
"width": 1024,
"guidance_scale": 3.5,
"max_sequence_length": 512,
},
},
ModelType.LTX_VIDEO_DEV: {
"backbone": "transformer",
"dataset": {
"name": "Gustavosta/Stable-Diffusion-Prompts",
"split": "train",
"column": "Prompt",
},
"inference_extra_args": {
"height": 512,
"width": 704,
"num_frames": 121,
"negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted",
},
},
ModelType.WAN22_T2V: {
"backbone": "transformer",
"dataset": {"name": "nkp37/OpenVid-1M", "split": "train", "column": "caption"},
"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"
),
},
},
}
@dataclass
class QuantizationConfig:
"""Configuration for model quantization."""
@@ -590,8 +420,8 @@ class Calibrator:
if self.model_type == ModelType.LTX_VIDEO_DEV:
# Special handling for LTX-Video
self._run_ltx_video_calibration(prompt_batch, extra_args)
elif self.model_type == ModelType.WAN22_T2V:
# Special handling for LTX-Video
elif self.model_type in [ModelType.WAN22_T2V_14b, ModelType.WAN22_T2V_5b]:
# Special handling for WAN video models
self._run_wan_video_calibration(prompt_batch, extra_args)
else:
common_args = {
@@ -606,23 +436,17 @@ class Calibrator:
def _run_wan_video_calibration(
self, prompt_batch: list[str], extra_args: dict[str, Any]
) -> None:
negative_prompt = extra_args["negative_prompt"]
height = extra_args["height"]
width = extra_args["width"]
num_frames = extra_args["num_frames"]
guidance_scale = extra_args["guidance_scale"]
guidance_scale_2 = extra_args["guidance_scale_2"]
kwargs = {}
kwargs["negative_prompt"] = extra_args["negative_prompt"]
kwargs["height"] = extra_args["height"]
kwargs["width"] = extra_args["width"]
kwargs["num_frames"] = extra_args["num_frames"]
kwargs["guidance_scale"] = extra_args["guidance_scale"]
if "guidance_scale_2" in extra_args:
kwargs["guidance_scale_2"] = extra_args["guidance_scale_2"]
kwargs["num_inference_steps"] = self.config.n_steps
self.pipe(
prompt=prompt_batch,
negative_prompt=negative_prompt,
height=height,
width=width,
num_frames=num_frames,
guidance_scale=guidance_scale,
guidance_scale_2=guidance_scale_2,
num_inference_steps=self.config.n_steps,
).frames # type: ignore[misc]
self.pipe(prompt=prompt_batch, **kwargs).frames # type: ignore[misc]
def _run_ltx_video_calibration(
self, prompt_batch: list[str], extra_args: dict[str, Any]
+7 -1
View File
@@ -73,9 +73,15 @@ def filter_func_ltx_video(name: str) -> bool:
return pattern.match(name) is not None
def filter_func_flux_dev(name: str) -> bool:
"""Filter function specifically for Flux-dev models."""
pattern = re.compile(r"(proj_out.*|.*(time_text_embed|context_embedder|x_embedder|norm_out).*)")
return pattern.match(name) is not None
def filter_func_wan_video(name: str) -> bool:
"""Filter function specifically for LTX-Video models."""
pattern = re.compile(r".*(patch_embedding|condition_embedder).*")
pattern = re.compile(r".*(patch_embedding|condition_embedder|proj_out).*")
return pattern.match(name) is not None