[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>
This commit is contained in:
Ajinkya Rasane
2026-09-18 17:22:28 +00:00
committed by GitHub
co-authored by Codex
parent 17ef5b6c39
commit cf1f48fa0f
16 changed files with 428 additions and 105 deletions
@@ -117,6 +117,17 @@ class DiffuserModel(NamedTuple):
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",
@@ -128,6 +139,17 @@ class DiffuserModel(NamedTuple):
),
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,
@@ -140,7 +162,9 @@ class DiffuserModel(NamedTuple):
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",
],
)
@@ -17,7 +17,6 @@
import pytest
import torch
import torch.nn as nn
from _test_utils.torch.quantization.onnx_export import TEST_MODELS, onnx_export_tester
@@ -40,7 +39,4 @@ def test_onnx_export_cuda(model_cls, num_bits, per_channel_quantization, constan
torch.manual_seed(0)
model = model_cls()
for _, module in model.named_modules():
if isinstance(module, nn.Conv2d) and num_bits == (4, 3):
pytest.skip("Conv2d with FP8 quantization is not supported yet")
onnx_export_tester(model, "cuda", num_bits, per_channel_quantization, constant_folding, dtype)
+234
View File
@@ -0,0 +1,234 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 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 importlib.util
import logging
import sys
from pathlib import Path
from unittest.mock import Mock
import pytest
import torch
from torch import nn
pytest.importorskip("onnx")
pytest.importorskip("onnx_graphsurgeon")
pytest.importorskip("diffusers")
import modelopt.torch.quantization as mtq
from examples.diffusers.quantization.onnx_utils import export as diffusion_export
from modelopt.torch.quantization.config import QuantizerAttributeConfig
from modelopt.torch.quantization.nn import TensorQuantizer
from modelopt.torch.quantization.plugins.diffusion import diffusers as diffusers_plugin
_QUANTIZATION_EXAMPLE = (
Path(__file__).resolve().parents[3] / "examples" / "diffusers" / "quantization"
)
_LOCAL_IMPORT_NAMES = (
"calib.plugin_calib",
"calib",
"calibration",
"config",
"models_utils",
"pipeline_manager",
"quantize_config",
"utils",
)
def _load_quantize_example():
spec = importlib.util.spec_from_file_location(
"diffusers_quantize_example", _QUANTIZATION_EXAMPLE / "quantize.py"
)
assert spec is not None and spec.loader is not None
original_modules = {
name: sys.modules.pop(name) for name in _LOCAL_IMPORT_NAMES if name in sys.modules
}
sys.path.insert(0, str(_QUANTIZATION_EXAMPLE))
try:
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
finally:
sys.path.pop(0)
for name in _LOCAL_IMPORT_NAMES:
sys.modules.pop(name, None)
sys.modules.update(original_modules)
return module
_quantize = _load_quantize_example()
ModelType = _quantize.ModelType
ModelConfig = _quantize.ModelConfig
QuantFormat = _quantize.QuantFormat
QuantizationConfig = _quantize.QuantizationConfig
Quantizer = _quantize.Quantizer
_infer_restored_quantization_format = _quantize._infer_restored_quantization_format
class _RecipeBackbone(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(16, 16, bias=False)
self.attn = nn.Module()
self.attn.to_q = nn.Linear(16, 16, bias=False)
self.attn.to_k = nn.Linear(16, 16, bias=False)
self.attn.to_v = nn.Linear(16, 16, bias=False)
self.conv = nn.Conv2d(4, 4, kernel_size=1, bias=False)
def _quantizer(*, num_bits, enabled=True, block_sizes=None):
quantizer = TensorQuantizer(
QuantizerAttributeConfig(num_bits=num_bits, axis=None, block_sizes=block_sizes)
)
quantizer.amax = torch.tensor(448.0)
if not enabled:
quantizer.disable()
return quantizer
_FP8_QUANTIZER_CONFIG = {"num_bits": (4, 3)}
_NVFP4_QUANTIZER_CONFIG = {
"num_bits": (2, 1),
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
}
@pytest.mark.parametrize("model_type", [ModelType.SDXL_BASE, ModelType.SDXL_TURBO])
def test_sdxl_fp4_recipe(model_type):
model = _RecipeBackbone()
config = Quantizer(
QuantizationConfig(format=QuantFormat.FP4),
ModelConfig(model_type=model_type),
logging.getLogger(__name__),
).get_quant_config(n_steps=1, backbone=model)
mtq.replace_quant_module(model)
mtq.set_quantizer_by_cfg(model, config["quant_cfg"])
for quantizer in (model.linear.input_quantizer, model.linear.weight_quantizer):
assert quantizer.is_enabled
assert quantizer.is_nvfp4_dynamic
assert quantizer.block_sizes[-1] == 16
for projection in (model.attn.to_q, model.attn.to_k, model.attn.to_v):
assert not projection.input_quantizer.is_enabled
assert not projection.weight_quantizer.is_enabled
for quantizer in (model.conv.input_quantizer, model.conv.weight_quantizer):
assert quantizer.is_enabled
assert quantizer.is_fp8
@pytest.mark.parametrize(
("format_config", "mha_config", "expected_format", "disable_fp8_mha"),
[
pytest.param(
_NVFP4_QUANTIZER_CONFIG,
_FP8_QUANTIZER_CONFIG,
QuantFormat.FP4,
False,
id="mixed-fp4",
),
pytest.param(
_FP8_QUANTIZER_CONFIG,
_FP8_QUANTIZER_CONFIG,
QuantFormat.FP8,
False,
id="fp8",
),
pytest.param(
{"num_bits": 8},
{**_FP8_QUANTIZER_CONFIG, "enabled": False},
QuantFormat.INT8,
True,
id="int8-disabled-fp8",
),
pytest.param(
{"num_bits": 8},
{"num_bits": 8},
QuantFormat.INT8,
True,
id="int8-mha",
),
],
)
def test_restored_quantizer_state_drives_format_and_fp8_mha(
monkeypatch, format_config, mha_config, expected_format, disable_fp8_mha
):
backbone = nn.Module()
backbone.quantizer = _quantizer(**format_config)
backbone.attention = nn.Module()
for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"):
setattr(backbone.attention, name, _quantizer(**mha_config))
backbone.attention.bmm2_output_quantizer = lambda output: output
fp8_sdpa = Mock(return_value=torch.empty(0))
monkeypatch.setattr(diffusers_plugin.FP8SDPA, "apply", fp8_sdpa)
monkeypatch.setattr(torch.onnx, "is_in_onnx_export", lambda: True)
assert _infer_restored_quantization_format([("transformer", backbone)]) == expected_format
diffusers_plugin._quantized_sdpa(backbone.attention, *(torch.empty(1) for _ in range(3)))
assert fp8_sdpa.call_args.args[-1] is disable_fp8_mha
def test_restore_infers_checkpoint_format_for_export(monkeypatch, tmp_path):
backbone = nn.Module()
backbone.quantizer = _quantizer(**_NVFP4_QUANTIZER_CONFIG)
pipeline_manager = Mock()
pipeline_manager.create_pipeline.return_value = object()
pipeline_manager.iter_backbones.return_value = [("transformer", backbone)]
export_manager = Mock()
monkeypatch.setattr(_quantize, "PipelineManager", lambda *args: pipeline_manager)
monkeypatch.setattr(_quantize, "ExportManager", lambda *args: export_manager)
monkeypatch.setattr(
sys,
"argv",
[
"quantize.py",
"--model",
"flux-schnell",
"--restore-from",
str(tmp_path),
"--onnx-dir",
str(tmp_path / "onnx"),
],
)
_quantize.main()
export_manager.restore_checkpoint.assert_called_once_with()
assert export_manager.export_onnx.call_args.args[-1] == QuantFormat.FP4
export_manager.export_hf_ckpt.assert_called_once()
def test_flux_fp8_export_saves_converted_rope_graph(monkeypatch, tmp_path):
original_model = Mock()
converted_model = Mock()
monkeypatch.setattr(
diffusion_export,
"generate_dummy_kwargs_and_dynamic_axes_and_shapes",
lambda *args: ({}, {}, None),
)
monkeypatch.setattr(diffusion_export, "onnx_export", lambda *args, **kwargs: None)
monkeypatch.setattr(diffusion_export.onnx, "load", lambda *args, **kwargs: original_model)
convert_rope_weight_type = Mock(return_value=converted_model)
monkeypatch.setattr(diffusion_export, "flux_convert_rope_weight_type", convert_rope_weight_type)
save_onnx = Mock()
monkeypatch.setattr(diffusion_export, "save_onnx", save_onnx)
diffusion_export.modelopt_export_sd(nn.Module(), tmp_path, "flux-dev", "fp8")
convert_rope_weight_type.assert_called_once_with(original_model)
save_onnx.assert_called_once_with(converted_model, tmp_path / "model.onnx")
@@ -42,6 +42,7 @@ from modelopt.onnx.quantization.qdq_utils import (
replace_zero_scale_with_smallest_nonzero,
)
from modelopt.onnx.quantization.quant_utils import pack_float32_to_4bit_cpp_based
from modelopt.onnx.utils import get_opset_version
def create_test_model_with_int4_dq_reshape_transpose_matmul(constant_scale: bool = False):
@@ -343,8 +344,7 @@ def create_test_model_with_nvfp4_qdq(with_transpose: bool = False):
value_info=value_info,
)
model = helper.make_model(graph)
return model
return helper.make_model(graph, opset_imports=[helper.make_opsetid("", 20)])
class TestQuantizeWeightsToInt4:
@@ -666,6 +666,8 @@ class TestFP4QDQTo2DQ:
# Run FP4QDQ to 2DQ conversion
converted_model = NVFP4QuantExporter.process_model(model)
assert get_opset_version(converted_model) == 23
# Verify TRT_FP4QDQ node is removed
fp4qdq_nodes = [node for node in converted_model.graph.node if node.op_type == "TRT_FP4QDQ"]
assert len(fp4qdq_nodes) == 0
@@ -39,6 +39,15 @@ from modelopt.torch.quantization.qtensor import NVFP4QTensor
from modelopt.torch.quantization.utils import is_quantized_linear
def _export_to_onnx(model, sample_input, **kwargs):
buffer = io.BytesIO()
if "enable_onnx_checker" in inspect.signature(torch.onnx.export).parameters:
kwargs["enable_onnx_checker"] = False
torch.onnx.export(model, sample_input, buffer, dynamo=False, **kwargs)
buffer.seek(0)
return onnx.load_model_from_string(buffer.read())
@pytest.mark.parametrize("model_cls", TEST_MODELS)
@pytest.mark.parametrize(
("num_bits", "per_channel_quantization", "constant_folding"),
@@ -59,6 +68,41 @@ def test_onnx_export_cpu(model_cls, num_bits, per_channel_quantization, constant
)
def test_fp8_conv_export_preserves_custom_qdq_and_kernel_shape():
model = torch.nn.Conv2d(3, 4, 3, bias=False).eval()
sample_input = torch.randn(1, 3, 8, 8)
model = mtq.quantize(
model,
mtq.FP8_DEFAULT_CFG,
forward_loop=lambda quantized_model: quantized_model(sample_input),
)
exported_model = _export_to_onnx(model, sample_input, opset_version=20)
producers = {output: node for node in exported_model.graph.node for output in node.output}
conv = next(node for node in exported_model.graph.node if node.op_type == "Conv")
for conv_input in conv.input[:2]:
dequantize = producers[conv_input]
quantize = producers[dequantize.input[0]]
assert dequantize.op_type == "TRT_FP8DequantizeLinear"
assert quantize.op_type == "TRT_FP8QuantizeLinear"
value_info = {value.name: value for value in exported_model.graph.value_info}
weight_dequantize = producers[conv.input[1]]
weight_quantize = producers[weight_dequantize.input[0]]
for value_name in (*weight_quantize.output, *weight_dequantize.output):
shape = [
dimension.dim_value for dimension in value_info[value_name].type.tensor_type.shape.dim
]
assert shape == [4, 3, 3, 3]
kernel_shape = next(
attribute for attribute in conv.attribute if attribute.name == "kernel_shape"
)
assert list(kernel_shape.ints) == [3, 3]
onnx.checker.check_model(exported_model)
def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch):
def forward_loop(model):
model(sample_input)
@@ -78,26 +122,14 @@ def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch):
module.input_quantizer.disable()
module.weight_quantizer._onnx_quantizer_type = "static"
buffer = io.BytesIO()
if "enable_onnx_checker" in inspect.signature(torch.onnx.export).parameters:
kwargs = {"enable_onnx_checker": False}
else:
kwargs = {}
torch.onnx.export(
exported_model = _export_to_onnx(
model,
sample_input,
buffer,
input_names=["input"],
output_names=["output"],
export_params=True,
opset_version=21,
dynamo=False,
**kwargs,
)
buffer.seek(0)
exported_model = onnx.load_model_from_string(buffer.read())
assert any(node.op_type == "TRT_FP4QDQ" for node in exported_model.graph.node)
converted_model = NVFP4QuantExporter.process_model(exported_model)