[5591371] Add performance guard for ONNX Autotune (#2318)

### What does this PR do?

Type of change: Bug fix

This PR prevents integrated ONNX Autotune from saving an INT8/FP8 result
that does not improve TensorRT latency. Autotune search models now use
the same FP16/BF16 conversion path as the delivered model. After
calibration and existing Q/DQ post-processing, the exact candidate is
benchmarked against its precision-matched no-Q/DQ baseline.

- Keep Q/DQ when the measured speedup meets
`Config.performance_threshold` (`1.02x` by default, inclusive).
- Otherwise save the high-precision no-Q/DQ fallback at the requested
output path and report `no_qdq`.
- Reject already-quantized Autotune inputs because they cannot produce a
true no-Q/DQ baseline.
- Leave quantization without Autotune, standalone uncalibrated Autotune,
pattern search, caches, and state schema v1 unchanged.

### Usage

```bash
python -m modelopt.onnx.quantization \
  --onnx_path=model.onnx \
  --quantize_mode=fp8 \
  --calibration_data_path=calibration.npz \
  --high_precision_dtype=fp16 \
  --autotune=default \
  --output_path=model.autotuned.onnx
```

The output contains either the accepted Q/DQ placement or the
high-precision fallback. The log reports `qdq` or `no_qdq`, the two
measured latencies, the speedup, and the threshold.

### Testing

- Ran the CPU-only ONNX Autotune, runtime-precision, and quantization
API suites with no GPU visible and CPU execution providers: 199 passed.
- Ran all applicable pre-commit hooks on the 12 changed files, including
Ruff, mypy, Bandit, license, and RST checks.
- On an RTX 6000 Ada GPU with TensorRT 10.8, ran explicit
`--autotune=default` on a synthetic `Conv(128→128) → Relu → MaxPool →
Gemm` graph. The calibrated guard retained two Q/DQ sites from its
paired measurement (`0.066 ms / 0.064 ms = 1.023x`, threshold `1.020x`).
The selected, baseline, and candidate models all built with `trtexec
--stronglyTyped` without an output-type error. Five alternating
follow-up trials also favored Q/DQ (`1.016x` median speedup).

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ✅
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
- Did you get Claude approval on this PR?: ❌

### Additional Information

Related to #439.

> 🤖 _Generated by Codex (AI agent)._


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **New Features**
- ONNX Autotune now benchmarks candidates at the requested runtime
precision and supports custom model transformations during export.
- Calibrated INT8/FP8 quantization is retained only when it meets the
configured performance threshold (default 1.02×); otherwise, the
high-precision model is saved without Q/DQ.

- **Bug Fixes**
  - Improved runtime-precision handling across INT8 and FP8 workflows.
- Improved validation of inputs, pre-quantized models, failures, and
temporary resources.

- **Documentation**
- Updated Autotune guidance and command-line help to explain runtime
precision, performance validation, and fallback behavior.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com>
Co-authored-by: Codex <codex@openai.com>
This commit is contained in:
Ajinkya Rasane
2026-09-03 12:20:27 -04:00
committed by GitHub
co-authored by Codex
parent 51cc5dbade
commit c49ce57d75
12 changed files with 729 additions and 119 deletions
+1
View File
@@ -8,6 +8,7 @@ Changelog
*Quantization*
- ONNX quantization with Autotune now benchmarks placements in the requested runtime precision and retains calibrated INT8/FP8 Q/DQ only when it meets the configured TensorRT speedup threshold (1.02x by default); otherwise it saves the high-precision no-Q/DQ model.
- Add a Muse Glimmer AutoQuantize recipe that searches language-model MLP projections, self-attention projections, and ``lm_head`` over W4A16 NVFP4 Four-Over-Six, FP8, and BF16 fallback at 5.5 effective bits while leaving the vision tower unquantized.
- Add ``examples/alpamayo/qad.py``, which runs quantization-aware distillation on the quantized Alpamayo checkpoint produced by ``examples/alpamayo/quantize.py``. It distills the quantized VLM against the original FP16 VLM with ``QADTrainer``, supports FSDP2 for multi-GPU runs, and ``--export`` reassembles the trained VLM into a full AlpamayoR1 checkpoint that ``AlpamayoR1.from_pretrained`` can reload.
- Add a calibration-free streaming Kimi-K3 converter and checkpoint-mirror recipe for NVFP4 routed experts with ``input_scale=1.0`` and 128x128 block-FP8 KDA/MLA attention weights. The converter operates shard-by-shard on the source checkpoint's packed MXFP4 experts instead of loading the 2.8T model through the in-memory ``hf_ptq.py`` path.
+2
View File
@@ -63,6 +63,8 @@ The command will:
Autotune searches for Q/DQ placement schemes that improve TensorRT runtime. It does not by itself define the full calibration and quantization policy for an accuracy-sensitive deployment. For end-to-end ONNX PTQ that starts from an unquantized model, run ONNX quantization with calibration data and enable ``--autotune`` there. See the `ONNX quantization Autotune options <_onnx_quantization.html#python-m-modelopt.onnx.quantization-autotune-only-applicable-when-autotune-is-set>`_.
End-to-end ONNX quantization benchmarks Autotune candidates in the requested runtime precision, then validates the exact calibrated output. It retains Q/DQ only when the output meets ``Config.performance_threshold`` (``1.02`` by default); otherwise the requested output path contains the high-precision no-Q/DQ model. The log reports which outcome was selected. Standalone Autotune remains an uncalibrated placement search.
**Output Files:**
Files are written under the output directory (default ``./autotuner_output``, or the path given by ``--output_dir``):
+2
View File
@@ -369,6 +369,8 @@ def get_parser() -> argparse.ArgumentParser:
choices=["quick", "default", "extensive"],
help=(
"If set, enable Autotune to detect optimal Q/DQ node placements according to TensorRT runtimes. "
"Candidates are benchmarked in the requested runtime precision, and calibrated Q/DQ is retained only "
"when it meets the 1.02x performance threshold; otherwise the output has no Q/DQ. "
"Available modes (presets 'schemes_per_region', 'warmup_runs', and 'timing_runs' values): "
" - 'quick': fewer schemes and benchmark runs for quick exploration; "
" - 'default': balanced, recommended for most cases; "
@@ -28,6 +28,7 @@ import dataclasses
import functools
import os
import random
from collections.abc import Callable
from datetime import datetime, timezone
import onnx
@@ -528,7 +529,10 @@ class QDQAutotunerBase:
quantized_node_indices.add(producer_idx)
nodes_to_quantize = [graph.nodes[i].name for i in quantized_node_indices]
op_types_to_quantize = list(get_autotuner_quantizable_ops())
op_types_to_quantize = get_autotuner_quantizable_ops()
if self.config.default_quant_type == "fp8":
op_types_to_quantize &= {"Conv", "Gemm", "MatMul", "Add"}
op_types_to_quantize = list(op_types_to_quantize)
# Inputs of quantized nodes NOT covered by Q/DQ (only non-constant producer inputs)
no_quantize_inputs: list[tuple[gs.Node, gs.Node, str]] = []
@@ -558,7 +562,11 @@ class QDQAutotunerBase:
@_requires_init
def export_onnx(
self, output_path: str | None = None, insert_qdq: bool = True, best: bool = False
self,
output_path: str | None = None,
insert_qdq: bool = True,
best: bool = False,
model_transform: Callable[[onnx.ModelProto], onnx.ModelProto] | None = None,
) -> bytes:
"""Export ONNX model with Q/DQ nodes inserted according to tested schemes.
@@ -606,6 +614,7 @@ class QDQAutotunerBase:
self.config,
insert_qdq=insert_qdq and bool(resolved_insertion_points),
needs_fp8_conversion=needs_fp8_conversion,
model_transform=model_transform,
)
model_bytes = model.SerializeToString()
@@ -16,6 +16,7 @@
"""Utilities for Q/DQ model export and insertion in ONNX autotune."""
import dataclasses
from collections.abc import Callable
import numpy as np
import onnx
@@ -318,6 +319,7 @@ def export_qdq_onnx(
*,
insert_qdq: bool = True,
needs_fp8_conversion: bool = False,
model_transform: Callable[[onnx.ModelProto], onnx.ModelProto] | None = None,
) -> onnx.ModelProto:
"""Export ONNX model with optional Q/DQ insertion and optional INT8→FP8 conversion.
@@ -329,6 +331,7 @@ def export_qdq_onnx(
config: Config for Q/DQ parameters and dtypes.
insert_qdq: If True, insert Q/DQ at resolved points before exporting.
needs_fp8_conversion: If True, build as INT8 then convert to FP8 (e.g. when config.default_quant_type is fp8).
model_transform: Optional transform applied before INT8-to-FP8 conversion.
Returns:
Exported ONNX ModelProto (with Q/DQ and/or FP8 as requested).
@@ -353,6 +356,9 @@ def export_qdq_onnx(
if insert_qdq and resolved_insertion_points:
fix_zero_point_initializers(model)
if model_transform is not None:
model = model_transform(model)
if needs_fp8_conversion:
logger.debug("Converting INT8 to FP8")
model = int8_to_fp8(model)
@@ -22,6 +22,7 @@ optimization of ONNX models using pattern-based region analysis and TensorRT per
import fnmatch
import shutil
import tempfile
from collections.abc import Callable
from pathlib import Path
import onnx
@@ -170,6 +171,7 @@ def region_pattern_autotuning_workflow(
qdq_baseline_model: str | None = None,
node_filter_list: list[str] | None = None,
verbose: bool = False,
model_transform: Callable[[onnx.ModelProto], onnx.ModelProto] | None = None,
) -> QDQAutotuner:
"""Run automated Q/DQ (Quantization/Dequantization) optimization on an ONNX model.
@@ -214,6 +216,7 @@ def region_pattern_autotuning_workflow(
node_filter_list: Optional list of wildcard patterns to filter ONNX nodes. Regions
without any matching nodes are skipped during autotuning (default: None)
verbose: Enable verbose logging in Config for detailed autotuner output (default: False)
model_transform: Optional transform applied to every model before benchmarking.
Returns:
QDQAutotuner instance after autotuning
@@ -265,6 +268,18 @@ def region_pattern_autotuning_workflow(
if state_path.exists():
logger.info(f"Resuming from checkpoint: {state_path}")
autotuner.load_state(str(state_path))
if model_transform is not None:
for pattern_schemes in autotuner.profiled_patterns:
for scheme in pattern_schemes.schemes:
scheme.latency_ms = float("inf")
scheme.error = False
scheme.profile_timestamp = None
if autotuner.pattern_cache is not None:
autotuner.pattern_cache.add_pattern_schemes(pattern_schemes)
autotuner.profiled_patterns.clear()
autotuner.baseline_latency_ms = None
autotuner.config = config
logger.info("Re-profiling checkpoint measurements for transformed models")
else:
logger.info("Starting new autotuning session")
@@ -286,7 +301,7 @@ def region_pattern_autotuning_workflow(
if autotuner.baseline_latency_ms is None:
logger.info("Measuring baseline (no Q/DQ)")
baseline_path = output_dir / "baseline.onnx"
autotuner.export_onnx(str(baseline_path), insert_qdq=False)
autotuner.export_onnx(str(baseline_path), insert_qdq=False, model_transform=model_transform)
baseline_log = logs_dir / "baseline.log"
baseline_latency = benchmark_onnx_model(str(baseline_path), str(baseline_log))
autotuner.submit(baseline_latency)
@@ -327,7 +342,9 @@ def region_pattern_autotuning_workflow(
break
schemes_tested += 1
model_bytes = autotuner.export_onnx(None, insert_qdq=True)
model_bytes = autotuner.export_onnx(
None, insert_qdq=True, model_transform=model_transform
)
test_log = logs_dir / f"region_{region.id}_scheme_{scheme_idx}.log"
flush_timing_cache = (iteration_count % 10) == 0
latency = benchmark_onnx_model(
@@ -351,7 +368,9 @@ def region_pattern_autotuning_workflow(
logger.info(f" Tested {schemes_tested} schemes")
region_model_path = models_dir / f"region_{region.id}_level_{region.level}.onnx"
autotuner.export_onnx(str(region_model_path), insert_qdq=True, best=True)
autotuner.export_onnx(
str(region_model_path), insert_qdq=True, best=True, model_transform=model_transform
)
logger.debug(f" Saved best model: {region_model_path.name}")
# Save state after each region (incremental, crash recovery)
@@ -363,7 +382,7 @@ def region_pattern_autotuning_workflow(
logger.info("Exporting final optimized model")
final_model_path = output_dir / "optimized_final.onnx"
autotuner.export_onnx(str(final_model_path), insert_qdq=True)
autotuner.export_onnx(str(final_model_path), insert_qdq=True, model_transform=model_transform)
final_log = logs_dir / "final.log"
final_latency = benchmark_onnx_model(str(final_model_path), str(final_log))
+17 -65
View File
@@ -28,23 +28,23 @@ from onnxruntime.quantization import CalibrationMethod
from onnxruntime.quantization.calibrate import CalibrationDataReader
import modelopt.onnx.utils as onnx_utils
from modelopt.onnx.autocast.convert import convert_to_f16
from modelopt.onnx.logging_config import configure_logging, logger
from modelopt.onnx.quantization.graph_utils import (
convert_fp16_io,
expand_node_names_from_patterns,
find_nodes_from_convs_to_exclude,
find_nodes_from_matmul_to_exclude,
find_nodes_to_exclude,
get_concat_eliminated_tensors,
get_tensor_producer_nodes,
insert_fp8_mha_casts,
remove_output_initializers,
remove_partial_input_qdq,
)
from modelopt.onnx.quantization.int8 import _find_nodes_to_quantize
from modelopt.onnx.quantization.ort_patching import _quantize_static as quantize_static
from modelopt.onnx.quantization.ort_utils import configure_ort
from modelopt.onnx.quantization.precision_utils import (
_convert_to_runtime_precision,
_upgrade_opset_21,
)
from modelopt.onnx.quantization.qdq_utils import has_qdq_nodes
@@ -128,36 +128,8 @@ def int8_to_fp8(onnx_model: onnx.ModelProto) -> onnx.ModelProto:
def upgrade_opset_21(onnx_model: onnx.ModelProto) -> onnx.ModelProto:
"""Modifies the ONNX graph such that it follows the opset 21 requirements.
This is necessary for FP8+FP16 quantization since FP8 QuantizeLinear/DequantizeLinear ops do not support FP16
scaling factors until opset 21.
"""
logger.info("Upgrading model to opset 21")
graph = gs.import_onnx(onnx_model)
for node in graph.nodes:
# QuantizeLinear/DequantizeLinear op with FP16 scales are only supported with empty domain
# and opset_import version=21.
if node.op in {"QuantizeLinear", "DequantizeLinear"}:
node.domain = ""
# ReduceMean op no longer has "axes" attribute in opset 21. Instead, it should be the second input tensor.
if node.op == "ReduceMean" and "axes" in node.attrs:
axes = gs.Constant(
name=node.name + "_axes", values=np.array(node.attrs["axes"], dtype=np.int64)
)
del node.attrs["axes"]
node.inputs.append(axes)
onnx_model = gs.export_onnx(graph)
# Set opset_import version to 21.
for opset_import in onnx_model.opset_import:
if opset_import.domain == "":
opset_import.version = 21
return onnx_model
"""Modify an FP8 graph to satisfy opset 21 requirements."""
return _upgrade_opset_21(onnx_model)
def quantize(
@@ -324,37 +296,17 @@ def quantize(
remove_partial_input_qdq(graph, no_quantize_inputs)
onnx_model = gs.export_onnx(graph)
if high_precision_dtype in ["fp16", "bf16"]:
# We need to convert float to float16/bfloat16 so as to speed up layers like LayerNorm or GroupNorm.
logger.info(f"Converting float tensors to {high_precision_dtype}")
graph = gs.import_onnx(onnx_model)
remove_output_initializers(graph, onnx_model.graph.initializer)
convert_fp16_io(graph)
onnx_model = gs.export_onnx(graph)
# Convert to fp16/bf16 model.
onnx_model = convert_to_f16(
onnx_model,
keep_io_types=not direct_io_types,
op_block_list=op_types_to_exclude_fp16 or [],
tensor_block_dict=custom_ops_to_cast_fp32 or {},
low_precision_type=high_precision_dtype,
trt_plugins=trt_extra_plugin_lib_paths,
opset=opset,
)
current_opsets = {opset.domain: opset.version for opset in onnx_model.opset_import}
opset_of_default_onnx_domain = current_opsets.get("", 0)
if opset_of_default_onnx_domain < 19:
# We need to convert the ONNX model to opset 19+ since FP8 QuantizeLinear/DequantizeLinear ops do not
# support FP16 scaling factors until opset 19. So, converting here to opset-21 (19+).
onnx_model = upgrade_opset_21(onnx_model)
if mha_accumulation_dtype == "fp32":
# Insert Cast nodes in MHA's BMM1 and BMM2's input and output tensors because
# The compiler only has FP32 accumulation kernels for FP8 MHAs.
logger.info("Inserting Cast nodes to enable FP8+FP16 MHA")
onnx_model = insert_fp8_mha_casts(onnx_model)
onnx_model = _convert_to_runtime_precision(
onnx_model,
quantize_mode="fp8",
high_precision_dtype=high_precision_dtype,
direct_io_types=direct_io_types,
op_types_to_exclude_fp16=op_types_to_exclude_fp16,
custom_ops_to_cast_fp32=custom_ops_to_cast_fp32,
trt_extra_plugin_lib_paths=trt_extra_plugin_lib_paths,
opset=opset,
mha_accumulation_dtype=mha_accumulation_dtype,
)
if nodes_to_quantize:
onnx_model = int8_to_fp8(onnx_model)
+11 -14
View File
@@ -27,7 +27,6 @@ from onnx_graphsurgeon.ir.node import Node
from onnxruntime.quantization import CalibrationMethod
from onnxruntime.quantization.calibrate import CalibrationDataReader
from modelopt.onnx.autocast.convert import convert_to_f16
from modelopt.onnx.logging_config import configure_logging, logger
from modelopt.onnx.quantization.calib_utils import import_scales_from_calib_cache
from modelopt.onnx.quantization.graph_utils import (
@@ -51,6 +50,7 @@ from modelopt.onnx.quantization.partitioning import (
find_quantizable_nodes,
get_skipped_output_layers,
)
from modelopt.onnx.quantization.precision_utils import _convert_to_runtime_precision
from modelopt.onnx.quantization.qdq_utils import has_qdq_nodes, replace_scale_values
@@ -300,19 +300,16 @@ def quantize(
if calibration_cache_path:
replace_scale_values(onnx_model.graph, act_scales_dict)
if high_precision_dtype in ["fp16", "bf16"]:
# We need to convert float to float16 so as to speed up layers like LayerNorm or GroupNorm.
logger.info(f"Converting float32 tensors to {high_precision_dtype}")
# Note: from convert_to_f16's perspective, high_precision_dtype is the precision to reduce to from FP32
onnx_model = convert_to_f16(
onnx_model,
keep_io_types=not direct_io_types,
op_block_list=op_types_to_exclude_fp16 or [],
tensor_block_dict=custom_ops_to_cast_fp32 or {},
low_precision_type=high_precision_dtype,
trt_plugins=trt_extra_plugin_lib_paths,
opset=opset,
)
onnx_model = _convert_to_runtime_precision(
onnx_model,
quantize_mode="int8",
high_precision_dtype=high_precision_dtype,
direct_io_types=direct_io_types,
op_types_to_exclude_fp16=op_types_to_exclude_fp16,
custom_ops_to_cast_fp32=custom_ops_to_cast_fp32,
trt_extra_plugin_lib_paths=trt_extra_plugin_lib_paths,
opset=opset,
)
if nodes_to_quantize:
logger.info(f"Quantization completed successfully in {time.time() - t_start} seconds")
@@ -0,0 +1,93 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Shared runtime-precision conversion for ONNX quantization."""
import numpy as np
import onnx
import onnx_graphsurgeon as gs
from modelopt.onnx.autocast.convert import convert_to_f16
from modelopt.onnx.logging_config import logger
from modelopt.onnx.quantization.graph_utils import (
convert_fp16_io,
insert_fp8_mha_casts,
remove_output_initializers,
)
def _upgrade_opset_21(model: onnx.ModelProto) -> onnx.ModelProto:
logger.info("Upgrading model to opset 21")
graph = gs.import_onnx(model)
for node in graph.nodes:
if node.op in {"QuantizeLinear", "DequantizeLinear"}:
node.domain = ""
if node.op == "ReduceMean" and "axes" in node.attrs:
axes = gs.Constant(
name=node.name + "_axes", values=np.array(node.attrs.pop("axes"), dtype=np.int64)
)
node.inputs.append(axes)
model = gs.export_onnx(graph)
for opset_import in model.opset_import:
if opset_import.domain == "":
opset_import.version = 21
return model
def _convert_to_runtime_precision(
model: onnx.ModelProto,
*,
quantize_mode: str,
high_precision_dtype: str,
direct_io_types: bool = False,
op_types_to_exclude_fp16: list[str] | None = None,
custom_ops_to_cast_fp32: dict | None = None,
trt_extra_plugin_lib_paths: list[str] | None = None,
opset: int | None = None,
mha_accumulation_dtype: str = "fp16",
) -> onnx.ModelProto:
"""Convert a quantized model to the precision used at runtime."""
if high_precision_dtype not in {"fp16", "bf16"}:
return model
logger.info(f"Converting float tensors to {high_precision_dtype}")
if quantize_mode == "fp8":
graph = gs.import_onnx(model)
remove_output_initializers(graph, model.graph.initializer)
convert_fp16_io(graph)
model = gs.export_onnx(graph)
model = convert_to_f16(
model,
keep_io_types=not direct_io_types,
op_block_list=op_types_to_exclude_fp16 or [],
tensor_block_dict=custom_ops_to_cast_fp32 or {},
low_precision_type=high_precision_dtype,
trt_plugins=trt_extra_plugin_lib_paths,
opset=opset,
)
if quantize_mode != "fp8":
return model
current_opsets = {opset.domain: opset.version for opset in model.opset_import}
if current_opsets.get("", 0) < 19:
model = _upgrade_opset_21(model)
if mha_accumulation_dtype == "fp32":
logger.info("Inserting Cast nodes to enable FP8+FP16 MHA")
model = insert_fp8_mha_casts(model)
return model
+180 -34
View File
@@ -31,11 +31,14 @@ This tool inserts Quantize-Dequantize (QDQ) nodes following compiler-friendly pa
model.
"""
import copy
import math
import os
import platform
import shutil
import tempfile
from collections.abc import Sequence
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Any
@@ -64,6 +67,7 @@ from modelopt.onnx.quantization.graph_utils import (
from modelopt.onnx.quantization.int4 import quantize as quantize_int4
from modelopt.onnx.quantization.int8 import quantize as quantize_int8
from modelopt.onnx.quantization.ort_utils import create_input_shapes_profile, update_trt_ep_support
from modelopt.onnx.quantization.precision_utils import _convert_to_runtime_precision
from modelopt.onnx.quantization.qdq_utils import (
qdq_to_dq,
remove_graph_input_q,
@@ -83,6 +87,35 @@ from modelopt.onnx.utils import (
__all__ = ["quantize"]
@dataclass
class _AutotuneContext:
ort_config: tuple[list[str], list[str], list[tuple[gs.Node, gs.Node, str]], list[str]]
baseline_model: onnx.ModelProto
performance_threshold: float
output_dir: Path
temporary_output_dir: tempfile.TemporaryDirectory | None = None
def cleanup(self) -> None:
if self.temporary_output_dir is not None:
self.temporary_output_dir.cleanup()
self.temporary_output_dir = None
def _run_with_autotune_cleanup(
context: _AutotuneContext | None, function: Callable[..., Any], *args: Any, **kwargs: Any
) -> Any:
try:
return function(*args, **kwargs)
except BaseException:
if context is not None:
context.cleanup()
raise
def _has_qdq_site(model: onnx.ModelProto) -> bool:
return any(node.op_type in {"QuantizeLinear", "DequantizeLinear"} for node in model.graph.node)
def _normalize_quantize_mode_for_opset(quantize_mode: str) -> str:
"""Map variants like "int4_awq", "int4_rtn", "nvfp4" to their base precision types for lookup purposes."""
mode_lower = quantize_mode.lower()
@@ -290,6 +323,11 @@ def _find_nodes_to_quantize_autotune(
quantize_mode: str,
trt_plugins: list[str] | None,
high_precision_dtype: str = "fp16",
direct_io_types: bool = False,
op_types_to_exclude_fp16: list[str] | None = None,
custom_ops_to_cast_fp32: dict | None = None,
opset: int | None = None,
mha_accumulation_dtype: str = "fp16",
output_dir: str | None = None,
num_schemes_per_region: int = 50,
pattern_cache_file: str | None = None,
@@ -302,7 +340,7 @@ def _find_nodes_to_quantize_autotune(
warmup_runs: int = 50,
timing_runs: int = 100,
trtexec_args: str | None = None,
) -> tuple[list[str], list[str], list[tuple[gs.Node, gs.Node, str]], list[str]]:
) -> _AutotuneContext:
"""Extracts quantization information from Autotune to provide ORT quantization."""
logger.info("Running Auto Q/DQ with TensorRT")
@@ -329,20 +367,101 @@ def _find_nodes_to_quantize_autotune(
if benchmark_instance is None:
raise RuntimeError("Failed to initialize TensorRT benchmark")
temporary_output_dir = None
if output_dir is None:
temporary_output_dir = tempfile.TemporaryDirectory(prefix="modelopt_autotune_")
resolved_output_dir = Path(temporary_output_dir.name)
else:
resolved_output_dir = Path(output_dir)
def model_transform(model: onnx.ModelProto) -> onnx.ModelProto:
return _convert_to_runtime_precision(
model,
quantize_mode=quantize_mode,
high_precision_dtype=high_precision_dtype,
direct_io_types=direct_io_types,
op_types_to_exclude_fp16=op_types_to_exclude_fp16,
custom_ops_to_cast_fp32=custom_ops_to_cast_fp32,
trt_extra_plugin_lib_paths=trt_plugins,
opset=opset,
mha_accumulation_dtype=mha_accumulation_dtype,
)
precision_map = {"fp16": "float16", "fp32": "float32", "bf16": "bfloat16"}
autotuner = region_pattern_autotuning_workflow(
onnx_model,
output_dir=Path(output_dir) if output_dir else None,
num_schemes_per_region=num_schemes_per_region,
pattern_cache_file=pattern_cache_file,
state_file=state_file,
quant_type=quantize_mode,
default_dq_dtype=precision_map[high_precision_dtype],
qdq_baseline_model=qdq_baseline_model,
node_filter_list=node_filter_list,
verbose=verbose,
try:
autotuner = region_pattern_autotuning_workflow(
onnx_model,
output_dir=resolved_output_dir,
num_schemes_per_region=num_schemes_per_region,
pattern_cache_file=pattern_cache_file,
state_file=state_file,
quant_type=quantize_mode,
default_dq_dtype=precision_map[high_precision_dtype],
qdq_baseline_model=qdq_baseline_model,
node_filter_list=node_filter_list,
verbose=verbose,
model_transform=model_transform,
)
return _AutotuneContext(
ort_config=autotuner.get_ort_quantization_config(),
baseline_model=model_transform(copy.deepcopy(onnx_model)),
performance_threshold=autotuner.config.performance_threshold,
output_dir=resolved_output_dir,
temporary_output_dir=temporary_output_dir,
)
except BaseException:
if temporary_output_dir is not None:
temporary_output_dir.cleanup()
raise
def _apply_autotune_final_guard(
candidate_model: onnx.ModelProto,
context: _AutotuneContext,
*,
use_external_data_format: bool,
) -> onnx.ModelProto:
"""Keep the calibrated artifact only when it beats its no-Q/DQ reference."""
# Keep Autotune's Torch and TensorRT imports off the ordinary quantization path.
from modelopt.onnx.quantization.autotune.workflows import benchmark_onnx_model
logs_dir = context.output_dir / "logs"
logs_dir.mkdir(parents=True, exist_ok=True)
baseline_path = context.output_dir / "calibrated_baseline.onnx"
candidate_path = context.output_dir / "calibrated_candidate.onnx"
baseline_model = copy.deepcopy(context.baseline_model)
candidate_model = copy.deepcopy(candidate_model)
save_onnx(baseline_model, str(baseline_path), use_external_data_format)
save_onnx(candidate_model, str(candidate_path), use_external_data_format)
baseline_latency = benchmark_onnx_model(
str(baseline_path), str(logs_dir / "calibrated_baseline.log")
)
return autotuner.get_ort_quantization_config()
if not math.isfinite(baseline_latency) or baseline_latency <= 0:
raise RuntimeError(
"Autotune could not establish a finite positive latency for the "
"high-precision no-Q/DQ reference"
)
candidate_latency = benchmark_onnx_model(
str(candidate_path), str(logs_dir / "calibrated_candidate.log")
)
candidate_is_valid = math.isfinite(candidate_latency) and candidate_latency > 0
speedup = baseline_latency / candidate_latency if candidate_is_valid else 0.0
keep_candidate = _has_qdq_site(candidate_model) and speedup >= context.performance_threshold
selection = "qdq" if keep_candidate else "no_qdq"
logger.info(
"Autotune selected %s: baseline=%.3f ms, candidate=%.3f ms, speedup=%.3fx, threshold=%.3fx",
selection,
baseline_latency,
candidate_latency,
speedup,
context.performance_threshold,
)
if not keep_candidate:
logger.warning("The saved Autotune artifact is not quantized.")
return candidate_model if keep_candidate else baseline_model
def quantize(
@@ -630,6 +749,9 @@ def quantize(
quantize_mode,
opset,
)
if autotune and _has_qdq_site(onnx_model):
raise ValueError("Autotune requires an unquantized source model without Q/DQ nodes")
original_calibration_eps = list(calibration_eps)
trt_plugins = update_trt_ep_support(calibration_eps, has_dds_op, has_custom_op, trt_plugins) # type: ignore[arg-type]
@@ -678,18 +800,19 @@ def quantize(
if calibrate_per_node and not calibration_shapes:
calibration_shapes = get_input_shapes(onnx_path)
autotune_context = None
if quantize_mode in ["fp8", "int8"]:
if autotune:
(
nodes_to_quantize_autotune,
op_types_to_quantize_autotune,
no_quantize_inputs,
op_types_needing_output_quant,
) = _find_nodes_to_quantize_autotune(
autotune_context = _find_nodes_to_quantize_autotune(
onnx_model,
quantize_mode,
trt_plugins,
high_precision_dtype,
direct_io_types=direct_io_types,
op_types_to_exclude_fp16=op_types_to_exclude_fp16,
custom_ops_to_cast_fp32=custom_ops_to_cast_fp32,
opset=opset,
mha_accumulation_dtype=mha_accumulation_dtype,
output_dir=autotune_output_dir,
num_schemes_per_region=autotune_num_schemes_per_region,
pattern_cache_file=autotune_pattern_cache_file,
@@ -703,6 +826,12 @@ def quantize(
timing_runs=autotune_timing_runs,
trtexec_args=autotune_trtexec_args,
)
(
nodes_to_quantize_autotune,
op_types_to_quantize_autotune,
no_quantize_inputs,
op_types_needing_output_quant,
) = autotune_context.ort_config
op_types_to_quantize = op_types_to_quantize or op_types_to_quantize_autotune
nodes_to_quantize = nodes_to_quantize or nodes_to_quantize_autotune
kwargs["no_quantize_inputs"] = no_quantize_inputs
@@ -710,7 +839,9 @@ def quantize(
kwargs["target_dla"] = target_dla
quantize_func = quantize_int8 if quantize_mode == "int8" else quantize_fp8
onnx_model = quantize_func(
onnx_model = _run_with_autotune_cleanup(
autotune_context,
quantize_func,
onnx_path=onnx_path,
calibration_method=calibration_method or "entropy",
calibration_data_reader=calibration_data_reader,
@@ -757,28 +888,43 @@ def quantize(
raise RuntimeError(f"Invalid quantization mode choice: {quantize_mode}")
if onnx_model:
# Fuse Q nodes for INT8/FP8 mode
if quantize_mode in ["int8", "fp8"]:
if dq_only:
onnx_model = qdq_to_dq(onnx_model)
onnx_model = _run_with_autotune_cleanup(autotune_context, qdq_to_dq, onnx_model)
if custom_ops_to_quantize:
# Remove DQ nodes from the input and Q from the output of the requested custom ops
onnx_model = remove_input_dq_and_output_q(
onnx_model, quantizable_custom_ops=custom_ops_to_quantize
onnx_model = _run_with_autotune_cleanup(
autotune_context,
remove_input_dq_and_output_q,
onnx_model,
quantizable_custom_ops=custom_ops_to_quantize,
)
if direct_io_types:
onnx_model = remove_graph_input_q(onnx_model)
onnx_model = _run_with_autotune_cleanup(
autotune_context, remove_graph_input_q, onnx_model
)
if autotune_context is not None:
onnx_model = _run_with_autotune_cleanup(
autotune_context,
_apply_autotune_final_guard,
onnx_model,
autotune_context,
use_external_data_format=use_external_data_format,
)
else:
# Remove redundant cast nodes in the quantized model
# Note. This is called within the qdq_to_dq function as well
remove_redundant_cast_nodes(onnx_model.graph)
# Collect and print stats of the quantized model
print_stat(gs.import_onnx(onnx_model))
graph = _run_with_autotune_cleanup(autotune_context, gs.import_onnx, onnx_model)
_run_with_autotune_cleanup(autotune_context, print_stat, graph)
_run_with_autotune_cleanup(
autotune_context, save_onnx, onnx_model, output_path, use_external_data_format
)
if autotune_context is not None and not _has_qdq_site(onnx_model):
logger.info(f"Autotune high-precision ONNX model is saved as {output_path}")
else:
logger.info(f"Quantized onnx model is saved as {output_path}")
# Save the quantized model to the output path
save_onnx(onnx_model, output_path, use_external_data_format)
logger.info(f"Quantized onnx model is saved as {output_path}")
if autotune_context is not None:
autotune_context.cleanup()
# Check if intermediate files should be deleted
if not keep_intermediate_files:
@@ -0,0 +1,195 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from types import SimpleNamespace
from unittest.mock import Mock
import onnx
import pytest
import modelopt.onnx.quantization.autotune.export_utils as export_utils
import modelopt.onnx.quantization.precision_utils as precision_utils
from modelopt.onnx.quantization.autotune import workflows
from modelopt.onnx.quantization.autotune.common import Config
@pytest.mark.parametrize(
("quantize_mode", "expected_events"),
[
("int8", ["convert"]),
("fp8", ["import", "remove_outputs", "convert_io", "export", "convert", "upgrade", "mha"]),
],
)
@pytest.mark.parametrize("direct_io_types", [False, True])
def test_runtime_precision_conversion_preserves_mode_steps(
monkeypatch, quantize_mode, expected_events, direct_io_types
):
source = onnx.ModelProto()
graph = object()
io_model = onnx.ModelProto()
converted = onnx.ModelProto()
converted.opset_import.add(domain="", version=17)
final = onnx.ModelProto()
events = []
monkeypatch.setattr(
precision_utils.gs, "import_onnx", lambda model: events.append("import") or graph
)
monkeypatch.setattr(
precision_utils,
"remove_output_initializers",
lambda graph, initializers: events.append("remove_outputs"),
)
monkeypatch.setattr(
precision_utils, "convert_fp16_io", lambda graph: events.append("convert_io")
)
monkeypatch.setattr(
precision_utils.gs, "export_onnx", lambda graph: events.append("export") or io_model
)
def convert(model, **kwargs):
events.append("convert")
assert model is (io_model if quantize_mode == "fp8" else source)
assert kwargs == {
"keep_io_types": not direct_io_types,
"op_block_list": ["Resize"],
"tensor_block_dict": {"Custom": {"inputs": [0]}},
"low_precision_type": "fp16",
"trt_plugins": ["plugin.so"],
"opset": 17,
}
return converted
monkeypatch.setattr(precision_utils, "convert_to_f16", convert)
monkeypatch.setattr(
precision_utils,
"_upgrade_opset_21",
lambda model: events.append("upgrade") or converted,
)
monkeypatch.setattr(
precision_utils,
"insert_fp8_mha_casts",
lambda model: events.append("mha") or final,
)
result = precision_utils._convert_to_runtime_precision(
source,
quantize_mode=quantize_mode,
high_precision_dtype="fp16",
direct_io_types=direct_io_types,
op_types_to_exclude_fp16=["Resize"],
custom_ops_to_cast_fp32={"Custom": {"inputs": [0]}},
trt_extra_plugin_lib_paths=["plugin.so"],
opset=17,
mha_accumulation_dtype="fp32",
)
assert events == expected_events
assert result is (final if quantize_mode == "fp8" else converted)
def test_export_transform_runs_between_int8_qdq_and_fp8(monkeypatch):
source = onnx.ModelProto()
source_bytes = source.SerializeToString()
graph = type("Graph", (), {"toposort": lambda self: None})()
int8_model, transformed, fp8_model = (onnx.ModelProto() for _ in range(3))
events = []
monkeypatch.setattr(export_utils.gs, "import_onnx", lambda model: graph)
monkeypatch.setattr(
export_utils.gs, "export_onnx", lambda graph: events.append("export") or int8_model
)
monkeypatch.setattr(
export_utils,
"insert_qdq_at_tensors",
lambda graph, points, config: events.append(f"insert_{config.default_quant_type}"),
)
monkeypatch.setattr(
export_utils,
"fix_zero_point_initializers",
lambda model: events.append("fix_zero_point"),
)
monkeypatch.setattr(
export_utils,
"int8_to_fp8",
lambda model: events.append("convert_fp8") or fp8_model,
)
def transform(model):
assert model is int8_model
events.append("transform")
return transformed
result = export_utils.export_qdq_onnx(
source,
{object()},
Config(default_quant_type="fp8"),
needs_fp8_conversion=True,
model_transform=transform,
)
assert events == ["insert_int8", "export", "fix_zero_point", "transform", "convert_fp8"]
assert result is fp8_model
assert source.SerializeToString() == source_bytes
def test_workflow_transforms_every_benchmark_export(monkeypatch, tmp_path):
stale_scheme = SimpleNamespace(latency_ms=0.5, error=True, profile_timestamp="old")
stale_pattern = SimpleNamespace(schemes=[stale_scheme])
autotuner = Mock(
regions=[SimpleNamespace(id=0, level=0)],
baseline_latency_ms=None,
current_profile_pattern_schemes=SimpleNamespace(schemes=[]),
)
autotuner.generate.return_value = 0
autotuner.export_onnx.return_value = onnx.ModelProto().SerializeToString()
def load_state(_):
autotuner.baseline_latency_ms = 0.5
autotuner.profiled_patterns = [stale_pattern]
autotuner.config = Config(default_quant_type="int8")
autotuner.load_state.side_effect = load_state
monkeypatch.setattr(workflows, "QDQAutotuner", lambda model: autotuner)
monkeypatch.setattr(workflows, "benchmark_onnx_model", lambda *args, **kwargs: 1.0)
def transform(model):
return model
state_path = tmp_path / "state.yaml"
state_path.touch()
workflows.region_pattern_autotuning_workflow(
onnx.ModelProto(),
output_dir=tmp_path,
state_file=str(state_path),
num_schemes_per_region=1,
quant_type="fp8",
model_transform=transform,
)
exports = autotuner.export_onnx.call_args_list
assert [(call.kwargs["insert_qdq"], call.kwargs.get("best", False)) for call in exports] == [
(False, False),
(True, False),
(True, True),
(True, False),
]
assert all(call.kwargs["model_transform"] is transform for call in exports)
assert stale_scheme.latency_ms == float("inf")
assert not stale_scheme.error
assert stale_scheme.profile_timestamp is None
autotuner.pattern_cache.add_pattern_schemes.assert_called_once_with(stale_pattern)
assert autotuner.profiled_patterns == []
assert autotuner.config.default_quant_type == "fp8"
@@ -15,17 +15,23 @@
"""Tests for ONNX quantization API handling."""
import copy
import importlib
import os
import tempfile
import onnx
import onnxruntime
import pytest
import torch
from _test_utils.onnx.lib_test_models import SimpleMLP, export_as_onnx
from onnx import TensorProto, helper
from packaging import version
import modelopt.onnx.quantization as moq
import modelopt.onnx.trt_utils as trt_utils
from modelopt.onnx.quantization.autotune import Config, QDQAutotuner
from modelopt.onnx.quantization.autotune.insertion_points import get_autotuner_quantizable_ops
from modelopt.onnx.utils import get_opset_version
# Mapping of quantization mode to minimum required opset
@@ -39,6 +45,36 @@ MIN_OPSET = {
ORT_VERSION_FOR_OPSET_22 = version.parse("1.23.0")
def _make_guard_models(site_ops=("QuantizeLinear", "DequantizeLinear")):
graph_input = helper.make_tensor_value_info("input", TensorProto.FLOAT, [1, 4])
graph_output = helper.make_tensor_value_info("output", TensorProto.FLOAT, [1, 4])
baseline = helper.make_model(
helper.make_graph(
[helper.make_node("Identity", ["input"], ["output"])],
"baseline",
[graph_input],
[graph_output],
),
opset_imports=[helper.make_opsetid("", 19)],
)
candidate = copy.deepcopy(baseline)
candidate.graph.node.extend(
helper.make_node(op_type, ["input"], [f"site_{index}"])
for index, op_type in enumerate(site_ops)
)
return baseline, candidate
def _make_guard_context(quantize_module, tmp_path, baseline):
tmp_path.mkdir(exist_ok=True)
return quantize_module._AutotuneContext(
ort_config=([], [], [], []),
baseline_model=baseline,
performance_threshold=1.02,
output_dir=tmp_path,
)
# Test scenarios: (scenario_name, export_opset_offset, request_opset_offset, expected_opset_offset)
# Offsets are relative to MIN_OPSET[quant_mode].
OPSET_SCENARIOS = [
@@ -74,6 +110,158 @@ def test_realign_input_shapes_profile_rejects_duplicate_calibration_eps():
)
@pytest.mark.parametrize("quant_type", ["int8", "fp8"])
def test_autotune_ort_op_types_match_quantization_mode(quant_type):
model, _ = _make_guard_models(())
autotuner = QDQAutotuner(model)
autotuner.initialize(Config(default_quant_type=quant_type))
expected = get_autotuner_quantizable_ops()
if quant_type == "fp8":
expected &= {"Conv", "Gemm", "MatMul", "Add"}
op_types = autotuner.get_ort_quantization_config()[1]
assert isinstance(op_types, list)
assert set(op_types) == expected
@pytest.mark.parametrize(
("candidate_latency", "site_ops", "expect_qdq"),
[
(101.0, ("QuantizeLinear", "DequantizeLinear"), False),
(100.0, ("QuantizeLinear", "DequantizeLinear"), True),
(99.0, ("QuantizeLinear", "DequantizeLinear"), True),
(90.0, ("QuantizeLinear",), True),
(90.0, ("DequantizeLinear",), True),
(90.0, (), False),
],
)
def test_autotune_final_guard_selection(
monkeypatch, tmp_path, candidate_latency, site_ops, expect_qdq
):
quantize_module = importlib.import_module("modelopt.onnx.quantization.quantize")
workflows = importlib.import_module("modelopt.onnx.quantization.autotune.workflows")
latencies = iter([102.0, candidate_latency])
monkeypatch.setattr(workflows, "benchmark_onnx_model", lambda *args: next(latencies))
baseline, candidate = _make_guard_models(site_ops)
selected = quantize_module._apply_autotune_final_guard(
candidate,
_make_guard_context(quantize_module, tmp_path, baseline),
use_external_data_format=False,
)
assert quantize_module._has_qdq_site(selected) is expect_qdq
@pytest.mark.parametrize("invalid_latency", [float("nan"), float("inf"), float("-inf"), 0, -1])
@pytest.mark.parametrize("invalid_model", ["baseline", "candidate"])
def test_autotune_final_guard_invalid_measurements(
monkeypatch, tmp_path, invalid_latency, invalid_model
):
quantize_module = importlib.import_module("modelopt.onnx.quantization.quantize")
workflows = importlib.import_module("modelopt.onnx.quantization.autotune.workflows")
latencies = [invalid_latency] if invalid_model == "baseline" else [102.0, invalid_latency]
monkeypatch.setattr(workflows, "benchmark_onnx_model", lambda *args: latencies.pop(0))
baseline, candidate = _make_guard_models()
context = _make_guard_context(quantize_module, tmp_path, baseline)
if invalid_model == "baseline":
with pytest.raises(RuntimeError, match="finite positive latency"):
quantize_module._apply_autotune_final_guard(
candidate, context, use_external_data_format=False
)
else:
selected = quantize_module._apply_autotune_final_guard(
candidate, context, use_external_data_format=False
)
assert not quantize_module._has_qdq_site(selected)
def test_autotune_rejects_prequantized_source(monkeypatch, tmp_path):
quantize_module = importlib.import_module("modelopt.onnx.quantization.quantize")
onnx_path = tmp_path / "model.onnx"
onnx_path.write_bytes(b"")
_, qdq_model = _make_guard_models()
preprocessed = (str(onnx_path), qdq_model, [], False, False, False, {}, {})
monkeypatch.setattr(quantize_module, "_preprocess_onnx", lambda *args, **kwargs: preprocessed)
with pytest.raises(ValueError, match="unquantized source model"):
quantize_module.quantize(
str(onnx_path),
output_path=str(tmp_path / "output.onnx"),
quantize_mode="fp8",
calibration_data_reader=object(),
calibration_eps=["cpu"],
autotune=True,
)
def test_autotune_tempdir_is_cleaned_after_failure(tmp_path):
quantize_module = importlib.import_module("modelopt.onnx.quantization.quantize")
temporary_output_dir = tempfile.TemporaryDirectory(dir=tmp_path)
temporary_path = temporary_output_dir.name
baseline, _ = _make_guard_models()
context = _make_guard_context(quantize_module, tmp_path, baseline)
context.temporary_output_dir = temporary_output_dir
with pytest.raises(RuntimeError, match="failed"):
quantize_module._run_with_autotune_cleanup(
context, lambda: (_ for _ in ()).throw(RuntimeError("failed"))
)
assert not os.path.exists(temporary_path)
def test_fp8_autotune_subthreshold_result_uses_precision_matched_fallback(monkeypatch, tmp_path):
monkeypatch.setattr(trt_utils, "TRT_PYTHON_AVAILABLE", False)
quantize_module = importlib.import_module("modelopt.onnx.quantization.quantize")
workflows = importlib.import_module("modelopt.onnx.quantization.autotune.workflows")
onnx_path = tmp_path / "model.onnx"
output_path = tmp_path / "autotuned.onnx"
autotune_dir = tmp_path / "autotune"
autotune_dir.mkdir()
export_as_onnx(SimpleMLP(), torch.randn(2, 16, 16), onnx_filename=str(onnx_path), opset=19)
def fake_find_nodes(model, quantize_mode, trt_plugins, high_precision_dtype, **kwargs):
nodes = [node.name for node in model.graph.node if node.op_type in {"Gemm", "MatMul"}]
baseline = quantize_module._convert_to_runtime_precision(
copy.deepcopy(model),
quantize_mode=quantize_mode,
high_precision_dtype=high_precision_dtype,
direct_io_types=kwargs["direct_io_types"],
op_types_to_exclude_fp16=kwargs["op_types_to_exclude_fp16"],
custom_ops_to_cast_fp32=kwargs["custom_ops_to_cast_fp32"],
trt_extra_plugin_lib_paths=trt_plugins,
opset=kwargs["opset"],
mha_accumulation_dtype=kwargs["mha_accumulation_dtype"],
)
return quantize_module._AutotuneContext(
ort_config=(nodes, ["Gemm", "MatMul"], [], []),
baseline_model=baseline,
performance_threshold=1.02,
output_dir=autotune_dir,
)
latencies = iter([100.0, 99.0])
monkeypatch.setattr(quantize_module, "_find_nodes_to_quantize_autotune", fake_find_nodes)
monkeypatch.setattr(workflows, "benchmark_onnx_model", lambda *args: next(latencies))
moq.quantize(
str(onnx_path),
output_path=str(output_path),
quantize_mode="fp8",
calibration_eps=["cpu"],
autotune=True,
autotune_output_dir=str(autotune_dir),
)
candidate = onnx.load(autotune_dir / "calibrated_candidate.onnx")
assert quantize_module._has_qdq_site(candidate)
selected = onnx.load(output_path)
assert not quantize_module._has_qdq_site(selected)
assert selected.graph.input[0].type.tensor_type.elem_type == TensorProto.FLOAT
assert selected.graph.output[0].type.tensor_type.elem_type == TensorProto.FLOAT
def test_quantize_infers_input_profiles_after_ep_support_update(monkeypatch, tmp_path):
quantize_module = importlib.import_module("modelopt.onnx.quantization.quantize")
onnx_path = tmp_path / "model.onnx"