Skip to content
17 changes: 17 additions & 0 deletions modelopt/onnx/autocast/convert.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
nodes.
"""

from copy import deepcopy

import numpy as np
import onnx

Expand All @@ -44,6 +46,17 @@
LATEST_IR_VERSION_SUPPORTED_BY_ORT = 10


def _capture_network_io_metadata(
model: onnx.ModelProto, keep_io_types: bool
) -> dict[str, list[onnx.ValueInfoProto]] | None:
if not keep_io_types:
return None
return {
"input": [deepcopy(io) for io in model.graph.input],
"output": [deepcopy(io) for io in model.graph.output],
}


def convert_to_mixed_precision(
onnx_path: str,
low_precision_type: str = "fp16",
Expand Down Expand Up @@ -97,6 +110,7 @@ def convert_to_mixed_precision(
# Load and process model
model = onnx.load(onnx_path, load_external_data=True)
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)

# Get original model's opset version
original_opset = onnx_utils.get_opset_version(model)
Expand Down Expand Up @@ -172,6 +186,7 @@ def convert_to_mixed_precision(
init_conversion_max_bytes=init_conversion_max_bytes,
custom_ops=graph_sanitizer.custom_ops,
use_standalone_type_inference=use_standalone_type_inference,
original_network_io_metadata=original_network_io_metadata,
)

# Obtain reference data
Expand Down Expand Up @@ -227,6 +242,7 @@ def convert_to_f16(
INT4 requires 21, NVFP4 requires 23).
"""
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)

# Check Q/DQ precision types in the model and determine required opset
qdq_precisions = get_qdq_precisions(model)
Expand Down Expand Up @@ -285,6 +301,7 @@ def convert_to_f16(
custom_ops=sanitizer.custom_ops,
tensor_block_dict=tensor_block_dict,
use_standalone_type_inference=use_standalone_type_inference,
original_network_io_metadata=original_network_io_metadata,
)
high_precision_nodes = [node.name for node in model.graph.node if node.op_type in op_block_list]
low_precision_nodes = [
Expand Down
271 changes: 268 additions & 3 deletions modelopt/onnx/autocast/precisionconverter.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ def __init__(
trt_plugins: list[str] | None = [],
tensor_block_dict: dict[str, dict[str, list[int]]] = {},
use_standalone_type_inference: bool = False,
original_network_io_metadata: dict[str, list[onnx.ValueInfoProto]] | None = None,
) -> None:
"""Initialize PrecisionConverter.

Expand All @@ -116,6 +117,7 @@ def __init__(
trt_plugins: List of custom TensorRT plugin library paths in .so format (compiled shared library).
tensor_block_dict: Dictionary of tensors (operation type and I/O indices) that should remain in FP32.
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.
"""
self.model = deepcopy(model)
self.value_info_map = value_info_map
Expand All @@ -139,6 +141,17 @@ 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],
}
if original_network_io_metadata is None
else {
field: [deepcopy(value) for value in values]
for field, values in original_network_io_metadata.items()
}
)
self.min_opset = min_opset
self.max_ir_version = max_ir_version
self.trt_plugins = trt_plugins
Expand Down Expand Up @@ -276,6 +289,8 @@ def convert(
# Remove redundant casts
self._cleanup()

self._restore_original_io_metadata()

self._sanity_check()

return self.model
Expand All @@ -289,6 +304,120 @@ 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 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

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:]
Comment thread
coderabbitai[bot] marked this conversation as resolved.

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 = []

Expand Down Expand Up @@ -397,19 +526,47 @@ def _propagate_cast_type_through_nodes(node, np_type, iter=1):
out.dtype = np_type

# Set the output shape
if not out.shape:
if isinstance(inp, gs.Constant):
if out.shape is 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)
_propagate_cast_type_through_nodes(node, np_type)

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
Expand Down Expand Up @@ -1043,6 +1200,7 @@ def _cleanup(self):
# Remove redundant casts
self._remove_redundant_casts()
self._deduplicate_network_output_producers()
self._refresh_gathernd_pre_cast_declarations()

