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:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user