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,203 @@
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
+ """Sparsity mode descriptor."""
17
+
18
+ from torch import nn
19
+
20
+ from modelopt.torch.opt.config import ModeloptBaseConfig
21
+ from modelopt.torch.opt.conversion import ApplyModeError
22
+ from modelopt.torch.opt.dynamic import DynamicSpace
23
+ from modelopt.torch.opt.mode import (
24
+ ConvertEntrypoint,
25
+ ConvertReturnType,
26
+ MetadataDict,
27
+ ModeDescriptor,
28
+ RestoreEntrypoint,
29
+ UpdateEntrypoint,
30
+ _ModeRegistryCls,
31
+ )
32
+ from modelopt.torch.opt.searcher import BaseSearcher
33
+ from modelopt.torch.utils import compare_dict, unwrap_model
34
+
35
+ from .config import ExportSparseConfig, SparseGPTConfig, SparseMagnitudeConfig
36
+ from .magnitude import MagnitudeSearcher
37
+ from .module import SpDMRegistry
38
+ from .sparsegpt import SparseGPTSearcher
39
+
40
+ SparsityModeRegistry = _ModeRegistryCls("sparsity")
41
+
42
+
43
+ def convert_sparse_model(model: nn.Module, config: ModeloptBaseConfig) -> ConvertReturnType:
44
+ """Function for converting a model to a sparsity meta-model."""
45
+ # we use the search space utility here with a custom registry to convert the model
46
+ dynamic_space = DynamicSpace(model)
47
+ dynamic_space.convert_to_dynamic(config.model_dump(), SpDMRegistry)
48
+
49
+ return dynamic_space.model, {"subnet_config": DynamicSpace(model).config()}
50
+
51
+
52
+ def restore_sparse_model(
53
+ model: nn.Module, config: ModeloptBaseConfig, metadata: MetadataDict
54
+ ) -> nn.Module:
55
+ """Function for restoring a previously convert model to a sparsity meta-model."""
56
+ model, _ = convert_sparse_model(model, config)
57
+
58
+ if "subnet_config" in metadata:
59
+ DynamicSpace(model).select(metadata["subnet_config"])
60
+
61
+ return model
62
+
63
+
64
+ def update_sparse_metadata(
65
+ model: nn.Module, config: ModeloptBaseConfig, metadata: MetadataDict
66
+ ) -> None:
67
+ """Update subnet config to current subnet config of model."""
68
+ metadata["subnet_config"] = DynamicSpace(model).config()
69
+
70
+
71
+ def export_sparse(model: nn.Module, config: ExportSparseConfig) -> ConvertReturnType:
72
+ """Export a sparse model to a regular model."""
73
+ # sanity check to avoid DP/DDP here in the entrypoint
74
+ model = unwrap_model(model, raise_error=True)
75
+
76
+ # store config from model if we can find it for a future convert/restore process
77
+ metadata = {"subnet_config": DynamicSpace(model).config()}
78
+
79
+ # export model in-place
80
+ model = DynamicSpace(model).export(SpDMRegistry)
81
+
82
+ return model, metadata
83
+
84
+
85
+ def restore_export_sparse(
86
+ model: nn.Module, config: ExportSparseConfig, metadata: MetadataDict
87
+ ) -> nn.Module:
88
+ """Restore & export a sparse model to a regular model."""
89
+ # select activated/deactivated sparse modules
90
+ DynamicSpace(model).select(metadata["subnet_config"])
91
+
92
+ # run export
93
+ model, metadata_new = export_sparse(model, config)
94
+
95
+ # double check metadata
96
+ unmatched_keys = compare_dict(metadata, metadata_new)
97
+ if unmatched_keys:
98
+ raise ApplyModeError(f"Unmatched metadata={unmatched_keys}!")
99
+
100
+ return model
101
+
102
+
103
+ @SparsityModeRegistry.register_mode
104
+ class SparseMagnitudeModeDescriptor(ModeDescriptor):
105
+ """Class to define and describe magnitude-based sparsification."""
106
+
107
+ @property
108
+ def name(self) -> str:
109
+ """Returns the name of the mode."""
110
+ return "sparse_magnitude"
111
+
112
+ @property
113
+ def config_class(self) -> type[ModeloptBaseConfig]:
114
+ """Specifies the config class for the mode."""
115
+ return SparseMagnitudeConfig
116
+
117
+ @property
118
+ def next_modes(self) -> set[str] | None:
119
+ """Specifies the next modes for the mode."""
120
+ return {"export_sparse", "kd_loss", "quantize"}
121
+
122
+ @property
123
+ def export_mode(self) -> str | None:
124
+ """The mode that corresponds to the export mode of this mode."""
125
+ return "export_sparse"
126
+
127
+ @property
128
+ def search_algorithm(self) -> type[BaseSearcher]:
129
+ """Specifies the search algorithm for the mode."""
130
+ return MagnitudeSearcher
131
+
132
+ @property
133
+ def convert(self) -> ConvertEntrypoint:
134
+ """The mode's entrypoint for converting a model."""
135
+ return convert_sparse_model
136
+
137
+ @property
138
+ def restore(self) -> RestoreEntrypoint:
139
+ """The mode's entrypoint for restoring a model."""
140
+ return restore_sparse_model
141
+
142
+ @property
143
+ def update_for_save(self) -> UpdateEntrypoint:
144
+ """The mode's entrypoint for updating the models metadata."""
145
+ return update_sparse_metadata
146
+
147
+ @property
148
+ def update_for_new_mode(self) -> UpdateEntrypoint:
149
+ """The mode's entrypoint for updating the models metadata."""
150
+ return update_sparse_metadata
151
+
152
+
153
+ @SparsityModeRegistry.register_mode
154
+ class SparseGPTModeDescriptor(SparseMagnitudeModeDescriptor):
155
+ """Class to define and describe sparsification based on SparseGPT."""
156
+
157
+ @property
158
+ def name(self) -> str:
159
+ """Returns the name of the mode."""
160
+ return "sparsegpt"
161
+
162
+ @property
163
+ def config_class(self) -> type[ModeloptBaseConfig]:
164
+ """Specifies the config class for the mode."""
165
+ return SparseGPTConfig
166
+
167
+ @property
168
+ def search_algorithm(self) -> type[BaseSearcher]:
169
+ """Specifies the search algorithm for the mode."""
170
+ return SparseGPTSearcher
171
+
172
+
173
+ @SparsityModeRegistry.register_mode
174
+ class ExportSparseModeDescriptor(ModeDescriptor):
175
+ """Class to describe the ``"export_sparse"`` mode.
176
+
177
+ The properties of this mode can be inspected via the source code.
178
+ """
179
+
180
+ @property
181
+ def name(self) -> str:
182
+ """Returns the value (str representation) of the mode."""
183
+ return "export_sparse"
184
+
185
+ @property
186
+ def config_class(self) -> type[ModeloptBaseConfig]:
187
+ """Specifies the config class for the mode."""
188
+ return ExportSparseConfig
189
+
190
+ @property
191
+ def is_export_mode(self) -> bool:
192
+ """Specifies if this mode is an export mode."""
193
+ return True
194
+
195
+ @property
196
+ def convert(self) -> ConvertEntrypoint:
197
+ """The mode's entrypoint for converting a model."""
198
+ return export_sparse
199
+
200
+ @property
201
+ def restore(self) -> RestoreEntrypoint:
202
+ """The mode's entrypoint for restoring a model."""
203
+ return restore_export_sparse
@@ -0,0 +1,87 @@
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
+ """Dynamic class for all sparse modules."""
17
+
18
+ import torch
19
+ from torch import nn
20
+
21
+ from modelopt.torch.opt.dynamic import DynamicModule, _DMRegistryCls
22
+ from modelopt.torch.opt.hparam import Hparam
23
+
24
+ __all__ = ["SpDMRegistry", "SparseModule"]
25
+
26
+
27
+ SpDMRegistry = _DMRegistryCls(prefix="Sparse") # global instance for the sparsity registry
28
+
29
+
30
+ @SpDMRegistry.register({nn.Linear: "nn.Linear", nn.Conv2d: "nn.Conv2d"})
31
+ class SparseModule(DynamicModule):
32
+ """Base dynamic class for all sparse modules."""
33
+
34
+ @staticmethod
35
+ def _get_weight(mod: "SparseModule", weight: torch.Tensor) -> torch.Tensor:
36
+ if mod.is_sparse and mod._weight_mask is not None:
37
+ masked_weight = weight * mod._weight_mask
38
+ # Quick workaround for the custom attribute for Megatron.
39
+ # TODO: maybe we need a more general way for customized attributes
40
+ if hasattr(weight, "main_grad"):
41
+ masked_weight.main_grad = weight.main_grad
42
+ return masked_weight
43
+ return weight
44
+
45
+ def _setup(self):
46
+ # define hparam to check if sparsity is activated
47
+ hp = Hparam([0, -1], original=0)
48
+ hp.active = 0
49
+ self._register_hparam("is_sparse", hp)
50
+
51
+ # define the sparse mask here (don't pre-allocate memory to maximize memory savings)
52
+ self._register_temp_attribute("_weight_mask", None, lambda m, n, v: m.register_buffer(n, v))
53
+
54
+ # register dynamic attributes of the class
55
+ self._register_dynamic_attribute("weight", self._get_weight)
56
+
57
+ def modify(self, *args, **kwargs):
58
+ """Initialize the sparsity mask when this is called.
59
+
60
+ Note that for any module that is not frozen via ``None`` in the rules, this function will be
61
+ called. Hence, we use this function to initialize the sparsity mask only when necessary.
62
+ """
63
+ hp = self.get_hparam("is_sparse")
64
+ if -1 in hp.choices and self._weight_mask is None:
65
+ hp.active = -1
66
+ self._weight_mask = torch.ones_like(self.weight, dtype=torch.bool)
67
+
68
+ def set_mask(self, value: torch.BoolTensor | None):
69
+ """Set the active sparse mask of the module weights."""
70
+ if value is None:
71
+ self._weight_mask = None
72
+ return
73
+
74
+ # sanity checks on the mask
75
+ w_shape = self.weight.shape
76
+ assert value is not None, "Mask cannot be None."
77
+ assert value.shape == w_shape, f"Mask must have shape {w_shape}, got {value.shape} instead."
78
+ assert value.dtype == torch.bool, f"Mask must be of type torch.bool, but got {value.dtype}."
79
+
80
+ # assign mask
81
+ with torch.no_grad():
82
+ if torch.all(value):
83
+ self._weight_mask = None
84
+ elif self._weight_mask is None:
85
+ self._weight_mask = value.detach().clone().to(self.weight.device)
86
+ else:
87
+ self._weight_mask.copy_(value.to(self._weight_mask.device))
@@ -0,0 +1,27 @@
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
+ """Handles sparsity plugins for third-party modules.
17
+
18
+ Currently, we support plugins for
19
+
20
+ - :meth:`megatron<modelopt.torch.sparsity.plugins.megatron>`
21
+
22
+ """
23
+
24
+ from modelopt.torch.utils import import_plugin
25
+
26
+ with import_plugin("megatron"):
27
+ from .megatron import *
@@ -0,0 +1,93 @@
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
+ """Support sparsify and save/resore for Megatron."""
17
+
18
+ import megatron.core.transformer.mlp as megatron_mlp
19
+ from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
20
+ from megatron.core.transformer.utils import make_sharded_tensors_for_checkpoint
21
+
22
+ from modelopt.torch.opt.plugins.megatron import _MegatronMLP
23
+
24
+ from ..config import SparseGPTConfig, SparseMagnitudeConfig
25
+ from ..module import SparseModule, SpDMRegistry
26
+
27
+
28
+ class _MegatronParallelLinear(SparseModule):
29
+ def _get_shard_axis_dict(self, state_dict):
30
+ raise NotImplementedError
31
+
32
+ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None):
33
+ sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets)
34
+
35
+ sparse_state_dict = {
36
+ k: v
37
+ for k, v in self.state_dict(prefix="", keep_vars=True).items()
38
+ if k == "_weight_mask"
39
+ }
40
+
41
+ sharded_axis_dict = self._get_shard_axis_dict(sparse_state_dict)
42
+
43
+ if sparse_state_dict:
44
+ sharded_state_dict.update(
45
+ **make_sharded_tensors_for_checkpoint(
46
+ sparse_state_dict, prefix, sharded_axis_dict, sharded_offsets
47
+ )
48
+ )
49
+
50
+ return sharded_state_dict
51
+
52
+
53
+ @SpDMRegistry.register(
54
+ {ColumnParallelLinear: "megatron.core.tensor_parallel.layers.ColumnParallelLinear"}
55
+ )
56
+ class _MegatronColumnParallelLinear(_MegatronParallelLinear):
57
+ def _get_shard_axis_dict(self, state_dict):
58
+ return {"_weight_mask": 0}
59
+
60
+
61
+ @SpDMRegistry.register(
62
+ {RowParallelLinear: "megatron.core.tensor_parallel.layers.RowParallelLinear"}
63
+ )
64
+ class _MegatronRowParallelLinear(_MegatronParallelLinear):
65
+ def _get_shard_axis_dict(self, state_dict):
66
+ return {"_weight_mask": 1}
67
+
68
+
69
+ @SpDMRegistry.register({megatron_mlp.MLP: "megatron.core.transformer.mlp.MLP"})
70
+ class _SparseMegatronMLP(_MegatronMLP):
71
+ """Module to support special handling of `linear_fc1` in `sharded_state_dict()` of MCore `MLP`."""
72
+
73
+ _modelopt_state_keys = [r"\._weight_mask$"]
74
+
75
+
76
+ def _get_extra_rules():
77
+ """Get the extra rules for megatron."""
78
+ return {
79
+ "megatron.core.tensor_parallel.layers.ColumnParallelLinear": {
80
+ "*": {},
81
+ "*output_layer*": None,
82
+ },
83
+ "megatron.core.tensor_parallel.layers.RowParallelLinear": {
84
+ "*": {},
85
+ "*output_layer*": None,
86
+ },
87
+ "megatron.core.transformer.mlp.MLP": {},
88
+ }
89
+
90
+
91
+ # Update the default rules
92
+ SparseMagnitudeConfig.register_default(_get_extra_rules())
93
+ SparseGPTConfig.register_default(_get_extra_rules())
@@ -0,0 +1,84 @@
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
+ """Searcher interface for sparsity algorithms."""
17
+
18
+ from abc import abstractmethod
19
+ from collections.abc import Iterator
20
+
21
+ import torch
22
+ import torch.nn as nn
23
+
24
+ from modelopt.torch.opt.searcher import BaseSearcher, SearchConfig, SearchStateDict
25
+ from modelopt.torch.utils import print_rank_0
26
+
27
+ from . import magnitude as asp
28
+ from .module import SparseModule
29
+
30
+
31
+ class BaseSparseSearcher(BaseSearcher):
32
+ """A generic sparse mask searching algorithm."""
33
+
34
+ _pattern_2_4 = "2:4 sparsity"
35
+
36
+ @property
37
+ def default_search_config(self) -> SearchConfig:
38
+ """Get the default config for the searcher."""
39
+ return {**super().default_search_config, "pattern": self._pattern_2_4}
40
+
41
+ @property
42
+ def default_state_dict(self) -> SearchStateDict:
43
+ """Return default state dict."""
44
+ return {}
45
+
46
+ def sanitize_search_config(self, config: SearchConfig | None) -> SearchStateDict:
47
+ """Sanitize the search config dict."""
48
+ config_sanitized = super().sanitize_search_config(config)
49
+
50
+ # sanity check of sparsity format
51
+ is_nm_prune, n, m = asp.get_nmprune_info(config_sanitized["pattern"])
52
+ assert is_nm_prune and n == 2 and m == 4, (
53
+ f"Unsupported pattern {self.config['pattern']} for sparsity"
54
+ )
55
+
56
+ return config_sanitized
57
+
58
+ @abstractmethod
59
+ def _check_weight_size(self, weight, mod_name) -> bool:
60
+ """Check if the weight size is supported by the algorithm."""
61
+ raise NotImplementedError
62
+
63
+ @abstractmethod
64
+ def _compute_mask(self, module: SparseModule) -> torch.BoolTensor:
65
+ """Compute the mask and update weight for a given module."""
66
+ raise NotImplementedError
67
+
68
+ def _named_sparsifiable_modules(self) -> Iterator[tuple[str, nn.Module]]:
69
+ """Get the named sparsifiable modules."""
70
+ for name, module in self.model.named_modules():
71
+ if (
72
+ isinstance(module, SparseModule)
73
+ and module.is_sparse
74
+ and self._check_weight_size(module.weight, name)
75
+ ):
76
+ yield name, module
77
+
78
+ def run_search(self):
79
+ """Search for sparse mask."""
80
+ for name, module in self._named_sparsifiable_modules():
81
+ # compute the mask (and potentially weight update inside compute_mask)
82
+ print_rank_0(f"Searching for sparse mask and weight update for module {name}.")
83
+ with torch.no_grad():
84
+ module.set_mask(self._compute_mask(module))