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,120 @@
|
|
|
1
|
+
# SPDX-FileCopyrightText: Copyright (c) 2023-2025 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
|
+
"""Utilities to surgeon ONNX graph after export."""
|
|
17
|
+
|
|
18
|
+
import re
|
|
19
|
+
import time
|
|
20
|
+
|
|
21
|
+
import onnx
|
|
22
|
+
import onnx_graphsurgeon as gs
|
|
23
|
+
import torch
|
|
24
|
+
from onnx_graphsurgeon.ir.tensor import LazyValues
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def clear_inputs(node: gs.Node | gs.Tensor):
|
|
28
|
+
"""Clear all inputs for a node or tensor in ONNX."""
|
|
29
|
+
for i in node.inputs:
|
|
30
|
+
i.outputs.clear()
|
|
31
|
+
node.inputs.clear()
|
|
32
|
+
return node
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def clear_outputs(node: gs.Node | gs.Tensor):
|
|
36
|
+
"""Clear all outputs for a node or tensor in ONNX."""
|
|
37
|
+
for o in node.outputs:
|
|
38
|
+
o.inputs.clear()
|
|
39
|
+
node.outputs.clear()
|
|
40
|
+
return node
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def extract_layer_id(name: str):
|
|
44
|
+
"""Extract layer id from certain ONNX layer name.
|
|
45
|
+
|
|
46
|
+
Parameters:
|
|
47
|
+
name: str
|
|
48
|
+
The name of ONNX layer. e.g. /model/layer.0/q_proj/...
|
|
49
|
+
|
|
50
|
+
Returns:
|
|
51
|
+
The layer id for the layer as int. In the example above, it returns 0
|
|
52
|
+
"""
|
|
53
|
+
match = re.search(r"layers\.(\d+)", name)
|
|
54
|
+
if match:
|
|
55
|
+
return int(match.group(1))
|
|
56
|
+
raise Exception(f"{name} does not contain layer info!")
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def no_none_elements(elements: list):
|
|
60
|
+
"""Check if all elements in the list are not None."""
|
|
61
|
+
return all(i is not None for i in elements)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def fold_fp8_qdq_to_dq(graph: gs.Graph):
|
|
65
|
+
"""Convert FP32/FP16 weights of the given ONNX model to FP8 weights.
|
|
66
|
+
|
|
67
|
+
Even though modelopt supports FP8 onnx export, the weights are represented in fp32 + QDQ.
|
|
68
|
+
The storage is therefore very bad. In this function,
|
|
69
|
+
Q nodes will get removed from the weights and have only DQ nodes with those converted FP8
|
|
70
|
+
weights in the output model.
|
|
71
|
+
|
|
72
|
+
Parameters:
|
|
73
|
+
graph: gs.Graph.
|
|
74
|
+
|
|
75
|
+
Returns:
|
|
76
|
+
gs.Graph with only DQ nodes for weights and same QDQ nodes for activations.
|
|
77
|
+
"""
|
|
78
|
+
start_time = time.time()
|
|
79
|
+
print("Replacing all (fp32 weights + fp8 QDQ) with (fp8 weights + DQ)...")
|
|
80
|
+
# Fold constants is required since the scale is not constant yet.
|
|
81
|
+
graph.cleanup().toposort().fold_constants().cleanup()
|
|
82
|
+
|
|
83
|
+
for node in graph.nodes:
|
|
84
|
+
if node.op == "TRT_FP8QuantizeLinear":
|
|
85
|
+
# Should not remove input QDQ
|
|
86
|
+
if not isinstance(node.inputs[0], gs.Constant):
|
|
87
|
+
continue
|
|
88
|
+
|
|
89
|
+
weights = node.inputs[0]
|
|
90
|
+
scale = node.inputs[1]
|
|
91
|
+
torch_weights = torch.from_numpy(weights.values)
|
|
92
|
+
torch_scale = torch.from_numpy(scale.values)
|
|
93
|
+
quantizer_name = scale.name.rsplit("/", 1)[0]
|
|
94
|
+
dq_op = node.outputs[0].outputs[0]
|
|
95
|
+
assert dq_op.op == "TRT_FP8DequantizeLinear", (
|
|
96
|
+
f"QDQ does not occur in pairs. You reached {dq_op.op}"
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
# Replace it with Dequantize with FP8 weights. This is a WAR because numpy does not support fp8.
|
|
100
|
+
numpy_weights = (
|
|
101
|
+
(torch_weights / torch_scale).to(torch.float8_e4m3fn).view(torch.uint8).numpy()
|
|
102
|
+
)
|
|
103
|
+
tensor = onnx.TensorProto()
|
|
104
|
+
tensor.data_type = onnx.TensorProto.FLOAT8E4M3FN
|
|
105
|
+
tensor.dims.extend(numpy_weights.shape)
|
|
106
|
+
tensor.raw_data = numpy_weights.tobytes()
|
|
107
|
+
values = LazyValues(tensor)
|
|
108
|
+
onnx_weights_fp8 = gs.Constant(quantizer_name + "/fp8_weights", values)
|
|
109
|
+
|
|
110
|
+
node.outputs.clear()
|
|
111
|
+
# DQ Op is separated out
|
|
112
|
+
dq_op.inputs[0] = onnx_weights_fp8
|
|
113
|
+
dq_op.op = "DequantizeLinear"
|
|
114
|
+
dq_op.outputs[0].dtype = dq_op.inputs[1].dtype
|
|
115
|
+
|
|
116
|
+
graph.cleanup().toposort()
|
|
117
|
+
end_time = time.time()
|
|
118
|
+
print(f"fp8 qdq replaced with only dq completed in {end_time - start_time}s.")
|
|
119
|
+
|
|
120
|
+
return graph
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
# SPDX-FileCopyrightText: Copyright (c) 2023-2025 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
|
+
"""Logging configuration for modelopt ONNX quantization."""
|
|
17
|
+
|
|
18
|
+
import logging
|
|
19
|
+
import os
|
|
20
|
+
import sys
|
|
21
|
+
|
|
22
|
+
# Create a parent logger for all ONNX components
|
|
23
|
+
logger = logging.getLogger("modelopt.onnx")
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def configure_logging(level=logging.INFO, log_file=None):
|
|
27
|
+
"""Configure logging for all ONNX components.
|
|
28
|
+
|
|
29
|
+
Args:
|
|
30
|
+
level: The logging level to use (default: logging.INFO)
|
|
31
|
+
log_file: Optional path to a log file. If provided, logs will be written to this file
|
|
32
|
+
in addition to stdout (default: None)
|
|
33
|
+
"""
|
|
34
|
+
# Set level for the parent logger and all child loggers
|
|
35
|
+
logger.setLevel(level)
|
|
36
|
+
|
|
37
|
+
# Remove any existing handlers to ensure clean configuration
|
|
38
|
+
for handler in logger.handlers[:]:
|
|
39
|
+
logger.removeHandler(handler)
|
|
40
|
+
|
|
41
|
+
formatter = logging.Formatter("[modelopt][onnx] - %(levelname)s - %(message)s")
|
|
42
|
+
|
|
43
|
+
# Add file handler if log_file is specified
|
|
44
|
+
if log_file:
|
|
45
|
+
try:
|
|
46
|
+
# Create directory if it doesn't exist
|
|
47
|
+
log_dir = os.path.dirname(log_file)
|
|
48
|
+
if log_dir:
|
|
49
|
+
os.makedirs(log_dir, exist_ok=True)
|
|
50
|
+
|
|
51
|
+
file_handler = logging.FileHandler(log_file)
|
|
52
|
+
file_handler.setFormatter(formatter)
|
|
53
|
+
logger.addHandler(file_handler)
|
|
54
|
+
logger.info(f"Logging configured to write to file: {log_file}")
|
|
55
|
+
except Exception as e:
|
|
56
|
+
print(
|
|
57
|
+
f"[modelopt][onnx] - ERROR - Failed to setup file logging to {log_file}: {e!s}",
|
|
58
|
+
file=sys.stderr,
|
|
59
|
+
)
|
|
60
|
+
print("[modelopt][onnx] - INFO - Falling back to console logging.", file=sys.stderr)
|
|
61
|
+
|
|
62
|
+
# Setup handler to print log in stdout
|
|
63
|
+
console_handler = logging.StreamHandler(sys.stdout)
|
|
64
|
+
console_handler.setFormatter(formatter)
|
|
65
|
+
logger.addHandler(console_handler)
|
|
66
|
+
|
|
67
|
+
# Prevent log messages from propagating to the root logger
|
|
68
|
+
logger.propagate = False
|
|
69
|
+
|
|
70
|
+
# Ensure all child loggers inherit the level setting
|
|
71
|
+
for name in logging.root.manager.loggerDict:
|
|
72
|
+
if name.startswith("modelopt.onnx"):
|
|
73
|
+
logging.getLogger(name).setLevel(level)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
# Configure with default settings if not already configured
|
|
77
|
+
configure_logging()
|
|
@@ -0,0 +1,300 @@
|
|
|
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
|
+
"""Utility functions to categorize onnx ops."""
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def is_unary_op(op_type: str):
|
|
20
|
+
"""Returns whether the given op is a unary operator or not."""
|
|
21
|
+
return op_type in [
|
|
22
|
+
"Neg",
|
|
23
|
+
"Sqrt",
|
|
24
|
+
"Abs",
|
|
25
|
+
"Log",
|
|
26
|
+
"Exp",
|
|
27
|
+
"Not",
|
|
28
|
+
"Cast",
|
|
29
|
+
"Floor",
|
|
30
|
+
"Ceil",
|
|
31
|
+
"Round",
|
|
32
|
+
"Erf",
|
|
33
|
+
"Gelu",
|
|
34
|
+
"Sin",
|
|
35
|
+
"Cos",
|
|
36
|
+
"Atan",
|
|
37
|
+
"Sign",
|
|
38
|
+
"IsNaN",
|
|
39
|
+
"IsInf",
|
|
40
|
+
"Log",
|
|
41
|
+
"LeakyRelu",
|
|
42
|
+
# "Relu",
|
|
43
|
+
"Elu",
|
|
44
|
+
"Tanh",
|
|
45
|
+
"Sigmoid",
|
|
46
|
+
# "BatchNormalization",
|
|
47
|
+
"Softmax",
|
|
48
|
+
"Softplus",
|
|
49
|
+
"InstanceNormalization",
|
|
50
|
+
"CumSum",
|
|
51
|
+
]
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def is_binary_op(op_type: str):
|
|
55
|
+
"""Returns whether the given op is a binary operator or not."""
|
|
56
|
+
return op_type in [
|
|
57
|
+
"Add",
|
|
58
|
+
"Sub",
|
|
59
|
+
"Mul",
|
|
60
|
+
"Pow",
|
|
61
|
+
"Div",
|
|
62
|
+
"Min",
|
|
63
|
+
"Max",
|
|
64
|
+
"Greater",
|
|
65
|
+
"GreaterOrEqual",
|
|
66
|
+
"Less",
|
|
67
|
+
"LessOrEqual",
|
|
68
|
+
"Equal",
|
|
69
|
+
"BitwiseOr",
|
|
70
|
+
"BitwiseAnd",
|
|
71
|
+
"BitwiseXor",
|
|
72
|
+
"BitShift",
|
|
73
|
+
]
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def is_fusible_reduction_op(op_type: str):
|
|
77
|
+
"""Returns whether the given op type is of reduction category and fusible by the compiler."""
|
|
78
|
+
return op_type in [
|
|
79
|
+
"ReduceMax",
|
|
80
|
+
"ReduceMin",
|
|
81
|
+
"ReduceMean",
|
|
82
|
+
"ReduceProd",
|
|
83
|
+
"ReduceSum",
|
|
84
|
+
"TopK", # Transformed to BottomK based on `largest` param
|
|
85
|
+
]
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def is_fusible_scaling_op(op_type: str):
|
|
89
|
+
"""Returns whether the given op type is of scaling category and fusible with input."""
|
|
90
|
+
return op_type in [
|
|
91
|
+
"SkipSimplifiedLayerNormalization",
|
|
92
|
+
"SimplifiedLayerNormalization",
|
|
93
|
+
"MatMul",
|
|
94
|
+
"Gemm",
|
|
95
|
+
"Mul",
|
|
96
|
+
]
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def is_copy_op(op_type: str):
|
|
100
|
+
"""Returns whether the given op is a copy operator or not."""
|
|
101
|
+
return op_type in [
|
|
102
|
+
"Flatten",
|
|
103
|
+
"Transpose",
|
|
104
|
+
"Concat",
|
|
105
|
+
"Split",
|
|
106
|
+
"Squeeze",
|
|
107
|
+
"Expand",
|
|
108
|
+
"ReverseSequence",
|
|
109
|
+
"Reshape",
|
|
110
|
+
"Tile",
|
|
111
|
+
"Gather",
|
|
112
|
+
"Slice",
|
|
113
|
+
"GatherElements",
|
|
114
|
+
"GatherND",
|
|
115
|
+
"ScatterElements",
|
|
116
|
+
"ScatterND",
|
|
117
|
+
"OneHot",
|
|
118
|
+
]
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def is_linear_op(op_type: str):
|
|
122
|
+
"""Returns whether the given op type is of Linear category or not."""
|
|
123
|
+
return op_type in ["Conv", "ConvTranspose", "Gemm", "MatMul"]
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def is_pointwise_or_elementwise_op(op_type: str):
|
|
127
|
+
"""Returns whether the given op type is of Pointwise or Elementwise category or not.
|
|
128
|
+
|
|
129
|
+
This considers only the fusible types.
|
|
130
|
+
"""
|
|
131
|
+
return is_unary_op(op_type) or is_binary_op(op_type)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def is_pooling_or_window_op(op_type: str):
|
|
135
|
+
"""Returns whether the given op type is of Pooling/Window category or not."""
|
|
136
|
+
return op_type in [
|
|
137
|
+
"AveragePool",
|
|
138
|
+
"GlobalAveragePool",
|
|
139
|
+
"MaxPool",
|
|
140
|
+
"GlobalMaxPool",
|
|
141
|
+
"GlobalLpPool",
|
|
142
|
+
"LpPool",
|
|
143
|
+
"MaxPoolGridSample",
|
|
144
|
+
"HammingWindow",
|
|
145
|
+
"BlackmanWindow",
|
|
146
|
+
"HannWindow",
|
|
147
|
+
]
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def is_normalization_op(op_type: str):
|
|
151
|
+
"""Returns whether the given op type is of Normalization category or not."""
|
|
152
|
+
return op_type in [
|
|
153
|
+
"BatchNormalization",
|
|
154
|
+
"InstanceNormalization",
|
|
155
|
+
"LRN",
|
|
156
|
+
"LpNormalization",
|
|
157
|
+
"GroupNormalization",
|
|
158
|
+
"LayerNormalization",
|
|
159
|
+
]
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def is_conversion_op(op_type: str):
|
|
163
|
+
"""Returns whether the given op type is of Conversion category or not."""
|
|
164
|
+
return op_type in ["Cast", "QuantizeLinear", "DequantizeLinear"]
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def is_non_reshape_copy_op(op_type: str):
|
|
168
|
+
"""Returns whether the given op is a non-reshape copy op or not."""
|
|
169
|
+
return is_copy_op(op_type) and (op_type != "Reshape")
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def is_irregular_mem_access_op(op_type: str):
|
|
173
|
+
"""Returns whether the given op type is of Irreggular mem access category or not."""
|
|
174
|
+
return op_type in [
|
|
175
|
+
"Gather",
|
|
176
|
+
"GatherElements",
|
|
177
|
+
"GatherND",
|
|
178
|
+
"MaxRoiPool",
|
|
179
|
+
"RoiAlign",
|
|
180
|
+
"ScatterND",
|
|
181
|
+
"ScatterElements",
|
|
182
|
+
"NonMaxSuppression",
|
|
183
|
+
]
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def is_generator_op(op_type: str):
|
|
187
|
+
"""Returns whether the given op type is of Generator category or not."""
|
|
188
|
+
return op_type in [
|
|
189
|
+
"Const",
|
|
190
|
+
"ConstOfShape",
|
|
191
|
+
"EyeLike",
|
|
192
|
+
"OneHot",
|
|
193
|
+
"Multinomial",
|
|
194
|
+
"RandomNormal",
|
|
195
|
+
"RandomUniform",
|
|
196
|
+
"Bernoulli",
|
|
197
|
+
]
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def is_modifier_op(op_type: str):
|
|
201
|
+
"""Returns whether the given op type is of Modifier category or not."""
|
|
202
|
+
return op_type in [
|
|
203
|
+
"Identity",
|
|
204
|
+
"Trilu",
|
|
205
|
+
"Expand",
|
|
206
|
+
"Pad",
|
|
207
|
+
"Dropout",
|
|
208
|
+
"TileDropout",
|
|
209
|
+
"Col2Im",
|
|
210
|
+
"MaxUnpool",
|
|
211
|
+
]
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
def is_sequence_op(op_type: str):
|
|
215
|
+
"""Returns whether the given op type is of Sequence category or not."""
|
|
216
|
+
return op_type in [
|
|
217
|
+
"SequenceAt",
|
|
218
|
+
"SequenceConstruct",
|
|
219
|
+
"SequenceEmpty",
|
|
220
|
+
"SequenceErase",
|
|
221
|
+
"SequenceInsert",
|
|
222
|
+
"SequenceLength",
|
|
223
|
+
]
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
def is_selection_op(op_type: str):
|
|
227
|
+
"""Returns whether the given op type is of Selection category or not."""
|
|
228
|
+
return op_type in ["Where", "Compress"]
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
def is_control_flow_op(op_type: str):
|
|
232
|
+
"""Returns whether the given op type is of Control Flow category or not."""
|
|
233
|
+
return op_type in ["If", "Loop"]
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def is_multiclass_op(op_type: str):
|
|
237
|
+
"""Returns whether the given op type is of Multiclass category or not."""
|
|
238
|
+
return op_type in ["Einsum"]
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
def is_recurrent_op(op_type: str):
|
|
242
|
+
"""Returns whether the given op type is of Recurrent category or not."""
|
|
243
|
+
return op_type in ["LSTM", "RNN", "GRU"]
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
def is_shape_op(op_type: str):
|
|
247
|
+
"""Returns whether the given op type is of Shape category or not."""
|
|
248
|
+
return op_type in ["Shape", "Size"]
|
|
249
|
+
|
|
250
|
+
|
|
251
|
+
def is_default_quantizable_op_by_ort(op_type: str):
|
|
252
|
+
"""Returns if ORT quantizes the op type by default.
|
|
253
|
+
|
|
254
|
+
Note. Subject to change with different ORT versions.
|
|
255
|
+
Note. Users can use nodes_to_quantize and/or op_types_to_quantize arguments to quantize
|
|
256
|
+
non-default operations.
|
|
257
|
+
Reference: https://github.com/microsoft/onnxruntime/blob/main/onnxruntime/python/tools/quantization/registry.py
|
|
258
|
+
"""
|
|
259
|
+
return op_type in [
|
|
260
|
+
"Conv",
|
|
261
|
+
"Gemm",
|
|
262
|
+
"ArgMax",
|
|
263
|
+
"Relu",
|
|
264
|
+
"Split",
|
|
265
|
+
"MaxPool",
|
|
266
|
+
"InstanceNormalization",
|
|
267
|
+
"Softmax",
|
|
268
|
+
"Where",
|
|
269
|
+
"Squeeze",
|
|
270
|
+
"GlobalAveragePool",
|
|
271
|
+
"Pad",
|
|
272
|
+
"Resize",
|
|
273
|
+
"ConvTranspose",
|
|
274
|
+
"Gather",
|
|
275
|
+
"Sigmoid",
|
|
276
|
+
"EmbedLayerNormalization",
|
|
277
|
+
"Reshape",
|
|
278
|
+
"Unsqueeze",
|
|
279
|
+
"Transpose",
|
|
280
|
+
"MatMul",
|
|
281
|
+
"Concat",
|
|
282
|
+
"Mul",
|
|
283
|
+
"Clip",
|
|
284
|
+
"Add",
|
|
285
|
+
"LeakyRelu",
|
|
286
|
+
"AveragePool",
|
|
287
|
+
]
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
def is_data_dependent_shape_op(op_type: str):
|
|
291
|
+
"""Returns whether the op type has data-dependent shapes (DDS).
|
|
292
|
+
|
|
293
|
+
DDS ops have output shapes that are only determined at runtime based on its input data.
|
|
294
|
+
Source: https://onnxruntime.ai/docs/execution-providers/TensorRT-ExecutionProvider.html#data-dependant-shape-dds-ops
|
|
295
|
+
"""
|
|
296
|
+
return op_type in [
|
|
297
|
+
"NonMaxSuppression",
|
|
298
|
+
"NonZero",
|
|
299
|
+
"RoiAlign",
|
|
300
|
+
]
|
|
@@ -0,0 +1,20 @@
|
|
|
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
|
+
"""Model optimization subpackage for onnx quantization."""
|
|
17
|
+
|
|
18
|
+
# TODO: Use high-level quantize instead of quantize_int4
|
|
19
|
+
from .int4 import quantize as quantize_int4
|
|
20
|
+
from .quantize import quantize
|