From ea8f814231bfff6acc373c97b0c5f56f3838d5c1 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Fri, 17 Jul 2026 01:06:59 +0000 Subject: [PATCH 1/7] Fix ONNX autocast metadata propagation Update ONNX autocast cleanup to keep folded Constant metadata in sync and infer GatherND output shapes in custom-op mode without running global ONNX shape inference. This preserves valid ONNX metadata after redundant Cast folding and avoids incorrect GatherND rank propagation when TensorRT plugin custom ops are present. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/autocast/precisionconverter.py | 31 +++++++- modelopt/onnx/utils.py | 19 +++++ .../onnx/autocast/test_precisionconverter.py | 77 +++++++++++++++++++ 3 files changed, 126 insertions(+), 1 deletion(-) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index 5474b57a062..25996daf49c 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -289,6 +289,33 @@ def _ensure_types_are_defined(self): def _propagate_types_shapes_custom_ops(self, model): """Propagate types and shapes after insertion of 'Cast' nodes or other graph modifications.""" logger.info("Propagating tensor shapes and types in model with custom ops.") + + def _get_shape(tensor): + if isinstance(tensor, gs.Constant): + return list(tensor.values.shape) + if not tensor.shape: + return None + return list(tensor.shape) + + def _infer_gathernd_op_shape(node): + if node.op != "GatherND" or len(node.inputs) < 2: + return None + + data_shape = _get_shape(node.inputs[0]) + indices_shape = _get_shape(node.inputs[1]) + if not data_shape or not indices_shape: + return None + + index_rank = indices_shape[-1] + batch_dims = node.attrs.get("batch_dims", 0) + if not isinstance(index_rank, int) or not isinstance(batch_dims, int): + return None + + suffix_start = batch_dims + index_rank + if suffix_start > len(data_shape): + return None + return indices_shape[:-1] + data_shape[suffix_start:] + graph = gs.import_onnx(model) traversed_tensors = [] @@ -398,7 +425,9 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): # Set the output shape if not out.shape: - if isinstance(inp, gs.Constant): + if shape := _infer_gathernd_op_shape(node): + out.shape = shape + elif isinstance(inp, gs.Constant): out.shape = inp.values.shape elif inp.inputs and inp.inputs[0].op == "Constant": out.shape = inp.inputs[0].attrs["value"].values.shape diff --git a/modelopt/onnx/utils.py b/modelopt/onnx/utils.py index 3b8edf76a84..c3efdcb63be 100644 --- a/modelopt/onnx/utils.py +++ b/modelopt/onnx/utils.py @@ -1532,6 +1532,22 @@ def _convert_constant_values(constant_node: onnx.NodeProto, cast_node: onnx.Node break +def _sync_value_info_elem_type(graph: onnx.GraphProto, tensor_name: str, elem_type: int) -> None: + """Synchronize declarations for a tensor whose producer dtype changed.""" + for value_info in list(graph.value_info) + list(graph.input) + list(graph.output): + tensor_type = value_info.type.tensor_type + if value_info.name == tensor_name and tensor_type.elem_type: + tensor_type.elem_type = elem_type + + for node in graph.node: + for attr in node.attribute: + if attr.type == onnx.AttributeProto.GRAPH: + _sync_value_info_elem_type(attr.g, tensor_name, elem_type) + elif attr.type == onnx.AttributeProto.GRAPHS: + for subgraph in attr.graphs: + _sync_value_info_elem_type(subgraph, tensor_name, elem_type) + + def remove_redundant_casts(onnx_model: onnx.ModelProto) -> onnx.ModelProto: """Removes both sequential casts and casts that don't change precision. @@ -1571,6 +1587,9 @@ def remove_redundant_casts(onnx_model: onnx.ModelProto) -> onnx.ModelProto: assert len(cast_producers) == 1 and cast_producers[0].op_type == "Constant" constant_producer = cast_producers[0] _convert_constant_values(constant_producer, node) + _sync_value_info_elem_type( + onnx_model.graph, constant_producer.output[0], get_cast_to_type(node) + ) _bypass_cast_node(onnx_model, node) logger.debug(f"Found foldable Constant->Cast pattern, removing {node.name}") diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index 4c0f19759e3..a7334743732 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -1854,3 +1854,80 @@ def test_if_subgraph_outer_scope_type_preservation( assert len(else_x_info) > 0, "X value_info should be preserved in else branch" assert then_x_info[0].type.tensor_type.elem_type != onnx.TensorProto.UNDEFINED assert else_x_info[0].type.tensor_type.elem_type != onnx.TensorProto.UNDEFINED + + +def test_folded_constant_cast_updates_value_info_type(): + const_tensor = numpy_helper.from_array( + np.array([1.0, 2.0], dtype=np.float32), name="const_value" + ) + const_node = helper.make_node( + "Constant", [], ["const_out"], name="const_node", value=const_tensor + ) + cast_node = helper.make_node( + "Cast", ["const_out"], ["cast_out"], name="cast_to_fp16", to=TensorProto.FLOAT16 + ) + identity_node = helper.make_node("Identity", ["cast_out"], ["Y"], name="identity") + + graph = helper.make_graph( + [const_node, cast_node, identity_node], + "constant_cast_value_info", + [], + [helper.make_tensor_value_info("Y", TensorProto.FLOAT16, [2])], + [], + value_info=[helper.make_tensor_value_info("const_out", TensorProto.FLOAT, [2])], + ) + model = helper.make_model(graph, producer_name="constant_cast_value_info") + model.opset_import[0].version = 19 + model.ir_version = 10 + + folded = onnx_utils.remove_redundant_casts(model) + + assert [node.op_type for node in folded.graph.node] == ["Constant", "Identity"] + const_out = next(vi for vi in folded.graph.value_info if vi.name == "const_out") + assert const_out.type.tensor_type.elem_type == TensorProto.FLOAT16 + onnx.shape_inference.infer_shapes(folded, strict_mode=True, check_type=True) + + +def test_custom_op_mode_uses_schema_shape_for_standard_gathernd(): + data = helper.make_tensor_value_info("data", TensorProto.FLOAT, [1, 4, 2]) + plugin_in = helper.make_tensor_value_info("plugin_in", TensorProto.FLOAT, [1, 4, 2]) + indices_init = numpy_helper.from_array( + np.array([[[0, 0], [3, 1]]], dtype=np.int64), name="indices" + ) + custom_node = helper.make_node( + "FakeTensorRTPlugin", ["plugin_in"], ["plugin_out"], name="fake_plugin" + ) + gather_node = helper.make_node( + "GatherND", + ["data", "indices"], + ["last_token_embed"], + name="shape_changing_gathernd", + batch_dims=1, + ) + graph = helper.make_graph( + [custom_node, gather_node], + "custom_op_gathernd_shape", + [data, plugin_in], + [ + helper.make_tensor_value_info("plugin_out", TensorProto.FLOAT, [1, 4, 2]), + helper.make_tensor_value_info("last_token_embed", TensorProto.FLOAT, None), + ], + [indices_init], + ) + model = helper.make_model(graph, producer_name="custom_op_gathernd_shape") + model.opset_import[0].version = 19 + model.ir_version = 10 + value_info_map, initializer_map, node_to_init_map = utils.setup_mappings(model) + + converter = PrecisionConverter( + model, + value_info_map, + initializer_map, + node_to_init_map, + keep_io_types=True, + custom_ops={"FakeTensorRTPlugin"}, + ) + propagated = converter._propagate_types_shapes_custom_ops(model) + + output = next(vi for vi in propagated.graph.output if vi.name == "last_token_embed") + assert [dim.dim_value for dim in output.type.tensor_type.shape.dim] == [1, 2] From ddd5b03f7089c181d5b59aa1bb1bcc70b19a52d3 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Fri, 17 Jul 2026 13:45:59 +0000 Subject: [PATCH 2/7] Handle scalar GatherND shape propagation Treat an empty inferred GatherND shape as a valid scalar shape instead of falling through to the generic input-shape fallback in custom-op propagation. Add a regression test for the rank-0 GatherND output case. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/autocast/precisionconverter.py | 3 +- .../onnx/autocast/test_precisionconverter.py | 42 +++++++++++++++++++ 2 files changed, 44 insertions(+), 1 deletion(-) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index 25996daf49c..faa85c252d5 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -425,7 +425,8 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): # Set the output shape if not out.shape: - if shape := _infer_gathernd_op_shape(node): + shape = _infer_gathernd_op_shape(node) + if shape is not None: out.shape = shape elif isinstance(inp, gs.Constant): out.shape = inp.values.shape diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index a7334743732..9468a9e6d1c 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -1931,3 +1931,45 @@ def test_custom_op_mode_uses_schema_shape_for_standard_gathernd(): output = next(vi for vi in propagated.graph.output if vi.name == "last_token_embed") assert [dim.dim_value for dim in output.type.tensor_type.shape.dim] == [1, 2] + + +def test_custom_op_mode_preserves_scalar_gathernd_shape(): + data = helper.make_tensor_value_info("data", TensorProto.FLOAT, [4]) + plugin_in = helper.make_tensor_value_info("plugin_in", TensorProto.FLOAT, [4]) + indices_init = numpy_helper.from_array(np.array([2], dtype=np.int64), name="indices") + custom_node = helper.make_node( + "FakeTensorRTPlugin", ["plugin_in"], ["plugin_out"], name="fake_plugin" + ) + gather_node = helper.make_node( + "GatherND", + ["data", "indices"], + ["selected_scalar"], + name="scalar_gathernd", + ) + graph = helper.make_graph( + [custom_node, gather_node], + "custom_op_scalar_gathernd_shape", + [data, plugin_in], + [ + helper.make_tensor_value_info("plugin_out", TensorProto.FLOAT, [4]), + helper.make_tensor_value_info("selected_scalar", TensorProto.FLOAT, None), + ], + [indices_init], + ) + model = helper.make_model(graph, producer_name="custom_op_scalar_gathernd_shape") + model.opset_import[0].version = 19 + model.ir_version = 10 + value_info_map, initializer_map, node_to_init_map = utils.setup_mappings(model) + + converter = PrecisionConverter( + model, + value_info_map, + initializer_map, + node_to_init_map, + keep_io_types=True, + custom_ops={"FakeTensorRTPlugin"}, + ) + propagated = converter._propagate_types_shapes_custom_ops(model) + + output = next(vi for vi in propagated.graph.output if vi.name == "selected_scalar") + assert [dim.dim_value for dim in output.type.tensor_type.shape.dim] == [] From c60259f83638ff381d99ba9418fb0c21fdce7ea0 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Fri, 17 Jul 2026 13:51:39 +0000 Subject: [PATCH 3/7] Simplify GatherND shape check Use a walrus expression with an explicit None check so scalar GatherND shapes remain valid while keeping the custom-op propagation branch concise. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/autocast/precisionconverter.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index faa85c252d5..3396045def4 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -425,8 +425,7 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): # Set the output shape if not out.shape: - shape = _infer_gathernd_op_shape(node) - if shape is not None: + if (shape := _infer_gathernd_op_shape(node)) is not None: out.shape = shape elif isinstance(inp, gs.Constant): out.shape = inp.values.shape From 879ce4911c7459aaf40d3208e118142e08eaa091 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Fri, 17 Jul 2026 13:56:09 +0000 Subject: [PATCH 4/7] Handle undefined Constant metadata declarations Synchronize matching ONNX tensor declarations even when the previous element type is TensorProto.UNDEFINED, and cover that Constant Cast folding case in the regression test. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/utils.py | 5 ++--- tests/unit/onnx/autocast/test_precisionconverter.py | 5 +++-- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/modelopt/onnx/utils.py b/modelopt/onnx/utils.py index c3efdcb63be..3f62376923c 100644 --- a/modelopt/onnx/utils.py +++ b/modelopt/onnx/utils.py @@ -1535,9 +1535,8 @@ def _convert_constant_values(constant_node: onnx.NodeProto, cast_node: onnx.Node def _sync_value_info_elem_type(graph: onnx.GraphProto, tensor_name: str, elem_type: int) -> None: """Synchronize declarations for a tensor whose producer dtype changed.""" for value_info in list(graph.value_info) + list(graph.input) + list(graph.output): - tensor_type = value_info.type.tensor_type - if value_info.name == tensor_name and tensor_type.elem_type: - tensor_type.elem_type = elem_type + if value_info.name == tensor_name and value_info.type.HasField("tensor_type"): + value_info.type.tensor_type.elem_type = elem_type for node in graph.node: for attr in node.attribute: diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index 9468a9e6d1c..856e3b42366 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -1856,7 +1856,8 @@ def test_if_subgraph_outer_scope_type_preservation( assert else_x_info[0].type.tensor_type.elem_type != onnx.TensorProto.UNDEFINED -def test_folded_constant_cast_updates_value_info_type(): +@pytest.mark.parametrize("value_info_elem_type", [TensorProto.FLOAT, TensorProto.UNDEFINED]) +def test_folded_constant_cast_updates_value_info_type(value_info_elem_type): const_tensor = numpy_helper.from_array( np.array([1.0, 2.0], dtype=np.float32), name="const_value" ) @@ -1874,7 +1875,7 @@ def test_folded_constant_cast_updates_value_info_type(): [], [helper.make_tensor_value_info("Y", TensorProto.FLOAT16, [2])], [], - value_info=[helper.make_tensor_value_info("const_out", TensorProto.FLOAT, [2])], + value_info=[helper.make_tensor_value_info("const_out", value_info_elem_type, [2])], ) model = helper.make_model(graph, producer_name="constant_cast_value_info") model.opset_import[0].version = 19 From deed3557b1bed06a60707bfe2e0d8036030aaed0 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Mon, 20 Jul 2026 18:07:42 +0000 Subject: [PATCH 5/7] Fix custom-op metadata for standard shape ops Preserve full public I/O metadata when keep_io_types=True and avoid applying the input-0 shape fallback to ordinary standard operators in custom-op mode. Infer shapes for standard Gather, Unsqueeze, Shape, and GatherND paths so strict ONNX checking passes without downstream metadata restoration. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/autocast/precisionconverter.py | 127 +++++++++++++++++- .../onnx/autocast/test_precisionconverter.py | 70 ++++++++++ 2 files changed, 194 insertions(+), 3 deletions(-) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index 3396045def4..def008f96c7 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -139,6 +139,10 @@ def __init__( 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], + "output": [deepcopy(io) for io in self.model.graph.output], + } self.min_opset = min_opset self.max_ir_version = max_ir_version self.trt_plugins = trt_plugins @@ -276,6 +280,8 @@ def convert( # Remove redundant casts self._cleanup() + self._restore_original_io_metadata() + self._sanity_check() return self.model @@ -293,10 +299,21 @@ def _propagate_types_shapes_custom_ops(self, model): def _get_shape(tensor): if isinstance(tensor, gs.Constant): return list(tensor.values.shape) - if not tensor.shape: + if tensor.shape is None: return None return list(tensor.shape) + def _get_const_values(tensor): + if isinstance(tensor, gs.Constant): + return tensor.values + if tensor.inputs and tensor.inputs[0].op == "Constant": + return tensor.inputs[0].attrs["value"].values + return None + + def _get_int_attr(node, attr_name, default): + value = node.attrs.get(attr_name, default) + return value if isinstance(value, int) else None + def _infer_gathernd_op_shape(node): if node.op != "GatherND" or len(node.inputs) < 2: return None @@ -316,6 +333,82 @@ def _infer_gathernd_op_shape(node): return None return indices_shape[:-1] + data_shape[suffix_start:] + def _infer_gather_op_shape(node): + if node.op != "Gather" or len(node.inputs) < 2: + return None + + data_shape = _get_shape(node.inputs[0]) + indices_shape = _get_shape(node.inputs[1]) + if data_shape is None or indices_shape is None: + return None + + axis = _get_int_attr(node, "axis", 0) + if axis is None: + return None + if axis < 0: + axis += len(data_shape) + if axis < 0 or axis >= len(data_shape): + return None + + return data_shape[:axis] + indices_shape + data_shape[axis + 1 :] + + def _infer_unsqueeze_op_shape(node): + if node.op != "Unsqueeze" or len(node.inputs) < 2: + return None + + data_shape = _get_shape(node.inputs[0]) + axes = _get_const_values(node.inputs[1]) + if data_shape is None or axes is None: + return None + + axes = [int(axis) for axis in np.asarray(axes).flatten()] + output_rank = len(data_shape) + len(axes) + normalized_axes = [] + for axis in axes: + if axis < 0: + axis += output_rank + if axis < 0 or axis >= output_rank: + return None + normalized_axes.append(axis) + + output_shape = list(data_shape) + for axis in sorted(normalized_axes): + output_shape.insert(axis, 1) + return output_shape + + def _infer_shape_op_shape(node): + if node.op != "Shape" or not node.inputs: + return None + + data_shape = _get_shape(node.inputs[0]) + if data_shape is None: + return None + + rank = len(data_shape) + start = _get_int_attr(node, "start", 0) + end = _get_int_attr(node, "end", rank) + if start is None or end is None: + return None + if start < 0: + start += rank + if end < 0: + end += rank + start = min(max(start, 0), rank) + end = min(max(end, 0), rank) + return [max(end - start, 0)] + + def _infer_standard_op_shape(node): + for infer_shape in ( + _infer_gathernd_op_shape, + _infer_gather_op_shape, + _infer_unsqueeze_op_shape, + _infer_shape_op_shape, + ): + shape = infer_shape(node) + if shape is not None: + return shape + return None + graph = gs.import_onnx(model) traversed_tensors = [] @@ -425,13 +518,13 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): # Set the output shape if not out.shape: - if (shape := _infer_gathernd_op_shape(node)) is not None: + if (shape := _infer_standard_op_shape(node)) is not None: out.shape = shape elif isinstance(inp, gs.Constant): out.shape = inp.values.shape elif inp.inputs and inp.inputs[0].op == "Constant": out.shape = inp.inputs[0].attrs["value"].values.shape - elif inp.shape: + elif node.op in self.custom_ops and inp.shape: out.shape = inp.shape # Propagate tensor types to the children nodes (until another Cast or Q node is met) @@ -439,6 +532,34 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): return gs.export_onnx(graph) + def _restore_original_io_metadata(self) -> None: + """Preserve complete public I/O metadata when keep_io_types=True.""" + if not self.keep_io_types: + return + + for io_field in ("input", "output"): + current_values = list(getattr(self.model.graph, io_field)) + original_values = self.original_network_io_metadata[io_field] + current_names = [value.name for value in current_values] + original_names = [value.name for value in original_values] + if current_names != original_names: + raise RuntimeError( + f"Cannot restore public graph {io_field} metadata because names changed: " + f"{original_names} -> {current_names}" + ) + for current_value, original_value in zip( + current_values, original_values, strict=True + ): + if ( + current_value.type.tensor_type.elem_type + != original_value.type.tensor_type.elem_type + ): + raise RuntimeError( + f"Cannot restore public graph {io_field} metadata for {current_value.name}: " + "element type changed" + ) + current_value.CopyFrom(original_value) + def _is_bf16(self, type: PrecisionTypes = None) -> bool: if type is None: type = self.low_precision_type diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index 856e3b42366..d80d552eeed 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -1974,3 +1974,73 @@ def test_custom_op_mode_preserves_scalar_gathernd_shape(): output = next(vi for vi in propagated.graph.output if vi.name == "selected_scalar") assert [dim.dim_value for dim in output.type.tensor_type.shape.dim] == [] + + +def test_custom_op_mode_uses_schema_shapes_for_standard_rank_changes(): + gather_data = helper.make_tensor_value_info("gather_data", TensorProto.FLOAT, [4]) + unsqueeze_data = helper.make_tensor_value_info("unsqueeze_data", TensorProto.FLOAT, [4]) + plugin_in = helper.make_tensor_value_info("plugin_in", TensorProto.FLOAT, [1]) + gather_index = numpy_helper.from_array(np.array(0, dtype=np.int64), "gather_index") + unsqueeze_axes = numpy_helper.from_array(np.array([0], dtype=np.int64), "unsqueeze_axes") + gather_node = helper.make_node( + "Gather", ["gather_data", "gather_index"], ["gather_y_pre_cast"], name="gather" + ) + gather_cast = helper.make_node( + "Cast", ["gather_y_pre_cast"], ["gather_y"], name="gather_cast", to=TensorProto.FLOAT + ) + unsqueeze_node = helper.make_node( + "Unsqueeze", + ["unsqueeze_data", "unsqueeze_axes"], + ["unsqueeze_y_pre_cast"], + name="unsqueeze", + ) + unsqueeze_cast = helper.make_node( + "Cast", + ["unsqueeze_y_pre_cast"], + ["unsqueeze_y"], + name="unsqueeze_cast", + to=TensorProto.FLOAT, + ) + custom_node = helper.make_node( + "FakePlugin", ["plugin_in"], ["plugin_y"], name="plugin", domain="test.plugins" + ) + graph = helper.make_graph( + [gather_node, gather_cast, unsqueeze_node, unsqueeze_cast, custom_node], + "custom_op_standard_rank_changes", + [gather_data, unsqueeze_data, plugin_in], + [ + helper.make_tensor_value_info("gather_y", TensorProto.FLOAT, []), + helper.make_tensor_value_info("unsqueeze_y", TensorProto.FLOAT, [1, 4]), + helper.make_tensor_value_info("plugin_y", TensorProto.FLOAT, [1]), + ], + [gather_index, unsqueeze_axes], + ) + model = helper.make_model( + graph, + producer_name="custom_op_standard_rank_changes", + opset_imports=[helper.make_opsetid("", 19), helper.make_opsetid("test.plugins", 1)], + ir_version=10, + ) + value_info_map, initializer_map, node_to_init_map = utils.setup_mappings(model) + + converter = PrecisionConverter( + model, + value_info_map, + initializer_map, + node_to_init_map, + keep_io_types=True, + custom_ops={"FakePlugin"}, + ) + propagated = converter._propagate_types_shapes_custom_ops(model) + + value_infos = {vi.name: vi for vi in [*propagated.graph.value_info, *propagated.graph.output]} + assert [dim.dim_value for dim in value_infos["gather_y_pre_cast"].type.tensor_type.shape.dim] == [] + assert [ + dim.dim_value for dim in value_infos["unsqueeze_y_pre_cast"].type.tensor_type.shape.dim + ] == [1, 4] + assert [dim.dim_value for dim in value_infos["gather_y"].type.tensor_type.shape.dim] == [] + assert [dim.dim_value for dim in value_infos["unsqueeze_y"].type.tensor_type.shape.dim] == [ + 1, + 4, + ] + onnx.checker.check_model(propagated, full_check=True) From c7d705dfcb84f4f2360909a4f2baa41e18d34ad1 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Mon, 20 Jul 2026 18:47:48 +0000 Subject: [PATCH 6/7] Preserve scalar shapes in custom-op propagation Treat only None as a missing tensor shape so known scalar shapes are not clobbered by custom-op fallback propagation. Add a regression for a custom op with scalar output metadata. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/autocast/precisionconverter.py | 2 +- .../onnx/autocast/test_precisionconverter.py | 34 +++++++++++++++++++ 2 files changed, 35 insertions(+), 1 deletion(-) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index def008f96c7..1afce008ff2 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -517,7 +517,7 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1): out.dtype = np_type # Set the output shape - if not out.shape: + if out.shape is None: if (shape := _infer_standard_op_shape(node)) is not None: out.shape = shape elif isinstance(inp, gs.Constant): diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index d80d552eeed..2c92d253691 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -2044,3 +2044,37 @@ def test_custom_op_mode_uses_schema_shapes_for_standard_rank_changes(): 4, ] onnx.checker.check_model(propagated, full_check=True) + + +def test_custom_op_mode_preserves_known_scalar_custom_op_shape(): + plugin_in = helper.make_tensor_value_info("plugin_in", TensorProto.FLOAT, [4]) + custom_node = helper.make_node( + "FakePlugin", ["plugin_in"], ["plugin_scalar"], name="plugin", domain="test.plugins" + ) + graph = helper.make_graph( + [custom_node], + "custom_op_known_scalar_shape", + [plugin_in], + [helper.make_tensor_value_info("plugin_scalar", TensorProto.FLOAT, [])], + [], + ) + model = helper.make_model( + graph, + producer_name="custom_op_known_scalar_shape", + opset_imports=[helper.make_opsetid("", 19), helper.make_opsetid("test.plugins", 1)], + ir_version=10, + ) + value_info_map, initializer_map, node_to_init_map = utils.setup_mappings(model) + + converter = PrecisionConverter( + model, + value_info_map, + initializer_map, + node_to_init_map, + keep_io_types=True, + custom_ops={"FakePlugin"}, + ) + propagated = converter._propagate_types_shapes_custom_ops(model) + + output = next(vi for vi in propagated.graph.output if vi.name == "plugin_scalar") + assert [dim.dim_value for dim in output.type.tensor_type.shape.dim] == [] From 83197e563ebeb6b34d758d95c90d072ff1f225e7 Mon Sep 17 00:00:00 2001 From: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> Date: Mon, 20 Jul 2026 18:50:39 +0000 Subject: [PATCH 7/7] Apply pre-commit formatting fixes Apply formatting requested by pre-commit for the scalar-shape propagation follow-up without changing behavior. Signed-off-by: Gwena Cunha <4861122+gcunhase@users.noreply.github.com> --- modelopt/onnx/autocast/precisionconverter.py | 4 +--- tests/unit/onnx/autocast/test_precisionconverter.py | 4 +++- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index 1afce008ff2..13a977f46ba 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -547,9 +547,7 @@ def _restore_original_io_metadata(self) -> None: f"Cannot restore public graph {io_field} metadata because names changed: " f"{original_names} -> {current_names}" ) - for current_value, original_value in zip( - current_values, original_values, strict=True - ): + for current_value, original_value in zip(current_values, original_values, strict=True): if ( current_value.type.tensor_type.elem_type != original_value.type.tensor_type.elem_type diff --git a/tests/unit/onnx/autocast/test_precisionconverter.py b/tests/unit/onnx/autocast/test_precisionconverter.py index 2c92d253691..435de5ce3b6 100644 --- a/tests/unit/onnx/autocast/test_precisionconverter.py +++ b/tests/unit/onnx/autocast/test_precisionconverter.py @@ -2034,7 +2034,9 @@ def test_custom_op_mode_uses_schema_shapes_for_standard_rank_changes(): propagated = converter._propagate_types_shapes_custom_ops(model) value_infos = {vi.name: vi for vi in [*propagated.graph.value_info, *propagated.graph.output]} - assert [dim.dim_value for dim in value_infos["gather_y_pre_cast"].type.tensor_type.shape.dim] == [] + assert [ + dim.dim_value for dim in value_infos["gather_y_pre_cast"].type.tensor_type.shape.dim + ] == [] assert [ dim.dim_value for dim in value_infos["unsqueeze_y_pre_cast"].type.tensor_type.shape.dim ] == [1, 4]