From 274878ee8bbd39251765e19e1f40d36d9f0f5a18 Mon Sep 17 00:00:00 2001 From: Ajinkya Rasane Date: Wed, 16 Sep 2026 16:46:58 -0400 Subject: [PATCH] Fix ONNX output cast I/O type preservation Signed-off-by: Ajinkya Rasane --- CHANGELOG.rst | 1 + modelopt/onnx/autocast/precisionconverter.py | 19 +++++++------ .../onnx/autocast/test_precisionconverter.py | 27 +++++++++++++++++++ 3 files changed, 39 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 01a0fbe1b47..f02e3250897 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -25,6 +25,7 @@ Changelog **Bug Fixes** +- Fix ONNX FP16 conversion failing to preserve public output types when type inference changes a graph output declaration before output casts are inserted. - 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 7db660b29b4..a78249fdca4 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -142,13 +142,6 @@ def __init__( self.low_precision_type = PRECISION_MAP[low_precision_type] self.high_precision_type = PRECISION_MAP["fp32"] - # Preserve original network inputs and outputs for sanity checks - self.original_network_io = { - io.name: io.type.tensor_type.elem_type for io in self.model.graph.input - } - self.original_network_io.update( - {io.name: io.type.tensor_type.elem_type for io in self.model.graph.output} - ) self.original_network_io_metadata = ( { "input": [deepcopy(io) for io in self.model.graph.input], @@ -160,6 +153,13 @@ def __init__( for field, values in original_network_io_metadata.items() } ) + # Preserve the public I/O types captured at the API boundary. Type inference may have + # changed the working model's declarations before the converter is initialized. + self.original_network_io = { + io.name: io.type.tensor_type.elem_type + for values in self.original_network_io_metadata.values() + for io in values + } self.min_opset = min_opset self.max_ir_version = max_ir_version self.trt_plugins = trt_plugins @@ -1440,7 +1440,10 @@ def _add_cast( # Update network output for output in self.model.graph.output: if output.name == tensor_name and ( - (self.keep_io_types and cast_to.onnx_type == output.type.tensor_type.elem_type) + ( + self.keep_io_types + and cast_to.onnx_type == self.original_network_io.get(tensor_name) + ) or ( not self.keep_io_types and cast_to.onnx_type == self.low_precision_type.onnx_type diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index d5ed09215fc..d3e0a786a65 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -2220,6 +2220,33 @@ def test_convert_to_f16_restores_public_io_metadata_from_entry_boundary(): onnx.checker.check_model(converted, full_check=True) +def test_convert_to_f16_preserves_declared_output_type_after_inference_changes_it(): + graph_input = helper.make_tensor_value_info("X", TensorProto.FLOAT16, [2, 3]) + graph_output = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [2, 3]) + node = helper.make_node("Identity", ["X"], ["Y"], name="Identity_0") + graph = helper.make_graph([node], "inferred_output_type", [graph_input], [graph_output]) + model = helper.make_model( + graph, + producer_name="inferred_output_type", + opset_imports=[helper.make_opsetid("", 19)], + ir_version=10, + ) + + converted = convert_to_f16( + model, keep_io_types=True, op_block_list=[], trt_plugins=[], opset=19 + ) + + output = next(vi for vi in converted.graph.output if vi.name == "Y") + assert output.type.tensor_type.elem_type == TensorProto.FLOAT + output_producers = [node for node in converted.graph.node if "Y" in node.output] + assert len(output_producers) == 1 + assert output_producers[0].op_type == "Cast" + assert next(attr.i for attr in output_producers[0].attribute if attr.name == "to") == ( + TensorProto.FLOAT + ) + onnx.checker.check_model(converted, full_check=True) + + def test_convert_to_f16_combines_op_and_node_exclusions(simple_model): model, *_ = simple_model converted = convert_to_f16(