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,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