mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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.
|
||||
|
||||
@@ -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``):
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user