nvidia-modelopt 0.35.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- modelopt/__init__.py +20 -0
- modelopt/deploy/__init__.py +16 -0
- modelopt/deploy/llm/__init__.py +30 -0
- modelopt/deploy/llm/generate.py +345 -0
- modelopt/onnx/__init__.py +34 -0
- modelopt/onnx/autocast/__init__.py +19 -0
- modelopt/onnx/autocast/__main__.py +187 -0
- modelopt/onnx/autocast/convert.py +202 -0
- modelopt/onnx/autocast/graphsanitizer.py +422 -0
- modelopt/onnx/autocast/logging_config.py +78 -0
- modelopt/onnx/autocast/nodeclassifier.py +416 -0
- modelopt/onnx/autocast/precisionconverter.py +1032 -0
- modelopt/onnx/autocast/referencerunner.py +159 -0
- modelopt/onnx/autocast/utils.py +117 -0
- modelopt/onnx/llm_export_utils/__init__.py +16 -0
- modelopt/onnx/llm_export_utils/export_utils.py +162 -0
- modelopt/onnx/llm_export_utils/quantization_utils.py +122 -0
- modelopt/onnx/llm_export_utils/surgeon_utils.py +120 -0
- modelopt/onnx/logging_config.py +77 -0
- modelopt/onnx/op_types.py +300 -0
- modelopt/onnx/quantization/__init__.py +20 -0
- modelopt/onnx/quantization/__main__.py +291 -0
- modelopt/onnx/quantization/calib_utils.py +178 -0
- modelopt/onnx/quantization/extensions.py +35 -0
- modelopt/onnx/quantization/fp8.py +375 -0
- modelopt/onnx/quantization/graph_utils.py +1479 -0
- modelopt/onnx/quantization/gs_patching.py +112 -0
- modelopt/onnx/quantization/int4.py +1347 -0
- modelopt/onnx/quantization/int8.py +283 -0
- modelopt/onnx/quantization/operators.py +121 -0
- modelopt/onnx/quantization/ort_patching.py +1766 -0
- modelopt/onnx/quantization/ort_utils.py +346 -0
- modelopt/onnx/quantization/partitioning.py +439 -0
- modelopt/onnx/quantization/qdq_utils.py +1349 -0
- modelopt/onnx/quantization/quant_utils.py +283 -0
- modelopt/onnx/quantization/quantize.py +528 -0
- modelopt/onnx/quantization/src/modelopt_round_and_pack_ext.cpp +127 -0
- modelopt/onnx/trt_utils.py +424 -0
- modelopt/onnx/utils.py +753 -0
- modelopt/torch/__init__.py +41 -0
- modelopt/torch/_deploy/__init__.py +23 -0
- modelopt/torch/_deploy/_runtime/__init__.py +29 -0
- modelopt/torch/_deploy/_runtime/common.py +72 -0
- modelopt/torch/_deploy/_runtime/ort_client.py +193 -0
- modelopt/torch/_deploy/_runtime/registry.py +83 -0
- modelopt/torch/_deploy/_runtime/runtime_client.py +154 -0
- modelopt/torch/_deploy/_runtime/tensorrt/constants.py +108 -0
- modelopt/torch/_deploy/_runtime/tensorrt/engine_builder.py +324 -0
- modelopt/torch/_deploy/_runtime/tensorrt/hw_param_config.py +51 -0
- modelopt/torch/_deploy/_runtime/tensorrt/layerwise_profiling.py +178 -0
- modelopt/torch/_deploy/_runtime/tensorrt/parse_trtexec_log.py +151 -0
- modelopt/torch/_deploy/_runtime/tensorrt/tensorrt_utils.py +200 -0
- modelopt/torch/_deploy/_runtime/trt_client.py +219 -0
- modelopt/torch/_deploy/compilation.py +142 -0
- modelopt/torch/_deploy/device_model.py +157 -0
- modelopt/torch/_deploy/profiling.py +193 -0
- modelopt/torch/_deploy/utils/__init__.py +18 -0
- modelopt/torch/_deploy/utils/onnx_optimizer.py +120 -0
- modelopt/torch/_deploy/utils/onnx_utils.py +58 -0
- modelopt/torch/_deploy/utils/torch_onnx.py +593 -0
- modelopt/torch/distill/__init__.py +28 -0
- modelopt/torch/distill/config.py +130 -0
- modelopt/torch/distill/distillation.py +84 -0
- modelopt/torch/distill/distillation_model.py +318 -0
- modelopt/torch/distill/loss_balancers.py +140 -0
- modelopt/torch/distill/losses.py +275 -0
- modelopt/torch/distill/mode.py +223 -0
- modelopt/torch/distill/plugins/__init__.py +24 -0
- modelopt/torch/distill/plugins/huggingface.py +110 -0
- modelopt/torch/distill/plugins/megatron.py +476 -0
- modelopt/torch/distill/registry.py +23 -0
- modelopt/torch/export/__init__.py +24 -0
- modelopt/torch/export/convert_hf_config.py +117 -0
- modelopt/torch/export/distribute.py +300 -0
- modelopt/torch/export/hf_config_map.py +84 -0
- modelopt/torch/export/layer_utils.py +1893 -0
- modelopt/torch/export/mcore_config_map.py +25 -0
- modelopt/torch/export/model_config.py +623 -0
- modelopt/torch/export/model_config_export.py +552 -0
- modelopt/torch/export/model_config_utils.py +392 -0
- modelopt/torch/export/model_utils.py +71 -0
- modelopt/torch/export/plugins/__init__.py +21 -0
- modelopt/torch/export/plugins/mcore_common.py +67 -0
- modelopt/torch/export/plugins/mcore_custom.py +530 -0
- modelopt/torch/export/plugins/mcore_deepseek.py +143 -0
- modelopt/torch/export/plugins/mcore_gptoss.py +71 -0
- modelopt/torch/export/plugins/mcore_llama.py +164 -0
- modelopt/torch/export/plugins/mcore_nemotron.py +90 -0
- modelopt/torch/export/plugins/mcore_qwen.py +99 -0
- modelopt/torch/export/plugins/megatron_importer.py +647 -0
- modelopt/torch/export/postprocess.py +847 -0
- modelopt/torch/export/quant_utils.py +1069 -0
- modelopt/torch/export/tensorrt_llm_type.py +42 -0
- modelopt/torch/export/tensorrt_llm_utils.py +405 -0
- modelopt/torch/export/transformer_engine.py +94 -0
- modelopt/torch/export/unified_export_hf.py +537 -0
- modelopt/torch/export/unified_export_megatron.py +1164 -0
- modelopt/torch/nas/__init__.py +27 -0
- modelopt/torch/nas/algorithms.py +856 -0
- modelopt/torch/nas/autonas.py +808 -0
- modelopt/torch/nas/conversion.py +119 -0
- modelopt/torch/nas/hparams/__init__.py +19 -0
- modelopt/torch/nas/hparams/concat.py +319 -0
- modelopt/torch/nas/hparams/container.py +48 -0
- modelopt/torch/nas/modules/__init__.py +21 -0
- modelopt/torch/nas/modules/container.py +99 -0
- modelopt/torch/nas/modules/conv.py +254 -0
- modelopt/torch/nas/modules/linear.py +76 -0
- modelopt/torch/nas/modules/norm.py +193 -0
- modelopt/torch/nas/modules/utils.py +61 -0
- modelopt/torch/nas/patch.py +261 -0
- modelopt/torch/nas/plugins/__init__.py +29 -0
- modelopt/torch/nas/plugins/megatron.py +1497 -0
- modelopt/torch/nas/plugins/torch.py +71 -0
- modelopt/torch/nas/plugins/transformer_engine.py +42 -0
- modelopt/torch/nas/plugins/transformers.py +210 -0
- modelopt/torch/nas/registry.py +23 -0
- modelopt/torch/nas/search_space.py +260 -0
- modelopt/torch/nas/traced_hp.py +149 -0
- modelopt/torch/nas/utils.py +408 -0
- modelopt/torch/opt/__init__.py +41 -0
- modelopt/torch/opt/_hooks.py +98 -0
- modelopt/torch/opt/config.py +383 -0
- modelopt/torch/opt/conversion.py +620 -0
- modelopt/torch/opt/dynamic.py +1367 -0
- modelopt/torch/opt/hparam.py +275 -0
- modelopt/torch/opt/mode.py +353 -0
- modelopt/torch/opt/plugins/__init__.py +38 -0
- modelopt/torch/opt/plugins/diffusers.py +40 -0
- modelopt/torch/opt/plugins/huggingface.py +162 -0
- modelopt/torch/opt/plugins/mcore_dist_checkpointing.py +211 -0
- modelopt/torch/opt/plugins/megatron.py +142 -0
- modelopt/torch/opt/plugins/peft.py +108 -0
- modelopt/torch/opt/plugins/transformers.py +160 -0
- modelopt/torch/opt/searcher.py +363 -0
- modelopt/torch/opt/utils.py +125 -0
- modelopt/torch/prune/__init__.py +26 -0
- modelopt/torch/prune/fastnas.py +376 -0
- modelopt/torch/prune/gradnas.py +349 -0
- modelopt/torch/prune/plugins/__init__.py +24 -0
- modelopt/torch/prune/plugins/mcore_minitron.py +263 -0
- modelopt/torch/prune/plugins/transformers.py +41 -0
- modelopt/torch/prune/pruning.py +211 -0
- modelopt/torch/quantization/__init__.py +26 -0
- modelopt/torch/quantization/algorithms.py +671 -0
- modelopt/torch/quantization/backends/__init__.py +20 -0
- modelopt/torch/quantization/backends/fp8_per_tensor_gemm.py +204 -0
- modelopt/torch/quantization/backends/gemm_registry.py +156 -0
- modelopt/torch/quantization/backends/nvfp4_gemm.py +201 -0
- modelopt/torch/quantization/backends/utils.py +28 -0
- modelopt/torch/quantization/calib/__init__.py +25 -0
- modelopt/torch/quantization/calib/bias.py +172 -0
- modelopt/torch/quantization/calib/calibrator.py +63 -0
- modelopt/torch/quantization/calib/histogram.py +434 -0
- modelopt/torch/quantization/calib/max.py +110 -0
- modelopt/torch/quantization/compress.py +227 -0
- modelopt/torch/quantization/config.py +1180 -0
- modelopt/torch/quantization/conversion.py +389 -0
- modelopt/torch/quantization/export_onnx.py +639 -0
- modelopt/torch/quantization/extensions.py +89 -0
- modelopt/torch/quantization/mode.py +428 -0
- modelopt/torch/quantization/model_calib.py +949 -0
- modelopt/torch/quantization/model_quant.py +477 -0
- modelopt/torch/quantization/nn/__init__.py +26 -0
- modelopt/torch/quantization/nn/functional.py +109 -0
- modelopt/torch/quantization/nn/modules/__init__.py +16 -0
- modelopt/torch/quantization/nn/modules/quant_activations.py +22 -0
- modelopt/torch/quantization/nn/modules/quant_batchnorm.py +24 -0
- modelopt/torch/quantization/nn/modules/quant_conv.py +123 -0
- modelopt/torch/quantization/nn/modules/quant_instancenorm.py +39 -0
- modelopt/torch/quantization/nn/modules/quant_linear.py +258 -0
- modelopt/torch/quantization/nn/modules/quant_module.py +211 -0
- modelopt/torch/quantization/nn/modules/quant_pooling.py +99 -0
- modelopt/torch/quantization/nn/modules/quant_rnn.py +527 -0
- modelopt/torch/quantization/nn/modules/tensor_quantizer.py +1293 -0
- modelopt/torch/quantization/plugins/__init__.py +73 -0
- modelopt/torch/quantization/plugins/accelerate.py +214 -0
- modelopt/torch/quantization/plugins/apex.py +52 -0
- modelopt/torch/quantization/plugins/attention.py +312 -0
- modelopt/torch/quantization/plugins/custom.py +192 -0
- modelopt/torch/quantization/plugins/diffusers.py +237 -0
- modelopt/torch/quantization/plugins/fairscale.py +45 -0
- modelopt/torch/quantization/plugins/huggingface.py +736 -0
- modelopt/torch/quantization/plugins/megatron.py +462 -0
- modelopt/torch/quantization/plugins/peft.py +148 -0
- modelopt/torch/quantization/plugins/transformer_engine.py +60 -0
- modelopt/torch/quantization/plugins/transformers.py +52 -0
- modelopt/torch/quantization/plugins/transformers_trainer.py +426 -0
- modelopt/torch/quantization/plugins/trl.py +27 -0
- modelopt/torch/quantization/plugins/vllm.py +199 -0
- modelopt/torch/quantization/qtensor/__init__.py +24 -0
- modelopt/torch/quantization/qtensor/base_qtensor.py +370 -0
- modelopt/torch/quantization/qtensor/fp8_tensor.py +151 -0
- modelopt/torch/quantization/qtensor/int4_tensor.py +130 -0
- modelopt/torch/quantization/qtensor/int8_tensor.py +124 -0
- modelopt/torch/quantization/qtensor/mxfp4_tensor.py +144 -0
- modelopt/torch/quantization/qtensor/nf4_tensor.py +199 -0
- modelopt/torch/quantization/qtensor/nvfp4_tensor.py +295 -0
- modelopt/torch/quantization/src/tensor_quant.cpp +77 -0
- modelopt/torch/quantization/src/tensor_quant.h +69 -0
- modelopt/torch/quantization/src/tensor_quant_fp8.cpp +48 -0
- modelopt/torch/quantization/src/tensor_quant_gpu.cu +361 -0
- modelopt/torch/quantization/src/tensor_quant_gpu_fp8.cu +122 -0
- modelopt/torch/quantization/src/tensor_quant_mx.cu +411 -0
- modelopt/torch/quantization/src/tensor_quant_mx.h +247 -0
- modelopt/torch/quantization/tensor_quant.py +816 -0
- modelopt/torch/quantization/triton/__init__.py +36 -0
- modelopt/torch/quantization/triton/fp4_kernel.py +303 -0
- modelopt/torch/quantization/utils.py +443 -0
- modelopt/torch/sparsity/__init__.py +19 -0
- modelopt/torch/sparsity/config.py +51 -0
- modelopt/torch/sparsity/magnitude.py +148 -0
- modelopt/torch/sparsity/mode.py +203 -0
- modelopt/torch/sparsity/module.py +87 -0
- modelopt/torch/sparsity/plugins/__init__.py +27 -0
- modelopt/torch/sparsity/plugins/megatron.py +93 -0
- modelopt/torch/sparsity/searcher.py +84 -0
- modelopt/torch/sparsity/sparsegpt.py +276 -0
- modelopt/torch/sparsity/sparsification.py +123 -0
- modelopt/torch/speculative/__init__.py +20 -0
- modelopt/torch/speculative/config.py +100 -0
- modelopt/torch/speculative/eagle/__init__.py +20 -0
- modelopt/torch/speculative/eagle/conversion.py +67 -0
- modelopt/torch/speculative/eagle/default_config.py +50 -0
- modelopt/torch/speculative/eagle/eagle_model.py +51 -0
- modelopt/torch/speculative/eagle/utils.py +91 -0
- modelopt/torch/speculative/medusa/__init__.py +19 -0
- modelopt/torch/speculative/medusa/conversion.py +59 -0
- modelopt/torch/speculative/medusa/medusa_model.py +32 -0
- modelopt/torch/speculative/mode.py +86 -0
- modelopt/torch/speculative/plugins/__init__.py +33 -0
- modelopt/torch/speculative/plugins/megatron_eagle.py +2119 -0
- modelopt/torch/speculative/plugins/megatron_medusa.py +312 -0
- modelopt/torch/speculative/plugins/transformers.py +1138 -0
- modelopt/torch/speculative/speculative_decoding.py +59 -0
- modelopt/torch/speculative/utils.py +364 -0
- modelopt/torch/trace/__init__.py +25 -0
- modelopt/torch/trace/analyzer.py +1368 -0
- modelopt/torch/trace/modules/__init__.py +19 -0
- modelopt/torch/trace/modules/concat.py +384 -0
- modelopt/torch/trace/modules/nn.py +153 -0
- modelopt/torch/trace/plugins/__init__.py +34 -0
- modelopt/torch/trace/plugins/megatron.py +36 -0
- modelopt/torch/trace/plugins/transformers.py +58 -0
- modelopt/torch/trace/symbols.py +545 -0
- modelopt/torch/trace/tracer.py +331 -0
- modelopt/torch/utils/__init__.py +28 -0
- modelopt/torch/utils/_pytree.py +133 -0
- modelopt/torch/utils/cpp_extension.py +91 -0
- modelopt/torch/utils/dataset_utils.py +484 -0
- modelopt/torch/utils/distributed.py +311 -0
- modelopt/torch/utils/graph.py +123 -0
- modelopt/torch/utils/image_processor.py +112 -0
- modelopt/torch/utils/import_utils.py +35 -0
- modelopt/torch/utils/list.py +54 -0
- modelopt/torch/utils/logging.py +127 -0
- modelopt/torch/utils/memory_monitor.py +146 -0
- modelopt/torch/utils/network.py +658 -0
- modelopt/torch/utils/perf.py +174 -0
- modelopt/torch/utils/plugins/__init__.py +27 -0
- modelopt/torch/utils/plugins/megatron_generate.py +299 -0
- modelopt/torch/utils/plugins/megatron_mmlu.py +152 -0
- modelopt/torch/utils/plugins/megatron_preprocess_data.py +200 -0
- modelopt/torch/utils/random.py +163 -0
- modelopt/torch/utils/speech_dataset_utils.py +134 -0
- modelopt/torch/utils/tensor.py +87 -0
- modelopt/torch/utils/vlm_dataset_utils.py +110 -0
- nvidia_modelopt-0.35.0.dist-info/METADATA +171 -0
- nvidia_modelopt-0.35.0.dist-info/RECORD +272 -0
- nvidia_modelopt-0.35.0.dist-info/WHEEL +5 -0
- nvidia_modelopt-0.35.0.dist-info/licenses/LICENSE +14 -0
- nvidia_modelopt-0.35.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,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))
|