From 51493d4e8fdc3fc1eec443b1d5c59f7a0369f4bc Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Fri, 11 Sep 2026 20:59:28 +0000 Subject: [PATCH 1/7] [6721556] Fix NVFP4 FP16 ONNX conversion Co-Authored-By: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- CHANGELOG.rst | 1 + modelopt/torch/_deploy/utils/torch_onnx.py | 11 ++-- .../deploy/utils/test_torch_onnx_utils.py | 52 +++++++++++++++++++ 3 files changed, 61 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 3e10c786753..012c9748ee6 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -25,6 +25,7 @@ Changelog **Bug Fixes** +- Fix FP16 ONNX exports of NVFP4-quantized models that failed TensorRT parsing due to mixed-precision elementwise inputs. - Fix ``--use_fsdp2`` HuggingFace checkpoint export gathering the whole model onto rank 0, which made export the dominant phase of a PTQ run and could exhaust host memory on large models. The model is now split into per-decoder-layer units dealt round-robin across ranks; each rank gathers every unit but keeps, packs, and writes only the ones it owns, so a rank buffers roughly ``model / world_size`` instead of the whole checkpoint, and rank 0 writes the combined index. Export configurations that cannot be split this way now raise instead of producing a mismatched checkpoint: FSDP2 combined with another DTensor parallelism (for example FSDP2 + tensor parallel on a 2-D mesh; HSDP is supported), models whose decoder layers cannot be discovered, a decoder layer object reused across layers, and a module that holds the decoder layers while owning parameters of its own. - Speed up ``mtq.quantize`` on FSDP2-sharded fused-MoE models. Promoting static-block weight quantizers gathered each expert's slice of the fused weight across ranks even though only quantizer state is read, adding a collective per expert to calibration. - Add FP8 and INT8 recipes that quantize timm ResNet shortcut inputs immediately before residual adds. The torch ONNX example now accepts PTQ and AutoQuantize recipes through ``--recipe`` and uses ``--qformat`` when no recipe is provided. ResNet supports only FP8 and INT8 because TensorRT has limited convolution kernel support; AutoQuantize and other quantization formats are no longer supported for ResNet. diff --git a/modelopt/torch/_deploy/utils/torch_onnx.py b/modelopt/torch/_deploy/utils/torch_onnx.py index b217f1188b3..77b837ee18b 100644 --- a/modelopt/torch/_deploy/utils/torch_onnx.py +++ b/modelopt/torch/_deploy/utils/torch_onnx.py @@ -590,9 +590,14 @@ def get_onnx_bytes_and_metadata( ) return onnx_model.to_bytes(), model_metadata - if weights_dtype == "fp16" and uses_fp8 and torch.bfloat16 in source_parameter_dtypes: + if ( + weights_dtype == "fp16" + and (uses_fp4 or uses_fp8) + and torch.bfloat16 in source_parameter_dtypes + ): + quantization_format = "NVFP4" if uses_fp4 else "FP8" raise ValueError( - "Converting a BF16 FP8 ONNX graph to FP16 is not supported yet " + f"Converting a BF16 {quantization_format} ONNX graph to FP16 is not supported yet " f"(source parameter dtypes: {source_parameter_dtype_names})" ) @@ -662,7 +667,7 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph = qdq_to_dq(onnx_opt_graph) if weights_dtype in ["fp16", "bf16"] and not is_bf16_fp8_noop: - if uses_other_unsupported_quantizer or uses_fp8: + if weights_dtype == "fp16" and (uses_fp4 or uses_other_unsupported_quantizer or uses_fp8): onnx_opt_graph = convert_float_to_float16( onnx_opt_graph, keep_io_types=False, diff --git a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py index 2fd1a1c8504..f9c49f00566 100644 --- a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py +++ b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py @@ -26,6 +26,7 @@ import torch.nn as nn from _test_utils.torch.deploy.lib_test_models import BaseDeployModel, get_deploy_models +import modelopt.torch._deploy.utils.torch_onnx as torch_onnx import modelopt.torch.quantization as mtq from modelopt.onnx.utils import get_batch_size_from_bytes, validate_batch_size from modelopt.torch._deploy.utils import ( @@ -299,6 +300,57 @@ def test_fp8_export_rejects_unsupported_dtype_conversion( assert not any(tmp_path.iterdir()) +def test_nvfp4_export_rejects_bf16_to_fp16(monkeypatch, tmp_path): + monkeypatch.setattr(tempfile, "tempdir", str(tmp_path)) + monkeypatch.setattr(torch_onnx, "is_fp4_quantized", lambda _: True) + model = nn.Linear(4, 4).eval().bfloat16() + + with pytest.raises( + ValueError, + match=r"Converting a BF16 NVFP4 ONNX graph to FP16.*torch.bfloat16", + ): + get_onnx_bytes_and_metadata( + model, + (torch.ones(1, 4, dtype=torch.bfloat16),), + weights_dtype="fp16", + ) + assert not any(tmp_path.iterdir()) + + +@pytest.mark.parametrize( + ("weights_dtype", "expected_converter"), + [("fp16", "onnxconverter"), ("bf16", "autocast")], +) +def test_nvfp4_export_selects_precision_converter(weights_dtype, expected_converter, monkeypatch): + calls = [] + + def record_onnxconverter(model, **kwargs): + calls.append("onnxconverter") + return model + + def record_autocast(model, **kwargs): + calls.append("autocast") + return model + + monkeypatch.setattr(torch_onnx, "is_fp4_quantized", lambda _: True) + monkeypatch.setattr( + torch_onnx, "configure_linear_module_onnx_quantizers", lambda _: nullcontext() + ) + monkeypatch.setattr(torch_onnx, "quantize_weights", lambda _, graph: graph) + monkeypatch.setattr(torch_onnx, "convert_float_to_float16", record_onnxconverter) + monkeypatch.setattr(torch_onnx, "convert_to_f16", record_autocast) + + model = nn.Linear(4, 4).eval() + get_onnx_bytes_and_metadata( + model, + (torch.ones(1, 4),), + weights_dtype=weights_dtype, + onnx_opset=23, + ) + + assert calls == [expected_converter] + + class SingleArgModel(nn.Module): def forward(self, x: torch.Tensor): return torch.add(x, x) - x From 1ffc34f89a28ba846dcea971d0af45cf6eccd868 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Fri, 11 Sep 2026 22:48:18 +0000 Subject: [PATCH 2/7] fix: preserve NVFP4 export precision metadata Co-Authored-By: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- modelopt/onnx/export/nvfp4_exporter.py | 8 +- modelopt/torch/_deploy/utils/torch_onnx.py | 30 ++++-- .../deploy/utils/test_torch_onnx_utils.py | 35 ++++--- .../quantization/test_onnx_export_cpu.py | 93 ++++++++++++++++++- 4 files changed, 142 insertions(+), 24 deletions(-) diff --git a/modelopt/onnx/export/nvfp4_exporter.py b/modelopt/onnx/export/nvfp4_exporter.py index 338e2725b14..512d051d108 100644 --- a/modelopt/onnx/export/nvfp4_exporter.py +++ b/modelopt/onnx/export/nvfp4_exporter.py @@ -331,7 +331,7 @@ def post_process(onnx_model: onnx.ModelProto) -> onnx.ModelProto: initializer_indices = { initializer.name: idx for idx, initializer in enumerate(graph.initializer) } - value_info_map = {vi.name: vi for vi in graph.value_info} + value_info_map = {vi.name: vi for vi in [*graph.value_info, *graph.output]} graph_inputs = {inp.name for inp in graph.input} cast_output_cache: dict[tuple[str, str], str] = {} @@ -373,6 +373,12 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): # Update the target node input to use the cast node output node.input[i] = cast_output_name + for output_name in node.output: + if output_name in value_info_map: + value_info_map[output_name].type.tensor_type.elem_type = onnx_dtype_map[ + precision_dtype + ] + precision_dtype = _get_precision_dtype() logger.debug(f"Using precision dtype: {precision_dtype}") diff --git a/modelopt/torch/_deploy/utils/torch_onnx.py b/modelopt/torch/_deploy/utils/torch_onnx.py index 77b837ee18b..7b797d2299e 100644 --- a/modelopt/torch/_deploy/utils/torch_onnx.py +++ b/modelopt/torch/_deploy/utils/torch_onnx.py @@ -35,6 +35,7 @@ from torch.nn.parallel import DataParallel, DistributedDataParallel from modelopt.onnx.autocast.convert import convert_to_f16 +from modelopt.onnx.autocast.graphsanitizer import GraphSanitizer from modelopt.onnx.export import ( FP8QuantExporter, INT4QuantExporter, @@ -52,9 +53,11 @@ fold_qdq_scale_fp16_to_fp32_casts, get_input_names, get_input_shapes, + get_min_opset_for_precisions, get_node_names, get_output_names, get_output_shapes, + get_qdq_precisions, infer_shapes, is_model_too_large_for_protobuf, remove_node_training_mode, @@ -507,10 +510,10 @@ def get_onnx_bytes_and_metadata( `torch.onnx.export `_. onnx_opset: The onnx opset version to use for exporting the model. dq_only: If True, the exported onnx model is converted to a dq_only model. - weights_dtype: Requested high-precision dtype for exported weights. For an FP8 model, - ``"bf16"`` is accepted only when every floating parameter is already BF16. This is - a weight-focused no-op, not a graph-wide conversion: floating buffers are not - considered for eligibility and may preserve higher-precision regions. + weights_dtype: Requested high-precision dtype for exported weights. For an FP8 or NVFP4 + model, ``"bf16"`` is accepted only when every floating parameter is already BF16. + This is a weight-focused no-op, not a graph-wide conversion: floating buffers are + not considered for eligibility and may preserve higher-precision regions. Returns: bytes: Onnx model in bytes. @@ -539,12 +542,19 @@ def get_onnx_bytes_and_metadata( uses_fp8 = is_fp8_quantized(model) uses_int8 = is_int8_quantized(model) uses_other_unsupported_quantizer = is_int4_quantized(model) or uses_mxfp8 or uses_int8 + is_bf16_fp4_noop = ( + weights_dtype == "bf16" + and source_parameter_dtypes == {torch.bfloat16} + and uses_fp4 + and not (uses_fp8 or uses_other_unsupported_quantizer) + ) is_bf16_fp8_noop = ( weights_dtype == "bf16" and source_parameter_dtypes == {torch.bfloat16} and uses_fp8 and not (uses_fp4 or uses_other_unsupported_quantizer) ) + is_bf16_quantized_noop = is_bf16_fp4_noop or is_bf16_fp8_noop # Standardize model args and also tensorize them so they also appear in the onnx graph! # Floats/ints are tensorized when they are provided, but not tensorized when they are not @@ -603,8 +613,8 @@ def get_onnx_bytes_and_metadata( if ( weights_dtype == "bf16" - and (uses_fp8 or uses_other_unsupported_quantizer) - and not is_bf16_fp8_noop + and (uses_fp4 or uses_fp8 or uses_other_unsupported_quantizer) + and not is_bf16_quantized_noop ): raise ValueError( "Converting a quantized ONNX graph to BF16 is not supported yet " @@ -663,10 +673,16 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph = quantize_weights(model, onnx_opt_graph) + if uses_fp4: + qdq_min_opset = get_min_opset_for_precisions(get_qdq_precisions(onnx_opt_graph)) + opset_sanitizer = GraphSanitizer(onnx_opt_graph, min_opset=qdq_min_opset) + opset_sanitizer.convert_opset() + onnx_opt_graph = opset_sanitizer.model + if dq_only: onnx_opt_graph = qdq_to_dq(onnx_opt_graph) - if weights_dtype in ["fp16", "bf16"] and not is_bf16_fp8_noop: + if weights_dtype in ["fp16", "bf16"] and not is_bf16_quantized_noop: if weights_dtype == "fp16" and (uses_fp4 or uses_other_unsupported_quantizer or uses_fp8): onnx_opt_graph = convert_float_to_float16( onnx_opt_graph, diff --git a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py index f9c49f00566..af61bcce553 100644 --- a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py +++ b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py @@ -300,28 +300,41 @@ def test_fp8_export_rejects_unsupported_dtype_conversion( assert not any(tmp_path.iterdir()) -def test_nvfp4_export_rejects_bf16_to_fp16(monkeypatch, tmp_path): +@pytest.mark.parametrize( + ("source_dtype", "weights_dtype"), + [(torch.bfloat16, "fp16"), (torch.float32, "bf16")], + ids=["bf16-to-fp16", "fp32-to-bf16"], +) +def test_nvfp4_export_rejects_unsupported_dtype_conversion( + source_dtype, weights_dtype, monkeypatch, tmp_path +): monkeypatch.setattr(tempfile, "tempdir", str(tmp_path)) monkeypatch.setattr(torch_onnx, "is_fp4_quantized", lambda _: True) - model = nn.Linear(4, 4).eval().bfloat16() + model = nn.Linear(4, 4).eval().to(source_dtype) with pytest.raises( ValueError, - match=r"Converting a BF16 NVFP4 ONNX graph to FP16.*torch.bfloat16", + match=rf"Converting .* to {weights_dtype.upper()}.*source parameter dtypes: {source_dtype}", ): get_onnx_bytes_and_metadata( model, - (torch.ones(1, 4, dtype=torch.bfloat16),), - weights_dtype="fp16", + (torch.ones(1, 4, dtype=source_dtype),), + weights_dtype=weights_dtype, ) assert not any(tmp_path.iterdir()) @pytest.mark.parametrize( - ("weights_dtype", "expected_converter"), - [("fp16", "onnxconverter"), ("bf16", "autocast")], + ("source_dtype", "weights_dtype", "expected_calls"), + [ + (torch.float32, "fp16", ["onnxconverter"]), + (torch.bfloat16, "bf16", []), + ], + ids=["fp32-to-fp16", "bf16-noop"], ) -def test_nvfp4_export_selects_precision_converter(weights_dtype, expected_converter, monkeypatch): +def test_nvfp4_export_selects_precision_converter( + source_dtype, weights_dtype, expected_calls, monkeypatch +): calls = [] def record_onnxconverter(model, **kwargs): @@ -340,15 +353,15 @@ def record_autocast(model, **kwargs): monkeypatch.setattr(torch_onnx, "convert_float_to_float16", record_onnxconverter) monkeypatch.setattr(torch_onnx, "convert_to_f16", record_autocast) - model = nn.Linear(4, 4).eval() + model = nn.Linear(4, 4).eval().to(source_dtype) get_onnx_bytes_and_metadata( model, - (torch.ones(1, 4),), + (torch.ones(1, 4, dtype=source_dtype),), weights_dtype=weights_dtype, onnx_opset=23, ) - assert calls == [expected_converter] + assert calls == expected_calls class SingleArgModel(nn.Module): diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index ce2ef626d63..3daa2ba6dab 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -35,6 +35,7 @@ from modelopt.onnx.export import NVFP4QuantExporter from modelopt.onnx.export.nvfp4_exporter import _encode_nvfp4_block_scale from modelopt.onnx.quantization.qdq_utils import fp4qdq_to_2dq +from modelopt.torch._deploy.utils import OnnxBytes, get_onnx_bytes_and_metadata from modelopt.torch.quantization.qtensor import NVFP4QTensor from modelopt.torch.quantization.utils import is_quantized_linear @@ -59,7 +60,17 @@ def test_onnx_export_cpu(model_cls, num_bits, per_channel_quantization, constant ) -def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch): +class _NVFP4LinearWithExplicitBias(torch.nn.Module): + def __init__(self, dtype): + super().__init__() + self.linear = torch.nn.Linear(16, 16, bias=False, dtype=dtype) + self.bias = torch.nn.Parameter(torch.ones(16, dtype=dtype)) + + def forward(self, inputs): + return self.linear(inputs) + self.bias + + +def _make_cpu_nvfp4_model(monkeypatch, model, sample_input, disable_input_quantizers=False): def forward_loop(model): model(sample_input) @@ -67,17 +78,23 @@ def cpu_dynamic_block_quantize(inputs, *args): return inputs monkeypatch.setattr(tensor_quant, "dynamic_block_quantize_op", cpu_dynamic_block_quantize) - - model = SimpleLinear().eval() - sample_input = model.get_input() model = mtq.quantize(model, mtq.NVFP4_DEFAULT_CFG, forward_loop=forward_loop) for module in model.modules(): assert not isinstance(module, torch.nn.Linear) or is_quantized_linear(module) if isinstance(module, torch.nn.Linear): - module.input_quantizer.disable() + if disable_input_quantizers: + module.input_quantizer.disable() module.weight_quantizer._onnx_quantizer_type = "static" + return model + + +def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch): + model = SimpleLinear().eval() + sample_input = model.get_input() + model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input, disable_input_quantizers=True) + buffer = io.BytesIO() if "enable_onnx_checker" in inspect.signature(torch.onnx.export).parameters: kwargs = {"enable_onnx_checker": False} @@ -105,6 +122,72 @@ def cpu_dynamic_block_quantize(inputs, *args): onnx.checker.check_model(converted_model) +@pytest.mark.parametrize( + ("source_dtype", "weights_dtype", "expected_dtype"), + [ + (torch.float32, "fp16", TensorProto.FLOAT16), + (torch.bfloat16, "bf16", TensorProto.BFLOAT16), + ], + ids=["fp32-to-fp16", "bf16-noop"], +) +def test_nvfp4_deploy_export_has_consistent_elementwise_types( + monkeypatch, source_dtype, weights_dtype, expected_dtype +): + model = _NVFP4LinearWithExplicitBias(source_dtype).eval() + sample_input = torch.ones(1, 2, 16, dtype=source_dtype) + model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) + + onnx_bytes, _ = get_onnx_bytes_and_metadata( + model, + (sample_input,), + weights_dtype=weights_dtype, + ) + exported_model = onnx.load_model_from_string( + OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + ) + + assert utils.get_opset_version(exported_model) >= 23 + onnx.checker.check_model(exported_model, full_check=True) + assert any( + node.op_type == "DequantizeLinear" + and any(attribute.name == "block_size" for attribute in node.attribute) + for node in exported_model.graph.node + ) + assert any(node.op_type == "TRT_FP4DynamicQuantize" for node in exported_model.graph.node) + + inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) + tensor_types = { + initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer + } + for value in [ + *inferred_model.graph.input, + *inferred_model.graph.value_info, + *inferred_model.graph.output, + ]: + if value.type.HasField("tensor_type"): + tensor_types[value.name] = value.type.tensor_type.elem_type + + floating_types = {TensorProto.FLOAT, TensorProto.FLOAT16, TensorProto.BFLOAT16} + elementwise_nodes = [ + node + for node in inferred_model.graph.node + if node.op_type in {"Add", "Sub", "Mul", "Div", "Pow"} + ] + assert any(node.op_type == "Add" for node in elementwise_nodes) + for node in elementwise_nodes: + input_types = [tensor_types[input_name] for input_name in node.input] + assert len(set(input_types) & floating_types) <= 1, ( + node.name, + [TensorProto.DataType.Name(input_type) for input_type in input_types], + ) + + add_node = next(node for node in elementwise_nodes if node.op_type == "Add") + assert [tensor_types[input_name] for input_name in add_node.input] == [ + expected_dtype, + expected_dtype, + ] + + @pytest.mark.parametrize( ("convert", "deprecated"), [ From c97ae4a5d2b566f361ba73d497e5951ac358a5ab Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Sat, 12 Sep 2026 04:42:35 +0000 Subject: [PATCH 3/7] fix: preserve NVFP4 FP32 graph boundaries Co-Authored-By: Codex Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- CHANGELOG.rst | 2 +- modelopt/onnx/export/nvfp4_exporter.py | 43 +++++++++++++--- .../quantization/test_onnx_export_cpu.py | 49 ++++++++++++++++++- 3 files changed, 84 insertions(+), 10 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 012c9748ee6..c56614fd579 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -25,7 +25,7 @@ Changelog **Bug Fixes** -- Fix FP16 ONNX exports of NVFP4-quantized models that failed TensorRT parsing due to mixed-precision elementwise inputs. +- Fix NVFP4 ONNX exports that failed ONNX or TensorRT parsing due to mixed-precision ``MatMul``, ``Gemm``, and elementwise inputs. Export now preserves FP32 graph boundaries around low-precision NVFP4 compute and rejects unsupported FP32-to-BF16 and BF16-to-FP16 conversions; use FP32-to-FP16 or native-BF16 export instead. - Fix ``--use_fsdp2`` HuggingFace checkpoint export gathering the whole model onto rank 0, which made export the dominant phase of a PTQ run and could exhaust host memory on large models. The model is now split into per-decoder-layer units dealt round-robin across ranks; each rank gathers every unit but keeps, packs, and writes only the ones it owns, so a rank buffers roughly ``model / world_size`` instead of the whole checkpoint, and rank 0 writes the combined index. Export configurations that cannot be split this way now raise instead of producing a mismatched checkpoint: FSDP2 combined with another DTensor parallelism (for example FSDP2 + tensor parallel on a 2-D mesh; HSDP is supported), models whose decoder layers cannot be discovered, a decoder layer object reused across layers, and a module that holds the decoder layers while owning parameters of its own. - Speed up ``mtq.quantize`` on FSDP2-sharded fused-MoE models. Promoting static-block weight quantizers gathered each expert's slice of the fused weight across ranks even though only quantizer state is read, adding a collective per expert to calibration. - Add FP8 and INT8 recipes that quantize timm ResNet shortcut inputs immediately before residual adds. The torch ONNX example now accepts PTQ and AutoQuantize recipes through ``--recipe`` and uses ``--qformat`` when no recipe is provided. ResNet supports only FP8 and INT8 because TensorRT has limited convolution kernel support; AutoQuantize and other quantization formats are no longer supported for ResNet. diff --git a/modelopt/onnx/export/nvfp4_exporter.py b/modelopt/onnx/export/nvfp4_exporter.py index 512d051d108..cd6768c5063 100644 --- a/modelopt/onnx/export/nvfp4_exporter.py +++ b/modelopt/onnx/export/nvfp4_exporter.py @@ -351,11 +351,15 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): assert maybe_matmul.op_type == "MatMul" node = maybe_matmul - # Create Cast nodes for each input of the target node except bias - for i, input_name in enumerate(node.input[:2]): + precision_onnx_dtype = onnx_dtype_map[precision_dtype] + cast_output_suffix = "bf16" if precision_dtype == "BFloat16" else "f16" + + compute_inputs = node.input[:3] if node.op_type == "Gemm" else node.input[:2] + for i, input_name in enumerate(compute_inputs): + if not input_name: + continue cast_output_name = cast_output_cache.get((input_name, precision_dtype)) if cast_output_name is None: - cast_output_suffix = "bf16" if precision_dtype == "BFloat16" else "f16" cast_output_name = f"{input_name}_{cast_output_suffix}" cast_output_cache[(input_name, precision_dtype)] = cast_output_name @@ -364,7 +368,7 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): "Cast", inputs=[input_name], # Original input of the target node outputs=[cast_output_name], - to=onnx_dtype_map[precision_dtype], # Cast to FP16/BF16 + to=precision_onnx_dtype, ) # Insert the Cast node into the graph @@ -373,11 +377,34 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): # Update the target node input to use the cast node output node.input[i] = cast_output_name - for output_name in node.output: - if output_name in value_info_map: - value_info_map[output_name].type.tensor_type.elem_type = onnx_dtype_map[ - precision_dtype + for i, output_name in enumerate(node.output): + output_value_info = value_info_map.get(output_name) + if output_value_info is None: + continue + + output_dtype = output_value_info.type.tensor_type.elem_type + if precision_dtype == "BFloat16" or output_dtype == precision_onnx_dtype: + output_value_info.type.tensor_type.elem_type = precision_onnx_dtype + continue + + precision_output_name = f"{output_name}_{cast_output_suffix}_output" + precision_output_value_info = onnx.ValueInfoProto() + precision_output_value_info.CopyFrom(output_value_info) + precision_output_value_info.name = precision_output_name + precision_output_value_info.type.tensor_type.elem_type = precision_onnx_dtype + graph.value_info.append(precision_output_value_info) + value_info_map[precision_output_name] = precision_output_value_info + node.output[i] = precision_output_name + graph.node.extend( + [ + onnx.helper.make_node( + "Cast", + inputs=[precision_output_name], + outputs=[output_name], + to=output_dtype, + ) ] + ) precision_dtype = _get_precision_dtype() logger.debug(f"Using precision dtype: {precision_dtype}") diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index 3daa2ba6dab..e3d88cc0c74 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -125,10 +125,11 @@ def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch): @pytest.mark.parametrize( ("source_dtype", "weights_dtype", "expected_dtype"), [ + (torch.float32, "fp32", TensorProto.FLOAT), (torch.float32, "fp16", TensorProto.FLOAT16), (torch.bfloat16, "bf16", TensorProto.BFLOAT16), ], - ids=["fp32-to-fp16", "bf16-noop"], + ids=["fp32-preserved", "fp32-to-fp16", "bf16-noop"], ) def test_nvfp4_deploy_export_has_consistent_elementwise_types( monkeypatch, source_dtype, weights_dtype, expected_dtype @@ -188,6 +189,52 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( ] +@pytest.mark.parametrize( + ("source_dtype", "weights_dtype", "expected_gemm_dtype", "expected_output_dtype"), + [ + (torch.float32, "fp32", TensorProto.FLOAT16, TensorProto.FLOAT), + (torch.float32, "fp16", TensorProto.FLOAT16, TensorProto.FLOAT16), + (torch.bfloat16, "bf16", TensorProto.BFLOAT16, TensorProto.BFLOAT16), + ], + ids=["fp32-preserved", "fp32-to-fp16", "bf16-noop"], +) +def test_nvfp4_deploy_export_has_consistent_gemm_types( + monkeypatch, source_dtype, weights_dtype, expected_gemm_dtype, expected_output_dtype +): + model = torch.nn.Linear(16, 16, dtype=source_dtype).eval() + sample_input = torch.ones(2, 16, dtype=source_dtype) + model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) + + onnx_bytes, _ = get_onnx_bytes_and_metadata( + model, + (sample_input,), + weights_dtype=weights_dtype, + ) + exported_model = onnx.load_model_from_string( + OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + ) + + assert utils.get_opset_version(exported_model) >= 23 + onnx.checker.check_model(exported_model, full_check=True) + assert any(node.op_type == "TRT_FP4DynamicQuantize" for node in exported_model.graph.node) + inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) + tensor_types = { + initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer + } + for value in [ + *inferred_model.graph.input, + *inferred_model.graph.value_info, + *inferred_model.graph.output, + ]: + if value.type.HasField("tensor_type"): + tensor_types[value.name] = value.type.tensor_type.elem_type + + gemm_node = next(node for node in inferred_model.graph.node if node.op_type == "Gemm") + assert [tensor_types[input_name] for input_name in gemm_node.input] == [expected_gemm_dtype] * 3 + assert tensor_types[gemm_node.output[0]] == expected_gemm_dtype + assert tensor_types[inferred_model.graph.output[0].name] == expected_output_dtype + + @pytest.mark.parametrize( ("convert", "deprecated"), [ From 7038eb84d646cbeecdf14806d3074fc7e97b4786 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:49:43 +0000 Subject: [PATCH 4/7] fix: preserve mixed NVFP4 precision boundaries Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- modelopt/onnx/export/nvfp4_exporter.py | 19 ++++---- .../quantization/test_onnx_export_cpu.py | 46 +++++++++++++++++++ 2 files changed, 54 insertions(+), 11 deletions(-) diff --git a/modelopt/onnx/export/nvfp4_exporter.py b/modelopt/onnx/export/nvfp4_exporter.py index cd6768c5063..932f27791dd 100644 --- a/modelopt/onnx/export/nvfp4_exporter.py +++ b/modelopt/onnx/export/nvfp4_exporter.py @@ -335,14 +335,10 @@ def post_process(onnx_model: onnx.ModelProto) -> onnx.ModelProto: graph_inputs = {inp.name for inp in graph.input} cast_output_cache: dict[tuple[str, str], str] = {} - def _get_precision_dtype() -> str: - # Check initializers to determine the precision of the weights - precision_dtype = "Half" - for initializer in graph.initializer: - if initializer.data_type == 16: - precision_dtype = "BFloat16" - break # Assuming all weights are of the same precision - return precision_dtype + def _get_precision_dtype(weight_initializer: onnx.TensorProto) -> str: + return ( + "BFloat16" if weight_initializer.data_type == onnx.TensorProto.BFLOAT16 else "Half" + ) def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): # Change the input types to match weight precision (precision_dtype) @@ -383,6 +379,8 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): continue output_dtype = output_value_info.type.tensor_type.elem_type + # TRT_FP4QDQ leaves native-BF16 outputs annotated FLOAT; real FP32 boundaries are + # explicit Casts. FP16 can convert FP32 graphs, so it restores implicit boundaries. if precision_dtype == "BFloat16" or output_dtype == precision_onnx_dtype: output_value_info.type.tensor_type.elem_type = precision_onnx_dtype continue @@ -406,9 +404,6 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): ] ) - precision_dtype = _get_precision_dtype() - logger.debug(f"Using precision dtype: {precision_dtype}") - fp4_qdq_nodes = [node for node in graph.node if node.op_type == "TRT_FP4QDQ"] logger.debug(f"Found {len(fp4_qdq_nodes)} FP4QDQ nodes to convert") @@ -416,6 +411,8 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): idx = initializer_indices.get(node.input[0]) assert idx is not None, f"Initializer for weight '{node.input[0]}' not found." initializers_to_delete.append(graph.initializer[idx].name) + precision_dtype = _get_precision_dtype(graph.initializer[idx]) + logger.debug(f"Using precision dtype {precision_dtype} for {node.input[0]}") # Retrieve compressed data from node attributes block_size = node.attribute[0].i diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index e3d88cc0c74..2cd82ee1b0c 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -70,6 +70,18 @@ def forward(self, inputs): return self.linear(inputs) + self.bias +class _NVFP4MixedPrecisionLinear(torch.nn.Module): + def __init__(self): + super().__init__() + self.fp32_linear = torch.nn.Linear(16, 16, dtype=torch.float32) + self.bf16_linear = torch.nn.Linear(16, 16, dtype=torch.bfloat16) + + def forward(self, inputs): + fp32_output = self.fp32_linear(inputs) + bf16_output = self.bf16_linear(inputs.to(torch.bfloat16)).float() + return fp32_output + bf16_output + + def _make_cpu_nvfp4_model(monkeypatch, model, sample_input, disable_input_quantizers=False): def forward_loop(model): model(sample_input) @@ -189,6 +201,40 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( ] +def test_nvfp4_deploy_export_preserves_mixed_precision_boundaries(monkeypatch): + model = _NVFP4MixedPrecisionLinear().eval() + sample_input = torch.ones(1, 16) + model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) + + onnx_bytes, _ = get_onnx_bytes_and_metadata( + model, + (sample_input,), + weights_dtype="fp32", + ) + exported_model = onnx.load_model_from_string( + OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + ) + + onnx.checker.check_model(exported_model, full_check=True) + inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) + tensor_types = { + initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer + } + for value in [ + *inferred_model.graph.input, + *inferred_model.graph.value_info, + *inferred_model.graph.output, + ]: + if value.type.HasField("tensor_type"): + tensor_types[value.name] = value.type.tensor_type.elem_type + + add_node = next(node for node in inferred_model.graph.node if node.op_type == "Add") + assert [tensor_types[input_name] for input_name in add_node.input] == [ + TensorProto.FLOAT, + TensorProto.FLOAT, + ] + + @pytest.mark.parametrize( ("source_dtype", "weights_dtype", "expected_gemm_dtype", "expected_output_dtype"), [ From b14b55757c0c279b153b7458ff71c8bae28d6033 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Mon, 21 Sep 2026 17:20:32 +0000 Subject: [PATCH 5/7] fix: allow NVFP4 cross-precision ONNX export Support FP32 or mixed sources exported to BF16 and BF16 or mixed sources exported to FP16, including accompanying FP8 layers. Preserve fixed FP4/FP8 dynamic-quantizer output types before conversion so missing intermediate shapes do not trigger TensorRT parsing of mixed compute dtypes. Cover LayerNorm, Gemm, graph I/O, and FP8 scale boundaries. Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- CHANGELOG.rst | 2 +- modelopt/onnx/autocast/graphsanitizer.py | 30 +++++++-- modelopt/onnx/autocast/precisionconverter.py | 18 +++-- modelopt/onnx/export/nvfp4_exporter.py | 50 +++++++++++++- modelopt/torch/_deploy/utils/torch_onnx.py | 37 +++++++---- .../unit/onnx/autocast/test_graphsanitizer.py | 64 ++++++++++++++++-- .../deploy/utils/test_torch_onnx_utils.py | 28 +------- .../quantization/test_onnx_export_cpu.py | 66 ++++++++++++++++--- 8 files changed, 229 insertions(+), 66 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 59f70617bab..f471cd6f530 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -66,7 +66,7 @@ Changelog - Fix ``examples/megatron_bridge/export_quantized_megatron_to_hf.py`` storing the MoE router at Megatron's ``moe_router_dtype``, which is a routing *compute* dtype, not a storage one. The router now exports at the export ``dtype`` like every other unquantized weight, matching what ``hf_ptq.py`` and the released NVFP4 checkpoints contain; pass ``moe_router_dtype`` to ``export_mcore_gpt_to_hf`` explicitly if you want the old fp32 storage. - Fix unified Megatron export writing a second, unreferenced copy of the vocab embedding when a model with MTP layers is exported with pipeline parallelism. The duplicate was never loaded but inflated the checkpoint by the size of the embedding (about 1 GB for Qwen3.6-35B-A3B); re-export to reclaim the space. - Fail fast on non-finite AutoQuantize output gradients with an actionable error before accumulating sensitivity scores, without changing attention backend settings. -- Fix NVFP4 ONNX exports that failed ONNX or TensorRT parsing due to mixed-precision ``MatMul``, ``Gemm``, and elementwise inputs. Export now preserves FP32 graph boundaries around low-precision NVFP4 compute and rejects unsupported FP32-to-BF16 and BF16-to-FP16 conversions; use FP32-to-FP16 or native-BF16 export instead. +- Fix NVFP4 ONNX exports that failed ONNX or TensorRT parsing due to mixed-precision ``MatMul``, ``Gemm``, and elementwise inputs. Export preserves FP32 graph boundaries around low-precision NVFP4 compute and supports conversion of FP32, BF16, or mixed FP32/BF16 models to FP16 or BF16. - Fix ONNX INT8 entropy calibration failing or producing invalid quantization parameters for FP16 activations. - Fix ``--use_fsdp2`` HuggingFace checkpoint export gathering the whole model onto rank 0, which made export the dominant phase of a PTQ run and could exhaust host memory on large models. The model is now split into per-decoder-layer units dealt round-robin across ranks; each rank gathers every unit but keeps, packs, and writes only the ones it owns, so a rank buffers roughly ``model / world_size`` instead of the whole checkpoint, and rank 0 writes the combined index. Export configurations that cannot be split this way now raise instead of producing a mismatched checkpoint: FSDP2 combined with another DTensor parallelism (for example FSDP2 + tensor parallel on a 2-D mesh; HSDP is supported), models whose decoder layers cannot be discovered, a decoder layer object reused across layers, and a module that holds the decoder layers while owning parameters of its own. - Speed up ``mtq.quantize`` on FSDP2-sharded fused-MoE models. Promoting static-block weight quantizers gathered each expert's slice of the fused weight across ranks even though only quantizer state is read, adding a collective per expert to calibration. diff --git a/modelopt/onnx/autocast/graphsanitizer.py b/modelopt/onnx/autocast/graphsanitizer.py index 58bd678159c..4a30c2484ea 100644 --- a/modelopt/onnx/autocast/graphsanitizer.py +++ b/modelopt/onnx/autocast/graphsanitizer.py @@ -132,11 +132,33 @@ def find_custom_nodes(self) -> None: # Set TensorRT plugin domain info in the graph for ORT compatibility self.model = set_trt_plugin_domain(self.model, self.custom_ops) - # Infer types and shapes in the graph for ORT compatibility - _, all_tensor_info = get_custom_layers(self.onnx_path or self.model, self.trt_plugins) - self.model = infer_types_shapes_tensorrt( - self.model, self.trt_plugins, all_tensor_info=all_tensor_info + tensor_types = { + value.name: value.type.tensor_type.elem_type + for value in [ + *self.model.graph.input, + *self.model.graph.value_info, + *self.model.graph.output, + ] + } + custom_output_types = [ + [tensor_types.get(output, onnx.TensorProto.UNDEFINED) for output in node.output] + for node in self.model.graph.node + if node.op_type in self.custom_ops + ] + has_nvfp4_types = self.custom_ops == {"TRT_FP4DynamicQuantize"} and all( + output_types == [onnx.TensorProto.FLOAT4E2M1, onnx.TensorProto.FLOAT8E4M3FN] + for output_types in custom_output_types ) + # NVFP4 export declares its FP4 data and FP8 scale outputs before dtype conversion. + # Avoid parsing that graph in TensorRT until its compute dtypes are normalized. + if not has_nvfp4_types: + # Infer types and shapes in the graph for ORT compatibility + _, all_tensor_info = get_custom_layers( + self.onnx_path or self.model, self.trt_plugins + ) + self.model = infer_types_shapes_tensorrt( + self.model, self.trt_plugins, all_tensor_info=all_tensor_info + ) def remove_disconnected_outputs(self) -> None: """Remove disconnected outputs from the model.""" diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index a78249fdca4..158890595eb 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -218,7 +218,7 @@ def convert( # Convert inputs to reduced precision type if not self.keep_io_types: for input in self.model.graph.input: - if input.type.tensor_type.elem_type == self.high_precision_type.onnx_type: + if input.type.tensor_type.elem_type in ONNX_TYPES: input.type.tensor_type.elem_type = self.low_precision_type.onnx_type cast_down_tensors, cast_up_tensors, fp32_input_to_low_precision_node = ( @@ -416,6 +416,8 @@ def _infer_shape_op_shape(node): return [max(end - start, 0)] def _infer_standard_op_shape(node): + if node.op == "Cast" and node.inputs: + return _get_shape(node.inputs[0]) for infer_shape in ( _infer_gathernd_op_shape, _infer_gather_op_shape, @@ -430,6 +432,10 @@ def _infer_standard_op_shape(node): graph = gs.import_onnx(model) traversed_tensors = [] + def _get_onnx_dtype(dtype): + # GraphSurgeon preserves BF16/FP8/FP4 as ONNX enums when NumPy has no native type. + return dtype if isinstance(dtype, int) else helper.np_dtype_to_tensor_dtype(dtype) + def _get_np_type(node, inp, opset=onnx.defs.onnx_opset_version()): if node.op == "Cast": return helper.tensor_dtype_to_np_dtype(node.attrs["to"]) @@ -444,7 +450,9 @@ def _get_np_type(node, inp, opset=onnx.defs.onnx_opset_version()): elif node.op not in self.custom_ops: op_schema = onnx.defs.get_schema(node.op, opset) out_types = list(op_schema.outputs[0].types) - inp_type = f"tensor({'float' if inp.dtype == 'float32' else inp.dtype})" + inp_type = ( + f"tensor({TensorProto.DataType.Name(_get_onnx_dtype(inp.dtype)).lower()})" + ) return ( inp.dtype if inp_type in out_types @@ -456,8 +464,8 @@ def _get_np_type(node, inp, opset=onnx.defs.onnx_opset_version()): def _can_propagate_type(from_type, to_type): try: - from_type_onnx = helper.np_dtype_to_tensor_dtype(from_type) - to_type_onnx = helper.np_dtype_to_tensor_dtype(to_type) + from_type_onnx = _get_onnx_dtype(from_type) + to_type_onnx = _get_onnx_dtype(to_type) return ( from_type_onnx in [*ONNX_TYPES, onnx.TensorProto.UNDEFINED] and to_type_onnx in ONNX_TYPES @@ -509,7 +517,7 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): logger.debug( f"{indent}Updated type in {child_out.name} from {child_out.dtype} to {np_type}." ) - elif helper.np_dtype_to_tensor_dtype(np_type) in ONNX_TYPES: + elif _get_onnx_dtype(np_type) in ONNX_TYPES: child_out.dtype = np_type logger.debug( f"{indent}Updated type in {child_out.name} from 'None' to {np_type}." diff --git a/modelopt/onnx/export/nvfp4_exporter.py b/modelopt/onnx/export/nvfp4_exporter.py index 95dea44b4e1..d7a494ece3f 100644 --- a/modelopt/onnx/export/nvfp4_exporter.py +++ b/modelopt/onnx/export/nvfp4_exporter.py @@ -331,10 +331,55 @@ def post_process(onnx_model: onnx.ModelProto) -> onnx.ModelProto: initializer_indices = { initializer.name: idx for idx, initializer in enumerate(graph.initializer) } - value_info_map = {vi.name: vi for vi in [*graph.value_info, *graph.output]} + value_info_map = {vi.name: vi for vi in [*graph.input, *graph.value_info, *graph.output]} graph_inputs = {inp.name for inp in graph.input} cast_output_cache: dict[tuple[str, str], str] = {} + def _annotate_dynamic_quantize_outputs(node: onnx.NodeProto): + # These fixed quantized output types must survive precision conversion without + # invoking TensorRT on the not-yet-normalized graph. + input_value_info = value_info_map.get(node.input[0]) + input_shape = ( + input_value_info.type.tensor_type.shape + if input_value_info is not None + and input_value_info.type.tensor_type.HasField("shape") + else None + ) + + attributes = { + attribute.name: onnx.helper.get_attribute_value(attribute) + for attribute in node.attribute + } + axis = attributes.get("axis", -1) + block_size = attributes["block_size"] + scale_shape = None + if input_shape is not None: + scale_shape = onnx.TensorShapeProto() + scale_shape.CopyFrom(input_shape) + if axis < 0: + axis += len(scale_shape.dim) + if 0 <= axis < len(scale_shape.dim): + axis_dimension = scale_shape.dim[axis] + if axis_dimension.HasField("dim_value"): + axis_dimension.dim_value = ( + axis_dimension.dim_value + block_size - 1 + ) // block_size + else: + axis_dimension.Clear() + + for output_name, output_dtype, output_shape in ( + (node.output[0], onnx_dtype_map["Float4"], input_shape), + (node.output[1], onnx_dtype_map["Float8"], scale_shape), + ): + output_value_info = value_info_map.get(output_name) + if output_value_info is None: + output_value_info = graph.value_info.add() + value_info_map[output_name] = output_value_info + output_value_info.name = output_name + output_value_info.type.tensor_type.elem_type = output_dtype + if output_shape is not None: + output_value_info.type.tensor_type.shape.CopyFrom(output_shape) + def _get_precision_dtype(weight_initializer: onnx.TensorProto) -> str: return ( "BFloat16" if weight_initializer.data_type == onnx.TensorProto.BFLOAT16 else "Half" @@ -405,6 +450,9 @@ def _cast_input_dtypes(node: onnx.NodeProto, precision_dtype: str): ) fp4_qdq_nodes = [node for node in graph.node if node.op_type == "TRT_FP4QDQ"] + for node in graph.node: + if node.op_type == "TRT_FP4DynamicQuantize": + _annotate_dynamic_quantize_outputs(node) logger.debug(f"Found {len(fp4_qdq_nodes)} FP4QDQ nodes to convert") for node in fp4_qdq_nodes: diff --git a/modelopt/torch/_deploy/utils/torch_onnx.py b/modelopt/torch/_deploy/utils/torch_onnx.py index 7b797d2299e..8c38600bd81 100644 --- a/modelopt/torch/_deploy/utils/torch_onnx.py +++ b/modelopt/torch/_deploy/utils/torch_onnx.py @@ -510,10 +510,11 @@ def get_onnx_bytes_and_metadata( `torch.onnx.export `_. onnx_opset: The onnx opset version to use for exporting the model. dq_only: If True, the exported onnx model is converted to a dq_only model. - weights_dtype: Requested high-precision dtype for exported weights. For an FP8 or NVFP4 - model, ``"bf16"`` is accepted only when every floating parameter is already BF16. - This is a weight-focused no-op, not a graph-wide conversion: floating buffers are - not considered for eligibility and may preserve higher-precision regions. + weights_dtype: Requested high-precision dtype for exported weights. NVFP4 models support + conversion to FP16 or BF16, including mixed FP32/BF16 source parameters. + For an FP8-only model, ``"bf16"`` is accepted only when every floating + parameter is already BF16; this is a weight-focused no-op, not a graph-wide + conversion. Returns: bytes: Onnx model in bytes. @@ -542,11 +543,11 @@ def get_onnx_bytes_and_metadata( uses_fp8 = is_fp8_quantized(model) uses_int8 = is_int8_quantized(model) uses_other_unsupported_quantizer = is_int4_quantized(model) or uses_mxfp8 or uses_int8 + supports_nvfp4_conversion = uses_fp4 and not uses_other_unsupported_quantizer is_bf16_fp4_noop = ( weights_dtype == "bf16" and source_parameter_dtypes == {torch.bfloat16} - and uses_fp4 - and not (uses_fp8 or uses_other_unsupported_quantizer) + and supports_nvfp4_conversion ) is_bf16_fp8_noop = ( weights_dtype == "bf16" @@ -602,18 +603,19 @@ def get_onnx_bytes_and_metadata( if ( weights_dtype == "fp16" - and (uses_fp4 or uses_fp8) + and uses_fp8 + and not supports_nvfp4_conversion and torch.bfloat16 in source_parameter_dtypes ): - quantization_format = "NVFP4" if uses_fp4 else "FP8" raise ValueError( - f"Converting a BF16 {quantization_format} ONNX graph to FP16 is not supported yet " + "Converting a BF16 FP8 ONNX graph to FP16 is not supported yet " f"(source parameter dtypes: {source_parameter_dtype_names})" ) if ( weights_dtype == "bf16" - and (uses_fp4 or uses_fp8 or uses_other_unsupported_quantizer) + and (uses_fp8 or uses_other_unsupported_quantizer) + and not supports_nvfp4_conversion and not is_bf16_quantized_noop ): raise ValueError( @@ -683,7 +685,14 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph = qdq_to_dq(onnx_opt_graph) if weights_dtype in ["fp16", "bf16"] and not is_bf16_quantized_noop: - if weights_dtype == "fp16" and (uses_fp4 or uses_other_unsupported_quantizer or uses_fp8): + convert_nvfp4_with_autocast = supports_nvfp4_conversion and ( + weights_dtype == "bf16" or torch.bfloat16 in source_parameter_dtypes + ) + if ( + weights_dtype == "fp16" + and (uses_fp4 or uses_other_unsupported_quantizer or uses_fp8) + and not convert_nvfp4_with_autocast + ): onnx_opt_graph = convert_float_to_float16( onnx_opt_graph, keep_io_types=False, @@ -701,13 +710,15 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph = fold_qdq_scale_fp16_to_fp32_casts(onnx_opt_graph) else: onnx_opt_graph = convert_to_f16( - onnx_opt_graph, low_precision_type=weights_dtype, keep_io_types=False + onnx_opt_graph, + low_precision_type=weights_dtype, + keep_io_types=False, ) onnx_opt_graph = remove_redundant_casts(onnx_opt_graph) # Remove Cast nodes around Q/DQ for optimal TRT fusion - if uses_fp8: + if uses_fp8 and weights_dtype == "fp16": onnx_opt_graph = fold_q_fp16_to_fp32_casts(onnx_opt_graph) onnx_opt_graph = fold_dq_fp32_to_fp16_casts(onnx_opt_graph) diff --git a/tests/unit/onnx/autocast/test_graphsanitizer.py b/tests/unit/onnx/autocast/test_graphsanitizer.py index f42324a6338..7262e78a2e3 100644 --- a/tests/unit/onnx/autocast/test_graphsanitizer.py +++ b/tests/unit/onnx/autocast/test_graphsanitizer.py @@ -451,14 +451,60 @@ def test_sanitize_large_external_initializer_metadata(): assert [node.op_type for node in sanitizer.model.graph.node] == ["Identity"] -def test_find_custom_nodes_uses_source_model_path(tmp_path, monkeypatch): +@pytest.mark.parametrize( + ("op_type", "output_dtype", "output_shape", "requires_trt_inference"), + [ + pytest.param( + "TRT_FP4DynamicQuantize", + TensorProto.FLOAT4E2M1, + [1], + False, + id="dynamic-quantize-complete-metadata", + ), + pytest.param( + "TRT_FP4DynamicQuantize", + TensorProto.UNDEFINED, + [1], + True, + id="dynamic-quantize-missing-type", + ), + pytest.param( + "TRT_FP4DynamicQuantize", + TensorProto.FLOAT, + [1], + True, + id="dynamic-quantize-wrong-type", + ), + pytest.param( + "TRT_FP4DynamicQuantize", + TensorProto.FLOAT4E2M1, + None, + False, + id="dynamic-quantize-missing-shape", + ), + pytest.param( + "CustomOp", + TensorProto.FLOAT4E2M1, + [1], + True, + id="other-custom-op-complete-metadata", + ), + ], +) +def test_find_custom_nodes_uses_source_model_path( + tmp_path, monkeypatch, op_type, output_dtype, output_shape, requires_trt_inference +): x = helper.make_tensor_value_info("X", TensorProto.FLOAT, [1]) - y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1]) - custom_node = helper.make_node("CustomOp", [x.name], [y.name], name="custom") - graph = helper.make_graph([custom_node], "custom_graph", [x], [y]) + y = helper.make_tensor_value_info("Y", output_dtype, output_shape) + scale = helper.make_tensor_value_info("scale", TensorProto.FLOAT8E4M3FN, [1]) + custom_node = helper.make_node(op_type, [x.name], [y.name, scale.name], name="custom") + graph = helper.make_graph([custom_node], "custom_graph", [x], [y, scale]) model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 22)]) model_path = tmp_path / "custom.onnx" - tensor_info = {"Y": {"dtype": TensorProto.FLOAT, "shape": [1]}} + tensor_info = { + "Y": {"dtype": TensorProto.FLOAT4E2M1, "shape": [1]}, + "scale": {"dtype": TensorProto.FLOAT8E4M3FN, "shape": [1]}, + } get_custom_layers = Mock(return_value=([custom_node.name], tensor_info)) infer_types_shapes = Mock(return_value=model) @@ -469,5 +515,9 @@ def test_find_custom_nodes_uses_source_model_path(tmp_path, monkeypatch): sanitizer = GraphSanitizer(model, min_opset=22, onnx_path=str(model_path)) sanitizer.find_custom_nodes() - get_custom_layers.assert_called_once_with(str(model_path.resolve()), []) - infer_types_shapes.assert_called_once_with(model, [], all_tensor_info=tensor_info) + if requires_trt_inference: + get_custom_layers.assert_called_once_with(str(model_path.resolve()), []) + infer_types_shapes.assert_called_once_with(model, [], all_tensor_info=tensor_info) + else: + get_custom_layers.assert_not_called() + infer_types_shapes.assert_not_called() diff --git a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py index af61bcce553..06b431740bd 100644 --- a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py +++ b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py @@ -300,37 +300,15 @@ def test_fp8_export_rejects_unsupported_dtype_conversion( assert not any(tmp_path.iterdir()) -@pytest.mark.parametrize( - ("source_dtype", "weights_dtype"), - [(torch.bfloat16, "fp16"), (torch.float32, "bf16")], - ids=["bf16-to-fp16", "fp32-to-bf16"], -) -def test_nvfp4_export_rejects_unsupported_dtype_conversion( - source_dtype, weights_dtype, monkeypatch, tmp_path -): - monkeypatch.setattr(tempfile, "tempdir", str(tmp_path)) - monkeypatch.setattr(torch_onnx, "is_fp4_quantized", lambda _: True) - model = nn.Linear(4, 4).eval().to(source_dtype) - - with pytest.raises( - ValueError, - match=rf"Converting .* to {weights_dtype.upper()}.*source parameter dtypes: {source_dtype}", - ): - get_onnx_bytes_and_metadata( - model, - (torch.ones(1, 4, dtype=source_dtype),), - weights_dtype=weights_dtype, - ) - assert not any(tmp_path.iterdir()) - - @pytest.mark.parametrize( ("source_dtype", "weights_dtype", "expected_calls"), [ (torch.float32, "fp16", ["onnxconverter"]), + (torch.float32, "bf16", ["autocast"]), + (torch.bfloat16, "fp16", ["autocast"]), (torch.bfloat16, "bf16", []), ], - ids=["fp32-to-fp16", "bf16-noop"], + ids=["fp32-to-fp16", "fp32-to-bf16", "bf16-to-fp16", "bf16-noop"], ) def test_nvfp4_export_selects_precision_converter( source_dtype, weights_dtype, expected_calls, monkeypatch diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index 612a24e957f..4cf28f69b2c 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -17,6 +17,7 @@ import inspect import io +from copy import deepcopy import numpy as np import pytest @@ -108,10 +109,12 @@ class _NVFP4LinearWithExplicitBias(torch.nn.Module): def __init__(self, dtype): super().__init__() self.linear = torch.nn.Linear(16, 16, bias=False, dtype=dtype) + self.norm = torch.nn.LayerNorm(16, dtype=dtype) + self.projection = torch.nn.Linear(16, 16, bias=False, dtype=dtype) self.bias = torch.nn.Parameter(torch.ones(16, dtype=dtype)) def forward(self, inputs): - return self.linear(inputs) + self.bias + return self.projection(self.norm(self.linear(inputs))) + self.bias class _NVFP4MixedPrecisionLinear(torch.nn.Module): @@ -126,7 +129,9 @@ def forward(self, inputs): return fp32_output + bf16_output -def _make_cpu_nvfp4_model(monkeypatch, model, sample_input, disable_input_quantizers=False): +def _make_cpu_nvfp4_model( + monkeypatch, model, sample_input, disable_input_quantizers=False, quant_config=None +): def forward_loop(model): model(sample_input) @@ -134,7 +139,7 @@ def cpu_dynamic_block_quantize(inputs, *args): return inputs monkeypatch.setattr(tensor_quant, "dynamic_block_quantize_op", cpu_dynamic_block_quantize) - model = mtq.quantize(model, mtq.NVFP4_DEFAULT_CFG, forward_loop=forward_loop) + model = mtq.quantize(model, quant_config or mtq.NVFP4_DEFAULT_CFG, forward_loop=forward_loop) for module in model.modules(): assert not isinstance(module, torch.nn.Linear) or is_quantized_linear(module) @@ -171,9 +176,11 @@ def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch): [ (torch.float32, "fp32", TensorProto.FLOAT), (torch.float32, "fp16", TensorProto.FLOAT16), + (torch.float32, "bf16", TensorProto.BFLOAT16), + (torch.bfloat16, "fp16", TensorProto.FLOAT16), (torch.bfloat16, "bf16", TensorProto.BFLOAT16), ], - ids=["fp32-preserved", "fp32-to-fp16", "bf16-noop"], + ids=["fp32-preserved", "fp32-to-fp16", "fp32-to-bf16", "bf16-to-fp16", "bf16-noop"], ) def test_nvfp4_deploy_export_has_consistent_elementwise_types( monkeypatch, source_dtype, weights_dtype, expected_dtype @@ -181,6 +188,8 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( model = _NVFP4LinearWithExplicitBias(source_dtype).eval() sample_input = torch.ones(1, 2, 16, dtype=source_dtype) model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) + # Swin leaves LayerNorm input quantizers disabled on its high-rank activation paths. + model.norm.input_quantizer.disable() onnx_bytes, _ = get_onnx_bytes_and_metadata( model, @@ -193,6 +202,8 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( assert utils.get_opset_version(exported_model) >= 23 onnx.checker.check_model(exported_model, full_check=True) + assert exported_model.graph.input[0].type.tensor_type.elem_type == expected_dtype + assert exported_model.graph.output[0].type.tensor_type.elem_type == expected_dtype assert any( node.op_type == "DequantizeLinear" and any(attribute.name == "block_size" for attribute in node.attribute) @@ -212,6 +223,14 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( if value.type.HasField("tensor_type"): tensor_types[value.name] = value.type.tensor_type.elem_type + dynamic_quantize = next( + node for node in inferred_model.graph.node if node.op_type == "TRT_FP4DynamicQuantize" + ) + assert [tensor_types[output] for output in dynamic_quantize.output] == [ + TensorProto.FLOAT4E2M1, + TensorProto.FLOAT8E4M3FN, + ] + floating_types = {TensorProto.FLOAT, TensorProto.FLOAT16, TensorProto.BFLOAT16} elementwise_nodes = [ node @@ -233,21 +252,39 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( ] -def test_nvfp4_deploy_export_preserves_mixed_precision_boundaries(monkeypatch): +@pytest.mark.parametrize( + ("weights_dtype", "expected_dtype"), + [ + ("fp32", TensorProto.FLOAT), + ("fp16", TensorProto.FLOAT16), + ("bf16", TensorProto.BFLOAT16), + ], +) +@pytest.mark.parametrize("with_fp8_branch", [False, True]) +def test_nvfp4_deploy_export_preserves_mixed_precision_boundaries( + monkeypatch, weights_dtype, expected_dtype, with_fp8_branch +): model = _NVFP4MixedPrecisionLinear().eval() sample_input = torch.ones(1, 16) - model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) + quant_config = deepcopy(mtq.NVFP4_DEFAULT_CFG) + if with_fp8_branch: + quant_config["quant_cfg"].append( + {"quantizer_name": "bf16_linear*quantizer", "cfg": {"num_bits": (4, 3), "axis": None}} + ) + model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input, quant_config=quant_config) onnx_bytes, _ = get_onnx_bytes_and_metadata( model, (sample_input,), - weights_dtype="fp32", + weights_dtype=weights_dtype, ) exported_model = onnx.load_model_from_string( OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() ) onnx.checker.check_model(exported_model, full_check=True) + assert exported_model.graph.input[0].type.tensor_type.elem_type == expected_dtype + assert exported_model.graph.output[0].type.tensor_type.elem_type == expected_dtype inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) tensor_types = { initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer @@ -262,9 +299,16 @@ def test_nvfp4_deploy_export_preserves_mixed_precision_boundaries(monkeypatch): add_node = next(node for node in inferred_model.graph.node if node.op_type == "Add") assert [tensor_types[input_name] for input_name in add_node.input] == [ - TensorProto.FLOAT, - TensorProto.FLOAT, + expected_dtype, + expected_dtype, ] + if with_fp8_branch and weights_dtype != "fp32": + quantize_nodes = [ + node for node in inferred_model.graph.node if node.op_type == "QuantizeLinear" + ] + assert quantize_nodes + for node in quantize_nodes: + assert tensor_types[node.input[1]] == expected_dtype @pytest.mark.parametrize( @@ -272,9 +316,11 @@ def test_nvfp4_deploy_export_preserves_mixed_precision_boundaries(monkeypatch): [ (torch.float32, "fp32", TensorProto.FLOAT16, TensorProto.FLOAT), (torch.float32, "fp16", TensorProto.FLOAT16, TensorProto.FLOAT16), + (torch.float32, "bf16", TensorProto.BFLOAT16, TensorProto.BFLOAT16), + (torch.bfloat16, "fp16", TensorProto.FLOAT16, TensorProto.FLOAT16), (torch.bfloat16, "bf16", TensorProto.BFLOAT16, TensorProto.BFLOAT16), ], - ids=["fp32-preserved", "fp32-to-fp16", "bf16-noop"], + ids=["fp32-preserved", "fp32-to-fp16", "fp32-to-bf16", "bf16-to-fp16", "bf16-noop"], ) def test_nvfp4_deploy_export_has_consistent_gemm_types( monkeypatch, source_dtype, weights_dtype, expected_gemm_dtype, expected_output_dtype From 8ebd8710594f6b5752a59960d07681fc69782906 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Tue, 22 Sep 2026 00:09:35 +0000 Subject: [PATCH 6/7] refactor: unify NVFP4 ONNX precision conversion Use AutoCast policies for supported NVFP4 FP16 and BF16 exports, preserve native BF16 no-op and unrelated quantizer routes, and preserve the sign of tiny initializer values. Consolidate regression setup and verify packed payloads, dynamic converted shapes, and numerical limits. Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- CHANGELOG.rst | 3 + modelopt/onnx/autocast/precisionconverter.py | 2 +- modelopt/torch/_deploy/utils/torch_onnx.py | 19 +-- .../onnx/autocast/test_precisionconverter.py | 65 +++++++++ .../deploy/utils/test_torch_onnx_utils.py | 40 ++++-- .../quantization/test_onnx_export_cpu.py | 134 +++++++++--------- 6 files changed, 175 insertions(+), 88 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index f471cd6f530..4debcbe5864 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -23,6 +23,8 @@ Changelog **Backward Breaking Changes** +- Supported NVFP4 ONNX FP16 conversions now use AutoCast's precision policy, including FP16 ``Div`` where supported and initializer clamping to the FP16 range (including infinities), instead of the legacy converter's narrower range. Exported numerical results can change; revalidate model accuracy when upgrading. + - The ``modelopt.onnx.quantization.graph_utils`` module has been removed with no compatibility shim; update direct imports using this migration map: @@ -67,6 +69,7 @@ Changelog - Fix unified Megatron export writing a second, unreferenced copy of the vocab embedding when a model with MTP layers is exported with pipeline parallelism. The duplicate was never loaded but inflated the checkpoint by the size of the embedding (about 1 GB for Qwen3.6-35B-A3B); re-export to reclaim the space. - Fail fast on non-finite AutoQuantize output gradients with an actionable error before accumulating sensitivity scores, without changing attention backend settings. - Fix NVFP4 ONNX exports that failed ONNX or TensorRT parsing due to mixed-precision ``MatMul``, ``Gemm``, and elementwise inputs. Export preserves FP32 graph boundaries around low-precision NVFP4 compute and supports conversion of FP32, BF16, or mixed FP32/BF16 models to FP16 or BF16. +- Fix ONNX AutoCast changing tiny negative initializer values to positive values during precision conversion; underflow clamping now preserves their sign. - Fix ONNX INT8 entropy calibration failing or producing invalid quantization parameters for FP16 activations. - Fix ``--use_fsdp2`` HuggingFace checkpoint export gathering the whole model onto rank 0, which made export the dominant phase of a PTQ run and could exhaust host memory on large models. The model is now split into per-decoder-layer units dealt round-robin across ranks; each rank gathers every unit but keeps, packs, and writes only the ones it owns, so a rank buffers roughly ``model / world_size`` instead of the whole checkpoint, and rank 0 writes the combined index. Export configurations that cannot be split this way now raise instead of producing a mismatched checkpoint: FSDP2 combined with another DTensor parallelism (for example FSDP2 + tensor parallel on a 2-D mesh; HSDP is supported), models whose decoder layers cannot be discovered, a decoder layer object reused across layers, and a module that holds the decoder layers while owning parameters of its own. - Speed up ``mtq.quantize`` on FSDP2-sharded fused-MoE models. Promoting static-block weight quantizers gathered each expert's slice of the fused weight across ranks even though only quantizer state is read, adding a collective per expert to calibration. diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index 158890595eb..939380cd11a 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -1248,7 +1248,7 @@ def _convert_initializer_data( self._warned_values_clamp_min = True np_array = np.where( (np_array != 0.0) & (np.abs(np_array) < data_lowest), - data_lowest, + np.copysign(data_lowest, np_array), np_array, ) new_array = np_array.astype(to_type.numpy_type) diff --git a/modelopt/torch/_deploy/utils/torch_onnx.py b/modelopt/torch/_deploy/utils/torch_onnx.py index 8c38600bd81..8ff8d9ca3d8 100644 --- a/modelopt/torch/_deploy/utils/torch_onnx.py +++ b/modelopt/torch/_deploy/utils/torch_onnx.py @@ -544,18 +544,12 @@ def get_onnx_bytes_and_metadata( uses_int8 = is_int8_quantized(model) uses_other_unsupported_quantizer = is_int4_quantized(model) or uses_mxfp8 or uses_int8 supports_nvfp4_conversion = uses_fp4 and not uses_other_unsupported_quantizer - is_bf16_fp4_noop = ( + is_bf16_quantized_noop = ( weights_dtype == "bf16" and source_parameter_dtypes == {torch.bfloat16} - and supports_nvfp4_conversion + and (uses_fp4 or uses_fp8) + and not uses_other_unsupported_quantizer ) - is_bf16_fp8_noop = ( - weights_dtype == "bf16" - and source_parameter_dtypes == {torch.bfloat16} - and uses_fp8 - and not (uses_fp4 or uses_other_unsupported_quantizer) - ) - is_bf16_quantized_noop = is_bf16_fp4_noop or is_bf16_fp8_noop # Standardize model args and also tensorize them so they also appear in the onnx graph! # Floats/ints are tensorized when they are provided, but not tensorized when they are not @@ -685,13 +679,10 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph = qdq_to_dq(onnx_opt_graph) if weights_dtype in ["fp16", "bf16"] and not is_bf16_quantized_noop: - convert_nvfp4_with_autocast = supports_nvfp4_conversion and ( - weights_dtype == "bf16" or torch.bfloat16 in source_parameter_dtypes - ) if ( weights_dtype == "fp16" - and (uses_fp4 or uses_other_unsupported_quantizer or uses_fp8) - and not convert_nvfp4_with_autocast + and (uses_other_unsupported_quantizer or uses_fp8) + and not supports_nvfp4_conversion ): onnx_opt_graph = convert_float_to_float16( onnx_opt_graph, diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index d3e0a786a65..a873615c7bf 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -624,6 +624,71 @@ def test_clamping_fp16_initializers_out_of_range( assert np.all(add_init_fp32_array == add_init_out_of_range) +@pytest.mark.parametrize("use_standalone_type_inference", [True, False]) +def test_convert_to_f16_clamps_initializers_preserving_sign(use_standalone_type_inference): + values = np.array( + [ + 20000, + -20000, + 70000, + -70000, + 1e-10, + -1e-10, + 1e-7, + -1e-7, + 0.0, + -0.0, + np.inf, + -np.inf, + np.nan, + ], + dtype=np.float32, + ) + graph = helper.make_graph( + [helper.make_node("Add", ["X", "weight"], ["Y"], name="add")], + "initializer_limits", + [helper.make_tensor_value_info("X", TensorProto.FLOAT, [len(values)])], + [helper.make_tensor_value_info("Y", TensorProto.FLOAT, [len(values)])], + [numpy_helper.from_array(values, name="weight")], + ) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 19)], ir_version=10) + + converted = convert_to_f16( + model, + keep_io_types=False, + use_standalone_type_inference=use_standalone_type_inference, + ) + + weight = next(init for init in converted.graph.initializer if init.name == "weight") + assert weight.data_type == TensorProto.FLOAT16 + actual = numpy_helper.to_array(weight) + limits = np.finfo(np.float16) + expected = np.array( + [ + 20000, + -20000, + limits.max, + -limits.max, + limits.smallest_subnormal, + -limits.smallest_subnormal, + 1e-7, + -1e-7, + 0.0, + -0.0, + limits.max, + -limits.max, + np.nan, + ], + dtype=np.float16, + ) + np.testing.assert_array_equal(actual, expected) + # Numeric equality alone does not distinguish positive and negative zero. + np.testing.assert_array_equal(np.signbit(actual[:-1]), np.signbit(values[:-1])) + for value in (*converted.graph.input, *converted.graph.output): + assert value.type.tensor_type.elem_type == TensorProto.FLOAT16 + onnx.checker.check_model(converted, full_check=True) + + @pytest.mark.parametrize("use_standalone_type_inference", [True, False]) def test_bf16_no_clamping_initializers_out_of_range( model_with_multiple_consumers, use_standalone_type_inference diff --git a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py index 06b431740bd..78a70d88635 100644 --- a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py +++ b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py @@ -301,17 +301,34 @@ def test_fp8_export_rejects_unsupported_dtype_conversion( @pytest.mark.parametrize( - ("source_dtype", "weights_dtype", "expected_calls"), + ("source_dtype", "weights_dtype", "quantizers", "expected_calls"), [ - (torch.float32, "fp16", ["onnxconverter"]), - (torch.float32, "bf16", ["autocast"]), - (torch.bfloat16, "fp16", ["autocast"]), - (torch.bfloat16, "bf16", []), + (torch.float32, "fp16", ("fp4",), ["autocast"]), + (torch.float32, "bf16", ("fp4",), ["autocast"]), + (torch.bfloat16, "fp16", ("fp4",), ["autocast"]), + (torch.bfloat16, "bf16", ("fp4",), []), + (torch.float32, "fp16", ("fp4", "fp8"), ["autocast"]), + (torch.bfloat16, "bf16", ("fp4", "fp8"), []), + (torch.float32, "fp16", ("fp8",), ["onnxconverter"]), + (torch.float32, "fp16", ("fp4", "int4"), ["onnxconverter"]), + (torch.float32, "fp16", ("fp4", "int8"), ["onnxconverter"]), + (torch.float32, "fp16", ("fp4", "mxfp8"), ["onnxconverter"]), + ], + ids=[ + "fp32-to-fp16", + "fp32-to-bf16", + "bf16-to-fp16", + "bf16-noop", + "fp4-fp8-to-fp16", + "fp4-fp8-bf16-noop", + "fp8-legacy", + "fp4-int4-legacy", + "fp4-int8-legacy", + "fp4-mxfp8-legacy", ], - ids=["fp32-to-fp16", "fp32-to-bf16", "bf16-to-fp16", "bf16-noop"], ) -def test_nvfp4_export_selects_precision_converter( - source_dtype, weights_dtype, expected_calls, monkeypatch +def test_quantized_export_selects_precision_converter( + source_dtype, weights_dtype, quantizers, expected_calls, monkeypatch ): calls = [] @@ -323,7 +340,12 @@ def record_autocast(model, **kwargs): calls.append("autocast") return model - monkeypatch.setattr(torch_onnx, "is_fp4_quantized", lambda _: True) + for quantizer in ("fp4", "fp8", "int4", "int8", "mxfp8"): + monkeypatch.setattr( + torch_onnx, + f"is_{quantizer}_quantized", + lambda _, enabled=quantizer in quantizers: enabled, + ) monkeypatch.setattr( torch_onnx, "configure_linear_module_onnx_quantizers", lambda _: nullcontext() ) diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index 4cf28f69b2c..8d973c59bc5 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -30,6 +30,7 @@ from _test_utils.torch.quantization.onnx_export import TEST_MODELS, onnx_export_tester from onnx import TensorProto, helper, numpy_helper +import modelopt.torch._deploy.utils.torch_onnx as torch_onnx import modelopt.torch.quantization as mtq import modelopt.torch.quantization.tensor_quant as tensor_quant from modelopt.onnx import utils @@ -151,6 +152,31 @@ def cpu_dynamic_block_quantize(inputs, *args): return model +def _export_deploy_onnx_with_types(model, sample_input, weights_dtype, dynamic_axes=None): + onnx_bytes, _ = get_onnx_bytes_and_metadata( + model, + (sample_input,), + weights_dtype=weights_dtype, + dynamic_axes=dynamic_axes or {}, + ) + exported_model = onnx.load_model_from_string( + OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + ) + onnx.checker.check_model(exported_model, full_check=True) + inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) + tensor_types = { + initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer + } + for value in [ + *inferred_model.graph.input, + *inferred_model.graph.value_info, + *inferred_model.graph.output, + ]: + if value.type.HasField("tensor_type"): + tensor_types[value.name] = value.type.tensor_type.elem_type + return exported_model, tensor_types + + def test_nvfp4_exported_onnx_is_topologically_sorted(monkeypatch): model = SimpleLinear().eval() sample_input = model.get_input() @@ -191,17 +217,11 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( # Swin leaves LayerNorm input quantizers disabled on its high-rank activation paths. model.norm.input_quantizer.disable() - onnx_bytes, _ = get_onnx_bytes_and_metadata( - model, - (sample_input,), - weights_dtype=weights_dtype, - ) - exported_model = onnx.load_model_from_string( - OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + exported_model, tensor_types = _export_deploy_onnx_with_types( + model, sample_input, weights_dtype ) assert utils.get_opset_version(exported_model) >= 23 - onnx.checker.check_model(exported_model, full_check=True) assert exported_model.graph.input[0].type.tensor_type.elem_type == expected_dtype assert exported_model.graph.output[0].type.tensor_type.elem_type == expected_dtype assert any( @@ -211,20 +231,8 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( ) assert any(node.op_type == "TRT_FP4DynamicQuantize" for node in exported_model.graph.node) - inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) - tensor_types = { - initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer - } - for value in [ - *inferred_model.graph.input, - *inferred_model.graph.value_info, - *inferred_model.graph.output, - ]: - if value.type.HasField("tensor_type"): - tensor_types[value.name] = value.type.tensor_type.elem_type - dynamic_quantize = next( - node for node in inferred_model.graph.node if node.op_type == "TRT_FP4DynamicQuantize" + node for node in exported_model.graph.node if node.op_type == "TRT_FP4DynamicQuantize" ) assert [tensor_types[output] for output in dynamic_quantize.output] == [ TensorProto.FLOAT4E2M1, @@ -234,7 +242,7 @@ def test_nvfp4_deploy_export_has_consistent_elementwise_types( floating_types = {TensorProto.FLOAT, TensorProto.FLOAT16, TensorProto.BFLOAT16} elementwise_nodes = [ node - for node in inferred_model.graph.node + for node in exported_model.graph.node if node.op_type in {"Add", "Sub", "Mul", "Div", "Pow"} ] assert any(node.op_type == "Add" for node in elementwise_nodes) @@ -273,38 +281,53 @@ def test_nvfp4_deploy_export_preserves_mixed_precision_boundaries( ) model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input, quant_config=quant_config) - onnx_bytes, _ = get_onnx_bytes_and_metadata( + quantized_types = {TensorProto.FLOAT4E2M1, TensorProto.FLOAT8E4M3FN} + expected_payloads = {} + quantize_weights = torch_onnx.quantize_weights + + def quantized_payloads(graph): + return { + initializer.name: ( + initializer.data_type, + tuple(initializer.dims), + initializer.raw_data, + ) + for initializer in graph.graph.initializer + if initializer.data_type in quantized_types + } + + def capture_quantized_payloads(model, graph): + graph = quantize_weights(model, graph) + expected_payloads.update(quantized_payloads(graph)) + return graph + + monkeypatch.setattr(torch_onnx, "quantize_weights", capture_quantized_payloads) + exported_model, tensor_types = _export_deploy_onnx_with_types( model, - (sample_input,), - weights_dtype=weights_dtype, - ) - exported_model = onnx.load_model_from_string( - OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + sample_input, + weights_dtype, + dynamic_axes=( + {"inputs": {0: "batch"}, "out": {0: "batch"}} if weights_dtype != "fp32" else None + ), ) - - onnx.checker.check_model(exported_model, full_check=True) + assert {payload[0] for payload in expected_payloads.values()} == quantized_types + assert all(payload[2] for payload in expected_payloads.values()) + assert quantized_payloads(exported_model) == expected_payloads assert exported_model.graph.input[0].type.tensor_type.elem_type == expected_dtype assert exported_model.graph.output[0].type.tensor_type.elem_type == expected_dtype - inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) - tensor_types = { - initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer - } - for value in [ - *inferred_model.graph.input, - *inferred_model.graph.value_info, - *inferred_model.graph.output, - ]: - if value.type.HasField("tensor_type"): - tensor_types[value.name] = value.type.tensor_type.elem_type + if weights_dtype != "fp32": + for value in [*exported_model.graph.input, *exported_model.graph.output]: + assert value.type.tensor_type.shape.dim[0].dim_param == "batch" + assert value.type.tensor_type.shape.dim[1].dim_value == 16 - add_node = next(node for node in inferred_model.graph.node if node.op_type == "Add") + add_node = next(node for node in exported_model.graph.node if node.op_type == "Add") assert [tensor_types[input_name] for input_name in add_node.input] == [ expected_dtype, expected_dtype, ] if with_fp8_branch and weights_dtype != "fp32": quantize_nodes = [ - node for node in inferred_model.graph.node if node.op_type == "QuantizeLinear" + node for node in exported_model.graph.node if node.op_type == "QuantizeLinear" ] assert quantize_nodes for node in quantize_nodes: @@ -329,34 +352,17 @@ def test_nvfp4_deploy_export_has_consistent_gemm_types( sample_input = torch.ones(2, 16, dtype=source_dtype) model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) - onnx_bytes, _ = get_onnx_bytes_and_metadata( - model, - (sample_input,), - weights_dtype=weights_dtype, - ) - exported_model = onnx.load_model_from_string( - OnnxBytes.from_bytes(onnx_bytes).get_onnx_model_file_bytes() + exported_model, tensor_types = _export_deploy_onnx_with_types( + model, sample_input, weights_dtype ) assert utils.get_opset_version(exported_model) >= 23 - onnx.checker.check_model(exported_model, full_check=True) assert any(node.op_type == "TRT_FP4DynamicQuantize" for node in exported_model.graph.node) - inferred_model = onnx.shape_inference.infer_shapes(exported_model, strict_mode=True) - tensor_types = { - initializer.name: initializer.data_type for initializer in inferred_model.graph.initializer - } - for value in [ - *inferred_model.graph.input, - *inferred_model.graph.value_info, - *inferred_model.graph.output, - ]: - if value.type.HasField("tensor_type"): - tensor_types[value.name] = value.type.tensor_type.elem_type - gemm_node = next(node for node in inferred_model.graph.node if node.op_type == "Gemm") + gemm_node = next(node for node in exported_model.graph.node if node.op_type == "Gemm") assert [tensor_types[input_name] for input_name in gemm_node.input] == [expected_gemm_dtype] * 3 assert tensor_types[gemm_node.output[0]] == expected_gemm_dtype - assert tensor_types[inferred_model.graph.output[0].name] == expected_output_dtype + assert tensor_types[exported_model.graph.output[0].name] == expected_output_dtype @pytest.mark.parametrize( From 03b814c1a9a85bcaa60c4135ee1e48c30c853202 Mon Sep 17 00:00:00 2001 From: ajrasane <131806219+ajrasane@users.noreply.github.com> Date: Wed, 23 Sep 2026 19:54:04 +0000 Subject: [PATCH 7/7] Scope NVFP4 inference deferral and boundary coverage Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com> --- modelopt/onnx/autocast/convert.py | 7 +++- modelopt/onnx/autocast/graphsanitizer.py | 19 ++++++--- modelopt/onnx/autocast/precisionconverter.py | 6 ++- modelopt/torch/_deploy/utils/torch_onnx.py | 1 + .../unit/onnx/autocast/test_graphsanitizer.py | 34 ++++++++++++++-- .../onnx/autocast/test_precisionconverter.py | 27 +++++++------ .../deploy/utils/test_torch_onnx_utils.py | 3 ++ .../quantization/test_onnx_export_cpu.py | 39 +++++++++++++++++-- 8 files changed, 110 insertions(+), 26 deletions(-) diff --git a/modelopt/onnx/autocast/convert.py b/modelopt/onnx/autocast/convert.py index aef5f03cce8..ff96f1e4dd9 100644 --- a/modelopt/onnx/autocast/convert.py +++ b/modelopt/onnx/autocast/convert.py @@ -232,6 +232,8 @@ def convert_to_f16( use_standalone_type_inference: bool = False, opset: int | None = None, nodes_to_exclude: list[str] | None = None, + *, + defer_nvfp4_trt_inference: bool = False, ) -> onnx.ModelProto: """Convert model to mixed precision, using PrecisionConverter. @@ -252,6 +254,8 @@ def convert_to_f16( increased if Q/DQ nodes in the model require a higher version (e.g., FP8 requires 19, INT4 requires 21, NVFP4 requires 23). nodes_to_exclude: List of regex patterns to match node names that should remain in FP32. + defer_nvfp4_trt_inference: Defer TensorRT inference for annotated NVFP4 plugins until + compute dtypes are converted. The caller must validate the converted graph in TensorRT. """ assert low_precision_type in ["fp16", "bf16"], "low_precision_type must be either fp16 or bf16" original_network_io_metadata = _capture_network_io_metadata(model, keep_io_types) @@ -291,7 +295,7 @@ def convert_to_f16( trt_plugins=trt_plugins, max_ir_version=LATEST_IR_VERSION_SUPPORTED_BY_ORT, ) - sanitizer.find_custom_nodes() + sanitizer.find_custom_nodes(defer_nvfp4_trt_inference=defer_nvfp4_trt_inference) sanitizer.convert_opset() sanitizer.ensure_graph_name_exists() sanitizer.convert_fp64_to_fp32() @@ -314,6 +318,7 @@ def convert_to_f16( tensor_block_dict=tensor_block_dict, use_standalone_type_inference=use_standalone_type_inference, original_network_io_metadata=original_network_io_metadata, + defer_nvfp4_trt_inference=defer_nvfp4_trt_inference, ) node_name_rule = DisabledNodeNameRegexRule(nodes_to_exclude or []) high_precision_nodes = [ diff --git a/modelopt/onnx/autocast/graphsanitizer.py b/modelopt/onnx/autocast/graphsanitizer.py index 4a30c2484ea..9c5187b6b5c 100644 --- a/modelopt/onnx/autocast/graphsanitizer.py +++ b/modelopt/onnx/autocast/graphsanitizer.py @@ -67,13 +67,16 @@ def __init__( self.onnx_path = os.path.abspath(onnx_path) if onnx_path is not None else None self.external_data_dir = os.path.dirname(self.onnx_path) if self.onnx_path else "" - def sanitize(self) -> None: + def sanitize(self, *, defer_nvfp4_trt_inference: bool = False) -> None: """Sanitize the model graph. Currently, this finds decomposed LayerNorm patterns and replaces them with a single LayerNormalization operator. Additional functionality may be added in the future. + + Args: + defer_nvfp4_trt_inference: Forward NVFP4 inference deferral to custom-node discovery. """ - self.find_custom_nodes() + self.find_custom_nodes(defer_nvfp4_trt_inference=defer_nvfp4_trt_inference) self.remove_disconnected_outputs() self.convert_opset() self.replace_layernorm_pattern() @@ -119,11 +122,15 @@ def ensure_custom_ops_precision(self) -> None: ] logger.info("Ensured custom ops precision") - def find_custom_nodes(self) -> None: + def find_custom_nodes(self, *, defer_nvfp4_trt_inference: bool = False) -> None: """Find custom nodes in the model. Scans through all nodes in the graph and logs any nodes that use custom operators that are not part of the standard ONNX operator set. + + Args: + defer_nvfp4_trt_inference: Defer inference only for NVFP4 plugins with declared + FP4/FP8 output types; the caller must validate after precision conversion. """ self.custom_ops = { node.op_type for node in self.model.graph.node if node.op_type not in self.standard_ops @@ -149,9 +156,9 @@ def find_custom_nodes(self) -> None: output_types == [onnx.TensorProto.FLOAT4E2M1, onnx.TensorProto.FLOAT8E4M3FN] for output_types in custom_output_types ) - # NVFP4 export declares its FP4 data and FP8 scale outputs before dtype conversion. - # Avoid parsing that graph in TensorRT until its compute dtypes are normalized. - if not has_nvfp4_types: + # NVFP4 export explicitly defers parsing until its compute dtypes are normalized. + # Other callers retain TensorRT inference, even with the same plugin annotations. + if not (defer_nvfp4_trt_inference and has_nvfp4_types): # Infer types and shapes in the graph for ORT compatibility _, all_tensor_info = get_custom_layers( self.onnx_path or self.model, self.trt_plugins diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index 939380cd11a..23ed1c5c7cf 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -100,6 +100,8 @@ def __init__( use_standalone_type_inference: bool = False, original_network_io_metadata: dict[str, list[onnx.ValueInfoProto]] | None = None, sanitize_model: bool = True, + *, + defer_nvfp4_trt_inference: bool = False, ) -> None: """Initialize PrecisionConverter. @@ -120,9 +122,11 @@ def __init__( use_standalone_type_inference: Use standalone type inference instead of ONNX's infer_shapes. original_network_io_metadata: Original public input/output metadata captured at the API boundary. sanitize_model: Whether to sanitize the model before precision conversion. + defer_nvfp4_trt_inference: Defer annotated NVFP4 plugin inference during sanitization. """ self.model = deepcopy(model) self.sanitize_model = sanitize_model + self.defer_nvfp4_trt_inference = defer_nvfp4_trt_inference if sanitize_model: self.value_info_map = value_info_map self.initializer_map = initializer_map @@ -1852,7 +1856,7 @@ def _sanitize_model(self): trt_plugins=self.trt_plugins, max_ir_version=self.max_ir_version, ) - graph_sanitizer.sanitize() + graph_sanitizer.sanitize(defer_nvfp4_trt_inference=self.defer_nvfp4_trt_inference) self.model = graph_sanitizer.model # Update value_info_map and initializer_map after sanitizing model diff --git a/modelopt/torch/_deploy/utils/torch_onnx.py b/modelopt/torch/_deploy/utils/torch_onnx.py index 8ff8d9ca3d8..51da69fea7e 100644 --- a/modelopt/torch/_deploy/utils/torch_onnx.py +++ b/modelopt/torch/_deploy/utils/torch_onnx.py @@ -704,6 +704,7 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph, low_precision_type=weights_dtype, keep_io_types=False, + defer_nvfp4_trt_inference=supports_nvfp4_conversion, ) onnx_opt_graph = remove_redundant_casts(onnx_opt_graph) diff --git a/tests/unit/onnx/autocast/test_graphsanitizer.py b/tests/unit/onnx/autocast/test_graphsanitizer.py index 7262e78a2e3..c273e625068 100644 --- a/tests/unit/onnx/autocast/test_graphsanitizer.py +++ b/tests/unit/onnx/autocast/test_graphsanitizer.py @@ -451,20 +451,23 @@ def test_sanitize_large_external_initializer_metadata(): assert [node.op_type for node in sanitizer.model.graph.node] == ["Identity"] +@pytest.mark.parametrize("defer_nvfp4_trt_inference", [False, True]) @pytest.mark.parametrize( - ("op_type", "output_dtype", "output_shape", "requires_trt_inference"), + ("op_type", "output_dtype", "output_shape", "extra_custom_op", "requires_trt_inference"), [ pytest.param( "TRT_FP4DynamicQuantize", TensorProto.FLOAT4E2M1, [1], False, + False, id="dynamic-quantize-complete-metadata", ), pytest.param( "TRT_FP4DynamicQuantize", TensorProto.UNDEFINED, [1], + False, True, id="dynamic-quantize-missing-type", ), @@ -472,6 +475,7 @@ def test_sanitize_large_external_initializer_metadata(): "TRT_FP4DynamicQuantize", TensorProto.FLOAT, [1], + False, True, id="dynamic-quantize-wrong-type", ), @@ -480,25 +484,44 @@ def test_sanitize_large_external_initializer_metadata(): TensorProto.FLOAT4E2M1, None, False, + False, id="dynamic-quantize-missing-shape", ), pytest.param( "CustomOp", TensorProto.FLOAT4E2M1, [1], + False, True, id="other-custom-op-complete-metadata", ), + pytest.param( + "TRT_FP4DynamicQuantize", + TensorProto.FLOAT4E2M1, + [1], + True, + True, + id="dynamic-quantize-with-another-custom-op", + ), ], ) def test_find_custom_nodes_uses_source_model_path( - tmp_path, monkeypatch, op_type, output_dtype, output_shape, requires_trt_inference + tmp_path, + monkeypatch, + op_type, + output_dtype, + output_shape, + extra_custom_op, + requires_trt_inference, + defer_nvfp4_trt_inference, ): x = helper.make_tensor_value_info("X", TensorProto.FLOAT, [1]) y = helper.make_tensor_value_info("Y", output_dtype, output_shape) scale = helper.make_tensor_value_info("scale", TensorProto.FLOAT8E4M3FN, [1]) custom_node = helper.make_node(op_type, [x.name], [y.name, scale.name], name="custom") graph = helper.make_graph([custom_node], "custom_graph", [x], [y, scale]) + if extra_custom_op: + graph.node.append(helper.make_node("CustomOp", [x.name], ["extra"], name="extra_custom")) model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 22)]) model_path = tmp_path / "custom.onnx" tensor_info = { @@ -513,9 +536,12 @@ def test_find_custom_nodes_uses_source_model_path( monkeypatch.setattr(graphsanitizer, "infer_types_shapes_tensorrt", infer_types_shapes) sanitizer = GraphSanitizer(model, min_opset=22, onnx_path=str(model_path)) - sanitizer.find_custom_nodes() + if defer_nvfp4_trt_inference: + sanitizer.find_custom_nodes(defer_nvfp4_trt_inference=True) + else: + sanitizer.find_custom_nodes() - if requires_trt_inference: + if requires_trt_inference or not defer_nvfp4_trt_inference: get_custom_layers.assert_called_once_with(str(model_path.resolve()), []) infer_types_shapes.assert_called_once_with(model, [], all_tensor_info=tensor_info) else: diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index a873615c7bf..333350a653f 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -733,29 +733,33 @@ def test_bf16_conversion_accepts_fp16_initializer(use_standalone_type_inference) model.ir_version = LATEST_IR_VERSION_SUPPORTED_BY_ORT onnx.checker.check_model(model) - model, value_info_map, initializer_map, node_to_init_map = setup_mappings( - model, use_standalone_type_inference - ) - converter = PrecisionConverter( + converted_model = convert_to_f16( model, - value_info_map, - initializer_map, - node_to_init_map, low_precision_type="bf16", + keep_io_types=False, use_standalone_type_inference=use_standalone_type_inference, ) - - converted_model = converter.convert(high_precision_nodes=[], low_precision_nodes=["add"]) + onnx.checker.check_model(converted_model, full_check=True) + assert [value.name for value in converted_model.graph.input] == ["X"] + assert [value.name for value in converted_model.graph.output] == ["Y"] + for value in [*converted_model.graph.input, *converted_model.graph.output]: + assert value.type.tensor_type.elem_type == TensorProto.BFLOAT16 + assert [dimension.dim_value for dimension in value.type.tensor_type.shape.dim] == [2] converted_weight = next( init for init in converted_model.graph.initializer if init.name == "weight" ) assert converted_weight.data_type == TensorProto.BFLOAT16 - assert converted_model.graph.output[0].type.tensor_type.elem_type == TensorProto.BFLOAT16 + assert list(converted_weight.dims) == [2] np.testing.assert_allclose( onnx_utils.read_f16_tensor_as_fp32(converted_weight), np.array([1.0, 2.0], dtype=np.float32), ) + add_node = next(node for node in converted_model.graph.node if node.op_type == "Add") + cast_nodes = [node for node in converted_model.graph.node if node.op_type == "Cast"] + assert not cast_nodes + assert list(add_node.input) == ["X", "weight"] + assert list(add_node.output) == ["Y"] #################################################################################################### @@ -2332,7 +2336,8 @@ def test_convert_to_f16_combines_op_and_node_exclusions(simple_model): def test_convert_to_f16_refreshes_gathernd_pre_cast_declaration(monkeypatch): - def discover_test_plugins_without_trt(self): + def discover_test_plugins_without_trt(self, *, defer_nvfp4_trt_inference=False): + assert not defer_nvfp4_trt_inference self.custom_ops = { node.op_type for node in self.model.graph.node if node.domain == "test.plugins" } diff --git a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py index 78a70d88635..8074d47be96 100644 --- a/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py +++ b/tests/unit/torch/deploy/utils/test_torch_onnx_utils.py @@ -313,6 +313,7 @@ def test_fp8_export_rejects_unsupported_dtype_conversion( (torch.float32, "fp16", ("fp4", "int4"), ["onnxconverter"]), (torch.float32, "fp16", ("fp4", "int8"), ["onnxconverter"]), (torch.float32, "fp16", ("fp4", "mxfp8"), ["onnxconverter"]), + (torch.float32, "fp16", (), ["autocast"]), ], ids=[ "fp32-to-fp16", @@ -325,6 +326,7 @@ def test_fp8_export_rejects_unsupported_dtype_conversion( "fp4-int4-legacy", "fp4-int8-legacy", "fp4-mxfp8-legacy", + "unquantized-autocast", ], ) def test_quantized_export_selects_precision_converter( @@ -338,6 +340,7 @@ def record_onnxconverter(model, **kwargs): def record_autocast(model, **kwargs): calls.append("autocast") + assert kwargs["defer_nvfp4_trt_inference"] == ("fp4" in quantizers) return model for quantizer in ("fp4", "fp8", "int4", "int8", "mxfp8"): diff --git a/tests/unit/torch/quantization/test_onnx_export_cpu.py b/tests/unit/torch/quantization/test_onnx_export_cpu.py index 8d973c59bc5..71b0b6fc84e 100644 --- a/tests/unit/torch/quantization/test_onnx_export_cpu.py +++ b/tests/unit/torch/quantization/test_onnx_export_cpu.py @@ -119,15 +119,18 @@ def forward(self, inputs): class _NVFP4MixedPrecisionLinear(torch.nn.Module): - def __init__(self): + def __init__(self, return_boundaries=False): super().__init__() self.fp32_linear = torch.nn.Linear(16, 16, dtype=torch.float32) self.bf16_linear = torch.nn.Linear(16, 16, dtype=torch.bfloat16) + self.return_boundaries = return_boundaries def forward(self, inputs): fp32_output = self.fp32_linear(inputs) - bf16_output = self.bf16_linear(inputs.to(torch.bfloat16)).float() - return fp32_output + bf16_output + bf16_output = self.bf16_linear(inputs.to(torch.bfloat16)) + bf16_output_as_fp32 = bf16_output.float() + output = fp32_output + bf16_output_as_fp32 + return (output, bf16_output, bf16_output_as_fp32) if self.return_boundaries else output def _make_cpu_nvfp4_model( @@ -334,6 +337,36 @@ def capture_quantized_payloads(model, graph): assert tensor_types[node.input[1]] == expected_dtype +def test_nvfp4_deploy_export_preserves_bf16_graph_outputs(monkeypatch): + model = _NVFP4MixedPrecisionLinear(return_boundaries=True).eval() + sample_input = torch.ones(1, 16) + with torch.no_grad(): + source_outputs = model(sample_input) + assert [output.dtype for output in source_outputs] == [ + torch.float32, + torch.bfloat16, + torch.float32, + ] + model = _make_cpu_nvfp4_model(monkeypatch, model, sample_input) + exported_model, _ = _export_deploy_onnx_with_types(model, sample_input, "fp32") + assert [value.type.tensor_type.elem_type for value in exported_model.graph.output] == [ + TensorProto.FLOAT, + TensorProto.BFLOAT16, + TensorProto.FLOAT, + ] + + producers = {output: node for node in exported_model.graph.node for output in node.output} + native_output = exported_model.graph.output[1].name + cast_output = exported_model.graph.output[2].name + assert producers[native_output].op_type in {"Gemm", "MatMul"} + cast_node = producers[cast_output] + assert cast_node.op_type == "Cast" + assert list(cast_node.input) == [native_output] + assert helper.get_attribute_value(next(a for a in cast_node.attribute if a.name == "to")) == ( + TensorProto.FLOAT + ) + + @pytest.mark.parametrize( ("source_dtype", "weights_dtype", "expected_gemm_dtype", "expected_output_dtype"), [