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.
Files changed (272) hide show
  1. modelopt/__init__.py +20 -0
  2. modelopt/deploy/__init__.py +16 -0
  3. modelopt/deploy/llm/__init__.py +30 -0
  4. modelopt/deploy/llm/generate.py +345 -0
  5. modelopt/onnx/__init__.py +34 -0
  6. modelopt/onnx/autocast/__init__.py +19 -0
  7. modelopt/onnx/autocast/__main__.py +187 -0
  8. modelopt/onnx/autocast/convert.py +202 -0
  9. modelopt/onnx/autocast/graphsanitizer.py +422 -0
  10. modelopt/onnx/autocast/logging_config.py +78 -0
  11. modelopt/onnx/autocast/nodeclassifier.py +416 -0
  12. modelopt/onnx/autocast/precisionconverter.py +1032 -0
  13. modelopt/onnx/autocast/referencerunner.py +159 -0
  14. modelopt/onnx/autocast/utils.py +117 -0
  15. modelopt/onnx/llm_export_utils/__init__.py +16 -0
  16. modelopt/onnx/llm_export_utils/export_utils.py +162 -0
  17. modelopt/onnx/llm_export_utils/quantization_utils.py +122 -0
  18. modelopt/onnx/llm_export_utils/surgeon_utils.py +120 -0
  19. modelopt/onnx/logging_config.py +77 -0
  20. modelopt/onnx/op_types.py +300 -0
  21. modelopt/onnx/quantization/__init__.py +20 -0
  22. modelopt/onnx/quantization/__main__.py +291 -0
  23. modelopt/onnx/quantization/calib_utils.py +178 -0
  24. modelopt/onnx/quantization/extensions.py +35 -0
  25. modelopt/onnx/quantization/fp8.py +375 -0
  26. modelopt/onnx/quantization/graph_utils.py +1479 -0
  27. modelopt/onnx/quantization/gs_patching.py +112 -0
  28. modelopt/onnx/quantization/int4.py +1347 -0
  29. modelopt/onnx/quantization/int8.py +283 -0
  30. modelopt/onnx/quantization/operators.py +121 -0
  31. modelopt/onnx/quantization/ort_patching.py +1766 -0
  32. modelopt/onnx/quantization/ort_utils.py +346 -0
  33. modelopt/onnx/quantization/partitioning.py +439 -0
  34. modelopt/onnx/quantization/qdq_utils.py +1349 -0
  35. modelopt/onnx/quantization/quant_utils.py +283 -0
  36. modelopt/onnx/quantization/quantize.py +528 -0
  37. modelopt/onnx/quantization/src/modelopt_round_and_pack_ext.cpp +127 -0
  38. modelopt/onnx/trt_utils.py +424 -0
  39. modelopt/onnx/utils.py +753 -0
  40. modelopt/torch/__init__.py +41 -0
  41. modelopt/torch/_deploy/__init__.py +23 -0
  42. modelopt/torch/_deploy/_runtime/__init__.py +29 -0
  43. modelopt/torch/_deploy/_runtime/common.py +72 -0
  44. modelopt/torch/_deploy/_runtime/ort_client.py +193 -0
  45. modelopt/torch/_deploy/_runtime/registry.py +83 -0
  46. modelopt/torch/_deploy/_runtime/runtime_client.py +154 -0
  47. modelopt/torch/_deploy/_runtime/tensorrt/constants.py +108 -0
  48. modelopt/torch/_deploy/_runtime/tensorrt/engine_builder.py +324 -0
  49. modelopt/torch/_deploy/_runtime/tensorrt/hw_param_config.py +51 -0
  50. modelopt/torch/_deploy/_runtime/tensorrt/layerwise_profiling.py +178 -0
  51. modelopt/torch/_deploy/_runtime/tensorrt/parse_trtexec_log.py +151 -0
  52. modelopt/torch/_deploy/_runtime/tensorrt/tensorrt_utils.py +200 -0
  53. modelopt/torch/_deploy/_runtime/trt_client.py +219 -0
  54. modelopt/torch/_deploy/compilation.py +142 -0
  55. modelopt/torch/_deploy/device_model.py +157 -0
  56. modelopt/torch/_deploy/profiling.py +193 -0
  57. modelopt/torch/_deploy/utils/__init__.py +18 -0
  58. modelopt/torch/_deploy/utils/onnx_optimizer.py +120 -0
  59. modelopt/torch/_deploy/utils/onnx_utils.py +58 -0
  60. modelopt/torch/_deploy/utils/torch_onnx.py +593 -0
  61. modelopt/torch/distill/__init__.py +28 -0
  62. modelopt/torch/distill/config.py +130 -0
  63. modelopt/torch/distill/distillation.py +84 -0
  64. modelopt/torch/distill/distillation_model.py +318 -0
  65. modelopt/torch/distill/loss_balancers.py +140 -0
  66. modelopt/torch/distill/losses.py +275 -0
  67. modelopt/torch/distill/mode.py +223 -0
  68. modelopt/torch/distill/plugins/__init__.py +24 -0
  69. modelopt/torch/distill/plugins/huggingface.py +110 -0
  70. modelopt/torch/distill/plugins/megatron.py +476 -0
  71. modelopt/torch/distill/registry.py +23 -0
  72. modelopt/torch/export/__init__.py +24 -0
  73. modelopt/torch/export/convert_hf_config.py +117 -0
  74. modelopt/torch/export/distribute.py +300 -0
  75. modelopt/torch/export/hf_config_map.py +84 -0
  76. modelopt/torch/export/layer_utils.py +1893 -0
  77. modelopt/torch/export/mcore_config_map.py +25 -0
  78. modelopt/torch/export/model_config.py +623 -0
  79. modelopt/torch/export/model_config_export.py +552 -0
  80. modelopt/torch/export/model_config_utils.py +392 -0
  81. modelopt/torch/export/model_utils.py +71 -0
  82. modelopt/torch/export/plugins/__init__.py +21 -0
  83. modelopt/torch/export/plugins/mcore_common.py +67 -0
  84. modelopt/torch/export/plugins/mcore_custom.py +530 -0
  85. modelopt/torch/export/plugins/mcore_deepseek.py +143 -0
  86. modelopt/torch/export/plugins/mcore_gptoss.py +71 -0
  87. modelopt/torch/export/plugins/mcore_llama.py +164 -0
  88. modelopt/torch/export/plugins/mcore_nemotron.py +90 -0
  89. modelopt/torch/export/plugins/mcore_qwen.py +99 -0
  90. modelopt/torch/export/plugins/megatron_importer.py +647 -0
  91. modelopt/torch/export/postprocess.py +847 -0
  92. modelopt/torch/export/quant_utils.py +1069 -0
  93. modelopt/torch/export/tensorrt_llm_type.py +42 -0
  94. modelopt/torch/export/tensorrt_llm_utils.py +405 -0
  95. modelopt/torch/export/transformer_engine.py +94 -0
  96. modelopt/torch/export/unified_export_hf.py +537 -0
  97. modelopt/torch/export/unified_export_megatron.py +1164 -0
  98. modelopt/torch/nas/__init__.py +27 -0
  99. modelopt/torch/nas/algorithms.py +856 -0
  100. modelopt/torch/nas/autonas.py +808 -0
  101. modelopt/torch/nas/conversion.py +119 -0
  102. modelopt/torch/nas/hparams/__init__.py +19 -0
  103. modelopt/torch/nas/hparams/concat.py +319 -0
  104. modelopt/torch/nas/hparams/container.py +48 -0
  105. modelopt/torch/nas/modules/__init__.py +21 -0
  106. modelopt/torch/nas/modules/container.py +99 -0
  107. modelopt/torch/nas/modules/conv.py +254 -0
  108. modelopt/torch/nas/modules/linear.py +76 -0
  109. modelopt/torch/nas/modules/norm.py +193 -0
  110. modelopt/torch/nas/modules/utils.py +61 -0
  111. modelopt/torch/nas/patch.py +261 -0
  112. modelopt/torch/nas/plugins/__init__.py +29 -0
  113. modelopt/torch/nas/plugins/megatron.py +1497 -0
  114. modelopt/torch/nas/plugins/torch.py +71 -0
  115. modelopt/torch/nas/plugins/transformer_engine.py +42 -0
  116. modelopt/torch/nas/plugins/transformers.py +210 -0
  117. modelopt/torch/nas/registry.py +23 -0
  118. modelopt/torch/nas/search_space.py +260 -0
  119. modelopt/torch/nas/traced_hp.py +149 -0
  120. modelopt/torch/nas/utils.py +408 -0
  121. modelopt/torch/opt/__init__.py +41 -0
  122. modelopt/torch/opt/_hooks.py +98 -0
  123. modelopt/torch/opt/config.py +383 -0
  124. modelopt/torch/opt/conversion.py +620 -0
  125. modelopt/torch/opt/dynamic.py +1367 -0
  126. modelopt/torch/opt/hparam.py +275 -0
  127. modelopt/torch/opt/mode.py +353 -0
  128. modelopt/torch/opt/plugins/__init__.py +38 -0
  129. modelopt/torch/opt/plugins/diffusers.py +40 -0
  130. modelopt/torch/opt/plugins/huggingface.py +162 -0
  131. modelopt/torch/opt/plugins/mcore_dist_checkpointing.py +211 -0
  132. modelopt/torch/opt/plugins/megatron.py +142 -0
  133. modelopt/torch/opt/plugins/peft.py +108 -0
  134. modelopt/torch/opt/plugins/transformers.py +160 -0
  135. modelopt/torch/opt/searcher.py +363 -0
  136. modelopt/torch/opt/utils.py +125 -0
  137. modelopt/torch/prune/__init__.py +26 -0
  138. modelopt/torch/prune/fastnas.py +376 -0
  139. modelopt/torch/prune/gradnas.py +349 -0
  140. modelopt/torch/prune/plugins/__init__.py +24 -0
  141. modelopt/torch/prune/plugins/mcore_minitron.py +263 -0
  142. modelopt/torch/prune/plugins/transformers.py +41 -0
  143. modelopt/torch/prune/pruning.py +211 -0
  144. modelopt/torch/quantization/__init__.py +26 -0
  145. modelopt/torch/quantization/algorithms.py +671 -0
  146. modelopt/torch/quantization/backends/__init__.py +20 -0
  147. modelopt/torch/quantization/backends/fp8_per_tensor_gemm.py +204 -0
  148. modelopt/torch/quantization/backends/gemm_registry.py +156 -0
  149. modelopt/torch/quantization/backends/nvfp4_gemm.py +201 -0
  150. modelopt/torch/quantization/backends/utils.py +28 -0
  151. modelopt/torch/quantization/calib/__init__.py +25 -0
  152. modelopt/torch/quantization/calib/bias.py +172 -0
  153. modelopt/torch/quantization/calib/calibrator.py +63 -0
  154. modelopt/torch/quantization/calib/histogram.py +434 -0
  155. modelopt/torch/quantization/calib/max.py +110 -0
  156. modelopt/torch/quantization/compress.py +227 -0
  157. modelopt/torch/quantization/config.py +1180 -0
  158. modelopt/torch/quantization/conversion.py +389 -0
  159. modelopt/torch/quantization/export_onnx.py +639 -0
  160. modelopt/torch/quantization/extensions.py +89 -0
  161. modelopt/torch/quantization/mode.py +428 -0
  162. modelopt/torch/quantization/model_calib.py +949 -0
  163. modelopt/torch/quantization/model_quant.py +477 -0
  164. modelopt/torch/quantization/nn/__init__.py +26 -0
  165. modelopt/torch/quantization/nn/functional.py +109 -0
  166. modelopt/torch/quantization/nn/modules/__init__.py +16 -0
  167. modelopt/torch/quantization/nn/modules/quant_activations.py +22 -0
  168. modelopt/torch/quantization/nn/modules/quant_batchnorm.py +24 -0
  169. modelopt/torch/quantization/nn/modules/quant_conv.py +123 -0
  170. modelopt/torch/quantization/nn/modules/quant_instancenorm.py +39 -0
  171. modelopt/torch/quantization/nn/modules/quant_linear.py +258 -0
  172. modelopt/torch/quantization/nn/modules/quant_module.py +211 -0
  173. modelopt/torch/quantization/nn/modules/quant_pooling.py +99 -0
  174. modelopt/torch/quantization/nn/modules/quant_rnn.py +527 -0
  175. modelopt/torch/quantization/nn/modules/tensor_quantizer.py +1293 -0
  176. modelopt/torch/quantization/plugins/__init__.py +73 -0
  177. modelopt/torch/quantization/plugins/accelerate.py +214 -0
  178. modelopt/torch/quantization/plugins/apex.py +52 -0
  179. modelopt/torch/quantization/plugins/attention.py +312 -0
  180. modelopt/torch/quantization/plugins/custom.py +192 -0
  181. modelopt/torch/quantization/plugins/diffusers.py +237 -0
  182. modelopt/torch/quantization/plugins/fairscale.py +45 -0
  183. modelopt/torch/quantization/plugins/huggingface.py +736 -0
  184. modelopt/torch/quantization/plugins/megatron.py +462 -0
  185. modelopt/torch/quantization/plugins/peft.py +148 -0
  186. modelopt/torch/quantization/plugins/transformer_engine.py +60 -0
  187. modelopt/torch/quantization/plugins/transformers.py +52 -0
  188. modelopt/torch/quantization/plugins/transformers_trainer.py +426 -0
  189. modelopt/torch/quantization/plugins/trl.py +27 -0
  190. modelopt/torch/quantization/plugins/vllm.py +199 -0
  191. modelopt/torch/quantization/qtensor/__init__.py +24 -0
  192. modelopt/torch/quantization/qtensor/base_qtensor.py +370 -0
  193. modelopt/torch/quantization/qtensor/fp8_tensor.py +151 -0
  194. modelopt/torch/quantization/qtensor/int4_tensor.py +130 -0
  195. modelopt/torch/quantization/qtensor/int8_tensor.py +124 -0
  196. modelopt/torch/quantization/qtensor/mxfp4_tensor.py +144 -0
  197. modelopt/torch/quantization/qtensor/nf4_tensor.py +199 -0
  198. modelopt/torch/quantization/qtensor/nvfp4_tensor.py +295 -0
  199. modelopt/torch/quantization/src/tensor_quant.cpp +77 -0
  200. modelopt/torch/quantization/src/tensor_quant.h +69 -0
  201. modelopt/torch/quantization/src/tensor_quant_fp8.cpp +48 -0
  202. modelopt/torch/quantization/src/tensor_quant_gpu.cu +361 -0
  203. modelopt/torch/quantization/src/tensor_quant_gpu_fp8.cu +122 -0
  204. modelopt/torch/quantization/src/tensor_quant_mx.cu +411 -0
  205. modelopt/torch/quantization/src/tensor_quant_mx.h +247 -0
  206. modelopt/torch/quantization/tensor_quant.py +816 -0
  207. modelopt/torch/quantization/triton/__init__.py +36 -0
  208. modelopt/torch/quantization/triton/fp4_kernel.py +303 -0
  209. modelopt/torch/quantization/utils.py +443 -0
  210. modelopt/torch/sparsity/__init__.py +19 -0
  211. modelopt/torch/sparsity/config.py +51 -0
  212. modelopt/torch/sparsity/magnitude.py +148 -0
  213. modelopt/torch/sparsity/mode.py +203 -0
  214. modelopt/torch/sparsity/module.py +87 -0
  215. modelopt/torch/sparsity/plugins/__init__.py +27 -0
  216. modelopt/torch/sparsity/plugins/megatron.py +93 -0
  217. modelopt/torch/sparsity/searcher.py +84 -0
  218. modelopt/torch/sparsity/sparsegpt.py +276 -0
  219. modelopt/torch/sparsity/sparsification.py +123 -0
  220. modelopt/torch/speculative/__init__.py +20 -0
  221. modelopt/torch/speculative/config.py +100 -0
  222. modelopt/torch/speculative/eagle/__init__.py +20 -0
  223. modelopt/torch/speculative/eagle/conversion.py +67 -0
  224. modelopt/torch/speculative/eagle/default_config.py +50 -0
  225. modelopt/torch/speculative/eagle/eagle_model.py +51 -0
  226. modelopt/torch/speculative/eagle/utils.py +91 -0
  227. modelopt/torch/speculative/medusa/__init__.py +19 -0
  228. modelopt/torch/speculative/medusa/conversion.py +59 -0
  229. modelopt/torch/speculative/medusa/medusa_model.py +32 -0
  230. modelopt/torch/speculative/mode.py +86 -0
  231. modelopt/torch/speculative/plugins/__init__.py +33 -0
  232. modelopt/torch/speculative/plugins/megatron_eagle.py +2119 -0
  233. modelopt/torch/speculative/plugins/megatron_medusa.py +312 -0
  234. modelopt/torch/speculative/plugins/transformers.py +1138 -0
  235. modelopt/torch/speculative/speculative_decoding.py +59 -0
  236. modelopt/torch/speculative/utils.py +364 -0
  237. modelopt/torch/trace/__init__.py +25 -0
  238. modelopt/torch/trace/analyzer.py +1368 -0
  239. modelopt/torch/trace/modules/__init__.py +19 -0
  240. modelopt/torch/trace/modules/concat.py +384 -0
  241. modelopt/torch/trace/modules/nn.py +153 -0
  242. modelopt/torch/trace/plugins/__init__.py +34 -0
  243. modelopt/torch/trace/plugins/megatron.py +36 -0
  244. modelopt/torch/trace/plugins/transformers.py +58 -0
  245. modelopt/torch/trace/symbols.py +545 -0
  246. modelopt/torch/trace/tracer.py +331 -0
  247. modelopt/torch/utils/__init__.py +28 -0
  248. modelopt/torch/utils/_pytree.py +133 -0
  249. modelopt/torch/utils/cpp_extension.py +91 -0
  250. modelopt/torch/utils/dataset_utils.py +484 -0
  251. modelopt/torch/utils/distributed.py +311 -0
  252. modelopt/torch/utils/graph.py +123 -0
  253. modelopt/torch/utils/image_processor.py +112 -0
  254. modelopt/torch/utils/import_utils.py +35 -0
  255. modelopt/torch/utils/list.py +54 -0
  256. modelopt/torch/utils/logging.py +127 -0
  257. modelopt/torch/utils/memory_monitor.py +146 -0
  258. modelopt/torch/utils/network.py +658 -0
  259. modelopt/torch/utils/perf.py +174 -0
  260. modelopt/torch/utils/plugins/__init__.py +27 -0
  261. modelopt/torch/utils/plugins/megatron_generate.py +299 -0
  262. modelopt/torch/utils/plugins/megatron_mmlu.py +152 -0
  263. modelopt/torch/utils/plugins/megatron_preprocess_data.py +200 -0
  264. modelopt/torch/utils/random.py +163 -0
  265. modelopt/torch/utils/speech_dataset_utils.py +134 -0
  266. modelopt/torch/utils/tensor.py +87 -0
  267. modelopt/torch/utils/vlm_dataset_utils.py +110 -0
  268. nvidia_modelopt-0.35.0.dist-info/METADATA +171 -0
  269. nvidia_modelopt-0.35.0.dist-info/RECORD +272 -0
  270. nvidia_modelopt-0.35.0.dist-info/WHEEL +5 -0
  271. nvidia_modelopt-0.35.0.dist-info/licenses/LICENSE +14 -0
  272. 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