[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
@@ -28,12 +28,12 @@ python quantize.py \
#### FLUX-Dev|SDXL|SDXL-Turbo|LTX-Video FP8/FP4 [Script](./quantize.py)
*In our example code, FP4 is only supported for Flux. However, you can modify our script to enable FP4 format support for your own model.*
FP4 ONNX export is supported for Flux and SDXL.
```sh
python quantize.py \
--model {flux-dev|sdxl-1.0|sdxl-turbo|ltx-video-dev} --model-dtype {Half|BFloat16} --trt-high-precision-dtype {Half|BFloat16} \
--format {fp8|fp4} --batch-size 2 --calib-size {128|256} --quantize-mha \
--format {fp8|fp4} --batch-size 2 --calib-size {128|256} \
--n-steps 20 --quantized-torch-ckpt-save-path ./{MODEL_NAME}.pt --collect-method default \
--onnx-dir {ONNX_DIR}
```
@@ -31,6 +31,9 @@ NVFP4_DEFAULT_CONFIG = load_config(
NVFP4_FP8_MHA_CONFIG = load_config(
"configs/ptq/presets/diffusers/nvfp4_fp8_mha", schema_type=QuantizeConfig
).model_dump(exclude_unset=True)
NVFP4_FP8_CONV_CONFIG = load_config(
"configs/ptq/presets/diffusers/nvfp4_fp8_conv", schema_type=QuantizeConfig
).model_dump(exclude_unset=True)
def set_quant_config_attr(quant_config, trt_high_precision_dtype, quant_algo, **kwargs):
@@ -51,8 +51,6 @@ from modelopt.onnx.export import NVFP4QuantExporter
from modelopt.torch.quantization.export_onnx import configure_linear_module_onnx_quantizers
from modelopt.torch.utils import torch_to
from .fp8_onnx_graphsurgeon import convert_zp_fp8
MODEL_ID_TO_DYNAMIC_AXES = {
"sdxl-1.0": {
"sample": {0: "batch_size", 1: "num_channels", 2: "height", 3: "width"},
@@ -124,18 +122,6 @@ def flux_convert_rope_weight_type(onnx_graph):
return gs.export_onnx(graph)
def generate_fp8_scales(backbone):
# temporary solution due to a known bug in torch.onnx._dynamo_export
for _, module in backbone.named_modules():
if isinstance(module, (torch.nn.Linear, torch.nn.Conv2d)) and (
hasattr(module.input_quantizer, "_amax") and module.input_quantizer is not None
):
module.input_quantizer._num_bits = 8
module.weight_quantizer._num_bits = 8
module.input_quantizer._amax = module.input_quantizer._amax * (127 / 448.0)
module.weight_quantizer._amax = module.weight_quantizer._amax * (127 / 448.0)
def _gen_dummy_inp_and_dyn_shapes_sdxl(backbone, min_bs=1, opt_bs=1):
assert isinstance(backbone, UNet2DConditionModel) or isinstance(
backbone._orig_mod, UNet2DConditionModel
@@ -469,7 +455,6 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision):
tmp_subfolder = tempfile.mkdtemp(prefix="myapp_")
tmp_output = Path(f"{tmp_subfolder}/{model_file_name}")
q_output = Path(f"{onnx_dir}/{model_file_name}")
quantizer_context = (
configure_linear_module_onnx_quantizers(backbone) if precision == "fp4" else nullcontext()
)
@@ -536,16 +521,8 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision):
)
print(f"Saved at {tmp_output}")
onnx_model = onnx.load(str(tmp_output), load_external_data=True)
if precision == "fp8":
if not model_name.startswith("flux"):
graph = gs.import_onnx(onnx_model)
graph.cleanup().toposort()
onnx_model = gs.export_onnx(graph)
onnx_model = convert_zp_fp8(onnx_model)
graph = gs.import_onnx(onnx_model)
onnx_model = gs.export_onnx(graph.cleanup())
else:
flux_convert_rope_weight_type(onnx_model)
if precision == "fp8" and model_name.startswith("flux"):
onnx_model = flux_convert_rope_weight_type(onnx_model)
if precision == "fp4":
onnx_model = NVFP4QuantExporter.process_model(onnx_model)
save_onnx(onnx_model, q_output)
@@ -97,29 +97,6 @@ def insert_cast(graph, input_tensor, attrs):
next_node.inputs[idx] = output_tensor
def convert_zp_fp8(onnx_graph):
"""
Convert Q/DQ zero datatype from INT8 to FP8.
We use this WAR because FP8 Conv cannot be exported to ONNX directly.
The workaround is to first convert the FP8 QDQs into INT8 QDQs,
then modify the ONNX model afterward to change those INT8 QDQs back into FP8 QDQs.
"""
# Find all zero constant nodes
qdq_zero_nodes = set()
for node in onnx_graph.graph.node:
if node.op_type == "QuantizeLinear" and len(node.input) > 2:
qdq_zero_nodes.add(node.input[2])
print(f"[WAR], found {len(qdq_zero_nodes)} INT8 QDQ pairs, you can ignore this message..")
# Convert zero point datatype from INT8 to FP8.
for node in onnx_graph.graph.node:
if node.output[0] in qdq_zero_nodes:
node.attribute[0].t.data_type = onnx.TensorProto.FLOAT8E4M3FN
return onnx_graph
def cast_resize_io(graph):
"""
After all activations and weights are converted to fp16, we will
+41 -31
View File
@@ -27,6 +27,7 @@ from config import (
FP8_DEFAULT_CONFIG,
INT8_DEFAULT_CONFIG,
NVFP4_DEFAULT_CONFIG,
NVFP4_FP8_CONV_CONFIG,
NVFP4_FP8_MHA_CONFIG,
reset_set_int8_config,
set_quant_config_attr,
@@ -55,6 +56,9 @@ from utils import check_conv_and_mha, check_lora
import modelopt.torch.opt as mto
import modelopt.torch.quantization as mtq
from modelopt.torch.export import export_hf_checkpoint
from modelopt.torch.quantization.nn import TensorQuantizer
_SDXL_MODEL_TYPES = (ModelType.SDXL_BASE, ModelType.SDXL_TURBO)
def setup_logging(verbose: bool = False) -> logging.Logger:
@@ -130,7 +134,9 @@ class Quantizer:
elif self.config.format == QuantFormat.FP8:
base_cfg = FP8_DEFAULT_CONFIG
elif self.config.format == QuantFormat.FP4:
if self.model_config.model_type.value.startswith("flux"):
if self.model_config.model_type in _SDXL_MODEL_TYPES:
base_cfg = NVFP4_FP8_CONV_CONFIG
elif self.model_config.model_type.value.startswith("flux"):
base_cfg = NVFP4_FP8_MHA_CONFIG
else:
base_cfg = NVFP4_DEFAULT_CONFIG
@@ -271,23 +277,6 @@ class ExportManager:
self.logger = logger
self.pipeline_manager = pipeline_manager
def _has_conv_layers(self, model: torch.nn.Module) -> bool:
"""
Check if the model contains any convolutional layers.
Args:
model: Model to check
Returns:
True if model contains Conv layers, False otherwise
"""
for module in model.modules():
if isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)) and (
module.input_quantizer.is_enabled or module.weight_quantizer.is_enabled
):
return True
return False
def save_checkpoint(
self,
backbone: torch.nn.Module,
@@ -335,15 +324,10 @@ class ExportManager:
# Deferred: the ONNX stack (onnx, onnx_graphsurgeon, ...) is only needed
# for --onnx-dir exports; HF-checkpoint-only runs must not require it.
from onnx_utils.export import generate_fp8_scales, modelopt_export_sd
from onnx_utils.export import modelopt_export_sd
self.logger.info(f"Starting ONNX export to {self.config.onnx_dir}")
if quant_format == QuantFormat.FP8 and self._has_conv_layers(backbone):
self.logger.info(
"Detected quantizing conv layers in backbone. Generating FP8 scales..."
)
generate_fp8_scales(backbone)
self.logger.info("Preparing models for export...")
pipe.to("cpu")
torch.cuda.empty_cache()
@@ -457,7 +441,7 @@ def create_argument_parser() -> argparse.ArgumentParser:
%(prog)s --model ltx-video-dev --format fp8 --batch-size 1 --calib-size 32 --ltx-skip-upsampler
# Restore and export a previously quantized model
%(prog)s --model flux-schnell --restore-from checkpoint.pt --onnx-dir ./exports/
%(prog)s --model flux-schnell --restore-from ./checkpoints/ --onnx-dir ./exports/
""",
)
model_group = parser.add_argument_group("Model Configuration")
@@ -586,7 +570,9 @@ def create_argument_parser() -> argparse.ArgumentParser:
help="Directory for HuggingFace checkpoint export",
)
export_group.add_argument(
"--restore-from", type=str, help="Path to restore from previous checkpoint"
"--restore-from",
type=str,
help="Checkpoint directory; quantization format and MHA policy are restored automatically",
)
export_group.add_argument(
"--trt-high-precision-dtype",
@@ -600,6 +586,25 @@ def create_argument_parser() -> argparse.ArgumentParser:
return parser
def _infer_restored_quantization_format(
backbones: list[tuple[str, torch.nn.Module]],
) -> QuantFormat:
has_nvfp4 = False
has_fp8 = False
for _, backbone in backbones:
for module in backbone.modules():
if isinstance(module, TensorQuantizer) and module.is_enabled:
has_nvfp4 |= module.is_nvfp4_dynamic or module.is_nvfp4_static
has_fp8 |= module.is_fp8
if has_nvfp4:
return QuantFormat.FP4
if has_fp8:
return QuantFormat.FP8
return QuantFormat.INT8
def main() -> None:
from diffusers.models.normalization import RMSNorm as DiffuserRMSNorm
@@ -674,9 +679,9 @@ def main() -> None:
)
logger.info("Validating configurations...")
quant_config.validate()
export_config.validate()
if not export_config.restore_from:
quant_config.validate()
calib_config.validate()
pipeline_manager = PipelineManager(model_config, logger)
@@ -685,8 +690,12 @@ def main() -> None:
export_manager = ExportManager(export_config, logger, pipeline_manager)
if export_config.restore_from and export_config.restore_from.exists():
if export_config.restore_from:
export_manager.restore_checkpoint()
quant_config.format = _infer_restored_quantization_format(
list(pipeline_manager.iter_backbones())
)
logger.info(f"Detected restored quantization format: {quant_config.format.value}")
else:
logger.info("Initializing calibration...")
@@ -716,11 +725,12 @@ def main() -> None:
mtq.compress(backbone)
logger.info(f"{backbone_name} compression completed")
# For VAE backbones, skip check_conv_and_mha — the whole point
# of VAE quantization is to quantize Conv layers.
if backbone_name not in ("video_decoder", "vae"):
check_conv_and_mha(
backbone, quant_config.format == QuantFormat.FP4, quant_config.quantize_mha
backbone,
quant_config.format == QuantFormat.FP4
and model_config.model_type not in _SDXL_MODEL_TYPES,
quant_config.quantize_mha,
)
export_manager.save_checkpoint(backbone, backbone_name)
-3
View File
@@ -64,11 +64,8 @@ def check_conv_and_mha(backbone, if_fp4, quantize_mha):
):
if hasattr(module, attr):
getattr(module, attr).disable()
setattr(module, "_disable_fp8_mha", True)
print(f"Disabled Attention layer quantization for layer {name}")
else:
setattr(module, "_disable_fp8_mha", False)
def filter_func_ltx_video(name: str) -> bool: