[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
+6
View File
@@ -430,4 +430,10 @@ class NVFP4QuantExporter(ONNXQuantExporter):
utils.topologically_sort_graph_nodes(graph)
if fp4_qdq_nodes:
default_opset = next(
opset for opset in onnx_model.opset_import if opset.domain in {"", "ai.onnx"}
)
default_opset.version = max(default_opset.version, 23)
return onnx_model
+4 -1
View File
@@ -225,9 +225,12 @@ def _fp8_quantize(
"Constant",
value_t=torch.tensor(scale_inv).to(torch_dtype_map[inputs.type().scalarType()]),
)
return g.op("trt::TRT_FP8QuantizeLinear", inputs, scale).setType(
quantized = g.op("trt::TRT_FP8QuantizeLinear", inputs, scale).setType(
inputs.type().with_dtype(torch.uint8).with_sizes(output_shape)
)
# PyTorch runs shape inference before setType for custom ops, so refresh its reliability state.
torch._C._jit_pass_onnx_node_shape_type_inference(quantized.node(), g.params_dict, g.opset)
return quantized
def _fp8_dequantize(
@@ -141,6 +141,14 @@ def _quantized_sdpa(self, *args, **kwargs):
q_quantized_scale = self.q_bmm_quantizer._get_amax(query)
k_quantized_scale = self.k_bmm_quantizer._get_amax(key)
v_quantized_scale = self.v_bmm_quantizer._get_amax(value)
disable_fp8_mha = not all(
quantizer.is_enabled and quantizer.is_fp8
for quantizer in (
self.q_bmm_quantizer,
self.k_bmm_quantizer,
self.v_bmm_quantizer,
)
)
# We don't need to calibrate the output of softmax
return self.bmm2_output_quantizer(
@@ -155,7 +163,7 @@ def _quantized_sdpa(self, *args, **kwargs):
self.q_bmm_quantizer.trt_high_precision_dtype
if hasattr(self.q_bmm_quantizer, "trt_high_precision_dtype")
else "Half",
self._disable_fp8_mha if hasattr(self, "_disable_fp8_mha") else True,
disable_fp8_mha,
)
)