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,71 @@
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
+ """Custom mapping from GPT-OSS Hugging Face models to Megatron Core models."""
17
+
18
+ from .mcore_custom import (
19
+ COL_TP,
20
+ PACK_EP,
21
+ REPLICATE,
22
+ ROW_TP,
23
+ CustomModuleMapping,
24
+ NameRemapping,
25
+ PackNameRemappingGPT,
26
+ QKVMerging,
27
+ QKVSlicing,
28
+ UnpackNameRemappingGPT,
29
+ )
30
+
31
+ gptoss_causal_lm_export: dict[str, CustomModuleMapping | bool] = {
32
+ "word_embeddings": NameRemapping("model.embed_tokens."),
33
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."),
34
+ "linear_qkv": QKVSlicing("model.layers.{}.self_attn."),
35
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."),
36
+ "softmax_offset": NameRemapping("model.layers.{}.self_attn.sinks"),
37
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."),
38
+ "use_packed_local_experts": True,
39
+ "local_experts.linear_fc1": PackNameRemappingGPT(
40
+ "model.layers.{}.mlp.experts.gate_up_proj",
41
+ {"layer_type": "linear_fc1"},
42
+ ),
43
+ "local_experts.linear_fc2": PackNameRemappingGPT(
44
+ "model.layers.{}.mlp.experts.down_proj",
45
+ {"layer_type": "linear_fc2"},
46
+ ),
47
+ "router": NameRemapping("model.layers.{}.mlp.router."),
48
+ "final_layernorm": NameRemapping("model.norm."),
49
+ "output_layer": NameRemapping("lm_head."),
50
+ }
51
+
52
+ gptoss_causal_lm_import: dict[str, CustomModuleMapping | bool] = {
53
+ "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP),
54
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE),
55
+ "linear_qkv": QKVMerging("model.layers.{}.self_attn.", COL_TP),
56
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP),
57
+ "softmax_offset": NameRemapping("model.layers.{}.self_attn.sinks", COL_TP),
58
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE),
59
+ "router": NameRemapping("model.layers.{}.mlp.router.", REPLICATE),
60
+ "use_packed_local_experts": True,
61
+ "local_experts.linear_fc1_ep": UnpackNameRemappingGPT(
62
+ "model.layers.{}.mlp.experts.gate_up_proj",
63
+ PACK_EP | {"layer_type": "linear_fc1"},
64
+ ),
65
+ "local_experts.linear_fc2_ep": UnpackNameRemappingGPT(
66
+ "model.layers.{}.mlp.experts.down_proj",
67
+ PACK_EP | {"layer_type": "linear_fc2"},
68
+ ),
69
+ "final_layernorm": NameRemapping("model.norm.", REPLICATE),
70
+ "output_layer": NameRemapping("lm_head.", COL_TP),
71
+ }
@@ -0,0 +1,164 @@
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
+
17
+ """Custom mapping from Llama Hugging Face models to Megatron Core models."""
18
+
19
+ from .mcore_custom import (
20
+ COL_TP,
21
+ PACK_COL_ETP,
22
+ PACK_EP,
23
+ PACK_ROW_ETP,
24
+ REPLICATE,
25
+ ROW_TP,
26
+ CustomModuleMapping,
27
+ GatedMLPMerging,
28
+ GatedMLPSlicing,
29
+ NameRemapping,
30
+ PackNameRemapping,
31
+ QKVMerging,
32
+ QKVSlicing,
33
+ UnpackNameRemapping,
34
+ )
35
+
36
+ llama_causal_lm_export: dict[str, CustomModuleMapping] = {
37
+ "word_embeddings": NameRemapping("model.embed_tokens."),
38
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."),
39
+ "linear_qkv": QKVSlicing("model.layers.{}.self_attn."),
40
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."),
41
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."),
42
+ "linear_fc1": GatedMLPSlicing("model.layers.{}.mlp."),
43
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."),
44
+ "final_layernorm": NameRemapping("model.norm."),
45
+ "output_layer": NameRemapping("lm_head."),
46
+ }
47
+
48
+ llama4_causal_lm_export: dict[str, CustomModuleMapping | bool] = {
49
+ "word_embeddings": NameRemapping("language_model.model.embed_tokens."),
50
+ "input_layernorm": NameRemapping("language_model.model.layers.{}.input_layernorm."),
51
+ # self_attn
52
+ "linear_qkv": QKVSlicing("language_model.model.layers.{}.self_attn."),
53
+ "linear_proj": NameRemapping("language_model.model.layers.{}.self_attn.o_proj."),
54
+ # mlp
55
+ "pre_mlp_layernorm": NameRemapping("language_model.model.layers.{}.post_attention_layernorm."),
56
+ "shared_experts.linear_fc1": GatedMLPSlicing(
57
+ "language_model.model.layers.{}.feed_forward.shared_expert.",
58
+ ),
59
+ "shared_experts.linear_fc2": NameRemapping(
60
+ "language_model.model.layers.{}.feed_forward.shared_expert.down_proj.",
61
+ ),
62
+ # moe_layer
63
+ "router": NameRemapping("language_model.model.layers.{}.feed_forward.router."),
64
+ "use_packed_local_experts": True,
65
+ "local_experts.linear_fc1": PackNameRemapping(
66
+ "language_model.model.layers.{}.feed_forward.experts.gate_up_proj",
67
+ {"layer_type": "linear_fc1"},
68
+ ),
69
+ "local_experts.linear_fc2": PackNameRemapping(
70
+ "language_model.model.layers.{}.feed_forward.experts.down_proj",
71
+ {"layer_type": "linear_fc2"},
72
+ ),
73
+ "final_layernorm": NameRemapping("language_model.model.norm."),
74
+ "output_layer": NameRemapping("language_model.lm_head."),
75
+ }
76
+
77
+ medusa_llama_causal_lm_export: dict[str, CustomModuleMapping] = {
78
+ # MedusaForCausalLM support
79
+ "lm_head": NameRemapping(
80
+ "medusa_heads.{}.1."
81
+ ), # TODO: lm_head is hardcoded to .1 as currently only support using 1 layer in medusa head
82
+ # needs a fix
83
+ "linear": NameRemapping("medusa_heads.{}.{}.linear."),
84
+ }
85
+
86
+ eagle_llama_causal_lm_export: dict[str, CustomModuleMapping] = {
87
+ "word_embeddings": NameRemapping("embed_tokens."),
88
+ "enorm": NameRemapping("enorm."),
89
+ "hnorm": NameRemapping("hnorm."),
90
+ "fc": NameRemapping("fc."),
91
+ "input_layernorm": NameRemapping("layers.{}.input_layernorm."),
92
+ "linear_qkv": QKVSlicing("layers.{}.self_attn."),
93
+ "linear_proj": NameRemapping("layers.{}.self_attn.o_proj."),
94
+ "pre_mlp_layernorm": NameRemapping("layers.{}.post_attention_layernorm."),
95
+ "linear_fc1": GatedMLPSlicing("layers.{}.mlp."),
96
+ "linear_fc2": NameRemapping("layers.{}.mlp.down_proj."),
97
+ "final_layernorm": NameRemapping("norm."),
98
+ "d2t": NameRemapping("d2t"),
99
+ "output_layer": NameRemapping("lm_head."),
100
+ }
101
+
102
+ eagle3_llama_causal_lm_export: dict[str, CustomModuleMapping] = {
103
+ "word_embeddings": NameRemapping("embed_tokens."),
104
+ "enorm": NameRemapping("midlayer.input_layernorm."),
105
+ "fc": NameRemapping("fc."),
106
+ "input_layernorm": NameRemapping("midlayer.hidden_norm."),
107
+ "linear_qkv": QKVSlicing("midlayer.self_attn."),
108
+ "linear_proj": NameRemapping("midlayer.self_attn.o_proj."),
109
+ "pre_mlp_layernorm": NameRemapping("midlayer.post_attention_layernorm."),
110
+ "linear_fc1": GatedMLPSlicing("midlayer.mlp."),
111
+ "linear_fc2": NameRemapping("midlayer.mlp.down_proj."),
112
+ "final_layernorm": NameRemapping("norm."),
113
+ "d2t": NameRemapping("d2t"),
114
+ "output_layer": NameRemapping("lm_head."),
115
+ }
116
+
117
+
118
+ llama_causal_lm_import: dict[str, CustomModuleMapping] = {
119
+ "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP),
120
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE),
121
+ "linear_qkv": QKVMerging("model.layers.{}.self_attn.", COL_TP),
122
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP),
123
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE),
124
+ "linear_fc1": GatedMLPMerging("model.layers.{}.mlp.", COL_TP),
125
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj.", ROW_TP),
126
+ "final_layernorm": NameRemapping("model.norm.", REPLICATE),
127
+ "output_layer": NameRemapping("lm_head.", COL_TP),
128
+ }
129
+
130
+ llama4_causal_lm_import: dict[str, CustomModuleMapping | bool] = {
131
+ "word_embeddings": NameRemapping("language_model.model.embed_tokens.", COL_TP),
132
+ "input_layernorm": NameRemapping("language_model.model.layers.{}.input_layernorm.", REPLICATE),
133
+ "linear_qkv": QKVMerging("language_model.model.layers.{}.self_attn.", COL_TP),
134
+ "linear_proj": NameRemapping("language_model.model.layers.{}.self_attn.o_proj.", ROW_TP),
135
+ "pre_mlp_layernorm": NameRemapping(
136
+ "language_model.model.layers.{}.post_attention_layernorm.", REPLICATE
137
+ ),
138
+ "shared_experts.linear_fc1": GatedMLPMerging(
139
+ "language_model.model.layers.{}.feed_forward.shared_expert.", COL_TP
140
+ ),
141
+ "shared_experts.linear_fc2": NameRemapping(
142
+ "language_model.model.layers.{}.feed_forward.shared_expert.down_proj.", ROW_TP
143
+ ),
144
+ "router": NameRemapping("language_model.model.layers.{}.feed_forward.router.", REPLICATE),
145
+ "use_packed_local_experts": True,
146
+ "local_experts.linear_fc1_etp": UnpackNameRemapping(
147
+ "language_model.model.layers.{}.feed_forward.experts.gate_up_proj",
148
+ PACK_COL_ETP | {"layer_type": "linear_fc1"},
149
+ ),
150
+ "local_experts.linear_fc2_etp": UnpackNameRemapping(
151
+ "language_model.model.layers.{}.feed_forward.experts.down_proj",
152
+ PACK_ROW_ETP | {"layer_type": "linear_fc2"},
153
+ ),
154
+ "local_experts.linear_fc1_ep": UnpackNameRemapping(
155
+ "language_model.model.layers.{}.feed_forward.experts.gate_up_proj",
156
+ PACK_EP | {"layer_type": "linear_fc1"},
157
+ ),
158
+ "local_experts.linear_fc2_ep": UnpackNameRemapping(
159
+ "language_model.model.layers.{}.feed_forward.experts.down_proj",
160
+ PACK_EP | {"layer_type": "linear_fc2"},
161
+ ),
162
+ "final_layernorm": NameRemapping("language_model.model.norm.", REPLICATE),
163
+ "output_layer": NameRemapping("language_model.lm_head.", COL_TP),
164
+ }
@@ -0,0 +1,90 @@
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
+
17
+ """Custom mapping from Nemotron Hugging Face models to Megatron Core models."""
18
+
19
+ from .mcore_custom import (
20
+ COL_TP,
21
+ REPLICATE,
22
+ ROW_TP,
23
+ CustomModuleMapping,
24
+ NameRemapping,
25
+ QKVMerging,
26
+ QKVSlicing,
27
+ )
28
+
29
+ # Example on adding a new CausalLM.
30
+ nemotron_causal_lm_export: dict[str, CustomModuleMapping] = {
31
+ # NemotronForCausalLM is using square-relu where no gated handle is needed.
32
+ "word_embeddings": NameRemapping("model.embed_tokens."),
33
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."),
34
+ "linear_qkv": QKVSlicing("model.layers.{}.self_attn."),
35
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."),
36
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."),
37
+ # NemotronForCausalLM is using square-relu where no gated handle is needed.
38
+ "linear_fc1": NameRemapping("model.layers.{}.mlp.up_proj."),
39
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."),
40
+ "final_layernorm": NameRemapping("model.norm."),
41
+ "output_layer": NameRemapping("lm_head."),
42
+ }
43
+
44
+
45
+ nemotron_h_causal_lm_import: dict[str, CustomModuleMapping] = {
46
+ "word_embeddings": NameRemapping("backbone.embeddings.", COL_TP),
47
+ "final_norm": NameRemapping("backbone.norm_f.", REPLICATE),
48
+ "output_layer": NameRemapping("lm_head.", COL_TP),
49
+ # Mamba
50
+ "norm": NameRemapping("backbone.layers.{}.norm.", REPLICATE),
51
+ "mixer_norm": NameRemapping("backbone.layers.{}.mixer.norm.", REPLICATE),
52
+ "A_log": NameRemapping("backbone.layers.{}.mixer.A_log", REPLICATE),
53
+ "D": NameRemapping("backbone.layers.{}.mixer.D", REPLICATE),
54
+ "dt_bias": NameRemapping("backbone.layers.{}.mixer.dt_bias", REPLICATE),
55
+ "conv1d": NameRemapping("backbone.layers.{}.mixer.conv1d.", REPLICATE),
56
+ "in_proj": NameRemapping("backbone.layers.{}.mixer.in_proj.", COL_TP),
57
+ "out_proj": NameRemapping("backbone.layers.{}.mixer.out_proj.", ROW_TP),
58
+ # Attention
59
+ "input_layernorm": NameRemapping("backbone.layers.{}.norm.", REPLICATE),
60
+ "linear_qkv": QKVMerging("backbone.layers.{}.mixer.", COL_TP),
61
+ "linear_proj": NameRemapping("backbone.layers.{}.mixer.o_proj.", ROW_TP),
62
+ # MLP
63
+ "pre_mlp_layernorm": NameRemapping("backbone.layers.{}.norm.", REPLICATE),
64
+ "linear_fc1": NameRemapping("backbone.layers.{}.mixer.up_proj.", COL_TP),
65
+ "linear_fc2": NameRemapping("backbone.layers.{}.mixer.down_proj.", ROW_TP),
66
+ }
67
+
68
+
69
+ nemotron_h_causal_lm_export: dict[str, CustomModuleMapping] = {
70
+ "word_embeddings": NameRemapping("backbone.embeddings."),
71
+ "final_norm": NameRemapping("backbone.norm_f."),
72
+ "output_layer": NameRemapping("lm_head."),
73
+ # Mamba
74
+ "norm": NameRemapping("backbone.layers.{}.norm."),
75
+ "mixer_norm": NameRemapping("backbone.layers.{}.mixer.norm."),
76
+ "A_log": NameRemapping("backbone.layers.{}.mixer.A_log"),
77
+ "D": NameRemapping("backbone.layers.{}.mixer.D"),
78
+ "dt_bias": NameRemapping("backbone.layers.{}.mixer.dt_bias"),
79
+ "conv1d": NameRemapping("backbone.layers.{}.mixer.conv1d."),
80
+ "in_proj": NameRemapping("backbone.layers.{}.mixer.in_proj."),
81
+ "out_proj": NameRemapping("backbone.layers.{}.mixer.out_proj."),
82
+ # Attention
83
+ "input_layernorm": NameRemapping("backbone.layers.{}.norm."),
84
+ "linear_qkv": QKVSlicing("backbone.layers.{}.mixer."),
85
+ "linear_proj": NameRemapping("backbone.layers.{}.mixer.o_proj."),
86
+ # MLP
87
+ "pre_mlp_layernorm": NameRemapping("backbone.layers.{}.norm."),
88
+ "linear_fc1": NameRemapping("backbone.layers.{}.mixer.up_proj."),
89
+ "linear_fc2": NameRemapping("backbone.layers.{}.mixer.down_proj."),
90
+ }
@@ -0,0 +1,99 @@
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
+ """Custom mapping from Qwen Hugging Face models to Megatron Core models."""
17
+
18
+ from .mcore_custom import (
19
+ COL_ETP,
20
+ COL_TP,
21
+ REPLICATE,
22
+ ROW_ETP,
23
+ ROW_TP,
24
+ CustomModuleMapping,
25
+ GatedMLPMerging,
26
+ GatedMLPSlicing,
27
+ NameRemapping,
28
+ QKVMerging,
29
+ QKVSlicing,
30
+ )
31
+
32
+ qwen3_causal_lm_import: dict[str, CustomModuleMapping] = {
33
+ "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP),
34
+ "final_layernorm": NameRemapping("model.norm.", REPLICATE),
35
+ "output_layer": NameRemapping("lm_head.", COL_TP),
36
+ # Attention
37
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE),
38
+ "linear_qkv": QKVMerging("model.layers.{}.self_attn.", COL_TP),
39
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP),
40
+ "q_layernorm": NameRemapping("model.layers.{}.self_attn.q_norm.", REPLICATE),
41
+ "k_layernorm": NameRemapping("model.layers.{}.self_attn.k_norm.", REPLICATE),
42
+ # MLP
43
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE),
44
+ "linear_fc1": GatedMLPMerging("model.layers.{}.mlp.", COL_TP),
45
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj.", ROW_TP),
46
+ # MoE
47
+ "router": NameRemapping("model.layers.{}.mlp.gate.", REPLICATE),
48
+ "local_experts.linear_fc1": GatedMLPMerging("model.layers.{}.mlp.experts.{}.", COL_ETP),
49
+ "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj.", ROW_ETP),
50
+ }
51
+
52
+
53
+ qwen3_causal_lm_export: dict[str, CustomModuleMapping] = {
54
+ "word_embeddings": NameRemapping("model.embed_tokens."),
55
+ "final_layernorm": NameRemapping("model.norm."),
56
+ "output_layer": NameRemapping("lm_head."),
57
+ # Attention
58
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."),
59
+ "linear_qkv": QKVSlicing("model.layers.{}.self_attn."),
60
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."),
61
+ "q_layernorm": NameRemapping("model.layers.{}.self_attn.q_norm."),
62
+ "k_layernorm": NameRemapping("model.layers.{}.self_attn.k_norm."),
63
+ # MLP
64
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."),
65
+ "linear_fc1": GatedMLPSlicing("model.layers.{}.mlp."),
66
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."),
67
+ # MoE
68
+ "router": NameRemapping("model.layers.{}.mlp.gate."),
69
+ "local_experts.linear_fc1": GatedMLPSlicing("model.layers.{}.mlp.experts.{}."),
70
+ "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj."),
71
+ }
72
+
73
+ qwen25_causal_lm_import: dict[str, CustomModuleMapping] = {
74
+ "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP),
75
+ "final_layernorm": NameRemapping("model.norm.", REPLICATE),
76
+ "output_layer": NameRemapping("lm_head.", COL_TP),
77
+ # Attention
78
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE),
79
+ "linear_qkv": QKVMerging("model.layers.{}.self_attn.", COL_TP),
80
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP),
81
+ # MLP
82
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE),
83
+ "linear_fc1": GatedMLPMerging("model.layers.{}.mlp.", COL_TP),
84
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj.", ROW_TP),
85
+ }
86
+
87
+ qwen25_causal_lm_export: dict[str, CustomModuleMapping] = {
88
+ "word_embeddings": NameRemapping("model.embed_tokens."),
89
+ "final_layernorm": NameRemapping("model.norm."),
90
+ "output_layer": NameRemapping("lm_head."),
91
+ # Attention
92
+ "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."),
93
+ "linear_qkv": QKVSlicing("model.layers.{}.self_attn."),
94
+ "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."),
95
+ # MLP
96
+ "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."),
97
+ "linear_fc1": GatedMLPSlicing("model.layers.{}.mlp."),
98
+ "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."),
99
+ }