Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,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.

- ``examples/hf_ptq`` no longer detects MTP layers by name. Weights the loader could not place -- an MTP head, an auxiliary tower -- are identified from Transformers' own accounting: the model is loaded with ``from_pretrained(..., output_loading_info=True)`` and the reported ``unexpected_keys`` (present in the checkpoint, not in the model's architecture) are recorded on the model and carried into the export unchanged. Everything the loader *did* place goes through the normal export path. This removes ``load_mtp_weights``, ``mtp_layer_prefixes_from_checkpoint`` and their support matrix of MTP storage conventions, along with ``_add_mtp_exclusions`` and the pre-quantization ``enable: False`` entries ``hf_ptq`` appended to the recipe's ``quant_cfg``. Two consequences: MTP layers now follow the recipe like any other module instead of being force-excluded by the script -- matching ``examples/megatron_bridge``, which has no MTP-specific code at all -- and ``quantization_config.ignore`` can no longer claim a layer is unquantized that the export in fact quantized. Recipes importing ``configs/ptq/units/default_disabled_quantizers`` still disable ``mtp.*``, so their behaviour is unchanged; a recipe omitting that unit will now quantize an MTP the model actually built.

- ``examples/hf_ptq --vllm_fakequant_export`` now raises ``NotImplementedError`` when the checkpoint holds weights the model has no parameter for and a shard actually provides them (an MTP head, an auxiliary tower). The fake-quant exporter writes only model-backed state, so it would otherwise drop those weights silently -- and a fake-quant checkpoint is evaluated, where a missing head changes the score rather than failing loudly. Use the unified HF export, which carries them through. Buffers Transformers recomputes are not weights to lose: ``*.inv_freq`` is skipped even when a shard provides it, since older Llama/Mistral-lineage conversions do list it in the index and refusing an export over it would reject checkpoints that export correctly today. The check runs immediately after the model loads, not at export time, so an incompatible run fails before calibration rather than after it.
Expand Down Expand Up @@ -81,6 +83,8 @@ 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 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 HuggingFace checkpoint export failing with ``activation scaling factor 0.0 not positive`` when a dynamic-block quantizer (such as an NVFP4 input quantizer) ends calibration with an amax of zero because the calibration data never activated that layer or expert. Such a quantizer now exports a positive fallback scale and warns instead of crashing, matching what static quantizers already did; if you see the warning, check whether the layer is expected to be inactive and consider a larger calibration size.
- 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.
Expand Down
7 changes: 6 additions & 1 deletion modelopt/onnx/autocast/convert.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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)
Expand Down Expand Up @@ -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()
Expand All @@ -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 = [
Expand Down
43 changes: 36 additions & 7 deletions modelopt/onnx/autocast/graphsanitizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -132,11 +139,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 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
)
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."""
Expand Down
26 changes: 19 additions & 7 deletions modelopt/onnx/autocast/precisionconverter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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
Expand Down Expand Up @@ -218,7 +222,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 = (
Expand Down Expand Up @@ -416,6 +420,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,
Expand All @@ -430,6 +436,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"])
Expand All @@ -444,7 +454,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
Expand All @@ -456,8 +468,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
Expand Down Expand Up @@ -509,7 +521,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}."
Expand Down Expand Up @@ -1240,7 +1252,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)
Expand Down Expand Up @@ -1844,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
Expand Down
Loading
Loading