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

331 lines
9.7 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2023-2025 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 pathlib import Path
from typing import NamedTuple
import pytest
from _test_utils.examples.models import FLUX_SCHNELL_PATH, SD3_PATH, SDXL_PATH
from _test_utils.examples.run_command import run_example_command
from _test_utils.torch.misc import minimum_sm
class DiffuserModel(NamedTuple):
dtype: str
name: str
path: str
format_type: str
quant_algo: str
collect_method: str
def _run_cmd(self, script: str, *args: str) -> None:
cmd_args = [
"python",
script,
"--model",
self.name,
"--override-model-path",
self.path,
]
cmd_args.extend(args)
run_example_command(cmd_args, "diffusers/quantization")
def _format_args(self) -> list[str]:
return [
"--calib-size",
"4",
"--percentile",
"1.0",
"--alpha",
"0.8",
"--n-steps",
"2",
"--batch-size",
"2",
"--format",
self.format_type,
"--collect-method",
self.collect_method,
"--quant-algo",
self.quant_algo,
]
def quantize(self, tmp_path: Path) -> None:
self._run_cmd(
"quantize.py",
*self._format_args(),
"--trt-high-precision-dtype",
self.dtype,
"--quantized-torch-ckpt-save-path",
str(tmp_path / f"{self.name}_{self.format_type}.pt"),
"--onnx-dir",
str(tmp_path / f"{self.name}_{self.format_type}_onnx"),
)
def restore(self, tmp_path: Path) -> None:
self._run_cmd(
"quantize.py",
*self._format_args(),
"--trt-high-precision-dtype",
self.dtype,
"--restore-from",
str(tmp_path / f"{self.name}_{self.format_type}.pt"),
"--onnx-dir",
str(tmp_path / f"{self.name}_{self.format_type}_onnx"),
)
def inference(self, tmp_path: Path) -> None:
self._run_cmd(
"diffusion_trt.py",
"--onnx-load-path",
str(tmp_path / f"{self.name}_{self.format_type}_onnx/model.onnx"),
"--dq-only",
"--torch-autocast",
"--num-inference-steps",
"2",
)
@pytest.mark.parametrize(
"model",
[
DiffuserModel(
name="flux-schnell",
path=FLUX_SCHNELL_PATH,
dtype="BFloat16",
format_type="int8",
quant_algo="smoothquant",
collect_method="min-mean",
),
DiffuserModel(
name="sd3-medium",
path=SD3_PATH,
dtype="Half",
format_type="int8",
quant_algo="smoothquant",
collect_method="min-mean",
),
pytest.param(
DiffuserModel(
name="sd3-medium",
path=SD3_PATH,
dtype="Half",
format_type="fp8",
quant_algo="max",
collect_method="default",
),
marks=minimum_sm(89),
),
pytest.param(
DiffuserModel(
name="sdxl-1.0",
path=SDXL_PATH,
dtype="Half",
format_type="fp8",
quant_algo="max",
collect_method="default",
),
marks=minimum_sm(89),
),
pytest.param(
DiffuserModel(
name="sdxl-1.0",
path=SDXL_PATH,
dtype="Half",
format_type="fp4",
quant_algo="max",
collect_method="default",
),
marks=minimum_sm(100),
),
DiffuserModel(
name="sdxl-1.0",
path=SDXL_PATH,
dtype="Half",
format_type="int8",
quant_algo="smoothquant",
collect_method="min-mean",
),
],
ids=[
"flux_schnell_bf16_int8_smoothquant_3.0_min_mean",
"sd3_medium_fp16_int8_smoothquant_3.0_min_mean",
"sd3_medium_fp16_fp8_max_3.0_default",
"sdxl_1.0_fp16_fp8_max_3.0_default",
"sdxl_1.0_fp16_fp4_max_3.0_default",
"sdxl_1.0_fp16_int8_smoothquant_3.0_min_mean",
],
)
def test_diffusers_quantization(
model: DiffuserModel,
tmp_path: Path,
) -> None:
model.quantize(tmp_path)
model.restore(tmp_path)
model.inference(tmp_path)
class Wan22Model(NamedTuple):
model: str
backbone: str | None
format_type: str
quant_algo: str
collect_method: str
def _ckpt_path(self, tmp_path: Path) -> str:
stem = self.model.replace("wan2.2-t2v-", "")
parts = [stem, *([self.backbone] if self.backbone else []), self.format_type]
return str(tmp_path / f"wan22_{'_'.join(parts)}.pt")
def _common_args(self, tiny_wan22_path: str) -> list[str]:
cmd_args = [
"python",
"quantize.py",
"--model",
self.model,
"--override-model-path",
tiny_wan22_path,
"--format",
self.format_type,
"--quant-algo",
self.quant_algo,
"--collect-method",
self.collect_method,
"--model-dtype",
"BFloat16",
"--trt-high-precision-dtype",
"BFloat16",
"--calib-size",
"2",
"--batch-size",
"1",
"--n-steps",
"2",
# Tiny video dims — override MODEL_DEFAULTS for fast CI.
"--extra-param",
"height=16",
"--extra-param",
"width=16",
"--extra-param",
"num_frames=5",
]
if self.backbone is not None:
cmd_args.extend(["--backbone", self.backbone])
return cmd_args
def quantize(self, tiny_wan22_path: str, tmp_path: Path) -> None:
run_example_command(
[
*self._common_args(tiny_wan22_path),
"--quantized-torch-ckpt-save-path",
self._ckpt_path(tmp_path),
],
"diffusers/quantization",
)
def restore(self, tiny_wan22_path: str, tmp_path: Path) -> None:
run_example_command(
[*self._common_args(tiny_wan22_path), "--restore-from", self._ckpt_path(tmp_path)],
"diffusers/quantization",
)
# The VAE (``AutoencoderKLWan``) is shared between Wan 2.2 14B and 5B, so the
# Conv3D NVFP4 implicit-GEMM dispatch exercises the same kernel either way; we
# parametrize both ``--model`` values to also cover the ``quantize.py`` dispatch
# for each.
@pytest.mark.parametrize(
"wan_model",
[
Wan22Model("wan2.2-t2v-14b", None, "int8", "smoothquant", "min-mean"),
pytest.param(
Wan22Model("wan2.2-t2v-14b", None, "fp8", "max", "default"),
marks=minimum_sm(89),
),
pytest.param(
Wan22Model("wan2.2-t2v-14b", None, "fp4", "max", "default"),
marks=minimum_sm(89),
),
pytest.param(
Wan22Model("wan2.2-t2v-14b", "vae", "fp8", "max", "default"),
marks=minimum_sm(89),
),
pytest.param(
Wan22Model("wan2.2-t2v-14b", "vae", "fp4", "max", "default"),
marks=minimum_sm(89),
),
pytest.param(
Wan22Model("wan2.2-t2v-5b", "vae", "fp8", "max", "default"),
marks=minimum_sm(89),
),
pytest.param(
Wan22Model("wan2.2-t2v-5b", "vae", "fp4", "max", "default"),
marks=minimum_sm(89),
),
],
ids=[
"wan22_14b_transformer_int8_smoothquant",
"wan22_14b_transformer_fp8_max",
"wan22_14b_transformer_fp4_max",
"wan22_14b_vae_fp8_max",
"wan22_14b_vae_fp4_max",
"wan22_5b_vae_fp8_max",
"wan22_5b_vae_fp4_max",
],
)
def test_wan22_quantization(wan_model: Wan22Model, tiny_wan22_path: str, tmp_path: Path) -> None:
wan_model.quantize(tiny_wan22_path, tmp_path)
wan_model.restore(tiny_wan22_path, tmp_path)
@pytest.mark.parametrize(
("model_name", "model_path", "torch_compile"),
[
("flux-schnell", FLUX_SCHNELL_PATH, False),
("flux-schnell", FLUX_SCHNELL_PATH, True),
("sd3-medium", SD3_PATH, False),
("sd3-medium", SD3_PATH, True),
("sdxl-1.0", SDXL_PATH, False),
("sdxl-1.0", SDXL_PATH, True),
],
ids=[
"flux_schnell_torch",
"flux_schnell_torch_compile",
"sd3_medium_torch",
"sd3_medium_torch_compile",
"sdxl_1.0_torch",
"sdxl_1.0_torch_compile",
],
)
def test_diffusion_trt_torch(
model_name: str,
model_path: str,
torch_compile: bool,
) -> None:
cmd_args = [
"python",
"diffusion_trt.py",
"--model",
model_name,
"--override-model-path",
model_path,
"--torch",
"--num-inference-steps",
"2",
]
if torch_compile:
cmd_args.append("--torch-compile")
run_example_command(cmd_args, "diffusers/quantization")