nvidia-modelopt 0.35.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- modelopt/__init__.py +20 -0
- modelopt/deploy/__init__.py +16 -0
- modelopt/deploy/llm/__init__.py +30 -0
- modelopt/deploy/llm/generate.py +345 -0
- modelopt/onnx/__init__.py +34 -0
- modelopt/onnx/autocast/__init__.py +19 -0
- modelopt/onnx/autocast/__main__.py +187 -0
- modelopt/onnx/autocast/convert.py +202 -0
- modelopt/onnx/autocast/graphsanitizer.py +422 -0
- modelopt/onnx/autocast/logging_config.py +78 -0
- modelopt/onnx/autocast/nodeclassifier.py +416 -0
- modelopt/onnx/autocast/precisionconverter.py +1032 -0
- modelopt/onnx/autocast/referencerunner.py +159 -0
- modelopt/onnx/autocast/utils.py +117 -0
- modelopt/onnx/llm_export_utils/__init__.py +16 -0
- modelopt/onnx/llm_export_utils/export_utils.py +162 -0
- modelopt/onnx/llm_export_utils/quantization_utils.py +122 -0
- modelopt/onnx/llm_export_utils/surgeon_utils.py +120 -0
- modelopt/onnx/logging_config.py +77 -0
- modelopt/onnx/op_types.py +300 -0
- modelopt/onnx/quantization/__init__.py +20 -0
- modelopt/onnx/quantization/__main__.py +291 -0
- modelopt/onnx/quantization/calib_utils.py +178 -0
- modelopt/onnx/quantization/extensions.py +35 -0
- modelopt/onnx/quantization/fp8.py +375 -0
- modelopt/onnx/quantization/graph_utils.py +1479 -0
- modelopt/onnx/quantization/gs_patching.py +112 -0
- modelopt/onnx/quantization/int4.py +1347 -0
- modelopt/onnx/quantization/int8.py +283 -0
- modelopt/onnx/quantization/operators.py +121 -0
- modelopt/onnx/quantization/ort_patching.py +1766 -0
- modelopt/onnx/quantization/ort_utils.py +346 -0
- modelopt/onnx/quantization/partitioning.py +439 -0
- modelopt/onnx/quantization/qdq_utils.py +1349 -0
- modelopt/onnx/quantization/quant_utils.py +283 -0
- modelopt/onnx/quantization/quantize.py +528 -0
- modelopt/onnx/quantization/src/modelopt_round_and_pack_ext.cpp +127 -0
- modelopt/onnx/trt_utils.py +424 -0
- modelopt/onnx/utils.py +753 -0
- modelopt/torch/__init__.py +41 -0
- modelopt/torch/_deploy/__init__.py +23 -0
- modelopt/torch/_deploy/_runtime/__init__.py +29 -0
- modelopt/torch/_deploy/_runtime/common.py +72 -0
- modelopt/torch/_deploy/_runtime/ort_client.py +193 -0
- modelopt/torch/_deploy/_runtime/registry.py +83 -0
- modelopt/torch/_deploy/_runtime/runtime_client.py +154 -0
- modelopt/torch/_deploy/_runtime/tensorrt/constants.py +108 -0
- modelopt/torch/_deploy/_runtime/tensorrt/engine_builder.py +324 -0
- modelopt/torch/_deploy/_runtime/tensorrt/hw_param_config.py +51 -0
- modelopt/torch/_deploy/_runtime/tensorrt/layerwise_profiling.py +178 -0
- modelopt/torch/_deploy/_runtime/tensorrt/parse_trtexec_log.py +151 -0
- modelopt/torch/_deploy/_runtime/tensorrt/tensorrt_utils.py +200 -0
- modelopt/torch/_deploy/_runtime/trt_client.py +219 -0
- modelopt/torch/_deploy/compilation.py +142 -0
- modelopt/torch/_deploy/device_model.py +157 -0
- modelopt/torch/_deploy/profiling.py +193 -0
- modelopt/torch/_deploy/utils/__init__.py +18 -0
- modelopt/torch/_deploy/utils/onnx_optimizer.py +120 -0
- modelopt/torch/_deploy/utils/onnx_utils.py +58 -0
- modelopt/torch/_deploy/utils/torch_onnx.py +593 -0
- modelopt/torch/distill/__init__.py +28 -0
- modelopt/torch/distill/config.py +130 -0
- modelopt/torch/distill/distillation.py +84 -0
- modelopt/torch/distill/distillation_model.py +318 -0
- modelopt/torch/distill/loss_balancers.py +140 -0
- modelopt/torch/distill/losses.py +275 -0
- modelopt/torch/distill/mode.py +223 -0
- modelopt/torch/distill/plugins/__init__.py +24 -0
- modelopt/torch/distill/plugins/huggingface.py +110 -0
- modelopt/torch/distill/plugins/megatron.py +476 -0
- modelopt/torch/distill/registry.py +23 -0
- modelopt/torch/export/__init__.py +24 -0
- modelopt/torch/export/convert_hf_config.py +117 -0
- modelopt/torch/export/distribute.py +300 -0
- modelopt/torch/export/hf_config_map.py +84 -0
- modelopt/torch/export/layer_utils.py +1893 -0
- modelopt/torch/export/mcore_config_map.py +25 -0
- modelopt/torch/export/model_config.py +623 -0
- modelopt/torch/export/model_config_export.py +552 -0
- modelopt/torch/export/model_config_utils.py +392 -0
- modelopt/torch/export/model_utils.py +71 -0
- modelopt/torch/export/plugins/__init__.py +21 -0
- modelopt/torch/export/plugins/mcore_common.py +67 -0
- modelopt/torch/export/plugins/mcore_custom.py +530 -0
- modelopt/torch/export/plugins/mcore_deepseek.py +143 -0
- modelopt/torch/export/plugins/mcore_gptoss.py +71 -0
- modelopt/torch/export/plugins/mcore_llama.py +164 -0
- modelopt/torch/export/plugins/mcore_nemotron.py +90 -0
- modelopt/torch/export/plugins/mcore_qwen.py +99 -0
- modelopt/torch/export/plugins/megatron_importer.py +647 -0
- modelopt/torch/export/postprocess.py +847 -0
- modelopt/torch/export/quant_utils.py +1069 -0
- modelopt/torch/export/tensorrt_llm_type.py +42 -0
- modelopt/torch/export/tensorrt_llm_utils.py +405 -0
- modelopt/torch/export/transformer_engine.py +94 -0
- modelopt/torch/export/unified_export_hf.py +537 -0
- modelopt/torch/export/unified_export_megatron.py +1164 -0
- modelopt/torch/nas/__init__.py +27 -0
- modelopt/torch/nas/algorithms.py +856 -0
- modelopt/torch/nas/autonas.py +808 -0
- modelopt/torch/nas/conversion.py +119 -0
- modelopt/torch/nas/hparams/__init__.py +19 -0
- modelopt/torch/nas/hparams/concat.py +319 -0
- modelopt/torch/nas/hparams/container.py +48 -0
- modelopt/torch/nas/modules/__init__.py +21 -0
- modelopt/torch/nas/modules/container.py +99 -0
- modelopt/torch/nas/modules/conv.py +254 -0
- modelopt/torch/nas/modules/linear.py +76 -0
- modelopt/torch/nas/modules/norm.py +193 -0
- modelopt/torch/nas/modules/utils.py +61 -0
- modelopt/torch/nas/patch.py +261 -0
- modelopt/torch/nas/plugins/__init__.py +29 -0
- modelopt/torch/nas/plugins/megatron.py +1497 -0
- modelopt/torch/nas/plugins/torch.py +71 -0
- modelopt/torch/nas/plugins/transformer_engine.py +42 -0
- modelopt/torch/nas/plugins/transformers.py +210 -0
- modelopt/torch/nas/registry.py +23 -0
- modelopt/torch/nas/search_space.py +260 -0
- modelopt/torch/nas/traced_hp.py +149 -0
- modelopt/torch/nas/utils.py +408 -0
- modelopt/torch/opt/__init__.py +41 -0
- modelopt/torch/opt/_hooks.py +98 -0
- modelopt/torch/opt/config.py +383 -0
- modelopt/torch/opt/conversion.py +620 -0
- modelopt/torch/opt/dynamic.py +1367 -0
- modelopt/torch/opt/hparam.py +275 -0
- modelopt/torch/opt/mode.py +353 -0
- modelopt/torch/opt/plugins/__init__.py +38 -0
- modelopt/torch/opt/plugins/diffusers.py +40 -0
- modelopt/torch/opt/plugins/huggingface.py +162 -0
- modelopt/torch/opt/plugins/mcore_dist_checkpointing.py +211 -0
- modelopt/torch/opt/plugins/megatron.py +142 -0
- modelopt/torch/opt/plugins/peft.py +108 -0
- modelopt/torch/opt/plugins/transformers.py +160 -0
- modelopt/torch/opt/searcher.py +363 -0
- modelopt/torch/opt/utils.py +125 -0
- modelopt/torch/prune/__init__.py +26 -0
- modelopt/torch/prune/fastnas.py +376 -0
- modelopt/torch/prune/gradnas.py +349 -0
- modelopt/torch/prune/plugins/__init__.py +24 -0
- modelopt/torch/prune/plugins/mcore_minitron.py +263 -0
- modelopt/torch/prune/plugins/transformers.py +41 -0
- modelopt/torch/prune/pruning.py +211 -0
- modelopt/torch/quantization/__init__.py +26 -0
- modelopt/torch/quantization/algorithms.py +671 -0
- modelopt/torch/quantization/backends/__init__.py +20 -0
- modelopt/torch/quantization/backends/fp8_per_tensor_gemm.py +204 -0
- modelopt/torch/quantization/backends/gemm_registry.py +156 -0
- modelopt/torch/quantization/backends/nvfp4_gemm.py +201 -0
- modelopt/torch/quantization/backends/utils.py +28 -0
- modelopt/torch/quantization/calib/__init__.py +25 -0
- modelopt/torch/quantization/calib/bias.py +172 -0
- modelopt/torch/quantization/calib/calibrator.py +63 -0
- modelopt/torch/quantization/calib/histogram.py +434 -0
- modelopt/torch/quantization/calib/max.py +110 -0
- modelopt/torch/quantization/compress.py +227 -0
- modelopt/torch/quantization/config.py +1180 -0
- modelopt/torch/quantization/conversion.py +389 -0
- modelopt/torch/quantization/export_onnx.py +639 -0
- modelopt/torch/quantization/extensions.py +89 -0
- modelopt/torch/quantization/mode.py +428 -0
- modelopt/torch/quantization/model_calib.py +949 -0
- modelopt/torch/quantization/model_quant.py +477 -0
- modelopt/torch/quantization/nn/__init__.py +26 -0
- modelopt/torch/quantization/nn/functional.py +109 -0
- modelopt/torch/quantization/nn/modules/__init__.py +16 -0
- modelopt/torch/quantization/nn/modules/quant_activations.py +22 -0
- modelopt/torch/quantization/nn/modules/quant_batchnorm.py +24 -0
- modelopt/torch/quantization/nn/modules/quant_conv.py +123 -0
- modelopt/torch/quantization/nn/modules/quant_instancenorm.py +39 -0
- modelopt/torch/quantization/nn/modules/quant_linear.py +258 -0
- modelopt/torch/quantization/nn/modules/quant_module.py +211 -0
- modelopt/torch/quantization/nn/modules/quant_pooling.py +99 -0
- modelopt/torch/quantization/nn/modules/quant_rnn.py +527 -0
- modelopt/torch/quantization/nn/modules/tensor_quantizer.py +1293 -0
- modelopt/torch/quantization/plugins/__init__.py +73 -0
- modelopt/torch/quantization/plugins/accelerate.py +214 -0
- modelopt/torch/quantization/plugins/apex.py +52 -0
- modelopt/torch/quantization/plugins/attention.py +312 -0
- modelopt/torch/quantization/plugins/custom.py +192 -0
- modelopt/torch/quantization/plugins/diffusers.py +237 -0
- modelopt/torch/quantization/plugins/fairscale.py +45 -0
- modelopt/torch/quantization/plugins/huggingface.py +736 -0
- modelopt/torch/quantization/plugins/megatron.py +462 -0
- modelopt/torch/quantization/plugins/peft.py +148 -0
- modelopt/torch/quantization/plugins/transformer_engine.py +60 -0
- modelopt/torch/quantization/plugins/transformers.py +52 -0
- modelopt/torch/quantization/plugins/transformers_trainer.py +426 -0
- modelopt/torch/quantization/plugins/trl.py +27 -0
- modelopt/torch/quantization/plugins/vllm.py +199 -0
- modelopt/torch/quantization/qtensor/__init__.py +24 -0
- modelopt/torch/quantization/qtensor/base_qtensor.py +370 -0
- modelopt/torch/quantization/qtensor/fp8_tensor.py +151 -0
- modelopt/torch/quantization/qtensor/int4_tensor.py +130 -0
- modelopt/torch/quantization/qtensor/int8_tensor.py +124 -0
- modelopt/torch/quantization/qtensor/mxfp4_tensor.py +144 -0
- modelopt/torch/quantization/qtensor/nf4_tensor.py +199 -0
- modelopt/torch/quantization/qtensor/nvfp4_tensor.py +295 -0
- modelopt/torch/quantization/src/tensor_quant.cpp +77 -0
- modelopt/torch/quantization/src/tensor_quant.h +69 -0
- modelopt/torch/quantization/src/tensor_quant_fp8.cpp +48 -0
- modelopt/torch/quantization/src/tensor_quant_gpu.cu +361 -0
- modelopt/torch/quantization/src/tensor_quant_gpu_fp8.cu +122 -0
- modelopt/torch/quantization/src/tensor_quant_mx.cu +411 -0
- modelopt/torch/quantization/src/tensor_quant_mx.h +247 -0
- modelopt/torch/quantization/tensor_quant.py +816 -0
- modelopt/torch/quantization/triton/__init__.py +36 -0
- modelopt/torch/quantization/triton/fp4_kernel.py +303 -0
- modelopt/torch/quantization/utils.py +443 -0
- modelopt/torch/sparsity/__init__.py +19 -0
- modelopt/torch/sparsity/config.py +51 -0
- modelopt/torch/sparsity/magnitude.py +148 -0
- modelopt/torch/sparsity/mode.py +203 -0
- modelopt/torch/sparsity/module.py +87 -0
- modelopt/torch/sparsity/plugins/__init__.py +27 -0
- modelopt/torch/sparsity/plugins/megatron.py +93 -0
- modelopt/torch/sparsity/searcher.py +84 -0
- modelopt/torch/sparsity/sparsegpt.py +276 -0
- modelopt/torch/sparsity/sparsification.py +123 -0
- modelopt/torch/speculative/__init__.py +20 -0
- modelopt/torch/speculative/config.py +100 -0
- modelopt/torch/speculative/eagle/__init__.py +20 -0
- modelopt/torch/speculative/eagle/conversion.py +67 -0
- modelopt/torch/speculative/eagle/default_config.py +50 -0
- modelopt/torch/speculative/eagle/eagle_model.py +51 -0
- modelopt/torch/speculative/eagle/utils.py +91 -0
- modelopt/torch/speculative/medusa/__init__.py +19 -0
- modelopt/torch/speculative/medusa/conversion.py +59 -0
- modelopt/torch/speculative/medusa/medusa_model.py +32 -0
- modelopt/torch/speculative/mode.py +86 -0
- modelopt/torch/speculative/plugins/__init__.py +33 -0
- modelopt/torch/speculative/plugins/megatron_eagle.py +2119 -0
- modelopt/torch/speculative/plugins/megatron_medusa.py +312 -0
- modelopt/torch/speculative/plugins/transformers.py +1138 -0
- modelopt/torch/speculative/speculative_decoding.py +59 -0
- modelopt/torch/speculative/utils.py +364 -0
- modelopt/torch/trace/__init__.py +25 -0
- modelopt/torch/trace/analyzer.py +1368 -0
- modelopt/torch/trace/modules/__init__.py +19 -0
- modelopt/torch/trace/modules/concat.py +384 -0
- modelopt/torch/trace/modules/nn.py +153 -0
- modelopt/torch/trace/plugins/__init__.py +34 -0
- modelopt/torch/trace/plugins/megatron.py +36 -0
- modelopt/torch/trace/plugins/transformers.py +58 -0
- modelopt/torch/trace/symbols.py +545 -0
- modelopt/torch/trace/tracer.py +331 -0
- modelopt/torch/utils/__init__.py +28 -0
- modelopt/torch/utils/_pytree.py +133 -0
- modelopt/torch/utils/cpp_extension.py +91 -0
- modelopt/torch/utils/dataset_utils.py +484 -0
- modelopt/torch/utils/distributed.py +311 -0
- modelopt/torch/utils/graph.py +123 -0
- modelopt/torch/utils/image_processor.py +112 -0
- modelopt/torch/utils/import_utils.py +35 -0
- modelopt/torch/utils/list.py +54 -0
- modelopt/torch/utils/logging.py +127 -0
- modelopt/torch/utils/memory_monitor.py +146 -0
- modelopt/torch/utils/network.py +658 -0
- modelopt/torch/utils/perf.py +174 -0
- modelopt/torch/utils/plugins/__init__.py +27 -0
- modelopt/torch/utils/plugins/megatron_generate.py +299 -0
- modelopt/torch/utils/plugins/megatron_mmlu.py +152 -0
- modelopt/torch/utils/plugins/megatron_preprocess_data.py +200 -0
- modelopt/torch/utils/random.py +163 -0
- modelopt/torch/utils/speech_dataset_utils.py +134 -0
- modelopt/torch/utils/tensor.py +87 -0
- modelopt/torch/utils/vlm_dataset_utils.py +110 -0
- nvidia_modelopt-0.35.0.dist-info/METADATA +171 -0
- nvidia_modelopt-0.35.0.dist-info/RECORD +272 -0
- nvidia_modelopt-0.35.0.dist-info/WHEEL +5 -0
- nvidia_modelopt-0.35.0.dist-info/licenses/LICENSE +14 -0
- nvidia_modelopt-0.35.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,375 @@
|
|
|
1
|
+
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
2
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
3
|
+
#
|
|
4
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
5
|
+
# you may not use this file except in compliance with the License.
|
|
6
|
+
# You may obtain a copy of the License at
|
|
7
|
+
#
|
|
8
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
9
|
+
#
|
|
10
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
11
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
12
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
13
|
+
# See the License for the specific language governing permissions and
|
|
14
|
+
# limitations under the License.
|
|
15
|
+
|
|
16
|
+
"""Performs FP8 GEMM only quantization of an ONNX model, and returns the ONNX ModelProto."""
|
|
17
|
+
|
|
18
|
+
import os
|
|
19
|
+
import tempfile
|
|
20
|
+
import time
|
|
21
|
+
from functools import reduce
|
|
22
|
+
|
|
23
|
+
import numpy as np
|
|
24
|
+
import onnx
|
|
25
|
+
import onnx_graphsurgeon as gs
|
|
26
|
+
from onnx import numpy_helper
|
|
27
|
+
from onnx_graphsurgeon.ir.graph import Graph
|
|
28
|
+
from onnxruntime.quantization import CalibrationMethod
|
|
29
|
+
from onnxruntime.quantization.calibrate import CalibrationDataReader
|
|
30
|
+
|
|
31
|
+
import modelopt.onnx.utils as onnx_utils
|
|
32
|
+
from modelopt.onnx.autocast.convert import convert_to_f16
|
|
33
|
+
from modelopt.onnx.logging_config import configure_logging, logger
|
|
34
|
+
from modelopt.onnx.quantization.graph_utils import (
|
|
35
|
+
convert_fp16_io,
|
|
36
|
+
expand_node_names_from_patterns,
|
|
37
|
+
find_nodes_to_exclude,
|
|
38
|
+
get_concat_eliminated_tensors,
|
|
39
|
+
get_tensor_producer_nodes,
|
|
40
|
+
insert_fp8_mha_casts,
|
|
41
|
+
remove_output_initializers,
|
|
42
|
+
remove_partial_input_qdq,
|
|
43
|
+
)
|
|
44
|
+
from modelopt.onnx.quantization.int8 import _find_nodes_to_quantize
|
|
45
|
+
from modelopt.onnx.quantization.ort_patching import _quantize_static as quantize_static
|
|
46
|
+
from modelopt.onnx.quantization.ort_utils import configure_ort
|
|
47
|
+
from modelopt.onnx.quantization.qdq_utils import has_qdq_nodes
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _find_unsupported_fp8_convs_to_exclude(graph: Graph):
|
|
51
|
+
"""Find unsupported FP8 Conv nodes to exclude.
|
|
52
|
+
|
|
53
|
+
The input and output channel alignment requirement for FP8
|
|
54
|
+
conv kernels for input and output type FP8E4M3 should be both 16.
|
|
55
|
+
The filter size for FP8 conv kernels should be less than 32.
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
graph: Onnx model graph.
|
|
59
|
+
|
|
60
|
+
Returns:
|
|
61
|
+
List of Conv nodes.
|
|
62
|
+
"""
|
|
63
|
+
unsupported_conv_nodes = []
|
|
64
|
+
logger.info("Scanning for unsupported FP8 Conv nodes")
|
|
65
|
+
for node in graph.nodes:
|
|
66
|
+
if node.op == "Conv":
|
|
67
|
+
weight = node.inputs[1]
|
|
68
|
+
|
|
69
|
+
# If weight.shape is None, it means the weight is not a constant tensor.
|
|
70
|
+
# Skip the convs with non-constant weights.
|
|
71
|
+
if weight.shape is None:
|
|
72
|
+
logger.debug(f"Skipped quantizing conv: {node.name} due to non-constant weight")
|
|
73
|
+
unsupported_conv_nodes.append(node.name)
|
|
74
|
+
continue
|
|
75
|
+
|
|
76
|
+
assert 3 <= len(weight.shape) <= 5, (
|
|
77
|
+
f"Invalid weight shape {weight.shape}. Only 1D, 2D, and 3D convolutions are supported"
|
|
78
|
+
)
|
|
79
|
+
output_channel = weight.shape[0]
|
|
80
|
+
input_channel = weight.shape[1]
|
|
81
|
+
if output_channel % 16 != input_channel % 16:
|
|
82
|
+
logger.debug(f"Found unpaddable conv for FP8: {node.name}")
|
|
83
|
+
unsupported_conv_nodes.append(node.name)
|
|
84
|
+
continue
|
|
85
|
+
|
|
86
|
+
if output_channel < 16 or input_channel < 16:
|
|
87
|
+
logger.debug(f"Found Conv with I/O channel size less than 16: {node.name}")
|
|
88
|
+
unsupported_conv_nodes.append(node.name)
|
|
89
|
+
continue
|
|
90
|
+
|
|
91
|
+
filter_size = reduce(lambda x, y: x * y, weight.shape[2:])
|
|
92
|
+
if filter_size > 32:
|
|
93
|
+
logger.debug(f"Found large filter conv for FP8: {node.name}")
|
|
94
|
+
unsupported_conv_nodes.append(node.name)
|
|
95
|
+
|
|
96
|
+
logger.info(f"Found {len(unsupported_conv_nodes)} unsupported FP8 Conv nodes")
|
|
97
|
+
return unsupported_conv_nodes
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def int8_to_fp8(onnx_model: onnx.ModelProto) -> onnx.ModelProto:
|
|
101
|
+
"""Converts the INT8 quantized model to FP8 quantized model.
|
|
102
|
+
|
|
103
|
+
Note. This conversion works only for max calibrated INT8 models.
|
|
104
|
+
|
|
105
|
+
Args:
|
|
106
|
+
onnx_model: INT8 quantized ONNX model.
|
|
107
|
+
|
|
108
|
+
Returns:
|
|
109
|
+
FP8 quantized ONNX model.
|
|
110
|
+
"""
|
|
111
|
+
logger.info("Starting INT8 to FP8 conversion")
|
|
112
|
+
graph = onnx_model.graph
|
|
113
|
+
initializers = graph.initializer
|
|
114
|
+
tensor_producers = get_tensor_producer_nodes(graph)
|
|
115
|
+
processed_tensor = set()
|
|
116
|
+
initializer_indices = {
|
|
117
|
+
initializer.name: idx for idx, initializer in enumerate(graph.initializer)
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
def _int8_scale_to_fp8_scale(scale: np.ndarray, scale_name: str):
|
|
121
|
+
np_scale = onnx.numpy_helper.to_array(scale)
|
|
122
|
+
np_fp8_scale = (np_scale * 448.0) / 127.0
|
|
123
|
+
dtype = onnx.helper.tensor_dtype_to_np_dtype(scale.data_type)
|
|
124
|
+
return numpy_helper.from_array(np_fp8_scale.astype(dtype), scale_name)
|
|
125
|
+
|
|
126
|
+
def _update_tensor_type(tensor_name):
|
|
127
|
+
tensor = onnx_utils.get_tensor_by_name(onnx_model, tensor_name)
|
|
128
|
+
if tensor:
|
|
129
|
+
tensor.type.tensor_type.elem_type = onnx.TensorProto.FLOAT8E4M3FN
|
|
130
|
+
|
|
131
|
+
def _convert(node: onnx.NodeProto):
|
|
132
|
+
scale_name = node.input[1]
|
|
133
|
+
zero_point_name = node.input[2]
|
|
134
|
+
|
|
135
|
+
if scale_name not in processed_tensor:
|
|
136
|
+
scale_idx = initializer_indices.get(scale_name)
|
|
137
|
+
if scale_idx is not None:
|
|
138
|
+
scale = initializers[scale_idx]
|
|
139
|
+
fp8_scale = _int8_scale_to_fp8_scale(scale, scale_name)
|
|
140
|
+
initializers[scale_idx].CopyFrom(fp8_scale)
|
|
141
|
+
else:
|
|
142
|
+
producer_node = tensor_producers[scale_name]
|
|
143
|
+
scale = producer_node.attribute[0].t
|
|
144
|
+
fp8_scale = _int8_scale_to_fp8_scale(scale, scale_name)
|
|
145
|
+
producer_node.attribute[0].t.CopyFrom(fp8_scale)
|
|
146
|
+
processed_tensor.add(scale_name)
|
|
147
|
+
|
|
148
|
+
if zero_point_name not in processed_tensor:
|
|
149
|
+
zero_point_idx = initializer_indices.get(zero_point_name)
|
|
150
|
+
assert zero_point_idx is not None, (
|
|
151
|
+
f"Expected '{zero_point_name}' to be found in 'graph.initializer', but it was not present."
|
|
152
|
+
)
|
|
153
|
+
zero_point = initializers[zero_point_idx]
|
|
154
|
+
dtype = onnx.helper.tensor_dtype_to_np_dtype(zero_point.data_type)
|
|
155
|
+
vals = np.array(zero_point.int32_data, dtype=dtype).tobytes()
|
|
156
|
+
|
|
157
|
+
np_zero_point = onnx.helper.make_tensor(
|
|
158
|
+
zero_point_name, onnx.TensorProto.FLOAT8E4M3FN, zero_point.dims, vals, raw=True
|
|
159
|
+
)
|
|
160
|
+
initializers[zero_point_idx].CopyFrom(np_zero_point)
|
|
161
|
+
processed_tensor.add(zero_point_name)
|
|
162
|
+
|
|
163
|
+
# Update the Q input tensor type and the DQ output tensor type
|
|
164
|
+
if node.op_type == "QuantizeLinear":
|
|
165
|
+
for out in node.output:
|
|
166
|
+
_update_tensor_type(out)
|
|
167
|
+
if node.op_type == "DequantizeLinear":
|
|
168
|
+
_update_tensor_type(node.input[0])
|
|
169
|
+
|
|
170
|
+
# Iterate through the nodes and convert the scales and zero points
|
|
171
|
+
# Also update the Q input tensor type and the DQ output tensor type
|
|
172
|
+
for node in graph.node:
|
|
173
|
+
if node.op_type in ["DequantizeLinear", "QuantizeLinear"]:
|
|
174
|
+
_convert(node)
|
|
175
|
+
|
|
176
|
+
return onnx_model
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
def upgrade_opset_21(onnx_model: onnx.ModelProto) -> onnx.ModelProto:
|
|
180
|
+
"""Modifies the ONNX graph such that it follows the opset 21 requirements.
|
|
181
|
+
|
|
182
|
+
This is necessary for FP8+FP16 quantization since FP8 QuantizeLinear/DequantizeLinear ops do not support FP16
|
|
183
|
+
scaling factors until opset 21.
|
|
184
|
+
"""
|
|
185
|
+
logger.info("Upgrading model to opset 21")
|
|
186
|
+
graph = gs.import_onnx(onnx_model)
|
|
187
|
+
|
|
188
|
+
for node in graph.nodes:
|
|
189
|
+
# QuantizeLinear/DequantizeLinear op with FP16 scales are only supported with empty domain
|
|
190
|
+
# and opset_import version=21.
|
|
191
|
+
if node.op in {"QuantizeLinear", "DequantizeLinear"}:
|
|
192
|
+
node.domain = ""
|
|
193
|
+
|
|
194
|
+
# ReduceMean op no longer has "axes" attribute in opset 21. Instead, it should be the second input tensor.
|
|
195
|
+
if node.op == "ReduceMean" and "axes" in node.attrs:
|
|
196
|
+
axes = gs.Constant(
|
|
197
|
+
name=node.name + "_axes", values=np.array(node.attrs["axes"], dtype=np.int64)
|
|
198
|
+
)
|
|
199
|
+
del node.attrs["axes"]
|
|
200
|
+
node.inputs.append(axes)
|
|
201
|
+
|
|
202
|
+
onnx_model = gs.export_onnx(graph)
|
|
203
|
+
|
|
204
|
+
# Set opset_import version to 21.
|
|
205
|
+
for opset_import in onnx_model.opset_import:
|
|
206
|
+
if opset_import.domain == "":
|
|
207
|
+
opset_import.version = 21
|
|
208
|
+
|
|
209
|
+
return onnx_model
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def quantize(
|
|
213
|
+
onnx_path: str,
|
|
214
|
+
calibration_method: str = "entropy",
|
|
215
|
+
calibration_data_reader: CalibrationDataReader = None,
|
|
216
|
+
calibration_cache_path: str | None = None,
|
|
217
|
+
calibration_shapes: str | None = None,
|
|
218
|
+
calibration_eps: list[str] = ["cpu", "cuda:0", "trt"],
|
|
219
|
+
op_types_to_quantize: list[str] | None = None,
|
|
220
|
+
op_types_to_exclude: list[str] | None = None,
|
|
221
|
+
nodes_to_quantize: list[str] | None = None,
|
|
222
|
+
nodes_to_exclude: list[str] | None = None,
|
|
223
|
+
use_external_data_format: bool = False,
|
|
224
|
+
intermediate_generated_files: list[str] = [],
|
|
225
|
+
trt_extra_plugin_lib_paths: list[str] | None = None,
|
|
226
|
+
high_precision_dtype: str = "fp16",
|
|
227
|
+
mha_accumulation_dtype: str = "fp16",
|
|
228
|
+
passes: list[str] = ["concat_elimination"],
|
|
229
|
+
log_level: str = "INFO",
|
|
230
|
+
calibrate_per_node: bool = False,
|
|
231
|
+
custom_ops_to_quantize: list[str] = [],
|
|
232
|
+
**kwargs,
|
|
233
|
+
) -> onnx.ModelProto:
|
|
234
|
+
"""Applies FP8 GEMM only quantization to an ONNX file.
|
|
235
|
+
|
|
236
|
+
Currently, ['Conv', 'Gemm', 'MatMul', 'Residual-Add'] quantization is supported.
|
|
237
|
+
"""
|
|
238
|
+
configure_logging(level=log_level.upper())
|
|
239
|
+
logger.info("Starting FP8 quantization process")
|
|
240
|
+
t_start = time.time()
|
|
241
|
+
|
|
242
|
+
# Load the onnx graph
|
|
243
|
+
logger.info(f"Loading ONNX model from {onnx_path}")
|
|
244
|
+
onnx_model = onnx.load(onnx_path, load_external_data=True)
|
|
245
|
+
onnx_model = onnx_utils.infer_shapes(onnx_model)
|
|
246
|
+
graph = gs.import_onnx(onnx_model)
|
|
247
|
+
graph.toposort()
|
|
248
|
+
|
|
249
|
+
# If the model already has QDQ nodes, skip the quantization process
|
|
250
|
+
if has_qdq_nodes(onnx_model):
|
|
251
|
+
logger.info("Model already has QDQ nodes, skipping quantization")
|
|
252
|
+
return onnx_model
|
|
253
|
+
|
|
254
|
+
# The quantizable op types for FP8 are limited to Conv, Gemm, Matmul, and Residual-Add
|
|
255
|
+
fp8_supported_op_types = ["Gemm", "MatMul", "Conv", "Add"]
|
|
256
|
+
op_types_to_quantize = op_types_to_quantize or fp8_supported_op_types
|
|
257
|
+
if not set(op_types_to_quantize) <= set(fp8_supported_op_types):
|
|
258
|
+
raise RuntimeError(
|
|
259
|
+
f"Unsupported op types in fp8 mode: '{set(op_types_to_quantize) - set(fp8_supported_op_types)}'"
|
|
260
|
+
)
|
|
261
|
+
op_types_to_quantize.extend(list(custom_ops_to_quantize))
|
|
262
|
+
|
|
263
|
+
# Collect node names to exclude from quantization
|
|
264
|
+
nodes_to_exclude = find_nodes_to_exclude(graph, nodes_to_exclude, op_types_to_exclude) # type: ignore[arg-type]
|
|
265
|
+
nodes_to_exclude.extend(_find_unsupported_fp8_convs_to_exclude(graph))
|
|
266
|
+
|
|
267
|
+
# Change the default configuration of ORT quantization
|
|
268
|
+
op_types = {node.op for node in graph.nodes}
|
|
269
|
+
trt_guided_options, quantizable_op_types = configure_ort(
|
|
270
|
+
list(op_types),
|
|
271
|
+
op_types_to_quantize,
|
|
272
|
+
trt_extra_plugin_lib_paths,
|
|
273
|
+
calibration_eps,
|
|
274
|
+
calibrate_per_node,
|
|
275
|
+
custom_ops_to_quantize,
|
|
276
|
+
)
|
|
277
|
+
logger.info(
|
|
278
|
+
f"Quantizable op types in the model: {[t for t in op_types_to_quantize if t in op_types]}"
|
|
279
|
+
)
|
|
280
|
+
|
|
281
|
+
# Collect node names to include in quantization
|
|
282
|
+
no_quantize_inputs = []
|
|
283
|
+
nodes_to_quantize = expand_node_names_from_patterns(graph, nodes_to_quantize)
|
|
284
|
+
if not nodes_to_quantize:
|
|
285
|
+
quantizable_nodes, no_quantize_inputs = _find_nodes_to_quantize(
|
|
286
|
+
graph, quantizable_op_types, nodes_to_exclude
|
|
287
|
+
)
|
|
288
|
+
nodes_to_quantize = [node.name for node in quantizable_nodes]
|
|
289
|
+
|
|
290
|
+
# Update the list of nodes to quantize
|
|
291
|
+
nodes_to_quantize = [
|
|
292
|
+
node_name for node_name in nodes_to_quantize if node_name not in nodes_to_exclude
|
|
293
|
+
]
|
|
294
|
+
|
|
295
|
+
if not nodes_to_quantize:
|
|
296
|
+
logger.info("No node or node type is selected for quantization or model does not have them")
|
|
297
|
+
else:
|
|
298
|
+
logger.debug(f"Selected nodes to quantize: {nodes_to_quantize}")
|
|
299
|
+
|
|
300
|
+
if passes and "concat_elimination" in passes:
|
|
301
|
+
group_qdq_tensors = get_concat_eliminated_tensors(onnx_model, nodes_to_quantize)
|
|
302
|
+
if group_qdq_tensors:
|
|
303
|
+
trt_guided_options["group_qdq_tensors"] = group_qdq_tensors
|
|
304
|
+
logger.debug(f"Grouping QDQ tensors for concat elimination: {group_qdq_tensors}")
|
|
305
|
+
|
|
306
|
+
# Create a temp file for intermediate model
|
|
307
|
+
tmp_onnx_file, tmp_onnx_path = tempfile.mkstemp(suffix=".onnx")
|
|
308
|
+
os.close(tmp_onnx_file)
|
|
309
|
+
logger.debug(f"Created temporary file for intermediate model: {tmp_onnx_path}")
|
|
310
|
+
|
|
311
|
+
# Quantize in INT8 mode using ORT's Entropy or MinMax calibration method, with
|
|
312
|
+
# ActivationSymmetric as True. For MinMax, that's equivalent to max calibration.
|
|
313
|
+
logger.info(f"Starting INT8 quantization with '{calibration_method}' calibration")
|
|
314
|
+
quantize_static(
|
|
315
|
+
onnx_path,
|
|
316
|
+
tmp_onnx_path,
|
|
317
|
+
calibration_data_reader,
|
|
318
|
+
op_types_to_quantize=op_types_to_quantize,
|
|
319
|
+
nodes_to_quantize=nodes_to_quantize,
|
|
320
|
+
per_channel=True,
|
|
321
|
+
extra_options=trt_guided_options,
|
|
322
|
+
use_external_data_format=use_external_data_format,
|
|
323
|
+
calibrate_method=(
|
|
324
|
+
CalibrationMethod.Entropy
|
|
325
|
+
if calibration_method == "entropy"
|
|
326
|
+
# With ActivationSymmetric as True, MinMax calibration is equivalent to max calibration
|
|
327
|
+
else CalibrationMethod.MinMax
|
|
328
|
+
),
|
|
329
|
+
)
|
|
330
|
+
intermediate_generated_files.append(tmp_onnx_path)
|
|
331
|
+
if use_external_data_format:
|
|
332
|
+
intermediate_generated_files.append(tmp_onnx_path + ".data")
|
|
333
|
+
|
|
334
|
+
# Post-processing of the onnx model after ORT quantization
|
|
335
|
+
logger.info("Starting post-processing of quantized model")
|
|
336
|
+
onnx_model = onnx.load(tmp_onnx_path)
|
|
337
|
+
graph = gs.import_onnx(onnx_model)
|
|
338
|
+
remove_partial_input_qdq(graph, no_quantize_inputs)
|
|
339
|
+
onnx_model = gs.export_onnx(graph)
|
|
340
|
+
|
|
341
|
+
if high_precision_dtype in ["fp16", "bf16"]:
|
|
342
|
+
# We need to convert float to float16/bfloat16 so as to speed up layers like LayerNorm or GroupNorm.
|
|
343
|
+
logger.info(f"Converting float tensors to {high_precision_dtype}")
|
|
344
|
+
graph = gs.import_onnx(onnx_model)
|
|
345
|
+
remove_output_initializers(graph, onnx_model.graph.initializer)
|
|
346
|
+
convert_fp16_io(graph)
|
|
347
|
+
onnx_model = gs.export_onnx(graph)
|
|
348
|
+
|
|
349
|
+
# Convert to fp16/bf16 model.
|
|
350
|
+
onnx_model = convert_to_f16(
|
|
351
|
+
onnx_model,
|
|
352
|
+
keep_io_types=True,
|
|
353
|
+
op_block_list=["Resize"],
|
|
354
|
+
low_precision_type=high_precision_dtype,
|
|
355
|
+
trt_plugins=trt_extra_plugin_lib_paths,
|
|
356
|
+
)
|
|
357
|
+
|
|
358
|
+
current_opsets = {opset.domain: opset.version for opset in onnx_model.opset_import}
|
|
359
|
+
opset_of_default_onnx_domain = current_opsets.get("", 0)
|
|
360
|
+
if opset_of_default_onnx_domain < 19:
|
|
361
|
+
# We need to convert the ONNX model to opset 19+ since FP8 QuantizeLinear/DequantizeLinear ops do not
|
|
362
|
+
# support FP16 scaling factors until opset 19. So, converting here to opset-21 (19+).
|
|
363
|
+
onnx_model = upgrade_opset_21(onnx_model)
|
|
364
|
+
|
|
365
|
+
if mha_accumulation_dtype == "fp32":
|
|
366
|
+
# Insert Cast nodes in MHA's BMM1 and BMM2's input and output tensors because
|
|
367
|
+
# The compiler only has FP32 accumulation kernels for FP8 MHAs.
|
|
368
|
+
logger.info("Inserting Cast nodes to enable FP8+FP16 MHA")
|
|
369
|
+
onnx_model = insert_fp8_mha_casts(onnx_model)
|
|
370
|
+
|
|
371
|
+
if nodes_to_quantize:
|
|
372
|
+
onnx_model = int8_to_fp8(onnx_model)
|
|
373
|
+
logger.info(f"FP8 quantization completed in {time.time() - t_start:.2f} seconds")
|
|
374
|
+
|
|
375
|
+
return onnx_model
|