diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 46168b7fe0f..071f588371b 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -38,6 +38,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 ``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. - Fix ONNX INT8 entropy calibration failing or producing invalid quantization parameters for FP16 activations. 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(