def _cleanup_no_consumer_nodes(self):
network_outputs = {o.name for o in self.model.graph.output}
Expand Down Expand Up @@ -1184,6 +1342,113 @@ def _fix_network_output_names(self):
self.model
)

def _get_tensor_shape(self, tensor_name: str) -> list[int | str | None] | None:
initializer = next(
(value for value in self.model.graph.initializer if value.name == tensor_name), None
)
if initializer is not None:
return list(initializer.dims)

for value_info in (
*self.model.graph.input,
*self.model.graph.output,
*self.model.graph.value_info,
):
if value_info.name != tensor_name:
continue
tensor_type = value_info.type.tensor_type
if not tensor_type.HasField("shape"):
continue
shape = []
for dim in tensor_type.shape.dim:
if dim.HasField("dim_value"):
shape.append(dim.dim_value)
elif dim.HasField("dim_param"):
shape.append(dim.dim_param)
else:
shape.append(None)
return shape

producer_nodes = onnx_utils.get_producer_nodes(self.model, tensor_name)
if len(producer_nodes) == 1 and producer_nodes[0].op_type == "Cast":
return self._get_tensor_shape(producer_nodes[0].input[0])
return None

def _get_tensor_elem_type(self, tensor_name: str) -> int | None:
for value_info in (
*self.model.graph.input,
*self.model.graph.output,
*self.model.graph.value_info,
):
if value_info.name == tensor_name and value_info.type.HasField("tensor_type"):
elem_type = value_info.type.tensor_type.elem_type
if elem_type != onnx.TensorProto.UNDEFINED:
return elem_type
initializer = next(
(value for value in self.model.graph.initializer if value.name == tensor_name), None
)
return initializer.data_type if initializer is not None else None

def _refresh_gathernd_output_declaration(self, node_name: str, tensor_name: str) -> None:
node = next(
(
node
for node in self.model.graph.node
if node.name == node_name and tensor_name in node.output
),
None,
)
if node is None or node.op_type != "GatherND" or len(node.input) < 2:
return

data_shape = self._get_tensor_shape(node.input[0])
indices_shape = self._get_tensor_shape(node.input[1])
if data_shape is None or indices_shape is None or not indices_shape:
return
index_rank = indices_shape[-1]
batch_dims = next(
(attr.i for attr in node.attribute if attr.name == "batch_dims"),
0,
)
if not isinstance(index_rank, int) or not isinstance(batch_dims, int):
return
suffix_start = batch_dims + index_rank
if suffix_start > len(data_shape):
return

shape = indices_shape[:-1] + data_shape[suffix_start:]
elem_type = (
self._get_tensor_elem_type(tensor_name)
or self._get_tensor_elem_type(node.input[0])
or self.low_precision_type.onnx_type
)
value_info = self.value_info_map.get(tensor_name)
updated_value_info = helper.make_tensor_value_info(tensor_name, elem_type, shape)
if value_info is None:
self.model.graph.value_info.append(updated_value_info)
else:
value_info.CopyFrom(updated_value_info)

def _refresh_gathernd_pre_cast_declarations(self) -> None:
gathernd_pre_cast_outputs = [
(node.name, output)
for node in self.model.graph.node
if node.op_type == "GatherND"
for output in node.output
if output.endswith("_pre_cast")
]
if not gathernd_pre_cast_outputs:
return

self.value_info_map, self.initializer_map, self.node_to_init_map = utils.setup_mappings(
self.model
)
for node_name, tensor_name in gathernd_pre_cast_outputs:
self._refresh_gathernd_output_declaration(node_name, tensor_name)
self.value_info_map, self.initializer_map, self.node_to_init_map = utils.setup_mappings(
self.model
)

def _deduplicate_network_output_producers(self):
for output in self.model.graph.output:
producer_nodes = onnx_utils.get_producer_nodes(self.model, output.name)
Expand Down
Loading
Loading