From ccece7f916a4ce4db143d2189e092425fbac426d Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Fri, 4 Sep 2026 17:45:39 +0000 Subject: [PATCH 01/10] [5565357] Fix SDXL NVFP4 export and performance Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/ONNX-TRT-Deployment.md | 18 +- examples/diffusers/quantization/config.py | 3 + .../quantization/onnx_utils/export.py | 1082 +++++++++++++++-- examples/diffusers/quantization/quantize.py | 152 ++- examples/diffusers/quantization/utils.py | 250 +++- .../ptq/presets/diffusers/nvfp4_fp8_conv.yaml | 53 + tests/examples/diffusers/test_diffusers.py | 24 + .../test_diffusers_fp4_onnx_validation.py | 992 +++++++++++++++ .../examples/test_diffusers_fp4_validation.py | 652 ++++++++++ 9 files changed, 3073 insertions(+), 153 deletions(-) create mode 100644 modelopt_recipes/configs/ptq/presets/diffusers/nvfp4_fp8_conv.yaml create mode 100644 tests/unit/examples/test_diffusers_fp4_onnx_validation.py create mode 100644 tests/unit/examples/test_diffusers_fp4_validation.py diff --git a/examples/diffusers/quantization/ONNX-TRT-Deployment.md b/examples/diffusers/quantization/ONNX-TRT-Deployment.md index 57448b8a38e..bb0239d5c31 100644 --- a/examples/diffusers/quantization/ONNX-TRT-Deployment.md +++ b/examples/diffusers/quantization/ONNX-TRT-Deployment.md @@ -28,28 +28,30 @@ 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 the SDXL family. For SDXL, the FP4 recipe uses block-16 NVFP4 for non-QKV Linear/GEMM layers and FP8 for Conv2d layers. Attention `to_q`, `to_k`, and `to_v` projection Linears remain in the high-precision model dtype so TensorRT can preserve their horizontal projection fusion. The script selects this mixed-precision recipe automatically when `--format fp4` is used. ```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} ``` +Add `--quantize-mha` to opt in to FP8 MHA quantization; this does not quantize the QKV projection Linears. + We recommend using a device with a minimum of 48GB of combined CPU and GPU memory for exporting ONNX models. If not, please use CPU for ONNX export. ## Build the TRT engine for the Quantized ONNX Backbone > [!IMPORTANT] > TensorRT environment must be setup prior -- Please see [Pre-Requisites](../README.md#pre-requisites) -> INT8 requires **TensorRT version >= 9.2.0**. If you prefer to use the FP8 TensorRT, ensure you have **TensorRT version 10.2.0 or higher**. You can download the latest version of TensorRT at [here](https://developer.nvidia.com/tensorrt/download). Deployment of SVDQuant is currently not supported. +> INT8 requires **TensorRT version >= 9.2.0**. FP8 requires **TensorRT version 10.2.0 or higher**. FP4 requires a Blackwell GPU (SM100 or newer) and a TensorRT version with NVFP4 support. You can download the latest version of TensorRT [here](https://developer.nvidia.com/tensorrt/download). Deployment of SVDQuant is currently not supported. -Generate INT8/FP8 Backbone Engine +Generate INT8/FP8/FP4 Backbone Engine ```bash -# For SDXL +# For SDXL INT8, FP8, or FP4 trtexec --builderOptimizationLevel=4 --stronglyTyped --onnx=./model.onnx \ --minShapes=sample:2x4x128x128,timestep:1,encoder_hidden_states:2x77x2048,text_embeds:2x1280,time_ids:2x6 \ --optShapes=sample:16x4x128x128,timestep:1,encoder_hidden_states:16x77x2048,text_embeds:16x1280,time_ids:16x6 \ @@ -91,15 +93,15 @@ python demo_txt2img_xl.py "enchanted winter forest, soft diffuse light on a snow Note, it will take some time to build TRT engines for the first time -- Replace the fp16 backbone TRT engine with int8 engine generated in [Build the TRT engine for the Quantized ONNX Backbone](#build-the-trt-engine-for-the-quantized-onnx-backbone), e.g.,: +- Replace the FP16 backbone TensorRT engine with the quantized engine generated in [Build the TRT engine for the Quantized ONNX Backbone](#build-the-trt-engine-for-the-quantized-onnx-backbone), e.g.: ```sh cp -r {YOUR_UNETXL}.plan ./engine/ ``` -Note, the engines must be built on the same GPU, and ensure that the INT8 engine name matches the names of the FP16 engines to enable compatibility with the demoDiffusion pipeline. +The engines must be built on the same GPU, and the quantized engine name must match the FP16 engine name to enable compatibility with the demoDiffusion pipeline. -- Run the above txt2img example command again. You can compare the generated images and latency for fp16 vs int8. +- Run the above txt2img example command again. You can compare the generated images and latency for FP16 versus INT8, FP8, or FP4. Similarly, you could run end-to-end pipeline with Model Optimizer quantized backbone and corresponding examples in demoDiffusion with other diffusion models. ## Running the inference pipeline with DeviceModel diff --git a/examples/diffusers/quantization/config.py b/examples/diffusers/quantization/config.py index cb8fdf3a5da..3ab8f5db87c 100644 --- a/examples/diffusers/quantization/config.py +++ b/examples/diffusers/quantization/config.py @@ -31,6 +31,9 @@ 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): diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 5da795f0f48..7936bc1933d 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -32,9 +32,11 @@ import os import shutil import tempfile -from contextlib import nullcontext +import uuid +from contextlib import nullcontext, suppress from pathlib import Path +import numpy as np import onnx import onnx_graphsurgeon as gs import torch @@ -48,7 +50,9 @@ from torch.onnx import export as onnx_export from modelopt.onnx.export import NVFP4QuantExporter +from modelopt.onnx.quantization.graph_utils import get_tensor_consumer_nodes from modelopt.torch.quantization.export_onnx import configure_linear_module_onnx_quantizers +from modelopt.torch.quantization.nn.modules.quant_linear import RealQuantLinear from modelopt.torch.utils import torch_to from .fp8_onnx_graphsurgeon import convert_zp_fp8 @@ -124,16 +128,36 @@ def flux_convert_rope_weight_type(onnx_graph): return gs.export_onnx(graph) -def generate_fp8_scales(backbone): +def generate_fp8_scales(backbone, *, conv_only=False): # 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) + module_types = (torch.nn.Conv2d,) if conv_only else (torch.nn.Linear, torch.nn.Conv2d) + quantizer_states = [] + try: + for _, module in backbone.named_modules(): + if not isinstance(module, module_types): + continue + for quantizer_name in ("input_quantizer", "weight_quantizer"): + quantizer = getattr(module, quantizer_name, None) + if ( + quantizer is None + or not quantizer.is_enabled + or quantizer.num_bits != (4, 3) + or getattr(quantizer, "_amax", None) is None + ): + continue + quantizer_states.append((quantizer, quantizer._num_bits, quantizer._amax)) + quantizer._num_bits = 8 + quantizer._amax = quantizer._amax * (127 / 448.0) + except BaseException: + restore_fp8_scales(quantizer_states) + raise + return quantizer_states + + +def restore_fp8_scales(quantizer_states): + for quantizer, num_bits, amax in reversed(quantizer_states): + quantizer._num_bits = num_bits + quantizer._amax = amax def _gen_dummy_inp_and_dyn_shapes_sdxl(backbone, min_bs=1, opt_bs=1): @@ -451,102 +475,978 @@ def remove_nesting(trt_dynamic_shapes): ) -def save_onnx(onnx_model, output): +def _get_int_attribute(node, name): + for attribute in node.attribute: + if attribute.name == name and attribute.type == onnx.AttributeProto.INT: + return attribute.i + return None + + +_WEIGHT_PASSTHROUGH_OPS = {"Cast", "Flatten", "Identity", "Reshape", "Transpose"} + + +def _trace_initializer_source(tensor_name, producers, initializer_names): + visited = set() + while tensor_name not in initializer_names: + if tensor_name in visited: + return None + visited.add(tensor_name) + producer = producers.get(tensor_name) + if ( + producer is None + or producer.op_type not in _WEIGHT_PASSTHROUGH_OPS + or not producer.input + ): + return None + tensor_name = producer.input[0] + return tensor_name + + +def _trace_tensor_consumers( + tensor_names, + consumers, + graph_outputs, + terminal_ops, + terminal_input_index, + passthrough_ops=_WEIGHT_PASSTHROUGH_OPS, + allowed_cast_dtypes=None, +): + pending = list(tensor_names) + visited = set() + terminal_consumers = [] + invalid_consumers = set() + + while pending: + tensor_name = pending.pop() + if tensor_name in visited: + continue + visited.add(tensor_name) + + if tensor_name in graph_outputs: + invalid_consumers.add(f"graph output {tensor_name}") + tensor_consumers = consumers.get(tensor_name, []) + if not tensor_consumers: + invalid_consumers.add(f"unused tensor {tensor_name}") + continue + + for consumer in tensor_consumers: + input_indices = [ + index + for index, input_name in enumerate(consumer.input) + if input_name == tensor_name + ] + if consumer.op_type in passthrough_ops and input_indices == [0]: + if ( + consumer.op_type == "Cast" + and allowed_cast_dtypes is not None + and _get_int_attribute(consumer, "to") not in allowed_cast_dtypes + ): + invalid_consumers.add(consumer.name or consumer.op_type) + continue + if consumer.output: + pending.extend(consumer.output) + else: + invalid_consumers.add(consumer.name or consumer.op_type) + elif consumer.op_type in terminal_ops and input_indices == [terminal_input_index]: + terminal_consumers.append(consumer) + else: + invalid_consumers.add( + consumer.name or (consumer.output[0] if consumer.output else consumer.op_type) + ) + + return terminal_consumers, sorted(invalid_consumers) + + +def _trace_source_node(tensor_name, producers): + visited = set() + while tensor_name not in visited: + visited.add(tensor_name) + producer = producers.get(tensor_name) + if producer is None or producer.op_type not in _WEIGHT_PASSTHROUGH_OPS: + return producer + if not producer.input: + return None + tensor_name = producer.input[0] + return None + + +def _find_initializer_backed_qdq_weights(onnx_model, allow_fp8_conv): + initializer_names = {initializer.name for initializer in onnx_model.graph.initializer} + producers = {output: node for node in onnx_model.graph.node for output in node.output if output} + consumers = get_tensor_consumer_nodes(onnx_model.graph) + graph_outputs = {output.name for output in onnx_model.graph.output} + disallowed_consumers = set() + allowed_pairs = [] + + for node in onnx_model.graph.node: + if node.op_type != "DequantizeLinear" or not node.input: + continue + quantize_node = producers.get(node.input[0]) + if ( + quantize_node is None + or quantize_node.op_type != "QuantizeLinear" + or not quantize_node.input + or _trace_initializer_source(quantize_node.input[0], producers, initializer_names) + is None + ): + continue + + terminal_consumers, invalid_consumers = _trace_tensor_consumers( + node.output, consumers, graph_outputs, {"Conv"}, 1 + ) + if allow_fp8_conv and terminal_consumers and not invalid_consumers: + allowed_pairs.append((quantize_node, node, terminal_consumers)) + else: + disallowed_consumers.update(invalid_consumers) + disallowed_consumers.update( + consumer.name or (consumer.output[0] if consumer.output else consumer.op_type) + for consumer in terminal_consumers + ) + if not terminal_consumers and not invalid_consumers: + disallowed_consumers.add(node.name or node.output[0]) + + return allowed_pairs, sorted(disallowed_consumers) + + +def _get_tensor_dtype(tensor_name, initializers, producers): + initializer = initializers.get(tensor_name) + if initializer is not None: + return initializer.data_type + producer = producers.get(tensor_name) + if producer is None or producer.op_type != "Constant": + return None + for attribute in producer.attribute: + if attribute.name == "value" and attribute.type == onnx.AttributeProto.TENSOR: + return attribute.t.data_type + return None + + +def _get_effective_tensor_dtype(tensor_name, initializers, producers, declared_dtypes): + visited = set() + while tensor_name not in visited: + visited.add(tensor_name) + producer = producers.get(tensor_name) + if producer is not None and producer.op_type == "Cast": + return _get_int_attribute(producer, "to") + if tensor_name in declared_dtypes: + return declared_dtypes[tensor_name] + tensor_dtype = _get_tensor_dtype(tensor_name, initializers, producers) + if tensor_dtype is not None: + return tensor_dtype + if producer is None or producer.op_type not in _WEIGHT_PASSTHROUGH_OPS: + return None + if not producer.input: + return None + tensor_name = producer.input[0] + return None + + +def _get_constant_array(tensor_name, initializers, producers): + tensor = initializers.get(tensor_name) + if tensor is None: + producer = producers.get(tensor_name) + if producer is not None and producer.op_type == "Constant": + tensor = next( + ( + attribute.t + for attribute in producer.attribute + if attribute.name == "value" and attribute.type == onnx.AttributeProto.TENSOR + ), + None, + ) + if tensor is None: + return None + try: + return onnx.numpy_helper.to_array(tensor) + except (TypeError, ValueError): + return None + + +def _validate_positive_scalar(tensor_name, role, initializers, producers): + value = _get_constant_array(tensor_name, initializers, producers) + if value is None or value.size != 1 or not np.isfinite(value).all() or not (value > 0).all(): + return [f"{role} must be a finite positive scalar constant"] + return [] + + +def _validate_normalized_fp8_qdq_pair( + quantize_node, + dequantize_node, + pair_name, + initializers, + producers, + consumers, + expected_scale_dtype=None, +): + errors = [] + if len(quantize_node.input) != 3 or len(dequantize_node.input) != 3: + return [f"{pair_name} must use three-input FP8 Q/DQ nodes"] + if len(quantize_node.output) != 1 or dequantize_node.input[0] != quantize_node.output[0]: + errors.append(f"{pair_name} does not form a direct Q/DQ pair") + elif consumers.get(quantize_node.output[0], []) != [dequantize_node]: + errors.append(f"{pair_name} quantized tensor must be consumed only by its DQ") + + if quantize_node.input[1:] != dequantize_node.input[1:]: + errors.append(f"{pair_name} Q/DQ nodes do not share scale and zero point") + if any( + attribute.name == "axis" + for node in (quantize_node, dequantize_node) + for attribute in node.attribute + ): + errors.append(f"{pair_name} must use per-tensor FP8 Q/DQ without an axis") + + errors.extend( + _validate_positive_scalar( + quantize_node.input[1], f"{pair_name} FP8 scale", initializers, producers + ) + ) + scale_dtype = _get_tensor_dtype(quantize_node.input[1], initializers, producers) + if scale_dtype not in { + onnx.TensorProto.FLOAT, + onnx.TensorProto.FLOAT16, + onnx.TensorProto.BFLOAT16, + }: + errors.append(f"{pair_name} FP8 scale must use a floating-point dtype") + elif expected_scale_dtype is not None and scale_dtype != expected_scale_dtype: + errors.append(f"{pair_name} FP8 scale dtype does not match its quantized tensor") + for role, node in ( + ("QuantizeLinear", quantize_node), + ("DequantizeLinear", dequantize_node), + ): + zero_name = node.input[2] + if _get_tensor_dtype(zero_name, initializers, producers) != onnx.TensorProto.FLOAT8E4M3FN: + errors.append(f"{pair_name} {role} zero point is not FLOAT8E4M3FN") + continue + zero = _get_constant_array(zero_name, initializers, producers) + if zero is None or zero.size != 1 or not (zero == 0).all(): + errors.append(f"{pair_name} {role} zero point must be a scalar zero") + return errors + + +def _validate_normalized_fp8_qdq(onnx_model, qdq_records): + initializers = {initializer.name: initializer for initializer in onnx_model.graph.initializer} + producers = {output: node for node in onnx_model.graph.node for output in node.output if output} + consumers = get_tensor_consumer_nodes(onnx_model.graph) + graph_outputs = {output.name for output in onnx_model.graph.output} + initializer_names = set(initializers) + declared_dtypes = { + value.name: value.type.tensor_type.elem_type + for values in ( + onnx_model.graph.input, + onnx_model.graph.value_info, + onnx_model.graph.output, + ) + for value in values + if value.type.HasField("tensor_type") + } + errors = [] + fp8_conv_ids = { + id(consumer) for _, _, terminal_consumers in qdq_records for consumer in terminal_consumers + } + validated_activation_dq_ids = set() + + for quantize_node, dequantize_node, terminal_consumers in qdq_records: + pair_name = dequantize_node.name or dequantize_node.output[0] + if len(terminal_consumers) != 1: + errors.append(f"{pair_name} weight DQ must feed exactly one FP8 Conv input 1") + errors.extend( + _validate_normalized_fp8_qdq_pair( + quantize_node, + dequantize_node, + pair_name, + initializers, + producers, + consumers, + _get_effective_tensor_dtype( + quantize_node.input[0], initializers, producers, declared_dtypes + ), + ) + ) + + for conv_node in terminal_consumers: + conv_name = conv_node.name or conv_node.output[0] + activation_dq = _trace_source_node(conv_node.input[0], producers) + if activation_dq is None or activation_dq.op_type != "DequantizeLinear": + errors.append(f"{conv_name} has no FP8 activation Q/DQ on input 0") + continue + activation_q = producers.get(activation_dq.input[0]) if activation_dq.input else None + if activation_q is None or activation_q.op_type != "QuantizeLinear": + errors.append(f"{conv_name} has no FP8 activation QuantizeLinear on input 0") + continue + if ( + not activation_q.input + or _trace_initializer_source(activation_q.input[0], producers, initializer_names) + is not None + ): + errors.append(f"{conv_name} activation Q/DQ is initializer-backed") + continue + + if id(activation_dq) not in validated_activation_dq_ids: + activation_name = activation_dq.name or activation_dq.output[0] + errors.extend( + _validate_normalized_fp8_qdq_pair( + activation_q, + activation_dq, + activation_name, + initializers, + producers, + consumers, + _get_effective_tensor_dtype( + activation_q.input[0], initializers, producers, declared_dtypes + ), + ) + ) + activation_consumers, invalid_consumers = _trace_tensor_consumers( + activation_dq.output, consumers, graph_outputs, {"Conv"}, 0 + ) + if invalid_consumers: + errors.append( + f"{activation_name} has non-Conv activation consumers: " + + ", ".join(invalid_consumers[:5]) + ) + if len(activation_consumers) != 1 or id(conv_node) not in { + id(consumer) for consumer in activation_consumers + }: + errors.append( + f"{activation_name} must feed exactly one validated FP8 Conv input 0" + ) + unexpected_conv_ids = { + id(consumer) for consumer in activation_consumers + } - fp8_conv_ids + if unexpected_conv_ids: + errors.append( + f"{activation_name} reaches a Conv without a validated FP8 weight Q/DQ" + ) + validated_activation_dq_ids.add(id(activation_dq)) + return errors + + +def _validate_dynamic_fp4_activations( + onnx_model, expected_block_size, fp4_weight_terminal_consumers +): + initializers = {initializer.name: initializer for initializer in onnx_model.graph.initializer} + producers = {output: node for node in onnx_model.graph.node for output in node.output if output} + consumers = get_tensor_consumer_nodes(onnx_model.graph) + graph_outputs = {output.name for output in onnx_model.graph.output} + expected_terminal_ids = {id(node) for node in fp4_weight_terminal_consumers} + dynamic_terminal_counts = {} + initializer_or_constant_names = set(initializers) | { + output + for node in onnx_model.graph.node + if node.op_type == "Constant" + for output in node.output + } + errors = [] + dynamic_nodes = [ + node for node in onnx_model.graph.node if node.op_type == "TRT_FP4DynamicQuantize" + ] + + if not dynamic_nodes: + errors.append("no TRT_FP4DynamicQuantize activation nodes were exported") + + for node in dynamic_nodes: + node_name = node.name or (node.output[0] if node.output else "") + if len(node.input) != 2 or len(node.output) != 2: + errors.append(f"{node_name} must have two inputs and two outputs") + continue + if node.domain != "trt": + errors.append(f"{node_name} must use the trt domain") + if ( + _trace_initializer_source(node.input[0], producers, initializer_or_constant_names) + is not None + ): + errors.append(f"{node_name} input 0 must be a dynamic activation") + if _get_int_attribute(node, "block_size") != expected_block_size: + errors.append(f"{node_name} does not use block_size={expected_block_size}") + if _get_int_attribute(node, "axis") != -1: + errors.append(f"{node_name} does not use axis=-1") + if _get_int_attribute(node, "scale_type") != onnx.TensorProto.FLOAT8E4M3FN: + errors.append(f"{node_name} does not produce FLOAT8E4M3FN block scales") + if _get_tensor_dtype(node.input[1], initializers, producers) != onnx.TensorProto.FLOAT: + errors.append(f"{node_name} global scale is not FLOAT") + errors.extend( + _validate_positive_scalar( + node.input[1], f"{node_name} global scale", initializers, producers + ) + ) + + quantized_consumers = consumers.get(node.output[0], []) + if ( + len(quantized_consumers) != 1 + or quantized_consumers[0].op_type != "DequantizeLinear" + or not quantized_consumers[0].input + or quantized_consumers[0].input[0] != node.output[0] + ): + errors.append(f"{node_name} FP4 output must feed exactly one DequantizeLinear") + continue + activation_dq = quantized_consumers[0] + activation_dq_name = activation_dq.name or activation_dq.output[0] + if len(activation_dq.input) != 2: + errors.append(f"{activation_dq_name} must be a two-input DequantizeLinear") + continue + if _get_int_attribute(activation_dq, "block_size") != expected_block_size: + errors.append(f"{activation_dq_name} does not use block_size={expected_block_size}") + if _get_int_attribute(activation_dq, "axis") != -1: + errors.append(f"{activation_dq_name} does not use axis=-1") + + scale_consumers = consumers.get(node.output[1], []) + if ( + len(scale_consumers) != 1 + or scale_consumers[0].op_type != "DequantizeLinear" + or not scale_consumers[0].input + or scale_consumers[0].input[0] != node.output[1] + ): + errors.append(f"{node_name} FP8 scale output must feed exactly one DequantizeLinear") + continue + scale_dq = scale_consumers[0] + scale_dq_name = scale_dq.name or scale_dq.output[0] + if len(scale_dq.input) != 2: + errors.append(f"{scale_dq_name} must be a two-input DequantizeLinear") + continue + if any(attribute.name in {"axis", "block_size"} for attribute in scale_dq.attribute): + errors.append(f"{scale_dq_name} must not use axis or block_size") + if _get_tensor_dtype(scale_dq.input[1], initializers, producers) != onnx.TensorProto.FLOAT: + errors.append(f"{scale_dq_name} global scale is not FLOAT") + errors.extend( + _validate_positive_scalar( + scale_dq.input[1], f"{scale_dq_name} global scale", initializers, producers + ) + ) + quantize_scale = _get_constant_array(node.input[1], initializers, producers) + dequantize_scale = _get_constant_array(scale_dq.input[1], initializers, producers) + if ( + quantize_scale is not None + and dequantize_scale is not None + and not np.array_equal(quantize_scale, dequantize_scale) + ): + errors.append(f"{node_name} quantize and dequantize global scales do not match") + if not scale_dq.output or activation_dq.input[1] != scale_dq.output[0]: + errors.append(f"{activation_dq_name} is not scaled by {scale_dq_name}") + elif consumers.get(scale_dq.output[0], []) != [activation_dq]: + errors.append(f"{scale_dq_name} output must be consumed only by {activation_dq_name}") + + terminal_consumers, invalid_consumers = _trace_tensor_consumers( + activation_dq.output, + consumers, + graph_outputs, + {"Gemm", "MatMul"}, + 0, + {"Cast", "Identity"}, + {onnx.TensorProto.FLOAT16, onnx.TensorProto.BFLOAT16}, + ) + for consumer in terminal_consumers: + terminal_id = id(consumer) + dynamic_terminal_counts[terminal_id] = dynamic_terminal_counts.get(terminal_id, 0) + 1 + if not terminal_consumers: + errors.append(f"{activation_dq_name} does not reach a Gemm/MatMul activation input") + elif len(terminal_consumers) != 1: + errors.append( + f"{activation_dq_name} must feed exactly one Gemm/MatMul activation input" + ) + if invalid_consumers: + errors.append( + f"{activation_dq_name} has non-activation consumers: " + + ", ".join(invalid_consumers[:5]) + ) + + dynamic_terminal_ids = set(dynamic_terminal_counts) + duplicate = sum(count != 1 for count in dynamic_terminal_counts.values()) + if dynamic_terminal_ids != expected_terminal_ids or duplicate: + missing = len(expected_terminal_ids - dynamic_terminal_ids) + extra = len(dynamic_terminal_ids - expected_terminal_ids) + errors.append( + "dynamic NVFP4 activation paths do not match FLOAT4 weight consumers " + f"(missing={missing}, extra={extra}, duplicate={duplicate})" + ) + return errors + + +def _validate_raw_fp4_graph( + onnx_model, + expected_block_size=16, + *, + allow_fp8_conv=False, + expected_linear_count=None, + expected_fp8_conv_count=None, +): + initializer_names = {initializer.name for initializer in onnx_model.graph.initializer} + consumers = get_tensor_consumer_nodes(onnx_model.graph) + graph_outputs = {output.name for output in onnx_model.graph.output} + fp4_nodes = [node for node in onnx_model.graph.node if node.op_type == "TRT_FP4QDQ"] + errors = [] + + if not fp4_nodes: + errors.append("no TRT_FP4QDQ weight markers were exported") + if expected_linear_count is not None and len(fp4_nodes) != expected_linear_count: + errors.append( + f"found {len(fp4_nodes)} TRT_FP4QDQ weight markers, expected " + f"{expected_linear_count} enabled Linear pairs" + ) + + for node in fp4_nodes: + node_name = node.name or (node.output[0] if node.output else "") + if not node.input or node.input[0] not in initializer_names: + errors.append(f"{node_name} is not backed by a weight initializer") + block_size = _get_int_attribute(node, "block_size") + if block_size != expected_block_size: + errors.append( + f"{node_name} has block_size={block_size}, expected {expected_block_size}" + ) + terminal_consumers, invalid_consumers = _trace_tensor_consumers( + node.output, consumers, graph_outputs, {"Gemm", "MatMul"}, 1 + ) + if not terminal_consumers: + errors.append(f"{node_name} does not reach a Gemm/MatMul weight input") + if invalid_consumers: + errors.append( + f"{node_name} has non-weight consumers: " + ", ".join(invalid_consumers[:5]) + ) + + fp8_conv_pairs, disallowed_qdq_weights = _find_initializer_backed_qdq_weights( + onnx_model, allow_fp8_conv + ) + if allow_fp8_conv and not fp8_conv_pairs: + errors.append("no initializer-backed FP8 Conv weight Q/DQ was exported") + if expected_fp8_conv_count is not None and len(fp8_conv_pairs) != expected_fp8_conv_count: + errors.append( + f"found {len(fp8_conv_pairs)} initializer-backed FP8 Conv weight Q/DQ pairs, " + f"expected {expected_fp8_conv_count} enabled Conv2d pairs" + ) + if disallowed_qdq_weights: + errors.append( + "disallowed initializer-backed Q/DQ weight consumers: " + + ", ".join(disallowed_qdq_weights[:5]) + ) + + if errors: + raise ValueError("Invalid raw FP4 ONNX graph: " + "; ".join(errors)) + return len(fp4_nodes) + + +def _validate_final_fp4_graph( + onnx_model, + expected_weight_count, + expected_block_size=16, + *, + allow_fp8_conv=False, + expected_fp8_conv_count=None, +): + initializers = {initializer.name: initializer for initializer in onnx_model.graph.initializer} + producers = {output: node for node in onnx_model.graph.node for output in node.output if output} + consumers = get_tensor_consumer_nodes(onnx_model.graph) + graph_outputs = {output.name for output in onnx_model.graph.output} + fp4_initializer_names = { + name + for name, initializer in initializers.items() + if initializer.data_type == onnx.TensorProto.FLOAT4E2M1 + } + weight_dq_nodes = [ + node + for node in onnx_model.graph.node + if node.op_type == "DequantizeLinear" + and node.input + and node.input[0] in fp4_initializer_names + ] + errors = [] + weight_dq_ids = {id(node) for node in weight_dq_nodes} + fp4_weight_terminal_consumers = [] + + remaining_markers = [node for node in onnx_model.graph.node if node.op_type == "TRT_FP4QDQ"] + if remaining_markers: + errors.append(f"{len(remaining_markers)} TRT_FP4QDQ weight markers remain") + if len(fp4_initializer_names) != expected_weight_count: + errors.append( + f"found {len(fp4_initializer_names)} FLOAT4 weights, expected {expected_weight_count}" + ) + if len(weight_dq_nodes) != expected_weight_count: + errors.append( + f"found {len(weight_dq_nodes)} FLOAT4 weight DQ nodes, expected {expected_weight_count}" + ) + + for initializer_name in fp4_initializer_names: + direct_consumers = consumers.get(initializer_name, []) + if ( + len(direct_consumers) != 1 + or id(direct_consumers[0]) not in weight_dq_ids + or [ + index + for index, input_name in enumerate(direct_consumers[0].input) + if input_name == initializer_name + ] + != [0] + ): + errors.append( + f"FLOAT4 weight {initializer_name} must feed exactly one weight DequantizeLinear" + ) + + fp4_weight_names = set() + fp8_scale_names = set() + for node in weight_dq_nodes: + node_name = node.name or (node.output[0] if node.output else "") + fp4_weight_names.add(node.input[0]) + if len(node.input) != 2: + errors.append(f"{node_name} is not a two-input FLOAT4 DequantizeLinear") + continue + if _get_int_attribute(node, "block_size") != expected_block_size: + errors.append(f"{node_name} does not use block_size={expected_block_size}") + if _get_int_attribute(node, "axis") != -1: + errors.append(f"{node_name} does not use axis=-1") + + scale_dq = producers.get(node.input[1]) + if scale_dq is None or scale_dq.op_type != "DequantizeLinear": + errors.append(f"{node_name} is not scaled by a preceding DequantizeLinear") + continue + scale_dq_name = scale_dq.name or (scale_dq.output[0] if scale_dq.output else "") + if len(scale_dq.input) != 2: + errors.append(f"{scale_dq_name} must be a two-input DequantizeLinear") + continue + if any(attribute.name in {"axis", "block_size"} for attribute in scale_dq.attribute): + errors.append(f"{scale_dq_name} must not use axis or block_size") + + fp8_scale = initializers.get(scale_dq.input[0]) + global_scale = initializers.get(scale_dq.input[1]) + if fp8_scale is None or fp8_scale.data_type != onnx.TensorProto.FLOAT8E4M3FN: + errors.append(f"{node_name} does not use a FLOAT8E4M3FN block-scale initializer") + else: + fp8_scale_names.add(fp8_scale.name) + fp8_scale_consumers = consumers.get(fp8_scale.name, []) + if ( + len(fp8_scale_consumers) != 1 + or id(fp8_scale_consumers[0]) != id(scale_dq) + or [ + index + for index, input_name in enumerate(scale_dq.input) + if input_name == fp8_scale.name + ] + != [0] + ): + errors.append(f"{fp8_scale.name} must be consumed only by {scale_dq_name} input 0") + if global_scale is None or global_scale.data_type != onnx.TensorProto.FLOAT: + errors.append(f"{node_name} does not use a FLOAT global-scale initializer") + else: + errors.extend( + _validate_positive_scalar( + global_scale.name, + f"{scale_dq_name} global scale", + initializers, + producers, + ) + ) + if len(scale_dq.output) != 1 or node.input[1] != scale_dq.output[0]: + errors.append(f"{node_name} is not scaled by {scale_dq_name}") + else: + scale_output_consumers = consumers.get(scale_dq.output[0], []) + if ( + len(scale_output_consumers) != 1 + or id(scale_output_consumers[0]) != id(node) + or [ + index + for index, input_name in enumerate(node.input) + if input_name == scale_dq.output[0] + ] + != [1] + ): + errors.append( + f"{scale_dq_name} output must be consumed only by {node_name} input 1" + ) + + terminal_consumers, invalid_consumers = _trace_tensor_consumers( + node.output, consumers, graph_outputs, {"Gemm", "MatMul"}, 1 + ) + fp4_weight_terminal_consumers.extend(terminal_consumers) + if not terminal_consumers: + errors.append(f"{node_name} does not reach a Gemm/MatMul weight input") + if invalid_consumers: + errors.append( + f"{node_name} has non-weight consumers: " + ", ".join(invalid_consumers[:5]) + ) + + if len(fp4_weight_names) != expected_weight_count: + errors.append( + f"found {len(fp4_weight_names)} referenced FLOAT4 weights, " + f"expected {expected_weight_count}" + ) + + if len(fp8_scale_names) != expected_weight_count: + errors.append( + f"found {len(fp8_scale_names)} FLOAT8 block scales, expected {expected_weight_count}" + ) + + errors.extend( + _validate_dynamic_fp4_activations( + onnx_model, expected_block_size, fp4_weight_terminal_consumers + ) + ) + + fp8_conv_pairs, disallowed_qdq_weights = _find_initializer_backed_qdq_weights( + onnx_model, allow_fp8_conv + ) + if allow_fp8_conv: + if not fp8_conv_pairs: + errors.append("no initializer-backed FP8 Conv weight Q/DQ remains") + errors.extend(_validate_normalized_fp8_qdq(onnx_model, fp8_conv_pairs)) + if expected_fp8_conv_count is not None and len(fp8_conv_pairs) != expected_fp8_conv_count: + errors.append( + f"found {len(fp8_conv_pairs)} initializer-backed FP8 Conv weight Q/DQ pairs, " + f"expected {expected_fp8_conv_count} enabled Conv2d pairs" + ) + if disallowed_qdq_weights: + errors.append( + "disallowed initializer-backed Q/DQ weight consumers: " + + ", ".join(disallowed_qdq_weights[:5]) + ) + + if errors: + raise ValueError("Invalid final FP4 ONNX graph: " + "; ".join(errors)) + + +def _normalize_fp8_qdq(onnx_model): + graph = gs.import_onnx(onnx_model) + graph.cleanup().toposort() + onnx_model = convert_zp_fp8(gs.export_onnx(graph)) + graph = gs.import_onnx(onnx_model) + return gs.export_onnx(graph.cleanup().toposort()) + + +def _ensure_default_opset(onnx_model, minimum_version): + for opset_import in onnx_model.opset_import: + if opset_import.domain in {"", "ai.onnx"}: + opset_import.version = max(opset_import.version, minimum_version) + return + opset_import = onnx_model.opset_import.add() + opset_import.domain = "" + opset_import.version = minimum_version + + +def _process_fp4_onnx_graph( + onnx_model, + model_name, + expected_block_size=16, + *, + expected_linear_count=None, + expected_fp8_conv_count=None, +): + allow_fp8_conv = model_name in {"sdxl-1.0", "sdxl-turbo"} + expected_weight_count = _validate_raw_fp4_graph( + onnx_model, + expected_block_size, + allow_fp8_conv=allow_fp8_conv, + expected_linear_count=expected_linear_count, + expected_fp8_conv_count=expected_fp8_conv_count, + ) + if allow_fp8_conv: + onnx_model = _normalize_fp8_qdq(onnx_model) + onnx_model = NVFP4QuantExporter.process_model(onnx_model) + _ensure_default_opset(onnx_model, 23) + _validate_final_fp4_graph( + onnx_model, + expected_weight_count, + expected_block_size, + allow_fp8_conv=allow_fp8_conv, + expected_fp8_conv_count=expected_fp8_conv_count, + ) + return onnx_model + + +def _get_sdxl_fp4_expected_counts(backbone): + linear_count = 0 + conv_count = 0 + for module in backbone.modules(): + input_quantizer = getattr(module, "input_quantizer", None) + weight_quantizer = getattr(module, "weight_quantizer", None) + pair_enabled = all( + quantizer is not None and getattr(quantizer, "is_enabled", False) + for quantizer in (input_quantizer, weight_quantizer) + ) + if not pair_enabled: + continue + if isinstance(module, (torch.nn.Linear, RealQuantLinear)): + linear_count += 1 + elif isinstance(module, torch.nn.Conv2d): + conv_count += 1 + return linear_count, conv_count + + +def save_onnx(onnx_model, output, external_data_name=None): onnx.save( onnx_model, str(output), save_as_external_data=True, all_tensors_to_one_file=True, - location=output.name + "_data", + location=external_data_name or output.name + "_data", size_threshold=1024, ) print(f"ONNX model saved to {output}") -def modelopt_export_sd(backbone, onnx_dir, model_name, precision): +def _get_external_data_paths(output): + fallback = output.with_name(output.name + "_data") + if not output.exists(): + return set() + try: + onnx_model = onnx.load(str(output), load_external_data=False) + except Exception: + return {fallback} if fallback.exists() else set() + + output_parent = output.parent.resolve() + paths = set() + for initializer in onnx_model.graph.initializer: + for entry in initializer.external_data: + if entry.key != "location": + continue + path = (output.parent / entry.value).resolve() + if path.parent == output_parent and path != output.resolve(): + paths.add(path) + return paths + + +def _save_onnx_atomically(onnx_model, output): + staging_dir = Path(tempfile.mkdtemp(prefix=".modelopt-export-", dir=output.parent)) + staged_output = staging_dir / output.name + external_data_name = f"{output.name}_data.{uuid.uuid4().hex}" + staged_data = staging_dir / external_data_name + published_data = output.parent / external_data_name + old_data_paths = _get_external_data_paths(output) + had_previous_output = output.exists() + previous_output = staging_dir / "previous-model.onnx" + try: + if had_previous_output: + shutil.copy2(output, previous_output) + save_onnx(onnx_model, staged_output, external_data_name=external_data_name) + onnx.checker.check_model(str(staged_output)) + has_external_data = staged_data.exists() + try: + if has_external_data: + os.replace(staged_data, published_data) + os.replace(staged_output, output) + except BaseException: + if previous_output.exists(): + os.replace(previous_output, output) + elif not had_previous_output: + output.unlink(missing_ok=True) + published_data.unlink(missing_ok=True) + raise + + for old_data_path in old_data_paths - {published_data.resolve()}: + with suppress(OSError): + old_data_path.unlink(missing_ok=True) + finally: + shutil.rmtree(staging_dir, ignore_errors=True) + + +def modelopt_export_sd(backbone, onnx_dir, model_name, precision, expected_fp4_block_size=16): model_file_name = "model.onnx" os.makedirs(f"{onnx_dir}", exist_ok=True) - tmp_subfolder = tempfile.mkdtemp(prefix="myapp_") + tmp_subfolder = tempfile.mkdtemp(prefix=".modelopt-raw-", dir=onnx_dir) tmp_output = Path(f"{tmp_subfolder}/{model_file_name}") q_output = Path(f"{onnx_dir}/{model_file_name}") + strict_sdxl_fp4 = precision == "fp4" and model_name in {"sdxl-1.0", "sdxl-turbo"} + expected_linear_count = None + expected_fp8_conv_count = None + if strict_sdxl_fp4: + expected_linear_count, expected_fp8_conv_count = _get_sdxl_fp4_expected_counts(backbone) - quantizer_context = ( - configure_linear_module_onnx_quantizers(backbone) if precision == "fp4" else nullcontext() - ) - - dummy_kwargs, dynamic_axes, _ = generate_dummy_kwargs_and_dynamic_axes_and_shapes( - model_name, backbone - ) + try: + quantizer_context = ( + configure_linear_module_onnx_quantizers(backbone) + if precision == "fp4" + else nullcontext() + ) - if model_name in ["sdxl-1.0", "sdxl-turbo"]: - input_names = ["sample", "timestep", "encoder_hidden_states", "text_embeds", "time_ids"] - output_names = ["latent"] - elif model_name == "sd3-medium": - input_names = ["hidden_states", "encoder_hidden_states", "pooled_projections", "timestep"] - output_names = ["sample"] - elif model_name == "sd3.5-medium": - input_names = ["hidden_states", "encoder_hidden_states", "pooled_projections", "timestep"] - output_names = ["out_hidden_states"] - elif model_name in ["flux-dev", "flux-schnell"]: - input_names = [ - "hidden_states", - "encoder_hidden_states", - "pooled_projections", - "timestep", - "img_ids", - "txt_ids", - ] - if model_name == "flux-dev": - input_names.append("guidance") - output_names = ["latent"] - elif model_name == "ltx-video-dev": - input_names = [ - "hidden_states", - "encoder_hidden_states", - "timestep", - "encoder_attention_mask", - "video_coords", - ] - output_names = ["latent"] - elif model_name == "wan2.2-t2v-14b": - input_names = [ - "hidden_states", - "timestep", - "encoder_hidden_states", - ] - output_names = ["latent"] - else: - raise NotImplementedError(f"Unsupported model_id: {model_name}") - - do_constant_folding = True - opset_version = 20 - - with quantizer_context, torch.inference_mode(): - onnx_export( - backbone, - (), - f=tmp_output.as_posix(), - kwargs=dummy_kwargs, - input_names=input_names, - output_names=output_names, - dynamic_axes=dynamic_axes, - do_constant_folding=do_constant_folding, - opset_version=opset_version, - dynamo=False, + dummy_kwargs, dynamic_axes, _ = generate_dummy_kwargs_and_dynamic_axes_and_shapes( + model_name, backbone ) - 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()) + + if model_name in ["sdxl-1.0", "sdxl-turbo"]: + input_names = [ + "sample", + "timestep", + "encoder_hidden_states", + "text_embeds", + "time_ids", + ] + output_names = ["latent"] + elif model_name == "sd3-medium": + input_names = [ + "hidden_states", + "encoder_hidden_states", + "pooled_projections", + "timestep", + ] + output_names = ["sample"] + elif model_name == "sd3.5-medium": + input_names = [ + "hidden_states", + "encoder_hidden_states", + "pooled_projections", + "timestep", + ] + output_names = ["out_hidden_states"] + elif model_name in ["flux-dev", "flux-schnell"]: + input_names = [ + "hidden_states", + "encoder_hidden_states", + "pooled_projections", + "timestep", + "img_ids", + "txt_ids", + ] + if model_name == "flux-dev": + input_names.append("guidance") + output_names = ["latent"] + elif model_name == "ltx-video-dev": + input_names = [ + "hidden_states", + "encoder_hidden_states", + "timestep", + "encoder_attention_mask", + "video_coords", + ] + output_names = ["latent"] + elif model_name == "wan2.2-t2v-14b": + input_names = [ + "hidden_states", + "timestep", + "encoder_hidden_states", + ] + output_names = ["latent"] + else: + raise NotImplementedError(f"Unsupported model_id: {model_name}") + + with quantizer_context, torch.inference_mode(): + onnx_export( + backbone, + (), + f=tmp_output.as_posix(), + kwargs=dummy_kwargs, + input_names=input_names, + output_names=output_names, + dynamic_axes=dynamic_axes, + do_constant_folding=True, + opset_version=20, + dynamo=False, + ) + 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"): + onnx_model = _normalize_fp8_qdq(onnx_model) + else: + flux_convert_rope_weight_type(onnx_model) + if precision == "fp4": + if strict_sdxl_fp4: + onnx_model = _process_fp4_onnx_graph( + onnx_model, + model_name, + expected_fp4_block_size, + expected_linear_count=expected_linear_count, + expected_fp8_conv_count=expected_fp8_conv_count, + ) + else: + onnx_model = NVFP4QuantExporter.process_model(onnx_model) + if strict_sdxl_fp4: + _save_onnx_atomically(onnx_model, q_output) else: - flux_convert_rope_weight_type(onnx_model) - if precision == "fp4": - onnx_model = NVFP4QuantExporter.process_model(onnx_model) - save_onnx(onnx_model, q_output) - shutil.rmtree(tmp_subfolder, ignore_errors=True) + save_onnx(onnx_model, q_output) + finally: + shutil.rmtree(tmp_subfolder, ignore_errors=True) diff --git a/examples/diffusers/quantization/quantize.py b/examples/diffusers/quantization/quantize.py index 1d71c088652..5d3902610e3 100644 --- a/examples/diffusers/quantization/quantize.py +++ b/examples/diffusers/quantization/quantize.py @@ -27,6 +27,7 @@ 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, @@ -50,12 +51,19 @@ QuantFormat, QuantizationConfig, ) -from utils import check_conv_and_mha, check_lora +from utils import ( + check_conv_and_mha, + check_lora, + validate_fp8_mha_quantizers, + validate_nvfp4_quantizers, +) import modelopt.torch.opt as mto import modelopt.torch.quantization as mtq from modelopt.torch.export import export_hf_checkpoint +_SDXL_MODEL_TYPES = (ModelType.SDXL_BASE, ModelType.SDXL_TURBO) + def setup_logging(verbose: bool = False) -> logging.Logger: """ @@ -130,7 +138,9 @@ def get_quant_config(self, n_steps: int, backbone: torch.nn.Module) -> Any: 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,19 +281,21 @@ def __init__( 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 - """ + def _has_fp8_conv_layers(self, model: torch.nn.Module) -> bool: + """Check whether the model contains an enabled FP8 convolution.""" 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 + if not isinstance(module, torch.nn.Conv1d | torch.nn.Conv2d | torch.nn.Conv3d): + continue + + quantizers = ( + getattr(module, "input_quantizer", None), + getattr(module, "weight_quantizer", None), + ) + if any( + quantizer is not None + and quantizer.is_enabled + and getattr(quantizer, "is_fp8", False) + for quantizer in quantizers ): return True return False @@ -320,6 +332,7 @@ def export_onnx( backbone: torch.nn.Module, model_type: ModelType, quant_format: QuantFormat, + fp4_block_size: int = 16, ) -> None: """ Export model to ONNX format. @@ -329,32 +342,52 @@ def export_onnx( backbone: Model backbone model_type: Type of model quant_format: Quantization format + fp4_block_size: Expected NVFP4 block size """ if not self.config.onnx_dir: return # 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 generate_fp8_scales, modelopt_export_sd, restore_fp8_scales 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() - backbone.to("cuda") - # Export to ONNX - backbone.eval() - with torch.no_grad(): - self.logger.info("Exporting to ONNX...") - modelopt_export_sd( - backbone, str(self.config.onnx_dir), model_type.value, quant_format.value + quantizer_states = [] + try: + uses_fp8_conv_workaround = quant_format == QuantFormat.FP8 or ( + quant_format == QuantFormat.FP4 and model_type in _SDXL_MODEL_TYPES ) + if uses_fp8_conv_workaround and self._has_fp8_conv_layers(backbone): + self.logger.info( + "Detected quantizing conv layers in backbone. Generating FP8 scales..." + ) + if quant_format == QuantFormat.FP4: + quantizer_states = generate_fp8_scales(backbone, conv_only=True) + else: + quantizer_states = generate_fp8_scales(backbone) + self.logger.info("Preparing models for export...") + pipe.to("cpu") + torch.cuda.empty_cache() + backbone.to("cuda") + # Export to ONNX + backbone.eval() + with torch.no_grad(): + self.logger.info("Exporting to ONNX...") + export_kwargs = ( + {"expected_fp4_block_size": fp4_block_size} + if quant_format == QuantFormat.FP4 + else {} + ) + modelopt_export_sd( + backbone, + str(self.config.onnx_dir), + model_type.value, + quant_format.value, + **export_kwargs, + ) + finally: + restore_fp8_scales(quantizer_states) self.logger.info("ONNX export completed successfully") @@ -600,6 +633,31 @@ def create_argument_parser() -> argparse.ArgumentParser: return parser +def _finalize_backbone_quantization( + backbone: torch.nn.Module, + backbone_name: str, + quant_config: QuantizationConfig, + model_type: ModelType, + restored: bool, +) -> None: + if backbone_name in ("video_decoder", "vae"): + return + + is_sdxl_fp4 = quant_config.format == QuantFormat.FP4 and model_type in _SDXL_MODEL_TYPES + if restored and not is_sdxl_fp4: + return + if is_sdxl_fp4 and restored: + validate_fp8_mha_quantizers(backbone, quant_config.quantize_mha) + check_conv_and_mha(backbone, quant_config.format == QuantFormat.FP4, quant_config.quantize_mha) + if is_sdxl_fp4: + validate_nvfp4_quantizers( + backbone, + quant_config.block_size, + quant_config.quantize_mha, + validate_sdxl_mixed_recipe=True, + ) + + def main() -> None: from diffusers.models.normalization import RMSNorm as DiffuserRMSNorm @@ -685,7 +743,8 @@ def main() -> None: export_manager = ExportManager(export_config, logger, pipeline_manager) - if export_config.restore_from and export_config.restore_from.exists(): + restored = bool(export_config.restore_from and export_config.restore_from.exists()) + if restored: export_manager.restore_checkpoint() else: @@ -709,22 +768,34 @@ def forward_loop(mod): forward_loop, backbone_name=backbone_name, ) - - # Compress model weights if requested (only for FP8/FP4) if quant_config.compress: logger.info(f"Compressing {backbone_name} weights...") 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 - ) - + _finalize_backbone_quantization( + backbone, + backbone_name, + quant_config, + model_config.model_type, + restored=False, + ) export_manager.save_checkpoint(backbone, backbone_name) + if ( + restored + and quant_config.format == QuantFormat.FP4 + and model_config.model_type in _SDXL_MODEL_TYPES + ): + for backbone_name, backbone in pipeline_manager.iter_backbones(): + _finalize_backbone_quantization( + backbone, + backbone_name, + quant_config, + model_config.model_type, + restored=True, + ) + pipeline_manager.print_quant_summary() for backbone_name, backbone in pipeline_manager.iter_backbones(): @@ -733,6 +804,7 @@ def forward_loop(mod): backbone, model_config.model_type, quant_config.format, + fp4_block_size=quant_config.block_size, ) export_manager.export_hf_ckpt(pipe, model_config) diff --git a/examples/diffusers/quantization/utils.py b/examples/diffusers/quantization/utils.py index c3cfdcd5cdd..325179accc9 100644 --- a/examples/diffusers/quantization/utils.py +++ b/examples/diffusers/quantization/utils.py @@ -25,6 +25,8 @@ from diffusers.utils import load_image import modelopt.torch.quantization as mtq +from modelopt.torch.quantization.nn import TensorQuantizer +from modelopt.torch.quantization.nn.modules.quant_linear import RealQuantLinear from modelopt.torch.quantization.plugins.diffusion.diffusers import AttentionModuleMixin USE_PEFT = True @@ -44,24 +46,37 @@ def filter_func_default(name: str) -> bool: return pattern.match(name) is not None -def check_conv_and_mha(backbone, if_fp4, quantize_mha): - for name, module in backbone.named_modules(): - if isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d)) and if_fp4: - module.weight_quantizer.disable() - module.input_quantizer.disable() +_MHA_QUANTIZER_NAMES = ( + "q_bmm_quantizer", + "k_bmm_quantizer", + "v_bmm_quantizer", + "softmax_quantizer", + "bmm2_output_quantizer", +) +_REQUIRED_MHA_QUANTIZER_NAMES = _MHA_QUANTIZER_NAMES[:4] +_SDXL_FP16_PROJECTION_NAMES = frozenset(("to_q", "to_k", "to_v")) - print(f"Disabled NVFP4 Conv layer quantization for layer {name}") - elif isinstance(module, (Attention, AttentionModuleMixin)): +def check_conv_and_mha(backbone, if_fp4, quantize_mha): + for name, module in backbone.named_modules(): + if isinstance(module, torch.nn.Conv1d | torch.nn.Conv2d) and if_fp4: + nvfp4_quantizers = [ + quantizer + for quantizer in ( + getattr(module, "weight_quantizer", None), + getattr(module, "input_quantizer", None), + ) + if isinstance(quantizer, TensorQuantizer) and quantizer.is_nvfp4_dynamic + ] + for quantizer in nvfp4_quantizers: + quantizer.disable() + if nvfp4_quantizers: + print(f"Disabled NVFP4 Conv layer quantization for layer {name}") + + elif isinstance(module, Attention | AttentionModuleMixin): head_size = int(module.inner_dim / module.heads) if not quantize_mha or head_size % 16 != 0: - for attr in ( - "q_bmm_quantizer", - "k_bmm_quantizer", - "v_bmm_quantizer", - "softmax_quantizer", - "bmm2_output_quantizer", - ): + for attr in _MHA_QUANTIZER_NAMES: if hasattr(module, attr): getattr(module, attr).disable() setattr(module, "_disable_fp8_mha", True) @@ -71,6 +86,213 @@ def check_conv_and_mha(backbone, if_fp4, quantize_mha): setattr(module, "_disable_fp8_mha", False) +def _validate_finite_positive_amax(name, quantizer): + amax = quantizer.amax + if amax is None or amax.numel() != 1 or not torch.isfinite(amax).all() or not (amax > 0).all(): + raise ValueError(f"Quantizer '{name}' must have a finite positive calibrated amax.") + + +def _validate_finite_nonnegative_amax(name, quantizer): + amax = quantizer.amax + if amax is None or amax.numel() != 1 or not torch.isfinite(amax).all() or not (amax >= 0).all(): + raise ValueError(f"Quantizer '{name}' must have a finite nonnegative calibrated amax.") + + +def _validate_calibrated_fp8_quantizer(name, quantizer): + if not quantizer.is_fp8: + raise ValueError(f"Quantizer '{name}' must use per-tensor FP8.") + _validate_finite_positive_amax(name, quantizer) + + +def validate_fp8_mha_quantizers(backbone, quantize_mha): + """Validate that the restored or finalized FP8 MHA state matches its policy.""" + for name, module in backbone.named_modules(): + if not isinstance(module, Attention | AttentionModuleMixin): + continue + + head_size = int(module.inner_dim / module.heads) + mha_enabled = quantize_mha and head_size % 16 == 0 + for attr in _MHA_QUANTIZER_NAMES: + quantizer = getattr(module, attr, None) + if quantizer is None: + if mha_enabled: + expected_state = ( + "present and enabled" + if attr in _REQUIRED_MHA_QUANTIZER_NAMES + else "present and disabled" + ) + raise ValueError( + f"FP8 MHA for attention '{name}' requires '{attr}' to be {expected_state}." + ) + continue + if not isinstance(quantizer, TensorQuantizer): + raise ValueError( + f"Attention '{name}.{attr}' must be a TensorQuantizer, got " + f"{type(quantizer).__name__}." + ) + if not mha_enabled: + if quantizer.is_enabled: + reason = "disabled by configuration" if not quantize_mha else "unsupported" + raise ValueError( + f"Attention '{name}.{attr}' must be disabled because FP8 MHA is {reason}." + ) + continue + if attr in _REQUIRED_MHA_QUANTIZER_NAMES and not quantizer.is_enabled: + raise ValueError(f"FP8 MHA for attention '{name}' requires '{attr}' to be enabled.") + if attr == "bmm2_output_quantizer" and quantizer.is_enabled: + raise ValueError( + f"FP8 MHA for attention '{name}' requires '{attr}' to be disabled." + ) + if quantizer.is_enabled: + qualified_name = f"{name}.{attr}" + if attr == "softmax_quantizer": + if not quantizer.is_fp8: + raise ValueError(f"Quantizer '{qualified_name}' must use per-tensor FP8.") + else: + _validate_calibrated_fp8_quantizer(qualified_name, quantizer) + + +def _validate_fp4_quantizer_placement(backbone, quantize_mha, allow_fp8_conv): + for module_name, module in backbone.named_modules(): + for quantizer_name, quantizer in module.named_children(): + if not isinstance(quantizer, TensorQuantizer) or not quantizer.is_enabled: + continue + qualified_name = f"{module_name}.{quantizer_name}".lstrip(".") + if quantizer.is_nvfp4_dynamic: + if not isinstance( + module, torch.nn.Linear | RealQuantLinear + ) or quantizer_name not in ( + "input_quantizer", + "weight_quantizer", + ): + raise ValueError( + f"Enabled NVFP4 quantizer '{qualified_name}' is only supported on Linear " + "input and weight quantizers." + ) + continue + + if quantizer.is_fp8: + is_conv_quantizer = ( + allow_fp8_conv + and isinstance(module, torch.nn.Conv2d) + and quantizer_name in ("input_quantizer", "weight_quantizer") + ) + is_mha_quantizer = ( + quantize_mha + and isinstance(module, Attention | AttentionModuleMixin) + and quantizer_name in _MHA_QUANTIZER_NAMES + ) + if is_conv_quantizer or is_mha_quantizer: + continue + raise ValueError( + f"Enabled FP8 quantizer '{qualified_name}' is only supported on SDXL Conv2d " + "input/weight quantizers or opt-in MHA quantizers." + ) + + raise ValueError( + f"Enabled quantizer '{qualified_name}' has an unsupported format for FP4 export." + ) + + +def validate_nvfp4_quantizers( + backbone, expected_block_size, quantize_mha, validate_sdxl_mixed_recipe=False +): + """Validate the quantizer state required by NVFP4 ONNX export.""" + enabled_linear_pairs = 0 + for name, module in backbone.named_modules(): + if not isinstance(module, torch.nn.Linear | RealQuantLinear): + continue + + input_quantizer = getattr(module, "input_quantizer", None) + weight_quantizer = getattr(module, "weight_quantizer", None) + if input_quantizer is None and weight_quantizer is None: + continue + if not isinstance(input_quantizer, TensorQuantizer) or not isinstance( + weight_quantizer, TensorQuantizer + ): + raise ValueError( + f"NVFP4 Linear '{name}' must use TensorQuantizer instances for both input and weight." + ) + + input_enabled = input_quantizer.is_enabled + weight_enabled = weight_quantizer.is_enabled + if validate_sdxl_mixed_recipe and name.rsplit(".", 1)[-1] in _SDXL_FP16_PROJECTION_NAMES: + if input_enabled or weight_enabled: + raise ValueError( + f"SDXL attention projection '{name}' must keep input and weight quantizers " + "disabled for the NVFP4 mixed recipe. Recalibrate with the current SDXL " + "FP4 recipe." + ) + continue + + if not input_enabled and not weight_enabled: + continue + if input_enabled != weight_enabled: + raise ValueError( + f"NVFP4 Linear '{name}' must enable input and weight quantizers as a pair." + ) + if not input_quantizer.is_nvfp4_dynamic or not weight_quantizer.is_nvfp4_dynamic: + raise ValueError( + f"NVFP4 Linear '{name}' must use dynamic E2M1 quantizers with FP8 block scales " + "for both input and weight." + ) + _validate_finite_nonnegative_amax(f"{name}.input_quantizer", input_quantizer) + _validate_finite_nonnegative_amax(f"{name}.weight_quantizer", weight_quantizer) + + input_block_size = input_quantizer.block_sizes.get(-1) + weight_block_size = weight_quantizer.block_sizes.get(-1) + if input_block_size != expected_block_size or weight_block_size != expected_block_size: + raise ValueError( + f"NVFP4 Linear '{name}' requires block size {expected_block_size}; got " + f"input={input_block_size}, weight={weight_block_size}." + ) + enabled_linear_pairs += 1 + + if enabled_linear_pairs == 0: + raise ValueError( + "NVFP4 quantization requires at least one enabled Linear input/weight pair." + ) + + if validate_sdxl_mixed_recipe: + enabled_conv_pairs = 0 + for name, module in backbone.named_modules(): + if not isinstance(module, torch.nn.Conv2d): + continue + + input_quantizer = getattr(module, "input_quantizer", None) + weight_quantizer = getattr(module, "weight_quantizer", None) + if input_quantizer is None and weight_quantizer is None: + continue + if not isinstance(input_quantizer, TensorQuantizer) or not isinstance( + weight_quantizer, TensorQuantizer + ): + raise ValueError( + f"SDXL FP8 Conv2d '{name}' must use TensorQuantizer instances for both input " + "and weight." + ) + + input_enabled = input_quantizer.is_enabled + weight_enabled = weight_quantizer.is_enabled + if not input_enabled and not weight_enabled: + continue + if input_enabled != weight_enabled: + raise ValueError( + f"SDXL FP8 Conv2d '{name}' must enable input and weight quantizers as a pair." + ) + _validate_calibrated_fp8_quantizer(f"{name}.input_quantizer", input_quantizer) + _validate_calibrated_fp8_quantizer(f"{name}.weight_quantizer", weight_quantizer) + enabled_conv_pairs += 1 + + if enabled_conv_pairs == 0: + raise ValueError( + "SDXL NVFP4 quantization requires at least one enabled calibrated FP8 Conv2d " + "input/weight pair." + ) + + validate_fp8_mha_quantizers(backbone, quantize_mha) + _validate_fp4_quantizer_placement(backbone, quantize_mha, validate_sdxl_mixed_recipe) + + def filter_func_ltx_video(name: str) -> bool: """Filter function specifically for LTX-Video models.""" pattern = re.compile( diff --git a/modelopt_recipes/configs/ptq/presets/diffusers/nvfp4_fp8_conv.yaml b/modelopt_recipes/configs/ptq/presets/diffusers/nvfp4_fp8_conv.yaml new file mode 100644 index 00000000000..16e2ebbebd2 --- /dev/null +++ b/modelopt_recipes/configs/ptq/presets/diffusers/nvfp4_fp8_conv.yaml @@ -0,0 +1,53 @@ +# 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. + +# Diffusers SDXL preset with dynamic NVFP4 Linears and per-tensor FP8 Conv2d layers. + +# modelopt-schema: modelopt.torch.quantization.config.QuantizeConfig +imports: + base_disable_all: configs/ptq/units/base_disable_all + fp8: configs/numerics/fp8 + nvfp4: configs/numerics/nvfp4 + +algorithm: max +quant_cfg: + - $import: base_disable_all + - parent_class: nn.Linear + quantizer_name: '*weight_quantizer' + cfg: + $import: nvfp4 + - parent_class: nn.Linear + quantizer_name: '*input_quantizer' + cfg: + $import: nvfp4 + - parent_class: nn.Linear + quantizer_name: '*to_[qkv].input_quantizer' + enable: false + - parent_class: nn.Linear + quantizer_name: '*to_[qkv].weight_quantizer' + enable: false + - parent_class: nn.Conv2d + quantizer_name: '*weight_quantizer' + cfg: + $import: fp8 + - parent_class: nn.Conv2d + quantizer_name: '*input_quantizer' + cfg: + $import: fp8 + - quantizer_name: '*output_quantizer' + enable: false + - quantizer_name: '*softmax_quantizer' + cfg: + $import: fp8 diff --git a/tests/examples/diffusers/test_diffusers.py b/tests/examples/diffusers/test_diffusers.py index 15c5eb44934..b6daa51d49c 100644 --- a/tests/examples/diffusers/test_diffusers.py +++ b/tests/examples/diffusers/test_diffusers.py @@ -128,6 +128,28 @@ def inference(self, tmp_path: Path) -> None: ), marks=minimum_sm(89), ), + pytest.param( + DiffuserModel( + name="flux-schnell", + path=FLUX_SCHNELL_PATH, + dtype="BFloat16", + format_type="fp4", + quant_algo="max", + collect_method="default", + ), + marks=minimum_sm(100), + ), + 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, @@ -141,6 +163,8 @@ def inference(self, tmp_path: Path) -> None: "flux_schnell_bf16_int8_smoothquant_3.0_min_mean", "sd3_medium_fp16_int8_smoothquant_3.0_min_mean", "sdxl_1.0_fp16_fp8_max_3.0_default", + "flux_schnell_bf16_fp4_max_3.0_default", + "sdxl_1.0_fp16_fp4_max_3.0_default", "sdxl_1.0_fp16_int8_smoothquant_3.0_min_mean", ], ) diff --git a/tests/unit/examples/test_diffusers_fp4_onnx_validation.py b/tests/unit/examples/test_diffusers_fp4_onnx_validation.py new file mode 100644 index 00000000000..3578da5096b --- /dev/null +++ b/tests/unit/examples/test_diffusers_fp4_onnx_validation.py @@ -0,0 +1,992 @@ +# 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. + +from contextlib import nullcontext +from pathlib import Path + +import numpy as np +import pytest +import torch + +onnx = pytest.importorskip("onnx") +pytest.importorskip("onnx_graphsurgeon") +pytest.importorskip("diffusers") +from onnx import TensorProto, helper, numpy_helper + +from examples.diffusers.quantization.onnx_utils import export as diffusion_export +from modelopt.onnx.export import NVFP4QuantExporter + + +class _Quantizer: + def __init__(self, num_bits, *, enabled=True, amax=448.0): + self._num_bits = num_bits + self._amax = torch.tensor(amax) + self.is_enabled = enabled + + @property + def num_bits(self): + return self._num_bits + + +def _make_raw_fp4_model(*, qdq_consumer="Conv", marker_block_size=16, marker_initializer=True): + dynamic_scale = numpy_helper.from_array( + np.array(1.0 / 2688.0, dtype=np.float32), "dynamic_scale_value" + ) + nodes = [ + helper.make_node( + "Constant", + [], + ["dynamic_quantize_scale"], + name="dynamic_quantize_scale", + value=dynamic_scale, + ), + helper.make_node( + "Constant", + [], + ["dynamic_dequantize_scale"], + name="dynamic_dequantize_scale", + value=dynamic_scale, + ), + ] + initializers = [] + inputs = [helper.make_tensor_value_info("activation", TensorProto.FLOAT16, [1, 16])] + outputs = [helper.make_tensor_value_info("linear_output", TensorProto.FLOAT16, [1, 16])] + value_info = [helper.make_tensor_value_info("fp4_weight_dq", TensorProto.FLOAT16, [16, 16])] + + if marker_initializer: + initializers.append( + numpy_helper.from_array( + np.linspace(-1.0, 1.0, 16 * 16, dtype=np.float16).reshape(16, 16), + "fp4_weight", + ) + ) + nodes.extend( + [ + helper.make_node( + "TRT_FP4DynamicQuantize", + ["activation", "dynamic_quantize_scale"], + ["activation_fp4", "activation_scale_fp8"], + name="activation_quantize", + domain="trt", + axis=-1, + block_size=marker_block_size, + scale_type=TensorProto.FLOAT8E4M3FN, + ), + helper.make_node( + "DequantizeLinear", + ["activation_scale_fp8", "dynamic_dequantize_scale"], + ["activation_scale_dq"], + name="activation_scale_dequantize", + ), + helper.make_node( + "DequantizeLinear", + ["activation_fp4", "activation_scale_dq"], + ["activation_dq"], + name="activation_dequantize", + axis=-1, + block_size=marker_block_size, + ), + helper.make_node( + "Cast", + ["activation_dq"], + ["activation_dq_fp16"], + name="activation_cast", + to=TensorProto.FLOAT16, + ), + helper.make_node( + "TRT_FP4QDQ", + ["fp4_weight"], + ["fp4_weight_dq"], + name="fp4_weight_qdq", + domain="trt", + block_size=marker_block_size, + ), + helper.make_node( + "MatMul", + ["activation_dq_fp16", "fp4_weight_dq"], + ["linear_output"], + name="fp4_matmul", + ), + ] + ) + + scale = numpy_helper.from_array(np.array(0.25, dtype=np.float16), "qdq_scale_value") + zero = numpy_helper.from_array(np.array(0, dtype=np.int8), "qdq_zero_value") + nodes.extend( + [ + helper.make_node("Constant", [], ["qdq_scale"], name="qdq_scale", value=scale), + helper.make_node("Constant", [], ["qdq_zero"], name="qdq_zero", value=zero), + helper.make_node( + "QuantizeLinear", + ["qdq_weight", "qdq_scale", "qdq_zero"], + ["quantized_weight"], + name="weight_quantize", + ), + helper.make_node( + "DequantizeLinear", + ["quantized_weight", "qdq_scale", "qdq_zero"], + ["qdq_weight_dq"], + name="weight_dequantize", + ), + helper.make_node( + "QuantizeLinear", + ["qdq_activation", "qdq_scale", "qdq_zero"], + ["quantized_activation"], + name="activation_fp8_quantize", + ), + helper.make_node( + "DequantizeLinear", + ["quantized_activation", "qdq_scale", "qdq_zero"], + ["qdq_activation_dq"], + name="activation_fp8_dequantize", + ), + ] + ) + + if qdq_consumer == "Conv": + initializers.append( + numpy_helper.from_array(np.ones((1, 1, 1, 1), dtype=np.float16), "qdq_weight") + ) + inputs.append( + helper.make_tensor_value_info("qdq_activation", TensorProto.FLOAT16, [1, 1, 2, 2]) + ) + outputs.append( + helper.make_tensor_value_info("qdq_output", TensorProto.FLOAT16, [1, 1, 2, 2]) + ) + nodes.append( + helper.make_node( + "Conv", + ["qdq_activation_dq", "qdq_weight_dq"], + ["qdq_output"], + name="fp8_conv", + ) + ) + else: + initializers.append( + numpy_helper.from_array(np.ones((4, 4), dtype=np.float16), "qdq_weight") + ) + inputs.append(helper.make_tensor_value_info("qdq_activation", TensorProto.FLOAT16, [1, 4])) + outputs.append(helper.make_tensor_value_info("qdq_output", TensorProto.FLOAT16, [1, 4])) + nodes.append( + helper.make_node( + qdq_consumer, + ["qdq_activation_dq", "qdq_weight_dq"], + ["qdq_output"], + name=f"fp8_{qdq_consumer.lower()}", + ) + ) + + graph = helper.make_graph( + nodes, + "mixed_fp4_fp8", + inputs, + outputs, + initializers, + value_info=value_info, + ) + return helper.make_model( + graph, + opset_imports=[helper.make_opsetid("", 20), helper.make_opsetid("trt", 1)], + ) + + +def _tensor_dtype(onnx_model, tensor_name): + for initializer in onnx_model.graph.initializer: + if initializer.name == tensor_name: + return initializer.data_type + node = next(node for node in onnx_model.graph.node if tensor_name in node.output) + value = next(attribute for attribute in node.attribute if attribute.name == "value").t + return value.data_type + + +def _insert_passthrough_before_weight_quantize(onnx_model): + weight = next( + initializer + for initializer in onnx_model.graph.initializer + if initializer.name == "qdq_weight" + ) + weight.name = "qdq_weight_source" + quantize_node = next(node for node in onnx_model.graph.node if node.name == "weight_quantize") + onnx_model.graph.node.insert( + 0, + helper.make_node("Identity", ["qdq_weight_source"], ["qdq_weight"], name="weight_identity"), + ) + assert quantize_node.input[0] == "qdq_weight" + + +def _insert_weight_cast_before_quantize(onnx_model, dtype): + quantize_node = next(node for node in onnx_model.graph.node if node.name == "weight_quantize") + quantize_node.input[0] = "qdq_weight_cast" + onnx_model.graph.node.insert( + 0, + helper.make_node( + "Cast", + ["qdq_weight"], + ["qdq_weight_cast"], + name="weight_cast", + to=dtype, + ), + ) + + +def _make_quantized_backbone(linear_count, conv_count): + backbone = torch.nn.Module() + for index in range(linear_count): + linear = torch.nn.Linear(16, 16, bias=False) + linear.input_quantizer = _Quantizer((2, 1)) + linear.weight_quantizer = _Quantizer((2, 1)) + backbone.add_module(f"linear_{index}", linear) + for index in range(conv_count): + conv = torch.nn.Conv2d(1, 1, 1, bias=False) + conv.input_quantizer = _Quantizer((4, 3)) + conv.weight_quantizer = _Quantizer((4, 3)) + backbone.add_module(f"conv_{index}", conv) + return backbone + + +def _make_external_data_model(fill_value): + weight = numpy_helper.from_array( + np.full((1024,), fill_value, dtype=np.float32), "external_weight" + ) + graph = helper.make_graph( + [helper.make_node("Identity", ["external_weight"], ["output"])], + "external_data", + [], + [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1024])], + [weight], + ) + return helper.make_model(graph) + + +def test_fp8_scale_workaround_can_target_only_enabled_conv_quantizers(): + conv = torch.nn.Conv2d(1, 1, 1) + conv.input_quantizer = _Quantizer((4, 3)) + conv.weight_quantizer = _Quantizer((4, 3), enabled=False) + linear = torch.nn.Linear(1, 1) + linear.input_quantizer = _Quantizer((2, 1)) + linear.weight_quantizer = _Quantizer((4, 3)) + model = torch.nn.Sequential(conv, linear) + + diffusion_export.generate_fp8_scales(model, conv_only=True) + + assert conv.input_quantizer.num_bits == 8 + assert conv.input_quantizer._amax == 127.0 + assert conv.weight_quantizer.num_bits == (4, 3) + assert linear.input_quantizer.num_bits == (2, 1) + assert linear.weight_quantizer.num_bits == (4, 3) + + diffusion_export.generate_fp8_scales(model) + + assert linear.weight_quantizer.num_bits == 8 + assert linear.weight_quantizer._amax == 127.0 + + +def test_mixed_sdxl_graph_preserves_fp8_conv_and_lowers_exact_nvfp4_topology(): + raw_model = _make_raw_fp4_model() + + converted_model = diffusion_export._process_fp4_onnx_graph(raw_model, "sdxl-1.0") + + assert not any(node.op_type == "TRT_FP4QDQ" for node in converted_model.graph.node) + assert ( + sum( + initializer.data_type == TensorProto.FLOAT4E2M1 + for initializer in converted_model.graph.initializer + ) + == 1 + ) + assert _tensor_dtype(converted_model, "qdq_zero") == TensorProto.FLOAT8E4M3FN + assert any( + node.op_type == "Conv" and node.name == "fp8_conv" for node in converted_model.graph.node + ) + assert next(opset.version for opset in converted_model.opset_import if not opset.domain) >= 23 + onnx.checker.check_model(converted_model) + + +@pytest.mark.parametrize( + ("linear_count", "conv_count", "error"), + [ + (2, 1, "found 1 TRT_FP4QDQ weight markers, expected 2 enabled Linear pairs"), + ( + 1, + 2, + "found 1 initializer-backed FP8 Conv weight Q/DQ pairs, expected 2 enabled Conv2d pairs", + ), + ], +) +def test_raw_sdxl_graph_counts_match_enabled_quantizer_pairs(linear_count, conv_count, error): + expected_linear_count, expected_conv_count = diffusion_export._get_sdxl_fp4_expected_counts( + _make_quantized_backbone(linear_count, conv_count) + ) + + with pytest.raises(ValueError, match=error): + diffusion_export._process_fp4_onnx_graph( + _make_raw_fp4_model(), + "sdxl-1.0", + expected_linear_count=expected_linear_count, + expected_fp8_conv_count=expected_conv_count, + ) + + +def test_final_sdxl_graph_count_matches_enabled_conv_pairs(): + converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") + _, expected_conv_count = diffusion_export._get_sdxl_fp4_expected_counts( + _make_quantized_backbone(1, 2) + ) + + with pytest.raises( + ValueError, + match="found 1 initializer-backed FP8 Conv weight Q/DQ pairs, expected 2 enabled Conv2d pairs", + ): + diffusion_export._validate_final_fp4_graph( + converted_model, + expected_weight_count=1, + allow_fp8_conv=True, + expected_fp8_conv_count=expected_conv_count, + ) + + +@pytest.mark.parametrize("pre_q_passthrough", [False, True]) +def test_non_sdxl_fp4_graph_rejects_static_qdq_conv_weights(pre_q_passthrough): + model = _make_raw_fp4_model() + if pre_q_passthrough: + _insert_passthrough_before_weight_quantize(model) + + with pytest.raises(ValueError, match="disallowed initializer-backed Q/DQ"): + diffusion_export._process_fp4_onnx_graph(model, "flux-dev") + + +def test_sdxl_fp4_graph_requires_static_fp8_conv_weight_qdq(): + model = _make_raw_fp4_model() + quantize_node = next(node for node in model.graph.node if node.name == "weight_quantize") + quantize_node.input[0] = "qdq_activation" + + with pytest.raises(ValueError, match="no initializer-backed FP8 Conv weight Q/DQ"): + diffusion_export._process_fp4_onnx_graph(model, "sdxl-1.0") + + +@pytest.mark.parametrize( + ("marker_initializer", "marker_block_size", "error"), + [ + (False, 16, "not backed by a weight initializer"), + (True, 32, "block_size=32, expected 16"), + ], +) +def test_raw_fp4_validation_rejects_invalid_markers(marker_initializer, marker_block_size, error): + model = _make_raw_fp4_model( + marker_initializer=marker_initializer, marker_block_size=marker_block_size + ) + + with pytest.raises(ValueError, match=error): + diffusion_export._validate_raw_fp4_graph(model) + + +@pytest.mark.parametrize("consumer_op", ["Add", "Gemm", "MatMul"]) +def test_raw_fp4_validation_rejects_static_qdq_non_conv_weights(consumer_op): + model = _make_raw_fp4_model(qdq_consumer=consumer_op) + + with pytest.raises(ValueError, match="disallowed initializer-backed Q/DQ"): + diffusion_export._validate_raw_fp4_graph(model, allow_fp8_conv=True) + + +@pytest.mark.parametrize("consumer_op", ["Add", None]) +def test_raw_fp4_validation_rejects_marker_without_weight_consumer(consumer_op): + model = _make_raw_fp4_model() + matmul = next(node for node in model.graph.node if node.name == "fp4_matmul") + if consumer_op is None: + model.graph.node.remove(matmul) + else: + matmul.op_type = consumer_op + + with pytest.raises(ValueError, match="does not reach a Gemm/MatMul weight input"): + diffusion_export._validate_raw_fp4_graph(model, allow_fp8_conv=True) + + +@pytest.mark.parametrize( + ("corruption", "error"), + [ + ("block-scale-dtype", "FLOAT8E4M3FN block-scale initializer"), + ("axis", "does not use axis=-1"), + ("fp8-zero-dtype", "zero point is not FLOAT8E4M3FN"), + ("weight-consumer", "does not reach a Gemm/MatMul weight input"), + ], +) +def test_final_fp4_validation_rejects_invalid_double_dq(corruption, error): + raw_model = _make_raw_fp4_model() + expected_weight_count = diffusion_export._validate_raw_fp4_graph(raw_model, allow_fp8_conv=True) + normalized_model = diffusion_export._normalize_fp8_qdq(raw_model) + converted_model = NVFP4QuantExporter.process_model(normalized_model) + fp4_weight = next( + initializer.name + for initializer in converted_model.graph.initializer + if initializer.data_type == TensorProto.FLOAT4E2M1 + ) + weight_dq = next( + node + for node in converted_model.graph.node + if node.op_type == "DequantizeLinear" and node.input[0] == fp4_weight + ) + scale_dq = next( + node for node in converted_model.graph.node if weight_dq.input[1] in node.output + ) + if corruption == "block-scale-dtype": + fp8_scale = next( + initializer + for initializer in converted_model.graph.initializer + if initializer.name == scale_dq.input[0] + ) + fp8_scale.data_type = TensorProto.FLOAT16 + elif corruption == "axis": + axis = next(attribute for attribute in weight_dq.attribute if attribute.name == "axis") + axis.i = 0 + elif corruption == "fp8-zero-dtype": + zero_node = next(node for node in converted_model.graph.node if node.name == "qdq_zero") + value = next(attribute for attribute in zero_node.attribute if attribute.name == "value") + value.t.data_type = TensorProto.INT8 + else: + matmul = next(node for node in converted_model.graph.node if node.name == "fp4_matmul") + matmul.op_type = "Add" + + with pytest.raises(ValueError, match=error): + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count, allow_fp8_conv=True + ) + + +@pytest.mark.parametrize( + ("corruption", "error"), + [ + ("extra-input", "must be a two-input DequantizeLinear"), + ("axis", "must not use axis or block_size"), + ("block-size", "must not use axis or block_size"), + ("fp8-scale-fanout", "must be consumed only by"), + ("global-scale-dtype", "does not use a FLOAT global-scale initializer"), + ("global-scale-zero", "global scale must be a finite positive scalar constant"), + ("scale-output-fanout", "output must be consumed only by"), + ], +) +def test_final_fp4_validation_rejects_invalid_weight_scale_dq(corruption, error): + converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") + fp4_weight = next( + initializer.name + for initializer in converted_model.graph.initializer + if initializer.data_type == TensorProto.FLOAT4E2M1 + ) + weight_dq = next( + node + for node in converted_model.graph.node + if node.op_type == "DequantizeLinear" and node.input[0] == fp4_weight + ) + scale_dq = next( + node for node in converted_model.graph.node if weight_dq.input[1] in node.output + ) + if corruption == "extra-input": + scale_dq.input.append(scale_dq.input[1]) + elif corruption in {"axis", "block-size"}: + scale_dq.attribute.append(helper.make_attribute(corruption.replace("-", "_"), 0)) + elif corruption == "fp8-scale-fanout": + converted_model.graph.node.append( + helper.make_node( + "Identity", + [scale_dq.input[0]], + ["extra_block_scale_use"], + name="extra_block_scale_use", + ) + ) + elif corruption == "global-scale-dtype": + global_scale = next( + initializer + for initializer in converted_model.graph.initializer + if initializer.name == scale_dq.input[1] + ) + global_scale.data_type = TensorProto.FLOAT16 + elif corruption == "global-scale-zero": + global_scale = next( + initializer + for initializer in converted_model.graph.initializer + if initializer.name == scale_dq.input[1] + ) + global_scale.CopyFrom( + numpy_helper.from_array(np.array(0.0, dtype=np.float32), global_scale.name) + ) + else: + converted_model.graph.node.append( + helper.make_node( + "Identity", + [scale_dq.output[0]], + ["extra_weight_scale_use"], + name="extra_weight_scale_use", + ) + ) + + with pytest.raises(ValueError, match=error): + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count=1, allow_fp8_conv=True + ) + + +def test_final_fp4_validation_rejects_extra_float4_weight_consumer(): + raw_model = _make_raw_fp4_model() + expected_weight_count = diffusion_export._validate_raw_fp4_graph(raw_model, allow_fp8_conv=True) + converted_model = diffusion_export._process_fp4_onnx_graph(raw_model, "sdxl-1.0") + fp4_weight = next( + initializer.name + for initializer in converted_model.graph.initializer + if initializer.data_type == TensorProto.FLOAT4E2M1 + ) + converted_model.graph.node.append( + helper.make_node("Identity", [fp4_weight], ["extra_weight_use"], name="extra_weight_use") + ) + + with pytest.raises(ValueError, match="must feed exactly one weight DequantizeLinear"): + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count, allow_fp8_conv=True + ) + + +@pytest.mark.parametrize( + ("corruption", "error"), + [ + ("mismatched-scale", "do not share scale and zero point"), + ("axis", "without an axis"), + ("nonpositive-scale", "finite positive scalar constant"), + ("wrong-scale-dtype", "FP8 scale must use a floating-point dtype"), + ("nonzero-zero", "zero point must be a scalar zero"), + ("missing-activation", "has no FP8 activation Q/DQ"), + ("activation-fanout", "has non-Conv activation consumers"), + ("weight-fanout", "weight DQ must feed exactly one FP8 Conv input 1"), + ], +) +def test_final_fp4_validation_rejects_invalid_fp8_conv_qdq(corruption, error): + converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") + weight_quantize = next( + node for node in converted_model.graph.node if node.name == "weight_quantize" + ) + weight_dequantize = next( + node for node in converted_model.graph.node if node.name == "weight_dequantize" + ) + if corruption == "mismatched-scale": + converted_model.graph.initializer.append( + numpy_helper.from_array(np.array(0.5, dtype=np.float16), "other_fp8_scale") + ) + weight_dequantize.input[1] = "other_fp8_scale" + elif corruption == "axis": + weight_quantize.attribute.append(helper.make_attribute("axis", 0)) + elif corruption == "nonpositive-scale": + scale_node = next(node for node in converted_model.graph.node if node.name == "qdq_scale") + scale_value = next( + attribute for attribute in scale_node.attribute if attribute.name == "value" + ) + scale_value.t.CopyFrom( + numpy_helper.from_array(np.array(0.0, dtype=np.float16), "qdq_scale_value") + ) + elif corruption == "wrong-scale-dtype": + scale_node = next(node for node in converted_model.graph.node if node.name == "qdq_scale") + scale_value = next( + attribute for attribute in scale_node.attribute if attribute.name == "value" + ) + scale_value.t.CopyFrom( + numpy_helper.from_array(np.array(1, dtype=np.int32), "qdq_scale_value") + ) + elif corruption == "nonzero-zero": + zero_node = next(node for node in converted_model.graph.node if node.name == "qdq_zero") + zero_value = next( + attribute for attribute in zero_node.attribute if attribute.name == "value" + ) + zero_value.t.raw_data = b"\x38" + elif corruption == "missing-activation": + conv = next(node for node in converted_model.graph.node if node.name == "fp8_conv") + conv.input[0] = "qdq_activation" + elif corruption == "activation-fanout": + converted_model.graph.node.append( + helper.make_node( + "Identity", + ["qdq_activation_dq"], + ["extra_fp8_activation_use"], + name="extra_fp8_activation_use", + ) + ) + else: + converted_model.graph.node.append( + helper.make_node( + "Conv", + ["qdq_activation_dq", "qdq_weight_dq"], + ["extra_fp8_conv_output"], + name="extra_fp8_conv", + ) + ) + + with pytest.raises(ValueError, match=error): + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count=1, allow_fp8_conv=True + ) + + +@pytest.mark.parametrize( + ("corruption", "error"), + [ + ("block-size", "activation_quantize does not use block_size=16"), + ("mismatched-scale", "quantize and dequantize global scales do not match"), + ("wrong-domain", "must use the trt domain"), + ("static-input", "input 0 must be a dynamic activation"), + ("fanout", "must feed exactly one Gemm/MatMul activation input"), + ("bypass", "dynamic NVFP4 activation paths do not match FLOAT4 weight consumers"), + ("extra-scale-consumer", "FP8 scale output must feed exactly one DequantizeLinear"), + ], +) +def test_final_fp4_validation_rejects_invalid_dynamic_activation(corruption, error): + converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") + dynamic_quantize = next( + node for node in converted_model.graph.node if node.name == "activation_quantize" + ) + if corruption == "block-size": + block_size = next( + attribute for attribute in dynamic_quantize.attribute if attribute.name == "block_size" + ) + block_size.i = 32 + elif corruption == "mismatched-scale": + scale_node = next( + node for node in converted_model.graph.node if node.name == "dynamic_dequantize_scale" + ) + scale_value = next( + attribute for attribute in scale_node.attribute if attribute.name == "value" + ) + scale_value.t.CopyFrom( + numpy_helper.from_array(np.array(2.0 / 2688.0, dtype=np.float32), "scale_value") + ) + elif corruption == "wrong-domain": + dynamic_quantize.domain = "other" + elif corruption == "static-input": + dynamic_quantize.input[0] = "dynamic_quantize_scale" + elif corruption == "fanout": + converted_model.graph.node.append( + helper.make_node( + "MatMul", + ["activation_dq_fp16", "fp4_weight_dq"], + ["extra_linear_output"], + name="extra_fp4_matmul", + ) + ) + elif corruption == "bypass": + matmul = next(node for node in converted_model.graph.node if node.name == "fp4_matmul") + matmul.input[0] = "activation" + else: + converted_model.graph.node.append( + helper.make_node( + "Identity", + [dynamic_quantize.output[1]], + ["extra_dynamic_scale_use"], + name="extra_dynamic_scale_use", + ) + ) + + with pytest.raises(ValueError, match=error): + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count=1, allow_fp8_conv=True + ) + + +def test_final_fp4_validation_accepts_bfloat16_fp8_conv_scales(): + converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") + qdq_weight = next( + initializer + for initializer in converted_model.graph.initializer + if initializer.name == "qdq_weight" + ) + qdq_weight.data_type = TensorProto.BFLOAT16 + qdq_activation = next( + value for value in converted_model.graph.input if value.name == "qdq_activation" + ) + qdq_activation.type.tensor_type.elem_type = TensorProto.BFLOAT16 + scale_node = next(node for node in converted_model.graph.node if node.name == "qdq_scale") + scale_value = next(attribute for attribute in scale_node.attribute if attribute.name == "value") + scale_value.t.data_type = TensorProto.BFLOAT16 + + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count=1, allow_fp8_conv=True + ) + + +def test_final_fp4_validation_uses_cast_output_dtype_for_fp8_conv_scale(): + raw_model = _make_raw_fp4_model() + _insert_weight_cast_before_quantize(raw_model, TensorProto.BFLOAT16) + activation = next(value for value in raw_model.graph.input if value.name == "qdq_activation") + activation.type.tensor_type.elem_type = TensorProto.BFLOAT16 + scale_node = next(node for node in raw_model.graph.node if node.name == "qdq_scale") + scale_value = next(attribute for attribute in scale_node.attribute if attribute.name == "value") + scale_value.t.data_type = TensorProto.BFLOAT16 + + converted_model = diffusion_export._process_fp4_onnx_graph(raw_model, "sdxl-1.0") + + diffusion_export._validate_final_fp4_graph( + converted_model, expected_weight_count=1, allow_fp8_conv=True + ) + + +def test_non_sdxl_fp4_export_uses_permissive_generic_lowering(monkeypatch, tmp_path): + raw_model = _make_raw_fp4_model() + generic_calls = [] + saved_models = [] + + def fake_onnx_export(*args, f, **kwargs): + del args, kwargs + onnx.save(raw_model, f) + + def generic_process(cls, onnx_model): + del cls + generic_calls.append(onnx_model) + return onnx_model + + monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) + monkeypatch.setattr( + diffusion_export, + "generate_dummy_kwargs_and_dynamic_axes_and_shapes", + lambda *args: ({}, {}, {}), + ) + monkeypatch.setattr( + diffusion_export, + "_process_fp4_onnx_graph", + lambda *args: pytest.fail("strict SDXL processing must not run for Flux"), + ) + monkeypatch.setattr(NVFP4QuantExporter, "process_model", classmethod(generic_process)) + monkeypatch.setattr( + diffusion_export, + "save_onnx", + lambda onnx_model, output: saved_models.append((onnx_model, output)), + ) + monkeypatch.setattr( + diffusion_export, + "_save_onnx_atomically", + lambda *args: pytest.fail("atomic checked publication must be SDXL FP4-only"), + ) + + diffusion_export.modelopt_export_sd(torch.nn.Identity(), tmp_path, "flux-dev", "fp4") + + assert len(generic_calls) == 1 + assert len(saved_models) == 1 + assert ( + next(opset.version for opset in saved_models[0][0].opset_import if not opset.domain) == 20 + ) + + +def test_sdxl_fp4_export_threads_enabled_quantizer_counts(monkeypatch, tmp_path): + raw_model = _make_raw_fp4_model() + processed_counts = [] + saved_models = [] + backbone = _make_quantized_backbone(2, 3) + + def fake_onnx_export(*args, f, **kwargs): + del args, kwargs + onnx.save(raw_model, f) + + def fake_process(onnx_model, model_name, block_size, **kwargs): + del model_name, block_size + processed_counts.append(kwargs) + return onnx_model + + monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) + monkeypatch.setattr( + diffusion_export, + "generate_dummy_kwargs_and_dynamic_axes_and_shapes", + lambda *args: ({}, {}, {}), + ) + monkeypatch.setattr( + diffusion_export, "configure_linear_module_onnx_quantizers", lambda _: nullcontext() + ) + monkeypatch.setattr(diffusion_export, "_process_fp4_onnx_graph", fake_process) + monkeypatch.setattr( + diffusion_export, + "_save_onnx_atomically", + lambda onnx_model, output: saved_models.append((onnx_model, output)), + ) + monkeypatch.setattr( + diffusion_export, + "save_onnx", + lambda *args: pytest.fail("SDXL FP4 must use atomic checked publication"), + ) + + diffusion_export.modelopt_export_sd(backbone, tmp_path, "sdxl-1.0", "fp4") + + assert processed_counts == [{"expected_linear_count": 2, "expected_fp8_conv_count": 3}] + assert len(saved_models) == 1 + + +@pytest.mark.parametrize("precision", ["fp16", "fp8"]) +def test_non_fp4_exports_keep_legacy_save_path(monkeypatch, tmp_path, precision): + raw_model = _make_raw_fp4_model() + saved_models = [] + + def fake_onnx_export(*args, f, **kwargs): + del args, kwargs + onnx.save(raw_model, f) + + monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) + monkeypatch.setattr( + diffusion_export, + "generate_dummy_kwargs_and_dynamic_axes_and_shapes", + lambda *args: ({}, {}, {}), + ) + monkeypatch.setattr(diffusion_export, "_normalize_fp8_qdq", lambda model: model) + monkeypatch.setattr( + diffusion_export, + "save_onnx", + lambda onnx_model, output: saved_models.append((onnx_model, output)), + ) + monkeypatch.setattr( + diffusion_export, + "_save_onnx_atomically", + lambda *args: pytest.fail("atomic checked publication must be SDXL FP4-only"), + ) + + diffusion_export.modelopt_export_sd(torch.nn.Identity(), tmp_path, "sdxl-1.0", precision) + + assert len(saved_models) == 1 + assert ( + next(opset.version for opset in saved_models[0][0].opset_import if not opset.domain) == 20 + ) + + +def test_invalid_fp4_export_preserves_existing_output_and_cleans_raw_temp(monkeypatch, tmp_path): + output_dir = tmp_path / "onnx" + output_dir.mkdir() + output = output_dir / "model.onnx" + output_data = output_dir / "model.onnx_data" + output.write_bytes(b"previous-model") + output_data.write_bytes(b"previous-data") + raw_dirs = [] + + invalid_model = helper.make_model( + helper.make_graph( + [helper.make_node("Identity", ["input"], ["output"])], + "invalid_fp4", + [helper.make_tensor_value_info("input", TensorProto.FLOAT, [1])], + [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1])], + ) + ) + + def fake_onnx_export(*args, f, **kwargs): + del args, kwargs + raw_dirs.append(Path(f).parent) + onnx.save(invalid_model, f) + + monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) + monkeypatch.setattr( + diffusion_export, + "generate_dummy_kwargs_and_dynamic_axes_and_shapes", + lambda *args: ({}, {}, {}), + ) + + with pytest.raises(ValueError, match="no TRT_FP4QDQ weight markers"): + diffusion_export.modelopt_export_sd(torch.nn.Identity(), output_dir, "sdxl-1.0", "fp4") + + assert output.read_bytes() == b"previous-model" + assert output_data.read_bytes() == b"previous-data" + assert len(raw_dirs) == 1 + assert not raw_dirs[0].exists() + + +def test_staged_save_failure_preserves_existing_output(monkeypatch, tmp_path): + output = tmp_path / "model.onnx" + output_data = tmp_path / "model.onnx_data" + output.write_bytes(b"previous-model") + output_data.write_bytes(b"previous-data") + + def fail_after_partial_save(onnx_model, staged_output, external_data_name=None): + del onnx_model + staged_output.write_bytes(b"partial-model") + (staged_output.parent / external_data_name).write_bytes(b"partial-data") + raise RuntimeError("save failed") + + monkeypatch.setattr(diffusion_export, "save_onnx", fail_after_partial_save) + + with pytest.raises(RuntimeError, match="save failed"): + diffusion_export._save_onnx_atomically(object(), output) + + assert output.read_bytes() == b"previous-model" + assert output_data.read_bytes() == b"previous-data" + assert not list(tmp_path.glob(".modelopt-export-*")) + + +def test_staged_checker_failure_preserves_existing_output(monkeypatch, tmp_path): + output = tmp_path / "model.onnx" + output_data = tmp_path / "model.onnx_data" + output.write_bytes(b"previous-model") + output_data.write_bytes(b"previous-data") + checked_paths = [] + + def fail_check(staged_output): + checked_paths.append(Path(staged_output)) + raise RuntimeError("checker failed") + + monkeypatch.setattr(diffusion_export.onnx.checker, "check_model", fail_check) + + with pytest.raises(RuntimeError, match="checker failed"): + diffusion_export._save_onnx_atomically(_make_external_data_model(2.0), output) + + assert len(checked_paths) == 1 + assert checked_paths[0].parent.name.startswith(".modelopt-export-") + assert output.read_bytes() == b"previous-model" + assert output_data.read_bytes() == b"previous-data" + assert not list(tmp_path.glob("model.onnx_data.*")) + assert not list(tmp_path.glob(".modelopt-export-*")) + + +def test_publication_failure_restores_previous_output_and_removes_new_data(monkeypatch, tmp_path): + output = tmp_path / "model.onnx" + output_data = tmp_path / "model.onnx_data" + output.write_bytes(b"previous-model") + output_data.write_bytes(b"previous-data") + + real_replace = diffusion_export.os.replace + + def fail_model_publish(source, destination): + source = Path(source) + destination = Path(destination) + if source.name == output.name and source.parent.name.startswith(".modelopt-export-"): + raise RuntimeError("publish failed") + real_replace(source, destination) + + monkeypatch.setattr(diffusion_export.os, "replace", fail_model_publish) + + with pytest.raises(RuntimeError, match="publish failed"): + diffusion_export._save_onnx_atomically(_make_external_data_model(2.0), output) + + assert output.read_bytes() == b"previous-model" + assert output_data.read_bytes() == b"previous-data" + assert not list(tmp_path.glob("model.onnx_data.*")) + assert not list(tmp_path.glob(".modelopt-export-*")) + + +def test_atomic_save_publishes_versioned_data_then_removes_old_data(tmp_path): + output = tmp_path / "model.onnx" + old_data = tmp_path / "model.onnx_data" + diffusion_export.save_onnx(_make_external_data_model(1.0), output) + assert old_data.exists() + + diffusion_export._save_onnx_atomically(_make_external_data_model(2.0), output) + + stored_model = onnx.load(str(output), load_external_data=False) + locations = { + entry.value + for initializer in stored_model.graph.initializer + for entry in initializer.external_data + if entry.key == "location" + } + assert len(locations) == 1 + external_data_name = locations.pop() + assert external_data_name.startswith("model.onnx_data.") + assert (tmp_path / external_data_name).exists() + assert not old_data.exists() + assert np.all(numpy_helper.to_array(onnx.load(str(output)).graph.initializer[0]) == 2.0) + assert not list(tmp_path.glob(".modelopt-export-*")) diff --git a/tests/unit/examples/test_diffusers_fp4_validation.py b/tests/unit/examples/test_diffusers_fp4_validation.py new file mode 100644 index 00000000000..5bfc4347163 --- /dev/null +++ b/tests/unit/examples/test_diffusers_fp4_validation.py @@ -0,0 +1,652 @@ +# 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 logging +import sys +from pathlib import Path + +import pytest +import torch +from diffusers.models.attention_processor import Attention +from torch import nn + +import modelopt.torch.quantization as mtq +from modelopt.torch.quantization.config import QuantizerAttributeConfig +from modelopt.torch.quantization.nn import TensorQuantizer + +_QUANTIZATION_EXAMPLE = ( + Path(__file__).resolve().parents[3] / "examples" / "diffusers" / "quantization" +) +sys.path.insert(0, str(_QUANTIZATION_EXAMPLE)) + +import quantize as quantize_module +from models_utils import ModelType +from quantize import ExportManager, Quantizer, _finalize_backbone_quantization +from quantize_config import ExportConfig, ModelConfig, QuantFormat, QuantizationConfig +from utils import check_conv_and_mha, validate_nvfp4_quantizers + +from examples.diffusers.quantization import onnx_utils as diffusion_onnx_utils +from examples.diffusers.quantization.onnx_utils import export as diffusion_export + +_MHA_QUANTIZER_NAMES = ( + "q_bmm_quantizer", + "k_bmm_quantizer", + "v_bmm_quantizer", + "softmax_quantizer", + "bmm2_output_quantizer", +) + + +def _nvfp4_quantizer(block_size=16, enabled=True, amax=1.0): + quantizer = TensorQuantizer( + QuantizerAttributeConfig( + num_bits=(2, 1), + block_sizes={-1: block_size, "type": "dynamic", "scale_bits": (4, 3)}, + ) + ) + if amax is not None: + quantizer.amax = torch.tensor(amax) + if not enabled: + quantizer.disable() + return quantizer + + +def _fp8_quantizer(enabled=True, amax=1.0): + quantizer = TensorQuantizer(QuantizerAttributeConfig(num_bits=(4, 3), axis=None)) + if amax is not None: + quantizer.amax = torch.tensor(amax) + if not enabled: + quantizer.disable() + return quantizer + + +class _QuantizedLinear(nn.Linear): + def __init__(self, input_quantizer=None, weight_quantizer=None): + super().__init__(16, 16, bias=False) + self.input_quantizer = input_quantizer or _nvfp4_quantizer() + self.weight_quantizer = weight_quantizer or _nvfp4_quantizer() + + +class _Backbone(nn.Module): + def __init__(self, linear=None, conv=None, attention=None): + super().__init__() + self.linear = linear or _QuantizedLinear() + if conv is not None: + self.conv = conv + if attention is not None: + self.attention = attention + + +def _quantized_conv(conv_cls, quantizer_factory): + conv = conv_cls(4, 4, kernel_size=1, bias=False) + conv.input_quantizer = quantizer_factory() + conv.weight_quantizer = quantizer_factory() + return conv + + +def _attention(head_size=16, enabled=True, softmax_amax=1.0): + attention = Attention(query_dim=head_size, heads=1, dim_head=head_size) + for name in _MHA_QUANTIZER_NAMES: + quantizer_enabled = enabled and name != "bmm2_output_quantizer" + amax = softmax_amax if name == "softmax_quantizer" else 1.0 + setattr(attention, name, _fp8_quantizer(enabled=quantizer_enabled, amax=amax)) + return attention + + +@pytest.mark.parametrize("model_type", [ModelType.SDXL_BASE, ModelType.SDXL_TURBO]) +def test_sdxl_fp4_uses_nvfp4_linear_fp8_conv_recipe(model_type): + config = Quantizer( + QuantizationConfig(format=QuantFormat.FP4, block_size=32), + ModelConfig(model_type=model_type), + logging.getLogger(__name__), + ).get_quant_config(n_steps=1, backbone=nn.Module()) + + linear_entries = { + entry["quantizer_name"]: entry["cfg"] + for entry in config["quant_cfg"] + if entry.get("parent_class") == "nn.Linear" + and entry.get("quantizer_name") in ("*input_quantizer", "*weight_quantizer") + } + conv_entries = { + entry["quantizer_name"]: entry["cfg"] + for entry in config["quant_cfg"] + if entry.get("parent_class") == "nn.Conv2d" + } + + assert set(linear_entries) == {"*input_quantizer", "*weight_quantizer"} + assert set(conv_entries) == {"*input_quantizer", "*weight_quantizer"} + for cfg in linear_entries.values(): + assert cfg["num_bits"] == (2, 1) + assert cfg["block_sizes"][-1] == 32 + assert cfg["block_sizes"]["scale_bits"] == (4, 3) + for cfg in conv_entries.values(): + assert cfg["num_bits"] == (4, 3) + assert cfg["axis"] is None + assert "block_sizes" not in cfg + + +def test_sdxl_fp4_recipe_applies_only_to_linear_and_conv(): + model = nn.Module() + model.linear = nn.Linear(16, 16, bias=False) + model.attn = nn.Module() + model.attn.to_q = nn.Linear(16, 16, bias=False) + model.attn.to_k = nn.Linear(16, 16, bias=False) + model.attn.to_v = nn.Linear(16, 16, bias=False) + model.conv = nn.Conv2d(4, 4, kernel_size=1, bias=False) + model.norm = nn.LayerNorm(16) + config = Quantizer( + QuantizationConfig(format=QuantFormat.FP4), + ModelConfig(model_type=ModelType.SDXL_BASE), + 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"]) + + assert model.linear.input_quantizer.is_enabled + assert model.linear.input_quantizer.is_nvfp4_dynamic + assert model.linear.weight_quantizer.is_enabled + assert model.linear.weight_quantizer.is_nvfp4_dynamic + 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 + assert model.conv.input_quantizer.is_enabled + assert model.conv.input_quantizer.is_fp8 + assert model.conv.weight_quantizer.is_enabled + assert model.conv.weight_quantizer.is_fp8 + assert not model.norm.input_quantizer.is_enabled + + +def test_fp4_finalization_preserves_calibrated_fp8_conv(): + conv = _quantized_conv(nn.Conv2d, lambda: _fp8_quantizer(amax=123.0)) + backbone = _Backbone(conv=conv) + before = { + name: quantizer.amax.clone() + for name, quantizer in ( + ("input", conv.input_quantizer), + ("weight", conv.weight_quantizer), + ) + } + + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4), + ModelType.SDXL_BASE, + restored=False, + ) + + for name, quantizer in ( + ("input", conv.input_quantizer), + ("weight", conv.weight_quantizer), + ): + assert quantizer.is_enabled + assert quantizer.is_fp8 + assert torch.equal(quantizer.amax, before[name]) + + +@pytest.mark.parametrize("export_fails", [False, True]) +def test_onnx_fp8_scale_workaround_restores_state_before_hf_export( + monkeypatch, tmp_path, export_fails +): + conv = _quantized_conv(nn.Conv2d, lambda: _fp8_quantizer(amax=448.0)) + backbone = _Backbone(conv=conv) + original_state = { + name: (quantizer.num_bits, quantizer.amax) + for name, quantizer in ( + ("input", conv.input_quantizer), + ("weight", conv.weight_quantizer), + ) + } + manager = ExportManager( + ExportConfig(onnx_dir=tmp_path / "onnx", hf_ckpt_dir=tmp_path / "hf"), + logging.getLogger(__name__), + pipeline_manager=None, + ) + + class _Pipeline: + def to(self, *args, **kwargs): + return self + + def fake_onnx_export(*args, **kwargs): + del args, kwargs + assert conv.input_quantizer.num_bits == 8 + assert conv.weight_quantizer.num_bits == 8 + assert conv.input_quantizer.amax == 127.0 + assert conv.weight_quantizer.amax == 127.0 + if export_fails: + raise RuntimeError("export failed") + + def fake_hf_export(*args, **kwargs): + del args, kwargs + for name, quantizer in ( + ("input", conv.input_quantizer), + ("weight", conv.weight_quantizer), + ): + num_bits, amax = original_state[name] + assert quantizer.num_bits == num_bits + assert quantizer.amax is amax + + monkeypatch.setattr(backbone, "to", lambda *args, **kwargs: backbone) + monkeypatch.setattr(torch.cuda, "empty_cache", lambda: None) + monkeypatch.setitem(sys.modules, "onnx_utils", diffusion_onnx_utils) + monkeypatch.setitem(sys.modules, "onnx_utils.export", diffusion_export) + monkeypatch.setattr(diffusion_export, "modelopt_export_sd", fake_onnx_export) + monkeypatch.setattr(quantize_module, "export_hf_checkpoint", fake_hf_export) + + if export_fails: + with pytest.raises(RuntimeError, match="export failed"): + manager.export_onnx(_Pipeline(), backbone, ModelType.SDXL_BASE, QuantFormat.FP4) + else: + manager.export_onnx(_Pipeline(), backbone, ModelType.SDXL_BASE, QuantFormat.FP4) + manager.export_hf_ckpt(_Pipeline(), ModelConfig(model_type=ModelType.SDXL_BASE)) + + for name, quantizer in ( + ("input", conv.input_quantizer), + ("weight", conv.weight_quantizer), + ): + num_bits, amax = original_state[name] + assert quantizer.num_bits == num_bits + assert quantizer.amax is amax + + +@pytest.mark.parametrize("conv_cls", [nn.Conv1d, nn.Conv2d]) +def test_fp4_finalization_disables_unsupported_nvfp4_conv(conv_cls): + conv = _quantized_conv(conv_cls, _nvfp4_quantizer) + backbone = _Backbone(conv=conv) + backbone.fp8_conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) + + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4), + ModelType.SDXL_BASE, + restored=False, + ) + + assert not conv.input_quantizer.is_enabled + assert not conv.weight_quantizer.is_enabled + + +@pytest.mark.parametrize( + ("input_quantizer", "weight_quantizer", "match"), + [ + (_nvfp4_quantizer(), _nvfp4_quantizer(enabled=False), "enable input and weight"), + (_fp8_quantizer(), _nvfp4_quantizer(), "dynamic E2M1"), + (_nvfp4_quantizer(32), _nvfp4_quantizer(), "requires block size 16"), + ( + _nvfp4_quantizer(enabled=False), + _nvfp4_quantizer(enabled=False), + "at least one enabled Linear", + ), + ], + ids=["unpaired", "wrong-format", "wrong-block-size", "no-enabled-pair"], +) +def test_nvfp4_linear_validation_rejects_invalid_state(input_quantizer, weight_quantizer, match): + backbone = _Backbone( + linear=_QuantizedLinear( + input_quantizer=input_quantizer, + weight_quantizer=weight_quantizer, + ) + ) + + with pytest.raises(ValueError, match=match): + validate_nvfp4_quantizers(backbone, expected_block_size=16, quantize_mha=False) + + +def test_nvfp4_linear_validation_allows_disabled_exclusion_pair(): + backbone = nn.Module() + backbone.enabled = _QuantizedLinear() + backbone.excluded = _QuantizedLinear( + input_quantizer=_nvfp4_quantizer(enabled=False), + weight_quantizer=_nvfp4_quantizer(enabled=False), + ) + + validate_nvfp4_quantizers(backbone, expected_block_size=16, quantize_mha=False) + + +@pytest.mark.parametrize("projection_name", ["to_q", "to_k", "to_v"]) +def test_sdxl_fp4_validation_rejects_enabled_attention_projection(projection_name): + backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) + backbone.attn = nn.Module() + setattr(backbone.attn, projection_name, _QuantizedLinear()) + + with pytest.raises(ValueError, match="must keep input and weight quantizers disabled"): + validate_nvfp4_quantizers( + backbone, + expected_block_size=16, + quantize_mha=False, + validate_sdxl_mixed_recipe=True, + ) + + +@pytest.mark.parametrize( + ("input_enabled", "weight_enabled"), + [(True, False), (False, True)], + ids=["input-only", "weight-only"], +) +def test_sdxl_fp4_validation_rejects_partial_attention_projection(input_enabled, weight_enabled): + backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) + backbone.attn = nn.Module() + backbone.attn.to_q = _QuantizedLinear( + input_quantizer=_nvfp4_quantizer(enabled=input_enabled), + weight_quantizer=_nvfp4_quantizer(enabled=weight_enabled), + ) + + with pytest.raises(ValueError, match="must keep input and weight quantizers disabled"): + validate_nvfp4_quantizers( + backbone, + expected_block_size=16, + quantize_mha=False, + validate_sdxl_mixed_recipe=True, + ) + + +def test_sdxl_fp4_validation_allows_disabled_attention_projections(): + backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) + backbone.attn = nn.Module() + for projection_name in ("to_q", "to_k", "to_v"): + setattr( + backbone.attn, + projection_name, + _QuantizedLinear( + input_quantizer=_nvfp4_quantizer(enabled=False), + weight_quantizer=_nvfp4_quantizer(enabled=False), + ), + ) + + validate_nvfp4_quantizers( + backbone, + expected_block_size=16, + quantize_mha=False, + validate_sdxl_mixed_recipe=True, + ) + + +@pytest.mark.parametrize( + ("quantizer_name", "amax"), + [ + ("input_quantizer", None), + ("weight_quantizer", float("nan")), + ("weight_quantizer", -1.0), + ], +) +def test_restore_finalization_rejects_uncalibrated_nvfp4_linear(quantizer_name, amax): + linear = _QuantizedLinear() + setattr(linear, quantizer_name, _nvfp4_quantizer(amax=amax)) + backbone = _Backbone(linear=linear, conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) + + with pytest.raises(ValueError, match="must have a finite nonnegative calibrated amax"): + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4), + ModelType.SDXL_BASE, + restored=True, + ) + + +def test_nvfp4_linear_validation_accepts_zero_amax(): + backbone = _Backbone( + linear=_QuantizedLinear( + input_quantizer=_nvfp4_quantizer(amax=0.0), + weight_quantizer=_nvfp4_quantizer(amax=0.0), + ) + ) + + validate_nvfp4_quantizers(backbone, expected_block_size=16, quantize_mha=False) + + +@pytest.mark.parametrize( + ("quantizer_factory", "match"), + [ + (_nvfp4_quantizer, "only supported on Linear input and weight quantizers"), + (_fp8_quantizer, "only supported on SDXL Conv2d"), + ], + ids=["nvfp4-layernorm", "fp8-layernorm"], +) +def test_sdxl_fp4_validation_rejects_quantized_layernorm(quantizer_factory, match): + conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) + backbone = _Backbone(conv=conv) + backbone.norm = nn.LayerNorm(16) + backbone.norm.input_quantizer = quantizer_factory() + + with pytest.raises(ValueError, match=match): + validate_nvfp4_quantizers( + backbone, + expected_block_size=16, + quantize_mha=False, + validate_sdxl_mixed_recipe=True, + ) + + +def test_restore_finalization_rejects_enabled_nvfp4_layernorm(): + conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) + backbone = _Backbone(conv=conv) + backbone.norm = nn.LayerNorm(16) + backbone.norm.input_quantizer = _nvfp4_quantizer() + + with pytest.raises(ValueError, match="only supported on Linear input and weight quantizers"): + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4), + ModelType.SDXL_BASE, + restored=True, + ) + + +@pytest.mark.parametrize("restored", [False, True]) +def test_non_sdxl_fp4_finalization_remains_permissive(restored): + backbone = _Backbone() + backbone.norm = nn.LayerNorm(16) + backbone.norm.input_quantizer = _nvfp4_quantizer() + + _finalize_backbone_quantization( + backbone, + "transformer", + QuantizationConfig(format=QuantFormat.FP4), + ModelType.FLUX_DEV, + restored=restored, + ) + + assert backbone.norm.input_quantizer.is_enabled + + +@pytest.mark.parametrize( + ("input_factory", "weight_factory", "match"), + [ + ( + _fp8_quantizer, + lambda: _fp8_quantizer(enabled=False), + "enable input and weight quantizers as a pair", + ), + ( + _fp8_quantizer, + _nvfp4_quantizer, + "must use per-tensor FP8", + ), + ( + lambda: _fp8_quantizer(amax=None), + _fp8_quantizer, + "finite positive calibrated amax", + ), + ( + lambda: _fp8_quantizer(amax=0.0), + _fp8_quantizer, + "finite positive calibrated amax", + ), + ( + lambda: _fp8_quantizer(amax=float("nan")), + _fp8_quantizer, + "finite positive calibrated amax", + ), + ], + ids=["partial", "wrong-format", "missing-amax", "zero-amax", "nonfinite-amax"], +) +def test_sdxl_fp8_conv_validation_rejects_invalid_state(input_factory, weight_factory, match): + conv = nn.Conv2d(4, 4, kernel_size=1, bias=False) + conv.input_quantizer = input_factory() + conv.weight_quantizer = weight_factory() + backbone = _Backbone(conv=conv) + + with pytest.raises(ValueError, match=match): + validate_nvfp4_quantizers( + backbone, + expected_block_size=16, + quantize_mha=False, + validate_sdxl_mixed_recipe=True, + ) + + +def test_sdxl_fp8_conv_validation_requires_enabled_pair(): + conv = _quantized_conv(nn.Conv2d, lambda: _fp8_quantizer(enabled=False)) + backbone = _Backbone(conv=conv) + + with pytest.raises(ValueError, match="at least one enabled calibrated FP8 Conv2d"): + validate_nvfp4_quantizers( + backbone, + expected_block_size=16, + quantize_mha=False, + validate_sdxl_mixed_recipe=True, + ) + + +@pytest.mark.parametrize(("quantize_mha", "enabled"), [(False, False), (True, True)]) +def test_restore_finalization_accepts_matching_mha_state(quantize_mha, enabled): + attention = _attention(enabled=enabled) + conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) + backbone = _Backbone(conv=conv, attention=attention) + + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4, quantize_mha=quantize_mha), + ModelType.SDXL_BASE, + restored=True, + ) + + for name in _MHA_QUANTIZER_NAMES: + expected_enabled = enabled and name != "bmm2_output_quantizer" + assert getattr(attention, name).is_enabled is expected_enabled + + +def test_restore_finalization_accepts_uncalibrated_softmax_quantizer(): + attention = _attention(enabled=True, softmax_amax=None) + backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer), attention=attention) + + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4, quantize_mha=True), + ModelType.SDXL_BASE, + restored=True, + ) + + assert attention.softmax_quantizer.is_enabled + assert attention.softmax_quantizer.amax is None + + +@pytest.mark.parametrize( + ("quantize_mha", "head_size", "mutate", "match"), + [ + (False, 16, lambda attention: None, "must be disabled"), + ( + True, + 16, + lambda attention: attention.q_bmm_quantizer.disable(), + "requires 'q_bmm_quantizer' to be enabled", + ), + ( + True, + 16, + lambda attention: attention.softmax_quantizer.disable(), + "requires 'softmax_quantizer' to be enabled", + ), + ( + True, + 16, + lambda attention: attention.bmm2_output_quantizer.enable(), + "requires 'bmm2_output_quantizer' to be disabled", + ), + ( + True, + 16, + lambda attention: setattr(attention, "softmax_quantizer", _nvfp4_quantizer()), + "must use per-tensor FP8", + ), + ( + True, + 16, + lambda attention: setattr(attention, "q_bmm_quantizer", _fp8_quantizer(amax=None)), + "finite positive calibrated amax", + ), + (True, 8, lambda attention: None, "must be disabled because FP8 MHA is unsupported"), + ], + ids=[ + "on-to-off", + "off-to-on", + "missing-softmax", + "enabled-bmm2-output", + "wrong-format", + "missing-amax", + "unsupported-head-size", + ], +) +def test_restore_finalization_rejects_mha_state_mismatch(quantize_mha, head_size, mutate, match): + attention = _attention(head_size=head_size, enabled=True) + mutate(attention) + backbone = _Backbone(attention=attention) + + with pytest.raises(ValueError, match=match): + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4, quantize_mha=quantize_mha), + ModelType.SDXL_BASE, + restored=True, + ) + + +@pytest.mark.parametrize("quantize_mha", [False, True]) +def test_fresh_finalization_applies_and_validates_mha_policy(quantize_mha): + attention = _attention(enabled=True) + conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) + backbone = _Backbone(conv=conv, attention=attention) + + _finalize_backbone_quantization( + backbone, + "unet", + QuantizationConfig(format=QuantFormat.FP4, quantize_mha=quantize_mha), + ModelType.SDXL_BASE, + restored=False, + ) + + for name in _MHA_QUANTIZER_NAMES: + quantizer = getattr(attention, name) + expected_enabled = quantize_mha and name != "bmm2_output_quantizer" + assert quantizer.is_enabled is expected_enabled + if quantizer.is_enabled: + assert quantizer.is_fp8 + + +def test_check_conv_does_not_disable_non_nvfp4_quantizers(): + conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) + backbone = _Backbone(conv=conv) + + check_conv_and_mha(backbone, if_fp4=True, quantize_mha=False) + + assert conv.input_quantizer.is_enabled + assert conv.weight_quantizer.is_enabled From 320d4bb59ab52c96ef6ece18ab4ba3ac3d8ff94c Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:22:34 +0000 Subject: [PATCH 02/10] [5565357] Simplify SDXL NVFP4 recipe Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/ONNX-TRT-Deployment.md | 16 +- .../quantization/onnx_utils/export.py | 1032 ++--------------- examples/diffusers/quantization/quantize.py | 136 +-- examples/diffusers/quantization/utils.py | 250 +--- tests/examples/diffusers/test_diffusers.py | 12 - tests/unit/examples/test_diffusers_fp4.py | 243 ++++ .../test_diffusers_fp4_onnx_validation.py | 992 ---------------- .../examples/test_diffusers_fp4_validation.py | 652 ----------- 8 files changed, 390 insertions(+), 2943 deletions(-) create mode 100644 tests/unit/examples/test_diffusers_fp4.py delete mode 100644 tests/unit/examples/test_diffusers_fp4_onnx_validation.py delete mode 100644 tests/unit/examples/test_diffusers_fp4_validation.py diff --git a/examples/diffusers/quantization/ONNX-TRT-Deployment.md b/examples/diffusers/quantization/ONNX-TRT-Deployment.md index bb0239d5c31..33aab492380 100644 --- a/examples/diffusers/quantization/ONNX-TRT-Deployment.md +++ b/examples/diffusers/quantization/ONNX-TRT-Deployment.md @@ -28,7 +28,7 @@ python quantize.py \ #### FLUX-Dev|SDXL|SDXL-Turbo|LTX-Video FP8/FP4 [Script](./quantize.py) -FP4 ONNX export is supported for Flux and the SDXL family. For SDXL, the FP4 recipe uses block-16 NVFP4 for non-QKV Linear/GEMM layers and FP8 for Conv2d layers. Attention `to_q`, `to_k`, and `to_v` projection Linears remain in the high-precision model dtype so TensorRT can preserve their horizontal projection fusion. The script selects this mixed-precision recipe automatically when `--format fp4` is used. +FP4 ONNX export is supported for Flux and SDXL. SDXL uses block-16 NVFP4 for non-QKV Linear/GEMM layers and FP8 for Conv2d layers, while Q/K/V projection Linears remain in the model dtype to preserve TensorRT fusion. Add `--quantize-mha` to optionally quantize MHA with FP8. FP4 deployment requires a Blackwell GPU and TensorRT with NVFP4 support. ```sh python quantize.py \ @@ -38,20 +38,18 @@ python quantize.py \ --onnx-dir {ONNX_DIR} ``` -Add `--quantize-mha` to opt in to FP8 MHA quantization; this does not quantize the QKV projection Linears. - We recommend using a device with a minimum of 48GB of combined CPU and GPU memory for exporting ONNX models. If not, please use CPU for ONNX export. ## Build the TRT engine for the Quantized ONNX Backbone > [!IMPORTANT] > TensorRT environment must be setup prior -- Please see [Pre-Requisites](../README.md#pre-requisites) -> INT8 requires **TensorRT version >= 9.2.0**. FP8 requires **TensorRT version 10.2.0 or higher**. FP4 requires a Blackwell GPU (SM100 or newer) and a TensorRT version with NVFP4 support. You can download the latest version of TensorRT [here](https://developer.nvidia.com/tensorrt/download). Deployment of SVDQuant is currently not supported. +> INT8 requires **TensorRT version >= 9.2.0**. If you prefer to use the FP8 TensorRT, ensure you have **TensorRT version 10.2.0 or higher**. You can download the latest version of TensorRT at [here](https://developer.nvidia.com/tensorrt/download). Deployment of SVDQuant is currently not supported. -Generate INT8/FP8/FP4 Backbone Engine +Generate INT8/FP8 Backbone Engine ```bash -# For SDXL INT8, FP8, or FP4 +# For SDXL trtexec --builderOptimizationLevel=4 --stronglyTyped --onnx=./model.onnx \ --minShapes=sample:2x4x128x128,timestep:1,encoder_hidden_states:2x77x2048,text_embeds:2x1280,time_ids:2x6 \ --optShapes=sample:16x4x128x128,timestep:1,encoder_hidden_states:16x77x2048,text_embeds:16x1280,time_ids:16x6 \ @@ -93,15 +91,15 @@ python demo_txt2img_xl.py "enchanted winter forest, soft diffuse light on a snow Note, it will take some time to build TRT engines for the first time -- Replace the FP16 backbone TensorRT engine with the quantized engine generated in [Build the TRT engine for the Quantized ONNX Backbone](#build-the-trt-engine-for-the-quantized-onnx-backbone), e.g.: +- Replace the fp16 backbone TRT engine with int8 engine generated in [Build the TRT engine for the Quantized ONNX Backbone](#build-the-trt-engine-for-the-quantized-onnx-backbone), e.g.,: ```sh cp -r {YOUR_UNETXL}.plan ./engine/ ``` -The engines must be built on the same GPU, and the quantized engine name must match the FP16 engine name to enable compatibility with the demoDiffusion pipeline. +Note, the engines must be built on the same GPU, and ensure that the INT8 engine name matches the names of the FP16 engines to enable compatibility with the demoDiffusion pipeline. -- Run the above txt2img example command again. You can compare the generated images and latency for FP16 versus INT8, FP8, or FP4. +- Run the above txt2img example command again. You can compare the generated images and latency for fp16 vs int8. Similarly, you could run end-to-end pipeline with Model Optimizer quantized backbone and corresponding examples in demoDiffusion with other diffusion models. ## Running the inference pipeline with DeviceModel diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 7936bc1933d..57d04296e91 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -32,11 +32,9 @@ import os import shutil import tempfile -import uuid -from contextlib import nullcontext, suppress +from contextlib import contextmanager, nullcontext from pathlib import Path -import numpy as np import onnx import onnx_graphsurgeon as gs import torch @@ -50,9 +48,7 @@ from torch.onnx import export as onnx_export from modelopt.onnx.export import NVFP4QuantExporter -from modelopt.onnx.quantization.graph_utils import get_tensor_consumer_nodes from modelopt.torch.quantization.export_onnx import configure_linear_module_onnx_quantizers -from modelopt.torch.quantization.nn.modules.quant_linear import RealQuantLinear from modelopt.torch.utils import torch_to from .fp8_onnx_graphsurgeon import convert_zp_fp8 @@ -128,7 +124,17 @@ def flux_convert_rope_weight_type(onnx_graph): return gs.export_onnx(graph) -def generate_fp8_scales(backbone, *, conv_only=False): +def _has_enabled_conv(backbone): + for module in backbone.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 + + +@contextmanager +def _temporary_fp8_export_scales(backbone, conv_only=False): # temporary solution due to a known bug in torch.onnx._dynamo_export module_types = (torch.nn.Conv2d,) if conv_only else (torch.nn.Linear, torch.nn.Conv2d) quantizer_states = [] @@ -148,16 +154,11 @@ def generate_fp8_scales(backbone, *, conv_only=False): quantizer_states.append((quantizer, quantizer._num_bits, quantizer._amax)) quantizer._num_bits = 8 quantizer._amax = quantizer._amax * (127 / 448.0) - except BaseException: - restore_fp8_scales(quantizer_states) - raise - return quantizer_states - - -def restore_fp8_scales(quantizer_states): - for quantizer, num_bits, amax in reversed(quantizer_states): - quantizer._num_bits = num_bits - quantizer._amax = amax + yield + finally: + for quantizer, num_bits, amax in reversed(quantizer_states): + quantizer._num_bits = num_bits + quantizer._amax = amax def _gen_dummy_inp_and_dyn_shapes_sdxl(backbone, min_bs=1, opt_bs=1): @@ -475,729 +476,16 @@ def remove_nesting(trt_dynamic_shapes): ) -def _get_int_attribute(node, name): - for attribute in node.attribute: - if attribute.name == name and attribute.type == onnx.AttributeProto.INT: - return attribute.i - return None - - -_WEIGHT_PASSTHROUGH_OPS = {"Cast", "Flatten", "Identity", "Reshape", "Transpose"} - - -def _trace_initializer_source(tensor_name, producers, initializer_names): - visited = set() - while tensor_name not in initializer_names: - if tensor_name in visited: - return None - visited.add(tensor_name) - producer = producers.get(tensor_name) - if ( - producer is None - or producer.op_type not in _WEIGHT_PASSTHROUGH_OPS - or not producer.input - ): - return None - tensor_name = producer.input[0] - return tensor_name - - -def _trace_tensor_consumers( - tensor_names, - consumers, - graph_outputs, - terminal_ops, - terminal_input_index, - passthrough_ops=_WEIGHT_PASSTHROUGH_OPS, - allowed_cast_dtypes=None, -): - pending = list(tensor_names) - visited = set() - terminal_consumers = [] - invalid_consumers = set() - - while pending: - tensor_name = pending.pop() - if tensor_name in visited: - continue - visited.add(tensor_name) - - if tensor_name in graph_outputs: - invalid_consumers.add(f"graph output {tensor_name}") - tensor_consumers = consumers.get(tensor_name, []) - if not tensor_consumers: - invalid_consumers.add(f"unused tensor {tensor_name}") - continue - - for consumer in tensor_consumers: - input_indices = [ - index - for index, input_name in enumerate(consumer.input) - if input_name == tensor_name - ] - if consumer.op_type in passthrough_ops and input_indices == [0]: - if ( - consumer.op_type == "Cast" - and allowed_cast_dtypes is not None - and _get_int_attribute(consumer, "to") not in allowed_cast_dtypes - ): - invalid_consumers.add(consumer.name or consumer.op_type) - continue - if consumer.output: - pending.extend(consumer.output) - else: - invalid_consumers.add(consumer.name or consumer.op_type) - elif consumer.op_type in terminal_ops and input_indices == [terminal_input_index]: - terminal_consumers.append(consumer) - else: - invalid_consumers.add( - consumer.name or (consumer.output[0] if consumer.output else consumer.op_type) - ) - - return terminal_consumers, sorted(invalid_consumers) - - -def _trace_source_node(tensor_name, producers): - visited = set() - while tensor_name not in visited: - visited.add(tensor_name) - producer = producers.get(tensor_name) - if producer is None or producer.op_type not in _WEIGHT_PASSTHROUGH_OPS: - return producer - if not producer.input: - return None - tensor_name = producer.input[0] - return None - - -def _find_initializer_backed_qdq_weights(onnx_model, allow_fp8_conv): - initializer_names = {initializer.name for initializer in onnx_model.graph.initializer} - producers = {output: node for node in onnx_model.graph.node for output in node.output if output} - consumers = get_tensor_consumer_nodes(onnx_model.graph) - graph_outputs = {output.name for output in onnx_model.graph.output} - disallowed_consumers = set() - allowed_pairs = [] - - for node in onnx_model.graph.node: - if node.op_type != "DequantizeLinear" or not node.input: - continue - quantize_node = producers.get(node.input[0]) - if ( - quantize_node is None - or quantize_node.op_type != "QuantizeLinear" - or not quantize_node.input - or _trace_initializer_source(quantize_node.input[0], producers, initializer_names) - is None - ): - continue - - terminal_consumers, invalid_consumers = _trace_tensor_consumers( - node.output, consumers, graph_outputs, {"Conv"}, 1 - ) - if allow_fp8_conv and terminal_consumers and not invalid_consumers: - allowed_pairs.append((quantize_node, node, terminal_consumers)) - else: - disallowed_consumers.update(invalid_consumers) - disallowed_consumers.update( - consumer.name or (consumer.output[0] if consumer.output else consumer.op_type) - for consumer in terminal_consumers - ) - if not terminal_consumers and not invalid_consumers: - disallowed_consumers.add(node.name or node.output[0]) - - return allowed_pairs, sorted(disallowed_consumers) - - -def _get_tensor_dtype(tensor_name, initializers, producers): - initializer = initializers.get(tensor_name) - if initializer is not None: - return initializer.data_type - producer = producers.get(tensor_name) - if producer is None or producer.op_type != "Constant": - return None - for attribute in producer.attribute: - if attribute.name == "value" and attribute.type == onnx.AttributeProto.TENSOR: - return attribute.t.data_type - return None - - -def _get_effective_tensor_dtype(tensor_name, initializers, producers, declared_dtypes): - visited = set() - while tensor_name not in visited: - visited.add(tensor_name) - producer = producers.get(tensor_name) - if producer is not None and producer.op_type == "Cast": - return _get_int_attribute(producer, "to") - if tensor_name in declared_dtypes: - return declared_dtypes[tensor_name] - tensor_dtype = _get_tensor_dtype(tensor_name, initializers, producers) - if tensor_dtype is not None: - return tensor_dtype - if producer is None or producer.op_type not in _WEIGHT_PASSTHROUGH_OPS: - return None - if not producer.input: - return None - tensor_name = producer.input[0] - return None - - -def _get_constant_array(tensor_name, initializers, producers): - tensor = initializers.get(tensor_name) - if tensor is None: - producer = producers.get(tensor_name) - if producer is not None and producer.op_type == "Constant": - tensor = next( - ( - attribute.t - for attribute in producer.attribute - if attribute.name == "value" and attribute.type == onnx.AttributeProto.TENSOR - ), - None, - ) - if tensor is None: - return None - try: - return onnx.numpy_helper.to_array(tensor) - except (TypeError, ValueError): - return None - - -def _validate_positive_scalar(tensor_name, role, initializers, producers): - value = _get_constant_array(tensor_name, initializers, producers) - if value is None or value.size != 1 or not np.isfinite(value).all() or not (value > 0).all(): - return [f"{role} must be a finite positive scalar constant"] - return [] - - -def _validate_normalized_fp8_qdq_pair( - quantize_node, - dequantize_node, - pair_name, - initializers, - producers, - consumers, - expected_scale_dtype=None, -): - errors = [] - if len(quantize_node.input) != 3 or len(dequantize_node.input) != 3: - return [f"{pair_name} must use three-input FP8 Q/DQ nodes"] - if len(quantize_node.output) != 1 or dequantize_node.input[0] != quantize_node.output[0]: - errors.append(f"{pair_name} does not form a direct Q/DQ pair") - elif consumers.get(quantize_node.output[0], []) != [dequantize_node]: - errors.append(f"{pair_name} quantized tensor must be consumed only by its DQ") - - if quantize_node.input[1:] != dequantize_node.input[1:]: - errors.append(f"{pair_name} Q/DQ nodes do not share scale and zero point") - if any( - attribute.name == "axis" - for node in (quantize_node, dequantize_node) - for attribute in node.attribute - ): - errors.append(f"{pair_name} must use per-tensor FP8 Q/DQ without an axis") - - errors.extend( - _validate_positive_scalar( - quantize_node.input[1], f"{pair_name} FP8 scale", initializers, producers - ) - ) - scale_dtype = _get_tensor_dtype(quantize_node.input[1], initializers, producers) - if scale_dtype not in { - onnx.TensorProto.FLOAT, - onnx.TensorProto.FLOAT16, - onnx.TensorProto.BFLOAT16, - }: - errors.append(f"{pair_name} FP8 scale must use a floating-point dtype") - elif expected_scale_dtype is not None and scale_dtype != expected_scale_dtype: - errors.append(f"{pair_name} FP8 scale dtype does not match its quantized tensor") - for role, node in ( - ("QuantizeLinear", quantize_node), - ("DequantizeLinear", dequantize_node), - ): - zero_name = node.input[2] - if _get_tensor_dtype(zero_name, initializers, producers) != onnx.TensorProto.FLOAT8E4M3FN: - errors.append(f"{pair_name} {role} zero point is not FLOAT8E4M3FN") - continue - zero = _get_constant_array(zero_name, initializers, producers) - if zero is None or zero.size != 1 or not (zero == 0).all(): - errors.append(f"{pair_name} {role} zero point must be a scalar zero") - return errors - - -def _validate_normalized_fp8_qdq(onnx_model, qdq_records): - initializers = {initializer.name: initializer for initializer in onnx_model.graph.initializer} - producers = {output: node for node in onnx_model.graph.node for output in node.output if output} - consumers = get_tensor_consumer_nodes(onnx_model.graph) - graph_outputs = {output.name for output in onnx_model.graph.output} - initializer_names = set(initializers) - declared_dtypes = { - value.name: value.type.tensor_type.elem_type - for values in ( - onnx_model.graph.input, - onnx_model.graph.value_info, - onnx_model.graph.output, - ) - for value in values - if value.type.HasField("tensor_type") - } - errors = [] - fp8_conv_ids = { - id(consumer) for _, _, terminal_consumers in qdq_records for consumer in terminal_consumers - } - validated_activation_dq_ids = set() - - for quantize_node, dequantize_node, terminal_consumers in qdq_records: - pair_name = dequantize_node.name or dequantize_node.output[0] - if len(terminal_consumers) != 1: - errors.append(f"{pair_name} weight DQ must feed exactly one FP8 Conv input 1") - errors.extend( - _validate_normalized_fp8_qdq_pair( - quantize_node, - dequantize_node, - pair_name, - initializers, - producers, - consumers, - _get_effective_tensor_dtype( - quantize_node.input[0], initializers, producers, declared_dtypes - ), - ) - ) - - for conv_node in terminal_consumers: - conv_name = conv_node.name or conv_node.output[0] - activation_dq = _trace_source_node(conv_node.input[0], producers) - if activation_dq is None or activation_dq.op_type != "DequantizeLinear": - errors.append(f"{conv_name} has no FP8 activation Q/DQ on input 0") - continue - activation_q = producers.get(activation_dq.input[0]) if activation_dq.input else None - if activation_q is None or activation_q.op_type != "QuantizeLinear": - errors.append(f"{conv_name} has no FP8 activation QuantizeLinear on input 0") - continue - if ( - not activation_q.input - or _trace_initializer_source(activation_q.input[0], producers, initializer_names) - is not None - ): - errors.append(f"{conv_name} activation Q/DQ is initializer-backed") - continue - - if id(activation_dq) not in validated_activation_dq_ids: - activation_name = activation_dq.name or activation_dq.output[0] - errors.extend( - _validate_normalized_fp8_qdq_pair( - activation_q, - activation_dq, - activation_name, - initializers, - producers, - consumers, - _get_effective_tensor_dtype( - activation_q.input[0], initializers, producers, declared_dtypes - ), - ) - ) - activation_consumers, invalid_consumers = _trace_tensor_consumers( - activation_dq.output, consumers, graph_outputs, {"Conv"}, 0 - ) - if invalid_consumers: - errors.append( - f"{activation_name} has non-Conv activation consumers: " - + ", ".join(invalid_consumers[:5]) - ) - if len(activation_consumers) != 1 or id(conv_node) not in { - id(consumer) for consumer in activation_consumers - }: - errors.append( - f"{activation_name} must feed exactly one validated FP8 Conv input 0" - ) - unexpected_conv_ids = { - id(consumer) for consumer in activation_consumers - } - fp8_conv_ids - if unexpected_conv_ids: - errors.append( - f"{activation_name} reaches a Conv without a validated FP8 weight Q/DQ" - ) - validated_activation_dq_ids.add(id(activation_dq)) - return errors - - -def _validate_dynamic_fp4_activations( - onnx_model, expected_block_size, fp4_weight_terminal_consumers -): - initializers = {initializer.name: initializer for initializer in onnx_model.graph.initializer} - producers = {output: node for node in onnx_model.graph.node for output in node.output if output} - consumers = get_tensor_consumer_nodes(onnx_model.graph) - graph_outputs = {output.name for output in onnx_model.graph.output} - expected_terminal_ids = {id(node) for node in fp4_weight_terminal_consumers} - dynamic_terminal_counts = {} - initializer_or_constant_names = set(initializers) | { - output - for node in onnx_model.graph.node - if node.op_type == "Constant" - for output in node.output - } - errors = [] - dynamic_nodes = [ - node for node in onnx_model.graph.node if node.op_type == "TRT_FP4DynamicQuantize" - ] - - if not dynamic_nodes: - errors.append("no TRT_FP4DynamicQuantize activation nodes were exported") - - for node in dynamic_nodes: - node_name = node.name or (node.output[0] if node.output else "") - if len(node.input) != 2 or len(node.output) != 2: - errors.append(f"{node_name} must have two inputs and two outputs") - continue - if node.domain != "trt": - errors.append(f"{node_name} must use the trt domain") - if ( - _trace_initializer_source(node.input[0], producers, initializer_or_constant_names) - is not None - ): - errors.append(f"{node_name} input 0 must be a dynamic activation") - if _get_int_attribute(node, "block_size") != expected_block_size: - errors.append(f"{node_name} does not use block_size={expected_block_size}") - if _get_int_attribute(node, "axis") != -1: - errors.append(f"{node_name} does not use axis=-1") - if _get_int_attribute(node, "scale_type") != onnx.TensorProto.FLOAT8E4M3FN: - errors.append(f"{node_name} does not produce FLOAT8E4M3FN block scales") - if _get_tensor_dtype(node.input[1], initializers, producers) != onnx.TensorProto.FLOAT: - errors.append(f"{node_name} global scale is not FLOAT") - errors.extend( - _validate_positive_scalar( - node.input[1], f"{node_name} global scale", initializers, producers - ) - ) - - quantized_consumers = consumers.get(node.output[0], []) - if ( - len(quantized_consumers) != 1 - or quantized_consumers[0].op_type != "DequantizeLinear" - or not quantized_consumers[0].input - or quantized_consumers[0].input[0] != node.output[0] - ): - errors.append(f"{node_name} FP4 output must feed exactly one DequantizeLinear") - continue - activation_dq = quantized_consumers[0] - activation_dq_name = activation_dq.name or activation_dq.output[0] - if len(activation_dq.input) != 2: - errors.append(f"{activation_dq_name} must be a two-input DequantizeLinear") - continue - if _get_int_attribute(activation_dq, "block_size") != expected_block_size: - errors.append(f"{activation_dq_name} does not use block_size={expected_block_size}") - if _get_int_attribute(activation_dq, "axis") != -1: - errors.append(f"{activation_dq_name} does not use axis=-1") - - scale_consumers = consumers.get(node.output[1], []) - if ( - len(scale_consumers) != 1 - or scale_consumers[0].op_type != "DequantizeLinear" - or not scale_consumers[0].input - or scale_consumers[0].input[0] != node.output[1] - ): - errors.append(f"{node_name} FP8 scale output must feed exactly one DequantizeLinear") - continue - scale_dq = scale_consumers[0] - scale_dq_name = scale_dq.name or scale_dq.output[0] - if len(scale_dq.input) != 2: - errors.append(f"{scale_dq_name} must be a two-input DequantizeLinear") - continue - if any(attribute.name in {"axis", "block_size"} for attribute in scale_dq.attribute): - errors.append(f"{scale_dq_name} must not use axis or block_size") - if _get_tensor_dtype(scale_dq.input[1], initializers, producers) != onnx.TensorProto.FLOAT: - errors.append(f"{scale_dq_name} global scale is not FLOAT") - errors.extend( - _validate_positive_scalar( - scale_dq.input[1], f"{scale_dq_name} global scale", initializers, producers - ) - ) - quantize_scale = _get_constant_array(node.input[1], initializers, producers) - dequantize_scale = _get_constant_array(scale_dq.input[1], initializers, producers) - if ( - quantize_scale is not None - and dequantize_scale is not None - and not np.array_equal(quantize_scale, dequantize_scale) - ): - errors.append(f"{node_name} quantize and dequantize global scales do not match") - if not scale_dq.output or activation_dq.input[1] != scale_dq.output[0]: - errors.append(f"{activation_dq_name} is not scaled by {scale_dq_name}") - elif consumers.get(scale_dq.output[0], []) != [activation_dq]: - errors.append(f"{scale_dq_name} output must be consumed only by {activation_dq_name}") - - terminal_consumers, invalid_consumers = _trace_tensor_consumers( - activation_dq.output, - consumers, - graph_outputs, - {"Gemm", "MatMul"}, - 0, - {"Cast", "Identity"}, - {onnx.TensorProto.FLOAT16, onnx.TensorProto.BFLOAT16}, - ) - for consumer in terminal_consumers: - terminal_id = id(consumer) - dynamic_terminal_counts[terminal_id] = dynamic_terminal_counts.get(terminal_id, 0) + 1 - if not terminal_consumers: - errors.append(f"{activation_dq_name} does not reach a Gemm/MatMul activation input") - elif len(terminal_consumers) != 1: - errors.append( - f"{activation_dq_name} must feed exactly one Gemm/MatMul activation input" - ) - if invalid_consumers: - errors.append( - f"{activation_dq_name} has non-activation consumers: " - + ", ".join(invalid_consumers[:5]) - ) - - dynamic_terminal_ids = set(dynamic_terminal_counts) - duplicate = sum(count != 1 for count in dynamic_terminal_counts.values()) - if dynamic_terminal_ids != expected_terminal_ids or duplicate: - missing = len(expected_terminal_ids - dynamic_terminal_ids) - extra = len(dynamic_terminal_ids - expected_terminal_ids) - errors.append( - "dynamic NVFP4 activation paths do not match FLOAT4 weight consumers " - f"(missing={missing}, extra={extra}, duplicate={duplicate})" - ) - return errors - - -def _validate_raw_fp4_graph( - onnx_model, - expected_block_size=16, - *, - allow_fp8_conv=False, - expected_linear_count=None, - expected_fp8_conv_count=None, -): - initializer_names = {initializer.name for initializer in onnx_model.graph.initializer} - consumers = get_tensor_consumer_nodes(onnx_model.graph) - graph_outputs = {output.name for output in onnx_model.graph.output} - fp4_nodes = [node for node in onnx_model.graph.node if node.op_type == "TRT_FP4QDQ"] - errors = [] - - if not fp4_nodes: - errors.append("no TRT_FP4QDQ weight markers were exported") - if expected_linear_count is not None and len(fp4_nodes) != expected_linear_count: - errors.append( - f"found {len(fp4_nodes)} TRT_FP4QDQ weight markers, expected " - f"{expected_linear_count} enabled Linear pairs" - ) - - for node in fp4_nodes: - node_name = node.name or (node.output[0] if node.output else "") - if not node.input or node.input[0] not in initializer_names: - errors.append(f"{node_name} is not backed by a weight initializer") - block_size = _get_int_attribute(node, "block_size") - if block_size != expected_block_size: - errors.append( - f"{node_name} has block_size={block_size}, expected {expected_block_size}" - ) - terminal_consumers, invalid_consumers = _trace_tensor_consumers( - node.output, consumers, graph_outputs, {"Gemm", "MatMul"}, 1 - ) - if not terminal_consumers: - errors.append(f"{node_name} does not reach a Gemm/MatMul weight input") - if invalid_consumers: - errors.append( - f"{node_name} has non-weight consumers: " + ", ".join(invalid_consumers[:5]) - ) - - fp8_conv_pairs, disallowed_qdq_weights = _find_initializer_backed_qdq_weights( - onnx_model, allow_fp8_conv - ) - if allow_fp8_conv and not fp8_conv_pairs: - errors.append("no initializer-backed FP8 Conv weight Q/DQ was exported") - if expected_fp8_conv_count is not None and len(fp8_conv_pairs) != expected_fp8_conv_count: - errors.append( - f"found {len(fp8_conv_pairs)} initializer-backed FP8 Conv weight Q/DQ pairs, " - f"expected {expected_fp8_conv_count} enabled Conv2d pairs" - ) - if disallowed_qdq_weights: - errors.append( - "disallowed initializer-backed Q/DQ weight consumers: " - + ", ".join(disallowed_qdq_weights[:5]) - ) - - if errors: - raise ValueError("Invalid raw FP4 ONNX graph: " + "; ".join(errors)) - return len(fp4_nodes) - - -def _validate_final_fp4_graph( - onnx_model, - expected_weight_count, - expected_block_size=16, - *, - allow_fp8_conv=False, - expected_fp8_conv_count=None, -): - initializers = {initializer.name: initializer for initializer in onnx_model.graph.initializer} - producers = {output: node for node in onnx_model.graph.node for output in node.output if output} - consumers = get_tensor_consumer_nodes(onnx_model.graph) - graph_outputs = {output.name for output in onnx_model.graph.output} - fp4_initializer_names = { - name - for name, initializer in initializers.items() - if initializer.data_type == onnx.TensorProto.FLOAT4E2M1 - } - weight_dq_nodes = [ - node - for node in onnx_model.graph.node - if node.op_type == "DequantizeLinear" - and node.input - and node.input[0] in fp4_initializer_names - ] - errors = [] - weight_dq_ids = {id(node) for node in weight_dq_nodes} - fp4_weight_terminal_consumers = [] - - remaining_markers = [node for node in onnx_model.graph.node if node.op_type == "TRT_FP4QDQ"] - if remaining_markers: - errors.append(f"{len(remaining_markers)} TRT_FP4QDQ weight markers remain") - if len(fp4_initializer_names) != expected_weight_count: - errors.append( - f"found {len(fp4_initializer_names)} FLOAT4 weights, expected {expected_weight_count}" - ) - if len(weight_dq_nodes) != expected_weight_count: - errors.append( - f"found {len(weight_dq_nodes)} FLOAT4 weight DQ nodes, expected {expected_weight_count}" - ) - - for initializer_name in fp4_initializer_names: - direct_consumers = consumers.get(initializer_name, []) - if ( - len(direct_consumers) != 1 - or id(direct_consumers[0]) not in weight_dq_ids - or [ - index - for index, input_name in enumerate(direct_consumers[0].input) - if input_name == initializer_name - ] - != [0] - ): - errors.append( - f"FLOAT4 weight {initializer_name} must feed exactly one weight DequantizeLinear" - ) - - fp4_weight_names = set() - fp8_scale_names = set() - for node in weight_dq_nodes: - node_name = node.name or (node.output[0] if node.output else "") - fp4_weight_names.add(node.input[0]) - if len(node.input) != 2: - errors.append(f"{node_name} is not a two-input FLOAT4 DequantizeLinear") - continue - if _get_int_attribute(node, "block_size") != expected_block_size: - errors.append(f"{node_name} does not use block_size={expected_block_size}") - if _get_int_attribute(node, "axis") != -1: - errors.append(f"{node_name} does not use axis=-1") - - scale_dq = producers.get(node.input[1]) - if scale_dq is None or scale_dq.op_type != "DequantizeLinear": - errors.append(f"{node_name} is not scaled by a preceding DequantizeLinear") - continue - scale_dq_name = scale_dq.name or (scale_dq.output[0] if scale_dq.output else "") - if len(scale_dq.input) != 2: - errors.append(f"{scale_dq_name} must be a two-input DequantizeLinear") - continue - if any(attribute.name in {"axis", "block_size"} for attribute in scale_dq.attribute): - errors.append(f"{scale_dq_name} must not use axis or block_size") - - fp8_scale = initializers.get(scale_dq.input[0]) - global_scale = initializers.get(scale_dq.input[1]) - if fp8_scale is None or fp8_scale.data_type != onnx.TensorProto.FLOAT8E4M3FN: - errors.append(f"{node_name} does not use a FLOAT8E4M3FN block-scale initializer") - else: - fp8_scale_names.add(fp8_scale.name) - fp8_scale_consumers = consumers.get(fp8_scale.name, []) - if ( - len(fp8_scale_consumers) != 1 - or id(fp8_scale_consumers[0]) != id(scale_dq) - or [ - index - for index, input_name in enumerate(scale_dq.input) - if input_name == fp8_scale.name - ] - != [0] - ): - errors.append(f"{fp8_scale.name} must be consumed only by {scale_dq_name} input 0") - if global_scale is None or global_scale.data_type != onnx.TensorProto.FLOAT: - errors.append(f"{node_name} does not use a FLOAT global-scale initializer") - else: - errors.extend( - _validate_positive_scalar( - global_scale.name, - f"{scale_dq_name} global scale", - initializers, - producers, - ) - ) - if len(scale_dq.output) != 1 or node.input[1] != scale_dq.output[0]: - errors.append(f"{node_name} is not scaled by {scale_dq_name}") - else: - scale_output_consumers = consumers.get(scale_dq.output[0], []) - if ( - len(scale_output_consumers) != 1 - or id(scale_output_consumers[0]) != id(node) - or [ - index - for index, input_name in enumerate(node.input) - if input_name == scale_dq.output[0] - ] - != [1] - ): - errors.append( - f"{scale_dq_name} output must be consumed only by {node_name} input 1" - ) - - terminal_consumers, invalid_consumers = _trace_tensor_consumers( - node.output, consumers, graph_outputs, {"Gemm", "MatMul"}, 1 - ) - fp4_weight_terminal_consumers.extend(terminal_consumers) - if not terminal_consumers: - errors.append(f"{node_name} does not reach a Gemm/MatMul weight input") - if invalid_consumers: - errors.append( - f"{node_name} has non-weight consumers: " + ", ".join(invalid_consumers[:5]) - ) - - if len(fp4_weight_names) != expected_weight_count: - errors.append( - f"found {len(fp4_weight_names)} referenced FLOAT4 weights, " - f"expected {expected_weight_count}" - ) - - if len(fp8_scale_names) != expected_weight_count: - errors.append( - f"found {len(fp8_scale_names)} FLOAT8 block scales, expected {expected_weight_count}" - ) - - errors.extend( - _validate_dynamic_fp4_activations( - onnx_model, expected_block_size, fp4_weight_terminal_consumers - ) - ) - - fp8_conv_pairs, disallowed_qdq_weights = _find_initializer_backed_qdq_weights( - onnx_model, allow_fp8_conv +def save_onnx(onnx_model, output): + onnx.save( + onnx_model, + str(output), + save_as_external_data=True, + all_tensors_to_one_file=True, + location=output.name + "_data", + size_threshold=1024, ) - if allow_fp8_conv: - if not fp8_conv_pairs: - errors.append("no initializer-backed FP8 Conv weight Q/DQ remains") - errors.extend(_validate_normalized_fp8_qdq(onnx_model, fp8_conv_pairs)) - if expected_fp8_conv_count is not None and len(fp8_conv_pairs) != expected_fp8_conv_count: - errors.append( - f"found {len(fp8_conv_pairs)} initializer-backed FP8 Conv weight Q/DQ pairs, " - f"expected {expected_fp8_conv_count} enabled Conv2d pairs" - ) - if disallowed_qdq_weights: - errors.append( - "disallowed initializer-backed Q/DQ weight consumers: " - + ", ".join(disallowed_qdq_weights[:5]) - ) - - if errors: - raise ValueError("Invalid final FP4 ONNX graph: " + "; ".join(errors)) + print(f"ONNX model saved to {output}") def _normalize_fp8_qdq(onnx_model): @@ -1205,7 +493,7 @@ def _normalize_fp8_qdq(onnx_model): graph.cleanup().toposort() onnx_model = convert_zp_fp8(gs.export_onnx(graph)) graph = gs.import_onnx(onnx_model) - return gs.export_onnx(graph.cleanup().toposort()) + return gs.export_onnx(graph.cleanup()) def _ensure_default_opset(onnx_model, minimum_version): @@ -1213,207 +501,87 @@ def _ensure_default_opset(onnx_model, minimum_version): if opset_import.domain in {"", "ai.onnx"}: opset_import.version = max(opset_import.version, minimum_version) return + opset_import = onnx_model.opset_import.add() opset_import.domain = "" opset_import.version = minimum_version -def _process_fp4_onnx_graph( - onnx_model, - model_name, - expected_block_size=16, - *, - expected_linear_count=None, - expected_fp8_conv_count=None, -): - allow_fp8_conv = model_name in {"sdxl-1.0", "sdxl-turbo"} - expected_weight_count = _validate_raw_fp4_graph( - onnx_model, - expected_block_size, - allow_fp8_conv=allow_fp8_conv, - expected_linear_count=expected_linear_count, - expected_fp8_conv_count=expected_fp8_conv_count, - ) - if allow_fp8_conv: +def _process_fp4_onnx_graph(onnx_model, model_name): + if model_name in {"sdxl-1.0", "sdxl-turbo"}: onnx_model = _normalize_fp8_qdq(onnx_model) onnx_model = NVFP4QuantExporter.process_model(onnx_model) - _ensure_default_opset(onnx_model, 23) - _validate_final_fp4_graph( - onnx_model, - expected_weight_count, - expected_block_size, - allow_fp8_conv=allow_fp8_conv, - expected_fp8_conv_count=expected_fp8_conv_count, - ) + if model_name in {"sdxl-1.0", "sdxl-turbo"}: + _ensure_default_opset(onnx_model, 23) return onnx_model -def _get_sdxl_fp4_expected_counts(backbone): - linear_count = 0 - conv_count = 0 - for module in backbone.modules(): - input_quantizer = getattr(module, "input_quantizer", None) - weight_quantizer = getattr(module, "weight_quantizer", None) - pair_enabled = all( - quantizer is not None and getattr(quantizer, "is_enabled", False) - for quantizer in (input_quantizer, weight_quantizer) - ) - if not pair_enabled: - continue - if isinstance(module, (torch.nn.Linear, RealQuantLinear)): - linear_count += 1 - elif isinstance(module, torch.nn.Conv2d): - conv_count += 1 - return linear_count, conv_count - +def modelopt_export_sd(backbone, onnx_dir, model_name, precision): + model_file_name = "model.onnx" + os.makedirs(f"{onnx_dir}", exist_ok=True) + q_output = Path(f"{onnx_dir}/{model_file_name}") + is_sdxl_fp4 = precision == "fp4" and model_name in {"sdxl-1.0", "sdxl-turbo"} -def save_onnx(onnx_model, output, external_data_name=None): - onnx.save( - onnx_model, - str(output), - save_as_external_data=True, - all_tensors_to_one_file=True, - location=external_data_name or output.name + "_data", - size_threshold=1024, + quantizer_context = ( + configure_linear_module_onnx_quantizers(backbone) if precision == "fp4" else nullcontext() + ) + fp8_scale_context = ( + _temporary_fp8_export_scales(backbone, conv_only=is_sdxl_fp4) + if is_sdxl_fp4 or (precision == "fp8" and _has_enabled_conv(backbone)) + else nullcontext() ) - print(f"ONNX model saved to {output}") + dummy_kwargs, dynamic_axes, _ = generate_dummy_kwargs_and_dynamic_axes_and_shapes( + model_name, backbone + ) -def _get_external_data_paths(output): - fallback = output.with_name(output.name + "_data") - if not output.exists(): - return set() - try: - onnx_model = onnx.load(str(output), load_external_data=False) - except Exception: - return {fallback} if fallback.exists() else set() - - output_parent = output.parent.resolve() - paths = set() - for initializer in onnx_model.graph.initializer: - for entry in initializer.external_data: - if entry.key != "location": - continue - path = (output.parent / entry.value).resolve() - if path.parent == output_parent and path != output.resolve(): - paths.add(path) - return paths - - -def _save_onnx_atomically(onnx_model, output): - staging_dir = Path(tempfile.mkdtemp(prefix=".modelopt-export-", dir=output.parent)) - staged_output = staging_dir / output.name - external_data_name = f"{output.name}_data.{uuid.uuid4().hex}" - staged_data = staging_dir / external_data_name - published_data = output.parent / external_data_name - old_data_paths = _get_external_data_paths(output) - had_previous_output = output.exists() - previous_output = staging_dir / "previous-model.onnx" - try: - if had_previous_output: - shutil.copy2(output, previous_output) - save_onnx(onnx_model, staged_output, external_data_name=external_data_name) - onnx.checker.check_model(str(staged_output)) - has_external_data = staged_data.exists() - try: - if has_external_data: - os.replace(staged_data, published_data) - os.replace(staged_output, output) - except BaseException: - if previous_output.exists(): - os.replace(previous_output, output) - elif not had_previous_output: - output.unlink(missing_ok=True) - published_data.unlink(missing_ok=True) - raise - - for old_data_path in old_data_paths - {published_data.resolve()}: - with suppress(OSError): - old_data_path.unlink(missing_ok=True) - finally: - shutil.rmtree(staging_dir, ignore_errors=True) + if model_name in ["sdxl-1.0", "sdxl-turbo"]: + input_names = ["sample", "timestep", "encoder_hidden_states", "text_embeds", "time_ids"] + output_names = ["latent"] + elif model_name == "sd3-medium": + input_names = ["hidden_states", "encoder_hidden_states", "pooled_projections", "timestep"] + output_names = ["sample"] + elif model_name == "sd3.5-medium": + input_names = ["hidden_states", "encoder_hidden_states", "pooled_projections", "timestep"] + output_names = ["out_hidden_states"] + elif model_name in ["flux-dev", "flux-schnell"]: + input_names = [ + "hidden_states", + "encoder_hidden_states", + "pooled_projections", + "timestep", + "img_ids", + "txt_ids", + ] + if model_name == "flux-dev": + input_names.append("guidance") + output_names = ["latent"] + elif model_name == "ltx-video-dev": + input_names = [ + "hidden_states", + "encoder_hidden_states", + "timestep", + "encoder_attention_mask", + "video_coords", + ] + output_names = ["latent"] + elif model_name == "wan2.2-t2v-14b": + input_names = [ + "hidden_states", + "timestep", + "encoder_hidden_states", + ] + output_names = ["latent"] + else: + raise NotImplementedError(f"Unsupported model_id: {model_name}") + do_constant_folding = True + opset_version = 20 -def modelopt_export_sd(backbone, onnx_dir, model_name, precision, expected_fp4_block_size=16): - model_file_name = "model.onnx" - os.makedirs(f"{onnx_dir}", exist_ok=True) - tmp_subfolder = tempfile.mkdtemp(prefix=".modelopt-raw-", dir=onnx_dir) + tmp_subfolder = tempfile.mkdtemp(prefix="myapp_") tmp_output = Path(f"{tmp_subfolder}/{model_file_name}") - q_output = Path(f"{onnx_dir}/{model_file_name}") - strict_sdxl_fp4 = precision == "fp4" and model_name in {"sdxl-1.0", "sdxl-turbo"} - expected_linear_count = None - expected_fp8_conv_count = None - if strict_sdxl_fp4: - expected_linear_count, expected_fp8_conv_count = _get_sdxl_fp4_expected_counts(backbone) - try: - quantizer_context = ( - configure_linear_module_onnx_quantizers(backbone) - if precision == "fp4" - else nullcontext() - ) - - dummy_kwargs, dynamic_axes, _ = generate_dummy_kwargs_and_dynamic_axes_and_shapes( - model_name, backbone - ) - - if model_name in ["sdxl-1.0", "sdxl-turbo"]: - input_names = [ - "sample", - "timestep", - "encoder_hidden_states", - "text_embeds", - "time_ids", - ] - output_names = ["latent"] - elif model_name == "sd3-medium": - input_names = [ - "hidden_states", - "encoder_hidden_states", - "pooled_projections", - "timestep", - ] - output_names = ["sample"] - elif model_name == "sd3.5-medium": - input_names = [ - "hidden_states", - "encoder_hidden_states", - "pooled_projections", - "timestep", - ] - output_names = ["out_hidden_states"] - elif model_name in ["flux-dev", "flux-schnell"]: - input_names = [ - "hidden_states", - "encoder_hidden_states", - "pooled_projections", - "timestep", - "img_ids", - "txt_ids", - ] - if model_name == "flux-dev": - input_names.append("guidance") - output_names = ["latent"] - elif model_name == "ltx-video-dev": - input_names = [ - "hidden_states", - "encoder_hidden_states", - "timestep", - "encoder_attention_mask", - "video_coords", - ] - output_names = ["latent"] - elif model_name == "wan2.2-t2v-14b": - input_names = [ - "hidden_states", - "timestep", - "encoder_hidden_states", - ] - output_names = ["latent"] - else: - raise NotImplementedError(f"Unsupported model_id: {model_name}") - - with quantizer_context, torch.inference_mode(): + with quantizer_context, fp8_scale_context, torch.inference_mode(): onnx_export( backbone, (), @@ -1422,8 +590,8 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision, expected_fp4_b input_names=input_names, output_names=output_names, dynamic_axes=dynamic_axes, - do_constant_folding=True, - opset_version=20, + do_constant_folding=do_constant_folding, + opset_version=opset_version, dynamo=False, ) print(f"Saved at {tmp_output}") @@ -1434,19 +602,7 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision, expected_fp4_b else: flux_convert_rope_weight_type(onnx_model) if precision == "fp4": - if strict_sdxl_fp4: - onnx_model = _process_fp4_onnx_graph( - onnx_model, - model_name, - expected_fp4_block_size, - expected_linear_count=expected_linear_count, - expected_fp8_conv_count=expected_fp8_conv_count, - ) - else: - onnx_model = NVFP4QuantExporter.process_model(onnx_model) - if strict_sdxl_fp4: - _save_onnx_atomically(onnx_model, q_output) - else: - save_onnx(onnx_model, q_output) + onnx_model = _process_fp4_onnx_graph(onnx_model, model_name) + save_onnx(onnx_model, q_output) finally: shutil.rmtree(tmp_subfolder, ignore_errors=True) diff --git a/examples/diffusers/quantization/quantize.py b/examples/diffusers/quantization/quantize.py index 5d3902610e3..f423875ad42 100644 --- a/examples/diffusers/quantization/quantize.py +++ b/examples/diffusers/quantization/quantize.py @@ -51,12 +51,7 @@ QuantFormat, QuantizationConfig, ) -from utils import ( - check_conv_and_mha, - check_lora, - validate_fp8_mha_quantizers, - validate_nvfp4_quantizers, -) +from utils import check_conv_and_mha, check_lora import modelopt.torch.opt as mto import modelopt.torch.quantization as mtq @@ -281,25 +276,6 @@ def __init__( self.logger = logger self.pipeline_manager = pipeline_manager - def _has_fp8_conv_layers(self, model: torch.nn.Module) -> bool: - """Check whether the model contains an enabled FP8 convolution.""" - for module in model.modules(): - if not isinstance(module, torch.nn.Conv1d | torch.nn.Conv2d | torch.nn.Conv3d): - continue - - quantizers = ( - getattr(module, "input_quantizer", None), - getattr(module, "weight_quantizer", None), - ) - if any( - quantizer is not None - and quantizer.is_enabled - and getattr(quantizer, "is_fp8", False) - for quantizer in quantizers - ): - return True - return False - def save_checkpoint( self, backbone: torch.nn.Module, @@ -332,7 +308,6 @@ def export_onnx( backbone: torch.nn.Module, model_type: ModelType, quant_format: QuantFormat, - fp4_block_size: int = 16, ) -> None: """ Export model to ONNX format. @@ -342,52 +317,27 @@ def export_onnx( backbone: Model backbone model_type: Type of model quant_format: Quantization format - fp4_block_size: Expected NVFP4 block size """ if not self.config.onnx_dir: return # 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, restore_fp8_scales + from onnx_utils.export import modelopt_export_sd self.logger.info(f"Starting ONNX export to {self.config.onnx_dir}") - quantizer_states = [] - try: - uses_fp8_conv_workaround = quant_format == QuantFormat.FP8 or ( - quant_format == QuantFormat.FP4 and model_type in _SDXL_MODEL_TYPES + self.logger.info("Preparing models for export...") + pipe.to("cpu") + torch.cuda.empty_cache() + backbone.to("cuda") + # Export to ONNX + backbone.eval() + with torch.no_grad(): + self.logger.info("Exporting to ONNX...") + modelopt_export_sd( + backbone, str(self.config.onnx_dir), model_type.value, quant_format.value ) - if uses_fp8_conv_workaround and self._has_fp8_conv_layers(backbone): - self.logger.info( - "Detected quantizing conv layers in backbone. Generating FP8 scales..." - ) - if quant_format == QuantFormat.FP4: - quantizer_states = generate_fp8_scales(backbone, conv_only=True) - else: - quantizer_states = generate_fp8_scales(backbone) - self.logger.info("Preparing models for export...") - pipe.to("cpu") - torch.cuda.empty_cache() - backbone.to("cuda") - # Export to ONNX - backbone.eval() - with torch.no_grad(): - self.logger.info("Exporting to ONNX...") - export_kwargs = ( - {"expected_fp4_block_size": fp4_block_size} - if quant_format == QuantFormat.FP4 - else {} - ) - modelopt_export_sd( - backbone, - str(self.config.onnx_dir), - model_type.value, - quant_format.value, - **export_kwargs, - ) - finally: - restore_fp8_scales(quantizer_states) self.logger.info("ONNX export completed successfully") @@ -633,29 +583,17 @@ def create_argument_parser() -> argparse.ArgumentParser: return parser -def _finalize_backbone_quantization( - backbone: torch.nn.Module, - backbone_name: str, +def _restore_sdxl_fp4_policy( + pipeline_manager: PipelineManager, quant_config: QuantizationConfig, model_type: ModelType, - restored: bool, ) -> None: - if backbone_name in ("video_decoder", "vae"): + if quant_config.format != QuantFormat.FP4 or model_type not in _SDXL_MODEL_TYPES: return - is_sdxl_fp4 = quant_config.format == QuantFormat.FP4 and model_type in _SDXL_MODEL_TYPES - if restored and not is_sdxl_fp4: - return - if is_sdxl_fp4 and restored: - validate_fp8_mha_quantizers(backbone, quant_config.quantize_mha) - check_conv_and_mha(backbone, quant_config.format == QuantFormat.FP4, quant_config.quantize_mha) - if is_sdxl_fp4: - validate_nvfp4_quantizers( - backbone, - quant_config.block_size, - quant_config.quantize_mha, - validate_sdxl_mixed_recipe=True, - ) + for backbone_name, backbone in pipeline_manager.iter_backbones(): + if backbone_name not in ("video_decoder", "vae"): + check_conv_and_mha(backbone, False, quant_config.quantize_mha) def main() -> None: @@ -743,9 +681,9 @@ def main() -> None: export_manager = ExportManager(export_config, logger, pipeline_manager) - restored = bool(export_config.restore_from and export_config.restore_from.exists()) - if restored: + if export_config.restore_from and export_config.restore_from.exists(): export_manager.restore_checkpoint() + _restore_sdxl_fp4_policy(pipeline_manager, quant_config, model_config.model_type) else: logger.info("Initializing calibration...") @@ -768,33 +706,24 @@ def forward_loop(mod): forward_loop, backbone_name=backbone_name, ) + + # Compress model weights if requested (only for FP8/FP4) if quant_config.compress: logger.info(f"Compressing {backbone_name} weights...") mtq.compress(backbone) logger.info(f"{backbone_name} compression completed") - _finalize_backbone_quantization( - backbone, - backbone_name, - quant_config, - model_config.model_type, - restored=False, - ) - export_manager.save_checkpoint(backbone, backbone_name) + # 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 + and model_config.model_type not in _SDXL_MODEL_TYPES, + quant_config.quantize_mha, + ) - if ( - restored - and quant_config.format == QuantFormat.FP4 - and model_config.model_type in _SDXL_MODEL_TYPES - ): - for backbone_name, backbone in pipeline_manager.iter_backbones(): - _finalize_backbone_quantization( - backbone, - backbone_name, - quant_config, - model_config.model_type, - restored=True, - ) + export_manager.save_checkpoint(backbone, backbone_name) pipeline_manager.print_quant_summary() @@ -804,7 +733,6 @@ def forward_loop(mod): backbone, model_config.model_type, quant_config.format, - fp4_block_size=quant_config.block_size, ) export_manager.export_hf_ckpt(pipe, model_config) diff --git a/examples/diffusers/quantization/utils.py b/examples/diffusers/quantization/utils.py index 325179accc9..c3cfdcd5cdd 100644 --- a/examples/diffusers/quantization/utils.py +++ b/examples/diffusers/quantization/utils.py @@ -25,8 +25,6 @@ from diffusers.utils import load_image import modelopt.torch.quantization as mtq -from modelopt.torch.quantization.nn import TensorQuantizer -from modelopt.torch.quantization.nn.modules.quant_linear import RealQuantLinear from modelopt.torch.quantization.plugins.diffusion.diffusers import AttentionModuleMixin USE_PEFT = True @@ -46,37 +44,24 @@ def filter_func_default(name: str) -> bool: return pattern.match(name) is not None -_MHA_QUANTIZER_NAMES = ( - "q_bmm_quantizer", - "k_bmm_quantizer", - "v_bmm_quantizer", - "softmax_quantizer", - "bmm2_output_quantizer", -) -_REQUIRED_MHA_QUANTIZER_NAMES = _MHA_QUANTIZER_NAMES[:4] -_SDXL_FP16_PROJECTION_NAMES = frozenset(("to_q", "to_k", "to_v")) - - def check_conv_and_mha(backbone, if_fp4, quantize_mha): for name, module in backbone.named_modules(): - if isinstance(module, torch.nn.Conv1d | torch.nn.Conv2d) and if_fp4: - nvfp4_quantizers = [ - quantizer - for quantizer in ( - getattr(module, "weight_quantizer", None), - getattr(module, "input_quantizer", None), - ) - if isinstance(quantizer, TensorQuantizer) and quantizer.is_nvfp4_dynamic - ] - for quantizer in nvfp4_quantizers: - quantizer.disable() - if nvfp4_quantizers: - print(f"Disabled NVFP4 Conv layer quantization for layer {name}") - - elif isinstance(module, Attention | AttentionModuleMixin): + if isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d)) and if_fp4: + module.weight_quantizer.disable() + module.input_quantizer.disable() + + print(f"Disabled NVFP4 Conv layer quantization for layer {name}") + + elif isinstance(module, (Attention, AttentionModuleMixin)): head_size = int(module.inner_dim / module.heads) if not quantize_mha or head_size % 16 != 0: - for attr in _MHA_QUANTIZER_NAMES: + for attr in ( + "q_bmm_quantizer", + "k_bmm_quantizer", + "v_bmm_quantizer", + "softmax_quantizer", + "bmm2_output_quantizer", + ): if hasattr(module, attr): getattr(module, attr).disable() setattr(module, "_disable_fp8_mha", True) @@ -86,213 +71,6 @@ def check_conv_and_mha(backbone, if_fp4, quantize_mha): setattr(module, "_disable_fp8_mha", False) -def _validate_finite_positive_amax(name, quantizer): - amax = quantizer.amax - if amax is None or amax.numel() != 1 or not torch.isfinite(amax).all() or not (amax > 0).all(): - raise ValueError(f"Quantizer '{name}' must have a finite positive calibrated amax.") - - -def _validate_finite_nonnegative_amax(name, quantizer): - amax = quantizer.amax - if amax is None or amax.numel() != 1 or not torch.isfinite(amax).all() or not (amax >= 0).all(): - raise ValueError(f"Quantizer '{name}' must have a finite nonnegative calibrated amax.") - - -def _validate_calibrated_fp8_quantizer(name, quantizer): - if not quantizer.is_fp8: - raise ValueError(f"Quantizer '{name}' must use per-tensor FP8.") - _validate_finite_positive_amax(name, quantizer) - - -def validate_fp8_mha_quantizers(backbone, quantize_mha): - """Validate that the restored or finalized FP8 MHA state matches its policy.""" - for name, module in backbone.named_modules(): - if not isinstance(module, Attention | AttentionModuleMixin): - continue - - head_size = int(module.inner_dim / module.heads) - mha_enabled = quantize_mha and head_size % 16 == 0 - for attr in _MHA_QUANTIZER_NAMES: - quantizer = getattr(module, attr, None) - if quantizer is None: - if mha_enabled: - expected_state = ( - "present and enabled" - if attr in _REQUIRED_MHA_QUANTIZER_NAMES - else "present and disabled" - ) - raise ValueError( - f"FP8 MHA for attention '{name}' requires '{attr}' to be {expected_state}." - ) - continue - if not isinstance(quantizer, TensorQuantizer): - raise ValueError( - f"Attention '{name}.{attr}' must be a TensorQuantizer, got " - f"{type(quantizer).__name__}." - ) - if not mha_enabled: - if quantizer.is_enabled: - reason = "disabled by configuration" if not quantize_mha else "unsupported" - raise ValueError( - f"Attention '{name}.{attr}' must be disabled because FP8 MHA is {reason}." - ) - continue - if attr in _REQUIRED_MHA_QUANTIZER_NAMES and not quantizer.is_enabled: - raise ValueError(f"FP8 MHA for attention '{name}' requires '{attr}' to be enabled.") - if attr == "bmm2_output_quantizer" and quantizer.is_enabled: - raise ValueError( - f"FP8 MHA for attention '{name}' requires '{attr}' to be disabled." - ) - if quantizer.is_enabled: - qualified_name = f"{name}.{attr}" - if attr == "softmax_quantizer": - if not quantizer.is_fp8: - raise ValueError(f"Quantizer '{qualified_name}' must use per-tensor FP8.") - else: - _validate_calibrated_fp8_quantizer(qualified_name, quantizer) - - -def _validate_fp4_quantizer_placement(backbone, quantize_mha, allow_fp8_conv): - for module_name, module in backbone.named_modules(): - for quantizer_name, quantizer in module.named_children(): - if not isinstance(quantizer, TensorQuantizer) or not quantizer.is_enabled: - continue - qualified_name = f"{module_name}.{quantizer_name}".lstrip(".") - if quantizer.is_nvfp4_dynamic: - if not isinstance( - module, torch.nn.Linear | RealQuantLinear - ) or quantizer_name not in ( - "input_quantizer", - "weight_quantizer", - ): - raise ValueError( - f"Enabled NVFP4 quantizer '{qualified_name}' is only supported on Linear " - "input and weight quantizers." - ) - continue - - if quantizer.is_fp8: - is_conv_quantizer = ( - allow_fp8_conv - and isinstance(module, torch.nn.Conv2d) - and quantizer_name in ("input_quantizer", "weight_quantizer") - ) - is_mha_quantizer = ( - quantize_mha - and isinstance(module, Attention | AttentionModuleMixin) - and quantizer_name in _MHA_QUANTIZER_NAMES - ) - if is_conv_quantizer or is_mha_quantizer: - continue - raise ValueError( - f"Enabled FP8 quantizer '{qualified_name}' is only supported on SDXL Conv2d " - "input/weight quantizers or opt-in MHA quantizers." - ) - - raise ValueError( - f"Enabled quantizer '{qualified_name}' has an unsupported format for FP4 export." - ) - - -def validate_nvfp4_quantizers( - backbone, expected_block_size, quantize_mha, validate_sdxl_mixed_recipe=False -): - """Validate the quantizer state required by NVFP4 ONNX export.""" - enabled_linear_pairs = 0 - for name, module in backbone.named_modules(): - if not isinstance(module, torch.nn.Linear | RealQuantLinear): - continue - - input_quantizer = getattr(module, "input_quantizer", None) - weight_quantizer = getattr(module, "weight_quantizer", None) - if input_quantizer is None and weight_quantizer is None: - continue - if not isinstance(input_quantizer, TensorQuantizer) or not isinstance( - weight_quantizer, TensorQuantizer - ): - raise ValueError( - f"NVFP4 Linear '{name}' must use TensorQuantizer instances for both input and weight." - ) - - input_enabled = input_quantizer.is_enabled - weight_enabled = weight_quantizer.is_enabled - if validate_sdxl_mixed_recipe and name.rsplit(".", 1)[-1] in _SDXL_FP16_PROJECTION_NAMES: - if input_enabled or weight_enabled: - raise ValueError( - f"SDXL attention projection '{name}' must keep input and weight quantizers " - "disabled for the NVFP4 mixed recipe. Recalibrate with the current SDXL " - "FP4 recipe." - ) - continue - - if not input_enabled and not weight_enabled: - continue - if input_enabled != weight_enabled: - raise ValueError( - f"NVFP4 Linear '{name}' must enable input and weight quantizers as a pair." - ) - if not input_quantizer.is_nvfp4_dynamic or not weight_quantizer.is_nvfp4_dynamic: - raise ValueError( - f"NVFP4 Linear '{name}' must use dynamic E2M1 quantizers with FP8 block scales " - "for both input and weight." - ) - _validate_finite_nonnegative_amax(f"{name}.input_quantizer", input_quantizer) - _validate_finite_nonnegative_amax(f"{name}.weight_quantizer", weight_quantizer) - - input_block_size = input_quantizer.block_sizes.get(-1) - weight_block_size = weight_quantizer.block_sizes.get(-1) - if input_block_size != expected_block_size or weight_block_size != expected_block_size: - raise ValueError( - f"NVFP4 Linear '{name}' requires block size {expected_block_size}; got " - f"input={input_block_size}, weight={weight_block_size}." - ) - enabled_linear_pairs += 1 - - if enabled_linear_pairs == 0: - raise ValueError( - "NVFP4 quantization requires at least one enabled Linear input/weight pair." - ) - - if validate_sdxl_mixed_recipe: - enabled_conv_pairs = 0 - for name, module in backbone.named_modules(): - if not isinstance(module, torch.nn.Conv2d): - continue - - input_quantizer = getattr(module, "input_quantizer", None) - weight_quantizer = getattr(module, "weight_quantizer", None) - if input_quantizer is None and weight_quantizer is None: - continue - if not isinstance(input_quantizer, TensorQuantizer) or not isinstance( - weight_quantizer, TensorQuantizer - ): - raise ValueError( - f"SDXL FP8 Conv2d '{name}' must use TensorQuantizer instances for both input " - "and weight." - ) - - input_enabled = input_quantizer.is_enabled - weight_enabled = weight_quantizer.is_enabled - if not input_enabled and not weight_enabled: - continue - if input_enabled != weight_enabled: - raise ValueError( - f"SDXL FP8 Conv2d '{name}' must enable input and weight quantizers as a pair." - ) - _validate_calibrated_fp8_quantizer(f"{name}.input_quantizer", input_quantizer) - _validate_calibrated_fp8_quantizer(f"{name}.weight_quantizer", weight_quantizer) - enabled_conv_pairs += 1 - - if enabled_conv_pairs == 0: - raise ValueError( - "SDXL NVFP4 quantization requires at least one enabled calibrated FP8 Conv2d " - "input/weight pair." - ) - - validate_fp8_mha_quantizers(backbone, quantize_mha) - _validate_fp4_quantizer_placement(backbone, quantize_mha, validate_sdxl_mixed_recipe) - - def filter_func_ltx_video(name: str) -> bool: """Filter function specifically for LTX-Video models.""" pattern = re.compile( diff --git a/tests/examples/diffusers/test_diffusers.py b/tests/examples/diffusers/test_diffusers.py index b6daa51d49c..ff73093b5a7 100644 --- a/tests/examples/diffusers/test_diffusers.py +++ b/tests/examples/diffusers/test_diffusers.py @@ -128,17 +128,6 @@ def inference(self, tmp_path: Path) -> None: ), marks=minimum_sm(89), ), - pytest.param( - DiffuserModel( - name="flux-schnell", - path=FLUX_SCHNELL_PATH, - dtype="BFloat16", - format_type="fp4", - quant_algo="max", - collect_method="default", - ), - marks=minimum_sm(100), - ), pytest.param( DiffuserModel( name="sdxl-1.0", @@ -163,7 +152,6 @@ def inference(self, tmp_path: Path) -> None: "flux_schnell_bf16_int8_smoothquant_3.0_min_mean", "sd3_medium_fp16_int8_smoothquant_3.0_min_mean", "sdxl_1.0_fp16_fp8_max_3.0_default", - "flux_schnell_bf16_fp4_max_3.0_default", "sdxl_1.0_fp16_fp4_max_3.0_default", "sdxl_1.0_fp16_int8_smoothquant_3.0_min_mean", ], diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py new file mode 100644 index 00000000000..edb6834f652 --- /dev/null +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -0,0 +1,243 @@ +# 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 logging +import sys +from pathlib import Path + +import numpy as np +import pytest +import torch +from torch import nn + +onnx = pytest.importorskip("onnx") +pytest.importorskip("onnx_graphsurgeon") +pytest.importorskip("diffusers") +from diffusers.models.attention_processor import Attention +from onnx import TensorProto, helper, numpy_helper + +import modelopt.torch.quantization as mtq +from modelopt.torch.quantization.config import QuantizerAttributeConfig +from modelopt.torch.quantization.nn import TensorQuantizer + +_QUANTIZATION_EXAMPLE = ( + Path(__file__).resolve().parents[3] / "examples" / "diffusers" / "quantization" +) +sys.path.insert(0, str(_QUANTIZATION_EXAMPLE)) + +from models_utils import ModelType +from quantize import Quantizer, _restore_sdxl_fp4_policy +from quantize_config import ModelConfig, QuantFormat, QuantizationConfig + +from examples.diffusers.quantization.onnx_utils import export as diffusion_export + + +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) + + +@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 + + +def test_restore_reapplies_sdxl_fp4_mha_policy(): + attention = Attention(query_dim=16, heads=1, dim_head=16) + skipped_attention = Attention(query_dim=16, heads=1, dim_head=16) + attention._disable_fp8_mha = True + skipped_attention._disable_fp8_mha = True + + class _PipelineManager: + def iter_backbones(self): + return (("unet", attention), ("vae", skipped_attention)) + + _restore_sdxl_fp4_policy( + _PipelineManager(), + QuantizationConfig(format=QuantFormat.FP4, quantize_mha=True), + ModelType.SDXL_BASE, + ) + + assert not attention._disable_fp8_mha + assert skipped_attention._disable_fp8_mha + + +def _fp8_quantizer(*, enabled=True): + quantizer = TensorQuantizer(QuantizerAttributeConfig(num_bits=(4, 3), axis=None)) + quantizer.amax = torch.tensor(448.0) + if not enabled: + quantizer.disable() + return quantizer + + +def _add_fp8_quantizers(module, *, enabled=True): + module.input_quantizer = _fp8_quantizer(enabled=enabled) + module.weight_quantizer = _fp8_quantizer(enabled=enabled) + + +@pytest.mark.parametrize("raises", [False, True]) +def test_temporary_fp8_conv_export_scales_restore_state(raises): + model = nn.Module() + model.conv = nn.Conv2d(1, 1, 1) + model.disabled_conv = nn.Conv2d(1, 1, 1) + model.linear = nn.Linear(1, 1) + _add_fp8_quantizers(model.conv) + _add_fp8_quantizers(model.disabled_conv, enabled=False) + _add_fp8_quantizers(model.linear) + linear_only = nn.Sequential(nn.Linear(1, 1)) + _add_fp8_quantizers(linear_only[0]) + assert not diffusion_export._has_enabled_conv(linear_only) + assert diffusion_export._has_enabled_conv(model) + changed = (model.conv.input_quantizer, model.conv.weight_quantizer) + unchanged = ( + model.disabled_conv.input_quantizer, + model.disabled_conv.weight_quantizer, + model.linear.input_quantizer, + model.linear.weight_quantizer, + ) + original_state = { + quantizer: (quantizer._num_bits, quantizer._amax) for quantizer in changed + unchanged + } + + def run_export(): + with diffusion_export._temporary_fp8_export_scales(model, conv_only=True): + for quantizer in changed: + assert quantizer.num_bits == 8 + assert quantizer.amax == 127.0 + for quantizer in unchanged: + assert (quantizer._num_bits, quantizer._amax) == original_state[quantizer] + if raises: + raise RuntimeError("export failed") + + if raises: + with pytest.raises(RuntimeError, match="export failed"): + run_export() + else: + run_export() + + for quantizer, (num_bits, amax) in original_state.items(): + assert quantizer._num_bits == num_bits + assert quantizer._amax is amax + + +def _make_mixed_fp4_fp8_model(): + fp4_weight = numpy_helper.from_array( + np.linspace(-1.0, 1.0, 16 * 16, dtype=np.float16).reshape(16, 16), "fp4_weight" + ) + fp8_weight = numpy_helper.from_array(np.ones((1, 1, 1, 1), dtype=np.float16), "fp8_weight") + scale = numpy_helper.from_array(np.array(0.25, dtype=np.float16), "fp8_scale_value") + zero = numpy_helper.from_array(np.array(0, dtype=np.int8), "fp8_zero_value") + nodes = [ + helper.make_node( + "TRT_FP4QDQ", + ["fp4_weight"], + ["fp4_weight_dq"], + name="fp4_weight_qdq", + domain="trt", + block_size=16, + ), + helper.make_node( + "MatMul", ["linear_input", "fp4_weight_dq"], ["linear_output"], name="fp4_matmul" + ), + helper.make_node("Constant", [], ["fp8_scale"], name="fp8_scale", value=scale), + helper.make_node("Constant", [], ["fp8_zero"], name="fp8_zero", value=zero), + helper.make_node( + "QuantizeLinear", + ["fp8_weight", "fp8_scale", "fp8_zero"], + ["fp8_weight_q"], + name="fp8_weight_quantize", + ), + helper.make_node( + "DequantizeLinear", + ["fp8_weight_q", "fp8_scale", "fp8_zero"], + ["fp8_weight_dq"], + name="fp8_weight_dequantize", + ), + helper.make_node( + "QuantizeLinear", + ["conv_input", "fp8_scale", "fp8_zero"], + ["conv_input_q"], + name="fp8_activation_quantize", + ), + helper.make_node( + "DequantizeLinear", + ["conv_input_q", "fp8_scale", "fp8_zero"], + ["conv_input_dq"], + name="fp8_activation_dequantize", + ), + helper.make_node( + "Conv", ["conv_input_dq", "fp8_weight_dq"], ["conv_output"], name="fp8_conv" + ), + ] + graph = helper.make_graph( + nodes, + "mixed_fp4_fp8", + [ + helper.make_tensor_value_info("linear_input", TensorProto.FLOAT16, [1, 16]), + helper.make_tensor_value_info("conv_input", TensorProto.FLOAT16, [1, 1, 2, 2]), + ], + [ + helper.make_tensor_value_info("linear_output", TensorProto.FLOAT16, [1, 16]), + helper.make_tensor_value_info("conv_output", TensorProto.FLOAT16, [1, 1, 2, 2]), + ], + [fp4_weight, fp8_weight], + value_info=[helper.make_tensor_value_info("fp4_weight_dq", TensorProto.FLOAT16, [16, 16])], + ) + return helper.make_model( + graph, + opset_imports=[helper.make_opsetid("", 20), helper.make_opsetid("trt", 1)], + ) + + +def _constant_dtype(model, name): + node = next(node for node in model.graph.node if node.name == name) + return next(attribute for attribute in node.attribute if attribute.name == "value").t.data_type + + +def test_mixed_sdxl_fp4_graph_postprocessing(): + model = diffusion_export._process_fp4_onnx_graph(_make_mixed_fp4_fp8_model(), "sdxl-1.0") + + assert not any(node.op_type == "TRT_FP4QDQ" for node in model.graph.node) + assert any(tensor.data_type == TensorProto.FLOAT4E2M1 for tensor in model.graph.initializer) + assert _constant_dtype(model, "fp8_zero") == TensorProto.FLOAT8E4M3FN + assert any(node.op_type == "Conv" and node.name == "fp8_conv" for node in model.graph.node) + assert sum(node.op_type == "QuantizeLinear" for node in model.graph.node) == 2 + assert next(opset.version for opset in model.opset_import if not opset.domain) >= 23 + onnx.checker.check_model(model) diff --git a/tests/unit/examples/test_diffusers_fp4_onnx_validation.py b/tests/unit/examples/test_diffusers_fp4_onnx_validation.py deleted file mode 100644 index 3578da5096b..00000000000 --- a/tests/unit/examples/test_diffusers_fp4_onnx_validation.py +++ /dev/null @@ -1,992 +0,0 @@ -# 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. - -from contextlib import nullcontext -from pathlib import Path - -import numpy as np -import pytest -import torch - -onnx = pytest.importorskip("onnx") -pytest.importorskip("onnx_graphsurgeon") -pytest.importorskip("diffusers") -from onnx import TensorProto, helper, numpy_helper - -from examples.diffusers.quantization.onnx_utils import export as diffusion_export -from modelopt.onnx.export import NVFP4QuantExporter - - -class _Quantizer: - def __init__(self, num_bits, *, enabled=True, amax=448.0): - self._num_bits = num_bits - self._amax = torch.tensor(amax) - self.is_enabled = enabled - - @property - def num_bits(self): - return self._num_bits - - -def _make_raw_fp4_model(*, qdq_consumer="Conv", marker_block_size=16, marker_initializer=True): - dynamic_scale = numpy_helper.from_array( - np.array(1.0 / 2688.0, dtype=np.float32), "dynamic_scale_value" - ) - nodes = [ - helper.make_node( - "Constant", - [], - ["dynamic_quantize_scale"], - name="dynamic_quantize_scale", - value=dynamic_scale, - ), - helper.make_node( - "Constant", - [], - ["dynamic_dequantize_scale"], - name="dynamic_dequantize_scale", - value=dynamic_scale, - ), - ] - initializers = [] - inputs = [helper.make_tensor_value_info("activation", TensorProto.FLOAT16, [1, 16])] - outputs = [helper.make_tensor_value_info("linear_output", TensorProto.FLOAT16, [1, 16])] - value_info = [helper.make_tensor_value_info("fp4_weight_dq", TensorProto.FLOAT16, [16, 16])] - - if marker_initializer: - initializers.append( - numpy_helper.from_array( - np.linspace(-1.0, 1.0, 16 * 16, dtype=np.float16).reshape(16, 16), - "fp4_weight", - ) - ) - nodes.extend( - [ - helper.make_node( - "TRT_FP4DynamicQuantize", - ["activation", "dynamic_quantize_scale"], - ["activation_fp4", "activation_scale_fp8"], - name="activation_quantize", - domain="trt", - axis=-1, - block_size=marker_block_size, - scale_type=TensorProto.FLOAT8E4M3FN, - ), - helper.make_node( - "DequantizeLinear", - ["activation_scale_fp8", "dynamic_dequantize_scale"], - ["activation_scale_dq"], - name="activation_scale_dequantize", - ), - helper.make_node( - "DequantizeLinear", - ["activation_fp4", "activation_scale_dq"], - ["activation_dq"], - name="activation_dequantize", - axis=-1, - block_size=marker_block_size, - ), - helper.make_node( - "Cast", - ["activation_dq"], - ["activation_dq_fp16"], - name="activation_cast", - to=TensorProto.FLOAT16, - ), - helper.make_node( - "TRT_FP4QDQ", - ["fp4_weight"], - ["fp4_weight_dq"], - name="fp4_weight_qdq", - domain="trt", - block_size=marker_block_size, - ), - helper.make_node( - "MatMul", - ["activation_dq_fp16", "fp4_weight_dq"], - ["linear_output"], - name="fp4_matmul", - ), - ] - ) - - scale = numpy_helper.from_array(np.array(0.25, dtype=np.float16), "qdq_scale_value") - zero = numpy_helper.from_array(np.array(0, dtype=np.int8), "qdq_zero_value") - nodes.extend( - [ - helper.make_node("Constant", [], ["qdq_scale"], name="qdq_scale", value=scale), - helper.make_node("Constant", [], ["qdq_zero"], name="qdq_zero", value=zero), - helper.make_node( - "QuantizeLinear", - ["qdq_weight", "qdq_scale", "qdq_zero"], - ["quantized_weight"], - name="weight_quantize", - ), - helper.make_node( - "DequantizeLinear", - ["quantized_weight", "qdq_scale", "qdq_zero"], - ["qdq_weight_dq"], - name="weight_dequantize", - ), - helper.make_node( - "QuantizeLinear", - ["qdq_activation", "qdq_scale", "qdq_zero"], - ["quantized_activation"], - name="activation_fp8_quantize", - ), - helper.make_node( - "DequantizeLinear", - ["quantized_activation", "qdq_scale", "qdq_zero"], - ["qdq_activation_dq"], - name="activation_fp8_dequantize", - ), - ] - ) - - if qdq_consumer == "Conv": - initializers.append( - numpy_helper.from_array(np.ones((1, 1, 1, 1), dtype=np.float16), "qdq_weight") - ) - inputs.append( - helper.make_tensor_value_info("qdq_activation", TensorProto.FLOAT16, [1, 1, 2, 2]) - ) - outputs.append( - helper.make_tensor_value_info("qdq_output", TensorProto.FLOAT16, [1, 1, 2, 2]) - ) - nodes.append( - helper.make_node( - "Conv", - ["qdq_activation_dq", "qdq_weight_dq"], - ["qdq_output"], - name="fp8_conv", - ) - ) - else: - initializers.append( - numpy_helper.from_array(np.ones((4, 4), dtype=np.float16), "qdq_weight") - ) - inputs.append(helper.make_tensor_value_info("qdq_activation", TensorProto.FLOAT16, [1, 4])) - outputs.append(helper.make_tensor_value_info("qdq_output", TensorProto.FLOAT16, [1, 4])) - nodes.append( - helper.make_node( - qdq_consumer, - ["qdq_activation_dq", "qdq_weight_dq"], - ["qdq_output"], - name=f"fp8_{qdq_consumer.lower()}", - ) - ) - - graph = helper.make_graph( - nodes, - "mixed_fp4_fp8", - inputs, - outputs, - initializers, - value_info=value_info, - ) - return helper.make_model( - graph, - opset_imports=[helper.make_opsetid("", 20), helper.make_opsetid("trt", 1)], - ) - - -def _tensor_dtype(onnx_model, tensor_name): - for initializer in onnx_model.graph.initializer: - if initializer.name == tensor_name: - return initializer.data_type - node = next(node for node in onnx_model.graph.node if tensor_name in node.output) - value = next(attribute for attribute in node.attribute if attribute.name == "value").t - return value.data_type - - -def _insert_passthrough_before_weight_quantize(onnx_model): - weight = next( - initializer - for initializer in onnx_model.graph.initializer - if initializer.name == "qdq_weight" - ) - weight.name = "qdq_weight_source" - quantize_node = next(node for node in onnx_model.graph.node if node.name == "weight_quantize") - onnx_model.graph.node.insert( - 0, - helper.make_node("Identity", ["qdq_weight_source"], ["qdq_weight"], name="weight_identity"), - ) - assert quantize_node.input[0] == "qdq_weight" - - -def _insert_weight_cast_before_quantize(onnx_model, dtype): - quantize_node = next(node for node in onnx_model.graph.node if node.name == "weight_quantize") - quantize_node.input[0] = "qdq_weight_cast" - onnx_model.graph.node.insert( - 0, - helper.make_node( - "Cast", - ["qdq_weight"], - ["qdq_weight_cast"], - name="weight_cast", - to=dtype, - ), - ) - - -def _make_quantized_backbone(linear_count, conv_count): - backbone = torch.nn.Module() - for index in range(linear_count): - linear = torch.nn.Linear(16, 16, bias=False) - linear.input_quantizer = _Quantizer((2, 1)) - linear.weight_quantizer = _Quantizer((2, 1)) - backbone.add_module(f"linear_{index}", linear) - for index in range(conv_count): - conv = torch.nn.Conv2d(1, 1, 1, bias=False) - conv.input_quantizer = _Quantizer((4, 3)) - conv.weight_quantizer = _Quantizer((4, 3)) - backbone.add_module(f"conv_{index}", conv) - return backbone - - -def _make_external_data_model(fill_value): - weight = numpy_helper.from_array( - np.full((1024,), fill_value, dtype=np.float32), "external_weight" - ) - graph = helper.make_graph( - [helper.make_node("Identity", ["external_weight"], ["output"])], - "external_data", - [], - [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1024])], - [weight], - ) - return helper.make_model(graph) - - -def test_fp8_scale_workaround_can_target_only_enabled_conv_quantizers(): - conv = torch.nn.Conv2d(1, 1, 1) - conv.input_quantizer = _Quantizer((4, 3)) - conv.weight_quantizer = _Quantizer((4, 3), enabled=False) - linear = torch.nn.Linear(1, 1) - linear.input_quantizer = _Quantizer((2, 1)) - linear.weight_quantizer = _Quantizer((4, 3)) - model = torch.nn.Sequential(conv, linear) - - diffusion_export.generate_fp8_scales(model, conv_only=True) - - assert conv.input_quantizer.num_bits == 8 - assert conv.input_quantizer._amax == 127.0 - assert conv.weight_quantizer.num_bits == (4, 3) - assert linear.input_quantizer.num_bits == (2, 1) - assert linear.weight_quantizer.num_bits == (4, 3) - - diffusion_export.generate_fp8_scales(model) - - assert linear.weight_quantizer.num_bits == 8 - assert linear.weight_quantizer._amax == 127.0 - - -def test_mixed_sdxl_graph_preserves_fp8_conv_and_lowers_exact_nvfp4_topology(): - raw_model = _make_raw_fp4_model() - - converted_model = diffusion_export._process_fp4_onnx_graph(raw_model, "sdxl-1.0") - - assert not any(node.op_type == "TRT_FP4QDQ" for node in converted_model.graph.node) - assert ( - sum( - initializer.data_type == TensorProto.FLOAT4E2M1 - for initializer in converted_model.graph.initializer - ) - == 1 - ) - assert _tensor_dtype(converted_model, "qdq_zero") == TensorProto.FLOAT8E4M3FN - assert any( - node.op_type == "Conv" and node.name == "fp8_conv" for node in converted_model.graph.node - ) - assert next(opset.version for opset in converted_model.opset_import if not opset.domain) >= 23 - onnx.checker.check_model(converted_model) - - -@pytest.mark.parametrize( - ("linear_count", "conv_count", "error"), - [ - (2, 1, "found 1 TRT_FP4QDQ weight markers, expected 2 enabled Linear pairs"), - ( - 1, - 2, - "found 1 initializer-backed FP8 Conv weight Q/DQ pairs, expected 2 enabled Conv2d pairs", - ), - ], -) -def test_raw_sdxl_graph_counts_match_enabled_quantizer_pairs(linear_count, conv_count, error): - expected_linear_count, expected_conv_count = diffusion_export._get_sdxl_fp4_expected_counts( - _make_quantized_backbone(linear_count, conv_count) - ) - - with pytest.raises(ValueError, match=error): - diffusion_export._process_fp4_onnx_graph( - _make_raw_fp4_model(), - "sdxl-1.0", - expected_linear_count=expected_linear_count, - expected_fp8_conv_count=expected_conv_count, - ) - - -def test_final_sdxl_graph_count_matches_enabled_conv_pairs(): - converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") - _, expected_conv_count = diffusion_export._get_sdxl_fp4_expected_counts( - _make_quantized_backbone(1, 2) - ) - - with pytest.raises( - ValueError, - match="found 1 initializer-backed FP8 Conv weight Q/DQ pairs, expected 2 enabled Conv2d pairs", - ): - diffusion_export._validate_final_fp4_graph( - converted_model, - expected_weight_count=1, - allow_fp8_conv=True, - expected_fp8_conv_count=expected_conv_count, - ) - - -@pytest.mark.parametrize("pre_q_passthrough", [False, True]) -def test_non_sdxl_fp4_graph_rejects_static_qdq_conv_weights(pre_q_passthrough): - model = _make_raw_fp4_model() - if pre_q_passthrough: - _insert_passthrough_before_weight_quantize(model) - - with pytest.raises(ValueError, match="disallowed initializer-backed Q/DQ"): - diffusion_export._process_fp4_onnx_graph(model, "flux-dev") - - -def test_sdxl_fp4_graph_requires_static_fp8_conv_weight_qdq(): - model = _make_raw_fp4_model() - quantize_node = next(node for node in model.graph.node if node.name == "weight_quantize") - quantize_node.input[0] = "qdq_activation" - - with pytest.raises(ValueError, match="no initializer-backed FP8 Conv weight Q/DQ"): - diffusion_export._process_fp4_onnx_graph(model, "sdxl-1.0") - - -@pytest.mark.parametrize( - ("marker_initializer", "marker_block_size", "error"), - [ - (False, 16, "not backed by a weight initializer"), - (True, 32, "block_size=32, expected 16"), - ], -) -def test_raw_fp4_validation_rejects_invalid_markers(marker_initializer, marker_block_size, error): - model = _make_raw_fp4_model( - marker_initializer=marker_initializer, marker_block_size=marker_block_size - ) - - with pytest.raises(ValueError, match=error): - diffusion_export._validate_raw_fp4_graph(model) - - -@pytest.mark.parametrize("consumer_op", ["Add", "Gemm", "MatMul"]) -def test_raw_fp4_validation_rejects_static_qdq_non_conv_weights(consumer_op): - model = _make_raw_fp4_model(qdq_consumer=consumer_op) - - with pytest.raises(ValueError, match="disallowed initializer-backed Q/DQ"): - diffusion_export._validate_raw_fp4_graph(model, allow_fp8_conv=True) - - -@pytest.mark.parametrize("consumer_op", ["Add", None]) -def test_raw_fp4_validation_rejects_marker_without_weight_consumer(consumer_op): - model = _make_raw_fp4_model() - matmul = next(node for node in model.graph.node if node.name == "fp4_matmul") - if consumer_op is None: - model.graph.node.remove(matmul) - else: - matmul.op_type = consumer_op - - with pytest.raises(ValueError, match="does not reach a Gemm/MatMul weight input"): - diffusion_export._validate_raw_fp4_graph(model, allow_fp8_conv=True) - - -@pytest.mark.parametrize( - ("corruption", "error"), - [ - ("block-scale-dtype", "FLOAT8E4M3FN block-scale initializer"), - ("axis", "does not use axis=-1"), - ("fp8-zero-dtype", "zero point is not FLOAT8E4M3FN"), - ("weight-consumer", "does not reach a Gemm/MatMul weight input"), - ], -) -def test_final_fp4_validation_rejects_invalid_double_dq(corruption, error): - raw_model = _make_raw_fp4_model() - expected_weight_count = diffusion_export._validate_raw_fp4_graph(raw_model, allow_fp8_conv=True) - normalized_model = diffusion_export._normalize_fp8_qdq(raw_model) - converted_model = NVFP4QuantExporter.process_model(normalized_model) - fp4_weight = next( - initializer.name - for initializer in converted_model.graph.initializer - if initializer.data_type == TensorProto.FLOAT4E2M1 - ) - weight_dq = next( - node - for node in converted_model.graph.node - if node.op_type == "DequantizeLinear" and node.input[0] == fp4_weight - ) - scale_dq = next( - node for node in converted_model.graph.node if weight_dq.input[1] in node.output - ) - if corruption == "block-scale-dtype": - fp8_scale = next( - initializer - for initializer in converted_model.graph.initializer - if initializer.name == scale_dq.input[0] - ) - fp8_scale.data_type = TensorProto.FLOAT16 - elif corruption == "axis": - axis = next(attribute for attribute in weight_dq.attribute if attribute.name == "axis") - axis.i = 0 - elif corruption == "fp8-zero-dtype": - zero_node = next(node for node in converted_model.graph.node if node.name == "qdq_zero") - value = next(attribute for attribute in zero_node.attribute if attribute.name == "value") - value.t.data_type = TensorProto.INT8 - else: - matmul = next(node for node in converted_model.graph.node if node.name == "fp4_matmul") - matmul.op_type = "Add" - - with pytest.raises(ValueError, match=error): - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count, allow_fp8_conv=True - ) - - -@pytest.mark.parametrize( - ("corruption", "error"), - [ - ("extra-input", "must be a two-input DequantizeLinear"), - ("axis", "must not use axis or block_size"), - ("block-size", "must not use axis or block_size"), - ("fp8-scale-fanout", "must be consumed only by"), - ("global-scale-dtype", "does not use a FLOAT global-scale initializer"), - ("global-scale-zero", "global scale must be a finite positive scalar constant"), - ("scale-output-fanout", "output must be consumed only by"), - ], -) -def test_final_fp4_validation_rejects_invalid_weight_scale_dq(corruption, error): - converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") - fp4_weight = next( - initializer.name - for initializer in converted_model.graph.initializer - if initializer.data_type == TensorProto.FLOAT4E2M1 - ) - weight_dq = next( - node - for node in converted_model.graph.node - if node.op_type == "DequantizeLinear" and node.input[0] == fp4_weight - ) - scale_dq = next( - node for node in converted_model.graph.node if weight_dq.input[1] in node.output - ) - if corruption == "extra-input": - scale_dq.input.append(scale_dq.input[1]) - elif corruption in {"axis", "block-size"}: - scale_dq.attribute.append(helper.make_attribute(corruption.replace("-", "_"), 0)) - elif corruption == "fp8-scale-fanout": - converted_model.graph.node.append( - helper.make_node( - "Identity", - [scale_dq.input[0]], - ["extra_block_scale_use"], - name="extra_block_scale_use", - ) - ) - elif corruption == "global-scale-dtype": - global_scale = next( - initializer - for initializer in converted_model.graph.initializer - if initializer.name == scale_dq.input[1] - ) - global_scale.data_type = TensorProto.FLOAT16 - elif corruption == "global-scale-zero": - global_scale = next( - initializer - for initializer in converted_model.graph.initializer - if initializer.name == scale_dq.input[1] - ) - global_scale.CopyFrom( - numpy_helper.from_array(np.array(0.0, dtype=np.float32), global_scale.name) - ) - else: - converted_model.graph.node.append( - helper.make_node( - "Identity", - [scale_dq.output[0]], - ["extra_weight_scale_use"], - name="extra_weight_scale_use", - ) - ) - - with pytest.raises(ValueError, match=error): - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count=1, allow_fp8_conv=True - ) - - -def test_final_fp4_validation_rejects_extra_float4_weight_consumer(): - raw_model = _make_raw_fp4_model() - expected_weight_count = diffusion_export._validate_raw_fp4_graph(raw_model, allow_fp8_conv=True) - converted_model = diffusion_export._process_fp4_onnx_graph(raw_model, "sdxl-1.0") - fp4_weight = next( - initializer.name - for initializer in converted_model.graph.initializer - if initializer.data_type == TensorProto.FLOAT4E2M1 - ) - converted_model.graph.node.append( - helper.make_node("Identity", [fp4_weight], ["extra_weight_use"], name="extra_weight_use") - ) - - with pytest.raises(ValueError, match="must feed exactly one weight DequantizeLinear"): - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count, allow_fp8_conv=True - ) - - -@pytest.mark.parametrize( - ("corruption", "error"), - [ - ("mismatched-scale", "do not share scale and zero point"), - ("axis", "without an axis"), - ("nonpositive-scale", "finite positive scalar constant"), - ("wrong-scale-dtype", "FP8 scale must use a floating-point dtype"), - ("nonzero-zero", "zero point must be a scalar zero"), - ("missing-activation", "has no FP8 activation Q/DQ"), - ("activation-fanout", "has non-Conv activation consumers"), - ("weight-fanout", "weight DQ must feed exactly one FP8 Conv input 1"), - ], -) -def test_final_fp4_validation_rejects_invalid_fp8_conv_qdq(corruption, error): - converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") - weight_quantize = next( - node for node in converted_model.graph.node if node.name == "weight_quantize" - ) - weight_dequantize = next( - node for node in converted_model.graph.node if node.name == "weight_dequantize" - ) - if corruption == "mismatched-scale": - converted_model.graph.initializer.append( - numpy_helper.from_array(np.array(0.5, dtype=np.float16), "other_fp8_scale") - ) - weight_dequantize.input[1] = "other_fp8_scale" - elif corruption == "axis": - weight_quantize.attribute.append(helper.make_attribute("axis", 0)) - elif corruption == "nonpositive-scale": - scale_node = next(node for node in converted_model.graph.node if node.name == "qdq_scale") - scale_value = next( - attribute for attribute in scale_node.attribute if attribute.name == "value" - ) - scale_value.t.CopyFrom( - numpy_helper.from_array(np.array(0.0, dtype=np.float16), "qdq_scale_value") - ) - elif corruption == "wrong-scale-dtype": - scale_node = next(node for node in converted_model.graph.node if node.name == "qdq_scale") - scale_value = next( - attribute for attribute in scale_node.attribute if attribute.name == "value" - ) - scale_value.t.CopyFrom( - numpy_helper.from_array(np.array(1, dtype=np.int32), "qdq_scale_value") - ) - elif corruption == "nonzero-zero": - zero_node = next(node for node in converted_model.graph.node if node.name == "qdq_zero") - zero_value = next( - attribute for attribute in zero_node.attribute if attribute.name == "value" - ) - zero_value.t.raw_data = b"\x38" - elif corruption == "missing-activation": - conv = next(node for node in converted_model.graph.node if node.name == "fp8_conv") - conv.input[0] = "qdq_activation" - elif corruption == "activation-fanout": - converted_model.graph.node.append( - helper.make_node( - "Identity", - ["qdq_activation_dq"], - ["extra_fp8_activation_use"], - name="extra_fp8_activation_use", - ) - ) - else: - converted_model.graph.node.append( - helper.make_node( - "Conv", - ["qdq_activation_dq", "qdq_weight_dq"], - ["extra_fp8_conv_output"], - name="extra_fp8_conv", - ) - ) - - with pytest.raises(ValueError, match=error): - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count=1, allow_fp8_conv=True - ) - - -@pytest.mark.parametrize( - ("corruption", "error"), - [ - ("block-size", "activation_quantize does not use block_size=16"), - ("mismatched-scale", "quantize and dequantize global scales do not match"), - ("wrong-domain", "must use the trt domain"), - ("static-input", "input 0 must be a dynamic activation"), - ("fanout", "must feed exactly one Gemm/MatMul activation input"), - ("bypass", "dynamic NVFP4 activation paths do not match FLOAT4 weight consumers"), - ("extra-scale-consumer", "FP8 scale output must feed exactly one DequantizeLinear"), - ], -) -def test_final_fp4_validation_rejects_invalid_dynamic_activation(corruption, error): - converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") - dynamic_quantize = next( - node for node in converted_model.graph.node if node.name == "activation_quantize" - ) - if corruption == "block-size": - block_size = next( - attribute for attribute in dynamic_quantize.attribute if attribute.name == "block_size" - ) - block_size.i = 32 - elif corruption == "mismatched-scale": - scale_node = next( - node for node in converted_model.graph.node if node.name == "dynamic_dequantize_scale" - ) - scale_value = next( - attribute for attribute in scale_node.attribute if attribute.name == "value" - ) - scale_value.t.CopyFrom( - numpy_helper.from_array(np.array(2.0 / 2688.0, dtype=np.float32), "scale_value") - ) - elif corruption == "wrong-domain": - dynamic_quantize.domain = "other" - elif corruption == "static-input": - dynamic_quantize.input[0] = "dynamic_quantize_scale" - elif corruption == "fanout": - converted_model.graph.node.append( - helper.make_node( - "MatMul", - ["activation_dq_fp16", "fp4_weight_dq"], - ["extra_linear_output"], - name="extra_fp4_matmul", - ) - ) - elif corruption == "bypass": - matmul = next(node for node in converted_model.graph.node if node.name == "fp4_matmul") - matmul.input[0] = "activation" - else: - converted_model.graph.node.append( - helper.make_node( - "Identity", - [dynamic_quantize.output[1]], - ["extra_dynamic_scale_use"], - name="extra_dynamic_scale_use", - ) - ) - - with pytest.raises(ValueError, match=error): - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count=1, allow_fp8_conv=True - ) - - -def test_final_fp4_validation_accepts_bfloat16_fp8_conv_scales(): - converted_model = diffusion_export._process_fp4_onnx_graph(_make_raw_fp4_model(), "sdxl-1.0") - qdq_weight = next( - initializer - for initializer in converted_model.graph.initializer - if initializer.name == "qdq_weight" - ) - qdq_weight.data_type = TensorProto.BFLOAT16 - qdq_activation = next( - value for value in converted_model.graph.input if value.name == "qdq_activation" - ) - qdq_activation.type.tensor_type.elem_type = TensorProto.BFLOAT16 - scale_node = next(node for node in converted_model.graph.node if node.name == "qdq_scale") - scale_value = next(attribute for attribute in scale_node.attribute if attribute.name == "value") - scale_value.t.data_type = TensorProto.BFLOAT16 - - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count=1, allow_fp8_conv=True - ) - - -def test_final_fp4_validation_uses_cast_output_dtype_for_fp8_conv_scale(): - raw_model = _make_raw_fp4_model() - _insert_weight_cast_before_quantize(raw_model, TensorProto.BFLOAT16) - activation = next(value for value in raw_model.graph.input if value.name == "qdq_activation") - activation.type.tensor_type.elem_type = TensorProto.BFLOAT16 - scale_node = next(node for node in raw_model.graph.node if node.name == "qdq_scale") - scale_value = next(attribute for attribute in scale_node.attribute if attribute.name == "value") - scale_value.t.data_type = TensorProto.BFLOAT16 - - converted_model = diffusion_export._process_fp4_onnx_graph(raw_model, "sdxl-1.0") - - diffusion_export._validate_final_fp4_graph( - converted_model, expected_weight_count=1, allow_fp8_conv=True - ) - - -def test_non_sdxl_fp4_export_uses_permissive_generic_lowering(monkeypatch, tmp_path): - raw_model = _make_raw_fp4_model() - generic_calls = [] - saved_models = [] - - def fake_onnx_export(*args, f, **kwargs): - del args, kwargs - onnx.save(raw_model, f) - - def generic_process(cls, onnx_model): - del cls - generic_calls.append(onnx_model) - return onnx_model - - monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) - monkeypatch.setattr( - diffusion_export, - "generate_dummy_kwargs_and_dynamic_axes_and_shapes", - lambda *args: ({}, {}, {}), - ) - monkeypatch.setattr( - diffusion_export, - "_process_fp4_onnx_graph", - lambda *args: pytest.fail("strict SDXL processing must not run for Flux"), - ) - monkeypatch.setattr(NVFP4QuantExporter, "process_model", classmethod(generic_process)) - monkeypatch.setattr( - diffusion_export, - "save_onnx", - lambda onnx_model, output: saved_models.append((onnx_model, output)), - ) - monkeypatch.setattr( - diffusion_export, - "_save_onnx_atomically", - lambda *args: pytest.fail("atomic checked publication must be SDXL FP4-only"), - ) - - diffusion_export.modelopt_export_sd(torch.nn.Identity(), tmp_path, "flux-dev", "fp4") - - assert len(generic_calls) == 1 - assert len(saved_models) == 1 - assert ( - next(opset.version for opset in saved_models[0][0].opset_import if not opset.domain) == 20 - ) - - -def test_sdxl_fp4_export_threads_enabled_quantizer_counts(monkeypatch, tmp_path): - raw_model = _make_raw_fp4_model() - processed_counts = [] - saved_models = [] - backbone = _make_quantized_backbone(2, 3) - - def fake_onnx_export(*args, f, **kwargs): - del args, kwargs - onnx.save(raw_model, f) - - def fake_process(onnx_model, model_name, block_size, **kwargs): - del model_name, block_size - processed_counts.append(kwargs) - return onnx_model - - monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) - monkeypatch.setattr( - diffusion_export, - "generate_dummy_kwargs_and_dynamic_axes_and_shapes", - lambda *args: ({}, {}, {}), - ) - monkeypatch.setattr( - diffusion_export, "configure_linear_module_onnx_quantizers", lambda _: nullcontext() - ) - monkeypatch.setattr(diffusion_export, "_process_fp4_onnx_graph", fake_process) - monkeypatch.setattr( - diffusion_export, - "_save_onnx_atomically", - lambda onnx_model, output: saved_models.append((onnx_model, output)), - ) - monkeypatch.setattr( - diffusion_export, - "save_onnx", - lambda *args: pytest.fail("SDXL FP4 must use atomic checked publication"), - ) - - diffusion_export.modelopt_export_sd(backbone, tmp_path, "sdxl-1.0", "fp4") - - assert processed_counts == [{"expected_linear_count": 2, "expected_fp8_conv_count": 3}] - assert len(saved_models) == 1 - - -@pytest.mark.parametrize("precision", ["fp16", "fp8"]) -def test_non_fp4_exports_keep_legacy_save_path(monkeypatch, tmp_path, precision): - raw_model = _make_raw_fp4_model() - saved_models = [] - - def fake_onnx_export(*args, f, **kwargs): - del args, kwargs - onnx.save(raw_model, f) - - monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) - monkeypatch.setattr( - diffusion_export, - "generate_dummy_kwargs_and_dynamic_axes_and_shapes", - lambda *args: ({}, {}, {}), - ) - monkeypatch.setattr(diffusion_export, "_normalize_fp8_qdq", lambda model: model) - monkeypatch.setattr( - diffusion_export, - "save_onnx", - lambda onnx_model, output: saved_models.append((onnx_model, output)), - ) - monkeypatch.setattr( - diffusion_export, - "_save_onnx_atomically", - lambda *args: pytest.fail("atomic checked publication must be SDXL FP4-only"), - ) - - diffusion_export.modelopt_export_sd(torch.nn.Identity(), tmp_path, "sdxl-1.0", precision) - - assert len(saved_models) == 1 - assert ( - next(opset.version for opset in saved_models[0][0].opset_import if not opset.domain) == 20 - ) - - -def test_invalid_fp4_export_preserves_existing_output_and_cleans_raw_temp(monkeypatch, tmp_path): - output_dir = tmp_path / "onnx" - output_dir.mkdir() - output = output_dir / "model.onnx" - output_data = output_dir / "model.onnx_data" - output.write_bytes(b"previous-model") - output_data.write_bytes(b"previous-data") - raw_dirs = [] - - invalid_model = helper.make_model( - helper.make_graph( - [helper.make_node("Identity", ["input"], ["output"])], - "invalid_fp4", - [helper.make_tensor_value_info("input", TensorProto.FLOAT, [1])], - [helper.make_tensor_value_info("output", TensorProto.FLOAT, [1])], - ) - ) - - def fake_onnx_export(*args, f, **kwargs): - del args, kwargs - raw_dirs.append(Path(f).parent) - onnx.save(invalid_model, f) - - monkeypatch.setattr(diffusion_export, "onnx_export", fake_onnx_export) - monkeypatch.setattr( - diffusion_export, - "generate_dummy_kwargs_and_dynamic_axes_and_shapes", - lambda *args: ({}, {}, {}), - ) - - with pytest.raises(ValueError, match="no TRT_FP4QDQ weight markers"): - diffusion_export.modelopt_export_sd(torch.nn.Identity(), output_dir, "sdxl-1.0", "fp4") - - assert output.read_bytes() == b"previous-model" - assert output_data.read_bytes() == b"previous-data" - assert len(raw_dirs) == 1 - assert not raw_dirs[0].exists() - - -def test_staged_save_failure_preserves_existing_output(monkeypatch, tmp_path): - output = tmp_path / "model.onnx" - output_data = tmp_path / "model.onnx_data" - output.write_bytes(b"previous-model") - output_data.write_bytes(b"previous-data") - - def fail_after_partial_save(onnx_model, staged_output, external_data_name=None): - del onnx_model - staged_output.write_bytes(b"partial-model") - (staged_output.parent / external_data_name).write_bytes(b"partial-data") - raise RuntimeError("save failed") - - monkeypatch.setattr(diffusion_export, "save_onnx", fail_after_partial_save) - - with pytest.raises(RuntimeError, match="save failed"): - diffusion_export._save_onnx_atomically(object(), output) - - assert output.read_bytes() == b"previous-model" - assert output_data.read_bytes() == b"previous-data" - assert not list(tmp_path.glob(".modelopt-export-*")) - - -def test_staged_checker_failure_preserves_existing_output(monkeypatch, tmp_path): - output = tmp_path / "model.onnx" - output_data = tmp_path / "model.onnx_data" - output.write_bytes(b"previous-model") - output_data.write_bytes(b"previous-data") - checked_paths = [] - - def fail_check(staged_output): - checked_paths.append(Path(staged_output)) - raise RuntimeError("checker failed") - - monkeypatch.setattr(diffusion_export.onnx.checker, "check_model", fail_check) - - with pytest.raises(RuntimeError, match="checker failed"): - diffusion_export._save_onnx_atomically(_make_external_data_model(2.0), output) - - assert len(checked_paths) == 1 - assert checked_paths[0].parent.name.startswith(".modelopt-export-") - assert output.read_bytes() == b"previous-model" - assert output_data.read_bytes() == b"previous-data" - assert not list(tmp_path.glob("model.onnx_data.*")) - assert not list(tmp_path.glob(".modelopt-export-*")) - - -def test_publication_failure_restores_previous_output_and_removes_new_data(monkeypatch, tmp_path): - output = tmp_path / "model.onnx" - output_data = tmp_path / "model.onnx_data" - output.write_bytes(b"previous-model") - output_data.write_bytes(b"previous-data") - - real_replace = diffusion_export.os.replace - - def fail_model_publish(source, destination): - source = Path(source) - destination = Path(destination) - if source.name == output.name and source.parent.name.startswith(".modelopt-export-"): - raise RuntimeError("publish failed") - real_replace(source, destination) - - monkeypatch.setattr(diffusion_export.os, "replace", fail_model_publish) - - with pytest.raises(RuntimeError, match="publish failed"): - diffusion_export._save_onnx_atomically(_make_external_data_model(2.0), output) - - assert output.read_bytes() == b"previous-model" - assert output_data.read_bytes() == b"previous-data" - assert not list(tmp_path.glob("model.onnx_data.*")) - assert not list(tmp_path.glob(".modelopt-export-*")) - - -def test_atomic_save_publishes_versioned_data_then_removes_old_data(tmp_path): - output = tmp_path / "model.onnx" - old_data = tmp_path / "model.onnx_data" - diffusion_export.save_onnx(_make_external_data_model(1.0), output) - assert old_data.exists() - - diffusion_export._save_onnx_atomically(_make_external_data_model(2.0), output) - - stored_model = onnx.load(str(output), load_external_data=False) - locations = { - entry.value - for initializer in stored_model.graph.initializer - for entry in initializer.external_data - if entry.key == "location" - } - assert len(locations) == 1 - external_data_name = locations.pop() - assert external_data_name.startswith("model.onnx_data.") - assert (tmp_path / external_data_name).exists() - assert not old_data.exists() - assert np.all(numpy_helper.to_array(onnx.load(str(output)).graph.initializer[0]) == 2.0) - assert not list(tmp_path.glob(".modelopt-export-*")) diff --git a/tests/unit/examples/test_diffusers_fp4_validation.py b/tests/unit/examples/test_diffusers_fp4_validation.py deleted file mode 100644 index 5bfc4347163..00000000000 --- a/tests/unit/examples/test_diffusers_fp4_validation.py +++ /dev/null @@ -1,652 +0,0 @@ -# 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 logging -import sys -from pathlib import Path - -import pytest -import torch -from diffusers.models.attention_processor import Attention -from torch import nn - -import modelopt.torch.quantization as mtq -from modelopt.torch.quantization.config import QuantizerAttributeConfig -from modelopt.torch.quantization.nn import TensorQuantizer - -_QUANTIZATION_EXAMPLE = ( - Path(__file__).resolve().parents[3] / "examples" / "diffusers" / "quantization" -) -sys.path.insert(0, str(_QUANTIZATION_EXAMPLE)) - -import quantize as quantize_module -from models_utils import ModelType -from quantize import ExportManager, Quantizer, _finalize_backbone_quantization -from quantize_config import ExportConfig, ModelConfig, QuantFormat, QuantizationConfig -from utils import check_conv_and_mha, validate_nvfp4_quantizers - -from examples.diffusers.quantization import onnx_utils as diffusion_onnx_utils -from examples.diffusers.quantization.onnx_utils import export as diffusion_export - -_MHA_QUANTIZER_NAMES = ( - "q_bmm_quantizer", - "k_bmm_quantizer", - "v_bmm_quantizer", - "softmax_quantizer", - "bmm2_output_quantizer", -) - - -def _nvfp4_quantizer(block_size=16, enabled=True, amax=1.0): - quantizer = TensorQuantizer( - QuantizerAttributeConfig( - num_bits=(2, 1), - block_sizes={-1: block_size, "type": "dynamic", "scale_bits": (4, 3)}, - ) - ) - if amax is not None: - quantizer.amax = torch.tensor(amax) - if not enabled: - quantizer.disable() - return quantizer - - -def _fp8_quantizer(enabled=True, amax=1.0): - quantizer = TensorQuantizer(QuantizerAttributeConfig(num_bits=(4, 3), axis=None)) - if amax is not None: - quantizer.amax = torch.tensor(amax) - if not enabled: - quantizer.disable() - return quantizer - - -class _QuantizedLinear(nn.Linear): - def __init__(self, input_quantizer=None, weight_quantizer=None): - super().__init__(16, 16, bias=False) - self.input_quantizer = input_quantizer or _nvfp4_quantizer() - self.weight_quantizer = weight_quantizer or _nvfp4_quantizer() - - -class _Backbone(nn.Module): - def __init__(self, linear=None, conv=None, attention=None): - super().__init__() - self.linear = linear or _QuantizedLinear() - if conv is not None: - self.conv = conv - if attention is not None: - self.attention = attention - - -def _quantized_conv(conv_cls, quantizer_factory): - conv = conv_cls(4, 4, kernel_size=1, bias=False) - conv.input_quantizer = quantizer_factory() - conv.weight_quantizer = quantizer_factory() - return conv - - -def _attention(head_size=16, enabled=True, softmax_amax=1.0): - attention = Attention(query_dim=head_size, heads=1, dim_head=head_size) - for name in _MHA_QUANTIZER_NAMES: - quantizer_enabled = enabled and name != "bmm2_output_quantizer" - amax = softmax_amax if name == "softmax_quantizer" else 1.0 - setattr(attention, name, _fp8_quantizer(enabled=quantizer_enabled, amax=amax)) - return attention - - -@pytest.mark.parametrize("model_type", [ModelType.SDXL_BASE, ModelType.SDXL_TURBO]) -def test_sdxl_fp4_uses_nvfp4_linear_fp8_conv_recipe(model_type): - config = Quantizer( - QuantizationConfig(format=QuantFormat.FP4, block_size=32), - ModelConfig(model_type=model_type), - logging.getLogger(__name__), - ).get_quant_config(n_steps=1, backbone=nn.Module()) - - linear_entries = { - entry["quantizer_name"]: entry["cfg"] - for entry in config["quant_cfg"] - if entry.get("parent_class") == "nn.Linear" - and entry.get("quantizer_name") in ("*input_quantizer", "*weight_quantizer") - } - conv_entries = { - entry["quantizer_name"]: entry["cfg"] - for entry in config["quant_cfg"] - if entry.get("parent_class") == "nn.Conv2d" - } - - assert set(linear_entries) == {"*input_quantizer", "*weight_quantizer"} - assert set(conv_entries) == {"*input_quantizer", "*weight_quantizer"} - for cfg in linear_entries.values(): - assert cfg["num_bits"] == (2, 1) - assert cfg["block_sizes"][-1] == 32 - assert cfg["block_sizes"]["scale_bits"] == (4, 3) - for cfg in conv_entries.values(): - assert cfg["num_bits"] == (4, 3) - assert cfg["axis"] is None - assert "block_sizes" not in cfg - - -def test_sdxl_fp4_recipe_applies_only_to_linear_and_conv(): - model = nn.Module() - model.linear = nn.Linear(16, 16, bias=False) - model.attn = nn.Module() - model.attn.to_q = nn.Linear(16, 16, bias=False) - model.attn.to_k = nn.Linear(16, 16, bias=False) - model.attn.to_v = nn.Linear(16, 16, bias=False) - model.conv = nn.Conv2d(4, 4, kernel_size=1, bias=False) - model.norm = nn.LayerNorm(16) - config = Quantizer( - QuantizationConfig(format=QuantFormat.FP4), - ModelConfig(model_type=ModelType.SDXL_BASE), - 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"]) - - assert model.linear.input_quantizer.is_enabled - assert model.linear.input_quantizer.is_nvfp4_dynamic - assert model.linear.weight_quantizer.is_enabled - assert model.linear.weight_quantizer.is_nvfp4_dynamic - 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 - assert model.conv.input_quantizer.is_enabled - assert model.conv.input_quantizer.is_fp8 - assert model.conv.weight_quantizer.is_enabled - assert model.conv.weight_quantizer.is_fp8 - assert not model.norm.input_quantizer.is_enabled - - -def test_fp4_finalization_preserves_calibrated_fp8_conv(): - conv = _quantized_conv(nn.Conv2d, lambda: _fp8_quantizer(amax=123.0)) - backbone = _Backbone(conv=conv) - before = { - name: quantizer.amax.clone() - for name, quantizer in ( - ("input", conv.input_quantizer), - ("weight", conv.weight_quantizer), - ) - } - - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4), - ModelType.SDXL_BASE, - restored=False, - ) - - for name, quantizer in ( - ("input", conv.input_quantizer), - ("weight", conv.weight_quantizer), - ): - assert quantizer.is_enabled - assert quantizer.is_fp8 - assert torch.equal(quantizer.amax, before[name]) - - -@pytest.mark.parametrize("export_fails", [False, True]) -def test_onnx_fp8_scale_workaround_restores_state_before_hf_export( - monkeypatch, tmp_path, export_fails -): - conv = _quantized_conv(nn.Conv2d, lambda: _fp8_quantizer(amax=448.0)) - backbone = _Backbone(conv=conv) - original_state = { - name: (quantizer.num_bits, quantizer.amax) - for name, quantizer in ( - ("input", conv.input_quantizer), - ("weight", conv.weight_quantizer), - ) - } - manager = ExportManager( - ExportConfig(onnx_dir=tmp_path / "onnx", hf_ckpt_dir=tmp_path / "hf"), - logging.getLogger(__name__), - pipeline_manager=None, - ) - - class _Pipeline: - def to(self, *args, **kwargs): - return self - - def fake_onnx_export(*args, **kwargs): - del args, kwargs - assert conv.input_quantizer.num_bits == 8 - assert conv.weight_quantizer.num_bits == 8 - assert conv.input_quantizer.amax == 127.0 - assert conv.weight_quantizer.amax == 127.0 - if export_fails: - raise RuntimeError("export failed") - - def fake_hf_export(*args, **kwargs): - del args, kwargs - for name, quantizer in ( - ("input", conv.input_quantizer), - ("weight", conv.weight_quantizer), - ): - num_bits, amax = original_state[name] - assert quantizer.num_bits == num_bits - assert quantizer.amax is amax - - monkeypatch.setattr(backbone, "to", lambda *args, **kwargs: backbone) - monkeypatch.setattr(torch.cuda, "empty_cache", lambda: None) - monkeypatch.setitem(sys.modules, "onnx_utils", diffusion_onnx_utils) - monkeypatch.setitem(sys.modules, "onnx_utils.export", diffusion_export) - monkeypatch.setattr(diffusion_export, "modelopt_export_sd", fake_onnx_export) - monkeypatch.setattr(quantize_module, "export_hf_checkpoint", fake_hf_export) - - if export_fails: - with pytest.raises(RuntimeError, match="export failed"): - manager.export_onnx(_Pipeline(), backbone, ModelType.SDXL_BASE, QuantFormat.FP4) - else: - manager.export_onnx(_Pipeline(), backbone, ModelType.SDXL_BASE, QuantFormat.FP4) - manager.export_hf_ckpt(_Pipeline(), ModelConfig(model_type=ModelType.SDXL_BASE)) - - for name, quantizer in ( - ("input", conv.input_quantizer), - ("weight", conv.weight_quantizer), - ): - num_bits, amax = original_state[name] - assert quantizer.num_bits == num_bits - assert quantizer.amax is amax - - -@pytest.mark.parametrize("conv_cls", [nn.Conv1d, nn.Conv2d]) -def test_fp4_finalization_disables_unsupported_nvfp4_conv(conv_cls): - conv = _quantized_conv(conv_cls, _nvfp4_quantizer) - backbone = _Backbone(conv=conv) - backbone.fp8_conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) - - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4), - ModelType.SDXL_BASE, - restored=False, - ) - - assert not conv.input_quantizer.is_enabled - assert not conv.weight_quantizer.is_enabled - - -@pytest.mark.parametrize( - ("input_quantizer", "weight_quantizer", "match"), - [ - (_nvfp4_quantizer(), _nvfp4_quantizer(enabled=False), "enable input and weight"), - (_fp8_quantizer(), _nvfp4_quantizer(), "dynamic E2M1"), - (_nvfp4_quantizer(32), _nvfp4_quantizer(), "requires block size 16"), - ( - _nvfp4_quantizer(enabled=False), - _nvfp4_quantizer(enabled=False), - "at least one enabled Linear", - ), - ], - ids=["unpaired", "wrong-format", "wrong-block-size", "no-enabled-pair"], -) -def test_nvfp4_linear_validation_rejects_invalid_state(input_quantizer, weight_quantizer, match): - backbone = _Backbone( - linear=_QuantizedLinear( - input_quantizer=input_quantizer, - weight_quantizer=weight_quantizer, - ) - ) - - with pytest.raises(ValueError, match=match): - validate_nvfp4_quantizers(backbone, expected_block_size=16, quantize_mha=False) - - -def test_nvfp4_linear_validation_allows_disabled_exclusion_pair(): - backbone = nn.Module() - backbone.enabled = _QuantizedLinear() - backbone.excluded = _QuantizedLinear( - input_quantizer=_nvfp4_quantizer(enabled=False), - weight_quantizer=_nvfp4_quantizer(enabled=False), - ) - - validate_nvfp4_quantizers(backbone, expected_block_size=16, quantize_mha=False) - - -@pytest.mark.parametrize("projection_name", ["to_q", "to_k", "to_v"]) -def test_sdxl_fp4_validation_rejects_enabled_attention_projection(projection_name): - backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) - backbone.attn = nn.Module() - setattr(backbone.attn, projection_name, _QuantizedLinear()) - - with pytest.raises(ValueError, match="must keep input and weight quantizers disabled"): - validate_nvfp4_quantizers( - backbone, - expected_block_size=16, - quantize_mha=False, - validate_sdxl_mixed_recipe=True, - ) - - -@pytest.mark.parametrize( - ("input_enabled", "weight_enabled"), - [(True, False), (False, True)], - ids=["input-only", "weight-only"], -) -def test_sdxl_fp4_validation_rejects_partial_attention_projection(input_enabled, weight_enabled): - backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) - backbone.attn = nn.Module() - backbone.attn.to_q = _QuantizedLinear( - input_quantizer=_nvfp4_quantizer(enabled=input_enabled), - weight_quantizer=_nvfp4_quantizer(enabled=weight_enabled), - ) - - with pytest.raises(ValueError, match="must keep input and weight quantizers disabled"): - validate_nvfp4_quantizers( - backbone, - expected_block_size=16, - quantize_mha=False, - validate_sdxl_mixed_recipe=True, - ) - - -def test_sdxl_fp4_validation_allows_disabled_attention_projections(): - backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) - backbone.attn = nn.Module() - for projection_name in ("to_q", "to_k", "to_v"): - setattr( - backbone.attn, - projection_name, - _QuantizedLinear( - input_quantizer=_nvfp4_quantizer(enabled=False), - weight_quantizer=_nvfp4_quantizer(enabled=False), - ), - ) - - validate_nvfp4_quantizers( - backbone, - expected_block_size=16, - quantize_mha=False, - validate_sdxl_mixed_recipe=True, - ) - - -@pytest.mark.parametrize( - ("quantizer_name", "amax"), - [ - ("input_quantizer", None), - ("weight_quantizer", float("nan")), - ("weight_quantizer", -1.0), - ], -) -def test_restore_finalization_rejects_uncalibrated_nvfp4_linear(quantizer_name, amax): - linear = _QuantizedLinear() - setattr(linear, quantizer_name, _nvfp4_quantizer(amax=amax)) - backbone = _Backbone(linear=linear, conv=_quantized_conv(nn.Conv2d, _fp8_quantizer)) - - with pytest.raises(ValueError, match="must have a finite nonnegative calibrated amax"): - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4), - ModelType.SDXL_BASE, - restored=True, - ) - - -def test_nvfp4_linear_validation_accepts_zero_amax(): - backbone = _Backbone( - linear=_QuantizedLinear( - input_quantizer=_nvfp4_quantizer(amax=0.0), - weight_quantizer=_nvfp4_quantizer(amax=0.0), - ) - ) - - validate_nvfp4_quantizers(backbone, expected_block_size=16, quantize_mha=False) - - -@pytest.mark.parametrize( - ("quantizer_factory", "match"), - [ - (_nvfp4_quantizer, "only supported on Linear input and weight quantizers"), - (_fp8_quantizer, "only supported on SDXL Conv2d"), - ], - ids=["nvfp4-layernorm", "fp8-layernorm"], -) -def test_sdxl_fp4_validation_rejects_quantized_layernorm(quantizer_factory, match): - conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) - backbone = _Backbone(conv=conv) - backbone.norm = nn.LayerNorm(16) - backbone.norm.input_quantizer = quantizer_factory() - - with pytest.raises(ValueError, match=match): - validate_nvfp4_quantizers( - backbone, - expected_block_size=16, - quantize_mha=False, - validate_sdxl_mixed_recipe=True, - ) - - -def test_restore_finalization_rejects_enabled_nvfp4_layernorm(): - conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) - backbone = _Backbone(conv=conv) - backbone.norm = nn.LayerNorm(16) - backbone.norm.input_quantizer = _nvfp4_quantizer() - - with pytest.raises(ValueError, match="only supported on Linear input and weight quantizers"): - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4), - ModelType.SDXL_BASE, - restored=True, - ) - - -@pytest.mark.parametrize("restored", [False, True]) -def test_non_sdxl_fp4_finalization_remains_permissive(restored): - backbone = _Backbone() - backbone.norm = nn.LayerNorm(16) - backbone.norm.input_quantizer = _nvfp4_quantizer() - - _finalize_backbone_quantization( - backbone, - "transformer", - QuantizationConfig(format=QuantFormat.FP4), - ModelType.FLUX_DEV, - restored=restored, - ) - - assert backbone.norm.input_quantizer.is_enabled - - -@pytest.mark.parametrize( - ("input_factory", "weight_factory", "match"), - [ - ( - _fp8_quantizer, - lambda: _fp8_quantizer(enabled=False), - "enable input and weight quantizers as a pair", - ), - ( - _fp8_quantizer, - _nvfp4_quantizer, - "must use per-tensor FP8", - ), - ( - lambda: _fp8_quantizer(amax=None), - _fp8_quantizer, - "finite positive calibrated amax", - ), - ( - lambda: _fp8_quantizer(amax=0.0), - _fp8_quantizer, - "finite positive calibrated amax", - ), - ( - lambda: _fp8_quantizer(amax=float("nan")), - _fp8_quantizer, - "finite positive calibrated amax", - ), - ], - ids=["partial", "wrong-format", "missing-amax", "zero-amax", "nonfinite-amax"], -) -def test_sdxl_fp8_conv_validation_rejects_invalid_state(input_factory, weight_factory, match): - conv = nn.Conv2d(4, 4, kernel_size=1, bias=False) - conv.input_quantizer = input_factory() - conv.weight_quantizer = weight_factory() - backbone = _Backbone(conv=conv) - - with pytest.raises(ValueError, match=match): - validate_nvfp4_quantizers( - backbone, - expected_block_size=16, - quantize_mha=False, - validate_sdxl_mixed_recipe=True, - ) - - -def test_sdxl_fp8_conv_validation_requires_enabled_pair(): - conv = _quantized_conv(nn.Conv2d, lambda: _fp8_quantizer(enabled=False)) - backbone = _Backbone(conv=conv) - - with pytest.raises(ValueError, match="at least one enabled calibrated FP8 Conv2d"): - validate_nvfp4_quantizers( - backbone, - expected_block_size=16, - quantize_mha=False, - validate_sdxl_mixed_recipe=True, - ) - - -@pytest.mark.parametrize(("quantize_mha", "enabled"), [(False, False), (True, True)]) -def test_restore_finalization_accepts_matching_mha_state(quantize_mha, enabled): - attention = _attention(enabled=enabled) - conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) - backbone = _Backbone(conv=conv, attention=attention) - - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4, quantize_mha=quantize_mha), - ModelType.SDXL_BASE, - restored=True, - ) - - for name in _MHA_QUANTIZER_NAMES: - expected_enabled = enabled and name != "bmm2_output_quantizer" - assert getattr(attention, name).is_enabled is expected_enabled - - -def test_restore_finalization_accepts_uncalibrated_softmax_quantizer(): - attention = _attention(enabled=True, softmax_amax=None) - backbone = _Backbone(conv=_quantized_conv(nn.Conv2d, _fp8_quantizer), attention=attention) - - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4, quantize_mha=True), - ModelType.SDXL_BASE, - restored=True, - ) - - assert attention.softmax_quantizer.is_enabled - assert attention.softmax_quantizer.amax is None - - -@pytest.mark.parametrize( - ("quantize_mha", "head_size", "mutate", "match"), - [ - (False, 16, lambda attention: None, "must be disabled"), - ( - True, - 16, - lambda attention: attention.q_bmm_quantizer.disable(), - "requires 'q_bmm_quantizer' to be enabled", - ), - ( - True, - 16, - lambda attention: attention.softmax_quantizer.disable(), - "requires 'softmax_quantizer' to be enabled", - ), - ( - True, - 16, - lambda attention: attention.bmm2_output_quantizer.enable(), - "requires 'bmm2_output_quantizer' to be disabled", - ), - ( - True, - 16, - lambda attention: setattr(attention, "softmax_quantizer", _nvfp4_quantizer()), - "must use per-tensor FP8", - ), - ( - True, - 16, - lambda attention: setattr(attention, "q_bmm_quantizer", _fp8_quantizer(amax=None)), - "finite positive calibrated amax", - ), - (True, 8, lambda attention: None, "must be disabled because FP8 MHA is unsupported"), - ], - ids=[ - "on-to-off", - "off-to-on", - "missing-softmax", - "enabled-bmm2-output", - "wrong-format", - "missing-amax", - "unsupported-head-size", - ], -) -def test_restore_finalization_rejects_mha_state_mismatch(quantize_mha, head_size, mutate, match): - attention = _attention(head_size=head_size, enabled=True) - mutate(attention) - backbone = _Backbone(attention=attention) - - with pytest.raises(ValueError, match=match): - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4, quantize_mha=quantize_mha), - ModelType.SDXL_BASE, - restored=True, - ) - - -@pytest.mark.parametrize("quantize_mha", [False, True]) -def test_fresh_finalization_applies_and_validates_mha_policy(quantize_mha): - attention = _attention(enabled=True) - conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) - backbone = _Backbone(conv=conv, attention=attention) - - _finalize_backbone_quantization( - backbone, - "unet", - QuantizationConfig(format=QuantFormat.FP4, quantize_mha=quantize_mha), - ModelType.SDXL_BASE, - restored=False, - ) - - for name in _MHA_QUANTIZER_NAMES: - quantizer = getattr(attention, name) - expected_enabled = quantize_mha and name != "bmm2_output_quantizer" - assert quantizer.is_enabled is expected_enabled - if quantizer.is_enabled: - assert quantizer.is_fp8 - - -def test_check_conv_does_not_disable_non_nvfp4_quantizers(): - conv = _quantized_conv(nn.Conv2d, _fp8_quantizer) - backbone = _Backbone(conv=conv) - - check_conv_and_mha(backbone, if_fp4=True, quantize_mha=False) - - assert conv.input_quantizer.is_enabled - assert conv.weight_quantizer.is_enabled From e9d9e6de92562dd9290fd9dc04ef24ac3904b309 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Thu, 10 Sep 2026 19:24:10 +0000 Subject: [PATCH 03/10] [5565357] Address Diffusers FP4 review feedback Restore the quantization policy after checkpoint loading for all supported Diffusers model families and make the FP8 export workaround predicate explicit. Add CPU coverage for the default FP8 scale path and isolate example-module imports in the focused test. Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/onnx_utils/export.py | 2 +- examples/diffusers/quantization/quantize.py | 32 ++-- tests/unit/examples/test_diffusers_fp4.py | 169 ++++++++++++++---- 3 files changed, 154 insertions(+), 49 deletions(-) diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 57d04296e91..8041ef2a122 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -147,7 +147,7 @@ def _temporary_fp8_export_scales(backbone, conv_only=False): if ( quantizer is None or not quantizer.is_enabled - or quantizer.num_bits != (4, 3) + or not quantizer.is_fp8 or getattr(quantizer, "_amax", None) is None ): continue diff --git a/examples/diffusers/quantization/quantize.py b/examples/diffusers/quantization/quantize.py index f423875ad42..4ab95d4a1bc 100644 --- a/examples/diffusers/quantization/quantize.py +++ b/examples/diffusers/quantization/quantize.py @@ -583,17 +583,20 @@ def create_argument_parser() -> argparse.ArgumentParser: return parser -def _restore_sdxl_fp4_policy( - pipeline_manager: PipelineManager, +def _apply_quantization_policy( + backbone: torch.nn.Module, + backbone_name: str, quant_config: QuantizationConfig, model_type: ModelType, ) -> None: - if quant_config.format != QuantFormat.FP4 or model_type not in _SDXL_MODEL_TYPES: + if backbone_name in ("video_decoder", "vae"): return - for backbone_name, backbone in pipeline_manager.iter_backbones(): - if backbone_name not in ("video_decoder", "vae"): - check_conv_and_mha(backbone, False, quant_config.quantize_mha) + check_conv_and_mha( + backbone, + quant_config.format == QuantFormat.FP4 and model_type not in _SDXL_MODEL_TYPES, + quant_config.quantize_mha, + ) def main() -> None: @@ -683,7 +686,10 @@ def main() -> None: if export_config.restore_from and export_config.restore_from.exists(): export_manager.restore_checkpoint() - _restore_sdxl_fp4_policy(pipeline_manager, quant_config, model_config.model_type) + for backbone_name, backbone in pipeline_manager.iter_backbones(): + _apply_quantization_policy( + backbone, backbone_name, quant_config, model_config.model_type + ) else: logger.info("Initializing calibration...") @@ -713,15 +719,9 @@ def forward_loop(mod): 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 - and model_config.model_type not in _SDXL_MODEL_TYPES, - quant_config.quantize_mha, - ) + _apply_quantization_policy( + backbone, backbone_name, quant_config, model_config.model_type + ) export_manager.save_checkpoint(backbone, backbone_name) diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index edb6834f652..d811b36e4df 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -13,6 +13,7 @@ # 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 @@ -29,19 +30,52 @@ from onnx import TensorProto, helper, numpy_helper 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 _QUANTIZATION_EXAMPLE = ( Path(__file__).resolve().parents[3] / "examples" / "diffusers" / "quantization" ) -sys.path.insert(0, str(_QUANTIZATION_EXAMPLE)) +_LOCAL_IMPORT_NAMES = ( + "calib.plugin_calib", + "calib", + "calibration", + "config", + "models_utils", + "pipeline_manager", + "quantize_config", + "utils", +) -from models_utils import ModelType -from quantize import Quantizer, _restore_sdxl_fp4_policy -from quantize_config import ModelConfig, QuantFormat, QuantizationConfig -from examples.diffusers.quantization.onnx_utils import export as diffusion_export +def _load_quantize_example(): + script = _QUANTIZATION_EXAMPLE / "quantize.py" + spec = importlib.util.spec_from_file_location("diffusers_quantize_example", script) + 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 +_apply_quantization_policy = _quantize._apply_quantization_policy class _RecipeBackbone(nn.Module): @@ -79,37 +113,66 @@ def test_sdxl_fp4_recipe(model_type): assert quantizer.is_fp8 -def test_restore_reapplies_sdxl_fp4_mha_policy(): - attention = Attention(query_dim=16, heads=1, dim_head=16) - skipped_attention = Attention(query_dim=16, heads=1, dim_head=16) - attention._disable_fp8_mha = True - skipped_attention._disable_fp8_mha = True +def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True): + quantizer = TensorQuantizer(QuantizerAttributeConfig(num_bits=num_bits, axis=None)) + if calibrated: + quantizer.amax = torch.tensor(448.0) + if not enabled: + quantizer.disable() + return quantizer - class _PipelineManager: - def iter_backbones(self): - return (("unet", attention), ("vae", skipped_attention)) - _restore_sdxl_fp4_policy( - _PipelineManager(), - QuantizationConfig(format=QuantFormat.FP4, quantize_mha=True), - ModelType.SDXL_BASE, +def _add_quantizers(module, *, num_bits=(4, 3), enabled=True, calibrated=True): + module.input_quantizer = _quantizer(num_bits=num_bits, enabled=enabled, calibrated=calibrated) + module.weight_quantizer = _quantizer(num_bits=num_bits, enabled=enabled, calibrated=calibrated) + + +@pytest.mark.parametrize( + ("model_type", "quant_format", "backbone_name", "conv_enabled"), + [ + (ModelType.SDXL_BASE, QuantFormat.FP4, "unet", True), + (ModelType.FLUX_DEV, QuantFormat.FP4, "transformer", False), + (ModelType.SD3_MEDIUM, QuantFormat.FP8, "transformer", True), + ], + ids=["sdxl-fp4", "flux-fp4", "sd3-fp8"], +) +def test_apply_quantization_policy(model_type, quant_format, backbone_name, conv_enabled): + backbone = nn.Module() + backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) + backbone.conv = nn.Conv2d(1, 1, 1) + _add_quantizers(backbone.conv) + backbone.attention._disable_fp8_mha = True + + _apply_quantization_policy( + backbone, + backbone_name, + QuantizationConfig(format=quant_format, quantize_mha=True), + model_type, ) - assert not attention._disable_fp8_mha - assert skipped_attention._disable_fp8_mha + assert backbone.attention._disable_fp8_mha is False + assert backbone.conv.input_quantizer.is_enabled is conv_enabled + assert backbone.conv.weight_quantizer.is_enabled is conv_enabled -def _fp8_quantizer(*, enabled=True): - quantizer = TensorQuantizer(QuantizerAttributeConfig(num_bits=(4, 3), axis=None)) - quantizer.amax = torch.tensor(448.0) - if not enabled: - quantizer.disable() - return quantizer +@pytest.mark.parametrize("backbone_name", ["vae", "video_decoder"]) +def test_apply_quantization_policy_skips_vae_backbones(backbone_name): + backbone = nn.Module() + backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) + backbone.conv = nn.Conv2d(1, 1, 1) + _add_quantizers(backbone.conv) + backbone.attention._disable_fp8_mha = True + _apply_quantization_policy( + backbone, + backbone_name, + QuantizationConfig(format=QuantFormat.FP4, quantize_mha=True), + ModelType.FLUX_DEV, + ) -def _add_fp8_quantizers(module, *, enabled=True): - module.input_quantizer = _fp8_quantizer(enabled=enabled) - module.weight_quantizer = _fp8_quantizer(enabled=enabled) + assert backbone.attention._disable_fp8_mha + assert backbone.conv.input_quantizer.is_enabled + assert backbone.conv.weight_quantizer.is_enabled @pytest.mark.parametrize("raises", [False, True]) @@ -118,11 +181,11 @@ def test_temporary_fp8_conv_export_scales_restore_state(raises): model.conv = nn.Conv2d(1, 1, 1) model.disabled_conv = nn.Conv2d(1, 1, 1) model.linear = nn.Linear(1, 1) - _add_fp8_quantizers(model.conv) - _add_fp8_quantizers(model.disabled_conv, enabled=False) - _add_fp8_quantizers(model.linear) + _add_quantizers(model.conv) + _add_quantizers(model.disabled_conv, enabled=False) + _add_quantizers(model.linear) linear_only = nn.Sequential(nn.Linear(1, 1)) - _add_fp8_quantizers(linear_only[0]) + _add_quantizers(linear_only[0]) assert not diffusion_export._has_enabled_conv(linear_only) assert diffusion_export._has_enabled_conv(model) changed = (model.conv.input_quantizer, model.conv.weight_quantizer) @@ -157,6 +220,48 @@ def run_export(): assert quantizer._amax is amax +@pytest.mark.parametrize( + ("module_type", "module_args"), + [(nn.Linear, (1, 1)), (nn.Conv2d, (1, 1, 1))], + ids=["linear", "conv2d"], +) +@pytest.mark.parametrize( + ("num_bits", "enabled", "calibrated", "expected_scaled"), + [ + ((4, 3), True, True, True), + ((4, 3), False, True, False), + ((4, 3), True, False, False), + (8, True, True, False), + ], + ids=["enabled-fp8", "disabled-fp8", "uncalibrated-fp8", "enabled-int8"], +) +def test_temporary_fp8_export_scales_filters_quantizers( + module_type, module_args, num_bits, enabled, calibrated, expected_scaled +): + model = nn.Sequential(module_type(*module_args)) + _add_quantizers(model[0], num_bits=num_bits, enabled=enabled, calibrated=calibrated) + quantizers = (model[0].input_quantizer, model[0].weight_quantizer) + original_state = { + quantizer: (quantizer._num_bits, getattr(quantizer, "_amax", None)) + for quantizer in quantizers + } + + with diffusion_export._temporary_fp8_export_scales(model, conv_only=False): + for quantizer in quantizers: + original_num_bits, original_amax = original_state[quantizer] + if expected_scaled: + assert quantizer.num_bits == 8 + assert quantizer.amax == 127.0 + assert quantizer._amax is not original_amax + else: + assert quantizer._num_bits == original_num_bits + assert getattr(quantizer, "_amax", None) is original_amax + + for quantizer, (original_num_bits, original_amax) in original_state.items(): + assert quantizer._num_bits == original_num_bits + assert getattr(quantizer, "_amax", None) is original_amax + + def _make_mixed_fp4_fp8_model(): fp4_weight = numpy_helper.from_array( np.linspace(-1.0, 1.0, 16 * 16, dtype=np.float16).reshape(16, 16), "fp4_weight" From 373eebdc2a812527465f8dfd6edc278384f41203 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Thu, 10 Sep 2026 23:36:36 +0000 Subject: [PATCH 04/10] [5565357] Fix Diffusers FP8 export handling Apply temporary FP8 export scales to enabled Conv1d and Conv3d quantizers and persist the converted Flux RoPE graph. Add CPU regression coverage for both review findings. Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/onnx_utils/export.py | 8 +++-- tests/unit/examples/test_diffusers_fp4.py | 36 +++++++++++++++++-- 2 files changed, 40 insertions(+), 4 deletions(-) diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 8041ef2a122..ac54a7df36d 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -136,7 +136,11 @@ def _has_enabled_conv(backbone): @contextmanager def _temporary_fp8_export_scales(backbone, conv_only=False): # temporary solution due to a known bug in torch.onnx._dynamo_export - module_types = (torch.nn.Conv2d,) if conv_only else (torch.nn.Linear, torch.nn.Conv2d) + module_types = ( + (torch.nn.Conv2d,) + if conv_only + else (torch.nn.Linear, torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d) + ) quantizer_states = [] try: for _, module in backbone.named_modules(): @@ -600,7 +604,7 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision): if not model_name.startswith("flux"): onnx_model = _normalize_fp8_qdq(onnx_model) else: - flux_convert_rope_weight_type(onnx_model) + onnx_model = flux_convert_rope_weight_type(onnx_model) if precision == "fp4": onnx_model = _process_fp4_onnx_graph(onnx_model, model_name) save_onnx(onnx_model, q_output) diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index d811b36e4df..71ab05b13f9 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -222,8 +222,13 @@ def run_export(): @pytest.mark.parametrize( ("module_type", "module_args"), - [(nn.Linear, (1, 1)), (nn.Conv2d, (1, 1, 1))], - ids=["linear", "conv2d"], + [ + (nn.Linear, (1, 1)), + (nn.Conv1d, (1, 1, 1)), + (nn.Conv2d, (1, 1, 1)), + (nn.Conv3d, (1, 1, 1)), + ], + ids=["linear", "conv1d", "conv2d", "conv3d"], ) @pytest.mark.parametrize( ("num_bits", "enabled", "calibrated", "expected_scaled"), @@ -262,6 +267,33 @@ def test_temporary_fp8_export_scales_filters_quantizers( assert getattr(quantizer, "_amax", None) is original_amax +def test_flux_export_saves_converted_rope_model(monkeypatch, tmp_path): + original_model = object() + converted_model = object() + saved_models = [] + + 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) + + def convert_rope_weight_type(model): + assert model is original_model + return converted_model + + monkeypatch.setattr(diffusion_export, "flux_convert_rope_weight_type", convert_rope_weight_type) + monkeypatch.setattr( + diffusion_export, "save_onnx", lambda model, path: saved_models.append((model, path)) + ) + + diffusion_export.modelopt_export_sd(nn.Module(), tmp_path, "flux-dev", "fp8") + + assert saved_models == [(converted_model, tmp_path / "model.onnx")] + + def _make_mixed_fp4_fp8_model(): fp4_weight = numpy_helper.from_array( np.linspace(-1.0, 1.0, 16 * 16, dtype=np.float16).reshape(16, 16), "fp4_weight" From 56e32794b2d161e2f2ab8eededfb6d16fd2155c2 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Fri, 11 Sep 2026 19:13:58 +0000 Subject: [PATCH 05/10] [5565357] Preserve restored Diffusers quantization policy Co-Authored-By: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- examples/diffusers/quantization/quantize.py | 54 +++++++-- tests/unit/examples/test_diffusers_fp4.py | 124 +++++++++++++++++++- 2 files changed, 166 insertions(+), 12 deletions(-) diff --git a/examples/diffusers/quantization/quantize.py b/examples/diffusers/quantization/quantize.py index 4ab95d4a1bc..cdb20ecc8e7 100644 --- a/examples/diffusers/quantization/quantize.py +++ b/examples/diffusers/quantization/quantize.py @@ -56,6 +56,7 @@ 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) @@ -440,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") @@ -569,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", @@ -599,6 +602,43 @@ def _apply_quantization_policy( ) +def _restore_quantization_policy( + backbones: list[tuple[str, torch.nn.Module]], +) -> QuantFormat: + has_nvfp4 = False + has_fp8 = False + + for backbone_name, backbone in backbones: + for module in backbone.modules(): + if isinstance(module, TensorQuantizer): + has_nvfp4 |= module.is_nvfp4_dynamic or module.is_nvfp4_static + has_fp8 |= module.is_fp8 + + if backbone_name in ("video_decoder", "vae"): + continue + + for module in backbone.modules(): + q_quantizer = getattr(module, "q_bmm_quantizer", None) + k_quantizer = getattr(module, "k_bmm_quantizer", None) + v_quantizer = getattr(module, "v_bmm_quantizer", None) + if not ( + isinstance(q_quantizer, TensorQuantizer) + and isinstance(k_quantizer, TensorQuantizer) + and isinstance(v_quantizer, TensorQuantizer) + ): + continue + module._disable_fp8_mha = not all( + quantizer.is_enabled and quantizer.is_fp8 + for quantizer in (q_quantizer, k_quantizer, v_quantizer) + ) + + 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 @@ -673,9 +713,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) @@ -686,10 +726,10 @@ def main() -> None: if export_config.restore_from and export_config.restore_from.exists(): export_manager.restore_checkpoint() - for backbone_name, backbone in pipeline_manager.iter_backbones(): - _apply_quantization_policy( - backbone, backbone_name, quant_config, model_config.model_type - ) + quant_config.format = _restore_quantization_policy( + list(pipeline_manager.iter_backbones()) + ) + logger.info(f"Detected restored quantization format: {quant_config.format.value}") else: logger.info("Initializing calibration...") diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index 71ab05b13f9..1af2f26d281 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -17,6 +17,7 @@ import logging import sys from pathlib import Path +from unittest.mock import Mock import numpy as np import pytest @@ -76,6 +77,7 @@ def _load_quantize_example(): QuantizationConfig = _quantize.QuantizationConfig Quantizer = _quantize.Quantizer _apply_quantization_policy = _quantize._apply_quantization_policy +_restore_quantization_policy = _quantize._restore_quantization_policy class _RecipeBackbone(nn.Module): @@ -113,8 +115,10 @@ def test_sdxl_fp4_recipe(model_type): assert quantizer.is_fp8 -def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True): - quantizer = TensorQuantizer(QuantizerAttributeConfig(num_bits=num_bits, axis=None)) +def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True, block_sizes=None): + quantizer = TensorQuantizer( + QuantizerAttributeConfig(num_bits=num_bits, axis=None, block_sizes=block_sizes) + ) if calibrated: quantizer.amax = torch.tensor(448.0) if not enabled: @@ -122,9 +126,119 @@ def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True): return quantizer -def _add_quantizers(module, *, num_bits=(4, 3), enabled=True, calibrated=True): - module.input_quantizer = _quantizer(num_bits=num_bits, enabled=enabled, calibrated=calibrated) - module.weight_quantizer = _quantizer(num_bits=num_bits, enabled=enabled, calibrated=calibrated) +def _add_quantizers(module, *, num_bits=(4, 3), enabled=True, calibrated=True, block_sizes=None): + module.input_quantizer = _quantizer( + num_bits=num_bits, + enabled=enabled, + calibrated=calibrated, + block_sizes=block_sizes, + ) + module.weight_quantizer = _quantizer( + num_bits=num_bits, + enabled=enabled, + calibrated=calibrated, + block_sizes=block_sizes, + ) + + +@pytest.mark.parametrize("mha_enabled", [True, False]) +def test_restore_policy_preserves_mha_state(mha_enabled): + backbone = nn.Module() + backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) + quantizers = [] + for name in ( + "q_bmm_quantizer", + "k_bmm_quantizer", + "v_bmm_quantizer", + "softmax_quantizer", + "bmm2_output_quantizer", + ): + quantizer = _quantizer(enabled=mha_enabled) + setattr(backbone.attention, name, quantizer) + quantizers.append(quantizer) + + restored_format = _restore_quantization_policy([("transformer", backbone)]) + + assert restored_format == QuantFormat.FP8 + assert all(quantizer.is_enabled is mha_enabled for quantizer in quantizers) + assert backbone.attention._disable_fp8_mha is not mha_enabled + + +@pytest.mark.parametrize("expected_format", list(QuantFormat)) +def test_restore_policy_infers_quantization_format(expected_format): + backbone = nn.Module() + + if expected_format == QuantFormat.FP4: + backbone.conv = nn.Conv2d(1, 1, 1) + _add_quantizers(backbone.conv) + backbone.linear = nn.Linear(16, 16) + _add_quantizers( + backbone.linear, + num_bits=(2, 1), + block_sizes={-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + ) + else: + backbone.linear = nn.Linear(16, 16) + _add_quantizers( + backbone.linear, + num_bits=(4, 3) if expected_format == QuantFormat.FP8 else 8, + ) + + assert _restore_quantization_policy([("transformer", backbone)]) == expected_format + + +def test_restore_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_path): + backbone = nn.Module() + backbone.conv = nn.Conv2d(1, 1, 1) + _add_quantizers(backbone.conv) + backbone.linear = nn.Linear(16, 16) + _add_quantizers( + backbone.linear, + num_bits=(2, 1), + block_sizes={-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + ) + backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) + mha_quantizers = [] + for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): + quantizer = _quantizer() + setattr(backbone.attention, name, quantizer) + mha_quantizers.append(quantizer) + + pipeline_manager = Mock() + pipeline_manager.create_pipeline.return_value = object() + pipeline_manager.iter_backbones.side_effect = lambda: iter([("transformer", backbone)]) + export_manager = Mock() + + monkeypatch.setattr(_quantize, "PipelineManager", lambda *args: pipeline_manager) + monkeypatch.setattr(_quantize, "ExportManager", lambda *args: export_manager) + monkeypatch.setattr(torch.nn, "RMSNorm", torch.nn.RMSNorm) + monkeypatch.setattr( + torch.nn.modules.normalization, + "RMSNorm", + torch.nn.modules.normalization.RMSNorm, + ) + monkeypatch.setattr( + sys, + "argv", + [ + "quantize.py", + "--model", + "flux-schnell", + "--restore-from", + str(tmp_path), + "--onnx-dir", + str(tmp_path / "onnx"), + "--hf-ckpt-dir", + str(tmp_path / "hf"), + ], + ) + + _quantize.main() + + assert export_manager.export_onnx.call_args.args[-1] == QuantFormat.FP4 + export_manager.export_hf_ckpt.assert_called_once() + assert all(quantizer.is_enabled for quantizer in mha_quantizers) + assert backbone.attention._disable_fp8_mha is False @pytest.mark.parametrize( From 4c63f6851ecc694dfca501f2a4e078d759a71a22 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Fri, 11 Sep 2026 21:11:10 +0000 Subject: [PATCH 06/10] [5565357] Simplify SDXL NVFP4 implementation Reduce restore, export, and focused-test complexity while preserving the validated mixed-precision recipe. Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/onnx_utils/export.py | 36 +- examples/diffusers/quantization/quantize.py | 36 +- tests/unit/examples/test_diffusers_fp4.py | 417 +++++------------- 3 files changed, 139 insertions(+), 350 deletions(-) diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index ac54a7df36d..29d0c84de30 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -30,7 +30,6 @@ # limitations under the License. import os -import shutil import tempfile from contextlib import contextmanager, nullcontext from pathlib import Path @@ -125,12 +124,11 @@ def flux_convert_rope_weight_type(onnx_graph): def _has_enabled_conv(backbone): - for module in backbone.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 + return any( + isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)) + and (module.input_quantizer.is_enabled or module.weight_quantizer.is_enabled) + for module in backbone.modules() + ) @contextmanager @@ -143,7 +141,7 @@ def _temporary_fp8_export_scales(backbone, conv_only=False): ) quantizer_states = [] try: - for _, module in backbone.named_modules(): + for module in backbone.modules(): if not isinstance(module, module_types): continue for quantizer_name in ("input_quantizer", "weight_quantizer"): @@ -501,14 +499,13 @@ def _normalize_fp8_qdq(onnx_model): def _ensure_default_opset(onnx_model, minimum_version): - for opset_import in onnx_model.opset_import: - if opset_import.domain in {"", "ai.onnx"}: - opset_import.version = max(opset_import.version, minimum_version) - return - - opset_import = onnx_model.opset_import.add() - opset_import.domain = "" - opset_import.version = minimum_version + opset_import = next( + (item for item in onnx_model.opset_import if item.domain in {"", "ai.onnx"}), None + ) + if opset_import is None: + opset_import = onnx_model.opset_import.add() + opset_import.domain = "" + opset_import.version = max(opset_import.version, minimum_version) def _process_fp4_onnx_graph(onnx_model, model_name): @@ -582,9 +579,8 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision): do_constant_folding = True opset_version = 20 - tmp_subfolder = tempfile.mkdtemp(prefix="myapp_") - tmp_output = Path(f"{tmp_subfolder}/{model_file_name}") - try: + with tempfile.TemporaryDirectory(prefix="myapp_", ignore_cleanup_errors=True) as tmp_subfolder: + tmp_output = Path(tmp_subfolder) / model_file_name with quantizer_context, fp8_scale_context, torch.inference_mode(): onnx_export( backbone, @@ -608,5 +604,3 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision): if precision == "fp4": onnx_model = _process_fp4_onnx_graph(onnx_model, model_name) save_onnx(onnx_model, q_output) - finally: - shutil.rmtree(tmp_subfolder, ignore_errors=True) diff --git a/examples/diffusers/quantization/quantize.py b/examples/diffusers/quantization/quantize.py index cdb20ecc8e7..0eace465353 100644 --- a/examples/diffusers/quantization/quantize.py +++ b/examples/diffusers/quantization/quantize.py @@ -586,22 +586,6 @@ def create_argument_parser() -> argparse.ArgumentParser: return parser -def _apply_quantization_policy( - backbone: torch.nn.Module, - backbone_name: str, - quant_config: QuantizationConfig, - model_type: ModelType, -) -> None: - if backbone_name in ("video_decoder", "vae"): - return - - check_conv_and_mha( - backbone, - quant_config.format == QuantFormat.FP4 and model_type not in _SDXL_MODEL_TYPES, - quant_config.quantize_mha, - ) - - def _restore_quantization_policy( backbones: list[tuple[str, torch.nn.Module]], ) -> QuantFormat: @@ -610,14 +594,12 @@ def _restore_quantization_policy( for backbone_name, backbone in backbones: for module in backbone.modules(): - if isinstance(module, TensorQuantizer): + 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 backbone_name in ("video_decoder", "vae"): - continue - - for module in backbone.modules(): + if backbone_name in ("video_decoder", "vae"): + continue q_quantizer = getattr(module, "q_bmm_quantizer", None) k_quantizer = getattr(module, "k_bmm_quantizer", None) v_quantizer = getattr(module, "v_bmm_quantizer", None) @@ -724,7 +706,7 @@ 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 = _restore_quantization_policy( list(pipeline_manager.iter_backbones()) @@ -759,9 +741,13 @@ def forward_loop(mod): mtq.compress(backbone) logger.info(f"{backbone_name} compression completed") - _apply_quantization_policy( - backbone, backbone_name, quant_config, model_config.model_type - ) + if backbone_name not in ("video_decoder", "vae"): + check_conv_and_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) diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index 1af2f26d281..d09037b5983 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -19,7 +19,6 @@ from pathlib import Path from unittest.mock import Mock -import numpy as np import pytest import torch from torch import nn @@ -27,8 +26,6 @@ onnx = pytest.importorskip("onnx") pytest.importorskip("onnx_graphsurgeon") pytest.importorskip("diffusers") -from diffusers.models.attention_processor import Attention -from onnx import TensorProto, helper, numpy_helper import modelopt.torch.quantization as mtq from examples.diffusers.quantization.onnx_utils import export as diffusion_export @@ -51,8 +48,9 @@ def _load_quantize_example(): - script = _QUANTIZATION_EXAMPLE / "quantize.py" - spec = importlib.util.spec_from_file_location("diffusers_quantize_example", script) + 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 = { @@ -76,7 +74,6 @@ def _load_quantize_example(): QuantFormat = _quantize.QuantFormat QuantizationConfig = _quantize.QuantizationConfig Quantizer = _quantize.Quantizer -_apply_quantization_policy = _quantize._apply_quantization_policy _restore_quantization_policy = _quantize._restore_quantization_policy @@ -91,6 +88,22 @@ def __init__(self): self.conv = nn.Conv2d(4, 4, kernel_size=1, bias=False) +def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True, block_sizes=None): + quantizer = TensorQuantizer( + QuantizerAttributeConfig(num_bits=num_bits, axis=None, block_sizes=block_sizes) + ) + if calibrated: + quantizer.amax = torch.tensor(448.0) + if not enabled: + quantizer.disable() + return quantizer + + +def _add_quantizers(module, **kwargs): + module.input_quantizer = _quantizer(**kwargs) + module.weight_quantizer = _quantizer(**kwargs) + + @pytest.mark.parametrize("model_type", [ModelType.SDXL_BASE, ModelType.SDXL_TURBO]) def test_sdxl_fp4_recipe(model_type): model = _RecipeBackbone() @@ -115,89 +128,55 @@ def test_sdxl_fp4_recipe(model_type): assert quantizer.is_fp8 -def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True, block_sizes=None): - quantizer = TensorQuantizer( - QuantizerAttributeConfig(num_bits=num_bits, axis=None, block_sizes=block_sizes) - ) - if calibrated: - quantizer.amax = torch.tensor(448.0) - if not enabled: - quantizer.disable() - return quantizer - - -def _add_quantizers(module, *, num_bits=(4, 3), enabled=True, calibrated=True, block_sizes=None): - module.input_quantizer = _quantizer( - num_bits=num_bits, - enabled=enabled, - calibrated=calibrated, - block_sizes=block_sizes, - ) - module.weight_quantizer = _quantizer( - num_bits=num_bits, - enabled=enabled, - calibrated=calibrated, - block_sizes=block_sizes, - ) - - -@pytest.mark.parametrize("mha_enabled", [True, False]) -def test_restore_policy_preserves_mha_state(mha_enabled): - backbone = nn.Module() - backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) - quantizers = [] - for name in ( - "q_bmm_quantizer", - "k_bmm_quantizer", - "v_bmm_quantizer", - "softmax_quantizer", - "bmm2_output_quantizer", - ): - quantizer = _quantizer(enabled=mha_enabled) - setattr(backbone.attention, name, quantizer) - quantizers.append(quantizer) - - restored_format = _restore_quantization_policy([("transformer", backbone)]) - - assert restored_format == QuantFormat.FP8 - assert all(quantizer.is_enabled is mha_enabled for quantizer in quantizers) - assert backbone.attention._disable_fp8_mha is not mha_enabled - - -@pytest.mark.parametrize("expected_format", list(QuantFormat)) -def test_restore_policy_infers_quantization_format(expected_format): +@pytest.mark.parametrize( + ("scenario", "expected_format", "mha_enabled"), + [ + ("mixed-fp4", QuantFormat.FP4, True), + ("fp8", QuantFormat.FP8, True), + ("int8-disabled-fp8-mha", QuantFormat.INT8, False), + ], +) +def test_restore_policy_uses_enabled_checkpoint_state(scenario, expected_format, mha_enabled): backbone = nn.Module() - - if expected_format == QuantFormat.FP4: - backbone.conv = nn.Conv2d(1, 1, 1) - _add_quantizers(backbone.conv) - backbone.linear = nn.Linear(16, 16) + backbone.linear = nn.Linear(16, 16) + if scenario == "mixed-fp4": _add_quantizers( backbone.linear, num_bits=(2, 1), block_sizes={-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, ) + backbone.conv = nn.Conv2d(1, 1, 1) + _add_quantizers(backbone.conv) + elif scenario == "fp8": + _add_quantizers(backbone.linear) else: - backbone.linear = nn.Linear(16, 16) - _add_quantizers( - backbone.linear, - num_bits=(4, 3) if expected_format == QuantFormat.FP8 else 8, - ) + _add_quantizers(backbone.linear, num_bits=8) + + backbone.attention = nn.Module() + for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): + setattr(backbone.attention, name, _quantizer(enabled=mha_enabled)) + quantizers = [module for module in backbone.modules() if isinstance(module, TensorQuantizer)] + state = {q: (q.is_enabled, q._num_bits, q._amax) for q in quantizers} + + restored_format = _restore_quantization_policy([("transformer", backbone)]) - assert _restore_quantization_policy([("transformer", backbone)]) == expected_format + assert restored_format == expected_format + assert backbone.attention._disable_fp8_mha is not mha_enabled + for quantizer, (enabled, num_bits, amax) in state.items(): + assert quantizer.is_enabled is enabled + assert quantizer._num_bits == num_bits + assert quantizer._amax is amax def test_restore_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_path): backbone = nn.Module() - backbone.conv = nn.Conv2d(1, 1, 1) - _add_quantizers(backbone.conv) backbone.linear = nn.Linear(16, 16) _add_quantizers( backbone.linear, num_bits=(2, 1), block_sizes={-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, ) - backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) + backbone.attention = nn.Module() mha_quantizers = [] for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): quantizer = _quantizer() @@ -208,15 +187,10 @@ def test_restore_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_ pipeline_manager.create_pipeline.return_value = object() pipeline_manager.iter_backbones.side_effect = lambda: iter([("transformer", backbone)]) export_manager = Mock() - monkeypatch.setattr(_quantize, "PipelineManager", lambda *args: pipeline_manager) monkeypatch.setattr(_quantize, "ExportManager", lambda *args: export_manager) monkeypatch.setattr(torch.nn, "RMSNorm", torch.nn.RMSNorm) - monkeypatch.setattr( - torch.nn.modules.normalization, - "RMSNorm", - torch.nn.modules.normalization.RMSNorm, - ) + monkeypatch.setattr(torch.nn.modules.normalization, "RMSNorm", torch.nn.RMSNorm) monkeypatch.setattr( sys, "argv", @@ -228,164 +202,68 @@ def test_restore_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_ str(tmp_path), "--onnx-dir", str(tmp_path / "onnx"), - "--hf-ckpt-dir", - str(tmp_path / "hf"), ], ) _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() assert all(quantizer.is_enabled for quantizer in mha_quantizers) assert backbone.attention._disable_fp8_mha is False -@pytest.mark.parametrize( - ("model_type", "quant_format", "backbone_name", "conv_enabled"), - [ - (ModelType.SDXL_BASE, QuantFormat.FP4, "unet", True), - (ModelType.FLUX_DEV, QuantFormat.FP4, "transformer", False), - (ModelType.SD3_MEDIUM, QuantFormat.FP8, "transformer", True), - ], - ids=["sdxl-fp4", "flux-fp4", "sd3-fp8"], -) -def test_apply_quantization_policy(model_type, quant_format, backbone_name, conv_enabled): - backbone = nn.Module() - backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) - backbone.conv = nn.Conv2d(1, 1, 1) - _add_quantizers(backbone.conv) - backbone.attention._disable_fp8_mha = True - - _apply_quantization_policy( - backbone, - backbone_name, - QuantizationConfig(format=quant_format, quantize_mha=True), - model_type, - ) - - assert backbone.attention._disable_fp8_mha is False - assert backbone.conv.input_quantizer.is_enabled is conv_enabled - assert backbone.conv.weight_quantizer.is_enabled is conv_enabled - - -@pytest.mark.parametrize("backbone_name", ["vae", "video_decoder"]) -def test_apply_quantization_policy_skips_vae_backbones(backbone_name): - backbone = nn.Module() - backbone.attention = Attention(query_dim=16, heads=1, dim_head=16) - backbone.conv = nn.Conv2d(1, 1, 1) - _add_quantizers(backbone.conv) - backbone.attention._disable_fp8_mha = True - - _apply_quantization_policy( - backbone, - backbone_name, - QuantizationConfig(format=QuantFormat.FP4, quantize_mha=True), - ModelType.FLUX_DEV, - ) - - assert backbone.attention._disable_fp8_mha - assert backbone.conv.input_quantizer.is_enabled - assert backbone.conv.weight_quantizer.is_enabled - - -@pytest.mark.parametrize("raises", [False, True]) -def test_temporary_fp8_conv_export_scales_restore_state(raises): +def test_temporary_fp8_export_scales_filter_and_restore_on_error(): model = nn.Module() - model.conv = nn.Conv2d(1, 1, 1) - model.disabled_conv = nn.Conv2d(1, 1, 1) - model.linear = nn.Linear(1, 1) - _add_quantizers(model.conv) - _add_quantizers(model.disabled_conv, enabled=False) - _add_quantizers(model.linear) - linear_only = nn.Sequential(nn.Linear(1, 1)) - _add_quantizers(linear_only[0]) - assert not diffusion_export._has_enabled_conv(linear_only) - assert diffusion_export._has_enabled_conv(model) - changed = (model.conv.input_quantizer, model.conv.weight_quantizer) - unchanged = ( - model.disabled_conv.input_quantizer, - model.disabled_conv.weight_quantizer, - model.linear.input_quantizer, - model.linear.weight_quantizer, - ) - original_state = { - quantizer: (quantizer._num_bits, quantizer._amax) for quantizer in changed + unchanged - } - - def run_export(): - with diffusion_export._temporary_fp8_export_scales(model, conv_only=True): - for quantizer in changed: - assert quantizer.num_bits == 8 - assert quantizer.amax == 127.0 - for quantizer in unchanged: - assert (quantizer._num_bits, quantizer._amax) == original_state[quantizer] - if raises: - raise RuntimeError("export failed") - - if raises: - with pytest.raises(RuntimeError, match="export failed"): - run_export() - else: - run_export() - - for quantizer, (num_bits, amax) in original_state.items(): - assert quantizer._num_bits == num_bits - assert quantizer._amax is amax - - -@pytest.mark.parametrize( - ("module_type", "module_args"), - [ - (nn.Linear, (1, 1)), - (nn.Conv1d, (1, 1, 1)), - (nn.Conv2d, (1, 1, 1)), - (nn.Conv3d, (1, 1, 1)), - ], - ids=["linear", "conv1d", "conv2d", "conv3d"], -) -@pytest.mark.parametrize( - ("num_bits", "enabled", "calibrated", "expected_scaled"), - [ - ((4, 3), True, True, True), - ((4, 3), False, True, False), - ((4, 3), True, False, False), - (8, True, True, False), - ], - ids=["enabled-fp8", "disabled-fp8", "uncalibrated-fp8", "enabled-int8"], -) -def test_temporary_fp8_export_scales_filters_quantizers( - module_type, module_args, num_bits, enabled, calibrated, expected_scaled -): - model = nn.Sequential(module_type(*module_args)) - _add_quantizers(model[0], num_bits=num_bits, enabled=enabled, calibrated=calibrated) - quantizers = (model[0].input_quantizer, model[0].weight_quantizer) - original_state = { - quantizer: (quantizer._num_bits, getattr(quantizer, "_amax", None)) - for quantizer in quantizers + modules = { + "linear": nn.Linear(1, 1), + "conv1d": nn.Conv1d(1, 1, 1), + "conv2d": nn.Conv2d(1, 1, 1), + "conv3d": nn.Conv3d(1, 1, 1), + "disabled": nn.Conv2d(1, 1, 1), + "uncalibrated": nn.Linear(1, 1), + "int8": nn.Conv2d(1, 1, 1), } + for name, module in modules.items(): + setattr(model, name, module) + _add_quantizers( + module, + enabled=name != "disabled", + calibrated=name != "uncalibrated", + num_bits=8 if name == "int8" else (4, 3), + ) + quantizers = [module for module in model.modules() if isinstance(module, TensorQuantizer)] + state = {q: (q._num_bits, getattr(q, "_amax", None)) for q in quantizers} - with diffusion_export._temporary_fp8_export_scales(model, conv_only=False): - for quantizer in quantizers: - original_num_bits, original_amax = original_state[quantizer] - if expected_scaled: - assert quantizer.num_bits == 8 - assert quantizer.amax == 127.0 - assert quantizer._amax is not original_amax - else: - assert quantizer._num_bits == original_num_bits - assert getattr(quantizer, "_amax", None) is original_amax - - for quantizer, (original_num_bits, original_amax) in original_state.items(): - assert quantizer._num_bits == original_num_bits - assert getattr(quantizer, "_amax", None) is original_amax + for conv_only, changed_names in ( + (True, {"conv2d"}), + (False, {"linear", "conv1d", "conv2d", "conv3d"}), + ): + with ( + pytest.raises(RuntimeError, match="export failed"), + diffusion_export._temporary_fp8_export_scales(model, conv_only=conv_only), + ): + for name, module in modules.items(): + for quantizer in (module.input_quantizer, module.weight_quantizer): + if name in changed_names: + assert quantizer.num_bits == 8 + assert quantizer.amax == 127.0 + else: + assert ( + quantizer._num_bits, + getattr(quantizer, "_amax", None), + ) == state[quantizer] + raise RuntimeError("export failed") + + for quantizer, (num_bits, amax) in state.items(): + assert quantizer._num_bits == num_bits + assert getattr(quantizer, "_amax", None) is amax def test_flux_export_saves_converted_rope_model(monkeypatch, tmp_path): original_model = object() converted_model = object() - saved_models = [] - monkeypatch.setattr( diffusion_export, "generate_dummy_kwargs_and_dynamic_axes_and_shapes", @@ -393,102 +271,33 @@ def test_flux_export_saves_converted_rope_model(monkeypatch, tmp_path): ) monkeypatch.setattr(diffusion_export, "onnx_export", lambda *args, **kwargs: None) monkeypatch.setattr(diffusion_export.onnx, "load", lambda *args, **kwargs: original_model) - - def convert_rope_weight_type(model): - assert model is original_model - return converted_model - - monkeypatch.setattr(diffusion_export, "flux_convert_rope_weight_type", convert_rope_weight_type) monkeypatch.setattr( - diffusion_export, "save_onnx", lambda model, path: saved_models.append((model, path)) + diffusion_export, "flux_convert_rope_weight_type", lambda model: converted_model ) + save_onnx = Mock() + monkeypatch.setattr(diffusion_export, "save_onnx", save_onnx) diffusion_export.modelopt_export_sd(nn.Module(), tmp_path, "flux-dev", "fp8") - assert saved_models == [(converted_model, tmp_path / "model.onnx")] - + save_onnx.assert_called_once_with(converted_model, tmp_path / "model.onnx") -def _make_mixed_fp4_fp8_model(): - fp4_weight = numpy_helper.from_array( - np.linspace(-1.0, 1.0, 16 * 16, dtype=np.float16).reshape(16, 16), "fp4_weight" - ) - fp8_weight = numpy_helper.from_array(np.ones((1, 1, 1, 1), dtype=np.float16), "fp8_weight") - scale = numpy_helper.from_array(np.array(0.25, dtype=np.float16), "fp8_scale_value") - zero = numpy_helper.from_array(np.array(0, dtype=np.int8), "fp8_zero_value") - nodes = [ - helper.make_node( - "TRT_FP4QDQ", - ["fp4_weight"], - ["fp4_weight_dq"], - name="fp4_weight_qdq", - domain="trt", - block_size=16, - ), - helper.make_node( - "MatMul", ["linear_input", "fp4_weight_dq"], ["linear_output"], name="fp4_matmul" - ), - helper.make_node("Constant", [], ["fp8_scale"], name="fp8_scale", value=scale), - helper.make_node("Constant", [], ["fp8_zero"], name="fp8_zero", value=zero), - helper.make_node( - "QuantizeLinear", - ["fp8_weight", "fp8_scale", "fp8_zero"], - ["fp8_weight_q"], - name="fp8_weight_quantize", - ), - helper.make_node( - "DequantizeLinear", - ["fp8_weight_q", "fp8_scale", "fp8_zero"], - ["fp8_weight_dq"], - name="fp8_weight_dequantize", - ), - helper.make_node( - "QuantizeLinear", - ["conv_input", "fp8_scale", "fp8_zero"], - ["conv_input_q"], - name="fp8_activation_quantize", - ), - helper.make_node( - "DequantizeLinear", - ["conv_input_q", "fp8_scale", "fp8_zero"], - ["conv_input_dq"], - name="fp8_activation_dequantize", - ), - helper.make_node( - "Conv", ["conv_input_dq", "fp8_weight_dq"], ["conv_output"], name="fp8_conv" - ), - ] - graph = helper.make_graph( - nodes, - "mixed_fp4_fp8", - [ - helper.make_tensor_value_info("linear_input", TensorProto.FLOAT16, [1, 16]), - helper.make_tensor_value_info("conv_input", TensorProto.FLOAT16, [1, 1, 2, 2]), - ], - [ - helper.make_tensor_value_info("linear_output", TensorProto.FLOAT16, [1, 16]), - helper.make_tensor_value_info("conv_output", TensorProto.FLOAT16, [1, 1, 2, 2]), - ], - [fp4_weight, fp8_weight], - value_info=[helper.make_tensor_value_info("fp4_weight_dq", TensorProto.FLOAT16, [16, 16])], - ) - return helper.make_model( - graph, - opset_imports=[helper.make_opsetid("", 20), helper.make_opsetid("trt", 1)], - ) +@pytest.mark.parametrize("model_name", ["sdxl-1.0", "sdxl-turbo"]) +def test_sdxl_fp4_processing_order_and_opset(monkeypatch, model_name): + model = onnx.ModelProto() + model.opset_import.add(domain="", version=20) + calls = [] -def _constant_dtype(model, name): - node = next(node for node in model.graph.node if node.name == name) - return next(attribute for attribute in node.attribute if attribute.name == "value").t.data_type + def record(name): + def process(current_model): + calls.append(name) + return current_model + return process -def test_mixed_sdxl_fp4_graph_postprocessing(): - model = diffusion_export._process_fp4_onnx_graph(_make_mixed_fp4_fp8_model(), "sdxl-1.0") + monkeypatch.setattr(diffusion_export, "_normalize_fp8_qdq", record("fp8")) + monkeypatch.setattr(diffusion_export.NVFP4QuantExporter, "process_model", record("nvfp4")) - assert not any(node.op_type == "TRT_FP4QDQ" for node in model.graph.node) - assert any(tensor.data_type == TensorProto.FLOAT4E2M1 for tensor in model.graph.initializer) - assert _constant_dtype(model, "fp8_zero") == TensorProto.FLOAT8E4M3FN - assert any(node.op_type == "Conv" and node.name == "fp8_conv" for node in model.graph.node) - assert sum(node.op_type == "QuantizeLinear" for node in model.graph.node) == 2 - assert next(opset.version for opset in model.opset_import if not opset.domain) >= 23 - onnx.checker.check_model(model) + assert diffusion_export._process_fp4_onnx_graph(model, model_name) is model + assert calls == ["fp8", "nvfp4"] + assert model.opset_import[0].version == 23 From cb7fc32a0cf6cef74b90683159bf8dbe901ad17b Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Sat, 12 Sep 2026 00:07:24 +0000 Subject: [PATCH 07/10] [5565357] Preserve FP8 custom-op shapes Co-Authored-By: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/onnx_utils/export.py | 117 ++++-------------- modelopt/torch/quantization/export_onnx.py | 5 +- .../quantization/test_onnx_export_cuda.py | 4 - tests/unit/examples/test_diffusers_fp4.py | 89 ++----------- .../quantization/test_onnx_export_cpu.py | 50 ++++++++ 5 files changed, 93 insertions(+), 172 deletions(-) diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 29d0c84de30..37895eeb7b5 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -30,8 +30,9 @@ # limitations under the License. import os +import shutil import tempfile -from contextlib import contextmanager, nullcontext +from contextlib import nullcontext from pathlib import Path import onnx @@ -50,8 +51,6 @@ 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"}, @@ -123,46 +122,6 @@ def flux_convert_rope_weight_type(onnx_graph): return gs.export_onnx(graph) -def _has_enabled_conv(backbone): - return any( - isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)) - and (module.input_quantizer.is_enabled or module.weight_quantizer.is_enabled) - for module in backbone.modules() - ) - - -@contextmanager -def _temporary_fp8_export_scales(backbone, conv_only=False): - # temporary solution due to a known bug in torch.onnx._dynamo_export - module_types = ( - (torch.nn.Conv2d,) - if conv_only - else (torch.nn.Linear, torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d) - ) - quantizer_states = [] - try: - for module in backbone.modules(): - if not isinstance(module, module_types): - continue - for quantizer_name in ("input_quantizer", "weight_quantizer"): - quantizer = getattr(module, quantizer_name, None) - if ( - quantizer is None - or not quantizer.is_enabled - or not quantizer.is_fp8 - or getattr(quantizer, "_amax", None) is None - ): - continue - quantizer_states.append((quantizer, quantizer._num_bits, quantizer._amax)) - quantizer._num_bits = 8 - quantizer._amax = quantizer._amax * (127 / 448.0) - yield - finally: - for quantizer, num_bits, amax in reversed(quantizer_states): - quantizer._num_bits = num_bits - quantizer._amax = amax - - 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 @@ -490,14 +449,6 @@ def save_onnx(onnx_model, output): print(f"ONNX model saved to {output}") -def _normalize_fp8_qdq(onnx_model): - graph = gs.import_onnx(onnx_model) - graph.cleanup().toposort() - onnx_model = convert_zp_fp8(gs.export_onnx(graph)) - graph = gs.import_onnx(onnx_model) - return gs.export_onnx(graph.cleanup()) - - def _ensure_default_opset(onnx_model, minimum_version): opset_import = next( (item for item in onnx_model.opset_import if item.domain in {"", "ai.onnx"}), None @@ -508,29 +459,17 @@ def _ensure_default_opset(onnx_model, minimum_version): opset_import.version = max(opset_import.version, minimum_version) -def _process_fp4_onnx_graph(onnx_model, model_name): - if model_name in {"sdxl-1.0", "sdxl-turbo"}: - onnx_model = _normalize_fp8_qdq(onnx_model) - onnx_model = NVFP4QuantExporter.process_model(onnx_model) - if model_name in {"sdxl-1.0", "sdxl-turbo"}: - _ensure_default_opset(onnx_model, 23) - return onnx_model - - def modelopt_export_sd(backbone, onnx_dir, model_name, precision): model_file_name = "model.onnx" os.makedirs(f"{onnx_dir}", exist_ok=True) + tmp_subfolder = tempfile.mkdtemp(prefix="myapp_") + tmp_output = Path(f"{tmp_subfolder}/{model_file_name}") q_output = Path(f"{onnx_dir}/{model_file_name}") is_sdxl_fp4 = precision == "fp4" and model_name in {"sdxl-1.0", "sdxl-turbo"} quantizer_context = ( configure_linear_module_onnx_quantizers(backbone) if precision == "fp4" else nullcontext() ) - fp8_scale_context = ( - _temporary_fp8_export_scales(backbone, conv_only=is_sdxl_fp4) - if is_sdxl_fp4 or (precision == "fp8" and _has_enabled_conv(backbone)) - else nullcontext() - ) dummy_kwargs, dynamic_axes, _ = generate_dummy_kwargs_and_dynamic_axes_and_shapes( model_name, backbone @@ -579,28 +518,26 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision): do_constant_folding = True opset_version = 20 - with tempfile.TemporaryDirectory(prefix="myapp_", ignore_cleanup_errors=True) as tmp_subfolder: - tmp_output = Path(tmp_subfolder) / model_file_name - with quantizer_context, fp8_scale_context, torch.inference_mode(): - onnx_export( - backbone, - (), - f=tmp_output.as_posix(), - kwargs=dummy_kwargs, - input_names=input_names, - output_names=output_names, - dynamic_axes=dynamic_axes, - do_constant_folding=do_constant_folding, - opset_version=opset_version, - dynamo=False, - ) - 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"): - onnx_model = _normalize_fp8_qdq(onnx_model) - else: - onnx_model = flux_convert_rope_weight_type(onnx_model) - if precision == "fp4": - onnx_model = _process_fp4_onnx_graph(onnx_model, model_name) - save_onnx(onnx_model, q_output) + with quantizer_context, torch.inference_mode(): + onnx_export( + backbone, + (), + f=tmp_output.as_posix(), + kwargs=dummy_kwargs, + input_names=input_names, + output_names=output_names, + dynamic_axes=dynamic_axes, + do_constant_folding=do_constant_folding, + opset_version=opset_version, + dynamo=False, + ) + print(f"Saved at {tmp_output}") + onnx_model = onnx.load(str(tmp_output), load_external_data=True) + if precision == "fp8" and model_name.startswith("flux"): + flux_convert_rope_weight_type(onnx_model) + if precision == "fp4": + onnx_model = NVFP4QuantExporter.process_model(onnx_model) + if is_sdxl_fp4: + _ensure_default_opset(onnx_model, 23) + save_onnx(onnx_model, q_output) + shutil.rmtree(tmp_subfolder, ignore_errors=True) diff --git a/modelopt/torch/quantization/export_onnx.py b/modelopt/torch/quantization/export_onnx.py index e5778c3c96b..2ad3383e572 100644 --- a/modelopt/torch/quantization/export_onnx.py +++ b/modelopt/torch/quantization/export_onnx.py @@ -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( diff --git a/tests/gpu/torch/quantization/test_onnx_export_cuda.py b/tests/gpu/torch/quantization/test_onnx_export_cuda.py index 300abc52e9e..39f422c75a9 100644 --- a/tests/gpu/torch/quantization/test_onnx_export_cuda.py +++ b/tests/gpu/torch/quantization/test_onnx_export_cuda.py @@ -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) diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index d09037b5983..3e2d891673d 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -214,90 +214,25 @@ def test_restore_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_ assert backbone.attention._disable_fp8_mha is False -def test_temporary_fp8_export_scales_filter_and_restore_on_error(): - model = nn.Module() - modules = { - "linear": nn.Linear(1, 1), - "conv1d": nn.Conv1d(1, 1, 1), - "conv2d": nn.Conv2d(1, 1, 1), - "conv3d": nn.Conv3d(1, 1, 1), - "disabled": nn.Conv2d(1, 1, 1), - "uncalibrated": nn.Linear(1, 1), - "int8": nn.Conv2d(1, 1, 1), - } - for name, module in modules.items(): - setattr(model, name, module) - _add_quantizers( - module, - enabled=name != "disabled", - calibrated=name != "uncalibrated", - num_bits=8 if name == "int8" else (4, 3), - ) - quantizers = [module for module in model.modules() if isinstance(module, TensorQuantizer)] - state = {q: (q._num_bits, getattr(q, "_amax", None)) for q in quantizers} - - for conv_only, changed_names in ( - (True, {"conv2d"}), - (False, {"linear", "conv1d", "conv2d", "conv3d"}), - ): - with ( - pytest.raises(RuntimeError, match="export failed"), - diffusion_export._temporary_fp8_export_scales(model, conv_only=conv_only), - ): - for name, module in modules.items(): - for quantizer in (module.input_quantizer, module.weight_quantizer): - if name in changed_names: - assert quantizer.num_bits == 8 - assert quantizer.amax == 127.0 - else: - assert ( - quantizer._num_bits, - getattr(quantizer, "_amax", None), - ) == state[quantizer] - raise RuntimeError("export failed") - - for quantizer, (num_bits, amax) in state.items(): - assert quantizer._num_bits == num_bits - assert getattr(quantizer, "_amax", None) is amax - - -def test_flux_export_saves_converted_rope_model(monkeypatch, tmp_path): - original_model = object() - converted_model = object() +def test_sdxl_fp4_export_lowers_nvfp4_at_opset_23(monkeypatch, tmp_path): + model = onnx.ModelProto() + model.opset_import.add(domain="", version=20) 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) - monkeypatch.setattr( - diffusion_export, "flux_convert_rope_weight_type", lambda model: converted_model - ) + monkeypatch.setattr(diffusion_export.onnx, "load", lambda *args, **kwargs: model) + process_model = Mock(side_effect=lambda current_model: current_model) + monkeypatch.setattr(diffusion_export.NVFP4QuantExporter, "process_model", process_model) save_onnx = Mock() monkeypatch.setattr(diffusion_export, "save_onnx", save_onnx) - diffusion_export.modelopt_export_sd(nn.Module(), tmp_path, "flux-dev", "fp8") - - save_onnx.assert_called_once_with(converted_model, tmp_path / "model.onnx") + diffusion_export.modelopt_export_sd(nn.Module(), tmp_path, "sdxl-1.0", "fp4") - -@pytest.mark.parametrize("model_name", ["sdxl-1.0", "sdxl-turbo"]) -def test_sdxl_fp4_processing_order_and_opset(monkeypatch, model_name): - model = onnx.ModelProto() - model.opset_import.add(domain="", version=20) - calls = [] - - def record(name): - def process(current_model): - calls.append(name) - return current_model - - return process - - monkeypatch.setattr(diffusion_export, "_normalize_fp8_qdq", record("fp8")) - monkeypatch.setattr(diffusion_export.NVFP4QuantExporter, "process_model", record("nvfp4")) - - assert diffusion_export._process_fp4_onnx_graph(model, model_name) is model - assert calls == ["fp8", "nvfp4"] - assert model.opset_import[0].version == 23 + process_model.assert_called_once_with(model) + save_onnx.assert_called_once_with(model, tmp_path / "model.onnx") + assert ( + next(opset.version for opset in model.opset_import if opset.domain in {"", "ai.onnx"}) == 23 + ) diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index ce2ef626d63..8ea90407a60 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -59,6 +59,56 @@ 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), + ) + + buffer = io.BytesIO() + if "enable_onnx_checker" in inspect.signature(torch.onnx.export).parameters: + kwargs = {"enable_onnx_checker": False} + else: + kwargs = {} + torch.onnx.export( + model, + sample_input, + buffer, + opset_version=20, + dynamo=False, + **kwargs, + ) + + buffer.seek(0) + exported_model = onnx.load_model_from_string(buffer.read()) + 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) From 7161de66588c658ff8ea822a549b9a835062f15b Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Sat, 12 Sep 2026 02:04:30 +0000 Subject: [PATCH 08/10] [5565357] Simplify Diffusers quantized export Centralize NVFP4 opset handling, infer restored quantization state directly, and derive the FP8 MHA policy from its quantizers. Simplify the focused recipe and ONNX export coverage. Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/ONNX-TRT-Deployment.md | 2 +- .../quantization/onnx_utils/export.py | 14 -- examples/diffusers/quantization/quantize.py | 22 +-- examples/diffusers/quantization/utils.py | 3 - modelopt/onnx/export/nvfp4_exporter.py | 6 + .../plugins/diffusion/diffusers.py | 10 +- tests/unit/examples/test_diffusers_fp4.py | 134 +++++++----------- .../unit/onnx/quantization/test_qdq_utils.py | 6 +- .../quantization/test_onnx_export_cpu.py | 40 ++---- 9 files changed, 88 insertions(+), 149 deletions(-) diff --git a/examples/diffusers/quantization/ONNX-TRT-Deployment.md b/examples/diffusers/quantization/ONNX-TRT-Deployment.md index 33aab492380..b6933c88f72 100644 --- a/examples/diffusers/quantization/ONNX-TRT-Deployment.md +++ b/examples/diffusers/quantization/ONNX-TRT-Deployment.md @@ -28,7 +28,7 @@ python quantize.py \ #### FLUX-Dev|SDXL|SDXL-Turbo|LTX-Video FP8/FP4 [Script](./quantize.py) -FP4 ONNX export is supported for Flux and SDXL. SDXL uses block-16 NVFP4 for non-QKV Linear/GEMM layers and FP8 for Conv2d layers, while Q/K/V projection Linears remain in the model dtype to preserve TensorRT fusion. Add `--quantize-mha` to optionally quantize MHA with FP8. FP4 deployment requires a Blackwell GPU and TensorRT with NVFP4 support. +FP4 ONNX export is supported for Flux and SDXL. ```sh python quantize.py \ diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 37895eeb7b5..18846d84733 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -449,24 +449,12 @@ def save_onnx(onnx_model, output): print(f"ONNX model saved to {output}") -def _ensure_default_opset(onnx_model, minimum_version): - opset_import = next( - (item for item in onnx_model.opset_import if item.domain in {"", "ai.onnx"}), None - ) - if opset_import is None: - opset_import = onnx_model.opset_import.add() - opset_import.domain = "" - opset_import.version = max(opset_import.version, minimum_version) - - def modelopt_export_sd(backbone, onnx_dir, model_name, precision): model_file_name = "model.onnx" os.makedirs(f"{onnx_dir}", exist_ok=True) tmp_subfolder = tempfile.mkdtemp(prefix="myapp_") tmp_output = Path(f"{tmp_subfolder}/{model_file_name}") q_output = Path(f"{onnx_dir}/{model_file_name}") - is_sdxl_fp4 = precision == "fp4" and model_name in {"sdxl-1.0", "sdxl-turbo"} - quantizer_context = ( configure_linear_module_onnx_quantizers(backbone) if precision == "fp4" else nullcontext() ) @@ -537,7 +525,5 @@ def modelopt_export_sd(backbone, onnx_dir, model_name, precision): flux_convert_rope_weight_type(onnx_model) if precision == "fp4": onnx_model = NVFP4QuantExporter.process_model(onnx_model) - if is_sdxl_fp4: - _ensure_default_opset(onnx_model, 23) save_onnx(onnx_model, q_output) shutil.rmtree(tmp_subfolder, ignore_errors=True) diff --git a/examples/diffusers/quantization/quantize.py b/examples/diffusers/quantization/quantize.py index 0eace465353..9262adbaac1 100644 --- a/examples/diffusers/quantization/quantize.py +++ b/examples/diffusers/quantization/quantize.py @@ -586,34 +586,18 @@ def create_argument_parser() -> argparse.ArgumentParser: return parser -def _restore_quantization_policy( +def _infer_restored_quantization_format( backbones: list[tuple[str, torch.nn.Module]], ) -> QuantFormat: has_nvfp4 = False has_fp8 = False - for backbone_name, backbone in backbones: + 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 backbone_name in ("video_decoder", "vae"): - continue - q_quantizer = getattr(module, "q_bmm_quantizer", None) - k_quantizer = getattr(module, "k_bmm_quantizer", None) - v_quantizer = getattr(module, "v_bmm_quantizer", None) - if not ( - isinstance(q_quantizer, TensorQuantizer) - and isinstance(k_quantizer, TensorQuantizer) - and isinstance(v_quantizer, TensorQuantizer) - ): - continue - module._disable_fp8_mha = not all( - quantizer.is_enabled and quantizer.is_fp8 - for quantizer in (q_quantizer, k_quantizer, v_quantizer) - ) - if has_nvfp4: return QuantFormat.FP4 if has_fp8: @@ -708,7 +692,7 @@ def main() -> None: if export_config.restore_from: export_manager.restore_checkpoint() - quant_config.format = _restore_quantization_policy( + quant_config.format = _infer_restored_quantization_format( list(pipeline_manager.iter_backbones()) ) logger.info(f"Detected restored quantization format: {quant_config.format.value}") diff --git a/examples/diffusers/quantization/utils.py b/examples/diffusers/quantization/utils.py index c3cfdcd5cdd..b7a79e49e70 100644 --- a/examples/diffusers/quantization/utils.py +++ b/examples/diffusers/quantization/utils.py @@ -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: diff --git a/modelopt/onnx/export/nvfp4_exporter.py b/modelopt/onnx/export/nvfp4_exporter.py index 338e2725b14..42598af5475 100644 --- a/modelopt/onnx/export/nvfp4_exporter.py +++ b/modelopt/onnx/export/nvfp4_exporter.py @@ -430,4 +430,10 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): 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 diff --git a/modelopt/torch/quantization/plugins/diffusion/diffusers.py b/modelopt/torch/quantization/plugins/diffusion/diffusers.py index f2f6a702479..e92c775c05b 100644 --- a/modelopt/torch/quantization/plugins/diffusion/diffusers.py +++ b/modelopt/torch/quantization/plugins/diffusion/diffusers.py @@ -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, ) ) diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index 3e2d891673d..16b121ce3f7 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -23,14 +23,14 @@ import torch from torch import nn -onnx = pytest.importorskip("onnx") +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" @@ -74,7 +74,7 @@ def _load_quantize_example(): QuantFormat = _quantize.QuantFormat QuantizationConfig = _quantize.QuantizationConfig Quantizer = _quantize.Quantizer -_restore_quantization_policy = _quantize._restore_quantization_policy +_infer_restored_quantization_format = _quantize._infer_restored_quantization_format class _RecipeBackbone(nn.Module): @@ -88,20 +88,21 @@ def __init__(self): self.conv = nn.Conv2d(4, 4, kernel_size=1, bias=False) -def _quantizer(*, num_bits=(4, 3), enabled=True, calibrated=True, block_sizes=None): +def _quantizer(*, num_bits, enabled=True, block_sizes=None): quantizer = TensorQuantizer( QuantizerAttributeConfig(num_bits=num_bits, axis=None, block_sizes=block_sizes) ) - if calibrated: - quantizer.amax = torch.tensor(448.0) + quantizer.amax = torch.tensor(448.0) if not enabled: quantizer.disable() return quantizer -def _add_quantizers(module, **kwargs): - module.input_quantizer = _quantizer(**kwargs) - module.weight_quantizer = _quantizer(**kwargs) +_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]) @@ -129,68 +130,67 @@ def test_sdxl_fp4_recipe(model_type): @pytest.mark.parametrize( - ("scenario", "expected_format", "mha_enabled"), + ("format_config", "mha_config", "expected_format", "disable_fp8_mha"), [ - ("mixed-fp4", QuantFormat.FP4, True), - ("fp8", QuantFormat.FP8, True), - ("int8-disabled-fp8-mha", QuantFormat.INT8, False), + 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_restore_policy_uses_enabled_checkpoint_state(scenario, expected_format, mha_enabled): +def test_restored_quantizer_state_drives_format_and_fp8_mha( + monkeypatch, format_config, mha_config, expected_format, disable_fp8_mha +): backbone = nn.Module() - backbone.linear = nn.Linear(16, 16) - if scenario == "mixed-fp4": - _add_quantizers( - backbone.linear, - num_bits=(2, 1), - block_sizes={-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, - ) - backbone.conv = nn.Conv2d(1, 1, 1) - _add_quantizers(backbone.conv) - elif scenario == "fp8": - _add_quantizers(backbone.linear) - else: - _add_quantizers(backbone.linear, num_bits=8) - + 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(enabled=mha_enabled)) - quantizers = [module for module in backbone.modules() if isinstance(module, TensorQuantizer)] - state = {q: (q.is_enabled, q._num_bits, q._amax) for q in quantizers} + setattr(backbone.attention, name, _quantizer(**mha_config)) + backbone.attention.bmm2_output_quantizer = lambda output: output - restored_format = _restore_quantization_policy([("transformer", backbone)]) + 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 restored_format == expected_format - assert backbone.attention._disable_fp8_mha is not mha_enabled - for quantizer, (enabled, num_bits, amax) in state.items(): - assert quantizer.is_enabled is enabled - assert quantizer._num_bits == num_bits - assert quantizer._amax is amax + 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_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_path): +def test_restore_infers_checkpoint_format_for_export(monkeypatch, tmp_path): backbone = nn.Module() - backbone.linear = nn.Linear(16, 16) - _add_quantizers( - backbone.linear, - num_bits=(2, 1), - block_sizes={-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, - ) - backbone.attention = nn.Module() - mha_quantizers = [] - for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"): - quantizer = _quantizer() - setattr(backbone.attention, name, quantizer) - mha_quantizers.append(quantizer) + backbone.quantizer = _quantizer(**_NVFP4_QUANTIZER_CONFIG) pipeline_manager = Mock() pipeline_manager.create_pipeline.return_value = object() - pipeline_manager.iter_backbones.side_effect = lambda: iter([("transformer", backbone)]) + 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(torch.nn, "RMSNorm", torch.nn.RMSNorm) - monkeypatch.setattr(torch.nn.modules.normalization, "RMSNorm", torch.nn.RMSNorm) monkeypatch.setattr( sys, "argv", @@ -210,29 +210,3 @@ def test_restore_defaults_reach_exports_with_checkpoint_policy(monkeypatch, tmp_ 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() - assert all(quantizer.is_enabled for quantizer in mha_quantizers) - assert backbone.attention._disable_fp8_mha is False - - -def test_sdxl_fp4_export_lowers_nvfp4_at_opset_23(monkeypatch, tmp_path): - model = onnx.ModelProto() - model.opset_import.add(domain="", version=20) - 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: model) - process_model = Mock(side_effect=lambda current_model: current_model) - monkeypatch.setattr(diffusion_export.NVFP4QuantExporter, "process_model", process_model) - save_onnx = Mock() - monkeypatch.setattr(diffusion_export, "save_onnx", save_onnx) - - diffusion_export.modelopt_export_sd(nn.Module(), tmp_path, "sdxl-1.0", "fp4") - - process_model.assert_called_once_with(model) - save_onnx.assert_called_once_with(model, tmp_path / "model.onnx") - assert ( - next(opset.version for opset in model.opset_import if opset.domain in {"", "ai.onnx"}) == 23 - ) diff --git a/tests/unit/onnx/quantization/test_qdq_utils.py b/tests/unit/onnx/quantization/test_qdq_utils.py index 4b1e69ec538..ebba851ff62 100644 --- a/tests/unit/onnx/quantization/test_qdq_utils.py +++ b/tests/unit/onnx/quantization/test_qdq_utils.py @@ -34,6 +34,7 @@ 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): @@ -335,8 +336,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: @@ -607,6 +607,8 @@ def test_fp4qdq_conversion(self, with_transpose): # 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 diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index 8ea90407a60..9b470f7b019 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -39,6 +39,15 @@ 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"), @@ -68,22 +77,7 @@ def test_fp8_conv_export_preserves_custom_qdq_and_kernel_shape(): forward_loop=lambda quantized_model: quantized_model(sample_input), ) - buffer = io.BytesIO() - if "enable_onnx_checker" in inspect.signature(torch.onnx.export).parameters: - kwargs = {"enable_onnx_checker": False} - else: - kwargs = {} - torch.onnx.export( - model, - sample_input, - buffer, - opset_version=20, - dynamo=False, - **kwargs, - ) - - buffer.seek(0) - exported_model = onnx.load_model_from_string(buffer.read()) + 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") @@ -128,26 +122,14 @@ def cpu_dynamic_block_quantize(inputs, *args): 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) From 034fe23ec86a061973b38b4266137bfa5cb914cd Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Sat, 12 Sep 2026 02:58:48 +0000 Subject: [PATCH 09/10] [5565357] Complete FP8 export follow-up Preserve the converted Flux graph, remove the obsolete FP8 zero-point rewrite, and cover the SD3 FP8 export path. Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- .../quantization/onnx_utils/export.py | 2 +- .../onnx_utils/fp8_onnx_graphsurgeon.py | 23 ------------------- tests/examples/diffusers/test_diffusers.py | 12 ++++++++++ tests/unit/examples/test_diffusers_fp4.py | 22 ++++++++++++++++++ 4 files changed, 35 insertions(+), 24 deletions(-) diff --git a/examples/diffusers/quantization/onnx_utils/export.py b/examples/diffusers/quantization/onnx_utils/export.py index 18846d84733..246dd1a9881 100644 --- a/examples/diffusers/quantization/onnx_utils/export.py +++ b/examples/diffusers/quantization/onnx_utils/export.py @@ -522,7 +522,7 @@ 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" and model_name.startswith("flux"): - flux_convert_rope_weight_type(onnx_model) + 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) diff --git a/examples/diffusers/quantization/onnx_utils/fp8_onnx_graphsurgeon.py b/examples/diffusers/quantization/onnx_utils/fp8_onnx_graphsurgeon.py index 7194e672635..7904d27e0a3 100644 --- a/examples/diffusers/quantization/onnx_utils/fp8_onnx_graphsurgeon.py +++ b/examples/diffusers/quantization/onnx_utils/fp8_onnx_graphsurgeon.py @@ -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 diff --git a/tests/examples/diffusers/test_diffusers.py b/tests/examples/diffusers/test_diffusers.py index ff73093b5a7..979894819de 100644 --- a/tests/examples/diffusers/test_diffusers.py +++ b/tests/examples/diffusers/test_diffusers.py @@ -117,6 +117,17 @@ def inference(self, tmp_path: Path) -> None: 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", @@ -151,6 +162,7 @@ def inference(self, tmp_path: Path) -> None: 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", diff --git a/tests/unit/examples/test_diffusers_fp4.py b/tests/unit/examples/test_diffusers_fp4.py index 16b121ce3f7..16023e509ad 100644 --- a/tests/unit/examples/test_diffusers_fp4.py +++ b/tests/unit/examples/test_diffusers_fp4.py @@ -28,6 +28,7 @@ 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 @@ -210,3 +211,24 @@ def test_restore_infers_checkpoint_format_for_export(monkeypatch, tmp_path): 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") From 4f18eaae918bf430b30448e76ca398b8649a54d8 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Sat, 12 Sep 2026 03:13:36 +0000 Subject: [PATCH 10/10] [5565357] Document shared ONNX export changes Record the NVFP4 opset, FP8 shape inference, and Diffusers attention policy changes in the changelog. Co-authored-by: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- CHANGELOG.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 5a20007eef7..94ed497ced7 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -12,6 +12,7 @@ Changelog **Bug Fixes** +- Fix shared ONNX export metadata and Diffusers attention policy: every ``NVFP4QuantExporter`` post-process now upgrades the default-domain opset to at least 23, all FP8 custom-op exports re-run ONNX shape/type inference after setting output metadata, and quantized SDPA derives FP8 MHA enablement from the live Q/K/V quantizers instead of honoring a caller-set ``_disable_fp8_mha`` attribute. - Fix ``megatron_generate`` dropping the VLM vision inputs (``pixel_values`` / ``image_grid_thw`` / ``image_sizes``) after the first generated token when KV-cache decoding is off, including the automatic fallback under sequence parallelism, which made generation silently ignore the image. No other ModelOpt feature is affected. 0.47.0 (2026-09-xx)