Files
Ajinkya RasaneandCodex cf1f48fa0f [5565357] Fix SDXL NVFP4 export and performance (#2336)
### 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>
2026-09-18 17:22:28 +00:00

193 lines
7.2 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 os
import re
from pathlib import Path
import torch
import torch.nn.functional as F
from datasets import load_dataset
from diffusers.models.attention_processor import Attention, AttnProcessor
from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear
from diffusers.utils import load_image
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.plugins.diffusion.diffusers import AttentionModuleMixin
USE_PEFT = True
try:
from peft.tuners.lora.layer import Conv2d as PEFTLoRAConv2d
from peft.tuners.lora.layer import Linear as PEFTLoRALinear
except ModuleNotFoundError:
USE_PEFT = False
# Model-specific filter functions for quantization
def filter_func_default(name: str) -> bool:
"""Default filter function for general models."""
pattern = re.compile(
r".*(time_emb_proj|time_embedding|conv_in|conv_out|conv_shortcut|add_embedding|pos_embed|time_text_embed|context_embedder|norm_out|x_embedder).*"
)
return pattern.match(name) is not None
def check_conv_and_mha(backbone, if_fp4, quantize_mha):
for name, module in backbone.named_modules():
if isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d)) and if_fp4:
module.weight_quantizer.disable()
module.input_quantizer.disable()
print(f"Disabled NVFP4 Conv layer quantization for layer {name}")
elif isinstance(module, (Attention, AttentionModuleMixin)):
head_size = int(module.inner_dim / module.heads)
if not quantize_mha or head_size % 16 != 0:
for attr in (
"q_bmm_quantizer",
"k_bmm_quantizer",
"v_bmm_quantizer",
"softmax_quantizer",
"bmm2_output_quantizer",
):
if hasattr(module, attr):
getattr(module, attr).disable()
print(f"Disabled Attention layer quantization for layer {name}")
def filter_func_ltx_video(name: str) -> bool:
"""Filter function specifically for LTX-Video models."""
pattern = re.compile(
r".*(proj_in|time_embed|caption_projection|proj_out|patchify_proj|adaln_single|transformer_blocks\.(0|1|2|45|46|47)\.).*"
)
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|time_guidance_embed|stream_modulation).*)"
)
return pattern.match(name) is not None
def filter_func_ltx2_vae(name: str) -> bool:
"""Filter for LTX-2 VAE: keeps only conv1/conv2 in up_blocks resnets."""
keep = re.compile(r".*up_blocks\.\d+\.resnets\.\d+\.conv[12](?:\.|$)")
return not keep.match(name)
def filter_func_wan_vae(name: str) -> bool:
"""Filter for Wan 2.2 VAE: keeps only conv1/conv2 in resnet blocks."""
keep = re.compile(
r".*(down_blocks\.\d+\.(?:resnets\.\d+\.)?conv[12]"
r"|mid_block\.resnets\.\d+\.conv[12]"
r"|up_blocks\.\d+\.resnets\.\d+\.conv[12])(?:\.|$)"
)
return not keep.match(name)
def filter_func_wan_video(name: str) -> bool:
"""Filter function specifically for WAN-Video models."""
pattern = re.compile(
r".*(patch_embedding|condition_embedder|proj_out|blocks\.(0|1|2|37|38|39)\.).*"
)
return pattern.match(name) is not None
# Qwen-Image's transformer has 60 ``transformer_blocks``. The recipe quantizes
# only those blocks while keeping the first two and last two -- and everything
# outside ``transformer_blocks`` -- in original precision. The model-agnostic,
# config-driven form of this recipe (deriving the block count from the model)
# lives in quantize.py; this name-only filter covers the plain FP8/NVFP4 path
# for the full 60-block Qwen-Image transformer.
QWEN_IMAGE_NUM_TRANSFORMER_BLOCKS = 60
_QWEN_IMAGE_BLOCK_RE = re.compile(r"(?:^|\.)transformer_blocks\.(\d+)(?:\.|$)")
def filter_func_qwen_image(name: str) -> bool:
"""Filter function specifically for Qwen-Image models.
Returns ``True`` for modules to keep in original precision (quantization
disabled): everything outside ``transformer_blocks``, plus the first two and
last two transformer blocks.
"""
match = _QWEN_IMAGE_BLOCK_RE.search(name)
if match is None:
return True
block_idx = int(match.group(1))
return block_idx < 2 or block_idx >= QWEN_IMAGE_NUM_TRANSFORMER_BLOCKS - 2
def load_calib_prompts(
batch_size,
calib_data_path: str | Path = "Gustavosta/Stable-Diffusion-Prompts",
split="train",
column="Prompt",
) -> list[list[str]]:
prompt_list: list[str] = []
if isinstance(calib_data_path, Path):
with open(calib_data_path) as f:
prompt_list = f.readlines()
else:
dataset = load_dataset(calib_data_path)
prompt_list = list(dataset[split][column])
return [prompt_list[i : i + batch_size] for i in range(0, len(prompt_list), batch_size)]
def load_calib_images(folder_path):
images = []
for filename in os.listdir(folder_path):
img_path = os.path.join(folder_path, filename)
if os.path.isfile(img_path):
image = load_image(img_path)
if image is not None:
images.append(image)
return images
def set_fmha(unet):
for name, module in unet.named_modules():
if isinstance(module, Attention):
module.set_processor(AttnProcessor())
def check_lora(unet):
for name, module in unet.named_modules():
if isinstance(module, (LoRACompatibleConv, LoRACompatibleLinear)):
assert module.lora_layer is None, (
f"To quantize {name}, LoRA layer should be fused/merged. Please"
" fuse the LoRA layer before quantization."
)
elif USE_PEFT and isinstance(module, (PEFTLoRAConv2d, PEFTLoRALinear)):
assert module.merged, (
f"To quantize {name}, LoRA layer should be fused/merged. Please"
" fuse the LoRA layer before quantization."
)
def fp8_mha_disable(backbone, quantized_mha_output: bool = True):
def mha_filter_func(name):
pattern = re.compile(
r".*(q_bmm_quantizer|k_bmm_quantizer|v_bmm_quantizer|softmax_quantizer).*"
if quantized_mha_output
else r".*(q_bmm_quantizer|k_bmm_quantizer|v_bmm_quantizer|softmax_quantizer|bmm2_output_quantizer).*"
)
return pattern.match(name) is not None
if hasattr(F, "scaled_dot_product_attention"):
mtq.disable_quantizer(backbone, mha_filter_func)