mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user