From 4c611e47a60084a86e1de7e48690a692a1b8170c Mon Sep 17 00:00:00 2001 From: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com> Date: Thu, 31 Jul 2025 22:48:09 +0530 Subject: [PATCH] Update files on GitHub Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com> --- .pre-commit-config.yaml | 1 + CHANGELOG.rst | 15 +- .../windows/_installation_with_olive.rst | 12 +- docs/source/guides/4_distillation.rst | 19 + examples/deepseek/ptq.py | 28 +- examples/deepseek/quantize_fp8_to_nvfp4.sh | 6 + examples/diffusers/quantization/config.py | 2 +- examples/llm_ptq/example_utils.py | 10 +- examples/llm_ptq/hf_ptq.py | 50 +- .../1_FP4-FP8_PTQ_Min-Max_Calibration.ipynb | 302 ++++ .../notebooks/2_PTQ_AWQ_Calibration.ipynb | 311 ++++ .../notebooks/3_PTQ_AutoQuantization.ipynb | 430 +++++ examples/llm_qat/README.md | 2 +- examples/llm_qat/main.py | 9 +- examples/nemo_run/qat/README.md | 30 + examples/nemo_run/qat/nemo_qat_flow.py | 17 + examples/onnx_ptq/README.md | 12 +- examples/onnx_ptq/docker/Dockerfile | 37 + examples/onnx_ptq/docker/build.sh | 131 ++ examples/onnx_ptq/llm_export.py | 348 ++++ examples/speculative_decoding/README.md | 4 +- examples/speculative_decoding/eagle_utils.py | 18 +- examples/speculative_decoding/medusa_utils.py | 11 +- .../speculative_decoding/server_generate.py | 24 +- .../vlm_ptq/scripts/huggingface_example.sh | 1 - examples/windows/README.md | 4 +- examples/windows/accuracy_benchmark/README.md | 2 +- examples/windows/onnx_ptq/genai_llm/README.md | 17 +- .../windows/onnx_ptq/genai_llm/quantize.py | 2 +- examples/windows/onnx_ptq/whisper/README.md | 2 + modelopt/__init__.py | 22 - modelopt/onnx/autocast/precisionconverter.py | 10 +- .../onnx/llm_export/utils/export_utils.py | 162 ++ .../onnx/llm_export/utils/surgeon_utils.py | 120 ++ modelopt/onnx/quantization/__main__.py | 11 +- modelopt/onnx/quantization/calib_utils.py | 8 + modelopt/onnx/quantization/fp8.py | 76 +- modelopt/onnx/quantization/graph_utils.py | 83 +- modelopt/onnx/quantization/int4.py | 49 +- modelopt/onnx/quantization/int8.py | 10 +- modelopt/onnx/quantization/operators.py | 25 + modelopt/onnx/quantization/ort_patching.py | 7 + modelopt/onnx/quantization/ort_utils.py | 12 +- modelopt/onnx/quantization/partitioning.py | 13 +- modelopt/onnx/quantization/qdq_utils.py | 136 +- modelopt/onnx/quantization/quantize.py | 84 +- modelopt/onnx/trt_utils.py | 100 ++ modelopt/onnx/utils.py | 18 +- modelopt/torch/__init__.py | 2 +- modelopt/torch/_deploy/utils/torch_onnx.py | 8 +- modelopt/torch/distill/distillation_model.py | 6 +- modelopt/torch/distill/losses.py | 135 +- modelopt/torch/export/convert_hf_config.py | 24 +- modelopt/torch/export/plugins/__init__.py | 5 + modelopt/torch/export/plugins/mcore_common.py | 5 +- modelopt/torch/export/plugins/mcore_custom.py | 315 +++- .../torch/export/plugins/mcore_deepseek.py | 366 ++-- modelopt/torch/export/plugins/mcore_llama.py | 179 +- .../torch/export/plugins/mcore_nemotron.py | 111 +- modelopt/torch/export/plugins/mcore_qwen.py | 88 +- .../torch/export/plugins/megatron_importer.py | 529 ++++++ modelopt/torch/export/quant_utils.py | 47 + .../torch/export/unified_export_megatron.py | 744 +++----- modelopt/torch/nas/hparams/concat.py | 46 +- modelopt/torch/nas/plugins/megatron.py | 251 +-- modelopt/torch/nas/search_space.py | 16 +- modelopt/torch/nas/utils.py | 8 +- modelopt/torch/opt/config.py | 6 +- modelopt/torch/opt/dynamic.py | 12 +- modelopt/torch/opt/hparam.py | 3 + modelopt/torch/opt/plugins/__init__.py | 3 + modelopt/torch/opt/plugins/huggingface.py | 12 +- .../opt/plugins/mcore_dist_checkpointing.py | 54 +- .../opt/plugins/megatron_model_config.py | 62 + modelopt/torch/opt/utils.py | 4 +- .../torch/prune/plugins/mcore_gpt_minitron.py | 25 +- modelopt/torch/quantization/algorithms.py | 671 +++++++ modelopt/torch/quantization/compress.py | 1 - modelopt/torch/quantization/config.py | 250 ++- modelopt/torch/quantization/export_onnx.py | 2 +- modelopt/torch/quantization/mode.py | 2 +- modelopt/torch/quantization/model_quant.py | 2 +- .../quantization/nn/modules/quant_linear.py | 15 + .../quantization/nn/modules/quant_module.py | 38 +- .../nn/modules/tensor_quantizer.py | 27 +- .../torch/quantization/plugins/__init__.py | 3 + .../torch/quantization/plugins/huggingface.py | 2 +- .../torch/quantization/plugins/megatron.py | 4 + modelopt/torch/quantization/plugins/peft.py | 64 +- .../plugins/transformers_trainer.py | 3 + modelopt/torch/quantization/plugins/vllm.py | 199 +++ .../torch/quantization/qtensor/__init__.py | 3 +- .../quantization/qtensor/base_qtensor.py | 2 +- .../quantization/qtensor/nvfp4_tensor.py | 282 +++ modelopt/torch/quantization/tensor_quant.py | 30 +- .../torch/quantization/triton/fp4_kernel.py | 17 +- modelopt/torch/speculative/config.py | 7 + .../torch/speculative/eagle/conversion.py | 1 + .../torch/speculative/eagle/eagle_model.py | 8 + .../torch/speculative/plugins/__init__.py | 7 +- .../{megatron.py => megatron_eagle.py} | 1568 ++++++----------- .../speculative/plugins/megatron_medusa.py | 312 ++++ .../torch/speculative/plugins/transformers.py | 2 + modelopt/torch/speculative/utils.py | 193 +- modelopt/torch/trace/plugins/megatron.py | 15 +- modelopt/torch/utils/dataset_utils.py | 5 +- .../torch/utils/plugins/__init__.py | 34 +- .../torch/utils/plugins/megatron_generate.py | 207 +++ modelopt/torch/utils/plugins/megatron_mmlu.py | 150 ++ pyproject.toml | 3 +- setup.py | 22 +- tests/_test_utils/import_helper.py | 22 +- tests/_test_utils/model.py | 16 + tests/_test_utils/ptq_utils.py | 105 ++ .../torch_dist/plugins/megatron_common.py | 147 +- .../torch_model/transformers_models.py | 18 + tests/examples/cnn_qat/test_resnet50.py | 77 + tests/examples/diffusers/test_diffusers.py | 159 ++ .../diffusers/test_flux_quantization.py | 56 - .../diffusers/test_sd3_quantization.py | 65 - .../diffusers/test_sdxl_quantization.py | 107 -- tests/examples/llm_ptq/test_bart.py | 31 - tests/examples/llm_ptq/test_llama.py | 145 -- tests/examples/llm_ptq/test_llm_ptq.py | 163 ++ tests/examples/llm_ptq/test_mixtral.py | 31 - tests/examples/llm_ptq/test_t5.py | 42 - tests/examples/llm_ptq/test_whisper.py | 44 - ...y => test_quantize_onnx_torch_int4_awq.py} | 4 + .../export/test_unified_export_megatron.py | 17 +- ...y => test_megatron_gpt_dynamic_modules.py} | 52 +- .../test_mcore_gpt_minitron_pruning.py | 12 +- .../quantization/plugins/test_megatron.py | 8 +- .../test_speculative_megatron_modules.py | 187 +- ...uantize_int4.py => test_quantize_zint4.py} | 4 + tests/unit/torch/distill/test_distill.py | 17 + .../torch/opt/plugins/test_hf_patching.py | 66 + .../unit/torch/quantization/test_autoquant.py | 2 +- .../test_transformers_attention_symbols.py | 5 +- tests/unit/torch/trace/test_symbol.py | 26 +- 139 files changed, 8483 insertions(+), 3165 deletions(-) create mode 100644 examples/llm_ptq/notebooks/1_FP4-FP8_PTQ_Min-Max_Calibration.ipynb create mode 100644 examples/llm_ptq/notebooks/2_PTQ_AWQ_Calibration.ipynb create mode 100644 examples/llm_ptq/notebooks/3_PTQ_AutoQuantization.ipynb create mode 100644 examples/onnx_ptq/docker/Dockerfile create mode 100755 examples/onnx_ptq/docker/build.sh create mode 100644 examples/onnx_ptq/llm_export.py create mode 100644 modelopt/onnx/llm_export/utils/export_utils.py create mode 100644 modelopt/onnx/llm_export/utils/surgeon_utils.py create mode 100644 modelopt/torch/export/plugins/megatron_importer.py create mode 100644 modelopt/torch/opt/plugins/megatron_model_config.py create mode 100644 modelopt/torch/quantization/algorithms.py create mode 100644 modelopt/torch/quantization/plugins/vllm.py create mode 100644 modelopt/torch/quantization/qtensor/nvfp4_tensor.py rename modelopt/torch/speculative/plugins/{megatron.py => megatron_eagle.py} (52%) create mode 100644 modelopt/torch/speculative/plugins/megatron_medusa.py rename tests/examples/diffusers/conftest.py => modelopt/torch/utils/plugins/__init__.py (51%) create mode 100644 modelopt/torch/utils/plugins/megatron_generate.py create mode 100644 modelopt/torch/utils/plugins/megatron_mmlu.py create mode 100644 tests/_test_utils/ptq_utils.py create mode 100644 tests/examples/cnn_qat/test_resnet50.py create mode 100644 tests/examples/diffusers/test_diffusers.py delete mode 100644 tests/examples/diffusers/test_flux_quantization.py delete mode 100644 tests/examples/diffusers/test_sd3_quantization.py delete mode 100644 tests/examples/diffusers/test_sdxl_quantization.py delete mode 100644 tests/examples/llm_ptq/test_bart.py delete mode 100644 tests/examples/llm_ptq/test_llama.py create mode 100644 tests/examples/llm_ptq/test_llm_ptq.py delete mode 100644 tests/examples/llm_ptq/test_mixtral.py delete mode 100644 tests/examples/llm_ptq/test_t5.py delete mode 100644 tests/examples/llm_ptq/test_whisper.py rename tests/gpu/onnx/{test_onnx_torch_int4_awq.py => test_quantize_onnx_torch_int4_awq.py} (94%) rename tests/gpu/torch/nas/plugins/{test_megatron_dynamic_modules.py => test_megatron_gpt_dynamic_modules.py} (88%) rename tests/unit/onnx/{test_quantize_int4.py => test_quantize_zint4.py} (95%) create mode 100644 tests/unit/torch/opt/plugins/test_hf_patching.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 24ac9a001..b48532cc4 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -96,6 +96,7 @@ repos: modelopt/torch/quantization/plugins/attention.py| modelopt/torch/speculative/eagle/utils.py| modelopt/torch/speculative/plugins/transformers.py| + modelopt/torch/utils/plugins/megatron_mmlu.py| examples/chained_optimizations/bert_prune_distill_quantize.py| examples/deepseek/quantize_to_nvfp4.py| examples/deepseek/ptq.py| diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 157ec8f1f..f57ae87b6 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -1,6 +1,19 @@ Model Optimizer Changelog (Linux) ================================= +0.35 (2025-08-xx) +^^^^^^^^^^^^^^^^^ + +**Backward Breaking Changes** + +**Deprecations** + +**New Features** + +- (Experimental) Add quantization support for custom TensorRT op in ONNX models. +- Add support for Minifinetuning (MFT; https://arxiv.org/abs/2506.15702) self-corrective distillation, which enables training on small datasets with severely mitigated catastrophic forgetting. +- Add tree decoding support for Megatron Eagle models. + 0.33 (2025-07-14) ^^^^^^^^^^^^^^^^^ @@ -20,7 +33,7 @@ Model Optimizer Changelog (Linux) - Add per node calibration support in ONNX quantization. - ModelOpt now supports quantization of tensor-parallel sharded Huggingface transformer models. This requires ``transformers>=4.52.0``. - Support quantization of FSDP2 wrapped models and add FSDP2 support in the ``llm_qat`` example. -- Add NeMo 2 Simplified Flow examples for quantization aware training/distillation (QAT/QAD), speculative decoding, pruning & distilllation. +- Add NeMo 2 Simplified Flow examples for quantization aware training/distillation (QAT/QAD), speculative decoding, pruning & distillation. 0.31 (2025-06-04) ^^^^^^^^^^^^^^^^^ diff --git a/docs/source/getting_started/windows/_installation_with_olive.rst b/docs/source/getting_started/windows/_installation_with_olive.rst index 9b9a8cf01..977a29f16 100644 --- a/docs/source/getting_started/windows/_installation_with_olive.rst +++ b/docs/source/getting_started/windows/_installation_with_olive.rst @@ -24,8 +24,9 @@ Setup Steps for Olive with ModelOpt-Windows $ pip install onnxruntime-genai-directml>=0.4.0 $ pip install onnxruntime-directml==1.20.0 + - Above onnxruntime and onnxruntime-genai packages enable Olive workflow with DirectML Execution-Provider (EP). To use other EPs, install corresponding packages. - Additionally, ensure that dependencies for TensorRT Model Optimizer - Windows are met as mentioned in the :ref:`Install-Page-Standalone-Windows`. + - Additionally, ensure that dependencies for TensorRT Model Optimizer - Windows are met as mentioned in the :ref:`Install-Page-Standalone-Windows`. **2. Configure Olive for TensorRT Model Optimizer – Windows** @@ -36,7 +37,11 @@ Setup Steps for Olive with ModelOpt-Windows - **Add Other Passes:** Add additional passes to the Olive configuration file as needed for the desired Olive workflow of your input model. [Refer `phi3 `_ Olive example] -**4. Run the Optimization** +**4. Install other dependencies** + + - Install other requirements as needed by the Olive scripts and config. + +**5. Run the Optimization** - **Execute Optimization:** To start the optimization process, run the following commands: @@ -56,4 +61,5 @@ Setup Steps for Olive with ModelOpt-Windows **Note**: -#. Currently, the TensorRT-Model Optimizer - Windows only supports Onnx Runtime GenAI based models in the Olive workflow. +#. Currently, the TensorRT-Model Optimizer - Windows only supports Onnx Runtime GenAI based LLM models in the Olive workflow. +#. To try out different LLMs and EPs in the Olive workflow of ModelOpt-Windows, refer the details provided in `phi3 `_ Olive example. diff --git a/docs/source/guides/4_distillation.rst b/docs/source/guides/4_distillation.rst index 773e689bb..4bf228ecd 100644 --- a/docs/source/guides/4_distillation.rst +++ b/docs/source/guides/4_distillation.rst @@ -62,6 +62,10 @@ Example usage: meta model. Thus, the same callable must be available in the namespace when restoring via the :meth:`mto.restore ` utility. +.. tip:: + When training the student on a small corpus of ground truth data, consider using :class:`MFTLoss ` for to perform Minifinetuning in lieu of the standard + :class:`LogitsDistillationLoss `. This will allow the student to learn from the teacher's distribution while adapting to the new data, improving the specialization of the new data without overwriting teacher's general knowledge. + .. note:: As the model is not of the same class anymore, calling ``type()`` on the model after conversion will not work as expected. @@ -124,6 +128,9 @@ maps or logits) which the teacher has already mastered. This can serve multiple **C.** Module replacement: One can replace a single module within a model with a more efficient one and use distillation on its original outputs to effectively re-integrate it into the whole model. + **D.** Minimal modification without catastrophic forgetting: A variant of distillation, called Minifinetuning, + allows for training a model on even small datasets without losing the original model's knowledge. + Student ^^^^^^^ @@ -192,3 +199,15 @@ ground truth labels may be. .. _1: https://arxiv.org/abs/1803.03635 + +Minifinetuning +^^^^^^^^^^^^^^ + +Minifinetuning is a technique that allows for training a model on even small datasets without losing the original +model's knowledge. This is achieved by algorithmic modification of the teacher's distribution depending on its +performance on the new dataset. The goal is to ensure that the separation between the correct and incorrect argmax +tokens is large enough, which can be controlled by a threshold parameter. ModelOpt provides a pre-defined loss function +for this purpose, called :class:`MFTDistillationLoss `, which can +be used in place of the standard :class:`LogitsDistillationLoss `. +More information about the technique can be found in the original paper: +`Minifinetuning: Low-Data Generation Domain Adaptation through Corrective Self-Distillation `_. diff --git a/examples/deepseek/ptq.py b/examples/deepseek/ptq.py index 48cd33eb3..376620f07 100644 --- a/examples/deepseek/ptq.py +++ b/examples/deepseek/ptq.py @@ -56,7 +56,13 @@ import modelopt.torch.quantization as mtq from modelopt.torch.export.model_config import KV_CACHE_FP8 from modelopt.torch.export.quant_utils import get_quant_config from modelopt.torch.quantization.nn import TensorQuantizer +from modelopt.torch.quantization.utils import ( + is_quantized_column_parallel_linear, + is_quantized_parallel_linear, + is_quantized_row_parallel_linear, +) from modelopt.torch.utils.dataset_utils import get_dataset_dataloader +from modelopt.torch.utils.distributed import ParallelState sys.path.append(str(Path(__file__).resolve().parent / "DeepSeek-V3/inference")) import model as deekseep_model @@ -105,6 +111,11 @@ def monkey_patch_deepseek_model(): def _setup(self): self.input_quantizer = TensorQuantizer() self.weight_quantizer = TensorQuantizer() + # Use TP parallel state + self._parallel_state = ParallelState(data_parallel_group=-1, tensor_parallel_group=None) + self._is_column_parallel = True + + assert is_quantized_column_parallel_linear(self) def forward(self, x: torch.Tensor) -> torch.Tensor: y = linear( @@ -124,6 +135,11 @@ def monkey_patch_deepseek_model(): def _setup(self): self.input_quantizer = TensorQuantizer() self.weight_quantizer = TensorQuantizer() + # Use TP parallel state + self._parallel_state = ParallelState(data_parallel_group=-1, tensor_parallel_group=None) + self._is_row_parallel = True + + assert is_quantized_row_parallel_linear(self) def forward(self, x: torch.Tensor) -> torch.Tensor: y = linear( @@ -146,6 +162,10 @@ def monkey_patch_deepseek_model(): def _setup(self): self.input_quantizer = TensorQuantizer() self.weight_quantizer = TensorQuantizer() + # No parallel state. + self._parallel_state = ParallelState(data_parallel_group=-1, tensor_parallel_group=-1) + + assert not is_quantized_parallel_linear(self) def forward(self, x: torch.Tensor) -> torch.Tensor: y = linear( @@ -238,6 +258,9 @@ def ptq( ## handle DeepSeek model structures transformer = model.model if hasattr(model, "model") else model + # make sure all processes are ready before starting the calibration + dist.barrier() + ## quant config mtq_cfg = getattr(mtq, quant_cfg) @@ -332,9 +355,12 @@ if __name__ == "__main__": parser.add_argument("--calib_size", type=int, default=512, help="samples for calibration.") parser.add_argument("--enable_fp8_kvcache", type=bool, default=True, help="enable fp8 kvcache.") parser.add_argument("--enable_wo_quant", action="store_true", help="enable MLA wo quant.") + parser.add_argument("--trust_remote_code", action="store_true", help="trust remote code.") args = parser.parse_args() model = load_deepseek_model(args.config, args.model_path, args.batch_size) - tokenizer = AutoTokenizer.from_pretrained(args.model_path) + tokenizer = AutoTokenizer.from_pretrained( + args.model_path, trust_remote_code=args.trust_remote_code + ) model = ptq(model, tokenizer, args.quant_cfg, args.batch_size, args.calib_size) save_amax_and_quant_config(model, args.output_path, args.enable_fp8_kvcache) diff --git a/examples/deepseek/quantize_fp8_to_nvfp4.sh b/examples/deepseek/quantize_fp8_to_nvfp4.sh index de5fd4949..8dd8f4fcd 100755 --- a/examples/deepseek/quantize_fp8_to_nvfp4.sh +++ b/examples/deepseek/quantize_fp8_to_nvfp4.sh @@ -70,6 +70,12 @@ if [[ -z "$FP8_HF_PATH" ]]; then usage fi +# for KIMI-K2, copy tiktoken.model to tokenizer to the quantized checkpoint +if [[ -f "$FP8_HF_PATH/tiktoken.model" ]]; then + echo "tiktoken.model found in $FP8_HF_PATH" + cp $FP8_HF_PATH/tiktoken.model $FP4_PATH/ +fi + # Copy miscellaneous files to the quantized checkpoint mkdir -p $FP4_PATH cp $FP8_HF_PATH/*.json $FP8_HF_PATH/*.py $FP4_PATH/ diff --git a/examples/diffusers/quantization/config.py b/examples/diffusers/quantization/config.py index 1e78df4ed..d7e82524d 100644 --- a/examples/diffusers/quantization/config.py +++ b/examples/diffusers/quantization/config.py @@ -18,7 +18,7 @@ import torch from calib.plugin_calib import PercentileCalibrator from utils import filter_func -from modelopt.core.torch.quantization.config import NVFP4_FP8_MHA_CONFIG # noqa: F401 +from modelopt.torch.quantization.config import NVFP4_FP8_MHA_CONFIG # noqa: F401 FP8_DEFAULT_CONFIG = { "quant_cfg": { diff --git a/examples/llm_ptq/example_utils.py b/examples/llm_ptq/example_utils.py index b83885ad8..4538080a0 100755 --- a/examples/llm_ptq/example_utils.py +++ b/examples/llm_ptq/example_utils.py @@ -20,7 +20,13 @@ from typing import Any import torch from accelerate import infer_auto_device_map, init_empty_weights from accelerate.utils import get_max_memory -from transformers import AutoConfig, AutoModelForCausalLM, AutoProcessor, AutoTokenizer +from transformers import ( + AutoConfig, + AutoModelForCausalLM, + AutoProcessor, + AutoTokenizer, + Llama4ForConditionalGeneration, +) from modelopt.torch.utils.image_processor import MllamaImageProcessor @@ -225,7 +231,7 @@ def get_model( **model_kwargs, ) elif hf_config.model_type == "llama4": - model = AutoModelForCausalLM.from_pretrained( + model = Llama4ForConditionalGeneration.from_pretrained( ckpt_path, device_map=device_map, **model_kwargs, diff --git a/examples/llm_ptq/hf_ptq.py b/examples/llm_ptq/hf_ptq.py index eb0a7d3a5..b385ca3d6 100755 --- a/examples/llm_ptq/hf_ptq.py +++ b/examples/llm_ptq/hf_ptq.py @@ -85,7 +85,8 @@ def auto_quantize( # Check if all provided quantization formats are supported if args.export_fmt == "hf": assert all( - qformat in ["fp8", "int4_awq", "nvfp4", "nvfp4_awq", "w4a8_awq", "fp8_pb_wo"] + qformat + in ["fp8", "int4_awq", "nvfp4", "nvfp4_awq", "w4a8_awq", "fp8_pb_wo", "w4a8_mxfp4_fp8"] for qformat in qformat_list ), ( "One or more quantization formats provided are not supported for unified checkpoint export" @@ -110,9 +111,7 @@ def auto_quantize( # TRTLLM only support one quantization format or None (do not quantize, internally supported) quantization_formats=[QUANT_CFG_CHOICES[format] for format in qformat_list], num_calib_steps=len(calib_dataloader), - num_score_steps=min( - len(calib_dataloader), 128 // batch_size - ), # Limit the number of score steps to avoid long calibration time + num_score_steps=len(calib_dataloader), verbose=True, disabled_layers=["*lm_head*"], ) @@ -218,6 +217,7 @@ def main(args): "nvfp4_awq", "w4a8_awq", "fp8_pb_wo", + "w4a8_mxfp4_fp8", ] or args.kv_cache_qformat in KV_QUANT_CFG_CHOICES ), f"Quantization format {args.qformat} not supported for HF export path" @@ -263,6 +263,9 @@ def main(args): device = model.model.device processor = None tokenizer = None + + full_model = model + if model_type == "mllama": if args.dataset is None: args.dataset = "scienceqa" @@ -300,6 +303,13 @@ def main(args): # Left padding usually provides better calibration result. tokenizer.padding_side = "left" + # We only quantize the language model for VLMs other than the type supported above. + if hasattr(model, "language_model"): + assert model_type == "llama4", ( + "Only llama4 should reach here. Please uncomment this check if you are modelopt developers." + ) + model = model.language_model + if args.sparsity_fmt != "dense": if args.batch_size == 0: # Sparse algorithm takes more GPU memory so we reduce the batch_size by 4. @@ -335,10 +345,6 @@ def main(args): ) if args.batch_size == 0: - # TODO: Enable auto-batch size calculation for auto_quantize - assert args.auto_quantize_bits is None, ( - "auto_quantize requires batch_size to be specified, please specify batch_size." - ) # Calibration/sparsification will actually take much more memory than regular inference # due to intermediate tensors for fake quantization. Setting sample_memory_usage_ratio # to 2 to avoid OOM for AWQ/SmoothQuant fake quantization as it will take more memory than inference. @@ -358,10 +364,14 @@ def main(args): ) else: sample_input_single_batch = None + + run_auto_quant = args.auto_quantize_bits is not None + args.batch_size = get_max_batch_size( model, - sample_memory_usage_ratio=sample_memory_usage_ratio, + sample_memory_usage_ratio=sample_memory_usage_ratio if not run_auto_quant else 1.0, sample_input_single_batch=sample_input_single_batch, + enable_grad=run_auto_quant, ) args.batch_size = min(args.batch_size, args.calib_size) @@ -550,23 +560,9 @@ def main(args): ) elif args.export_fmt == "hf": export_hf_checkpoint( - model, + full_model, export_dir=export_path, ) - if model_type == "llama4": - # TRT-LLM expects the original model config instead of the config from text model, - # so we need to copy the original model config to the export path. - # Also we copy the preprocessor config to the export path. - from transformers import AutoConfig, AutoProcessor - - # Use HuggingFace API to handle both model IDs and local paths - AutoConfig.from_pretrained( - args.pyt_ckpt_path, trust_remote_code=args.trust_remote_code - ).save_pretrained(export_path) - - AutoProcessor.from_pretrained( - args.pyt_ckpt_path, trust_remote_code=args.trust_remote_code - ).save_pretrained(export_path) else: raise NotImplementedError(f"{args.export_fmt} not supported") @@ -639,12 +635,6 @@ if __name__ == "__main__": choices=KV_QUANT_CFG_CHOICES.keys(), help="Specify KV cache quantization format, default to fp8 if not provided", ) - parser.add_argument( - "--vlm", - help="Specify whether this is a visual-language model", - default=False, - action="store_true", - ) parser.add_argument( "--export_fmt", required=False, diff --git a/examples/llm_ptq/notebooks/1_FP4-FP8_PTQ_Min-Max_Calibration.ipynb b/examples/llm_ptq/notebooks/1_FP4-FP8_PTQ_Min-Max_Calibration.ipynb new file mode 100644 index 000000000..2c981c2f1 --- /dev/null +++ b/examples/llm_ptq/notebooks/1_FP4-FP8_PTQ_Min-Max_Calibration.ipynb @@ -0,0 +1,302 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "9161daa2-03a6-41cd-b349-410004ab37c4", + "metadata": {}, + "source": [ + "# Post-Training Quantization with Min-Max Calibration using TensorRT Model Optimizer PTQ\n", + "\n", + "This notebook demonstrates how to apply standard Post-Training Quantization (PTQ) using min-max calibration on an LLM—specifically meta-llama/Llama-3.1-8B-Instruct—with NVIDIA's TensorRT Model Optimizer (ModelOpt) PTQ toolkit. We walk through loading the model, calibrating it using a CNN/DailyMail dataset sample, applying FP8 quantization, generating outputs, and exporting the quantized model.\n", + "\n", + "Key Dependendancies: \n", + "- nvidia-modelopt\n", + "- torch\n", + "- transformers" + ] + }, + { + "cell_type": "markdown", + "id": "0972a570-2549-4494-b56e-9a901c5c2286", + "metadata": {}, + "source": [ + "## Standard FP4/FP8 Quantization with Min-Max Calibration" + ] + }, + { + "cell_type": "markdown", + "id": "657e4d08-e85f-4d2f-b22d-c11a15268c8b", + "metadata": {}, + "source": [ + "### 1. Import Dependencies\n", + "Import all necessary libraries:\n", + "\n", + "- `torch`: Used for tensor computation and model execution.\n", + "\n", + "- `modelopt.torch.quantization`: Core API for quantization using TensorRT ModelOpt PTQ.\n", + "\n", + "- `transformers`: Hugging Face interface to load and tokenize LLMs.\n", + "\n", + "- `get_dataset_dataloader` and `create_forward_loop`: Utilities to prepare calibration data and run calibration.\n", + "\n", + "- `login`: Required to download gated models (like Llama 3.1) from Hugging Face.\n", + "\n", + "💡 If you're using this in Colab or a restricted environment, make sure all packages are installed and CUDA is available." + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "c31ffa95-21cb-49a4-bf96-a0e626016466", + "metadata": {}, + "outputs": [], + "source": [ + "import torch\n", + "from huggingface_hub import login\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer\n", + "\n", + "import modelopt.torch.quantization as mtq\n", + "from modelopt.torch.utils.dataset_utils import create_forward_loop, get_dataset_dataloader" + ] + }, + { + "cell_type": "markdown", + "id": "f3a654c1-54ad-4597-a1e3-985822aac5ff", + "metadata": {}, + "source": [ + "### 2. Set Configurations and Login to Hugging Face\n", + "\n", + "Set the model you want to quantize (Llama-3.1-8B-Instruct) and the dataset to use for calibration (cnn_dailymail).\n", + "\n", + "- `batch_size` and `calib_samples` control how much data is used during calibration—more samples improve accuracy but - increase calibration time.\n", + "\n", + "🔐 You must `login()` with a valid Hugging Face token to access gated models. Get your token at hf.co/settings/tokens.\n", + "\n", + "🔁 You can substitute your own model or dataset as long as the inputs are compatible with the model's tokenizer." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "fc1a201f-4cb1-43bb-b0a9-c859ea6eb722", + "metadata": {}, + "outputs": [], + "source": [ + "model_name = \"meta-llama/Llama-3.1-8B-Instruct\"\n", + "dataset_name = \"cnn_dailymail\"\n", + "batch_size = 8\n", + "calib_samples = 512\n", + "\n", + "login()" + ] + }, + { + "cell_type": "markdown", + "id": "8b233aae-5341-4f9b-9183-3ccf94af0e89", + "metadata": {}, + "source": [ + "### 3. Load Model and Tokenizer\n", + "\n", + "- Load the model into GPU memory.\n", + "- Set `pad_token` to eos_token to prevent padding errors in decoder-only models like Llama.\n", + "\n", + "💡 Always check for token mismatch warnings in console when loading tokenizer.\n", + "🧠 Setting `pad_token` helps avoid errors during batch generation or dataset collation." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "79b6b892-dc12-4549-8841-13016da409c4", + "metadata": {}, + "outputs": [], + "source": [ + "model = AutoModelForCausalLM.from_pretrained(model_name).cuda()\n", + "tokenizer = AutoTokenizer.from_pretrained(model_name)\n", + "tokenizer.pad_token = tokenizer.eos_token" + ] + }, + { + "cell_type": "markdown", + "id": "bfb43104-4251-4634-9efc-312f9a5f0dc6", + "metadata": {}, + "source": [ + "### 4. Configure Dataloader\n", + "- Load a few batches of real-world text to extract representative activation ranges.\n", + "- The calibration dataset should reflect your expected inference use case for best results.\n", + "\n", + "⚠️ More samples = better accuracy, but takes longer. We recommend 512 samples or more. \n", + "🧪 Use your target task’s dataset (e.g., chat, summarization, code) for domain-specific calibration." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "d32f87c1-9d74-4b2b-9250-3709f81d0c58", + "metadata": {}, + "outputs": [], + "source": [ + "dataloader = get_dataset_dataloader(\n", + " dataset_name=dataset_name,\n", + " tokenizer=tokenizer,\n", + " batch_size=batch_size,\n", + " num_samples=calib_samples,\n", + " device=\"cuda\",\n", + ")" + ] + }, + { + "cell_type": "markdown", + "id": "825f84b1-9557-4a90-b0fe-a539a54f63da", + "metadata": {}, + "source": [ + "### 5. Create the Foward Loop\n", + "- Wraps your `dataloader` into a loop that feeds batches into the model.\n", + "- Required by `modelopt.quantize()` to perform calibration pass.\n", + "\n", + "🧰 You can create your own custom forward loop if you're doing multi-modal or conditional generation tasks.\n", + "🧠 ModelOpt expects this loop to return outputs so it can record activations for min/max stats." + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "id": "b774f65d-a6b5-409e-afc0-9d09589c9a90", + "metadata": {}, + "outputs": [], + "source": [ + "forward_loop = create_forward_loop(dataloader=dataloader)" + ] + }, + { + "cell_type": "markdown", + "id": "0b14930c-baa9-4ae8-9b2a-0a412949b204", + "metadata": {}, + "source": [ + "### 6. Set Quantization Configuration and Apply\n", + "- Apply FP8 quantization using the default min-max config provided by TensorRT ModelOpt.\n", + "- This pass captures the range of activations and applies a quantization transform.\n", + "- To change the quantization configuration, you simply need to change the value of the `quant_cfg` variable. For example, to change this from FP8 to NVFP4, you can set it to `mtq.NVFP4_DEFAULT_CFG`\n", + "\n", + "📏 Min-max calibration uses observed min and max values per tensor to set scaling ranges.\n", + "💡 You can experiment with other formats (e.g., FP4, INT8) by swapping out quant_cfg." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a3ce3b47-48ac-4a27-a5ed-351a10c104a9", + "metadata": {}, + "outputs": [], + "source": [ + "quant_cfg = mtq.FP8_DEFAULT_CFG # mtq.NVFP4_DEFAULT_CFG\n", + "model = mtq.quantize(model, quant_cfg, forward_loop=forward_loop)" + ] + }, + { + "cell_type": "markdown", + "id": "1d6ffc8f-7f04-4017-af53-54f4849646c5", + "metadata": {}, + "source": [ + "### 7. Quick Test of Quantized Model\n", + "- Test the quantized model with a simple prompt.\n", + "- This helps verify that quantization didn’t break forward generation or drastically harm output quality.\n", + "\n", + "✅ Expect slightly more variation or truncation in output compared to the original model, but it should still be coherent.\n", + "🧪 You can test on more complex prompts to evaluate qualitative performance further." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1a601bfc-0637-409d-bbf7-b97a2bbf6526", + "metadata": {}, + "outputs": [], + "source": [ + "model = torch.compile(model)\n", + "inputs = tokenizer(\"Hello world\", return_tensors=\"pt\").to(\"cuda\")\n", + "outputs = model.generate(**inputs, max_new_tokens=20)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1c77452d-67e2-41e0-818e-6cacdd9c3895", + "metadata": {}, + "outputs": [], + "source": [ + "print(tokenizer.decode(outputs[0], skip_special_tokens=True))" + ] + }, + { + "cell_type": "markdown", + "id": "2a5dc86d-9d09-45cc-96df-b9ddcebaf917", + "metadata": {}, + "source": [ + "### 8. Export Quantized Checkpoint\n", + "- Save the quantized model in Hugging Face-compatible format for reuse or deployment.\n", + "- Export includes weights and config files in standard structure.\n", + "\n", + "📁 This allows you to upload it to Hugging Face Hub or load later with from_pretrained() 🧰 You can also use this exported model with inference engines like vLLM, SGLang, or TensorRT-LLM." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "0bfa6720-6cc9-4668-a325-263035b17e7d", + "metadata": {}, + "outputs": [], + "source": [ + "from modelopt.torch.export import export_hf_checkpoint\n", + "\n", + "export_path = \"./quantized_model_min-max/\"\n", + "export_hf_checkpoint(model, export_dir=export_path)\n", + "tokenizer.save_pretrained(export_path)" + ] + }, + { + "cell_type": "markdown", + "id": "9934d62d-3b38-4efa-95a9-e5b27634a365", + "metadata": {}, + "source": [ + "# ✅ Conclusion & Key Takeaways\n", + " ✅ Min-max calibration is a fast and simple way to apply quantization with good performance tradeoffs.\n", + "\n", + " ✅ TensorRT-LLM ModelOpt PTQ abstracts away many of the complexities of quantization while still offering flexibility and export options.\n", + "\n", + " ✅ Using a representative dataset like cnn_dailymail improves calibration accuracy for summarization-style models.\n", + "\n", + " ✅ The quantized model remains Hugging Face-compatible—meaning it can be deployed or fine-tuned using existing tools.\n", + "\n", + " ✅ You can easily customize: The quantization format (e.g., INT8, FP4), Calibration samples and batch size, and Dataset/task alignment" + ] + }, + { + "cell_type": "markdown", + "id": "9ba7aee5", + "metadata": {}, + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.5" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/llm_ptq/notebooks/2_PTQ_AWQ_Calibration.ipynb b/examples/llm_ptq/notebooks/2_PTQ_AWQ_Calibration.ipynb new file mode 100644 index 000000000..c6218cfb7 --- /dev/null +++ b/examples/llm_ptq/notebooks/2_PTQ_AWQ_Calibration.ipynb @@ -0,0 +1,311 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "9161daa2-03a6-41cd-b349-410004ab37c4", + "metadata": {}, + "source": [ + "# Weight-Only INT4 Quantization with AWQ using TensorRT ModelOpt PTQ\n", + "\n", + "This notebook demonstrates how to apply weight-only INT4 quantization using the Activation-aware Weight Quantization (AWQ) technique via NVIDIA TensorRT-LLM Model Optimizer (ModelOpt) PTQ.\n", + "\n", + "Unlike standard min-max calibration, AWQ does not quantize activations—instead, it uses knowledge of activation ranges to inform how model weights are quantized.\n", + "\n", + "Key Dependendancies: \n", + "- nvidia-modelopt\n", + "- torch\n", + "- transformers" + ] + }, + { + "cell_type": "markdown", + "id": "0972a570-2549-4494-b56e-9a901c5c2286", + "metadata": {}, + "source": [ + "## Quantization with AWQ Quantization" + ] + }, + { + "cell_type": "markdown", + "id": "657e4d08-e85f-4d2f-b22d-c11a15268c8b", + "metadata": {}, + "source": [ + "### 1. Import Dependencies\n", + "Import all necessary libraries:\n", + "\n", + "- `torch`: Used for tensor computation and model execution.\n", + "\n", + "- `modelopt.torch.quantization`: Core API for quantization using TensorRT ModelOpt PTQ.\n", + "\n", + "- `transformers`: Hugging Face interface to load and tokenize LLMs.\n", + "\n", + "- `get_dataset_dataloader` and `create_forward_loop`: Utilities to prepare calibration data and run calibration.\n", + "\n", + "- `login`: Required to download gated models (like Llama 3.1) from Hugging Face.\n", + "\n", + "💡 If you're using this in Colab or a restricted environment, make sure all packages are installed and CUDA is available." + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "c31ffa95-21cb-49a4-bf96-a0e626016466", + "metadata": {}, + "outputs": [], + "source": [ + "import torch\n", + "from huggingface_hub import login\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer\n", + "\n", + "import modelopt.torch.quantization as mtq\n", + "from modelopt.torch.utils.dataset_utils import create_forward_loop, get_dataset_dataloader" + ] + }, + { + "cell_type": "markdown", + "id": "f3a654c1-54ad-4597-a1e3-985822aac5ff", + "metadata": {}, + "source": [ + "### 2. Set Configurations and Login to Hugging Face\n", + "\n", + "Set the model you want to quantize (Llama-3.1-8B-Instruct) and the dataset to use for calibration (cnn_dailymail).\n", + "\n", + "- `batch_size` and `calib_samples` control how much data is used during calibration—more samples improve accuracy but - increase calibration time.\n", + "\n", + "🔐 You must `login()` with a valid Hugging Face token to access gated models. Get your token at hf.co/settings/tokens.\n", + "\n", + "🔁 You can substitute your own model or dataset as long as the inputs are compatible with the model's tokenizer." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "fc1a201f-4cb1-43bb-b0a9-c859ea6eb722", + "metadata": {}, + "outputs": [], + "source": [ + "model_name = \"meta-llama/Llama-3.1-8B-Instruct\"\n", + "dataset_name = \"cnn_dailymail\"\n", + "batch_size = 8\n", + "calib_samples = 512\n", + "\n", + "login()" + ] + }, + { + "cell_type": "markdown", + "id": "8b233aae-5341-4f9b-9183-3ccf94af0e89", + "metadata": {}, + "source": [ + "### 3. Load Model and Tokenizer\n", + "\n", + "- Load the model into GPU memory.\n", + "- Set `pad_token` to eos_token to prevent padding errors in decoder-only models like Llama.\n", + "\n", + "💡 Always check for token mismatch warnings in console when loading tokenizer.\n", + "🧠 Setting `pad_token` helps avoid errors during batch generation or dataset collation." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "79b6b892-dc12-4549-8841-13016da409c4", + "metadata": {}, + "outputs": [], + "source": [ + "model = AutoModelForCausalLM.from_pretrained(model_name).cuda()\n", + "tokenizer = AutoTokenizer.from_pretrained(model_name)\n", + "tokenizer.pad_token = tokenizer.eos_token" + ] + }, + { + "cell_type": "markdown", + "id": "bfb43104-4251-4634-9efc-312f9a5f0dc6", + "metadata": {}, + "source": [ + "### 4. Configure Dataloader\n", + "- Load a few batches of real-world text to extract representative activation ranges.\n", + "- The calibration dataset should reflect your expected inference use case for best results.\n", + "\n", + "⚠️ More samples = better accuracy, but takes longer. We recommend 512 samples or more.\n", + "🧪 Use your target task’s dataset (e.g., chat, summarization, code) for domain-specific calibration." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "d32f87c1-9d74-4b2b-9250-3709f81d0c58", + "metadata": {}, + "outputs": [], + "source": [ + "dataloader = get_dataset_dataloader(\n", + " dataset_name=dataset_name,\n", + " tokenizer=tokenizer,\n", + " batch_size=batch_size,\n", + " num_samples=calib_samples,\n", + " device=\"cuda\",\n", + ")" + ] + }, + { + "cell_type": "markdown", + "id": "825f84b1-9557-4a90-b0fe-a539a54f63da", + "metadata": {}, + "source": [ + "### 5. Create the Foward Loop\n", + "- Wraps your `dataloader` into a loop that feeds batches into the model.\n", + "- Required by `modelopt.quantize()` to perform calibration pass.\n", + "\n", + "🧰 You can create your own custom forward loop if you're doing multi-modal or conditional generation tasks." + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "id": "b774f65d-a6b5-409e-afc0-9d09589c9a90", + "metadata": {}, + "outputs": [], + "source": [ + "forward_loop = create_forward_loop(dataloader=dataloader)" + ] + }, + { + "cell_type": "markdown", + "id": "0b14930c-baa9-4ae8-9b2a-0a412949b204", + "metadata": {}, + "source": [ + "### 6. Set Quantization Configuration and Apply\n", + "🔧 Retrieve and customize the AWQ (Activation-aware Weight Quantization) config for INT4 quantization.\n", + "- `mtq.INT4_AWQ_CFG` provides a pre-tuned config optimized for low-bit weight quantization with block-wise granularity.\n", + "- `block_sizes` control how quantization groups are split across dimensions. This affects compression ratio, memory layout, and accuracy.\n", + "- The last dimension (typically 128 or 64) defines the quantization block size for each row of weights.\n", + "\n", + "💡 You can experiment with smaller block sizes (e.g., 64 or 32) for better accuracy at the cost of less compression." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a3ce3b47-48ac-4a27-a5ed-351a10c104a9", + "metadata": {}, + "outputs": [], + "source": [ + "# Get default AWQ config and optionally adjust block size\n", + "quant_cfg = mtq.INT4_AWQ_CFG\n", + "weight_quantizer = quant_cfg[\"quant_cfg\"][\"*weight_quantizer\"]\n", + "if isinstance(weight_quantizer, list):\n", + " weight_quantizer = weight_quantizer[0]\n", + "weight_quantizer[\"block_sizes\"][-1] = 128 # Optional: override block size\n", + "\n", + "# Apply AWQ quantization\n", + "model = mtq.quantize(model, quant_cfg, forward_loop=forward_loop)" + ] + }, + { + "cell_type": "markdown", + "id": "1d6ffc8f-7f04-4017-af53-54f4849646c5", + "metadata": {}, + "source": [ + "### 7. Quick Test of Quantized Model\n", + "- Test the quantized model with a simple prompt.\n", + "- This helps verify that quantization didn’t break forward generation or drastically harm output quality.\n", + "\n", + "🧪 You can test on more complex prompts to evaluate qualitative performance further." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1a601bfc-0637-409d-bbf7-b97a2bbf6526", + "metadata": {}, + "outputs": [], + "source": [ + "model = torch.compile(model)\n", + "inputs = tokenizer(\"Hello world\", return_tensors=\"pt\").to(\"cuda\")\n", + "outputs = model.generate(**inputs, max_new_tokens=20)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1c77452d-67e2-41e0-818e-6cacdd9c3895", + "metadata": {}, + "outputs": [], + "source": [ + "print(tokenizer.decode(outputs[0], skip_special_tokens=True))" + ] + }, + { + "cell_type": "markdown", + "id": "2a5dc86d-9d09-45cc-96df-b9ddcebaf917", + "metadata": {}, + "source": [ + "### 8. Export Quantized Checkpoint\n", + "- Save the quantized model in Hugging Face-compatible format for reuse or deployment.\n", + "- Export includes weights and config files in standard structure.\n", + "\n", + "📁 This allows you to upload it to Hugging Face Hub or load later with from_pretrained() 🧰 You can also use this exported model with inference engines like vLLM, SGLang, or TensorRT-LLM." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "0bfa6720-6cc9-4668-a325-263035b17e7d", + "metadata": {}, + "outputs": [], + "source": [ + "from modelopt.torch.export import export_hf_checkpoint\n", + "\n", + "export_path = \"./quantized_model_awq/\"\n", + "export_hf_checkpoint(model, export_dir=export_path)\n", + "tokenizer.save_pretrained(export_path)" + ] + }, + { + "cell_type": "markdown", + "id": "9934d62d-3b38-4efa-95a9-e5b27634a365", + "metadata": {}, + "source": [ + "# ✅ Conclusion & Key Takeaways\n", + " ✅ AWQ (Activation-aware Weight Quantization) is an efficient, deployment-ready method for compressing large language models without quantizing activations.\n", + "\n", + " ✅ Using INT4 weight-only quantization, AWQ significantly reduces model memory footprint and improves inference throughput—ideal for GPU inference workloads.\n", + "\n", + " ✅ Block-wise quantization (e.g., block size = 128) enables hardware-friendly tensor layouts that optimize for tensor core utilization on NVIDIA GPUs.\n", + "\n", + " ✅ The TensorRT-LLM ModelOpt PTQ API provides a flexible and high-level interface for experimenting with quantization formats, including full customization of AWQ configs.\n", + "\n", + " ✅ Exported models remain compatible with Hugging Face interfaces, making them easy to use in production pipelines or deploy via inference frameworks like vLLM or TensorRT-LLM." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5c3e01cc-2149-4d39-a600-c213577e8080", + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.5" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/llm_ptq/notebooks/3_PTQ_AutoQuantization.ipynb b/examples/llm_ptq/notebooks/3_PTQ_AutoQuantization.ipynb new file mode 100644 index 000000000..1e8a33a4a --- /dev/null +++ b/examples/llm_ptq/notebooks/3_PTQ_AutoQuantization.ipynb @@ -0,0 +1,430 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "b80da483-f81a-408f-9755-6c1bc9d84119", + "metadata": {}, + "source": [ + "# AutoQuantization with TensorRT Model Optimizer PTQ\n", + "\n", + "This notebook demonstrates how to use ModelOpt PTQ's auto_quantize feature to perform automated mixed-precision quantization on the Meta-LLaMA-3-8B model. You'll define a target effective bit rate (e.g., 8.0), provide a search space of quantization formats, and optionally include KV cache quantization.\n", + "\n", + "The process automatically searches the quantization format and layer mapping that best satisfies the target bit constraint while minimizing accuracy loss—using loss-based scoring and real calibration data.\n", + "\n", + "Key Dependendancies: \n", + "- nvidia-modelopt\n", + "- torch\n", + "- transformers" + ] + }, + { + "cell_type": "markdown", + "id": "c7ae88b2-b51b-40fc-8ef3-4ff0bf37916b", + "metadata": {}, + "source": [ + "# Applying AutoQuantization" + ] + }, + { + "cell_type": "markdown", + "id": "2247af06-7752-4d28-8fa1-22af961130a3", + "metadata": {}, + "source": [ + "### 1. Import Dependencies\n", + "\n", + "Load general-purpose and reproducibility packages. Random seeds will be set for deterministic calibration, and the Hugging Face login is required to pull the model.\n", + "\n", + "- Import core ModelOpt utilities for quantization, calibration, and dataset handling.\n", + "- init_quantized_weights is used internally by some ModelOpt features for weight initialization—no need to call directly here, but it supports hybrid quantized weight loading when needed.\n", + "- Set fixed seeds for reproducibility. This ensures consistent calibration data selection and loss values during scoring." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "942eb87c-3e69-49dd-a7b6-933e5a700b61", + "metadata": {}, + "outputs": [], + "source": [ + "import random\n", + "import time\n", + "\n", + "import numpy as np\n", + "import torch\n", + "from huggingface_hub import login\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer\n", + "\n", + "import modelopt.torch.quantization as mtq\n", + "from modelopt.torch.utils.dataset_utils import (\n", + " create_forward_loop,\n", + " get_dataset_dataloader,\n", + " get_max_batch_size,\n", + ")\n", + "\n", + "SEED = 1234\n", + "random.seed(SEED)\n", + "np.random.seed(SEED)\n", + "torch.manual_seed(SEED)" + ] + }, + { + "cell_type": "markdown", + "id": "2f71edaf-e3b5-469a-a36c-b1b93c033c60", + "metadata": {}, + "source": [ + "### 2. Set Configurations and Login to Hugging Face\n", + "\n", + "Define all major tuning knobs:\n", + "\n", + "- `EFFECTIVE_BITS` is the average precision target across quantized layers.\n", + "- `Q_FORMATS` defines the list of quantization formats used during the AutoQuant search.\n", + "- `KV_FORMAT` allows optional quantization of KV cache, applied after main quant.\n", + "- `EXPORT_FMT` supports exporting to Hugging Face (`hf`) or TensorRT-LLM (`tensorrt_llm`).\n", + "- Also logs into Hugging Face\n", + "\n", + "💡 Try adding formats like `\"nvfp4\"`, `\"w4a8_awq\"` to explore tradeoffs." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a5814162-f577-4c05-8573-36aad32323c3", + "metadata": {}, + "outputs": [], + "source": [ + "MODEL_ID = \"meta-llama/Meta-Llama-3-8B\"\n", + "DATASET = \"cnn_dailymail\"\n", + "CALIB_SAMPLES = 512\n", + "EFFECTIVE_BITS = 6.0 # search target\n", + "Q_FORMATS = \"fp8,int4_awq\" # search space\n", + "KV_FORMAT = \"none\" # or \"none\" to skip\n", + "EXPORT_DIR = \"llama3_8b_autoq\" # output folder\n", + "EXPORT_FMT = \"tensorrt_llm\" # or \"hf\"\n", + "# ----------------------------\n", + "DEVICE = \"cuda\"\n", + "DTYPE = torch.float16 # keep default for faster search\n", + "\n", + "login()" + ] + }, + { + "cell_type": "markdown", + "id": "8e0a92e4-7bf1-485b-a283-3c69b64d2767", + "metadata": {}, + "source": [ + "### 3. Load Model and Tokenizer\n", + "\n", + "Load model and tokenizer into memory:\n", + "\n", + "- torch_dtype=torch.float16 is used to reduce memory during calibration.\n", + "- Left padding is preferred for decoder-only LLMs to better align prompt-token positions during calibration.\n", + "\n", + "⚠️ Ensure the pad_token is set for batching; some LLaMA variants may not have it by default." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "c3c91ffc-c73d-44df-b833-e773dca0457e", + "metadata": { + "scrolled": true + }, + "outputs": [], + "source": [ + "model = AutoModelForCausalLM.from_pretrained(MODEL_ID, torch_dtype=DTYPE).to(DEVICE)\n", + "tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)\n", + "tokenizer.padding_side = \"left\"\n", + "tokenizer.pad_token = tokenizer.eos_token\n", + "model.eval()" + ] + }, + { + "cell_type": "markdown", + "id": "5259ff10-3d05-44fa-b650-bcdbcadc74ad", + "metadata": {}, + "source": [ + "### 4. Configure Data Loader and Forward Loop\n", + "\n", + "- Use a small number of real samples to capture representative activations and enable loss-based scoring.\n", + "- include_labels=True is important for loss computation, which guides AutoQuant format decisions.\n", + "\n", + "⚙️ get_max_batch_size() estimates the largest batch that fits in memory given model size and hardware." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a90e25a0-f271-4ca4-ae53-f3f77a09fba7", + "metadata": {}, + "outputs": [], + "source": [ + "batch_size = min(get_max_batch_size(model), CALIB_SAMPLES)\n", + "calib_loader = get_dataset_dataloader(\n", + " dataset_name=DATASET,\n", + " tokenizer=tokenizer,\n", + " batch_size=batch_size,\n", + " num_samples=CALIB_SAMPLES,\n", + " device=DEVICE,\n", + " include_labels=True,\n", + ")\n", + "forward_loop = create_forward_loop(dataloader=calib_loader)\n", + "print(f\"Calibration batches: {len(calib_loader)} | Batch size: {batch_size}\")" + ] + }, + { + "cell_type": "markdown", + "id": "3056b7c2-aee1-4d63-b062-42430750104f", + "metadata": {}, + "source": [ + "### 5. Possible Quantization Configurations\n", + "\n", + "Define lookup tables for available quantization config presets. These are used to construct the format search space for AutoQuant.\n", + "\n", + "✅ You can freely extend these dictionaries to add custom formats or constraints." + ] + }, + { + "cell_type": "code", + "execution_count": 19, + "id": "01c135ff-6acb-4a0e-a447-0774d9d269cc", + "metadata": {}, + "outputs": [], + "source": [ + "QUANT_CFG = {\n", + " \"int8\": mtq.INT8_DEFAULT_CFG,\n", + " \"int8_sq\": mtq.INT8_SMOOTHQUANT_CFG,\n", + " \"fp8\": mtq.FP8_DEFAULT_CFG,\n", + " \"int4_awq\": mtq.INT4_AWQ_CFG,\n", + " \"nvfp4\": mtq.NVFP4_DEFAULT_CFG,\n", + " \"nvfp4_awq\": mtq.NVFP4_AWQ_LITE_CFG,\n", + " \"w4a8_awq\": mtq.W4A8_AWQ_BETA_CFG,\n", + "}\n", + "\n", + "KV_CFG = {\n", + " \"none\": None,\n", + " \"fp8\": mtq.FP8_KV_CFG[\"quant_cfg\"],\n", + " \"nvfp4\": mtq.NVFP4_KV_CFG[\"quant_cfg\"],\n", + " \"nvfp4_affine\": mtq.NVFP4_AFFINE_KV_CFG[\"quant_cfg\"],\n", + "}" + ] + }, + { + "cell_type": "markdown", + "id": "6d3c58b3-816d-47e9-89b4-65dc7362f22f", + "metadata": {}, + "source": [ + "### 6. Start AutoQuantization Search and Optimization\n", + "\n", + "Wrap the model’s native loss function. Required for format scoring during AutoQuant—smaller loss = better format match.\n", + "\n", + "Automatically search the best per-layer quantization format mapping:\n", + "\n", + "- Constraints guide the average bit precision.\n", + "- Loss is evaluated across candidate formats to preserve accuracy.\n", + "- disabled_layers=[\"*lm_head*\"] keeps the final layer unquantized (important for generation quality).\n", + "\n", + "🔍 Verbose mode shows layer-level decisions and scoring for each candidate format." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "759cf632-25e3-4ab0-a217-22843ea15138", + "metadata": { + "collapsed": true, + "jupyter": { + "outputs_hidden": true + } + }, + "outputs": [], + "source": [ + "def loss_fn(out, batch): # tiny wrapper around HF loss\n", + " return out.loss\n", + "\n", + "\n", + "print(\"🚧 Launching auto_quantize ...\")\n", + "t0 = time.time()\n", + "\n", + "model, _ = mtq.auto_quantize(\n", + " model,\n", + " constraints={\"effective_bits\": EFFECTIVE_BITS},\n", + " data_loader=calib_loader,\n", + " forward_step=lambda m, b: m(**b),\n", + " loss_func=loss_fn,\n", + " quantization_formats=[QUANT_CFG[q] for q in Q_FORMATS.split(\",\")],\n", + " num_calib_steps=len(calib_loader),\n", + " num_score_steps=len(calib_loader),\n", + " verbose=True,\n", + " disabled_layers=[\"*lm_head*\"], # keep LM head in fp16\n", + ")\n", + "print(f\"✅ Done in {time.time() - t0:.1f}s\")" + ] + }, + { + "cell_type": "markdown", + "id": "a4ba5f17-932f-41c8-a27f-c6b812a041a5", + "metadata": {}, + "source": [ + "### 7. [Optional] KV Cache AutoQuantization\n", + "\n", + "- This happens after main quantization.\n", + "- Only KV-specific quantizers are enabled during this pass.\n", + "\n", + "⚠️ Quantizing KV cache may affect generation performance and context retention—test thoroughly." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "21b805b1-8583-4fdb-afeb-b266821ebe26", + "metadata": {}, + "outputs": [], + "source": [ + "if KV_FORMAT != \"none\":\n", + " print(f\"Enabling KV cache quantization ⟶ {KV_FORMAT}\")\n", + " kv_cfg = KV_CFG[KV_FORMAT]\n", + "\n", + " # Plug only the KV quantizers\n", + " mtq.set_quantizer_by_cfg(model, quant_cfg=kv_cfg)\n", + "\n", + " # Calibrate **only** those quantizers\n", + " with mtq.set_quantizer_by_cfg_context(model, {\"*\": {\"enable\": False}, **kv_cfg}):\n", + " mtq.calibrate(model, algorithm=\"max\", forward_loop=forward_loop)\n", + "else:\n", + " print(\"KV cache left unquantized.\")" + ] + }, + { + "cell_type": "markdown", + "id": "eaf6f061-a1d8-45c1-a390-aa5730d5a60d", + "metadata": {}, + "source": [ + "### 8. Inspect the Quantized Layers\n", + "\n", + "Print a full summary of quantized layers, formats, and bit precision estimates—useful for debugging or profiling." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "b20343e9-2ae0-4434-a19d-d7b6e912ba74", + "metadata": { + "scrolled": true + }, + "outputs": [], + "source": [ + "mtq.print_quant_summary(model)" + ] + }, + { + "cell_type": "markdown", + "id": "6d5ad1e8-f200-4efa-aa23-dd7eec52516c", + "metadata": {}, + "source": [ + "### 9. Quick Test of Quantized Model\n", + "\n", + "Sanity check: Run a quick generation to verify the quantized model produces reasonable output.\n", + "\n", + "🧪 Consider adding prompts from your real use case to validate quality before deployment." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9992c873-02ec-4886-a47d-d38f8e4b405d", + "metadata": {}, + "outputs": [], + "source": [ + "sample = \"Tell me a short story about a quantized llama.\"\n", + "inputs = tokenizer(sample, return_tensors=\"pt\").to(DEVICE)\n", + "with torch.inference_mode():\n", + " gen_ids = model.generate(**inputs, max_new_tokens=50, do_sample=True)\n", + "print(tokenizer.decode(gen_ids[0], skip_special_tokens=True))" + ] + }, + { + "cell_type": "markdown", + "id": "e7598596-9eca-42a1-bf7f-665a8023c48c", + "metadata": {}, + "source": [ + "### 10. Export Model for TensorRT-LLM\n", + "\n", + "Export the quantized model to the desired format:\n", + "\n", + "- Use tensorrt_llm for high-performance deployment on NVIDIA accelerators.\n", + "- Use hf to reload with Hugging Face APIs or inference frameworks like vLLM.\n", + "\n", + "📁 Check the contents of the output folder to confirm all weights/configs are present." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "13ed5ad6-fd9f-44dd-b2c0-aefc46d1d38b", + "metadata": {}, + "outputs": [], + "source": [ + "from modelopt.torch.export import export_hf_checkpoint, export_tensorrt_llm_checkpoint\n", + "\n", + "if EXPORT_FMT == \"tensorrt_llm\":\n", + " export_tensorrt_llm_checkpoint(\n", + " model,\n", + " model_type=\"llama\",\n", + " export_dir=EXPORT_DIR,\n", + " inference_tensor_parallel=1,\n", + " inference_pipeline_parallel=1,\n", + " )\n", + "else:\n", + " export_hf_checkpoint(model, export_dir=EXPORT_DIR)\n", + "\n", + "print(f\"📦 Saved quantized model to → {EXPORT_DIR}\")" + ] + }, + { + "cell_type": "markdown", + "id": "638588fa-3781-4ac4-a77e-140331caba80", + "metadata": {}, + "source": [ + "# ✅ Conclusion & Key Takeaways\n", + " ✅ AutoQuant in TensorRT-LLM ModelOpt enables fast, automated mixed-precision quantization by searching across multiple formats (e.g., FP8, INT4-AWQ) to meet a user-defined effective bit constraint.\n", + "\n", + " ✅ Using a small calibration set with loss-based scoring, AutoQuant intelligently selects the optimal quantization format per layer—balancing model size, performance, and accuracy.\n", + "\n", + " ✅ The workflow supports flexible search spaces and fine-grained control over disabled layers, block sizes, and forward calibration passes.\n", + "\n", + " ✅ Optional KV cache quantization provides further memory and bandwidth savings, but should be enabled only after validating generation quality.\n", + "\n", + " ✅ Exported models are fully compatible with both Hugging Face and TensorRT-LLM inference runtimes—enabling rapid deployment across a wide range of applications." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "b04c029c-5f9c-4f0c-8db7-db7e9d8800ab", + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.5" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/llm_qat/README.md b/examples/llm_qat/README.md index f51adf3ca..902def953 100644 --- a/examples/llm_qat/README.md +++ b/examples/llm_qat/README.md @@ -150,7 +150,7 @@ for QAT over PTQ alone. #### NeMo QAT/QAD Simplified Flow Example -The [examples/nemo_run/qat](../nemo_run/qat) directory also contains an end-to-end NeMo QAT Simplified Flow example, which supports both QAT with cross-entropy loss and QAD (quantization-aware distillation) with knowledge-distillation loss between the full-precision teacher and quantized student models. +The [examples/nemo_run/qat](../nemo_run/qat) directory also contains an end-to-end NeMo QAT Simplified Flow example, which supports both QAT with cross-entropy loss and QAD (quantization-aware distillation) with knowledge-distillation loss between the full-precision teacher and quantized student models. Refer to [README](../nemo_run/qat/README.md) for more detail. #### Testing QAT model with LLM benchmarks for accuracy evaluation diff --git a/examples/llm_qat/main.py b/examples/llm_qat/main.py index ec3077ce1..8ebbf21eb 100644 --- a/examples/llm_qat/main.py +++ b/examples/llm_qat/main.py @@ -243,6 +243,7 @@ def train(): distill_kwargs["distill_config"] = distill_config trainer_cls = QADTrainer if training_args.distill else QATTrainer + training_args.lora_config = get_lora_config() trainer = trainer_cls( model=model, processing_class=tokenizer, @@ -253,14 +254,6 @@ def train(): **data_module, ) - # add lora adapter - if training_args.lora: - model.add_adapter(get_lora_config(), adapter_name="adapter") - - # compress model weights after lora adapter inserted to prevent training error - if checkpoint is None and quant_args.compress: - mtq.compress(model) - # There could be GPU memory leak during QAT causing OOM. This is a workaround to fix it. monkey_patch_training_step_to_fix_memory_leak(trainer) diff --git a/examples/nemo_run/qat/README.md b/examples/nemo_run/qat/README.md index 8d2fddfa7..4f11c5cd3 100644 --- a/examples/nemo_run/qat/README.md +++ b/examples/nemo_run/qat/README.md @@ -17,8 +17,38 @@ To run SFT properly you may also need to clone NeMo and Megatron-LM at the respe ### Running the Flow +#### QAT + From the `nemo_run` folder, launch the example with `python qat/nemo_qat_flow.py --model-name --finetune-recipe `. Available NeMo recipe names are listed [here](https://github.com/NVIDIA/NeMo/tree/main/nemo/collections/llm/recipes). To provide your own custom dataset, use the `--data-path` flag, otherwise the default [LIMA](https://huggingface.co/datasets/GAIR/lima) dataset will be used. +To perform QAT, run: + +```bash +python qat/nemo_qat_flow.py \ + --model-name meta-llama/Meta-Llama-3.1-8B-Instruct \ + --finetune-recipe llama31_8b \ + --algorithm fp8 \ + --chat-template llama_chat_template.txt \ + --experiment llama3_qat_nemo +``` + +> **_NOTE:_** To enable KV cache quantization, add `--enable-kv-cache` and specify qformat using `--kv-cache-qformat `. + +#### QAD + +In order to train using QAD, launch the example with `python qat/nemo_qat_flow.py --model-name --distill`. It will utilize `distillation_recipe` with quantized student model and full precision teacher model to train the quantized model. + +To perform QAD training, run: + +```bash +python qat/nemo_qat_flow.py \ + --model-name meta-llama/Meta-Llama-3.1-8B-Instruct \ + --distill \ + --algorithm fp8 \ + --chat-template llama_chat_template.txt \ + --experiment llama3_qad_nemo +``` + ### Custom Chat Template By default the script will use the model/tokenizer's chat template, which may not contain the `{% generation %}` and `{% endgeneration %}` tags around the assistant tokens which are needed to generate the assistant loss mask (see [this PR](https://github.com/huggingface/transformers/pull/30650)). To provide path to a custom chat template, use the `--chat-template ` flag. diff --git a/examples/nemo_run/qat/nemo_qat_flow.py b/examples/nemo_run/qat/nemo_qat_flow.py index 364dfb2d0..5b1894108 100644 --- a/examples/nemo_run/qat/nemo_qat_flow.py +++ b/examples/nemo_run/qat/nemo_qat_flow.py @@ -96,6 +96,18 @@ def get_parser(): help="Number of GPUs for quantization. Some models require a different number of GPUs for PTQ vs training.", default=8, ) + parser.add_argument( + "--kv-cache-qformat", + type=str, + default="fp8", + choices=["fp8", "nvfp4"], + help="KV-cache quantization format", + ) + parser.add_argument( + "--enable_kv_cache", help="Enables KV-cache quantization", action="store_true" + ) + parser.add_argument("--disable_kv_cache", dest="enable_kv_cache", action="store_false") + parser.set_defaults(enable_kv_cache=None) return parser @@ -181,6 +193,11 @@ if __name__ == "__main__": ptq_model_out, "--export_format", "nemo", + "--algorithm", + args.algorithm, + "--kv_cache_qformat", + args.kv_cache_qformat, + "--enable_kv_cache" if args.enable_kv_cache else "--disable_kv_cache", "-ctp", f"{args.ptq_gpus}", ], diff --git a/examples/onnx_ptq/README.md b/examples/onnx_ptq/README.md index ce1e45091..1ac2e7693 100644 --- a/examples/onnx_ptq/README.md +++ b/examples/onnx_ptq/README.md @@ -11,12 +11,18 @@ Note that this example is for ONNX model quantization. For end to end quantizati ### Linux -Please follow the main [README](../README.md#docker) to build the docker image with TensorRT 10.x pre-installed. +Build the Docker image (will be tagged `docker.io/library/onnx_ptq_examples:latest`)
-The container can be run with the following command: +This Docker image includes the latest publicly available TensorRT version, providing access to cutting-edge features and superior performance compared to the public modelopt_examples Docker [image](https://github.com/NVIDIA/TensorRT-Model-Optimizer/tree/main/docker/Dockerfile). ```bash -docker run --user 0:0 -it --gpus all --shm-size=2g -v /path/to/ImageNet/dataset:/workspace/imagenet docker.io/library/modelopt_examples:latest +./docker/build.sh +``` + +Run the docker image + +```bash +docker run --user 0:0 -it --gpus all --shm-size=2g -v /path/to/ImageNet/dataset:/workspace/imagenet docker.io/library/onnx_ptq_examples:latest ``` ### Prepare the example model diff --git a/examples/onnx_ptq/docker/Dockerfile b/examples/onnx_ptq/docker/Dockerfile new file mode 100644 index 000000000..6837797ba --- /dev/null +++ b/examples/onnx_ptq/docker/Dockerfile @@ -0,0 +1,37 @@ +FROM nvcr.io/nvidia/tensorrt:25.06-py3 + +ARG CMAKE_VERSION=3.28.0 + +ENV PIP_EXTRA_INDEX_URL="https://pypi.nvidia.com" \ + PIP_NO_CACHE_DIR=off + +RUN python -m pip install --upgrade pip \ + && pip install cmake==${CMAKE_VERSION} \ + && mkdir -p -m 0600 ~/.ssh \ + && ssh-keyscan github.com >> ~/.ssh/known_hosts + +WORKDIR /workspace + +RUN pip install tensorrt==10.12.0.36 && \ + export TRT_PATH=$(python -c "import tensorrt; import os; print(os.path.dirname(tensorrt.__file__))") && \ + export LD_LIBRARY_PATH="$TRT_PATH/lib:${LD_LIBRARY_PATH}" && \ + export PATH="$TRT_PATH/bin:${PATH}" + +# Update PATH variables for local TensorRT installation +ENV LD_LIBRARY_PATH="/workspace/TensorRT/lib:${LD_LIBRARY_PATH}" \ + PATH="/workspace/TensorRT/bin:${PATH}" + +# Copy application code and install requirements +COPY modelopt modelopt/modelopt +COPY examples/onnx_ptq modelopt/examples/onnx_ptq +COPY setup.py modelopt/setup.py +COPY pyproject.toml modelopt/pyproject.toml + +# Install onnx_ptq requirements +RUN pip install -r modelopt/examples/onnx_ptq/requirements.txt + +# Install modelopt +RUN pip install -e "./modelopt[hf,onnx]" + +# Allow users to run without root +RUN chmod -R 777 /workspace diff --git a/examples/onnx_ptq/docker/build.sh b/examples/onnx_ptq/docker/build.sh new file mode 100755 index 000000000..f1ac572ea --- /dev/null +++ b/examples/onnx_ptq/docker/build.sh @@ -0,0 +1,131 @@ +#!/bin/bash +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +set -euo pipefail # Exit on error, undefined vars, pipe failures + +# Default values +IMAGE_NAME="modelopt_onnx_examples:latest" +DOCKERFILE_PATH="examples/onnx_ptq/docker/Dockerfile" + +# Function to show usage +usage() { + cat << EOF +Usage: $0 [OPTIONS] + +Options: + -t, --tag IMAGE_NAME Docker image name (default: $IMAGE_NAME) + -h, --help Show this help message + +This script automatically detects whether you're running from: + • modelopt/ root directory + • modelopt/examples/onnx_ptq/ directory + +and builds the Docker image accordingly. +EOF + exit 1 +} + +# Parse arguments +while [[ $# -gt 0 ]]; do + case $1 in + -t|--tag) + [[ -n "${2:-}" ]] || { echo "Error: --tag requires a value"; exit 1; } + IMAGE_NAME="$2" + shift 2 + ;; + -h|--help) + usage + ;; + *) + echo "Error: Unknown option '$1'" + usage + ;; + esac +done + +# Function to find modelopt root directory +find_modelopt_root() { + local current_dir="$PWD" + + # Check current directory first + if [[ -f "setup.py" && -f "pyproject.toml" && -d "modelopt" ]]; then + echo "$current_dir" + return 0 + fi + + # Check parent directories (up to 3 levels) + for i in {1..3}; do + local parent_dir + parent_dir=$(dirname "$current_dir") + [[ "$parent_dir" == "$current_dir" ]] && break # Reached filesystem root + + if [[ -f "$parent_dir/setup.py" && -f "$parent_dir/pyproject.toml" && -d "$parent_dir/modelopt" ]]; then + echo "$parent_dir" + return 0 + fi + current_dir="$parent_dir" + done + + return 1 +} + +# Find modelopt root directory +echo "🔍 Locating modelopt root directory..." +if ROOT_DIR=$(find_modelopt_root); then + echo "✅ Found modelopt root: $ROOT_DIR" + cd "$ROOT_DIR" +else + cat << EOF +❌ Error: Cannot locate modelopt root directory. + +Expected structure: + modelopt/ + ├── setup.py + ├── pyproject.toml + ├── modelopt/ + └── examples/onnx_ptq/docker/ + +Please run this script from within the modelopt repository. +EOF + exit 1 +fi + +# Validate that Dockerfile exists +if [[ ! -f "$DOCKERFILE_PATH" ]]; then + echo "❌ Error: Dockerfile not found at $DOCKERFILE_PATH" + exit 1 +fi + +# Build Docker image +echo "🐳 Building Docker image..." +echo " • Image name: $IMAGE_NAME" +echo " • Build context: $(pwd)" +echo " • Dockerfile: $DOCKERFILE_PATH" +echo + +docker build \ + --file "$DOCKERFILE_PATH" \ + --tag "$IMAGE_NAME" \ + . \ + "$@" + +echo +echo "✅ Docker image built successfully: $IMAGE_NAME" +echo +echo "🚀 To run the container:" +echo " docker run --user 0:0 -it --gpus all --shm-size=2g \\" +echo " -v /path/to/ImageNet/dataset:/workspace/imagenet \\" +echo " $IMAGE_NAME" diff --git a/examples/onnx_ptq/llm_export.py b/examples/onnx_ptq/llm_export.py new file mode 100644 index 000000000..437895b20 --- /dev/null +++ b/examples/onnx_ptq/llm_export.py @@ -0,0 +1,348 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""This script is used to export a LLM model to ONNX and perform quantization.""" + +import argparse +import os +import shutil +import time + +import onnx +import onnx_graphsurgeon as gs +import torch +from packaging.version import Version +from transformers import AutoTokenizer + +from modelopt.onnx.llm_export.utils.export_utils import ( + ModelLoader, + WrapperModelForCausalLM, + llm_to_onnx, +) + + +def llm_arguments(): + """Parse the arguments for the llm export script.""" + parser = argparse.ArgumentParser() + parser.add_argument( + "--torch_dir", type=str, help="The folder of HF PyTorch model ckpt", required=False + ) + parser.add_argument( + "--dtype", + type=str, + default="fp16", + choices=["fp16", "fp8", "nvfp4"], + help="The precision of onnx export", + ) + + parser.add_argument( + "--lm_head", + type=str, + default="fp16", + choices=["fp16"], + help="The precision of lm_head. Currently only fp16 is tested and supported", + ) + parser.add_argument( + "--output_dir", + type=str, + help="The directory to store the generated ONNX model", + required=True, + ) + + parser.add_argument( + "--onnx_path", + type=str, + help="Pass this option when you have existing onnx to surgeon", + required=False, + ) + parser.add_argument( + "--save_original", + action="store_true", + default=False, + help="Save the original ONNX from torch.onnx.export without any modification", + ) + parser.add_argument( + "--dataset_dir", type=str, help="The path of dataset for quantization", required=False + ) + parser.add_argument( + "--config_path", + type=str, + help="The path of config.json, in case it is not with the PyTorch or ONNX file", + default=None, + ) + return parser + + +def get_config_path(args): + """Look for config.json. It is recommended to keep a copy per ONNX path. + + Args: + args: argparse.Namespace + + Returns: + str: The path of config.json + """ + if args.config_path and os.path.exists(args.config_path): + return args.config_path + if args.torch_dir: + torch_config = os.path.join(args.torch_dir, "config.json") + if os.path.exists(torch_config): + return torch_config + if args.onnx_path: + onnx_config = os.path.join(os.path.dirname(args.onnx_path), "config.json") + if os.path.exists(onnx_config): + return onnx_config + print("Warning: cannot find config.json. Please pass in --config_path.") + return None + + +def export_raw_llm( + model, + output_dir, + dtype, + config_path, + torch_dir, + lm_head_precision="fp16", + dataset_dir="", + wrapper_cls=WrapperModelForCausalLM, + extra_inputs={}, + extra_dyn_axes={}, +): + """Export raw llm model to ONNX and perform quantization. + + Args: + model: torch.nn.module + output_dir: str + dtype: str + config_path: str + torch_dir: str, Used for loading tokenizer for quantization + dataset_dir: str, Used for quantization + wrapper_cls: class, Used for wrapping the model + extra_inputs: dict, Used for extra inputs + extra_dyn_axes: dict, Used for extra dynamic axes + """ + os.makedirs(output_dir, exist_ok=True) + + if dtype == "fp16": + print("Loading fp16 ONNX model...") + + llm_to_onnx( + wrapper_cls(model), output_dir, extra_inputs=extra_inputs, extra_dyn_axes=extra_dyn_axes + ) + shutil.copy(config_path, os.path.join(output_dir, "config.json")) + + # Need to quantize model to fp8, int4 or nvfp4 + if dtype in ["fp8", "nvfp4"]: + # Avoid import modelopt when no quantization is needed + from modelopt.onnx.llm_export.utils.quantization_utils import quantize + from modelopt.torch.export import export_hf_checkpoint + from modelopt.torch.quantization.utils import is_quantized_linear + + tokenizer = AutoTokenizer.from_pretrained(torch_dir, trust_remote_code=True) + modelopt_state = os.path.join(torch_dir, "modelopt_state.pth") + if not os.path.exists(modelopt_state): + model = quantize(model, tokenizer, dtype, lm_head_precision, dataset_dir) + + if dtype == "nvfp4": + # This is required for nvfp4 ONNX export + for module in model.modules(): + assert not isinstance(module, torch.nn.Linear) or is_quantized_linear(module) + if isinstance(module, torch.nn.Linear): + module.input_quantizer._trt_high_precision_dtype = "Half" + module.input_quantizer._onnx_quantizer_type = "dynamic" + module.weight_quantizer._onnx_quantizer_type = "static" + + if dtype in {"fp8", "nvfp4"}: + print(f"Exporting {dtype} ONNX model from quantized PyTorch model...") + llm_to_onnx( + wrapper_cls( + model, + ), + output_dir, + extra_inputs=extra_inputs, + extra_dyn_axes=extra_dyn_axes, + ) + shutil.copy(config_path, os.path.join(output_dir, "config.json")) + + # Compress weights + quantized_model_dir = f"{output_dir}_{dtype}_quantized" + os.makedirs(quantized_model_dir, exist_ok=True) + with torch.inference_mode(): + export_hf_checkpoint(model, dtype=torch.float16, export_dir=quantized_model_dir) + + return model.state_dict() + + +def surgeon_llm( + raw_onnx_path, + output_dir, + dtype, + config_path, + lm_head_precision="fp16", +): + """Surgeon raw llm onnx to fit TRT. + + For example, insert quantization q/dq nodes. + + Args: + raw_onnx_path: str + output_dir: str + dtype: str + config_path: str + lm_head_precision: str + """ + + t0 = time.time() + graph = gs.import_onnx(onnx.load(raw_onnx_path)) + t1 = time.time() + print(f"Importing ONNX graph takes {t1 - t0}s.") + graph.fold_constants().cleanup().toposort() + + if dtype == "fp8" or lm_head_precision == "fp8": + from modelopt.onnx.llm_export.utils.surgeon_utils import fold_fp8_qdq_to_dq + + graph = fold_fp8_qdq_to_dq(graph) + + os.makedirs(output_dir, exist_ok=True) + t2 = time.time() + + onnx_model = gs.export_onnx(graph) + + if dtype == "nvfp4": + t4 = time.time() + from modelopt.onnx.quantization.qdq_utils import fp4qdq_to_2dq + + onnx_model = fp4qdq_to_2dq(onnx_model, verbose=True) + t5 = time.time() + print(f"nvfp4 qdq to 2 dqs inserted in {t5 - t4}.") + + output_onnx_name = f"{output_dir}/model.onnx" + print( + f"Saving ONNX files in {output_dir}. All existing ONNX in the folder will be overwritten." + ) + for filename in os.listdir(output_dir): + file_path = os.path.join(output_dir, filename) + try: + if ( + os.path.isfile(file_path) or os.path.islink(file_path) + ) and ".json" not in file_path: + os.unlink(file_path) + + except Exception as e: + print(f"Failed to delete {file_path}. Reason: {e}") + + onnx.save_model( + onnx_model, + output_onnx_name, + save_as_external_data=True, + all_tensors_to_one_file=True, + location="onnx_model.data", + convert_attribute=True, + ) + + if os.path.exists(config_path): + if config_path.endswith("config.json"): + shutil.copy(config_path, os.path.join(output_dir, "config.json")) + else: + shutil.copy( + os.path.join(config_path, "config.json"), os.path.join(output_dir, "config.json") + ) + + t3 = time.time() + print(f"Surgeon LLM completed in {t3 - t2}s.") + + +def check_dtype_support(args): + """Check whether the dtype is supported by DriveOS LLM SDK. + + Returns False if it is not supported because of: + 1. Modelopt < 0.23.0 does not support nvfp4 + """ + + def get_modelopt_version(): + try: + import modelopt + + return Version(modelopt.__version__) + except Exception as e: + print(f"Modelopt version cannot be parsed. Reason: {e!s}") + + if (args.dtype == "nvfp4") and get_modelopt_version() < Version("0.23.0"): + print( + "nvfp4 is not supported by installed modelopt version. Please upgrade to 0.23.0 or above for nvfp4 export." + ) + return False + + return True + + +def main(args): + """Main function to export the LLM model to ONNX.""" + assert args.torch_dir or args.onnx_path, ( + "You need to provide either --torch_dir or --onnx_path to process the export script." + ) + start_time = time.time() + + if not check_dtype_support(args): + return + + if args.onnx_path: + raw_onnx_path = args.onnx_path + + model_loader = ModelLoader( + args.torch_dir, + args.config_path, + ) + + if args.torch_dir: + # Exporting ONNX from PyTorch model + model = model_loader.load_model() + onnx_dir = args.output_dir + "_raw" if args.save_original else args.output_dir + # Surgeon graph based on precision + raw_onnx_path = f"{onnx_dir}/model.onnx" + extra_inputs, extra_dyn_axes = {}, {} + export_raw_llm( + model=model, + output_dir=onnx_dir, + dtype=args.dtype, + config_path=args.config_path, + torch_dir=args.torch_dir, + lm_head_precision=args.lm_head, + dataset_dir=args.dataset_dir, + wrapper_cls=WrapperModelForCausalLM, + extra_inputs=extra_inputs, + extra_dyn_axes=extra_dyn_axes, + ) + + # Providing the config path to config.json results in a hf validation error for internvl_chat. + surgeon_llm( + raw_onnx_path=raw_onnx_path, + output_dir=args.output_dir, + dtype=args.dtype, + config_path=args.config_path, + lm_head_precision=args.lm_head, + ) + + end_time = time.time() + print( + f"LLM ONNX saved to {args.output_dir} with {args.dtype} precision in {end_time - start_time}s." + ) + + +if __name__ == "__main__": + parser = llm_arguments() + args = parser.parse_args() + args.config_path = get_config_path(args) + main(args) diff --git a/examples/speculative_decoding/README.md b/examples/speculative_decoding/README.md index 5dde2efa2..f9b430e2a 100644 --- a/examples/speculative_decoding/README.md +++ b/examples/speculative_decoding/README.md @@ -23,13 +23,13 @@ Then, we adapt the fine-tuning data by calling this server. In this example, we ```sh git clone https://huggingface.co/datasets/nvidia/Daring-Anteater -python3 server_generate.py --data_path Daring-Anteater/train.jsonl --output_path finetune/data.jsonl --max_token 512 +python3 server_generate.py --data_path Daring-Anteater/train.jsonl --output_path finetune/data.jsonl --max_token 512 --chat ``` To add a system prompt, use the `--system_prompt` argument: ```sh -python3 server_generate.py --data_path Daring-Anteater/train.jsonl --output_path finetune/data.jsonl --max_token 512 --system_prompt +python3 server_generate.py --data_path Daring-Anteater/train.jsonl --output_path finetune/data.jsonl --max_token 512 --chat --system_prompt ``` #### SLURM Prepare Data diff --git a/examples/speculative_decoding/eagle_utils.py b/examples/speculative_decoding/eagle_utils.py index ba95064cd..8ac292a37 100644 --- a/examples/speculative_decoding/eagle_utils.py +++ b/examples/speculative_decoding/eagle_utils.py @@ -42,12 +42,24 @@ def preprocess(examples, tokenizer): for i in range(len(examples)): messages = [] source = examples[i]["conversations"] - if source[0]["from"].lower() != "user": + + # Detect format: either role/content or from/value + def get_role_content(item): + if "role" in item and "content" in item: + return item["role"], item["content"] + elif "from" in item and "value" in item: + return item["from"], item["value"] + else: + raise ValueError(f"Unknown conversation format: {item}") + + first_role, _ = get_role_content(source[0]) + if first_role.lower() != "user": # Skip the first one if it is not from human source = source[1:] for j, sentence in enumerate(source): - assert sentence["from"].lower() == roles[j % 2], f"{i}" - messages.append({"role": sentence["from"].lower(), "content": sentence["value"]}) + role, content = get_role_content(sentence) + assert role.lower() == roles[j % 2], f"{i}" + messages.append({"role": role.lower(), "content": content}) conversation = tokenizer.apply_chat_template( messages, tokenize=False, diff --git a/examples/speculative_decoding/medusa_utils.py b/examples/speculative_decoding/medusa_utils.py index d9c8214a4..30dc238c3 100644 --- a/examples/speculative_decoding/medusa_utils.py +++ b/examples/speculative_decoding/medusa_utils.py @@ -45,7 +45,16 @@ IGNORE_TOKEN_ID = LabelSmoother.ignore_index def change_format(conversations): chat = [] for conversation in conversations: - turn = {"role": conversation["from"].lower(), "content": conversation["value"].lower()} + # Detect format: either role/content or from/value + if "role" in conversation and "content" in conversation: + role = conversation["role"] + content = conversation["content"] + elif "from" in conversation and "value" in conversation: + role = conversation["from"] + content = conversation["value"] + else: + raise ValueError(f"Unknown conversation format: {conversation}") + turn = {"role": role.lower(), "content": content.lower()} chat.append(turn) return chat diff --git a/examples/speculative_decoding/server_generate.py b/examples/speculative_decoding/server_generate.py index 057476629..72c9faa47 100644 --- a/examples/speculative_decoding/server_generate.py +++ b/examples/speculative_decoding/server_generate.py @@ -81,9 +81,24 @@ def generate_data(messages, idx, system_prompt): output_messages.append(system_message) for message in messages[::2]: - if message["role"].lower() != "user": + # Detect message format + if "from" in message and "value" in message: + role = message["from"].lower() + content = message["value"] + elif "role" in message and "content" in message: + role = message["role"].lower() + content = message["content"] + else: + raise ValueError(f"Message format not recognized: {message}") + + if role != "user": return - output_messages.append(message) + output_messages.append( + { + "role": "user", + "content": content, + } + ) try: response = client.chat.completions.create( model=model_name, @@ -156,6 +171,11 @@ if os.path.exists(args.output_path): break finished_ids = set(finished_ids) +# Ensure the output directory exists before writing to the output file +output_dir = os.path.dirname(args.output_path) +if output_dir and not os.path.exists(output_dir): + os.makedirs(output_dir, exist_ok=True) + if done: print("All conversations already generated") sys.exit() diff --git a/examples/vlm_ptq/scripts/huggingface_example.sh b/examples/vlm_ptq/scripts/huggingface_example.sh index 51c56fc02..7c6e23200 100755 --- a/examples/vlm_ptq/scripts/huggingface_example.sh +++ b/examples/vlm_ptq/scripts/huggingface_example.sh @@ -177,7 +177,6 @@ if [[ $TASKS =~ "build" ]] || [[ ! -d "$ENGINE_DIR" ]] || [[ ! $(ls -A $ENGINE_D --inference_tensor_parallel=$TP \ --inference_pipeline_parallel=$PP \ --export_fmt=$EXPORT_FORMAT \ - --vlm \ $PTQ_ARGS else echo "Quantized model config $MODEL_CONFIG exists, skipping the quantization stage" diff --git a/examples/windows/README.md b/examples/windows/README.md index f013d8d8b..1dab8de24 100644 --- a/examples/windows/README.md +++ b/examples/windows/README.md @@ -5,7 +5,7 @@ #### A Library to Quantize and Compress Deep Learning Models for Optimized Inference on Native Windows RTX GPUs [![Documentation](https://img.shields.io/badge/Documentation-latest-brightgreen.svg?style=flat)](https://nvidia.github.io/TensorRT-Model-Optimizer/) -[![version](https://img.shields.io/badge/v0.27.0-orange?label=Release)](https://pypi.org/project/nvidia-modelopt/0.27.0/) +[![version](https://img.shields.io/badge/v0.33.0-orange?label=Release)](https://pypi.org/project/nvidia-modelopt/) [![license](https://img.shields.io/badge/License-Apache%202.0-blue)](../../LICENSE) [Examples](#examples) | @@ -59,7 +59,7 @@ pip install onnxruntime-genai-directml>=0.4.0 pip install onnxruntime-directml==1.20.0 ``` -For more details, please refer to the [detailed installation instructions](https://nvidia.github.io/TensorRT-Model-Optimizer/getting_started/2_installation.html). +For more details, please refer to the [detailed installation instructions](https://nvidia.github.io/TensorRT-Model-Optimizer/getting_started/windows/_installation_for_Windows.html). ## Techniques diff --git a/examples/windows/accuracy_benchmark/README.md b/examples/windows/accuracy_benchmark/README.md index 94aea0eba..386c5d444 100644 --- a/examples/windows/accuracy_benchmark/README.md +++ b/examples/windows/accuracy_benchmark/README.md @@ -29,7 +29,7 @@ This repository provides scripts, popular third-party benchmarks, and instructio The MMLU benchmark assesses LLM performance across a wide range of tasks, producing a score between 0 and 1, where a higher score indicates better accuracy. Please refer the [MMLU Paper](https://arxiv.org/abs/2009.03300) for more details on this. -### MMLU Setup +### Setup The table below lists the setup steps to prepare your environment for evaluating LLMs using the MMLU benchmark. diff --git a/examples/windows/onnx_ptq/genai_llm/README.md b/examples/windows/onnx_ptq/genai_llm/README.md index 6e764acb9..5012c638d 100644 --- a/examples/windows/onnx_ptq/genai_llm/README.md +++ b/examples/windows/onnx_ptq/genai_llm/README.md @@ -8,7 +8,7 @@ This example takes an ONNX model as input, along with the necessary quantization ### Setup -1. Install ModelOpt-Windows. Refer [installation instructions](../README.md). +1. Install ModelOpt-Windows. Refer [installation instructions](../../README.md). 1. Install required dependencies @@ -43,7 +43,7 @@ The table below lists key command-line arguments of the ONNX PTQ example script. |---------------------------|------------------------------------------------------|-------------------------------------------------------------| | `--calib_size` | 32 (default), 64, 128 | Specifies the calibration size. | | `--dataset` | cnn (default), pilevel | Choose calibration dataset: cnn_dailymail or pile-val. | -| `--algo` | awq_lite (default), awq_clip | Select the quantization algorithm. | +| `--algo` | awq_lite (default), awq_clip, rtn, rtn_dq | Select the quantization algorithm. | | `--onnx_path` | input .onnx file path | Path to the input ONNX model. | | `--output_path` | output .onnx file path | Path to save the quantized ONNX model. | | `--use_zero_point` | True, False (default) | Enable zero-point based quantization. | @@ -54,7 +54,7 @@ The table below lists key command-line arguments of the ONNX PTQ example script. | `--awqclip_alpha_step` | 0.05 (default) | Step-size for AWQ weight clipping, user-defined | | `--awqclip_alpha_min` | 0.5 (default) | Minimum AWQ weight-clipping threshold, user-defined | | `--awqclip_bsz_col` | 1024 (default) | Chunk size in columns during weight clipping, user-defined | -| `--calibration_eps` | dml, cuda, cpu, NvTensorRtRtx (default: [dml,cpu]) | List of calibration endpoints. | +| `--calibration_eps` | dml, cuda, cpu, NvTensorRtRtx (default: [dml,cpu]) | List of execution-providers to use for session run during calibration | Run the following command to view all available parameters in the script: @@ -62,11 +62,20 @@ Run the following command to view all available parameters in the script: python quantize.py --help ``` +Note: + +1. For the `algo` argument, we have following options to choose form: awq_lite, awq_clip, rtn, rtn_dq. + - The 'awq_lite' option does core AWQ scale search and INT4 quantization. + - The 'awq_clip' option primarily does weight clipping and INT4 quantization. + - The 'rtn' option does INT4 RTN quantization with Q->DQ nodes for weights. + - The 'rtn_dq' option does INT4 RTN quantization with only DQ nodes for weights. +1. RTN algorithm doesn't use calibration-data. + Please refer to `quantize.py` for further details on command-line parameters. ### Evaluate the Quantized Model -To evaluate the quantized model, please refer to the [accuracy benchmarking](../accuracy_benchmark/README.md) and [onnxruntime-genai performance benchmarking](https://github.com/microsoft/onnxruntime-genai/tree/main/benchmark/python). +To evaluate the quantized model, please refer to the [accuracy benchmarking](../../accuracy_benchmark/README.md) and [onnxruntime-genai performance benchmarking](https://github.com/microsoft/onnxruntime-genai/tree/main/benchmark/python). ### Deployment diff --git a/examples/windows/onnx_ptq/genai_llm/quantize.py b/examples/windows/onnx_ptq/genai_llm/quantize.py index 1c2b965bf..8ec5d8b05 100644 --- a/examples/windows/onnx_ptq/genai_llm/quantize.py +++ b/examples/windows/onnx_ptq/genai_llm/quantize.py @@ -401,7 +401,7 @@ if __name__ == "__main__": "--algo", type=str, default="awq_lite", - help="Device for calibration data", + help="Algorithm or calibration-method to use. Choose from [awq_lite, awq_clip, rtn, rtn_dq]", ) parser.add_argument( "--dataset", diff --git a/examples/windows/onnx_ptq/whisper/README.md b/examples/windows/onnx_ptq/whisper/README.md index 794f321d1..73ca9f4c4 100644 --- a/examples/windows/onnx_ptq/whisper/README.md +++ b/examples/windows/onnx_ptq/whisper/README.md @@ -73,6 +73,8 @@ python .\whisper_optimum_ort_inference.py --model_name=openai/whisper-large --on ``` +A sample audio file (.wav) is provided with this example (*file*: `demo.wav`). + ## Quantization script The script `whisper_onnx_quantization.py` supports various quantization schemes for the given ONNX whisper model. diff --git a/modelopt/__init__.py b/modelopt/__init__.py index 1244a36f2..d0a778c03 100644 --- a/modelopt/__init__.py +++ b/modelopt/__init__.py @@ -15,28 +15,6 @@ """Nvidia TensorRT Model Optimizer (modelopt).""" -import sys as _sys -import warnings as _warnings from importlib.metadata import version as _version __version__ = _version("nvidia-modelopt") - - -try: - # Import from local source if available - from . import core - - __core_version__ = __version__ -except ImportError: - # Import from nvidia-modelopt-core wheel installation - import modelopt_core as core # type: ignore[no-redef] - - _sys.modules["modelopt.core"] = core - __core_version__ = _version("nvidia-modelopt-core") - -# Versions need to be the same for compatibility -if __version__.split(".")[:2] != __core_version__.split(".")[:2]: - _warnings.warn( - f"Version mismatch between nvidia-modelopt ({__version__}) and nvidia-modelopt-core" - f" ({__core_version__}). Please ensure both versions are the same for compatibility." - ) diff --git a/modelopt/onnx/autocast/precisionconverter.py b/modelopt/onnx/autocast/precisionconverter.py index dd294c6f9..37d20a47a 100644 --- a/modelopt/onnx/autocast/precisionconverter.py +++ b/modelopt/onnx/autocast/precisionconverter.py @@ -175,11 +175,15 @@ class PrecisionConverter: # Populate type information with inferred types self.model = self._propagate_types_shapes_custom_ops(self.model) else: - # Clear type information for intermediates and outputs + # Clear type/shape information for intermediates and outputs for vi in self.model.graph.value_info: vi.type.tensor_type.elem_type = onnx.TensorProto.UNDEFINED + for idx, d in enumerate(vi.type.tensor_type.shape.dim): + vi.type.tensor_type.shape.dim[idx].dim_param = "unk" for out in self.model.graph.output: out.type.tensor_type.elem_type = onnx.TensorProto.UNDEFINED + for idx, d in enumerate(out.type.tensor_type.shape.dim): + out.type.tensor_type.shape.dim[idx].dim_param = "unk" # Populate type information with inferred types self.model = onnx_utils.infer_shapes(self.model, strict_mode=True, check_type=False) # Sanity check: Verify type correctness @@ -230,9 +234,9 @@ class PrecisionConverter: ) def _propagate_cast_type_through_nodes(node, np_type, iter=1): - # Return if node is of cast or quantize type (from iter=2) + # Return if node is of cast type (from iter=2) indent = " " * iter - if iter > 1 and any(op in node.op.lower() for op in ["cast", "quantize"]): + if iter > 1 and any(op in node.op.lower() for op in ["cast"]): return out = node.outputs[0] diff --git a/modelopt/onnx/llm_export/utils/export_utils.py b/modelopt/onnx/llm_export/utils/export_utils.py new file mode 100644 index 000000000..b495eafdb --- /dev/null +++ b/modelopt/onnx/llm_export/utils/export_utils.py @@ -0,0 +1,162 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Utilities for exporting LLM models to ONNX.""" + +import json +import os +import time +from enum import Enum + +import torch +from transformers import DynamicCache + + +class RopeType(Enum): + """Rope type enum.""" + + K_NONE = 0 + K_ROPE_ROTATE_GPTJ = 1 + K_ROPE_ROTATE_NEOX = 2 + K_MROPE = 3 + + +class ModelLoader: + """A class to handle HuggingFace model loading and configuration.""" + + def __init__(self, torch_dir, config_path): + """Initialize the ModelLoader.""" + self.config_path = config_path + self.torch_dir = torch_dir + self.model_type = self.get_model_type() + self.hf_model = None + self.rope_type = RopeType.K_ROPE_ROTATE_NEOX + + def get_model_type(self): + """Get model type from config file.""" + with open(self.config_path) as f: + return json.load(f).get("model_type") + + def load_model(self): + """Load HuggingFace model based on model type.""" + print(f"Loading HF model from {self.torch_dir} with model type {self.model_type}") + from transformers import AutoModelForCausalLM + + self.hf_model = AutoModelForCausalLM.from_pretrained( + self.torch_dir, torch_dtype=torch.float16, trust_remote_code=True + ) + + return self.hf_model.eval().cuda() + + def get_rope_type(self): + """Get rope type.""" + return self.rope_type + + +class WrapperModelForCausalLM(torch.nn.Module): + """Wrapper Model to ensure all models have the same I/O.""" + + def __init__(self, model): + """Initialize the WrapperModelForCausalLM.""" + super().__init__() + try: + self.model = model.model + except Exception: + self.model = model + self.lm_head = model.lm_head + self.config = model.config + + def forward( + self, + input_ids, + past_key_values, + ): + """Forward pass.""" + past_key_values = DynamicCache.from_legacy_cache(past_key_values) + outputs = self.model(input_ids=input_ids, past_key_values=past_key_values, use_cache=True) + hidden_states = outputs[0] + past_key_values = outputs.past_key_values.to_legacy_cache() + logits = self.lm_head(hidden_states) + return logits, past_key_values + + +def llm_to_onnx(model, output_dir, extra_inputs={}, extra_dyn_axes={}): + """Export the WrapperModelForCausalLM to ONNX with fixed I/O names and shape definitions and save to `output_dir`. + + Parameters: + model: torch.Module + output_dir: str, the output_dir of the original ONNX. + extra_inputs: dict, append additional inputs after kv_cache. Usually for VL models + extra_dyn_axes: dict. Usually for VL models + """ + start_time = time.time() + config = model.config + num_layers = config.num_hidden_layers + num_attention_heads = config.num_attention_heads + num_key_value_heads = config.num_key_value_heads + hidden_size = config.hidden_size + hidden_size_per_layer = hidden_size // num_attention_heads + + dummy_bs = 1 + dummy_len = 10 + dummy_input_ids = torch.randint(100, (dummy_bs, dummy_len), dtype=torch.int64).cuda() + input_names = ["input_ids"] + output_names = ["logits"] + dynamic_axes = {"input_ids": {0: "batch_size", 1: "seq_len"}} + dummy_kv_cache = () + for i in range(num_layers): + dummy_k = torch.rand( + (dummy_bs, num_key_value_heads, dummy_len, hidden_size_per_layer), dtype=torch.float16 + ).cuda() + dummy_v = torch.rand( + (dummy_bs, num_key_value_heads, dummy_len, hidden_size_per_layer), dtype=torch.float16 + ).cuda() + dummy_kv_cache = (*dummy_kv_cache, (dummy_k, dummy_v)) + input_names.extend([f"past_key_values.{i}.key", f"past_key_values.{i}.value"]) + output_names.extend([f"present_key_values.{i}.key", f"present_key_values.{i}.value"]) + input_dynamic_axes = {0: "batch_size", 2: "past_len"} + dynamic_axes[f"past_key_values.{i}.key"] = input_dynamic_axes + dynamic_axes[f"past_key_values.{i}.value"] = input_dynamic_axes + + torch_to_onnx( + model, + (dummy_input_ids, {"past_key_values": dummy_kv_cache, **extra_inputs}), + output_dir, + "model.onnx", + input_names=input_names + list(extra_inputs.keys()), + output_names=output_names, + dynamic_axes=dynamic_axes | extra_dyn_axes, + ) + + end_time = time.time() + print( + f"Native ONNX Export from torch completed in {end_time - start_time}s. ONNX file is saved to {output_dir}." + ) + + +def torch_to_onnx(model, inputs, onnx_dir, onnx_name, input_names, output_names, dynamic_axes): + """Export the model to ONNX.""" + os.makedirs(onnx_dir, exist_ok=True) + with torch.inference_mode(): + torch.onnx.export( + model, + inputs, + f"{onnx_dir}/{onnx_name}", + input_names=input_names, + output_names=output_names, + dynamic_axes=dynamic_axes, + opset_version=19, + do_constant_folding=True, + ) diff --git a/modelopt/onnx/llm_export/utils/surgeon_utils.py b/modelopt/onnx/llm_export/utils/surgeon_utils.py new file mode 100644 index 000000000..664ed0f32 --- /dev/null +++ b/modelopt/onnx/llm_export/utils/surgeon_utils.py @@ -0,0 +1,120 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Utilities to surgeon ONNX graph after export.""" + +import re +import time + +import onnx +import onnx_graphsurgeon as gs +import torch +from onnx_graphsurgeon.ir.tensor import LazyValues + + +def clear_inputs(node: gs.Node | gs.Tensor): + """Clear all inputs for a node or tensor in ONNX.""" + for i in node.inputs: + i.outputs.clear() + node.inputs.clear() + return node + + +def clear_outputs(node: gs.Node | gs.Tensor): + """Clear all outputs for a node or tensor in ONNX.""" + for o in node.outputs: + o.inputs.clear() + node.outputs.clear() + return node + + +def extract_layer_id(name: str): + """Extract layer id from certain ONNX layer name. + + Parameters: + name: str + The name of ONNX layer. e.g. /model/layer.0/q_proj/... + + Returns: + The layer id for the layer as int. In the example above, it returns 0 + """ + match = re.search(r"layers\.(\d+)", name) + if match: + return int(match.group(1)) + raise Exception(f"{name} does not contain layer info!") + + +def no_none_elements(elements: list): + """Check if all elements in the list are not None.""" + return all(i is not None for i in elements) + + +def fold_fp8_qdq_to_dq(graph: gs.Graph): + """Convert FP32/FP16 weights of the given ONNX model to FP8 weights. + + Even though modelopt supports FP8 onnx export, the weights are represented in fp32 + QDQ. + The storage is therefore very bad. In this function, + Q nodes will get removed from the weights and have only DQ nodes with those converted FP8 + weights in the output model. + + Parameters: + graph: gs.Graph. + + Returns: + gs.Graph with only DQ nodes for weights and same QDQ nodes for activations. + """ + start_time = time.time() + print("Replacing all (fp32 weights + fp8 QDQ) with (fp8 weights + DQ)...") + # Fold constants is required since the scale is not constant yet. + graph.cleanup().toposort().fold_constants().cleanup() + + for node in graph.nodes: + if node.op == "TRT_FP8QuantizeLinear": + # Should not remove input QDQ + if not isinstance(node.inputs[0], gs.Constant): + continue + + weights = node.inputs[0] + scale = node.inputs[1] + torch_weights = torch.from_numpy(weights.values) + torch_scale = torch.from_numpy(scale.values) + quantizer_name = scale.name.rsplit("/", 1)[0] + dq_op = node.outputs[0].outputs[0] + assert dq_op.op == "TRT_FP8DequantizeLinear", ( + f"QDQ does not occur in pairs. You reached {dq_op.op}" + ) + + # Replace it with Dequantize with FP8 weights. This is a WAR because numpy does not support fp8. + numpy_weights = ( + (torch_weights / torch_scale).to(torch.float8_e4m3fn).view(torch.uint8).numpy() + ) + tensor = onnx.TensorProto() + tensor.data_type = onnx.TensorProto.FLOAT8E4M3FN + tensor.dims.extend(numpy_weights.shape) + tensor.raw_data = numpy_weights.tobytes() + values = LazyValues(tensor) + onnx_weights_fp8 = gs.Constant(quantizer_name + "/fp8_weights", values) + + node.outputs.clear() + # DQ Op is separated out + dq_op.inputs[0] = onnx_weights_fp8 + dq_op.op = "DequantizeLinear" + dq_op.outputs[0].dtype = dq_op.inputs[1].dtype + + graph.cleanup().toposort() + end_time = time.time() + print(f"fp8 qdq replaced with only dq completed in {end_time - start_time}s.") + + return graph diff --git a/modelopt/onnx/quantization/__main__.py b/modelopt/onnx/quantization/__main__.py index 2935d27ad..f6ef3f292 100644 --- a/modelopt/onnx/quantization/__main__.py +++ b/modelopt/onnx/quantization/__main__.py @@ -43,8 +43,8 @@ def get_parser() -> argparse.ArgumentParser: type=str, choices=["max", "entropy", "awq_clip", "rtn_dq"], help=( - "Calibration method choices for fp8: {max (default)}, " - "int8: {entropy (default), max}, int4: {awq_clip (default), rtn_dq}." + "Calibration method choices for int8/fp8: {entropy (default), max}, " + "int4: {awq_clip (default), rtn_dq}." ), ) group.add_argument( @@ -170,8 +170,11 @@ def get_parser() -> argparse.ArgumentParser: nargs="+", help=( "A space-separated list indicating the precision for each custom op. " - "Each item should have the format :, where precision can be fp32 (default) or fp16. " - "For example: op_type_1:fp16 op_type_2:fp32." + "Each item should have the format : (all inputs and outputs have the same precision) " + "or :[,,...]:[,,...] " + "(inputs and outputs can have different precisions), where precision can be fp32 (default), " + "fp16, int8, or fp8. Note that int8/fp8 should be the same as the quantization mode. " + "For example: op_type_1:fp16 op_type_2:[int8,fp32]:[int8]." ), ) argparser.add_argument( diff --git a/modelopt/onnx/quantization/calib_utils.py b/modelopt/onnx/quantization/calib_utils.py index 56e66b9d9..5ce9b1278 100644 --- a/modelopt/onnx/quantization/calib_utils.py +++ b/modelopt/onnx/quantization/calib_utils.py @@ -110,6 +110,10 @@ class CalibrationDataProvider(CalibrationDataReader): assert len(self.calibration_data_list) > 0, "Calibration data list is empty!" return self.calibration_data_list[0] + def rewind(self): + """Rewinds the data reader to the first index.""" + self.calibration_data_reader = iter(self.calibration_data_list) + class RandomDataProvider(CalibrationDataReader): """Calibration data reader class with random data provider.""" @@ -142,6 +146,10 @@ class RandomDataProvider(CalibrationDataReader): assert len(self.calibration_data_list) > 0, "Calibration data list is empty!" return self.calibration_data_list[0] + def rewind(self): + """Rewinds the data reader to the first index.""" + self.calibration_data_reader = iter(self.calibration_data_list) + def import_scales_from_calib_cache(cache_path: str) -> dict[str, float]: """Reads TensorRT calibration cache and returns as dictionary. diff --git a/modelopt/onnx/quantization/fp8.py b/modelopt/onnx/quantization/fp8.py index 9758574fb..c30914c5f 100755 --- a/modelopt/onnx/quantization/fp8.py +++ b/modelopt/onnx/quantization/fp8.py @@ -32,18 +32,16 @@ import modelopt.onnx.utils as onnx_utils from modelopt.onnx.autocast.convert import convert_to_f16 from modelopt.onnx.logging_config import configure_logging, logger from modelopt.onnx.quantization.graph_utils import ( - build_non_residual_input_map, convert_fp16_io, expand_node_names_from_patterns, find_nodes_to_exclude, get_concat_eliminated_tensors, - get_resize_scales, get_tensor_producer_nodes, insert_fp8_mha_casts, remove_output_initializers, remove_partial_input_qdq, - replace_resize_scales, ) +from modelopt.onnx.quantization.int8 import _find_nodes_to_quantize from modelopt.onnx.quantization.ort_patching import _quantize_static as quantize_static from modelopt.onnx.quantization.ort_utils import configure_ort from modelopt.onnx.quantization.qdq_utils import has_qdq_nodes @@ -99,19 +97,18 @@ def _find_unsupported_fp8_convs_to_exclude(graph: Graph): return unsupported_conv_nodes -def int8_to_fp8(onnx_path: str) -> onnx.ModelProto: +def int8_to_fp8(onnx_model: onnx.ModelProto) -> onnx.ModelProto: """Converts the INT8 quantized model to FP8 quantized model. Note. This conversion works only for max calibrated INT8 models. Args: - onnx_path: Path to the INT8 quantized ONNX model. + onnx_model: INT8 quantized ONNX model. Returns: FP8 quantized ONNX model. """ logger.info("Starting INT8 to FP8 conversion") - onnx_model = onnx.load(onnx_path, load_external_data=True) graph = onnx_model.graph initializers = graph.initializer tensor_producers = get_tensor_producer_nodes(graph) @@ -126,6 +123,11 @@ def int8_to_fp8(onnx_path: str) -> onnx.ModelProto: dtype = onnx.helper.tensor_dtype_to_np_dtype(scale.data_type) return numpy_helper.from_array(np_fp8_scale.astype(dtype), scale_name) + def _update_tensor_type(tensor_name): + tensor = onnx_utils.get_tensor_by_name(onnx_model, tensor_name) + if tensor: + tensor.type.tensor_type.elem_type = onnx.TensorProto.FLOAT8E4M3FN + def _convert(node: onnx.NodeProto): scale_name = node.input[1] zero_point_name = node.input[2] @@ -158,7 +160,15 @@ def int8_to_fp8(onnx_path: str) -> onnx.ModelProto: initializers[zero_point_idx].CopyFrom(np_zero_point) processed_tensor.add(zero_point_name) + # Update the Q input tensor type and the DQ output tensor type + if node.op_type == "QuantizeLinear": + for out in node.output: + _update_tensor_type(out) + if node.op_type == "DequantizeLinear": + _update_tensor_type(node.input[0]) + # Iterate through the nodes and convert the scales and zero points + # Also update the Q input tensor type and the DQ output tensor type for node in graph.node: if node.op_type in ["DequantizeLinear", "QuantizeLinear"]: _convert(node) @@ -201,7 +211,7 @@ def upgrade_opset_21(onnx_model: onnx.ModelProto) -> onnx.ModelProto: def quantize( onnx_path: str, - calibration_method: str = "max", + calibration_method: str = "entropy", calibration_data_reader: CalibrationDataReader = None, calibration_cache_path: str | None = None, calibration_shapes: str | None = None, @@ -218,6 +228,7 @@ def quantize( passes: list[str] = ["concat_elimination"], log_level: str = "INFO", calibrate_per_node: bool = False, + custom_ops_to_quantize: list[str] = [], **kwargs, ) -> onnx.ModelProto: """Applies FP8 GEMM only quantization to an ONNX file. @@ -228,9 +239,6 @@ def quantize( logger.info("Starting FP8 quantization process") t_start = time.time() - if calibration_method != "max": - raise RuntimeError("Only the max calibration method is supported for FP8 quantization.") - # Load the onnx graph logger.info(f"Loading ONNX model from {onnx_path}") onnx_model = onnx.load(onnx_path, load_external_data=True) @@ -243,22 +251,28 @@ def quantize( logger.info("Model already has QDQ nodes, skipping quantization") return onnx_model - # The quantizable op types for FP8 are limited to Conv, Gemm, and Matmul - fp8_supported_op_types = ["Gemm", "MatMul", "Conv"] + # The quantizable op types for FP8 are limited to Conv, Gemm, Matmul, and Residual-Add + fp8_supported_op_types = ["Gemm", "MatMul", "Conv", "Add"] op_types_to_quantize = op_types_to_quantize or fp8_supported_op_types if not set(op_types_to_quantize) <= set(fp8_supported_op_types): raise RuntimeError( f"Unsupported op types in fp8 mode: '{set(op_types_to_quantize) - set(fp8_supported_op_types)}'" ) + op_types_to_quantize.extend(list(custom_ops_to_quantize)) + + # Collect node names to exclude from quantization + nodes_to_exclude = find_nodes_to_exclude(graph, nodes_to_exclude, op_types_to_exclude) # type: ignore[arg-type] + nodes_to_exclude.extend(_find_unsupported_fp8_convs_to_exclude(graph)) # type: ignore[union-attr] # Change the default configuration of ORT quantization op_types = {node.op for node in graph.nodes} - trt_guided_options, _ = configure_ort( + trt_guided_options, quantizable_op_types = configure_ort( list(op_types), op_types_to_quantize, trt_extra_plugin_lib_paths, calibration_eps, calibrate_per_node, + custom_ops_to_quantize, ) logger.info( f"Quantizable op types in the model: {[t for t in op_types_to_quantize if t in op_types]}" @@ -268,16 +282,10 @@ def quantize( no_quantize_inputs = [] nodes_to_quantize = expand_node_names_from_patterns(graph, nodes_to_quantize) if not nodes_to_quantize: - nodes_to_quantize = [node.name for node in graph.nodes if node.op in op_types_to_quantize] - _, no_quantize_inputs = build_non_residual_input_map(graph) - if no_quantize_inputs: - op_types_to_quantize.append("Add") - add_nodes = [dst.name for _, dst, _ in no_quantize_inputs] - nodes_to_quantize.extend(add_nodes) - - # Collect node names to exclude from quantization - nodes_to_exclude = find_nodes_to_exclude(graph, nodes_to_exclude, op_types_to_exclude) # type: ignore[arg-type] - nodes_to_exclude.extend(_find_unsupported_fp8_convs_to_exclude(graph)) # type: ignore[union-attr] + quantizable_nodes, no_quantize_inputs = _find_nodes_to_quantize( + graph, quantizable_op_types, nodes_to_exclude + ) + nodes_to_quantize = [node.name for node in quantizable_nodes] # Update the list of nodes to quantize nodes_to_quantize = [ @@ -300,9 +308,9 @@ def quantize( os.close(tmp_onnx_file) logger.debug(f"Created temporary file for intermediate model: {tmp_onnx_path}") - # Quantize in INT8 mode using ORT's MinMax calibration method, with - # ActivationSymmetric as True, which is equivalent to max calibration - logger.info("Starting INT8 quantization with MinMax calibration") + # Quantize in INT8 mode using ORT's Entropy or MinMax calibration method, with + # ActivationSymmetric as True. For MinMax, that's equivalent to max calibration. + logger.info(f"Starting INT8 quantization with '{calibration_method}' calibration") quantize_static( onnx_path, tmp_onnx_path, @@ -312,14 +320,20 @@ def quantize( per_channel=True, extra_options=trt_guided_options, use_external_data_format=use_external_data_format, - calibrate_method=CalibrationMethod.MinMax, + calibrate_method=( + CalibrationMethod.Entropy + if calibration_method == "entropy" + # With ActivationSymmetric as True, MinMax calibration is equivalent to max calibration + else CalibrationMethod.MinMax + ), ) intermediate_generated_files.append(tmp_onnx_path) if use_external_data_format: intermediate_generated_files.append(tmp_onnx_path + ".data") # Post-processing of the onnx model after ORT quantization - onnx_model = int8_to_fp8(tmp_onnx_path) + logger.info("Starting post-processing of quantized model") + onnx_model = onnx.load(tmp_onnx_path) graph = gs.import_onnx(onnx_model) remove_partial_input_qdq(graph, no_quantize_inputs) onnx_model = gs.export_onnx(graph) @@ -332,8 +346,6 @@ def quantize( convert_fp16_io(graph) onnx_model = gs.export_onnx(graph) - # Record the old fp32 scale value of Resize node. - resize_scale_inits = get_resize_scales(onnx_model) # Convert to fp16/bf16 model. onnx_model = convert_to_f16( onnx_model, @@ -342,8 +354,6 @@ def quantize( low_precision_type=high_precision_dtype, trt_plugins=trt_extra_plugin_lib_paths, ) - # Replace the fp16/bf16 scale with old fp32 scale. - onnx_model = replace_resize_scales(onnx_model, resize_scale_inits) current_opsets = {opset.domain: opset.version for opset in onnx_model.opset_import} opset_of_default_onnx_domain = current_opsets.get("", 0) @@ -358,5 +368,7 @@ def quantize( logger.info("Inserting Cast nodes to enable FP8+FP16 MHA") onnx_model = insert_fp8_mha_casts(onnx_model) + onnx_model = int8_to_fp8(onnx_model) + logger.info(f"FP8 quantization completed in {time.time() - t_start:.2f} seconds") return onnx_model diff --git a/modelopt/onnx/quantization/graph_utils.py b/modelopt/onnx/quantization/graph_utils.py index 4304f7483..b579137d9 100755 --- a/modelopt/onnx/quantization/graph_utils.py +++ b/modelopt/onnx/quantization/graph_utils.py @@ -185,6 +185,7 @@ def get_fusible_backbone(node: Node, graph: Graph) -> Node | None: ["Relu", "BiasAdd", "ConstMul", conv_type], ["BatchNormalization", "BiasAdd", conv_type], ["Relu", "BatchNormalization", "BiasAdd", conv_type], + ["MaxPool", "Relu", "BatchNormalization", "BiasAdd", conv_type], ] for idx, path_type in enumerate(fusible_linear_path_types): if has_path_type(node, graph, path_type, is_forward=False, wild_card_types=[]): @@ -193,6 +194,27 @@ def get_fusible_backbone(node: Node, graph: Graph) -> Node | None: return None +def get_tensor_from_name(graph: onnx.GraphProto, tensor_name: str) -> onnx.ValueInfoProto | None: + """Returns a ValueInfoProto given a tensor name. + + Args: + graph: ONNX model graph + tensor_name: String with tensor name. + + Returns: + onnx.ValueInfoProto: actual graph tensor. + """ + # Search in inputs + vi = next((vi for vi in graph.input if vi.name == tensor_name), None) + # If not found, search in outputs + if vi is None: + vi = next((vi for vi in graph.output if vi.name == tensor_name), None) + # If not found, search in value_info (intermediate tensors) + if vi is None: + vi = next((vi for vi in graph.value_info if vi.name == tensor_name), None) + return vi + + def get_tensor_producer_nodes( graph: onnx.GraphProto, ) -> dict[str, onnx.NodeProto]: @@ -889,11 +911,9 @@ def find_nodes_from_mha_to_exclude( return [*set(nodes_to_exclude)] # type: ignore[arg-type] -def add_fp16_fp32_cast(onnx_path, custom_ops_to_cast_to_fp16, use_external_data_format): - """Adds cast_to_fp16 nodes to the inputs of a layer and cast_to_fp32 to the outputs.""" - logger.info( - "Adding cast_to_fp16 nodes to the inputs of a layer and cast_to_fp32 to the outputs" - ) +def cast_custom_ops(onnx_model: onnx.ModelProto, ops_to_cast: dict) -> onnx.ModelProto: + """Adds cast_to_fp16 nodes to the inputs and cast_to_fp32 to the outputs of a layer in the requested indices.""" + logger.info("Casting custom ops in the requested inputs and outputs") name_dict = {} def _get_unique_name(old_name): @@ -903,6 +923,10 @@ def add_fp16_fp32_cast(onnx_path, custom_ops_to_cast_to_fp16, use_external_data_ name_dict[old_name] = name_dict[old_name] + 1 return old_name + "_" + str(name_dict[old_name]) + def _is_castable_tensor(tensor) -> bool: + castable_types = ["float16", "float32", "double"] + return tensor.dtype and tensor.dtype in castable_types + def _add_cast_node_inp(tensor, precision="fp16", suffix=""): if precision == "fp16": onnx_precision = int(onnx.TensorProto.FLOAT16) @@ -949,25 +973,32 @@ def add_fp16_fp32_cast(onnx_path, custom_ops_to_cast_to_fp16, use_external_data_ graph.nodes.append(cast_node) return cast_inp - graph = gs.import_onnx(onnx.load(onnx_path)) - castable_nodes = [n for n in graph.nodes if n.op in custom_ops_to_cast_to_fp16] + graph = gs.import_onnx(onnx_model) + castable_nodes = [n for n in graph.nodes if n.op in ops_to_cast] for node in castable_nodes: - # Cast all inputs to FP16 - for inp_idx, inp in enumerate(node.inputs): - cast_out = _add_cast_node_inp(inp) - node.inputs[inp_idx] = cast_out + inp_idxs = ops_to_cast[node.op]["inp"] + out_idxs = ops_to_cast[node.op]["out"] - # Cast all outputs from FP16 back to FP32 + # Cast relevant inputs to FP16 + for inp_idx, inp in enumerate(node.inputs): + if inp_idx in inp_idxs and _is_castable_tensor(inp): + cast_out = _add_cast_node_inp(inp) + node.inputs[inp_idx] = cast_out + + # Cast relevant outputs from FP16 back to FP32 for out_idx, out in enumerate(node.outputs): - cast_inp = _add_cast_node_out(out) - node.outputs[out_idx] = cast_inp + if out_idx in out_idxs and _is_castable_tensor(out): + cast_inp = _add_cast_node_out(out) + node.outputs[out_idx] = cast_inp graph.cleanup().toposort() - new_onnx_path = onnx_path.replace(".onnx", "_castFP16.onnx") - save_onnx(gs.export_onnx(graph), new_onnx_path, use_external_data_format) - return new_onnx_path + onnx_model = gs.export_onnx(graph) + # TODO: remove manual ir_version change once ORT supports ir_version 11 + onnx_model.ir_version = 10 + + return onnx_model def print_stat(graph: Graph) -> None: @@ -1232,24 +1263,6 @@ def get_resize_scales(onnx_model): return resize_scale_inits -def replace_resize_scales(onnx_model, resize_scale_inits): - """Replace Resize op's fp16 scale value with old fp32 scale.""" - if len(resize_scale_inits) == 0: - return onnx_model - - graph = gs.import_onnx(onnx_model) - for node in graph.nodes: - if node.op == "Resize" and node.name in resize_scale_inits: - cast_node = node.inputs[2].inputs[0] - scale = cast_node.inputs[0] - for new_init in onnx_model.graph.initializer: - if new_init.name == scale.name: - old_data_type, old_raw_data = resize_scale_inits[node.name] - new_init.data_type = old_data_type - new_init.raw_data = old_raw_data - return onnx_model - - def get_concat_eliminated_tensors( onnx_model: onnx.ModelProto, nodes_to_quantize: list[str], diff --git a/modelopt/onnx/quantization/int4.py b/modelopt/onnx/quantization/int4.py index a7cdbc002..800e63cbb 100644 --- a/modelopt/onnx/quantization/int4.py +++ b/modelopt/onnx/quantization/int4.py @@ -161,9 +161,9 @@ def dq_tensor(w: np.ndarray, s: np.ndarray, block_size: int, zp: np.ndarray = No def quantize_rtn( onnx_model: onnx.ModelProto, - gemm_io_type: onnx.TensorProto.DataType, block_size: int, dq_only: bool = False, + nodes_to_exclude: list[str] = [], ) -> onnx.ModelProto: """Quantizes `onnx_model` using the RTN (Round-to-Nearest) algorithm. @@ -178,11 +178,15 @@ def quantize_rtn( t_start = time.time() graph = gs.import_onnx(onnx_model) - gemm_nodes = [node for node in graph.nodes if node.op in ["Gemm", "MatMul"]] + nodes_to_exclude = expand_node_names_from_patterns(graph, nodes_to_exclude) + gemm_nodes = [ + node + for node in graph.nodes + if node.op in ["Gemm", "MatMul"] and node.name not in nodes_to_exclude + ] logger.info(f"Found {len(gemm_nodes)} Gemm/MatMul nodes to quantize") gemm_tensors = {} - act_tensors = [] for gemm in gemm_nodes: for in_tensor in gemm.inputs: if not isinstance(in_tensor, gs.Constant): @@ -191,27 +195,24 @@ def quantize_rtn( # 1D blocked quantization not supported. continue gemm_tensors[in_tensor.name] = in_tensor - act_tensors.append(gemm.inputs[0]) gemm_weights = {name: tensor.values for name, tensor in gemm_tensors.items()} logger.info(f"Found {len(gemm_weights)} quantizable weights") logger.info("Computing scales for gemm weights") scales = {} + gemm_io_type = {} for name, w in gemm_weights.items(): logger.debug(f"Computing scales for weight {name} of shape {w.shape}") s, zp = find_scales(np.asarray(w), block_size) assert zp is None, "zero-point is not enabled but zp is found non-None" scales[name] = s + gemm_io_type[name] = onnx.helper.np_dtype_to_tensor_dtype(cast("int", w.dtype)) # Change the scale type to the expected type, fp16 by default for name in scales: s = scales[name] - scales[name] = s.astype(onnx.helper.tensor_dtype_to_np_dtype(gemm_io_type)) - - # Change the input activation type to the expected type, fp16 by default - for act_tensor in act_tensors: - _change_input_type(onnx_model.graph, act_tensor.name, gemm_io_type) + scales[name] = s.astype(onnx.helper.tensor_dtype_to_np_dtype(gemm_io_type[name])) # Import the update graph graph = gs.import_onnx(onnx_model) @@ -350,19 +351,13 @@ def _find_quantizable_weights( ) -> list[tuple[onnx.ValueInfoProto, onnx.ValueInfoProto, bool, int]]: """Finds the quantizable weights from the graph.""" wa_pack = [] - gemm_nodes = [node for node in graph.node if node.op_type in ["Gemm", "MatMul"]] + gemm_nodes = [ + node + for node in graph.node + if node.op_type in ["Gemm", "MatMul"] and node.name not in nodes_to_exclude + ] initializer_idxs = {initializer.name: idx for idx, initializer in enumerate(graph.initializer)} for gemm in gemm_nodes: - exclude_this_node = False - - for i in range(len(nodes_to_exclude)): - if nodes_to_exclude[i] in gemm.name: - exclude_this_node = True - break - - if exclude_this_node: - continue - if gemm.input[0] in initializer_idxs: # Ex. two const input to MatMul_115 in fastvit0.onnx # Note. RTN algorithm will quantize these weights though @@ -884,6 +879,7 @@ def run_awq_scale_search_per_subgraph( def get_parent_child_nodes_map( graph: onnx.GraphProto, wa_pack: list[tuple[gs.Tensor, gs.Tensor, bool, int]], + nodes_to_exclude: list[str], ): """Get mapping of parent nodes to their MatMul/Gemm nodes with quantizable weights.""" parent_child_nodes_map = {} @@ -894,7 +890,7 @@ def get_parent_child_nodes_map( parent_name = output_name_to_node[act_tensor.name].name parent_child_nodes_map[parent_name] = [] for node in input_name_to_nodes[act_tensor.name]: - if node.op_type in ["Gemm", "MatMul"]: + if node.op_type in ["Gemm", "MatMul"] and node.name not in nodes_to_exclude: parent_child_nodes_map[parent_name].append(node) return parent_child_nodes_map, input_name_to_nodes @@ -934,7 +930,7 @@ def _quantize_awq_lite( wa_pack = _find_quantizable_weights(graph, nodes_to_exclude) if fuse_nodes: parent_child_nodes_map, input_name_to_nodes = get_parent_child_nodes_map( - onnx_model.graph, wa_pack + onnx_model.graph, wa_pack, nodes_to_exclude ) # Add input activations to graph output @@ -1141,11 +1137,13 @@ def _quantize_awq_lite( for inp in parent.input: if initializer_map.get(inp) is not None: tensor = initializer_map[inp] + old_dim = tensor.dims tensor_array = numpy_helper.to_array( tensor, base_dir=os.path.dirname(augmented_onnx_path), ) new_tensor = np.asarray(tensor_array) / input_scale + new_tensor = new_tensor.reshape(old_dim) new_tensor = numpy_helper.from_array(new_tensor.get(), tensor.name) # replace initializer with new scaled array tensor.CopyFrom(new_tensor) @@ -1277,8 +1275,6 @@ def quantize( block_size = 128 logger.info(f"Using default block size: {block_size}") - gemm_io_type: onnx.TensorProto.DataType = onnx.TensorProto.FLOAT - # set config params nodes_to_exclude = nodes_to_exclude or [] logger.debug(f"Excluding nodes matching patterns: {nodes_to_exclude}") @@ -1301,7 +1297,10 @@ def quantize( if calibration_method in ["rtn", "rtn_dq", "rtn_trt", "rtn_trt_dq"]: onnx_model = quantize_rtn( - onnx_model, gemm_io_type, block_size, dq_only="dq" in calibration_method + onnx_model, + block_size, + dq_only="dq" in calibration_method, + nodes_to_exclude=nodes_to_exclude, ) elif calibration_method in ["awq_lite", "awq_full"]: do_weight_clipping = False diff --git a/modelopt/onnx/quantization/int8.py b/modelopt/onnx/quantization/int8.py index 6b6d544d1..a13a3f7a3 100755 --- a/modelopt/onnx/quantization/int8.py +++ b/modelopt/onnx/quantization/int8.py @@ -90,16 +90,16 @@ def _find_nodes_to_quantize( ) quantizable_nodes = quantizable_kgen_heads + quantizable_partition_nodes - paritially_quantizable_nodes = [dst for _, dst, _ in no_quantize_inputs] + partially_quantizable_nodes = [dst for _, dst, _ in no_quantize_inputs] # Quantize all inputs of partially quantizable nodes by ORT # but remove QDQ from non-quantizable inputs in the post-processing step - quantizable_nodes.extend(paritially_quantizable_nodes) + quantizable_nodes.extend(partially_quantizable_nodes) quantizable_nodes.extend( find_quantizable_nodes(graph, quantizable_nodes, partitioned_nodes, quantizable_op_types) ) - skip_list = get_skipped_output_layers(graph, paritially_quantizable_nodes) + skip_list = get_skipped_output_layers(graph, partially_quantizable_nodes) quantizable_nodes = [node for node in quantizable_nodes if node.name not in skip_list] logger.info(f"Total number of quantizable nodes: {len(quantizable_nodes)}") @@ -124,6 +124,7 @@ def quantize( passes: list[str] = ["concat_elimination"], log_level: str = "INFO", calibrate_per_node: bool = False, + custom_ops_to_quantize: list[str] = [], **kwargs, ) -> onnx.ModelProto: """Applies INT8 quantization to an ONNX file using the compiler friendly heuristics. @@ -169,6 +170,8 @@ def quantize( # Change the default configuration of ORT quantization op_types_to_quantize = op_types_to_quantize or [] + if op_types_to_quantize: + op_types_to_quantize.extend(custom_ops_to_quantize) op_types = {node.op for node in graph.nodes} trt_guided_options, quantizable_op_types = configure_ort( list(op_types), @@ -176,6 +179,7 @@ def quantize( trt_extra_plugin_lib_paths, calibration_eps, calibrate_per_node, + custom_ops_to_quantize, ) logger.info(f"Quantizable op types: {[t for t in quantizable_op_types if t in op_types]}") diff --git a/modelopt/onnx/quantization/operators.py b/modelopt/onnx/quantization/operators.py index 783737f0e..b3c5ebc14 100644 --- a/modelopt/onnx/quantization/operators.py +++ b/modelopt/onnx/quantization/operators.py @@ -94,3 +94,28 @@ class QDQConvTranspose(QDQOperatorBase): self.quantizer.quantize_bias_tensor( node.name, node.input[2], node.input[0], node.input[1] ) + + +class QDQCustomOp(QDQOperatorBase): + """By default, ORT does not quantize custom ops. This module is intended to help with that. + + Note. QDQOperatorBase is not sufficient for dynamic input and output only quantization. + """ + + def __init__(self, onnx_quantizer, onnx_node): + """Normalization quantizer init.""" + super().__init__(onnx_quantizer, onnx_node) + + def quantize(self): + """Main function to quantize the custom ops.""" + node = self.node + + # Quantize only the dynamic inputs with type FLOAT or FLOAT16 + for inp in node.input: + if self.quantizer._is_tensor_quantizable(inp): + self.quantizer.quantize_activation_tensor(inp) + + # Quantize the outputs with type FLOAT or FLOAT16 + for out in node.output: + if self.quantizer._is_tensor_quantizable(out): + self.quantizer.quantize_activation_tensor(out) diff --git a/modelopt/onnx/quantization/ort_patching.py b/modelopt/onnx/quantization/ort_patching.py index f10d90966..2b89d1fb4 100755 --- a/modelopt/onnx/quantization/ort_patching.py +++ b/modelopt/onnx/quantization/ort_patching.py @@ -276,6 +276,13 @@ def _select_tensors_to_calibrate(calibrator, model: onnx.ModelProto): and (tensor_name not in initializer) ): tensors_to_calibrate.add(tensor_name) + for tensor_name in node.output: + if tensor_name in value_infos: + vi = value_infos[tensor_name] + if vi.type.HasField("tensor_type") and ( + vi.type.tensor_type.elem_type in tensor_type_to_calibrate + ): + tensors_to_calibrate.add(tensor_name) return tensors_to_calibrate, value_infos diff --git a/modelopt/onnx/quantization/ort_utils.py b/modelopt/onnx/quantization/ort_utils.py index 5da8a220b..b40fa9ce8 100755 --- a/modelopt/onnx/quantization/ort_utils.py +++ b/modelopt/onnx/quantization/ort_utils.py @@ -25,7 +25,7 @@ from onnxruntime.quantization.registry import QDQRegistry, QLinearOpsRegistry from packaging.version import Version from modelopt.onnx.logging_config import logger -from modelopt.onnx.quantization.operators import QDQConvTranspose, QDQNormalization +from modelopt.onnx.quantization.operators import QDQConvTranspose, QDQCustomOp, QDQNormalization from modelopt.onnx.quantization.ort_patching import patch_ort_modules @@ -248,6 +248,7 @@ def configure_ort( trt_extra_plugin_lib_paths: list[str] | None = None, calibration_eps: list[str] | None = None, calibrate_per_node: bool = False, + custom_ops_to_quantize: list[str] = [], ): """Configure and patches ORT to support ModelOpt ONNX quantization.""" logger.info("Configuring ORT for ModelOpt ONNX quantization") @@ -260,6 +261,8 @@ def configure_ort( QDQRegistry["HardSwish"] = ( QDQOperatorBase # Example: mobilenet_v3_opset17, efficientvit_b3_opset17 ) + for custom_op in custom_ops_to_quantize: + QDQRegistry[custom_op] = QDQCustomOp # Patch ORT modules to fix bugs and support some edge cases patch_ort_modules(calibrate_per_node) @@ -293,17 +296,20 @@ def configure_ort( del QDQRegistry[op_type] # Prepare TensorRT friendly quantization settings + no_output_quantization_op_types = [ + op_type for op_type in op_types if op_type not in custom_ops_to_quantize + ] if trt_extra_plugin_lib_paths is not None: trt_extra_plugin_lib_paths = ";".join(trt_extra_plugin_lib_paths) trt_guided_options = { "QuantizeBias": False, "ActivationSymmetric": True, - "OpTypesToExcludeOutputQuantization": op_types, # No output quantization + "OpTypesToExcludeOutputQuantization": no_output_quantization_op_types, # No output quantization "AddQDQPairToWeight": True, # Instead of quantizing the weights, add QDQ node "QDQOpTypePerChannelSupportToAxis": { "Conv": 0, # Cout axis for Conv: [Cout, Cin, k1, k2] "ConvTranspose": 1, # Cout axis for ConvTranspose: [Cin, Cout, k1, k2] - }, # per_channel should be True (modeopt default) + }, # per_channel should be True (modelopt default) "DedicatedQDQPair": False, "ForceQuantizeNoInputCheck": ( # By default, for some latent operators like MaxPool, Transpose, etc., diff --git a/modelopt/onnx/quantization/partitioning.py b/modelopt/onnx/quantization/partitioning.py index 6dfcb7dd8..5e86ef574 100644 --- a/modelopt/onnx/quantization/partitioning.py +++ b/modelopt/onnx/quantization/partitioning.py @@ -31,7 +31,7 @@ from modelopt.onnx.quantization.graph_utils import ( has_path_type, is_const_input, ) -from modelopt.onnx.utils import get_child_nodes, get_variable_inputs +from modelopt.onnx.utils import get_child_nodes, get_parent_nodes, get_variable_inputs def _build_fusible_partition( @@ -282,6 +282,16 @@ def find_quantizable_nodes( return False + def _has_quantizable_producer(node: Node, quantizable_node_set: set[str]) -> bool: + parents = get_parent_nodes(node) + for parent_node in parents: + if (parent_node.name in quantizable_node_set) or ( + (is_copy_op(parent_node.op) or parent_node.op == "Cast") + and _has_quantizable_producer(parent_node, quantizable_node_set) + ): + return True + return False + quantizable_nodes = [] pooling_and_window_ops = [] for node in graph.nodes: @@ -313,6 +323,7 @@ def find_quantizable_nodes( node for node in pooling_and_window_ops if _has_quantizable_consumer(node, quantizable_node_set) + or _has_quantizable_producer(node, quantizable_node_set) ) return quantizable_nodes diff --git a/modelopt/onnx/quantization/qdq_utils.py b/modelopt/onnx/quantization/qdq_utils.py index 0534ca3ab..fd1b449cb 100644 --- a/modelopt/onnx/quantization/qdq_utils.py +++ b/modelopt/onnx/quantization/qdq_utils.py @@ -29,6 +29,7 @@ from modelopt.onnx import utils from modelopt.onnx.logging_config import logger from modelopt.onnx.quantization.graph_utils import ( get_tensor_consumer_nodes, + get_tensor_from_name, get_tensor_producer_nodes, remove_redundant_cast_nodes, ) @@ -495,7 +496,7 @@ def _get_scale_and_zp( return scale_array, zp_array -def _get_succesive_consumers( +def _get_successive_consumers( node: onnx.NodeProto, tensor_consumers: dict[str, list[onnx.NodeProto]] ) -> tuple[onnx.NodeProto, onnx.NodeProto]: """Get the DequantizeLinear node and its consumer node for a given QuantizeLinear node. @@ -654,7 +655,7 @@ def qdq_to_dq(onnx_model: onnx.ModelProto) -> onnx.ModelProto: scale_array, zp_array = _get_scale_and_zp(node, initializers, tensor_producers) # Validate Q->DQ->Op pattern and get consumers - dq_node, quantized_node = _get_succesive_consumers(node, tensor_consumers) + dq_node, quantized_node = _get_successive_consumers(node, tensor_consumers) # Convert weight scaled = _convert_weight(weight_array, scale_array, zp_array, quantized_node) @@ -696,6 +697,137 @@ def qdq_to_dq(onnx_model: onnx.ModelProto) -> onnx.ModelProto: return onnx_model +def remove_input_dq_and_output_q( + onnx_model: onnx.ModelProto, quantizable_custom_ops: dict +) -> onnx.ModelProto: + """Remove DQ nodes from the input and Q from the output of quantized custom ops for TensorRT compatibility. + + TensorRT requires only Q nodes in the inputs and only DQ nodes in the outputs of custom ops. + For more information, see https://docs.nvidia.com/deeplearning/tensorrt/latest/inference-library/work-quantized-types.html#q-dq-interaction-with-plugins + + Args: + onnx_model: ONNX model protobuf to convert + quantizable_custom_ops: dictionary of custom ops and I/O indices to perform Q and DQ deletions as needed. + + Returns: + ONNX model protobuf with only Q in the inputs and only DQ in the outputs of custom ops. + + Raises: + ValueError: If the model is invalid or removal fails + RuntimeError: If graph operations fail + """ + logger.info("Deleting DQ nodes in the input and Q nodes in the output of custom ops.") + if not isinstance(onnx_model, onnx.ModelProto): + raise ValueError("Input must be an ONNX model protobuf") + + graph = onnx_model.graph + if not graph.node: + raise ValueError("Model graph is empty") + + initializers, tensor_producers, tensor_consumers = _get_graph_metadata(graph) + q_nodes = [ + (idx, node) for idx, node in enumerate(graph.node) if node.op_type == "QuantizeLinear" + ] + dq_nodes = [ + (idx, node) for idx, node in enumerate(graph.node) if node.op_type == "DequantizeLinear" + ] + q_indices = [] + dq_indices = [] + + # Remove DQ nodes in the input of custom ops + for node_idx, node in dq_nodes: + consumers = tensor_consumers[node.output[0]] + for inp_name in node.input: + logger.debug(f"Processing QDQ node for input {inp_name}") + + # Ignore initializers (scale, zero_point) + if inp_name in initializers: + continue + + try: + # Update the previous Q node output name, each DQ should only have one Q producer + q_node = tensor_producers[inp_name] + assert isinstance(q_node, onnx.NodeProto), ( + f"Expected producer {node.name} to be of type NodeProto" + ) + assert q_node.op_type == "QuantizeLinear", ( + f"Expected QuantizeLinear producer for {node.name}" + ) + + # Only remove DQs from the inputs of custom ops + if consumers[0].op_type not in quantizable_custom_ops: + continue + + # Rewire graph to connect Q with the node after DQ (skip DQ) + for consumer in consumers: + for cons_idx, cons_inp in enumerate(consumer.input): + if cons_inp == node.output[0]: + # If the input tensor is meant to be quantized, delete DQ. Otherwise, delete both Q/DQ. + if cons_idx in quantizable_custom_ops[consumer.op_type]["inp"]: + consumer.input[cons_idx] = q_node.output[0] + else: + q_node_prev = tensor_producers[q_node.input[0]] + consumer.input[cons_idx] = q_node_prev.output[0] + break + + # Track DequantizeLinear node indices for cleanup + dq_indices.append(node_idx) + + except Exception as e: + raise RuntimeError(f"Failed to convert node {node.name}: {e!s}") + + # Remove Q nodes in the output of custom ops + for node_idx, node in q_nodes: + for out_name in node.output: + logger.debug(f"Processing QDQ node for output {out_name}") + + try: + # Update the Q node output name, each Q should only have one DQ consumer + dq_node = tensor_consumers[out_name] + assert len(dq_node) == 1, f"Expected single consumer for {node.name}" + assert dq_node[0].op_type == "DequantizeLinear", ( + f"Expected DequantizeLinear producer for {node.name}" + ) + + # Only remove Qs from the output of custom ops + if ( + node.input[0] in initializers + or get_tensor_from_name(graph, node.input[0]) in graph.input + ): + continue + producer = tensor_producers[node.input[0]] + if producer.op_type not in quantizable_custom_ops: + continue + + # Rewire graph to connect the output of custom op to the input of DQ (skip Q) + # If the output tensor is meant to be quantized, delete Q. Otherwise, delete both Q/DQ. + if quantizable_custom_ops[producer.op_type]["out"]: + dq_node[0].input[0] = producer.output[0] + else: + dq_node_next = tensor_consumers[dq_node[0].output[0]] + dq_node_next[0].input[0] = producer.output[0] + + # Track QuantizeLinear node indices for cleanup + q_indices.append(node_idx) + + except Exception as e: + raise RuntimeError(f"Failed to convert node {node.name}: {e!s}") + + # Remove processed nodes + for node_idx in sorted(q_indices + dq_indices, reverse=True): + del graph.node[node_idx] + + logger.info( + f"Removed {len(q_indices)} Q node{'' if len(q_indices) == 1 else 's'} and" + f" {len(dq_indices)} DQ node{'' if len(dq_indices) == 1 else 's'}" + ) + + # TODO: remove manual ir_version change once ORT supports ir_version 11 + onnx_model.ir_version = 10 + + return onnx_model + + def quantize_weights_to_mxfp8( onnx_model: onnx.ModelProto, ) -> onnx.ModelProto: diff --git a/modelopt/onnx/quantization/quantize.py b/modelopt/onnx/quantization/quantize.py index f0071d1f0..67f5cfb58 100755 --- a/modelopt/onnx/quantization/quantize.py +++ b/modelopt/onnx/quantization/quantize.py @@ -50,7 +50,7 @@ from modelopt.onnx.quantization.calib_utils import ( ) from modelopt.onnx.quantization.fp8 import quantize as quantize_fp8 from modelopt.onnx.quantization.graph_utils import ( - add_fp16_fp32_cast, + cast_custom_ops, find_nodes_from_mha_to_exclude, print_stat, remove_redundant_cast_nodes, @@ -58,8 +58,8 @@ from modelopt.onnx.quantization.graph_utils import ( from modelopt.onnx.quantization.int4 import quantize as quantize_int4 from modelopt.onnx.quantization.int8 import quantize as quantize_int8 from modelopt.onnx.quantization.ort_utils import update_trt_ep_support -from modelopt.onnx.quantization.qdq_utils import qdq_to_dq -from modelopt.onnx.trt_utils import load_onnx_model +from modelopt.onnx.quantization.qdq_utils import qdq_to_dq, remove_input_dq_and_output_q +from modelopt.onnx.trt_utils import interpret_trt_plugins_precision_flag, load_onnx_model from modelopt.onnx.utils import duplicate_shared_constants, name_onnx_nodes, save_onnx __all__ = ["quantize"] @@ -74,7 +74,8 @@ def _preprocess_onnx( trt_plugins_precision: list[str] | None, override_shapes: str, simplify: bool = False, -) -> tuple[str, list[str], bool, bool, bool]: + quantize_mode: str = "int8", +) -> tuple[str, onnx.ModelProto, list[str], bool, bool, bool, dict]: logger.info(f"Preprocessing the model {onnx_path}") intermediate_generated_files = [] output_dir = os.path.dirname(output_path) @@ -172,29 +173,27 @@ def _preprocess_onnx( logger.info(f"Model is cloned to {onnx_path} after naming the nodes") intermediate_generated_files.append(onnx_path) - # If custom op precisions are given, check if they're fp16. If so, add cast_to_fp16 before all inputs and - # cast_to_fp32 after all outputs. + # If custom op precisions are given, add Cast or Q/DQ where appropriate. + custom_ops_to_quantize = {} if trt_plugins_precision: - logger.debug("Processing custom op precisions") - custom_ops_to_cast = [] - for trt_plugin_precision in trt_plugins_precision: - assert ":" in trt_plugin_precision, ( - "Plugin pre cision is incorrectly formatted." - " Please check that it's in the format :." - ) - op_type, precision = trt_plugin_precision.split(":") - if precision == "fp16": - custom_ops_to_cast.append(op_type) + custom_ops_to_cast, custom_ops_to_quantize = interpret_trt_plugins_precision_flag( + onnx_model, trt_plugins_precision, quantize_mode + ) if custom_ops_to_cast: - onnx_path = add_fp16_fp32_cast(onnx_path, custom_ops_to_cast, use_external_data_format) + onnx_model = cast_custom_ops(onnx_model, custom_ops_to_cast) + onnx_path = os.path.join(output_dir, f"{model_name}_castFP16.onnx") + save_onnx(onnx_model, onnx_path, use_external_data_format) + logger.info(f"Model is cloned to {onnx_path} after casting tensors to FP16") intermediate_generated_files.append(onnx_path) return ( onnx_path, + onnx_model, intermediate_generated_files, has_custom_op, has_dds_op, use_external_data_format, + custom_ops_to_quantize, ) @@ -239,8 +238,8 @@ def quantize( calibration_data: Calibration data, either a numpy array or list/dict of numpy arrays. calibration_method: - Calibration method choices. Options are int8: 'entropy' (default) and 'max', - fp8: 'max' (default) and int4: 'awq_clip' (default), 'awq_lite', 'awq_full' and 'rtn_dq'. + Calibration method choices. Options are int8/fp8: {'entropy' (default), 'max'} + and int4: {'awq_clip' (default), 'awq_lite', 'awq_full', 'rtn_dq'}. calibration_cache_path: Path to pre-calculated activation tensor ranges, also known as calibration cache. calibration_shapes: @@ -355,17 +354,24 @@ def quantize( # We need to preprocess the model with naming, weight duplication etc. enable_shared_constants_duplication = kwargs.get("enable_shared_constants_duplication", True) - onnx_path, intermediate_generated_files, has_custom_op, has_dds_op, use_external_data_format = ( - _preprocess_onnx( - onnx_path, - use_external_data_format, - output_path, - enable_shared_constants_duplication, - trt_plugins, - trt_plugins_precision, - override_shapes, # type: ignore[arg-type] - simplify, - ) + ( + onnx_path, + onnx_model, + intermediate_generated_files, + has_custom_op, + has_dds_op, + use_external_data_format, + custom_ops_to_quantize, + ) = _preprocess_onnx( + onnx_path, + use_external_data_format, + output_path, + enable_shared_constants_duplication, + trt_plugins, + trt_plugins_precision, + override_shapes, # type: ignore[arg-type] + simplify, + quantize_mode, ) trt_plugins = update_trt_ep_support(calibration_eps, has_dds_op, has_custom_op, trt_plugins) # type: ignore[arg-type] @@ -398,10 +404,9 @@ def quantize( if quantize_mode in ["fp8", "int8"]: quantize_func = quantize_int8 if quantize_mode == "int8" else quantize_fp8 - default_calibration_method = "entropy" if quantize_mode == "int8" else "max" onnx_model = quantize_func( onnx_path=onnx_path, - calibration_method=calibration_method or default_calibration_method, + calibration_method=calibration_method or "entropy", calibration_data_reader=calibration_data_reader, calibration_cache_path=calibration_cache_path, calibration_shapes=calibration_shapes, @@ -418,6 +423,7 @@ def quantize( passes=passes, log_level=log_level, calibrate_per_node=calibrate_per_node, + custom_ops_to_quantize=list(custom_ops_to_quantize.keys()), **kwargs, ) elif "int4" in quantize_mode: @@ -438,8 +444,18 @@ def quantize( if onnx_model: # Fuse Q nodes for INT8/FP8 mode - if quantize_mode in ["int8", "fp8"] and dq_only: - onnx_model = qdq_to_dq(onnx_model) + if quantize_mode in ["int8", "fp8"]: + if dq_only: + onnx_model = qdq_to_dq(onnx_model) + if custom_ops_to_quantize: + # Remove DQ nodes from the input and Q from the output of the requested custom ops + onnx_model = remove_input_dq_and_output_q( + onnx_model, quantizable_custom_ops=custom_ops_to_quantize + ) + # Sort nodes topologically + graph = gs.import_onnx(onnx_model) + graph.toposort().cleanup() + onnx_model = gs.export_onnx(graph) else: # Remove redundant cast nodes in the quantized model # Note. This is called within the qdq_to_dq function as well diff --git a/modelopt/onnx/trt_utils.py b/modelopt/onnx/trt_utils.py index 9c9cfa1db..c6fdf0d13 100644 --- a/modelopt/onnx/trt_utils.py +++ b/modelopt/onnx/trt_utils.py @@ -276,3 +276,103 @@ def load_onnx_model( static_shaped_onnx_path or onnx_path, use_external_data_format, ) + + +def interpret_trt_plugins_precision_flag( + onnx_model: onnx.ModelProto, + trt_plugins_precision: list[str], + quantize_mode: str, +) -> tuple[dict, dict]: + """Convert custom ops precision flag to dictionaries with custom op and I/O indices to be cast/quantized. + + Args: + onnx_model: ONNX model to detect with nodes need to be cast/quantized. + trt_plugins_precision: List indicating the precision for each custom op. + quantize_mode: String indicating the quantization mode. + + Returns: + Dictionary with custom ops to cast containing the I/O indices to cast. + Dictionary with custom ops to quantize containing the I/O indices to quantize. + """ + # If custom op precisions are given, check if they're supported (fp32, fp16, int8, fp8) + custom_ops_to_cast = {} + custom_ops_to_quantize = {} + supported_precisions = ["fp32", "fp16", "int8", "fp8"] + logger.debug("Processing custom op precisions") + + graph = gs.import_onnx(onnx_model) + + for trt_plugin_precision in trt_plugins_precision: + assert trt_plugin_precision.count(":") in [1, 2], ( + "Plugin precision is incorrectly formatted." + " Please check that it's in the format : or" + " :[,,...]:[,,...]." + ) + # Split only on the first ":" to get 'op_type' + op_type, precision = trt_plugin_precision.split(":", 1) + custom_op_nodes = [node for node in graph.nodes if node.op == op_type] + if not custom_op_nodes: + logger.warning(f"No nodes of type {op_type} were found. Skipping.") + continue + num_inps = max([len(node.inputs) for node in custom_op_nodes]) + num_outs = max([len(node.outputs) for node in custom_op_nodes]) + + # Now split the remainder of the string to get the I/O precisions + if trt_plugin_precision.count(":") == 1: + if precision not in supported_precisions: + logger.warning(f"Precision {precision} is not supported. Skipping.") + if precision == "fp16": + custom_ops_to_cast[op_type] = { + "inp": list(range(num_inps)), + "out": list(range(num_outs)), + } + if precision in ["int8", "fp8"]: + if precision != quantize_mode: + precision = quantize_mode + logger.warning( + f"Requested custom op precision ({precision}) is different than quantize mode: " + f"{quantize_mode}. Mixed {precision}+{quantize_mode} precision is not yet supported. " + f"Setting the custom op precision to be the same as quantize mode." + ) + custom_ops_to_quantize[op_type] = { + "inp": list(range(num_inps)), + "out": list(range(num_outs)), + } + else: + inp_precision, out_precision = precision.split(":") + inp_precision = inp_precision.strip("[]").split(",") + out_precision = out_precision.strip("[]").split(",") + if not all(p in supported_precisions for p in inp_precision + out_precision): + logger.warning( + f"One or more precisions in {inp_precision + out_precision} are not supported. Skipping those." + ) + assert len(inp_precision) == num_inps, ( + f"Number of inputs doesn't match expectation: {len(inp_precision)} vs {num_inps}." + ) + assert len(out_precision) == num_outs, ( + f"Number of outputs doesn't match expectation: {len(out_precision)} vs {num_outs}." + ) + + if any( + p in ["int8", "fp8"] and p != quantize_mode for p in inp_precision + out_precision + ): + logger.warning( + f"Requested custom op precision ('inp': {inp_precision}, 'out': {out_precision}) is different " + f"than quantize mode: {quantize_mode}. Such mixed precision is not yet supported. " + f"Setting the custom op precision to be the same as quantize mode." + ) + + # Will cast the inputs to FP16 and the outputs back to FP32 + inp_precision_cast = [i for i, p in enumerate(inp_precision) if p == "fp16"] + out_precision_cast = [i for i, p in enumerate(out_precision) if p in ["fp16", "fp32"]] + custom_ops_to_cast[op_type] = {"inp": inp_precision_cast, "out": out_precision_cast} + + # Will add Q/DQ nodes in the requested I/O indices + inp_precision_quant = [i for i, p in enumerate(inp_precision) if p in ["int8", "fp8"]] + out_precision_quant = [i for i, p in enumerate(out_precision) if p in ["int8", "fp8"]] + custom_ops_to_quantize[op_type] = { + "inp": inp_precision_quant, + "out": out_precision_quant, + } + + return custom_ops_to_cast, custom_ops_to_quantize diff --git a/modelopt/onnx/utils.py b/modelopt/onnx/utils.py index 2719cbe2e..e4b1aff3b 100644 --- a/modelopt/onnx/utils.py +++ b/modelopt/onnx/utils.py @@ -25,7 +25,7 @@ from typing import Any import numpy as np import onnx import onnx_graphsurgeon as gs -from onnx import numpy_helper +from onnx import ValueInfoProto, numpy_helper from onnx.helper import get_attribute_value from onnx_graphsurgeon import Constant, Node, Variable @@ -287,6 +287,22 @@ def _convert_types_to_np(types: dict[str, int] | list[int] | int) -> Any: return onnx.helper.tensor_dtype_to_np_dtype(types) +def get_tensor_by_name(onnx_model: onnx.ModelProto, tensor_name: str) -> ValueInfoProto | None: + """This function returns a tensor from its name. + + Args: + onnx_model: ONNX model. + tensor_name: tensor name. + + Returns: + tensor + """ + for tensor in onnx_model.graph.value_info: + if tensor.name == tensor_name: + return tensor + return None + + def gen_random_inputs( model: onnx.ModelProto, shapes_spec: str | None = None ) -> dict[str, np.ndarray]: diff --git a/modelopt/torch/__init__.py b/modelopt/torch/__init__.py index c6e8875e2..965d08a8d 100644 --- a/modelopt/torch/__init__.py +++ b/modelopt/torch/__init__.py @@ -32,7 +32,7 @@ if _Version(_torch_version) < _Version("2.5"): try: from transformers import __version__ as _transformers_version - if not (_Version("4.48") <= _Version(_transformers_version) < _Version("4.54")): + if not (_Version("4.48") <= _Version(_transformers_version) < _Version("5.0")): _warnings.warn( f"transformers version {_transformers_version} is incompatible with nvidia-modelopt and may cause issues. " "Please install recommended version with `pip install nvidia-modelopt[hf]` if working with HF models.", diff --git a/modelopt/torch/_deploy/utils/torch_onnx.py b/modelopt/torch/_deploy/utils/torch_onnx.py index 850919264..c2fd3bae4 100644 --- a/modelopt/torch/_deploy/utils/torch_onnx.py +++ b/modelopt/torch/_deploy/utils/torch_onnx.py @@ -28,6 +28,7 @@ import onnx import torch import torch.nn as nn from onnx import ModelProto +from onnxconverter_common import convert_float_to_float16 from packaging.version import Version from torch.nn.parallel import DataParallel, DistributedDataParallel @@ -459,7 +460,12 @@ def get_onnx_bytes_and_metadata( onnx_opt_graph = qdq_to_dq(onnx_opt_graph) if weights_dtype == "float16": - onnx_opt_graph = convert_to_f16(onnx_opt_graph, keep_io_types=False) + if is_fp4_quantized(model) or is_mxfp8_quantized(model): + onnx_opt_graph = convert_float_to_float16( + onnx_opt_graph, keep_io_types=False, disable_shape_infer=True + ) + else: + onnx_opt_graph = convert_to_f16(onnx_opt_graph, keep_io_types=False) # If the onnx model contains external data store the external tensors in one file and save the onnx model if has_external_data(onnx_save_path): diff --git a/modelopt/torch/distill/distillation_model.py b/modelopt/torch/distill/distillation_model.py index cd205bcf3..339ff2c3e 100644 --- a/modelopt/torch/distill/distillation_model.py +++ b/modelopt/torch/distill/distillation_model.py @@ -239,6 +239,7 @@ class DistillationModel(DynamicModule): student_loss: torch.Tensor | None = None, loss_reduction_fn: Callable | None = None, skip_balancer: bool = False, + labels: torch.Tensor | None = None, ) -> torch.Tensor | dict[str, torch.Tensor]: """Compute total loss for distillation backpropagation. @@ -247,6 +248,8 @@ class DistillationModel(DynamicModule): loss_reduction_fn: Callable to be called on each loss tensor prior to balancing. Useful for loss-masking situations where the callable changes arguments each iteration. skip_balancer: Whether or not to use loss balancer to reduce the loss dict into a scalar. + labels: Labels to be passed to the loss function, if needed. This is necessary for losses that + require labels, such as MFTLoss. Returns: If reduce is True, the scalar total loss weighted between ``student_loss`` and the distillation losses. @@ -265,7 +268,8 @@ class DistillationModel(DynamicModule): student_layer._intermediate_output = None teacher_layer._intermediate_output = None - loss = loss_fn(out_s, out_t) # Student is pred, Teacher is target + extra_kwargs = {"labels": labels} if labels is not None else {} + loss = loss_fn(out_s, out_t, **extra_kwargs) # Student is pred, Teacher is target if loss_reduction_fn is not None: # Needed in cases where a loss mask is used on non-scalar loss-fn outputs, prior to # reducing to a scalar loss value. diff --git a/modelopt/torch/distill/losses.py b/modelopt/torch/distill/losses.py index 3ebbdd790..258824bf0 100644 --- a/modelopt/torch/distill/losses.py +++ b/modelopt/torch/distill/losses.py @@ -22,7 +22,7 @@ import torch.nn as nn import torch.nn.functional as F from torch.nn.modules.loss import _Loss as Loss -__all__ = ["LogitsDistillationLoss", "MGDLoss"] +__all__ = ["LogitsDistillationLoss", "MFTLoss", "MGDLoss"] class LogitsDistillationLoss(Loss): @@ -40,8 +40,8 @@ class LogitsDistillationLoss(Loss): use your own reduction function afterwards, i.e. with loss masks. """ super().__init__() - self._temperature = temperature - self._reduction = reduction + self._temperature: float = temperature + self._reduction: str = reduction def forward(self, logits_s: torch.Tensor, logits_t: torch.Tensor) -> torch.Tensor: """Compute KD loss on student and teacher logits. @@ -70,6 +70,131 @@ class LogitsDistillationLoss(Loss): return kd_loss +class MFTLoss(Loss): + """KL-divergence loss with Minifinetuning threshold modification. + + This function implements the distillation loss found in the paper: https://arxiv.org/abs/2506.15702. + """ + + def __init__( + self, temperature: float = 1.0, threshold: float = 0.2, reduction: str = "batchmean" + ): + """Constructor. + + Args: + temperature: A value used to soften the logits_t and logits_s before computing the MFT loss on them. + reduction: How to reduce the final pointwise loss before returning. Pass ``"none"`` to + use your own reduction function afterwards, i.e. with loss masks. + threshold: A value used to correct the teacher's distribution. It is used to ensure that + the separation between the correct and incorrect argmax tokens is large enough. + The value should be in the range [0, 1]. Defaults to 0.2. + """ + super().__init__() + self._temperature: float = temperature + self._reduction: str = reduction + self._threshold: float = threshold + + def forward( + self, logits_s: torch.Tensor, logits_t: torch.Tensor, labels: torch.Tensor + ) -> torch.Tensor: + """Compute KD loss on student and teacher logits. + + Args: + logits_s: Student's logits, treated as prediction. + logits_t: Teacher's logits, treated as training target. + labels: Labels for the ground truth, used to prepare the corrected teacher distributions. + + .. note:: + + Assumes class logits dimension is last. + """ + soft_log_probs = F.log_softmax(logits_s / self._temperature, dim=-1) # (B, ..., C) + soft_log_probs = soft_log_probs.view(-1, soft_log_probs.size(-1)) # (new B, C) + + target_logits: torch.Tensor = logits_t / self._temperature # (B, ..., C) + target_logits = target_logits.view(-1, target_logits.size(-1)) # (new B, C) + soft_targets = self._prepare_corrected_distributions( + target_logits, labels, self._threshold, apply_threshold_to_all=True + ) + + kd_loss = F.kl_div( + soft_log_probs, soft_targets.detach(), reduction=self._reduction + ) # shape depends on reduction; "batchmean" would result in a scalar (1,) + + # Since the magnitudes of the gradients produced by the soft logits scale as 1/(T^2), + # multiplying them by T^2 ensures that the relative contributions of the logits + # remain roughly unchanged while experimenting with meta-parameters. + kd_loss *= self._temperature**2 + + return kd_loss + + def _prepare_corrected_distributions( + self, + logits: torch.Tensor, + labels: torch.Tensor, + threshold: float, + apply_threshold_to_all: bool = True, + ) -> torch.Tensor: + """Prepare the corrected distributions for MFT loss. + + Args: + logits: The logits from the teacher model, shape (batch, channels) # e.g. (batch_size * seq_len, vocab_size) + in case of LMs + labels: The ground truth labels, shape (batch) # e.g. (batch_size * seq_len) in case of LMs + threshold: The threshold value for the MFT correction. + apply_threshold_to_all: If True, apply the threshold correction to all tokens, + not just the incorrect argmax tokens. Defaults to True. + + Returns: + A tensor containing the corrected distributions, shape (batch_size * seq_len, vocab_size). + """ + # Ensure logits is a 2D tensor and labels is a 1D tensor + if logits.dim() != 2 or labels.dim() != 1: + raise ValueError("Logits must be a 2D tensor and labels must be a 1D tensor.") + # logits: (batch, channels) + # labels: (batch) + distribution = F.softmax(logits, dim=-1) # (batch, channels) + + argmax = distribution.argmax(dim=-1) # (batch,) + incorrect_argmax = argmax != labels # (batch,) + + p_argmax = torch.gather(distribution, 1, argmax.unsqueeze(1)).squeeze(1) # (batch,) + p_label = torch.gather(distribution, 1, labels.unsqueeze(1)).squeeze(1) # (batch,) + + # correction of the distribution at the tokens where the argmax is incorrect + mixin_factor = (p_argmax - p_label + threshold) / ( + 1 + p_argmax - p_label + 1e-7 + ) # (batch,) + adjusted_incorrect_distribution = distribution * ( + 1 - mixin_factor.unsqueeze(1) + ) # (batch, channels) + _ = adjusted_incorrect_distribution.scatter_add_( + 1, labels.unsqueeze(1), mixin_factor.unsqueeze(1) + ) # (batch, channels) + + if apply_threshold_to_all: + # correction of the distribution at the tokens where the argmax is correct but + # the separation may not be large enough + capped_targets = torch.where( + p_label > 1 - threshold, 1, p_label + threshold + ) # (batch,) + mixin_factor = (capped_targets - p_argmax) / (1 - p_argmax + 1e-7) # (batch,) + adjusted_correct_distribution = distribution * ( + 1 - mixin_factor.unsqueeze(1) + ) # (batch, channels) + _ = adjusted_correct_distribution.scatter_add_( + 1, labels.unsqueeze(1), mixin_factor.unsqueeze(1) + ) + else: + adjusted_correct_distribution = distribution + + return torch.where( + incorrect_argmax.unsqueeze(1), + adjusted_incorrect_distribution, + adjusted_correct_distribution, + ) # (batch, channels) + + class MGDLoss(Loss): """PyTorch version of Masked Generative Distillation. @@ -92,8 +217,8 @@ class MGDLoss(Loss): lambda_mgd: Masked ratio. Defaults to 0.65. """ super().__init__() - self._alpha_mgd = alpha_mgd - self._lambda_mgd = lambda_mgd + self._alpha_mgd: float = alpha_mgd + self._lambda_mgd: float = lambda_mgd if num_student_channels != num_teacher_channels: self.align = nn.Conv2d( diff --git a/modelopt/torch/export/convert_hf_config.py b/modelopt/torch/export/convert_hf_config.py index f1edf6276..86cee9580 100644 --- a/modelopt/torch/export/convert_hf_config.py +++ b/modelopt/torch/export/convert_hf_config.py @@ -15,13 +15,23 @@ """Convert modelopt quantization export config to align with llm-compressor config format.""" +from typing import Any -def convert_hf_quant_config_format(input_config: dict) -> dict: + +def convert_hf_quant_config_format(input_config: dict[str, Any]) -> dict[str, Any]: """Converts modelopt quantization config dictionary to align with llm-compressor config format. Args: input_config: The original quantization config dictionary. + Note: + The "targets" field specifies which PyTorch module types to quantize. Compressed-tensors + works with any PyTorch module type and uses dynamic matching against module.__class__.__name__. + Typically this includes "Linear" modules, but can also include "Embedding" and other types. + + See: https://github.com/neuralmagic/compressed-tensors/blob/fa6a48f1da6b47106912bcd25eba7171ba7cfec7/src/sparsetensors/quantization/quant_scheme.py#L29 + Example usage: https://github.com/neuralmagic/compressed-tensors/blob/9938a6ec6e10498d39a3071dfd1c40e3939ee80b/tests/test_quantization/lifecycle/test_apply.py#L118 + Example: .. code-block:: python @@ -55,7 +65,7 @@ def convert_hf_quant_config_format(input_config: dict) -> dict: "producer": {"name": "modelopt", "version": "0.29.0"}, } """ - new_config = {} + new_config: dict[str, Any] = {} original_quantization_details = input_config.get("quantization", {}) quant_algo_value = original_quantization_details.get("quant_algo") @@ -66,6 +76,7 @@ def convert_hf_quant_config_format(input_config: dict) -> dict: config_group_details = { "input_activations": {"dynamic": False, "num_bits": 8, "type": "float"}, "weights": {"dynamic": False, "num_bits": 8, "type": "float"}, + "targets": ["Linear"], } new_config["config_groups"] = {"group_0": config_group_details} elif quant_algo_value == "NVFP4": @@ -78,6 +89,7 @@ def convert_hf_quant_config_format(input_config: dict) -> dict: "group_size": group_size, }, "weights": {"dynamic": False, "num_bits": 4, "type": "float", "group_size": group_size}, + "targets": ["Linear"], } new_config["config_groups"] = {"group_0": config_group_details} @@ -90,12 +102,16 @@ def convert_hf_quant_config_format(input_config: dict) -> dict: kv_cache_quant_algo = original_quantization_details.get("kv_cache_quant_algo") if kv_cache_quant_algo: - new_config["kv_cache_scheme"] = kv_cache_quant_algo + if kv_cache_quant_algo == "FP8": + new_config["kv_cache_scheme"] = {"dynamic": False, "num_bits": 8, "type": "float"} + else: + # TODO: Handle other kv cache quantization algorithms + new_config["kv_cache_scheme"] = kv_cache_quant_algo producer_info = input_config.get("producer") if producer_info: new_config["producer"] = producer_info - new_config["quant_library"] = "modelopt" + new_config["quant_method"] = "modelopt" return new_config diff --git a/modelopt/torch/export/plugins/__init__.py b/modelopt/torch/export/plugins/__init__.py index 83c11aa83..a7bbc3fb3 100644 --- a/modelopt/torch/export/plugins/__init__.py +++ b/modelopt/torch/export/plugins/__init__.py @@ -14,3 +14,8 @@ # limitations under the License. """Export package plugin.""" + +from modelopt.torch.utils import import_plugin + +with import_plugin("megatron_importer"): + from .megatron_importer import * diff --git a/modelopt/torch/export/plugins/mcore_common.py b/modelopt/torch/export/plugins/mcore_common.py index d001b8ffd..f44a9e5f7 100644 --- a/modelopt/torch/export/plugins/mcore_common.py +++ b/modelopt/torch/export/plugins/mcore_common.py @@ -21,6 +21,8 @@ from .mcore_deepseek import deepseek_causal_lm_export, deepseek_causal_lm_import from .mcore_llama import ( eagle3_llama_causal_lm_export, eagle_llama_causal_lm_export, + llama4_causal_lm_export, + llama4_causal_lm_import, llama_causal_lm_export, llama_causal_lm_import, ) @@ -35,7 +37,7 @@ all_mcore_hf_export_mapping: dict[str, Any] = { "DeepseekV2ForCausalLM": deepseek_causal_lm_export, "DeepseekV3ForCausalLM": deepseek_causal_lm_export, "LlamaForCausalLM": llama_causal_lm_export, - "Llama4ForConditionalGeneration": {}, + "Llama4ForConditionalGeneration": llama4_causal_lm_export, "NemotronForCausalLM": nemotron_causal_lm_export, "NemotronHForCausalLM": nemotron_h_causal_lm_export, "LlamaForCausalLMEagle": eagle_llama_causal_lm_export, @@ -46,6 +48,7 @@ all_mcore_hf_export_mapping: dict[str, Any] = { all_mcore_hf_import_mapping: dict[str, Any] = { "LlamaForCausalLM": llama_causal_lm_import, + "Llama4ForConditionalGeneration": llama4_causal_lm_import, "DeepseekV2ForCausalLM": deepseek_causal_lm_import, "DeepseekV3ForCausalLM": deepseek_causal_lm_import, "NemotronHForCausalLM": nemotron_h_causal_lm_import, diff --git a/modelopt/torch/export/plugins/mcore_custom.py b/modelopt/torch/export/plugins/mcore_custom.py index 47c12a3aa..063bac5c6 100644 --- a/modelopt/torch/export/plugins/mcore_custom.py +++ b/modelopt/torch/export/plugins/mcore_custom.py @@ -14,28 +14,56 @@ # limitations under the License. -"""Custom mapping utility.""" +"""Custom Megatron mapping and safetensors utility.""" import json import os +from dataclasses import dataclass +from pathlib import Path from typing import Any import torch -from safetensors.torch import save_file +from safetensors.torch import safe_open, save_file from modelopt.torch.utils import import_plugin with import_plugin("megatron"): from megatron.core.parallel_state import ( + get_expert_model_parallel_rank, + get_expert_model_parallel_world_size, + get_expert_tensor_parallel_rank, + get_expert_tensor_parallel_world_size, get_pipeline_model_parallel_rank, get_pipeline_model_parallel_world_size, + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, ) MAX_SAFETENSOR_SIZE = 4000000000 PER_RANK_PARTITIONS = 10 -COL_PARALLEL = {"sharding_dim": 0} -ROW_PARALLEL = {"sharding_dim": 1} + +@dataclass +class ParallelConfig: + """The parallel configuration including sharding_dim and parallel_group. + + Args: + sharding_dim: The dimension (0 or 1) to shard the tensor. + parallel_group: The parallel group (TP or ETP) to shard the tensor. + """ + + sharding_dim: int | None = None + parallel_group: str | None = None + + +COL_TP = {"parallel_config": ParallelConfig(sharding_dim=0, parallel_group="TP")} +ROW_TP = {"parallel_config": ParallelConfig(sharding_dim=1, parallel_group="TP")} +COL_ETP = {"parallel_config": ParallelConfig(sharding_dim=0, parallel_group="ETP")} +ROW_ETP = {"parallel_config": ParallelConfig(sharding_dim=1, parallel_group="ETP")} +PACK_COL_ETP = {"parallel_config": ParallelConfig(sharding_dim=2, parallel_group="ETP")} +PACK_ROW_ETP = {"parallel_config": ParallelConfig(sharding_dim=1, parallel_group="ETP")} +PACK_EP = {"parallel_config": ParallelConfig(sharding_dim=0, parallel_group="EP")} +REPLICATE = {} class CustomModuleMapping: @@ -50,19 +78,106 @@ class CustomModuleMapping: self.func_kwargs = func_kwargs +class NameRemapping(CustomModuleMapping): + """A custom module mapping that renames of the modules.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that renames of the modules.""" + super().__init__( + func_name="name_remapping", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + +class QKVMerging(CustomModuleMapping): + """A custom module mapping that merges Q, K, V.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that merges Q, K, V.""" + super().__init__( + func_name="qkv_merging", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + +class GatedMLPMerging(CustomModuleMapping): + """A custom module mapping that merges gate_proj and up_proj.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that merges gate_proj and up_proj.""" + super().__init__( + func_name="gated_mlp_merging", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + +class QKVSlicing(CustomModuleMapping): + """A custom module mapping that slices Q, K, V.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that slices Q, K, V.""" + super().__init__( + func_name="qkv_slicing", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + +class GatedMLPSlicing(CustomModuleMapping): + """A custom module mapping that slices gate_proj and up_proj.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that slices gate_proj and up_proj.""" + super().__init__( + func_name="gated_mlp_slicing", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + +class PackNameRemapping(CustomModuleMapping): + """A custom module mapping that packs module after name remapping.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that packs and renames module.""" + super().__init__( + func_name="pack_name_remapping", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + +class UnpackNameRemapping(CustomModuleMapping): + """A custom module mapping that unpacks module after name remapping.""" + + def __init__(self, target_name_or_prefix: str = "", func_kwargs: dict[str, Any] = {}): + """Create a custom module mapping that unpacks module after name remapping.""" + super().__init__( + func_name="unpack_name_remapping", + target_name_or_prefix=target_name_or_prefix, + func_kwargs=func_kwargs, + ) + + def save_safetensors(state_dict, save_directory: str | os.PathLike): """Save safetensors with pipeline model parallel support.""" pp_rank = get_pipeline_model_parallel_rank() pp_size = get_pipeline_model_parallel_world_size() - all_safetensors = [{} for i in range(10)] + all_safetensors = [{} for i in range(PER_RANK_PARTITIONS)] local_idx = 0 local_total_size = 0 for key, val in state_dict.items(): tensor_size = val.numel() * val.element_size() - if local_total_size + tensor_size > MAX_SAFETENSOR_SIZE and local_idx < PER_RANK_PARTITIONS: + if ( + local_total_size + tensor_size > MAX_SAFETENSOR_SIZE + and local_idx < PER_RANK_PARTITIONS - 1 + ): local_idx += 1 local_total_size = 0 all_safetensors[local_idx][key] = val @@ -108,3 +223,191 @@ def save_safetensors(state_dict, save_directory: str | os.PathLike): with open(save_directory + "/model.safetensors.index.json", "w") as f: json.dump(safetensor_index, f, indent=4) + + +def _get_safetensors_file(pretrained_model_path: str | Path, key: str) -> Path | None: + """Given a tensor key return the safetensors file that contains this tensor if exists. + + Args: + pretrained_model_path: The path to the pretrained model. + key: The key of the tensor to get. + """ + safetensors_file = Path(pretrained_model_path) / "model.safetensors" + safetensors_index_file = Path(pretrained_model_path) / "model.safetensors.index.json" + if safetensors_file.is_file(): + pass + elif safetensors_index_file.is_file(): + with open(safetensors_index_file) as f: + safetensors_index = json.load(f) + safetensors_file = ( + (Path(pretrained_model_path) / safetensors_index["weight_map"][key]) + if key in safetensors_index["weight_map"] + else None + ) + else: + raise ValueError("Only safetensors (single of multi- files) are supported.") + return safetensors_file + + +def _get_safetensor_slices( + pretrained_model_path: str | Path, key: str, parallel_config: ParallelConfig | None = None +): + """Given a tensor key return the safetensors slice if exists. + + Depending on the parallel_config, the tensor will be sharded along the sharding_dim over + the size of the parallel group (equally divided over each rank in the group). + The sharding_dim can be 0 or 1, which corresponds to the column parallel or row parallel, + respectively. The parallel group can be tensor model parallel (TP) group or expert tensor + parallel (ETP) group. + + Args: + pretrained_model_path: The path to the pretrained model. + key: The key of the tensor to get. + parallel_config: The parallel configuration including sharding_dim and parallel_group. + """ + safetensors_file = _get_safetensors_file(pretrained_model_path, key) + + if safetensors_file is None: + return None + + with safe_open(safetensors_file, framework="pt") as f: + if key not in f.keys(): # noqa: SIM118 + tensor = None + elif parallel_config is None: + tensor = f.get_tensor(key) + else: + tensor_slice = f.get_slice(key) + assert tensor_slice is not None + shape = tensor_slice.get_shape() + assert len(shape) in (2, 3), f"Shape {shape} is not supported!" + # MCore tensor parallel model sharding + sharding_dim = parallel_config.sharding_dim + parallel_group = parallel_config.parallel_group + + if parallel_group == "EP": + # For packed EP case, Llama4 uses 3D tensor for local experts + tp_size = get_expert_tensor_parallel_world_size() + if tp_size > 1: + raise ValueError("Packed MoE import only supports ETP=1.") + if len(shape) != 3: + raise ValueError( + "Packed MoE import only supports 3D tensor in shape [num_experts, in_dim, out_dim]." + ) + if sharding_dim != 0: + raise ValueError("Packed MoE import only supports sharding_dim=0.") + ep_rank = get_expert_model_parallel_rank() + ep_size = get_expert_model_parallel_world_size() + per_rank_size = shape[sharding_dim] // ep_size + rank_offset = ep_rank * per_rank_size + if shape[sharding_dim] % ep_size > 0: + raise ValueError( + "{} {} @ dim {} is not a multiple of {}".format( + key, shape, sharding_dim, ep_size + ) + ) + tensor = tensor_slice[rank_offset : rank_offset + per_rank_size, :, :] + else: + if parallel_group == "TP": + tp_rank = get_tensor_model_parallel_rank() + tp_size = get_tensor_model_parallel_world_size() + elif parallel_group == "ETP": + tp_rank = get_expert_tensor_parallel_rank() + tp_size = get_expert_tensor_parallel_world_size() + else: + raise ValueError(f"Unsupported parallel group [tp|etp]: {parallel_group}") + per_rank_size = shape[sharding_dim] // tp_size + rank_offset = tp_rank * per_rank_size + + if shape[sharding_dim] % tp_size > 0: + raise ValueError( + "{} {} @ dim {} is not a multiple of {}".format( + key, shape, sharding_dim, tp_size + ) + ) + if len(shape) == 2: + if sharding_dim in (1, -1): + tensor = tensor_slice[:, rank_offset : rank_offset + per_rank_size] + else: + tensor = tensor_slice[rank_offset : rank_offset + per_rank_size, :] + elif len(shape) == 3: + # For packed ETP case, Llama4 uses 3D tensor for local experts + if sharding_dim == 1: + tensor = tensor_slice[:, rank_offset : rank_offset + per_rank_size, :] + elif sharding_dim == 2: + tensor = tensor_slice[:, :, rank_offset : rank_offset + per_rank_size] + else: + raise ValueError( + f"Unsupported sharding_dim: {sharding_dim} for shape: {shape}" + ) + else: + raise ValueError(f"Unsupported shape: {shape}") + return tensor + + +def _get_dsfp8_weight_scale_inv( + pretrained_model_path: str | Path, key: str, parallel_config: ParallelConfig | None = None +): + """Given a weight tensor key return weight_scale_inv if exists. + + Args: + pretrained_model_path: The path to the pretrained model. + key: The key of the tensor to get. + parallel_config: The parallel configuration including sharding_dim and parallel_group. + """ + if not key.endswith("weight"): + return None + return _get_safetensor_slices(pretrained_model_path, key + "_scale_inv", parallel_config) + + +def get_safetensor( + pretrained_model_path: str | Path, + key: str, + parallel_config: ParallelConfig | None = None, + dequantize: bool = False, +) -> torch.Tensor: + """Get a safetensor from the sharded checkpoint. + + Args: + pretrained_model_path: The path to the pretrained model. + key: The key of the tensor to get. + parallel_config: The parallel configuration including sharding_dim and parallel_group. + """ + tensor = _get_safetensor_slices(pretrained_model_path, key, parallel_config) + + if tensor is None: + raise ValueError(f"Key [{key}] does not exist!") + + if dequantize: + # DSFP8 (weight-only 128x128-blocked FP8) + if tensor.dtype is torch.float8_e4m3fn: + weight_scale_inv = _get_dsfp8_weight_scale_inv( + pretrained_model_path, key, parallel_config + ) + if weight_scale_inv is None: + raise ValueError(f"Fail to dequantize [{key}]! weight_sacle_inv not found!") + else: + # [TODO]: may need padding here + out_dim, in_dim = tensor.shape + blk_out_dim, blk_in_dim = weight_scale_inv.shape + + tensor = tensor.to(torch.bfloat16) + if out_dim % 128 > 0 or in_dim % 128 > 0: + padded_tensor = torch.zeros( + blk_out_dim * 128, blk_in_dim * 128, dtype=torch.bfloat16 + ) + padded_tensor[:out_dim, :in_dim] = tensor + else: + padded_tensor = tensor + + padded_tensor = padded_tensor.reshape(blk_out_dim, 128, blk_in_dim, 128) + weight_scale_inv = weight_scale_inv.to(torch.bfloat16) + weight_scale_inv = weight_scale_inv.reshape(blk_out_dim, 1, blk_in_dim, 1) + padded_tensor = padded_tensor * weight_scale_inv + padded_tensor = padded_tensor.view(blk_out_dim * 128, blk_in_dim * 128) + + if out_dim % 128 > 0 or in_dim % 128 > 0: + tensor = padded_tensor[:out_dim, :in_dim] + else: + tensor = padded_tensor + + return tensor.contiguous() diff --git a/modelopt/torch/export/plugins/mcore_deepseek.py b/modelopt/torch/export/plugins/mcore_deepseek.py index d4ce55860..d02259e35 100644 --- a/modelopt/torch/export/plugins/mcore_deepseek.py +++ b/modelopt/torch/export/plugins/mcore_deepseek.py @@ -16,278 +16,128 @@ """Custom mapping from DeepSeek Hugging Face models to Megatron Core models.""" -from .mcore_custom import COL_PARALLEL, ROW_PARALLEL, CustomModuleMapping +from .mcore_custom import ( + COL_ETP, + COL_TP, + REPLICATE, + ROW_ETP, + ROW_TP, + CustomModuleMapping, + GatedMLPMerging, + GatedMLPSlicing, + NameRemapping, +) deepseek_causal_lm_export: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens."), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + "word_embeddings": NameRemapping("model.embed_tokens."), + "final_layernorm": NameRemapping("model.norm."), + "output_layer": NameRemapping("lm_head."), # Multi-Latent Attention (V3 has lora on q as well) - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_q_proj": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.q_proj."), - "linear_q_down_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_a_proj." - ), - "linear_q_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_a_layernorm." - ), - "linear_q_up_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_b_proj." - ), - "linear_kv_down_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_a_proj_with_mqa." - ), - "linear_kv_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_a_layernorm." - ), - "linear_kv_up_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_b_proj." - ), - "linear_proj": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.o_proj."), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." - ), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."), + "linear_q_proj": NameRemapping("model.layers.{}.self_attn.q_proj."), + "linear_q_down_proj": NameRemapping("model.layers.{}.self_attn.q_a_proj."), + "linear_q_layernorm": NameRemapping("model.layers.{}.self_attn.q_a_layernorm."), + "linear_q_up_proj": NameRemapping("model.layers.{}.self_attn.q_b_proj."), + "linear_kv_down_proj": NameRemapping("model.layers.{}.self_attn.kv_a_proj_with_mqa."), + "linear_kv_layernorm": NameRemapping("model.layers.{}.self_attn.kv_a_layernorm."), + "linear_kv_up_proj": NameRemapping("model.layers.{}.self_attn.kv_b_proj."), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."), # MLP for dense layers - "linear_fc1": CustomModuleMapping("gated_mlp_slicing", "model.layers.{}.mlp."), - "linear_fc2": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.down_proj."), + "linear_fc1": GatedMLPSlicing("model.layers.{}.mlp."), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."), # MoE shared experts - "router": CustomModuleMapping( - "name_remapping", - "model.layers.{}.mlp.gate.", - {"mapping": {"expert_bias": "e_score_correction_bias"}}, - ), - "shared_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "model.layers.{}.mlp.shared_experts." - ), - "shared_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.shared_experts.down_proj." + "router": NameRemapping( + "model.layers.{}.mlp.gate.", {"mapping": {"expert_bias": "e_score_correction_bias"}} ), + "shared_experts.linear_fc1": GatedMLPSlicing("model.layers.{}.mlp.shared_experts."), + "shared_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.shared_experts.down_proj."), # MoE local experts - "local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "model.layers.{}.mlp.experts.{}." - ), - "local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.experts.{}.down_proj." - ), - # MedusaForCausalLM support - "medusa_heads.lm_head": CustomModuleMapping( - "name_remapping", "medusa_heads.{}.1." - ), # TODO: lm_head is hardcoded to .1 as currently only support using 1 layer in medusa head - # needs a fix - "medusa_heads.medusa_layers.linear": CustomModuleMapping( - "name_remapping", "medusa_heads.{}.{}.linear." - ), - # EagleForCausalLM support - "eagle_module.fc": CustomModuleMapping("name_remapping", "eagle_module.fc."), - "eagle_module.input_layernorm": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.input_layernorm." - ), - "eagle_module.linear_q_proj": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.q_proj." - ), - "eagle_module.linear_q_down_proj": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.q_a_proj." - ), - "eagle_module.linear_q_layernorm": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.q_a_layernorm." - ), - "eagle_module.linear_q_up_proj": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.q_b_proj." - ), - "eagle_module.linear_kv_down_proj": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.kv_a_proj_with_mqa." - ), - "eagle_module.linear_kv_layernorm": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.kv_a_layernorm." - ), - "eagle_module.linear_kv_up_proj": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.kv_b_proj." - ), - "eagle_module.linear_proj": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.self_attn.o_proj." - ), - "eagle_module.pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.post_attention_layernorm." - ), - "eagle_module.router": CustomModuleMapping( - "name_remapping", - "eagle_module.layers.{}.mlp.gate.", - {"mapping": {"expert_bias": "e_score_correction_bias"}}, - ), - "eagle_module.shared_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "eagle_module.layers.{}.mlp.shared_experts." - ), - "eagle_module.shared_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.mlp.shared_experts.down_proj." - ), - "eagle_module.local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "eagle_module.layers.{}.mlp.experts.{}." - ), - "eagle_module.local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.mlp.experts.{}.down_proj." - ), - "eagle_module.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "eagle_module.layers.{}.mlp." - ), - "eagle_module.linear_fc2": CustomModuleMapping( - "name_remapping", "eagle_module.layers.{}.mlp.down_proj." - ), - # MTPCausalLM support - "mtp.fc": CustomModuleMapping("name_remapping", "mtp.{}.fc."), - "mtp.enorm": CustomModuleMapping("name_remapping", "mtp.{}.enorm."), - "mtp.hnorm": CustomModuleMapping("name_remapping", "mtp.{}.hnorm."), - "mtp.input_layernorm": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.input_layernorm." - ), - "mtp.linear_q_proj": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.q_proj." - ), - "mtp.linear_q_down_proj": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.q_a_proj." - ), - "mtp.linear_q_layernorm": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.q_a_layernorm." - ), - "mtp.linear_q_up_proj": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.q_b_proj." - ), - "mtp.linear_kv_down_proj": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.kv_a_proj_with_mqa." - ), - "mtp.linear_kv_layernorm": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.kv_a_layernorm." - ), - "mtp.linear_kv_up_proj": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.self_attn.kv_b_proj." - ), - "mtp.linear_proj": CustomModuleMapping("name_remapping", "mtp.{}.layers.{}.self_attn.o_proj."), - "mtp.pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.post_attention_layernorm." - ), - "mtp.router": CustomModuleMapping( - "name_remapping", - "mtp.{}.layers.{}.mlp.gate.", - {"mapping": {"expert_bias": "e_score_correction_bias"}}, - ), - "mtp.shared_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "mtp.{}.layers.{}.mlp.shared_experts." - ), - "mtp.shared_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.mlp.shared_experts.down_proj." - ), - "mtp.local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "mtp.{}.layers.{}.mlp.experts.{}." - ), - "mtp.local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "mtp.{}.layers.{}.mlp.experts.{}.down_proj." - ), + "local_experts.linear_fc1": GatedMLPSlicing("model.layers.{}.mlp.experts.{}."), + "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj."), } + +eagle_mtp_deepseek_causal_lm_export: dict[str, CustomModuleMapping] = { + "fc": NameRemapping("eh_proj."), + "enorm": NameRemapping("enorm."), + "hnorm": NameRemapping("hnorm."), + "input_layernorm": NameRemapping("layers.{}.input_layernorm."), + "linear_q_proj": NameRemapping("layers.{}.self_attn.q_proj."), + "linear_q_down_proj": NameRemapping("layers.{}.self_attn.q_a_proj."), + "linear_q_layernorm": NameRemapping("layers.{}.self_attn.q_a_layernorm."), + "linear_q_up_proj": NameRemapping("layers.{}.self_attn.q_b_proj."), + "linear_kv_down_proj": NameRemapping("layers.{}.self_attn.kv_a_proj_with_mqa."), + "linear_kv_layernorm": NameRemapping("layers.{}.self_attn.kv_a_layernorm."), + "linear_kv_up_proj": NameRemapping("layers.{}.self_attn.kv_b_proj."), + "linear_proj": NameRemapping("layers.{}.self_attn.o_proj."), + "pre_mlp_layernorm": NameRemapping("layers.{}.post_attention_layernorm."), + "router": NameRemapping( + "layers.{}.mlp.gate.", {"mapping": {"expert_bias": "e_score_correction_bias"}} + ), + "shared_experts.linear_fc1": GatedMLPSlicing("layers.{}.mlp.shared_experts."), + "shared_experts.linear_fc2": NameRemapping("layers.{}.mlp.shared_experts.down_proj."), + "local_experts.linear_fc1": GatedMLPSlicing("layers.{}.mlp.experts.{}."), + "local_experts.linear_fc2": NameRemapping("layers.{}.mlp.experts.{}.down_proj."), +} + + deepseek_causal_lm_import = { - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens.", COL_PARALLEL), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head.", COL_PARALLEL), + "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP), + "final_layernorm": NameRemapping("model.norm.", REPLICATE), + "output_layer": NameRemapping("lm_head.", COL_TP), # Per-layer - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_q_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_proj.", COL_PARALLEL - ), - "linear_q_down_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_a_proj.", COL_PARALLEL - ), - "linear_q_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_a_layernorm." - ), - "linear_q_up_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_b_proj.", COL_PARALLEL - ), - "linear_kv_down_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_a_proj_with_mqa.", COL_PARALLEL - ), - "linear_kv_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_a_layernorm." - ), - "linear_kv_up_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_b_proj.", COL_PARALLEL - ), - "linear_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.o_proj.", ROW_PARALLEL - ), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." - ), - "linear_fc1": CustomModuleMapping("gated_mlp_merging", "model.layers.{}.mlp.", COL_PARALLEL), - "linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.down_proj.", ROW_PARALLEL + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE), + "linear_q_proj": NameRemapping("model.layers.{}.self_attn.q_proj.", COL_TP), + "linear_q_down_proj": NameRemapping("model.layers.{}.self_attn.q_a_proj.", REPLICATE), + "linear_q_layernorm": NameRemapping("model.layers.{}.self_attn.q_a_layernorm.", REPLICATE), + "linear_q_up_proj": NameRemapping("model.layers.{}.self_attn.q_b_proj.", COL_TP), + "linear_kv_down_proj": NameRemapping( + "model.layers.{}.self_attn.kv_a_proj_with_mqa.", REPLICATE ), + "linear_kv_layernorm": NameRemapping("model.layers.{}.self_attn.kv_a_layernorm.", REPLICATE), + "linear_kv_up_proj": NameRemapping("model.layers.{}.self_attn.kv_b_proj.", COL_TP), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE), + "linear_fc1": GatedMLPMerging("model.layers.{}.mlp.", COL_TP), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj.", ROW_TP), # MoE shared experts - "router": CustomModuleMapping( - "name_remapping", - "model.layers.{}.mlp.gate.", - {"mapping": {"expert_bias": "e_score_correction_bias"}}, + "router": NameRemapping( + "model.layers.{}.mlp.gate.", {"mapping": {"expert_bias": "e_score_correction_bias"}} ), - "shared_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_merging", "model.layers.{}.mlp.shared_experts.", COL_PARALLEL - ), - "shared_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.shared_experts.down_proj.", ROW_PARALLEL + "shared_experts.linear_fc1": GatedMLPMerging("model.layers.{}.mlp.shared_experts.", COL_TP), + "shared_experts.linear_fc2": NameRemapping( + "model.layers.{}.mlp.shared_experts.down_proj.", ROW_TP ), # MoE local experts - "local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_merging", "model.layers.{}.mlp.experts.{}.", COL_PARALLEL - ), - "local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.experts.{}.down_proj.", ROW_PARALLEL - ), - # MTP - "mtp.fc": CustomModuleMapping("name_remapping", "model.layers.{}.eh_proj."), - "mtp.enorm": CustomModuleMapping("name_remapping", "model.layers.{}.enorm."), - "mtp.hnorm": CustomModuleMapping("name_remapping", "model.layers.{}.hnorm."), - "mtp.input_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.input_layernorm." - ), - "mtp.linear_q_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_proj.", COL_PARALLEL - ), - "mtp.linear_q_down_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_a_proj.", COL_PARALLEL - ), - "mtp.linear_q_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_a_layernorm." - ), - "mtp.linear_q_up_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.q_b_proj.", COL_PARALLEL - ), - "mtp.linear_kv_down_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_a_proj_with_mqa.", COL_PARALLEL - ), - "mtp.linear_kv_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_a_layernorm." - ), - "mtp.linear_kv_up_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.kv_b_proj.", COL_PARALLEL - ), - "mtp.linear_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.o_proj.", ROW_PARALLEL - ), - "mtp.pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." - ), - "mtp.router": CustomModuleMapping( - "name_remapping", - "model.layers.{}.mlp.gate.", - {"mapping": {"expert_bias": "e_score_correction_bias"}}, - ), - "mtp.shared_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_merging", "model.layers.{}.mlp.shared_experts.", COL_PARALLEL - ), - "mtp.shared_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.shared_experts.down_proj.", ROW_PARALLEL - ), - "mtp.local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_merging", "model.layers.{}.mlp.experts.{}.", COL_PARALLEL - ), - "mtp.local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.experts.{}.down_proj.", ROW_PARALLEL - ), + "local_experts.linear_fc1": GatedMLPMerging("model.layers.{}.mlp.experts.{}.", COL_ETP), + "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj.", ROW_ETP), +} + + +eagle_mtp_deepseek_causal_lm_import: dict[str, CustomModuleMapping] = { + "fc": NameRemapping("model.layers.{}.eh_proj.", REPLICATE), + "enorm": NameRemapping("model.layers.{}.enorm.", REPLICATE), + "hnorm": NameRemapping("model.layers.{}.hnorm.", REPLICATE), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE), + "linear_q_proj": NameRemapping("model.layers.{}.self_attn.q_proj.", COL_TP), + "linear_q_down_proj": NameRemapping("model.layers.{}.self_attn.q_a_proj.", REPLICATE), + "linear_q_layernorm": NameRemapping("model.layers.{}.self_attn.q_a_layernorm.", REPLICATE), + "linear_q_up_proj": NameRemapping("model.layers.{}.self_attn.q_b_proj.", COL_TP), + "linear_kv_down_proj": NameRemapping( + "model.layers.{}.self_attn.kv_a_proj_with_mqa.", REPLICATE + ), + "linear_kv_layernorm": NameRemapping("model.layers.{}.self_attn.kv_a_layernorm.", REPLICATE), + "linear_kv_up_proj": NameRemapping("model.layers.{}.self_attn.kv_b_proj.", COL_TP), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE), + "router": NameRemapping( + "model.layers.{}.mlp.gate.", {"mapping": {"expert_bias": "e_score_correction_bias"}} + ), + "shared_experts.linear_fc1": GatedMLPMerging("model.layers.{}.mlp.shared_experts.", COL_TP), + "shared_experts.linear_fc2": NameRemapping( + "model.layers.{}.mlp.shared_experts.down_proj.", ROW_TP + ), + "local_experts.linear_fc1": GatedMLPMerging("model.layers.{}.mlp.experts.{}.", COL_ETP), + "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj.", ROW_ETP), } diff --git a/modelopt/torch/export/plugins/mcore_llama.py b/modelopt/torch/export/plugins/mcore_llama.py index feac8472a..aa48e9e99 100644 --- a/modelopt/torch/export/plugins/mcore_llama.py +++ b/modelopt/torch/export/plugins/mcore_llama.py @@ -16,84 +16,149 @@ """Custom mapping from Llama Hugging Face models to Megatron Core models.""" -from .mcore_custom import COL_PARALLEL, ROW_PARALLEL, CustomModuleMapping +from .mcore_custom import ( + COL_TP, + PACK_COL_ETP, + PACK_EP, + PACK_ROW_ETP, + REPLICATE, + ROW_TP, + CustomModuleMapping, + GatedMLPMerging, + GatedMLPSlicing, + NameRemapping, + PackNameRemapping, + QKVMerging, + QKVSlicing, + UnpackNameRemapping, +) llama_causal_lm_export: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens."), - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_qkv": CustomModuleMapping( - "qkv_slicing", - "model.layers.{}.self_attn.", + "word_embeddings": NameRemapping("model.embed_tokens."), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."), + "linear_qkv": QKVSlicing("model.layers.{}.self_attn."), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."), + "linear_fc1": GatedMLPSlicing("model.layers.{}.mlp."), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."), + "final_layernorm": NameRemapping("model.norm."), + "output_layer": NameRemapping("lm_head."), +} + +llama4_causal_lm_export: dict[str, CustomModuleMapping | bool] = { + "word_embeddings": NameRemapping("language_model.model.embed_tokens."), + "input_layernorm": NameRemapping("language_model.model.layers.{}.input_layernorm."), + # self_attn + "linear_qkv": QKVSlicing("language_model.model.layers.{}.self_attn."), + "linear_proj": NameRemapping("language_model.model.layers.{}.self_attn.o_proj."), + # mlp + "pre_mlp_layernorm": NameRemapping("language_model.model.layers.{}.post_attention_layernorm."), + "shared_experts.linear_fc1": GatedMLPSlicing( + "language_model.model.layers.{}.feed_forward.shared_expert.", ), - "linear_proj": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.o_proj."), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." + "shared_experts.linear_fc2": NameRemapping( + "language_model.model.layers.{}.feed_forward.shared_expert.down_proj.", ), - "linear_fc1": CustomModuleMapping("gated_mlp_slicing", "model.layers.{}.mlp."), - "linear_fc2": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.down_proj."), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + # moe_layer + "router": NameRemapping("language_model.model.layers.{}.feed_forward.router."), + "use_packed_local_experts": True, + "local_experts.linear_fc1": PackNameRemapping( + "language_model.model.layers.{}.feed_forward.experts.gate_up_proj", + {"layer_type": "linear_fc1"}, + ), + "local_experts.linear_fc2": PackNameRemapping( + "language_model.model.layers.{}.feed_forward.experts.down_proj", + {"layer_type": "linear_fc2"}, + ), + "final_layernorm": NameRemapping("language_model.model.norm."), + "output_layer": NameRemapping("language_model.lm_head."), } medusa_llama_causal_lm_export: dict[str, CustomModuleMapping] = { # MedusaForCausalLM support - "lm_head": CustomModuleMapping( - "name_remapping", "medusa_heads.{}.1." + "lm_head": NameRemapping( + "medusa_heads.{}.1." ), # TODO: lm_head is hardcoded to .1 as currently only support using 1 layer in medusa head # needs a fix - "linear": CustomModuleMapping("name_remapping", "medusa_heads.{}.{}.linear."), + "linear": NameRemapping("medusa_heads.{}.{}.linear."), } eagle_llama_causal_lm_export: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "embed_tokens."), - "enorm": CustomModuleMapping("name_remapping", "enorm."), - "hnorm": CustomModuleMapping("name_remapping", "hnorm."), - "fc": CustomModuleMapping("name_remapping", "fc."), - "input_layernorm": CustomModuleMapping("name_remapping", "layers.{}.input_layernorm."), - "linear_qkv": CustomModuleMapping("qkv_slicing", "layers.{}.self_attn."), - "linear_proj": CustomModuleMapping("name_remapping", "layers.{}.self_attn.o_proj."), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "layers.{}.post_attention_layernorm." - ), - "linear_fc1": CustomModuleMapping("gated_mlp_slicing", "layers.{}.mlp."), - "linear_fc2": CustomModuleMapping("name_remapping", "layers.{}.mlp.down_proj."), - "final_layernorm": CustomModuleMapping("name_remapping", "norm."), - "d2t": CustomModuleMapping("name_remapping", "d2t"), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + "word_embeddings": NameRemapping("embed_tokens."), + "enorm": NameRemapping("enorm."), + "hnorm": NameRemapping("hnorm."), + "fc": NameRemapping("fc."), + "input_layernorm": NameRemapping("layers.{}.input_layernorm."), + "linear_qkv": QKVSlicing("layers.{}.self_attn."), + "linear_proj": NameRemapping("layers.{}.self_attn.o_proj."), + "pre_mlp_layernorm": NameRemapping("layers.{}.post_attention_layernorm."), + "linear_fc1": GatedMLPSlicing("layers.{}.mlp."), + "linear_fc2": NameRemapping("layers.{}.mlp.down_proj."), + "final_layernorm": NameRemapping("norm."), + "d2t": NameRemapping("d2t"), + "output_layer": NameRemapping("lm_head."), } eagle3_llama_causal_lm_export: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "embed_tokens."), - "enorm": CustomModuleMapping("name_remapping", "midlayer.input_layernorm."), - "fc": CustomModuleMapping("name_remapping", "fc."), - "input_layernorm": CustomModuleMapping("name_remapping", "midlayer.hidden_norm."), - "linear_qkv": CustomModuleMapping("qkv_slicing", "midlayer.self_attn."), - "linear_proj": CustomModuleMapping("name_remapping", "midlayer.self_attn.o_proj."), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "midlayer.post_attention_layernorm." - ), - "linear_fc1": CustomModuleMapping("gated_mlp_slicing", "midlayer.mlp."), - "linear_fc2": CustomModuleMapping("name_remapping", "midlayer.mlp.down_proj."), - "final_layernorm": CustomModuleMapping("name_remapping", "norm."), - "d2t": CustomModuleMapping("name_remapping", "d2t"), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + "word_embeddings": NameRemapping("embed_tokens."), + "enorm": NameRemapping("midlayer.input_layernorm."), + "fc": NameRemapping("fc."), + "input_layernorm": NameRemapping("midlayer.hidden_norm."), + "linear_qkv": QKVSlicing("midlayer.self_attn."), + "linear_proj": NameRemapping("midlayer.self_attn.o_proj."), + "pre_mlp_layernorm": NameRemapping("midlayer.post_attention_layernorm."), + "linear_fc1": GatedMLPSlicing("midlayer.mlp."), + "linear_fc2": NameRemapping("midlayer.mlp.down_proj."), + "final_layernorm": NameRemapping("norm."), + "d2t": NameRemapping("d2t"), + "output_layer": NameRemapping("lm_head."), } llama_causal_lm_import: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens.", COL_PARALLEL), - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_qkv": CustomModuleMapping("qkv_merging", "model.layers.{}.self_attn.", COL_PARALLEL), - "linear_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.o_proj.", ROW_PARALLEL + "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE), + "linear_qkv": QKVMerging("model.layers.{}.self_attn.", COL_TP), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE), + "linear_fc1": GatedMLPMerging("model.layers.{}.mlp.", COL_TP), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj.", ROW_TP), + "final_layernorm": NameRemapping("model.norm.", REPLICATE), + "output_layer": NameRemapping("lm_head.", COL_TP), +} + +llama4_causal_lm_import: dict[str, CustomModuleMapping | bool] = { + "word_embeddings": NameRemapping("language_model.model.embed_tokens.", COL_TP), + "input_layernorm": NameRemapping("language_model.model.layers.{}.input_layernorm.", REPLICATE), + "linear_qkv": QKVMerging("language_model.model.layers.{}.self_attn.", COL_TP), + "linear_proj": NameRemapping("language_model.model.layers.{}.self_attn.o_proj.", ROW_TP), + "pre_mlp_layernorm": NameRemapping( + "language_model.model.layers.{}.post_attention_layernorm.", REPLICATE ), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." + "shared_experts.linear_fc1": GatedMLPMerging( + "language_model.model.layers.{}.feed_forward.shared_expert.", COL_TP ), - "linear_fc1": CustomModuleMapping("gated_mlp_merging", "model.layers.{}.mlp.", COL_PARALLEL), - "linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.down_proj.", ROW_PARALLEL + "shared_experts.linear_fc2": NameRemapping( + "language_model.model.layers.{}.feed_forward.shared_expert.down_proj.", ROW_TP ), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head.", COL_PARALLEL), + "router": NameRemapping("language_model.model.layers.{}.feed_forward.router.", REPLICATE), + "use_packed_local_experts": True, + "local_experts.linear_fc1_etp": UnpackNameRemapping( + "language_model.model.layers.{}.feed_forward.experts.gate_up_proj", + PACK_COL_ETP | {"layer_type": "linear_fc1"}, + ), + "local_experts.linear_fc2_etp": UnpackNameRemapping( + "language_model.model.layers.{}.feed_forward.experts.down_proj", + PACK_ROW_ETP | {"layer_type": "linear_fc2"}, + ), + "local_experts.linear_fc1_ep": UnpackNameRemapping( + "language_model.model.layers.{}.feed_forward.experts.gate_up_proj", + PACK_EP | {"layer_type": "linear_fc1"}, + ), + "local_experts.linear_fc2_ep": UnpackNameRemapping( + "language_model.model.layers.{}.feed_forward.experts.down_proj", + PACK_EP | {"layer_type": "linear_fc2"}, + ), + "final_layernorm": NameRemapping("language_model.model.norm.", REPLICATE), + "output_layer": NameRemapping("language_model.lm_head.", COL_TP), } diff --git a/modelopt/torch/export/plugins/mcore_nemotron.py b/modelopt/torch/export/plugins/mcore_nemotron.py index 53c019e02..752826bbc 100644 --- a/modelopt/torch/export/plugins/mcore_nemotron.py +++ b/modelopt/torch/export/plugins/mcore_nemotron.py @@ -16,82 +16,75 @@ """Custom mapping from Nemotron Hugging Face models to Megatron Core models.""" -from .mcore_custom import COL_PARALLEL, ROW_PARALLEL, CustomModuleMapping +from .mcore_custom import ( + COL_TP, + REPLICATE, + ROW_TP, + CustomModuleMapping, + NameRemapping, + QKVMerging, + QKVSlicing, +) # Example on adding a new CausalLM. nemotron_causal_lm_export: dict[str, CustomModuleMapping] = { # NemotronForCausalLM is using square-relu where no gated handle is needed. - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens."), - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_qkv": CustomModuleMapping( - "qkv_slicing", - "model.layers.{}.self_attn.", - ), - "linear_proj": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.o_proj."), - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." - ), + "word_embeddings": NameRemapping("model.embed_tokens."), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."), + "linear_qkv": QKVSlicing("model.layers.{}.self_attn."), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."), # NemotronForCausalLM is using square-relu where no gated handle is needed. - "linear_fc1": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.up_proj."), - "linear_fc2": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.down_proj."), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + "linear_fc1": NameRemapping("model.layers.{}.mlp.up_proj."), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."), + "final_layernorm": NameRemapping("model.norm."), + "output_layer": NameRemapping("lm_head."), } nemotron_h_causal_lm_import: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "backbone.embeddings.", COL_PARALLEL), - "final_norm": CustomModuleMapping("name_remapping", "backbone.norm_f."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head.", COL_PARALLEL), + "word_embeddings": NameRemapping("backbone.embeddings.", COL_TP), + "final_norm": NameRemapping("backbone.norm_f.", REPLICATE), + "output_layer": NameRemapping("lm_head.", COL_TP), # Mamba - "norm": CustomModuleMapping("name_remapping", "backbone.layers.{}.norm."), - "mixer_norm": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.norm."), - "A_log": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.A_log"), - "D": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.D"), - "dt_bias": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.dt_bias"), - "conv1d": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.conv1d."), - "in_proj": CustomModuleMapping( - "name_remapping", "backbone.layers.{}.mixer.in_proj.", COL_PARALLEL - ), - "out_proj": CustomModuleMapping( - "name_remapping", "backbone.layers.{}.mixer.out_proj.", ROW_PARALLEL - ), + "norm": NameRemapping("backbone.layers.{}.norm.", REPLICATE), + "mixer_norm": NameRemapping("backbone.layers.{}.mixer.norm.", REPLICATE), + "A_log": NameRemapping("backbone.layers.{}.mixer.A_log", REPLICATE), + "D": NameRemapping("backbone.layers.{}.mixer.D", REPLICATE), + "dt_bias": NameRemapping("backbone.layers.{}.mixer.dt_bias", REPLICATE), + "conv1d": NameRemapping("backbone.layers.{}.mixer.conv1d.", REPLICATE), + "in_proj": NameRemapping("backbone.layers.{}.mixer.in_proj.", COL_TP), + "out_proj": NameRemapping("backbone.layers.{}.mixer.out_proj.", ROW_TP), # Attention - "input_layernorm": CustomModuleMapping("name_remapping", "backbone.layers.{}.norm."), - "linear_qkv": CustomModuleMapping("qkv_merging", "backbone.layers.{}.mixer.", COL_PARALLEL), - "linear_proj": CustomModuleMapping( - "name_remapping", "backbone.layers.{}.mixer.o_proj.", ROW_PARALLEL - ), + "input_layernorm": NameRemapping("backbone.layers.{}.norm.", REPLICATE), + "linear_qkv": QKVMerging("backbone.layers.{}.mixer.", COL_TP), + "linear_proj": NameRemapping("backbone.layers.{}.mixer.o_proj.", ROW_TP), # MLP - "pre_mlp_layernorm": CustomModuleMapping("name_remapping", "backbone.layers.{}.norm."), - "linear_fc1": CustomModuleMapping( - "name_remapping", "backbone.layers.{}.mixer.up_proj.", COL_PARALLEL - ), - "linear_fc2": CustomModuleMapping( - "name_remapping", "backbone.layers.{}.mixer.down_proj.", ROW_PARALLEL - ), + "pre_mlp_layernorm": NameRemapping("backbone.layers.{}.norm.", REPLICATE), + "linear_fc1": NameRemapping("backbone.layers.{}.mixer.up_proj.", COL_TP), + "linear_fc2": NameRemapping("backbone.layers.{}.mixer.down_proj.", ROW_TP), } nemotron_h_causal_lm_export: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "backbone.embeddings."), - "final_norm": CustomModuleMapping("name_remapping", "backbone.norm_f."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + "word_embeddings": NameRemapping("backbone.embeddings."), + "final_norm": NameRemapping("backbone.norm_f."), + "output_layer": NameRemapping("lm_head."), # Mamba - "norm": CustomModuleMapping("name_remapping", "backbone.layers.{}.norm."), - "mixer_norm": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.norm."), - "A_log": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.A_log"), - "D": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.D"), - "dt_bias": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.dt_bias"), - "conv1d": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.conv1d."), - "in_proj": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.in_proj."), - "out_proj": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.out_proj."), + "norm": NameRemapping("backbone.layers.{}.norm."), + "mixer_norm": NameRemapping("backbone.layers.{}.mixer.norm."), + "A_log": NameRemapping("backbone.layers.{}.mixer.A_log"), + "D": NameRemapping("backbone.layers.{}.mixer.D"), + "dt_bias": NameRemapping("backbone.layers.{}.mixer.dt_bias"), + "conv1d": NameRemapping("backbone.layers.{}.mixer.conv1d."), + "in_proj": NameRemapping("backbone.layers.{}.mixer.in_proj."), + "out_proj": NameRemapping("backbone.layers.{}.mixer.out_proj."), # Attention - "input_layernorm": CustomModuleMapping("name_remapping", "backbone.layers.{}.norm."), - "linear_qkv": CustomModuleMapping("qkv_slicing", "backbone.layers.{}.mixer."), - "linear_proj": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.o_proj."), + "input_layernorm": NameRemapping("backbone.layers.{}.norm."), + "linear_qkv": QKVSlicing("backbone.layers.{}.mixer."), + "linear_proj": NameRemapping("backbone.layers.{}.mixer.o_proj."), # MLP - "pre_mlp_layernorm": CustomModuleMapping("name_remapping", "backbone.layers.{}.norm."), - "linear_fc1": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.up_proj."), - "linear_fc2": CustomModuleMapping("name_remapping", "backbone.layers.{}.mixer.down_proj."), + "pre_mlp_layernorm": NameRemapping("backbone.layers.{}.norm."), + "linear_fc1": NameRemapping("backbone.layers.{}.mixer.up_proj."), + "linear_fc2": NameRemapping("backbone.layers.{}.mixer.down_proj."), } diff --git a/modelopt/torch/export/plugins/mcore_qwen.py b/modelopt/torch/export/plugins/mcore_qwen.py index 6a75a4709..b266701b8 100644 --- a/modelopt/torch/export/plugins/mcore_qwen.py +++ b/modelopt/torch/export/plugins/mcore_qwen.py @@ -15,63 +15,57 @@ """Custom mapping from Qwen Hugging Face models to Megatron Core models.""" -from .mcore_custom import COL_PARALLEL, ROW_PARALLEL, CustomModuleMapping +from .mcore_custom import ( + COL_ETP, + COL_TP, + REPLICATE, + ROW_ETP, + ROW_TP, + CustomModuleMapping, + GatedMLPMerging, + GatedMLPSlicing, + NameRemapping, + QKVMerging, + QKVSlicing, +) qwen3_causal_lm_import: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens.", COL_PARALLEL), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head.", COL_PARALLEL), + "word_embeddings": NameRemapping("model.embed_tokens.", COL_TP), + "final_layernorm": NameRemapping("model.norm.", REPLICATE), + "output_layer": NameRemapping("lm_head.", COL_TP), # Attention - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_qkv": CustomModuleMapping("qkv_merging", "model.layers.{}.self_attn.", COL_PARALLEL), - "linear_proj": CustomModuleMapping( - "name_remapping", "model.layers.{}.self_attn.o_proj.", ROW_PARALLEL - ), - "q_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.q_norm."), - "k_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.k_norm."), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm.", REPLICATE), + "linear_qkv": QKVMerging("model.layers.{}.self_attn.", COL_TP), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj.", ROW_TP), + "q_layernorm": NameRemapping("model.layers.{}.self_attn.q_norm.", REPLICATE), + "k_layernorm": NameRemapping("model.layers.{}.self_attn.k_norm.", REPLICATE), # MLP - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." - ), - "linear_fc1": CustomModuleMapping("gated_mlp_merging", "model.layers.{}.mlp.", COL_PARALLEL), - "linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.down_proj.", ROW_PARALLEL - ), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm.", REPLICATE), + "linear_fc1": GatedMLPMerging("model.layers.{}.mlp.", COL_TP), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj.", ROW_TP), # MoE - "router": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.gate."), - "local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_merging", "model.layers.{}.mlp.experts.{}.", COL_PARALLEL - ), - "local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", - "model.layers.{}.mlp.experts.{}.down_proj.", - ROW_PARALLEL, - ), + "router": NameRemapping("model.layers.{}.mlp.gate.", REPLICATE), + "local_experts.linear_fc1": GatedMLPMerging("model.layers.{}.mlp.experts.{}.", COL_ETP), + "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj.", ROW_ETP), } qwen3_causal_lm_export: dict[str, CustomModuleMapping] = { - "word_embeddings": CustomModuleMapping("name_remapping", "model.embed_tokens."), - "final_layernorm": CustomModuleMapping("name_remapping", "model.norm."), - "output_layer": CustomModuleMapping("name_remapping", "lm_head."), + "word_embeddings": NameRemapping("model.embed_tokens."), + "final_layernorm": NameRemapping("model.norm."), + "output_layer": NameRemapping("lm_head."), # Attention - "input_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.input_layernorm."), - "linear_qkv": CustomModuleMapping("qkv_slicing", "model.layers.{}.self_attn."), - "linear_proj": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.o_proj."), - "q_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.q_norm."), - "k_layernorm": CustomModuleMapping("name_remapping", "model.layers.{}.self_attn.k_norm."), + "input_layernorm": NameRemapping("model.layers.{}.input_layernorm."), + "linear_qkv": QKVSlicing("model.layers.{}.self_attn."), + "linear_proj": NameRemapping("model.layers.{}.self_attn.o_proj."), + "q_layernorm": NameRemapping("model.layers.{}.self_attn.q_norm."), + "k_layernorm": NameRemapping("model.layers.{}.self_attn.k_norm."), # MLP - "pre_mlp_layernorm": CustomModuleMapping( - "name_remapping", "model.layers.{}.post_attention_layernorm." - ), - "linear_fc1": CustomModuleMapping("gated_mlp_slicing", "model.layers.{}.mlp."), - "linear_fc2": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.down_proj."), + "pre_mlp_layernorm": NameRemapping("model.layers.{}.post_attention_layernorm."), + "linear_fc1": GatedMLPSlicing("model.layers.{}.mlp."), + "linear_fc2": NameRemapping("model.layers.{}.mlp.down_proj."), # MoE - "router": CustomModuleMapping("name_remapping", "model.layers.{}.mlp.gate."), - "local_experts.linear_fc1": CustomModuleMapping( - "gated_mlp_slicing", "model.layers.{}.mlp.experts.{}." - ), - "local_experts.linear_fc2": CustomModuleMapping( - "name_remapping", "model.layers.{}.mlp.experts.{}.down_proj." - ), + "router": NameRemapping("model.layers.{}.mlp.gate."), + "local_experts.linear_fc1": GatedMLPSlicing("model.layers.{}.mlp.experts.{}."), + "local_experts.linear_fc2": NameRemapping("model.layers.{}.mlp.experts.{}.down_proj."), } diff --git a/modelopt/torch/export/plugins/megatron_importer.py b/modelopt/torch/export/plugins/megatron_importer.py new file mode 100644 index 000000000..8eea0dfe1 --- /dev/null +++ b/modelopt/torch/export/plugins/megatron_importer.py @@ -0,0 +1,529 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Code that export quantized Megatron Core models for deployment.""" + +import tempfile +from pathlib import Path + +import torch +import torch.distributed +from huggingface_hub import snapshot_download +from tqdm import tqdm + +from modelopt.torch.utils import import_plugin + +from .mcore_common import all_mcore_hf_import_mapping +from .mcore_custom import CustomModuleMapping, ParallelConfig, get_safetensor + +with import_plugin("transformers", verbose=False): + import transformers + +has_mcore = False +with import_plugin("megatron"): + from megatron.core.parallel_state import ( + get_expert_model_parallel_world_size, + get_tensor_model_parallel_world_size, + ) + from megatron.core.ssm.mamba_layer import MambaLayer + from megatron.core.transformer.identity_op import IdentityOp + from megatron.core.transformer.torch_norm import L2Norm + from megatron.core.transformer.transformer_layer import TransformerLayer + + has_mcore = True + + +class GPTModelImporter: + """Megatron Core GPTModel HuggingFace Importer. + + The Importer is created by `import_mcore_gpt_from_hf` to host attributes + and methods that import a Megatron Core GPTModel from a supported Hugging + Face model. + + Args: + model: The Megatron Core GPTModel instance. + pretrained_model_name_or_path: Can be either: the *model id* of a + pretrained model hosted inside a model repo on huggingface.co; or + a *directory* containing model weights saved using + [`~PreTrainedModel.save_pretrained`], e.g., `./my_model_directory/`. + dtype: The weights data type to export the unquantized layers. + """ + + weight_scale_name: str = "weight_scale_inv" + + def __init__( + self, + model: torch.nn.Module, + pretrained_model_name_or_path: str, + workspace_dir: str | None = None, + dtype=torch.bfloat16, + dequantize: bool = True, + trust_remote_code: bool = True, + verbose: bool = False, + ): + """Create a GPTModel importer instance.""" + self._hf_config = transformers.AutoConfig.from_pretrained( + pretrained_model_name_or_path, trust_remote_code=trust_remote_code + ) + pretrained_model_path = Path(pretrained_model_name_or_path) + if not pretrained_model_path.is_dir(): + if workspace_dir is None: + workspace_dir = tempfile.gettempdir() + pretrained_model_path = workspace_dir + "/" + pretrained_model_name_or_path + if torch.distributed.get_rank() == 0: + snapshot_download( + repo_id=pretrained_model_name_or_path, + local_dir=pretrained_model_path, + ) + torch.distributed.barrier() + self.arch = self._hf_config.architectures[0] + self.all_rules = self._populate_rule_book() + self.rules = self.all_rules[self.arch] + self.model = model + self.pretrained_model_path = pretrained_model_path + self.dtype = dtype + self.dequantize = dequantize + self.verbose = verbose + self.disable_tqdm = torch.distributed.get_rank() > 0 or verbose + + def _populate_rule_book(self): + """The rule book maps each state_dict key to a Callable.""" + all_rules = {} + + def _custom_mapping_to_lambda(mapping): + method_map = { + "name_remapping": self._name_remapping, + "qkv_merging": self._qkv_merging, + "gated_mlp_merging": self._gated_mlp_merging, + "unpack_name_remapping": self._unpack_name_remapping, + } + func = method_map[mapping.func_name] + prefix = mapping.target_name_or_prefix + func_kwargs = mapping.func_kwargs + return lambda m, *args: func(m, prefix.format(*args), **func_kwargs) + + for arch, mappings in all_mcore_hf_import_mapping.items(): + all_rules[arch] = { + k: _custom_mapping_to_lambda(v) if isinstance(v, CustomModuleMapping) else v + for (k, v) in mappings.items() + if isinstance(v, (CustomModuleMapping, bool)) + } + + return all_rules + + def _get_safetensor(self, key, parallel_config: ParallelConfig | None = None): + return get_safetensor( + self.pretrained_model_path, key, parallel_config, dequantize=self.dequantize + ) + + def _name_remapping( + self, + module, + prefix, + mapping={}, + parallel_config: ParallelConfig | None = None, + ): + if isinstance(module, torch.Tensor): + module.data.copy_(self._get_safetensor(prefix)) + return + + weight = module.state_dict().get("weight", None) + weight_scale = module.state_dict().get("weight_quantizer._scale", None) + + state_dict = {} + + if weight is None: + raise ValueError(f"{module!s} does not contain weight!") + else: + tensor = self._get_safetensor(prefix + "weight", parallel_config=parallel_config) + + if weight_scale is not None: + scale_name = prefix + self.weight_scale_name + if weight_scale.ndim > 0: + scale = self._get_safetensor(scale_name, parallel_config=parallel_config) + else: + scale = self._get_safetensor(scale_name) + scale = scale.to(weight_scale.dtype).to(device=weight_scale.device) + state_dict["weight_quantizer._scale"] = scale + + if tensor.shape != weight.shape: + expanded_tensor = torch.zeros(weight.shape, dtype=tensor.dtype) + expanded_tensor[: tensor.shape[0], : tensor.shape[1]] = tensor + tensor = expanded_tensor + state_dict["weight"] = tensor.view(dtype=weight.dtype).to(device=weight.device) + else: + state_dict["weight"] = tensor.to(dtype=self.dtype).to(device=weight.device) + + # Handle the rest of the state_dict. + for key, val in module.state_dict().items(): + if key in {"weight", "weight_quantizer._scale"}: + continue + elif "extra_state" in key: + state_dict[key] = val + else: + source_key = mapping.get(key, key) + tensor = self._get_safetensor(prefix + source_key, parallel_config=parallel_config) + state_dict[key] = tensor.to(dtype=self.dtype).to(device=val.device) + + module.load_state_dict(state_dict) + + def _gated_mlp_merging( + self, + module, + prefix, + gate_proj_name="gate_proj", + up_proj_name="up_proj", + parallel_config: ParallelConfig | None = None, + ): + weight = module.state_dict().get("weight", None) + weight_scale = module.state_dict().get("weight_quantizer._scale", None) + + state_dict = {} + + if weight is None: + raise ValueError(f"{module!s} does not contain weight!") + else: + gate_proj = self._get_safetensor( + prefix + gate_proj_name + ".weight", parallel_config=parallel_config + ) + up_proj = self._get_safetensor( + prefix + up_proj_name + ".weight", parallel_config=parallel_config + ) + tensor = torch.cat((gate_proj, up_proj), dim=0) + + if weight_scale is not None: + gate_scale_name = prefix + gate_proj_name + "." + self.weight_scale_name + up_scale_name = prefix + up_proj_name + "." + self.weight_scale_name + if weight_scale.ndim > 0: + gate_scale = self._get_safetensor(gate_scale_name, parallel_config=parallel_config) + up_scale = self._get_safetensor(up_scale_name, parallel_config=parallel_config) + scale = torch.cat((gate_scale, up_scale), dim=0) + else: + scale = self._get_safetensor(gate_scale_name) + # If source model is per tensor, compute a per tensor scale with max. + if scale.ndim > 0: + scale = scale.max(dim=0).max(dim=0) + state_dict["weight_quantizer._scale"] = scale.to(weight_scale.dtype).to( + device=weight_scale.device + ) + state_dict["weight"] = tensor.view(weight.dtype).to(device=weight.device) + else: + state_dict["weight"] = tensor.to(self.dtype).to(device=weight.device) + + module.load_state_dict(state_dict) + + def _qkv_merging( + self, + module, + prefix, + q_proj_name="q_proj", + k_proj_name="k_proj", + v_proj_name="v_proj", + parallel_config: ParallelConfig | None = None, + ): + config = module.config + hidden_size = config.hidden_size + num_query_groups = config.num_query_groups + head_num = config.num_attention_heads + head_size = config.kv_channels + + if parallel_config is not None: + tp_size = get_tensor_model_parallel_world_size() + assert head_num % tp_size == 0 + assert num_query_groups % tp_size == 0 + head_num = head_num // tp_size + num_query_groups = num_query_groups // tp_size + + heads_per_group = head_num // num_query_groups + qkv_total_dim = head_num + 2 * num_query_groups + q_slice = torch.cat( + [ + torch.arange((heads_per_group + 2) * i, (heads_per_group + 2) * i + heads_per_group) + for i in range(num_query_groups) + ] + ) + k_slice = torch.arange(heads_per_group, qkv_total_dim, (heads_per_group + 2)) + v_slice = torch.arange(heads_per_group + 1, qkv_total_dim, (heads_per_group + 2)) + + state_dict = {} + + weight = module.state_dict().get("weight", None) + weight_scale = module.state_dict().get("weight_quantizer._scale", None) + + if weight is None: + raise ValueError(f"{module!s} does not contain weight!") + + if weight_scale is not None: + q_scale_name = prefix + q_proj_name + "." + self.weight_scale_name + k_scale_name = prefix + k_proj_name + "." + self.weight_scale_name + v_scale_name = prefix + v_proj_name + "." + self.weight_scale_name + + if weight_scale.ndim > 0: + q_scale = self._get_safetensor(q_scale_name, parallel_config=parallel_config) + k_scale = self._get_safetensor(k_scale_name, parallel_config=parallel_config) + v_scale = self._get_safetensor(v_scale_name, parallel_config=parallel_config) + weight_scale[q_slice] = q_scale.to(weight_scale.dtype).to( + device=weight_scale.device + ) + weight_scale[k_slice] = k_scale.to(weight_scale.dtype).to( + device=weight_scale.device + ) + weight_scale[v_slice] = v_scale.to(weight_scale.dtype).to( + device=weight_scale.device + ) + else: + q_scale = self._get_safetensor(q_scale_name) + weight_scale = q_scale.to(weight_scale.dtype).to(device=weight_scale.device) + state_dict["weight_quantizer._scale"] = weight_scale + + q_proj = self._get_safetensor( + prefix + q_proj_name + ".weight", parallel_config=parallel_config + ) + k_proj = self._get_safetensor( + prefix + k_proj_name + ".weight", parallel_config=parallel_config + ) + v_proj = self._get_safetensor( + prefix + v_proj_name + ".weight", parallel_config=parallel_config + ) + q_proj = q_proj.reshape(-1, head_size, hidden_size) + k_proj = k_proj.reshape(-1, head_size, hidden_size) + v_proj = v_proj.reshape(-1, head_size, hidden_size) + tensor = weight.detach().clone().reshape([qkv_total_dim, head_size, hidden_size]) + + if weight_scale is not None: + tensor[q_slice] = q_proj.view(dtype=tensor.dtype).to(device=tensor.device) + tensor[k_slice] = k_proj.view(dtype=tensor.dtype).to(device=tensor.device) + tensor[v_slice] = v_proj.view(dtype=tensor.dtype).to(device=tensor.device) + else: + tensor[q_slice] = q_proj.to(dtype=tensor.dtype).to(device=tensor.device) + tensor[k_slice] = k_proj.to(dtype=tensor.dtype).to(device=tensor.device) + tensor[v_slice] = v_proj.to(dtype=tensor.dtype).to(device=tensor.device) + + state_dict["weight"] = tensor.reshape(-1, hidden_size) + + module.load_state_dict(state_dict) + + def _unpack_name_remapping( + self, + module, + prefix, + layer_type: str, + parallel_config: ParallelConfig | None = None, + ): + tensor = self._get_safetensor(prefix, parallel_config=parallel_config) + + for idx, sub_module in enumerate(module.children()): + state_dict = {} + linear_module = getattr(sub_module, layer_type) + weight = linear_module.state_dict().get("weight", None) + sub_tensor = tensor[idx] + if weight is None: + raise ValueError(f"{linear_module!s} does not contain weight!") + # TODO (yueshen): Handle weight_scale case + else: + # Transpose to match huggingface format with Mcore format + sub_tensor = sub_tensor.transpose(-1, -2) + state_dict["weight"] = sub_tensor.to(dtype=self.dtype).to(device=weight.device) + + for key, val in linear_module.state_dict().items(): + if key in {"weight", "weight_quantizer._scale"}: + continue + elif "extra_state" in key: + state_dict[key] = val + + linear_module.load_state_dict(state_dict) + + def _import_state_dict(self): + model = self.model + + layer_pbar = tqdm(model.decoder.layers, disable=self.disable_tqdm) + + # Embedding + if hasattr(model, "embedding"): + layer_pbar.set_description("Importing word embedding") + self.rules["word_embeddings"](model.embedding.word_embeddings) + + # Decoder layers + for layer in layer_pbar: + layer_id = layer.layer_number - 1 + + if isinstance(layer, MambaLayer): + if not isinstance(layer.norm, IdentityOp): + self.rules["norm"](layer.norm, layer_id) + + self.rules["mixer_norm"](layer.mixer.norm, layer_id) + self.rules["A_log"](layer.mixer.A_log, layer_id) + self.rules["D"](layer.mixer.D, layer_id) + self.rules["dt_bias"](layer.mixer.dt_bias, layer_id) + + self.rules["conv1d"](layer.mixer.conv1d, layer_id) + self.rules["in_proj"](layer.mixer.in_proj, layer_id) + self.rules["out_proj"](layer.mixer.out_proj, layer_id) + + elif isinstance(layer, TransformerLayer): + if not isinstance(layer.input_layernorm, IdentityOp): + self.rules["input_layernorm"](layer.input_layernorm, layer_id) + + attention = layer.self_attention + if not isinstance(attention, IdentityOp): + if "MLASelfAttention" in str(type(attention)): + if hasattr(attention, "linear_q_proj"): + layer_pbar.set_description("Importing MLA (without q LoRA)") + self.rules["linear_q_proj"](attention.linear_q_proj, layer_id) + else: + layer_pbar.set_description("Importing MLA (with q LoRA)") + self.rules["linear_q_down_proj"](attention.linear_q_down_proj, layer_id) + self.rules["linear_q_layernorm"](attention.q_layernorm, layer_id) + self.rules["linear_q_up_proj"](attention.linear_q_up_proj, layer_id) + self.rules["linear_kv_down_proj"](attention.linear_kv_down_proj, layer_id) + self.rules["linear_kv_layernorm"](attention.kv_layernorm, layer_id) + self.rules["linear_kv_up_proj"](attention.linear_kv_up_proj, layer_id) + self.rules["linear_proj"](attention.linear_proj, layer_id) + else: + layer_pbar.set_description("Importing GQA/MHA") + if attention.q_layernorm is not None and not isinstance( + attention.q_layernorm, (IdentityOp, L2Norm) + ): + self.rules["q_layernorm"](attention.q_layernorm, layer_id) + self.rules["k_layernorm"](attention.k_layernorm, layer_id) + self.rules["linear_qkv"](attention.linear_qkv, layer_id) + self.rules["linear_proj"](attention.linear_proj, layer_id) + + if not isinstance(layer.pre_mlp_layernorm, IdentityOp): + self.rules["pre_mlp_layernorm"](layer.pre_mlp_layernorm, layer_id) + + if not isinstance(layer.mlp, IdentityOp): + if "MoE" in str(type(layer.mlp)): + layer_pbar.set_description("Importing MoE") + self.rules["router"](layer.mlp.router, layer_id) + if ( + hasattr(layer.mlp, "shared_experts") + and layer.mlp.shared_experts is not None + ): + layer_pbar.set_description("Importing MoE shared experts") + fc1 = layer.mlp.shared_experts.linear_fc1 + fc2 = layer.mlp.shared_experts.linear_fc2 + self.rules["shared_experts.linear_fc1"](fc1, layer_id) + self.rules["shared_experts.linear_fc2"](fc2, layer_id) + if not self.rules.get("use_packed_local_experts", False): + for local_expert_id, expert in tqdm( + enumerate(layer.mlp.experts.local_experts), + desc="Importing MoE local experts", + leave=False, + disable=self.disable_tqdm, + ): + expert_id = layer.mlp.local_expert_indices[local_expert_id] + fc1 = expert.linear_fc1 + fc2 = expert.linear_fc2 + self.rules["local_experts.linear_fc1"](fc1, layer_id, expert_id) + self.rules["local_experts.linear_fc2"](fc2, layer_id, expert_id) + # We only support either EP or ETP for now + elif get_expert_model_parallel_world_size() > 1: + # EP supports for packed MoE + self.rules["local_experts.linear_fc1_ep"]( + layer.mlp.experts.local_experts, layer_id + ) + self.rules["local_experts.linear_fc2_ep"]( + layer.mlp.experts.local_experts, layer_id + ) + else: + # ETP supports for packed MoE + self.rules["local_experts.linear_fc1_etp"]( + layer.mlp.experts.local_experts, layer_id + ) + self.rules["local_experts.linear_fc2_etp"]( + layer.mlp.experts.local_experts, layer_id + ) + else: + layer_pbar.set_description("Importing MLP") + self.rules["linear_fc1"](layer.mlp.linear_fc1, layer_id) + self.rules["linear_fc2"](layer.mlp.linear_fc2, layer_id) + + if self.verbose: + print( + "{:3}/{:3} completes importing layer {:3}.".format( + torch.distributed.get_rank(), torch.distributed.get_world_size(), layer_id + ), + flush=True, + ) + + # Final layernorm + if hasattr(model.decoder, "final_layernorm") and model.decoder.final_layernorm: + self.rules["final_layernorm"](model.decoder.final_layernorm) + + if hasattr(model.decoder, "final_norm") and model.decoder.final_norm: + self.rules["final_norm"](model.decoder.final_norm) + + # Output layer + if hasattr(model, "output_layer") and not model.share_embeddings_and_output_weights: + self.rules["output_layer"](model.output_layer) + + # MTP + if hasattr(model, "mtp"): + # MTP is the last layer in DeepSeek V3/R1 + layer_id += 1 + for mtp in model.mtp: + self.rules["mtp.fc"](mtp.fc, layer_id) + self.rules["mtp.enorm"](mtp.enorm, layer_id) + self.rules["mtp.hnorm"](mtp.hnorm, layer_id) + self.rules["mtp.input_layernorm"](mtp.decoder.layers[0].input_layernorm, layer_id) + if hasattr(mtp.decoder.layers[0].self_attention, "linear_q_proj"): + self.rules["mtp.linear_q_proj"]( + mtp.decoder.layers[0].self_attention.linear_q_proj, layer_id + ) + else: + self.rules["mtp.linear_q_down_proj"]( + mtp.decoder.layers[0].self_attention.linear_q_down_proj, layer_id + ) + self.rules["mtp.linear_q_layernorm"]( + mtp.decoder.layers[0].self_attention.q_layernorm, layer_id + ) + self.rules["mtp.linear_q_up_proj"]( + mtp.decoder.layers[0].self_attention.linear_q_up_proj, layer_id + ) + self.rules["mtp.linear_kv_down_proj"]( + mtp.decoder.layers[0].self_attention.linear_kv_down_proj, layer_id + ) + self.rules["mtp.linear_kv_layernorm"]( + mtp.decoder.layers[0].self_attention.kv_layernorm, layer_id + ) + self.rules["mtp.linear_kv_up_proj"]( + mtp.decoder.layers[0].self_attention.linear_kv_up_proj, layer_id + ) + self.rules["mtp.linear_proj"]( + mtp.decoder.layers[0].self_attention.linear_proj, layer_id + ) + self.rules["mtp.pre_mlp_layernorm"]( + mtp.decoder.layers[0].pre_mlp_layernorm, layer_id + ) + self.rules["mtp.router"](mtp.decoder.layers[0].mlp.router, layer_id) + self.rules["mtp.shared_experts.linear_fc1"]( + mtp.decoder.layers[0].mlp.shared_experts.linear_fc1, layer_id + ) + self.rules["mtp.shared_experts.linear_fc2"]( + mtp.decoder.layers[0].mlp.shared_experts.linear_fc2, layer_id + ) + for expert_id, expert in tqdm( + enumerate(mtp.decoder.layers[0].mlp.experts.local_experts), + desc="Importing MoE local experts", + leave=False, + disable=self.disable_tqdm, + ): + self.rules["mtp.local_experts.linear_fc1"]( + expert.linear_fc1, layer_id, expert_id + ) + self.rules["mtp.local_experts.linear_fc2"]( + expert.linear_fc2, layer_id, expert_id + ) diff --git a/modelopt/torch/export/quant_utils.py b/modelopt/torch/export/quant_utils.py index 2f398054f..3d26c887f 100644 --- a/modelopt/torch/export/quant_utils.py +++ b/modelopt/torch/export/quant_utils.py @@ -979,6 +979,48 @@ def quantize_llama4_experts_for_hf_export(module: nn.Module): assert module.gate_up_proj_input_quantizer.is_enabled assert module.down_proj_input_quantizer.is_enabled + # Handle uncalibrated input quantizers that have None amax values + input_quantizers = [ + module.gate_up_proj_input_quantizer, + module.down_proj_input_quantizer, + ] + + # Only handle amax for non-dynamic quantizers + non_dynamic_quantizers = [q for q in input_quantizers if not getattr(q, "_dynamic", False)] + + if non_dynamic_quantizers: + # Find the maximum amax value from non-None quantizers + valid_amax_values = [ + quantizer.amax for quantizer in non_dynamic_quantizers if quantizer.amax is not None + ] + + device = module.gate_up_proj.device + + # If all quantizers have None amax, set a default value + if not valid_amax_values: + default_amax = torch.tensor(1.0, dtype=torch.float32, device=device) + warn( + "All input quantizers have None amax values. Setting default amax to 1.0. " + "This typically occurs when experts are not activated during calibration. " + "Consider increasing your calibration dataset size to ensure all experts are exercised." + ) + for quantizer in non_dynamic_quantizers: + if quantizer.amax is None: + quantizer.amax = default_amax.clone() + else: + # Set None amax values to the maximum of existing values + max_amax = torch.max(torch.stack(valid_amax_values)) + if max_amax.device != device: + max_amax = max_amax.to(device) + for quantizer in non_dynamic_quantizers: + if quantizer.amax is None: + warn( + f"Missing amax value for input quantizer. Setting it to {max_amax.item()} for export. " + "This typically occurs when certain experts are not activated during calibration. " + "Consider increasing your calibration dataset size to ensure all experts are exercised." + ) + quantizer.amax = max_amax.clone() + for weight_name in ["gate_up_proj", "down_proj"]: weight = getattr(module, weight_name) weight_quantizer = getattr(module, f"{weight_name}_weight_quantizer") @@ -1049,6 +1091,11 @@ def quantize_llama4_experts_for_hf_export(module: nn.Module): for input_name in ["gate_up_proj", "down_proj"]: input_quantizer = getattr(module, f"{input_name}_input_quantizer") + + # Skip processing for dynamic quantization since it doesn't have fixed amax + if getattr(input_quantizer, "_dynamic", False): + continue + if input_quantizer.num_bits == (4, 3): assert not input_quantizer.block_sizes diff --git a/modelopt/torch/export/unified_export_megatron.py b/modelopt/torch/export/unified_export_megatron.py index cf8dc8438..b88c480e5 100644 --- a/modelopt/torch/export/unified_export_megatron.py +++ b/modelopt/torch/export/unified_export_megatron.py @@ -27,6 +27,7 @@ from typing import Any from warnings import warn import torch +import torch.distributed import torch.nn as nn from huggingface_hub import snapshot_download from safetensors.torch import safe_open, save_file @@ -42,8 +43,9 @@ from .model_config import ( QUANTIZATION_FP8_PB_WO, QUANTIZATION_NVFP4, ) -from .plugins.mcore_common import all_mcore_hf_export_mapping, all_mcore_hf_import_mapping -from .plugins.mcore_custom import save_safetensors +from .plugins.mcore_common import all_mcore_hf_export_mapping +from .plugins.mcore_custom import CustomModuleMapping, save_safetensors +from .plugins.megatron_importer import GPTModelImporter from .quant_utils import ( get_activation_scaling_factor, get_kv_cache_dtype, @@ -57,19 +59,20 @@ from .quant_utils import ( with import_plugin("transformers", verbose=False): import transformers + from transformers import AutoProcessor has_mcore = False with import_plugin("megatron"): from megatron.core.models.gpt import GPTModel from megatron.core.models.mamba import MambaModel + from megatron.core.models.multimodal.llava_model import LLaVAModel from megatron.core.parallel_state import ( get_pipeline_model_parallel_rank, get_pipeline_model_parallel_world_size, - get_tensor_model_parallel_rank, - get_tensor_model_parallel_world_size, ) from megatron.core.ssm.mamba_layer import MambaLayer from megatron.core.transformer.identity_op import IdentityOp + from megatron.core.transformer.torch_norm import L2Norm from megatron.core.transformer.transformer_layer import TransformerLayer has_mcore = True @@ -185,7 +188,7 @@ class GPTModelExporter: trust_remote_code: bool = True, ): """Create a GPTModel exporter instance.""" - if not isinstance(model, GPTModel) and not isinstance(model, MambaModel): + if not isinstance(model, (GPTModel, MambaModel, LLaVAModel)): raise ValueError("Input to GPTModelExport must be a megatron.core.models.GPTModel!") self._state_dict = OrderedDict() @@ -202,11 +205,14 @@ class GPTModelExporter: self._hf_text_config.head_dim = model.config.kv_channels self._hf_text_config.num_attention_heads = model.config.num_attention_heads self._hf_text_config.num_key_value_heads = model.config.num_query_groups - self._hf_text_config.intermediate_size = model.config.ffn_hidden_size + self.is_multimodal = isinstance(model, LLaVAModel) + if not self.is_multimodal: + self._hf_text_config.intermediate_size = model.config.ffn_hidden_size self._hf_quant_config = None self._hf_extra_config = None self.export_extra_modules = export_extra_modules - self.model = model + self.is_multimodal = isinstance(model, LLaVAModel) + self.model = model.language_model if self.is_multimodal else model self.dtype = dtype self.trust_remote_code = trust_remote_code self.arch = self._hf_config.architectures[0] @@ -232,48 +238,51 @@ class GPTModelExporter: self.rules = self.all_rules[architectures] - # By default, we use Llama-3.1 - self._hf_extra_config = transformers.AutoConfig.from_pretrained( - "nvidia/Llama-3.1-8B-Instruct-FP8", trust_remote_code=self.trust_remote_code - ) + if torch.distributed.get_rank() == torch.distributed.get_world_size() - 1: + # By default, we use Llama-3.1 + self._hf_extra_config = transformers.AutoConfig.from_pretrained( + "nvidia/Llama-3.1-8B-Instruct-FP8", trust_remote_code=self.trust_remote_code + ) - eagle_config = { - "use_input_layernorm_in_first_layer": mode_cfg["config"][ - "use_input_layernorm_in_first_layer" - ], - "use_last_layernorm": mode_cfg["config"]["use_last_layernorm"], - "use_mtp_layernorm": mode_cfg["config"]["use_mtp_layernorm"], - "use_aux_hidden_state": mode_cfg["config"]["use_aux_hidden_state"], - "eagle_aux_hidden_state_layer_ids": model.eagle_aux_hidden_state_layer_ids, - } + eagle_config = { + "use_input_layernorm_in_first_layer": mode_cfg["config"][ + "use_input_layernorm_in_first_layer" + ], + "use_last_layernorm": mode_cfg["config"]["use_last_layernorm"], + "use_mtp_layernorm": mode_cfg["config"]["use_mtp_layernorm"], + "use_aux_hidden_state": mode_cfg["config"]["use_aux_hidden_state"], + "eagle_aux_hidden_state_layer_ids": model.eagle_aux_hidden_state_layer_ids, + } - eagle_config_update = { - "architectures": [architectures], - "head_dim": self._hf_text_config.head_dim, - "hidden_act": self._hf_text_config.hidden_act, - "hidden_size": self._hf_text_config.hidden_size, - "intermediate_size": self._hf_text_config.intermediate_size, - "max_position_embeddings": self._hf_text_config.max_position_embeddings, - "num_attention_heads": self._hf_text_config.num_attention_heads, - "num_key_value_heads": self._hf_text_config.num_key_value_heads, - "num_hidden_layers": mode_cfg["config"]["eagle_num_layers"], - "vocab_size": self._hf_text_config.vocab_size, - # Unset any special token ids given that the tokenizer can change here. - "bos_token_id": None, - "eos_token_id": None, - "pad_token_id": None, - "sep_token_id": None, - # The following attributes are EAGLE specific - "eagle_config": eagle_config, - } + eagle_config_update = { + "architectures": [architectures], + "head_dim": model.eagle_module.config.kv_channels, + "hidden_act": self._hf_text_config.hidden_act, + "hidden_size": self._hf_text_config.hidden_size, + "intermediate_size": model.eagle_module.config.ffn_hidden_size, + "max_position_embeddings": self._hf_text_config.max_position_embeddings, + "num_attention_heads": model.eagle_module.config.num_attention_heads, + "num_key_value_heads": model.eagle_module.config.num_query_groups, + "num_hidden_layers": mode_cfg["config"]["eagle_num_layers"], + "vocab_size": self._hf_text_config.vocab_size, + # Unset any special token ids given that the tokenizer can change here. + "bos_token_id": None, + "eos_token_id": None, + "pad_token_id": None, + "sep_token_id": None, + # The following attributes are EAGLE specific + "eagle_config": eagle_config, + } - # [TODO] (yeyu): there is also target_hidden_size - if mode_cfg["config"]["draft_vocab_size"] > 0: - eagle_config_update["draft_vocab_size"] = mode_cfg["config"]["draft_vocab_size"] - else: - eagle_config_update["draft_vocab_size"] = None + # [TODO] (yeyu): there is also target_hidden_size + if mode_cfg["config"]["draft_vocab_size"] > 0: + eagle_config_update["draft_vocab_size"] = mode_cfg["config"][ + "draft_vocab_size" + ] + else: + eagle_config_update["draft_vocab_size"] = None - self._hf_extra_config.update(eagle_config_update) + self._hf_extra_config.update(eagle_config_update) if mode == "mtp" and export_extra_modules: mtp_config = { @@ -292,7 +301,11 @@ class GPTModelExporter: } self._hf_config.mtp = mtp_config - def save_pretrained(self, save_directory: str | os.PathLike): + def save_pretrained( + self, + save_directory: str | os.PathLike, + pretrained_model_name_or_path: str | os.PathLike | None = None, + ): """Save a unified checkpoint which can be deploied by vLLM and TensorRT-LLM. Args: @@ -318,7 +331,7 @@ class GPTModelExporter: quantization = "NVFP4" # TODO (chenhany): need to handle Medusa and EAGLE meatadata - if torch.distributed.get_rank() == 0: + if torch.distributed.get_rank() == torch.distributed.get_world_size() - 1: if self.export_extra_modules and self._hf_extra_config is not None: # os.makedirs(save_directory, exist_ok=True) # with open(save_directory + "/config.json", 'w') as file: @@ -342,8 +355,17 @@ class GPTModelExporter: pass except TypeError: pass + try: + # Load and save preprocessor config from the original model + processor = AutoProcessor.from_pretrained( + self._hf_pretrained_model_name, trust_remote_code=self.trust_remote_code + ) + if hasattr(processor, "image_processor"): + processor.image_processor.save_pretrained(save_directory) + except (OSError, ValueError, ImportError): + pass - if torch.distributed.get_rank() == 0: + if torch.distributed.get_rank() == torch.distributed.get_world_size() - 1: hf_quant_config = { "producer": { "name": "modelopt", @@ -358,6 +380,75 @@ class GPTModelExporter: with open(save_directory + "/hf_quant_config.json", "w") as f: json.dump(hf_quant_config, f, indent=4) + if ( + torch.distributed.get_rank() == 0 + and self.is_multimodal + and pretrained_model_name_or_path is not None + ): + hf_checkpoint_path = Path(pretrained_model_name_or_path) + if not hf_checkpoint_path.is_dir(): + hf_checkpoint_path = tempfile.gettempdir() + "/" + pretrained_model_name_or_path + if not Path(hf_checkpoint_path).exists(): + snapshot_download( + repo_id=pretrained_model_name_or_path, + local_dir=hf_checkpoint_path, + ) + + safetensors_file = Path(hf_checkpoint_path) / "model.safetensors" + safetensors_index_file = Path(hf_checkpoint_path) / "model.safetensors.index.json" + + multimodal_state_dict = {} + + if safetensors_file.is_file(): + print(f"Loading multimodal components from single file: {safetensors_file}") + with safe_open(safetensors_file, framework="pt") as f: + multimodal_keys = [ + key + for key in f.keys() # noqa: SIM118 + if key.startswith(("multi_modal_projector", "vision_model")) + ] + for key in tqdm(multimodal_keys, desc="Loading multimodal tensors"): + multimodal_state_dict[key] = f.get_tensor(key) + + elif safetensors_index_file.is_file(): + print(f"Loading multimodal components from sharded model: {hf_checkpoint_path}") + with open(safetensors_index_file) as f: + safetensors_index = json.load(f) + + # For multimodal models, vision_model and multi_modal_projector are in the first shard + all_shard_files = sorted(set(safetensors_index["weight_map"].values())) + first_shard_file = all_shard_files[0] # e.g., "model-00001-of-00050.safetensors" + + # Load multimodal components from the first shard file + safetensors_filepath = Path(hf_checkpoint_path) / first_shard_file + print(f"Loading multimodal components from {first_shard_file}") + + with safe_open(safetensors_filepath, framework="pt") as f: + shard_keys = list(f.keys()) + multimodal_keys_in_shard = [ + k + for k in shard_keys + if k.startswith(("multi_modal_projector", "vision_model")) + ] + + if multimodal_keys_in_shard: + print( + f"Found {len(multimodal_keys_in_shard)} multimodal tensors in {first_shard_file}" + ) + for key in tqdm( + multimodal_keys_in_shard, desc="Loading multimodal tensors" + ): + multimodal_state_dict[key] = f.get_tensor(key) + else: + print(f"No multimodal components found in {first_shard_file}") + + else: + print(f"Warning: No safetensors files found in {hf_checkpoint_path}") + + print(f"Successfully loaded {len(multimodal_state_dict)} multimodal tensors") + # Add multimodal components to state_dict + state_dict.update(multimodal_state_dict) + # Barrier to ensure the export_dir has been created. torch.distributed.barrier() @@ -393,6 +484,7 @@ class GPTModelExporter: "name_remapping": self._name_remapping, "qkv_slicing": self._qkv_slicing, "gated_mlp_slicing": self._gated_mlp_slicing, + "pack_name_remapping": self._pack_name_remapping, } func = method_map[mapping.func_name] prefix = mapping.target_name_or_prefix @@ -400,7 +492,11 @@ class GPTModelExporter: return lambda m, *args: func(m, prefix.format(*args), **func_kwargs) for arch, mappings in all_mcore_hf_export_mapping.items(): - all_rules[arch] = {k: _custom_mapping_to_lambda(v) for (k, v) in mappings.items()} + all_rules[arch] = { + k: _custom_mapping_to_lambda(v) if isinstance(v, CustomModuleMapping) else v + for (k, v) in mappings.items() + if isinstance(v, (CustomModuleMapping, bool)) + } return all_rules @@ -608,6 +704,71 @@ class GPTModelExporter: self._state_dict[k_proj_key] = val.detach().clone() self._state_dict[v_proj_key] = val.detach().clone() + def _pack_name_remapping(self, module, prefix, layer_type=None): + """Pack name remapping into one tensor.""" + weight_list = [] + weight_scale_list = [] + weight_scale_2_list = [] + input_scale_list = [] + + for expert in module: + assert layer_type is not None, "layer_type is required for pack_name_remapping" + name_to_value, qformat, block_size = get_quantized_state( + getattr(expert, layer_type), self.dtype + ) + weight = name_to_value.pop("weight") + weight_scale, weight_scale_2 = self._get_weight_scales(name_to_value, qformat) + input_scale = ( + name_to_value.pop("input_scale") if "input_scale" in name_to_value else None + ) + + weight_list.append(weight) + weight_scale_list.append(weight_scale) + weight_scale_2_list.append(weight_scale_2) + input_scale_list.append(input_scale) + + merged_weight = torch.stack(weight_list, dim=0) + + # Transpose the last two dimensions to match HuggingFace format + # NeMo format: [num_experts, out_features, in_features] + # HF format: [num_experts, in_features, out_features] + merged_weight = merged_weight.transpose(-2, -1).contiguous() + + if weight_scale_2_list[0] is None: + merged_weight_scale_2 = None + if weight_scale_list[0] is not None: + merged_weight_scale = torch.max(torch.stack(weight_scale_list, dim=0), dim=0)[0] + else: + merged_weight_scale = None + else: + # NVFP4 + merged_weight_scale_2 = torch.max(torch.stack(weight_scale_2_list, dim=0), dim=0)[0] + merged_weight_scale = torch.stack(weight_scale_list, dim=0) + # Transpose the scaling factors to match the transposed weights + merged_weight_scale = merged_weight_scale.transpose(-2, -1).contiguous() + + if input_scale_list[0] is not None: + merged_input_scale = torch.max(torch.stack(input_scale_list, dim=0), dim=0)[0] + else: + merged_input_scale = None + + # Save the merged weights + if merged_weight_scale is None: + self._state_dict[prefix] = merged_weight + else: + self._state_dict[prefix] = to_quantized_weight( + merged_weight, + merged_weight_scale, + qformat, + merged_weight_scale_2, + block_size, + ) + self._state_dict[prefix + "_weight_scale"] = merged_weight_scale + if merged_weight_scale_2 is not None: + self._state_dict[prefix + "_weight_scale_2"] = merged_weight_scale_2 + if merged_input_scale is not None: + self._state_dict[prefix + "_input_scale"] = merged_input_scale + def _get_medusa_heads_state_dict(self): medusa_heads = getattr(self.model, "medusa_heads", None) if medusa_heads is None: @@ -832,7 +993,7 @@ class GPTModelExporter: self.rules["linear_proj"](layer.self_attention.linear_proj, layer_id) else: if layer.self_attention.q_layernorm is not None and not isinstance( - layer.self_attention.q_layernorm, IdentityOp + layer.self_attention.q_layernorm, (IdentityOp, L2Norm) ): self.rules["q_layernorm"](layer.self_attention.q_layernorm, layer_id) self.rules["k_layernorm"](layer.self_attention.k_layernorm, layer_id) @@ -855,12 +1016,21 @@ class GPTModelExporter: self.rules["shared_experts.linear_fc2"]( layer.mlp.shared_experts.linear_fc2, layer_id ) - for expert_id, expert in enumerate(layer.mlp.experts.local_experts): + if not self.rules.get("use_packed_local_experts", False): + for expert_id, expert in enumerate(layer.mlp.experts.local_experts): + self.rules["local_experts.linear_fc1"]( + expert.linear_fc1, layer_id, expert_id + ) + self.rules["local_experts.linear_fc2"]( + expert.linear_fc2, layer_id, expert_id + ) + else: + # For llama 4, in hf unified checkpoint, all local experts share one scale self.rules["local_experts.linear_fc1"]( - expert.linear_fc1, layer_id, expert_id + layer.mlp.experts.local_experts, layer_id ) self.rules["local_experts.linear_fc2"]( - expert.linear_fc2, layer_id, expert_id + layer.mlp.experts.local_experts, layer_id ) else: self.rules["linear_fc1"](layer.mlp.linear_fc1, layer_id) @@ -892,473 +1062,7 @@ def export_mcore_gpt_to_hf( exporter = GPTModelExporter( model, pretrained_model_name_or_path, export_extra_modules=export_extra_modules, dtype=dtype ) - exporter.save_pretrained(export_dir) - - -class GPTModelImporter: - """Megatron Core GPTModel HuggingFace Importer. - - The Importer is created by `import_mcore_gpt_from_hf` to host attributes - and methods that import a Megatron Core GPTModel from a supported Hugging - Face model. - - Args: - model: The Megatron Core GPTModel instance. - pretrained_model_name_or_path: Can be either: the *model id* of a - pretrained model hosted inside a model repo on huggingface.co; or - a *directory* containing model weights saved using - [`~PreTrainedModel.save_pretrained`], e.g., `./my_model_directory/`. - dtype: The weights data type to export the unquantized layers. - """ - - weight_scale_name: str = "weight_scale_inv" - - def __init__( - self, - model: torch.nn.Module, - pretrained_model_name_or_path: str, - workspace_dir: str | None = None, - dtype=torch.bfloat16, - trust_remote_code: bool = True, - ): - """Create a GPTModel importer instance.""" - self._hf_config = transformers.AutoConfig.from_pretrained( - pretrained_model_name_or_path, trust_remote_code=trust_remote_code - ) - pretrained_model_path = Path(pretrained_model_name_or_path) - if not pretrained_model_path.is_dir(): - if workspace_dir is None: - workspace_dir = tempfile.gettempdir() - pretrained_model_path = workspace_dir + "/" + pretrained_model_name_or_path - if torch.distributed.get_rank() == 0: - snapshot_download( - repo_id=pretrained_model_name_or_path, - local_dir=pretrained_model_path, - ) - torch.distributed.barrier() - self.arch = self._hf_config.architectures[0] - self.all_rules = self._populate_rule_book() - self.rules = self.all_rules[self.arch] - self.model = model - self.pretrained_model_path = pretrained_model_path - self.dtype = dtype - self.disable_tqdm = torch.distributed.get_rank() > 0 - - def _populate_rule_book(self): - """The rule book maps each state_dict key to a Callable.""" - all_rules = {} - - def _custom_mapping_to_lambda(mapping): - method_map = { - "name_remapping": self._name_remapping, - "qkv_merging": self._qkv_merging, - "gated_mlp_merging": self._gated_mlp_merging, - } - func = method_map[mapping.func_name] - prefix = mapping.target_name_or_prefix - func_kwargs = mapping.func_kwargs - return lambda m, *args: func(m, prefix.format(*args), **func_kwargs) - - for arch, mappings in all_mcore_hf_import_mapping.items(): - all_rules[arch] = {k: _custom_mapping_to_lambda(v) for (k, v) in mappings.items()} - - return all_rules - - def _get_safetensor(self, key, sharding_dim: int | None = None): - """Get a safetensor from the sharded checkpoint.""" - safetensors_file = Path(self.pretrained_model_path) / "model.safetensors" - safetensors_index_file = Path(self.pretrained_model_path) / "model.safetensors.index.json" - if safetensors_file.is_file(): - pass - elif safetensors_index_file.is_file(): - with open(safetensors_index_file) as f: - safetensors_index = json.load(f) - safetensors_file = ( - Path(self.pretrained_model_path) / safetensors_index["weight_map"][key] - ) - else: - raise ValueError("Only safetensors (single of multi- files) are supported.") - - with safe_open(safetensors_file, framework="pt") as f: - if sharding_dim is None: - tensor = f.get_tensor(key) - else: - tensor_slice = f.get_slice(key) - assert tensor_slice is not None - shape = tensor_slice.get_shape() - # MCore tensor parallel model sharding - tp_rank = get_tensor_model_parallel_rank() - tp_size = get_tensor_model_parallel_world_size() - per_rank_size = shape[sharding_dim] // tp_size - rank_offset = tp_rank * per_rank_size - assert len(shape) == 2 - assert shape[sharding_dim] % tp_size == 0 - if sharding_dim in (1, -1): - tensor = tensor_slice[:, rank_offset : rank_offset + per_rank_size] - else: - tensor = tensor_slice[rank_offset : rank_offset + per_rank_size, :] - return tensor - - def _get_tensor_parallel_shard(self, tensor: torch.Tensor, dim: int): - tp_rank = get_tensor_model_parallel_rank() - tp_size = get_tensor_model_parallel_world_size() - if tp_size == 1: - return tensor - return torch.chunk(tensor, tp_size, dim=dim)[tp_rank] - - def _name_remapping( - self, - module, - prefix, - mapping={}, - sharding_dim: int | None = None, - ): - if isinstance(module, torch.Tensor): - module.data.copy_(self._get_safetensor(prefix)) - return - - weight = module.state_dict().get("weight", None) - weight_scale = module.state_dict().get("weight_quantizer._scale", None) - - state_dict = {} - - if weight is None: - raise ValueError(f"{module!s} does not contain weight!") - else: - tensor = self._get_safetensor(prefix + "weight", sharding_dim=sharding_dim) - - if weight_scale is not None: - scale_name = prefix + self.weight_scale_name - if weight_scale.ndim > 0: - scale = self._get_safetensor(scale_name, sharding_dim=sharding_dim) - else: - scale = self._get_safetensor(scale_name) - scale = scale.to(weight_scale.dtype).to(device=weight_scale.device) - state_dict["weight_quantizer._scale"] = scale - - if tensor.shape != weight.shape: - expanded_tensor = torch.zeros(weight.shape, dtype=tensor.dtype) - expanded_tensor[: tensor.shape[0], : tensor.shape[1]] = tensor - tensor = expanded_tensor - state_dict["weight"] = tensor.view(dtype=weight.dtype).to(device=weight.device) - else: - state_dict["weight"] = tensor.to(dtype=self.dtype).to(device=weight.device) - - # Handle the rest of the state_dict. - for key, val in module.state_dict().items(): - if key in {"weight", "weight_quantizer._scale"}: - continue - elif "extra_state" in key: - state_dict[key] = val - else: - source_key = mapping.get(key, key) - tensor = self._get_safetensor(prefix + source_key, sharding_dim=sharding_dim) - state_dict[key] = tensor.to(dtype=self.dtype).to(device=val.device) - - module.load_state_dict(state_dict) - - def _gated_mlp_merging( - self, - module, - prefix, - gate_proj_name="gate_proj", - up_proj_name="up_proj", - sharding_dim: int | None = None, - ): - weight = module.state_dict().get("weight", None) - weight_scale = module.state_dict().get("weight_quantizer._scale", None) - - state_dict = {} - - if weight is None: - raise ValueError(f"{module!s} does not contain weight!") - else: - gate_proj = self._get_safetensor( - prefix + gate_proj_name + ".weight", sharding_dim=sharding_dim - ) - up_proj = self._get_safetensor( - prefix + up_proj_name + ".weight", sharding_dim=sharding_dim - ) - tensor = torch.cat((gate_proj, up_proj), dim=0) - - if weight_scale is not None: - gate_scale_name = prefix + gate_proj_name + "." + self.weight_scale_name - up_scale_name = prefix + up_proj_name + "." + self.weight_scale_name - if weight_scale.ndim > 0: - gate_scale = self._get_safetensor(gate_scale_name, sharding_dim=sharding_dim) - up_scale = self._get_safetensor(up_scale_name, sharding_dim=sharding_dim) - scale = torch.cat((gate_scale, up_scale), dim=0) - else: - scale = self._get_safetensor(gate_scale_name) - # If source model is per tensor, compute a per tensor scale with max. - if scale.ndim > 0: - scale = scale.max(dim=0).max(dim=0) - state_dict["weight_quantizer._scale"] = scale.to(weight_scale.dtype).to( - device=weight_scale.device - ) - state_dict["weight"] = tensor.view(weight.dtype).to(device=weight.device) - else: - state_dict["weight"] = tensor.to(self.dtype).to(device=weight.device) - - module.load_state_dict(state_dict) - - def _qkv_merging( - self, - module, - prefix, - q_proj_name="q_proj", - k_proj_name="k_proj", - v_proj_name="v_proj", - sharding_dim: int | None = None, - ): - config = module.config - hidden_size = config.hidden_size - num_query_groups = config.num_query_groups - head_num = config.num_attention_heads - head_size = config.kv_channels - - if sharding_dim is not None: - tp_size = get_tensor_model_parallel_world_size() - assert head_num % tp_size == 0 - assert num_query_groups % tp_size == 0 - head_num = head_num // tp_size - num_query_groups = num_query_groups // tp_size - - heads_per_group = head_num // num_query_groups - qkv_total_dim = head_num + 2 * num_query_groups - q_slice = torch.cat( - [ - torch.arange((heads_per_group + 2) * i, (heads_per_group + 2) * i + heads_per_group) - for i in range(num_query_groups) - ] - ) - k_slice = torch.arange(heads_per_group, qkv_total_dim, (heads_per_group + 2)) - v_slice = torch.arange(heads_per_group + 1, qkv_total_dim, (heads_per_group + 2)) - - state_dict = {} - - weight = module.state_dict().get("weight", None) - weight_scale = module.state_dict().get("weight_quantizer._scale", None) - - if weight is None: - raise ValueError(f"{module!s} does not contain weight!") - - if weight_scale is not None: - q_scale_name = prefix + q_proj_name + "." + self.weight_scale_name - k_scale_name = prefix + k_proj_name + "." + self.weight_scale_name - v_scale_name = prefix + v_proj_name + "." + self.weight_scale_name - - if weight_scale.ndim > 0: - q_scale = self._get_safetensor(q_scale_name, sharding_dim=sharding_dim) - k_scale = self._get_safetensor(k_scale_name, sharding_dim=sharding_dim) - v_scale = self._get_safetensor(v_scale_name, sharding_dim=sharding_dim) - weight_scale[q_slice] = q_scale.to(weight_scale.dtype).to( - device=weight_scale.device - ) - weight_scale[k_slice] = k_scale.to(weight_scale.dtype).to( - device=weight_scale.device - ) - weight_scale[v_slice] = v_scale.to(weight_scale.dtype).to( - device=weight_scale.device - ) - else: - q_scale = self._get_safetensor(q_scale_name) - weight_scale = q_scale.to(weight_scale.dtype).to(device=weight_scale.device) - state_dict["weight_quantizer._scale"] = weight_scale - - q_proj = self._get_safetensor(prefix + q_proj_name + ".weight", sharding_dim=sharding_dim) - k_proj = self._get_safetensor(prefix + k_proj_name + ".weight", sharding_dim=sharding_dim) - v_proj = self._get_safetensor(prefix + v_proj_name + ".weight", sharding_dim=sharding_dim) - q_proj = q_proj.reshape(-1, head_size, hidden_size) - k_proj = k_proj.reshape(-1, head_size, hidden_size) - v_proj = v_proj.reshape(-1, head_size, hidden_size) - tensor = weight.detach().clone().reshape([qkv_total_dim, head_size, hidden_size]) - - if weight_scale is not None: - tensor[q_slice] = q_proj.view(dtype=tensor.dtype).to(device=tensor.device) - tensor[k_slice] = k_proj.view(dtype=tensor.dtype).to(device=tensor.device) - tensor[v_slice] = v_proj.view(dtype=tensor.dtype).to(device=tensor.device) - else: - tensor[q_slice] = q_proj.to(dtype=tensor.dtype).to(device=tensor.device) - tensor[k_slice] = k_proj.to(dtype=tensor.dtype).to(device=tensor.device) - tensor[v_slice] = v_proj.to(dtype=tensor.dtype).to(device=tensor.device) - - state_dict["weight"] = tensor.reshape(-1, hidden_size) - - module.load_state_dict(state_dict) - - def _import_state_dict(self): - model = self.model - - layer_pbar = tqdm(model.decoder.layers, disable=self.disable_tqdm) - - # Embedding - if hasattr(model, "embedding"): - layer_pbar.set_description("Importing word embedding") - self.rules["word_embeddings"](model.embedding.word_embeddings) - - # Decoder layers - for layer in layer_pbar: - layer_id = layer.layer_number - 1 - - if isinstance(layer, MambaLayer): - if not isinstance(layer.norm, IdentityOp): - self.rules["norm"](layer.norm, layer_id) - - self.rules["mixer_norm"](layer.mixer.norm, layer_id) - self.rules["A_log"](layer.mixer.A_log, layer_id) - self.rules["D"](layer.mixer.D, layer_id) - self.rules["dt_bias"](layer.mixer.dt_bias, layer_id) - - self.rules["conv1d"](layer.mixer.conv1d, layer_id) - self.rules["in_proj"](layer.mixer.in_proj, layer_id) - self.rules["out_proj"](layer.mixer.out_proj, layer_id) - - elif isinstance(layer, TransformerLayer): - if not isinstance(layer.input_layernorm, IdentityOp): - self.rules["input_layernorm"](layer.input_layernorm, layer_id) - - if not isinstance(layer.self_attention, IdentityOp): - if "MLASelfAttention" in str(type(layer.self_attention)): - if hasattr(layer.self_attention, "linear_q_proj"): - layer_pbar.set_description("Importing MLA (without q LoRA)") - self.rules["linear_q_proj"]( - layer.self_attention.linear_q_proj, layer_id - ) - else: - layer_pbar.set_description("Importing MLA (with q LoRA)") - self.rules["linear_q_down_proj"]( - layer.self_attention.linear_q_down_proj, layer_id - ) - self.rules["linear_q_layernorm"]( - layer.self_attention.q_layernorm, layer_id - ) - self.rules["linear_q_up_proj"]( - layer.self_attention.linear_q_up_proj, layer_id - ) - self.rules["linear_kv_down_proj"]( - layer.self_attention.linear_kv_down_proj, layer_id - ) - self.rules["linear_kv_layernorm"]( - layer.self_attention.kv_layernorm, layer_id - ) - self.rules["linear_kv_up_proj"]( - layer.self_attention.linear_kv_up_proj, layer_id - ) - self.rules["linear_proj"](layer.self_attention.linear_proj, layer_id) - else: - layer_pbar.set_description("Importing GQA/MHA") - if layer.self_attention.q_layernorm is not None and not isinstance( - layer.self_attention.q_layernorm, IdentityOp - ): - self.rules["q_layernorm"](layer.self_attention.q_layernorm, layer_id) - self.rules["k_layernorm"](layer.self_attention.k_layernorm, layer_id) - self.rules["linear_qkv"](layer.self_attention.linear_qkv, layer_id) - self.rules["linear_proj"](layer.self_attention.linear_proj, layer_id) - - if not isinstance(layer.pre_mlp_layernorm, IdentityOp): - self.rules["pre_mlp_layernorm"](layer.pre_mlp_layernorm, layer_id) - - if not isinstance(layer.mlp, IdentityOp): - if "MoE" in str(type(layer.mlp)): - layer_pbar.set_description("Importing MoE") - self.rules["router"](layer.mlp.router, layer_id) - if ( - hasattr(layer.mlp, "shared_experts") - and layer.mlp.shared_experts is not None - ): - layer_pbar.set_description("Importing MoE shared experts") - self.rules["shared_experts.linear_fc1"]( - layer.mlp.shared_experts.linear_fc1, layer_id - ) - self.rules["shared_experts.linear_fc2"]( - layer.mlp.shared_experts.linear_fc2, layer_id - ) - for expert_id, expert in tqdm( - enumerate(layer.mlp.experts.local_experts), - desc="Importing MoE local experts", - leave=False, - disable=self.disable_tqdm, - ): - self.rules["local_experts.linear_fc1"]( - expert.linear_fc1, layer_id, expert_id - ) - self.rules["local_experts.linear_fc2"]( - expert.linear_fc2, layer_id, expert_id - ) - else: - layer_pbar.set_description("Importing MLP") - self.rules["linear_fc1"](layer.mlp.linear_fc1, layer_id) - self.rules["linear_fc2"](layer.mlp.linear_fc2, layer_id) - - # Final layernorm - if hasattr(model.decoder, "final_layernorm") and model.decoder.final_layernorm: - self.rules["final_layernorm"](model.decoder.final_layernorm) - - if hasattr(model.decoder, "final_norm") and model.decoder.final_norm: - self.rules["final_norm"](model.decoder.final_norm) - - # Output layer - if hasattr(model, "output_layer") and not model.share_embeddings_and_output_weights: - self.rules["output_layer"](model.output_layer) - - # MTP - if hasattr(model, "mtp"): - # MTP is the last layer in DeepSeek V3/R1 - layer_id += 1 - for mtp in model.mtp: - self.rules["mtp.fc"](mtp.fc, layer_id) - self.rules["mtp.enorm"](mtp.enorm, layer_id) - self.rules["mtp.hnorm"](mtp.hnorm, layer_id) - self.rules["mtp.input_layernorm"](mtp.decoder.layers[0].input_layernorm, layer_id) - if hasattr(mtp.decoder.layers[0].self_attention, "linear_q_proj"): - self.rules["mtp.linear_q_proj"]( - mtp.decoder.layers[0].self_attention.linear_q_proj, layer_id - ) - else: - self.rules["mtp.linear_q_down_proj"]( - mtp.decoder.layers[0].self_attention.linear_q_down_proj, layer_id - ) - self.rules["mtp.linear_q_layernorm"]( - mtp.decoder.layers[0].self_attention.q_layernorm, layer_id - ) - self.rules["mtp.linear_q_up_proj"]( - mtp.decoder.layers[0].self_attention.linear_q_up_proj, layer_id - ) - self.rules["mtp.linear_kv_down_proj"]( - mtp.decoder.layers[0].self_attention.linear_kv_down_proj, layer_id - ) - self.rules["mtp.linear_kv_layernorm"]( - mtp.decoder.layers[0].self_attention.kv_layernorm, layer_id - ) - self.rules["mtp.linear_kv_up_proj"]( - mtp.decoder.layers[0].self_attention.linear_kv_up_proj, layer_id - ) - self.rules["mtp.linear_proj"]( - mtp.decoder.layers[0].self_attention.linear_proj, layer_id - ) - self.rules["mtp.pre_mlp_layernorm"]( - mtp.decoder.layers[0].pre_mlp_layernorm, layer_id - ) - self.rules["mtp.router"](mtp.decoder.layers[0].mlp.router, layer_id) - self.rules["mtp.shared_experts.linear_fc1"]( - mtp.decoder.layers[0].mlp.shared_experts.linear_fc1, layer_id - ) - self.rules["mtp.shared_experts.linear_fc2"]( - mtp.decoder.layers[0].mlp.shared_experts.linear_fc2, layer_id - ) - for expert_id, expert in tqdm( - enumerate(mtp.decoder.layers[0].mlp.experts.local_experts), - desc="Importing MoE local experts", - leave=False, - disable=self.disable_tqdm, - ): - self.rules["mtp.local_experts.linear_fc1"]( - expert.linear_fc1, layer_id, expert_id - ) - self.rules["mtp.local_experts.linear_fc2"]( - expert.linear_fc2, layer_id, expert_id - ) + exporter.save_pretrained(export_dir, pretrained_model_name_or_path) def import_mcore_gpt_from_hf( diff --git a/modelopt/torch/nas/hparams/concat.py b/modelopt/torch/nas/hparams/concat.py index 59d2d7272..531db989e 100644 --- a/modelopt/torch/nas/hparams/concat.py +++ b/modelopt/torch/nas/hparams/concat.py @@ -19,6 +19,7 @@ from collections import defaultdict from collections.abc import Callable, Iterator from itertools import product from math import prod +from warnings import warn import numpy as np import torch @@ -51,6 +52,8 @@ class ConcatTracedHp(TracedHp): # use itertools.product and iterate over all combinations yield from product(*all_choices) return + else: + warn(f"ConcatTracedHp: {n_combos=} is larger than {n_max=}. Pruning combinations.") # otherwise, we use a pruned set of combinations based on immediate vicinity of each value @@ -104,6 +107,12 @@ class ConcatTracedHp(TracedHp): sum_to_combo[s] = combo_full return sum_to_combo + def _set_hp_start_idx(self) -> None: + """Compute start indices for input hps.""" + # NOTE: the last index is the length of the concatenated hp + hp_start_idx = np.concatenate(([0], np.cumsum([hp.max for hp in self._inputs]))) + self._hp_start_idx = torch.asarray(hp_start_idx, dtype=torch.long) + @property def active(self) -> int: """Return the sum of active values of all hparams.""" @@ -273,9 +282,38 @@ class ConcatTracedHp(TracedHp): hp_in.choices = [combo[i] for combo in sum_to_combo.values()] hp_in._is_configurable = False - # compute start indices for input hp - # NOTE: the last index is the length of the concatenated hp - hp_start_idx = np.concatenate(([0], np.cumsum([hp.max for hp in self._inputs]))) - self._hp_start_idx = torch.asarray(hp_start_idx, dtype=torch.long) + self._set_hp_start_idx() return mapping + + def reset_choices(self) -> None: + """Reset the choices of the concat hparam. + + Useful if we want to reset choices after input hparam choices are changed during modify(). + """ + self._sum_to_combo = self._get_sum_to_combo() + self._set_hp_start_idx() + with self._force_configurable(): + self.choices = list(self._sum_to_combo) + + +def build_concat_hp(inputs: list[TracedHp]): + """Initialize a non-configurable concat hparam from a list of input hparams. + + One key difference from ConcatTracedHp via tracing is that in ConcatTracedHp, the input hparams + are not configurable, and only the concatenated hparam is configurable. In build_concat_hp, its the opposite. + + This is useful for building concat hparams from a list of configurable input hparams instead of + tracing (e.g. for megatron language model DynamicModule). + """ + concat_hp = object.__new__(ConcatTracedHp) + concat_hp._inputs = inputs + concat_hp._sum_to_combo = concat_hp._get_sum_to_combo() + concat_hp._set_hp_start_idx() + + choices = list(concat_hp._sum_to_combo) + concat_hp.__init__(choices) # type: ignore[misc] + concat_hp._is_configurable = False + concat_hp._importance_estimators = None + + return concat_hp diff --git a/modelopt/torch/nas/plugins/megatron.py b/modelopt/torch/nas/plugins/megatron.py index cd0854d43..e95078d47 100644 --- a/modelopt/torch/nas/plugins/megatron.py +++ b/modelopt/torch/nas/plugins/megatron.py @@ -15,8 +15,8 @@ """Plugin to add NAS/Pruning support for megatron-core GPT model.""" -from collections.abc import Callable, Sequence from typing import Any +from warnings import warn import torch import torch.nn as nn @@ -48,10 +48,8 @@ from megatron.core.transformer.mlp import MLP from megatron.core.transformer.transformer_layer import TransformerLayer from modelopt.torch.opt.dynamic import DynamicModule -from modelopt.torch.opt.hparam import HPType from modelopt.torch.opt.searcher import ConstraintsDict from modelopt.torch.opt.utils import named_hparams -from modelopt.torch.trace import Symbol from modelopt.torch.utils import distributed as dist from modelopt.torch.utils import ( get_module_device, @@ -68,12 +66,15 @@ from ..algorithms import ( ConstraintsFunc, ConstraintsRes, ) +from ..hparams.concat import build_concat_hp from ..modules import _DynamicLayerNorm from ..modules.utils import get_sliced_tensor, get_sliced_tensor_by_slices from ..registry import DMRegistry from ..search_space import SampleFunc from ..traced_hp import TracedHp +SUPPORTED_MODELS = {GPTModel: "megatron.core.models.gpt.GPTModel"} + try: from megatron.core.extensions.transformer_engine import TEDotProductAttention @@ -81,7 +82,16 @@ try: except ImportError: HAS_TE = False -__all__ = ["_DynamicGPTModel", "drop_mcore_gpt_layers"] +try: + from megatron.core.models.mamba import MambaModel + + SUPPORTED_MODELS[MambaModel] = "megatron.core.models.mamba.MambaModel" + + HAS_MAMBA = True +except ImportError: + HAS_MAMBA = False + +__all__ = ["drop_mcore_gpt_layers", "drop_mcore_language_model_layers"] class _DynamicParallelLinear(DynamicModule): @@ -141,9 +151,8 @@ class _DynamicVocabParallelEmbedding(DynamicModule): self._register_hparam("embedding_dim", TracedHp(list(range(1, self.embedding_dim + 1)))) self._register_dynamic_attribute("weight", self._get_weight) - def _get_weight( - self, mod: "_DynamicVocabParallelEmbedding", weight: torch.Tensor - ) -> torch.Tensor: + @staticmethod + def _get_weight(mod: "_DynamicVocabParallelEmbedding", weight: torch.Tensor) -> torch.Tensor: """Return the weight tensor of the embedding layer.""" return get_sliced_tensor(mod, weight, None, "embedding_dim") @@ -166,52 +175,6 @@ class _DynamicFusedLayerNorm(_DynamicLayerNorm): self._register_dynamic_attribute("hidden_size", self._get_normalized_shape) -class RepeatedTracedHp(TracedHp): - """An hparam repeated N number of times to form a longer hparam. - - One key difference from ConcatTracedHp is that in ConcatTracedHp, the input hparams are not configurable, - and only the concatenated hparam is configurable. In RepeatedTracedHp, its the opposite. - """ - - def __init__(self, hparam: TracedHp, num_repeats: int) -> None: - """Initialize the repeated hparam.""" - self._hparam = hparam - self._num_repeats = num_repeats - choices = [c * self._num_repeats for c in self._hparam.choices] - original = self._hparam.original * self._num_repeats - super().__init__(choices, original) - self._is_configurable = False - self._importance_estimators = None - - @property # type: ignore[misc] - def active(self) -> int: - """Return the active value of the hparam.""" - assert isinstance(self._hparam.active, int) - return self._hparam.active * self._num_repeats - - @property - def active_slice(self) -> TracedHp.ActiveSlice: - """Return the currently active sorted indices or slice corresponding to the active value.""" - hp_active_slice = self._hparam.active_slice - if isinstance(hp_active_slice, slice): - hp_active_slice = torch.LongTensor(range(hp_active_slice.stop)) - - active_slice = torch.cat( - [hp_active_slice + i * self._hparam.max for i in range(self._num_repeats)] - ) - return active_slice - - @property # type: ignore[misc] - def choices(self) -> Sequence[HPType]: - """Return available choices.""" - return [c * self._num_repeats for c in self._hparam.choices] - - def _resolve_dependencies( - self, sym: Symbol, get_hp: Callable[[Symbol], TracedHp] - ) -> dict[Symbol, TracedHp]: - raise NotImplementedError("RepeatedTracedHp does not support `_resolve_dependencies`!") - - @DMRegistry.register({MLP: "megatron.core.transformer.mlp.MLP"}) class _DynamicMLP(DynamicModule): """A ``megatron.core.transformer.mlp.MLP`` layer with dynamic hyperparams.""" @@ -226,7 +189,7 @@ class _DynamicMLP(DynamicModule): ffn_hidden_size = TracedHp(list(range(1, self.config.ffn_hidden_size + 1))) fc1_output_size = ( - RepeatedTracedHp(ffn_hidden_size, 2) + build_concat_hp([ffn_hidden_size] * 2) if self.config.gated_linear_unit else ffn_hidden_size ) @@ -648,41 +611,30 @@ class _DynamicSelfAttention(DynamicModule): return self -@DMRegistry.register( - {TransformerLayer: "megatron.core.transformer.transformer_layer.TransformerLayer"} -) -class _DynamicTransformerLayer(DynamicModule): - """A ``megatron.core.transformer.transformer_layer.TransformerLayer`` layer with dynamic hyperparams.""" +class MambaTransformerLayerMixin(nn.Module): + """A mixin for MambaLayer and TransformerLayer to share the same logic.""" - def _setup(self): - # Convert the layernorms, self-attention, and mlp layers to dynamic modules - self.input_layernorm = DMRegistry.convert(self.input_layernorm) - self.self_attention = DMRegistry.convert(self.self_attention) - self.pre_mlp_layernorm = DMRegistry.convert(self.pre_mlp_layernorm) - self.mlp = DMRegistry.convert(self.mlp) - - # Register forward hook to collect activations for importance estimation + def _setup_mixin(self): + """Setup the mixin.""" self._register_temp_attribute("_scores", 0.0) self.hook_handle = self.register_forward_hook( self._layer_imp_forward_hook, with_kwargs=True ) - def set_hidden_size_hp(self, hidden_size: TracedHp) -> None: - self.input_layernorm.num_features = hidden_size - self.self_attention.linear_qkv.input_size = hidden_size - self.self_attention.linear_proj.output_size = hidden_size - self.pre_mlp_layernorm.num_features = hidden_size - self.mlp.linear_fc1.input_size = hidden_size - self.mlp.linear_fc2.output_size = hidden_size + def _export_mixin(self): + """Export the mixin.""" + self.hook_handle.remove() def _layer_imp_forward_hook(self, module, args, kwargs, output) -> None: """Hook to collect cosine similarity between input and output to rank layers for depth pruning.""" hidden_states = kwargs["hidden_states"] if "hidden_states" in kwargs else args[0] - output, _ = output # [seq_len, batch_size, hidden_size] + if isinstance(self, TransformerLayer): + output, _ = output # [seq_len, batch_size, hidden_size] # Dont aggregate activations from non-max subnets (e.g. from profiling) - if hidden_states.shape[-1] != self.input_layernorm.get_hparam("num_features").max: + # NOTE: max_hidden_size is set in set_hidden_size_hp for both DyamicModule classes below! + if hidden_states.shape[-1] != self.max_hidden_size: return with torch.no_grad(): @@ -692,48 +644,87 @@ class _DynamicTransformerLayer(DynamicModule): global_score = reduce_from_tensor_model_parallel_region(score).item() self._scores += global_score # aggregate sum instead of mean of scores for simplicity + +@DMRegistry.register( + {TransformerLayer: "megatron.core.transformer.transformer_layer.TransformerLayer"} +) +class _DynamicTransformerLayer(DynamicModule, MambaTransformerLayerMixin): + """A ``megatron.core.transformer.transformer_layer.TransformerLayer`` layer with dynamic hyperparams.""" + + def _setup(self): + # Convert the layernorms, self-attention, and mlp layers to dynamic modules + # NOTE: Mamba stack layers have either Attention or MLP, not both unlike GPT models + if isinstance(self.self_attention, SelfAttention): + self.input_layernorm = DMRegistry.convert(self.input_layernorm) + self.self_attention = DMRegistry.convert(self.self_attention) + if isinstance(self.mlp, MLP): + self.pre_mlp_layernorm = DMRegistry.convert(self.pre_mlp_layernorm) + self.mlp = DMRegistry.convert(self.mlp) + + # Register forward hook to collect activations for importance estimation + self._setup_mixin() + + def set_hidden_size_hp(self, hidden_size: TracedHp) -> None: + if isinstance(self.self_attention, SelfAttention): + self.input_layernorm.num_features = hidden_size + self.self_attention.linear_qkv.input_size = hidden_size + self.self_attention.linear_proj.output_size = hidden_size + if isinstance(self.mlp, MLP): + self.pre_mlp_layernorm.num_features = hidden_size + self.mlp.linear_fc1.input_size = hidden_size + self.mlp.linear_fc2.output_size = hidden_size + + self._register_temp_attribute("max_hidden_size", hidden_size.max) + def modify( self, *, num_heads_per_group_divisor: int = 1, num_query_groups_divisor: int = 1, ffn_hidden_size_divisor: int = 1, + **kwargs, # Unused hparams ) -> None: # Modify SelfAttention hparams - for hp_name, divisor in [ - ("num_heads_per_group", num_heads_per_group_divisor), - ("num_query_groups", num_query_groups_divisor), - ]: - hp = self.self_attention.get_hparam(hp_name) - choices = {int(make_divisible(c, divisor)) for c in hp.choices} # type: ignore[arg-type] - hp.choices = list(set(hp.choices) & choices | {hp.original}) + if isinstance(self.self_attention, SelfAttention): + for hp_name, divisor in [ + ("num_heads_per_group", num_heads_per_group_divisor), + ("num_query_groups", num_query_groups_divisor), + ]: + hp = self.self_attention.get_hparam(hp_name) + choices = {int(make_divisible(c, divisor)) for c in hp.choices} + hp.choices = list(set(hp.choices) & choices | {hp.original}) # Modify MLP hparams - hp_mlp = self.mlp.get_hparam("ffn_hidden_size") - choices = {int(make_divisible(c, ffn_hidden_size_divisor)) for c in hp_mlp.choices} # type: ignore[arg-type] - hp_mlp.choices = list(set(hp_mlp.choices) & choices | {hp_mlp.original}) + if isinstance(self.mlp, MLP): + hp_mlp = self.mlp.get_hparam("ffn_hidden_size") + choices = {int(make_divisible(c, ffn_hidden_size_divisor)) for c in hp_mlp.choices} + hp_mlp.choices = list(set(hp_mlp.choices) & choices | {hp_mlp.original}) def export(self): """Export the dynamic module to a torch.nn.Module.""" - self.hook_handle.remove() - self.input_layernorm.export() - self.self_attention.export() - self.pre_mlp_layernorm.export() - self.mlp.export() + self._export_mixin() + if isinstance(self.self_attention, SelfAttention): + self.input_layernorm.export() + self.self_attention.export() + if isinstance(self.mlp, MLP): + self.pre_mlp_layernorm.export() + self.mlp.export() super().export() return self def freeze(self): """Freeze the dynamic module.""" super().freeze() - self.input_layernorm.freeze() - self.self_attention.freeze() - self.pre_mlp_layernorm.freeze() - self.mlp.freeze() + if isinstance(self.self_attention, SelfAttention): + self.input_layernorm.freeze() + self.self_attention.freeze() + if isinstance(self.mlp, MLP): + self.pre_mlp_layernorm.freeze() + self.mlp.freeze() -@DMRegistry.register({GPTModel: "megatron.core.models.gpt.GPTModel"}) -class _DynamicGPTModel(DynamicModule): +@DMRegistry.register(SUPPORTED_MODELS) +class _DynamicMCoreLanguageModel(DynamicModule): """A ``megatron.core.models.gpt.GPTModel`` model with dynamic hyperparams.""" def _setup(self): @@ -764,9 +755,19 @@ class _DynamicGPTModel(DynamicModule): self.decoder.layers[i] = DMRegistry.convert(self.decoder.layers[i]) self.decoder.layers[i].set_hidden_size_hp(hidden_size) + # NOTE: GPTModel has final_layernorm, MambaModel has final_norm + self._register_temp_attribute( + "final_norm_attr_name", + "final_layernorm" if hasattr(self.decoder, "final_layernorm") else "final_norm", + ) + if is_pipeline_last_stage(): - self.decoder.final_layernorm = DMRegistry.convert(self.decoder.final_layernorm) - self.decoder.final_layernorm.num_features = hidden_size + setattr( + self.decoder, + self.final_norm_attr_name, + DMRegistry.convert(getattr(self.decoder, self.final_norm_attr_name)), + ) + getattr(self.decoder, self.final_norm_attr_name).num_features = hidden_size self.output_layer = DMRegistry.convert(self.output_layer) self.output_layer.input_size = hidden_size self.output_layer.get_hparam("output_size").choices = [self.output_layer.output_size] @@ -775,13 +776,20 @@ class _DynamicGPTModel(DynamicModule): self._register_temp_attribute("_activations", {}) self.hook_handles = [] for layer in self.decoder.layers: - self.hook_handles.append( - layer.input_layernorm.register_forward_hook(self._emb_layernorm_forward_hook) - ) - self.hook_handles.append( - layer.pre_mlp_layernorm.register_forward_hook(self._emb_layernorm_forward_hook) - ) - hidden_size.register_importance(self._estimate_importance) # type: ignore[union-attr] + if isinstance(layer, TransformerLayer): + if isinstance(layer.self_attention, SelfAttention): + self.hook_handles.append( + layer.input_layernorm.register_forward_hook( + self._emb_layernorm_forward_hook + ) + ) + if isinstance(layer.mlp, MLP): + self.hook_handles.append( + layer.pre_mlp_layernorm.register_forward_hook( + self._emb_layernorm_forward_hook + ) + ) + hidden_size.register_importance(self._estimate_hidden_size_importance) # type: ignore[union-attr] def _emb_layernorm_forward_hook(self, module, input, output) -> None: """Hook to collect activations for importance estimation. @@ -804,7 +812,7 @@ class _DynamicGPTModel(DynamicModule): else: self._activations[module] += activations - def _estimate_importance(self) -> TracedHp.Importance: + def _estimate_hidden_size_importance(self) -> TracedHp.Importance: """Return the activation magnitude-based importance of the hidden_size.""" assert self._activations, "No activations collected for importance estimation." aggregated_activations = [ @@ -868,7 +876,7 @@ class _DynamicGPTModel(DynamicModule): # sort layers by scores and drop the lowest ones sorted_layers = sorted(layer_scores.items(), key=lambda x: x[1], reverse=True) layers_to_drop = [layer for layer, _ in sorted_layers[num_layers_hp.active :]] # type: ignore[misc] - drop_mcore_gpt_layers(self, layers_to_drop=layers_to_drop) + drop_mcore_language_model_layers(self, layers_to_drop=layers_to_drop) def export(self) -> torch.nn.Module: """Export the dynamic module to a torch.nn.Module.""" @@ -886,7 +894,7 @@ class _DynamicGPTModel(DynamicModule): for layer in self.decoder.layers: layer.export() if is_pipeline_last_stage(): - self.decoder.final_layernorm.export() + getattr(self.decoder, self.final_norm_attr_name).export() self.output_layer.export() super().export() return self @@ -898,23 +906,27 @@ class _DynamicGPTModel(DynamicModule): layer.freeze() -def drop_mcore_gpt_layers(model: nn.Module, *, layers_to_drop: list[int]) -> None: +def drop_mcore_language_model_layers(model: nn.Module, *, layers_to_drop: list[int]) -> None: """Remove given layers (1-indexed) of the model (works with TP and/or PP). - If model is a wrapper around GPTModel, we unwrap it to get the actual GPTModel. + If model is a wrapper around GPTModel or MambaModel, it will be unwrapped. """ - # NOTE: If this function is invoked from _DynamicGPTModel during export, model.config.num_layers is already updated + # NOTE: If this function is invoked from _DynamicMCoreLanguageModel during export, + # model.config.num_layers is already updated layers_to_drop = sorted(layers_to_drop) assert layers_to_drop[0] >= 1, ( f"Layers to drop should be in range 1 to {model.config.num_layers}, got {layers_to_drop}." ) + supported_model_types = tuple(SUPPORTED_MODELS.keys()) for m in model.modules(): - if isinstance(m, GPTModel): + if isinstance(m, supported_model_types): model = m break - assert isinstance(model, GPTModel), f"Model should have {GPTModel} submodule, got {model}" - print_rank_0(f"Dropping layers {layers_to_drop} from {GPTModel}.") + assert isinstance(model, supported_model_types), ( + f"Model should have one of {supported_model_types} submodule, got {model}" + ) + print_rank_0(f"Dropping layers {layers_to_drop} from {type(model)}.") # get the number of layers remaining in each pp rank layers_remaining_per_pp = torch.zeros( @@ -961,6 +973,15 @@ def drop_mcore_gpt_layers(model: nn.Module, *, layers_to_drop: list[int]) -> Non model.config.num_layers = new_num_layers +def drop_mcore_gpt_layers(model: nn.Module, *, layers_to_drop: list[int]) -> None: + """[DEPRECATED] Remove given layers (1-indexed) of the model (works with TP and/or PP).""" + warn( + "`drop_mcore_gpt_layers` is deprecated in favor of `drop_mcore_language_model_layers`.", + DeprecationWarning, + ) + drop_mcore_language_model_layers(model, layers_to_drop=layers_to_drop) + + class MegatronConstraintsFunc(ConstraintsFunc): """A Functor class to check if sub-net satisfied all provided constraints. diff --git a/modelopt/torch/nas/search_space.py b/modelopt/torch/nas/search_space.py index 9499f1e9d..1984d8b9e 100644 --- a/modelopt/torch/nas/search_space.py +++ b/modelopt/torch/nas/search_space.py @@ -24,7 +24,8 @@ import torch.nn as nn from modelopt.torch.opt import RulesDict from modelopt.torch.opt.dynamic import DynamicModule, DynamicSpace, _DMRegistryCls from modelopt.torch.trace import Symbol, SymMap, analyze_symbols -from modelopt.torch.utils import random +from modelopt.torch.utils import print_rank_0, random +from modelopt.torch.utils.distributed import is_master, rank from .registry import DMRegistry from .traced_hp import TracedHp, TracedHpRegistry @@ -127,12 +128,17 @@ class SearchSpace(DynamicSpace): return self.config() @torch.no_grad() - def sort_parameters(self, hps_to_sort: set[str] | None = None) -> None: + def sort_parameters(self, hps_to_sort: set[str] | None = None, verbose: bool = False) -> None: """A graph propagation based parameter sorting algorithm. Args: hps_to_sort: A set of hparam names to sort. If not provided or empty, all hparams will be sorted. + verbose: Whether to print the search space and hparam importances. """ + print_rank_0("Sorting parameters...") + if verbose: + self.print_summary() + # get config and set to max config = self.config() self.sample(sample_func=max) @@ -149,6 +155,8 @@ class SearchSpace(DynamicSpace): # compute order from importance and enforce it order = torch.argsort(importance, descending=True) hp.enforce_order(order) + if verbose: + print(f"Sorted {name} for rank {rank()} with {importance=}") # now that we have enforced an order we can force reassign all parameters/buffers! for _, mod in self.named_dynamic_modules(): @@ -163,7 +171,9 @@ class SearchSpace(DynamicSpace): def print_summary(self, skipped_hparams: list[str] = ["kernel_size"]) -> None: """Print a summary of the search space.""" - print("\nSearch Space Summary:\n{:-^100}".format("")) + if not is_master(): + return + print(f"\nSearch Space Summary for rank {rank()}:\n{'-' * 100}") hp_visited = set() # Only highlight configurable hparams once for name, hp in self.named_hparams(): if not any(name.endswith(s) for s in skipped_hparams): diff --git a/modelopt/torch/nas/utils.py b/modelopt/torch/nas/utils.py index 9ea3eec93..56b103fec 100644 --- a/modelopt/torch/nas/utils.py +++ b/modelopt/torch/nas/utils.py @@ -184,13 +184,17 @@ def _reset_before_sample(model: nn.Module): PatchManager.get_manager(model).reset_before_sample() -def sort_parameters(model: nn.Module, hps_to_sort: set[str] | None = None) -> None: +def sort_parameters( + model: nn.Module, hps_to_sort: set[str] | None = None, verbose: bool = False +) -> None: """Sort the parameters of the model according to the stored importances. Args: + model: A model that contains DynamicModule(s). hps_to_sort: A set of hparam names to sort. If not provided or empty, all hparams will be sorted. + verbose: Whether to print the search space and hparam importances. """ - _SearchSpaceUnwrapped(model).sort_parameters(hps_to_sort) + _SearchSpaceUnwrapped(model).sort_parameters(hps_to_sort, verbose) def print_search_space_summary( diff --git a/modelopt/torch/opt/config.py b/modelopt/torch/opt/config.py index 12efbcf61..a4899bbd3 100644 --- a/modelopt/torch/opt/config.py +++ b/modelopt/torch/opt/config.py @@ -84,10 +84,10 @@ class ModeloptBaseConfig(BaseModel): """Get the field name from the given key (can be name or alias of field).""" assert isinstance(key, str), f"key must be a string, got {type(key)}" - if key in self.model_fields or key in self._iterable_model_extra: + if key in type(self).model_fields or key in self._iterable_model_extra: return key else: - for name, field_info in self.model_fields.items(): + for name, field_info in type(self).model_fields.items(): if field_info.alias == key: return name raise AttributeError(f"Key {key} not found in the config.") @@ -121,7 +121,7 @@ class ModeloptBaseConfig(BaseModel): def __iter__(self) -> Iterator[str]: """Iterate over aliases (or name if alias is not defined) of fields.""" - for field_name, field_info in self.model_fields.items(): + for field_name, field_info in type(self).model_fields.items(): yield field_info.alias or field_name yield from self._iterable_model_extra diff --git a/modelopt/torch/opt/dynamic.py b/modelopt/torch/opt/dynamic.py index a74680a24..8b355b947 100644 --- a/modelopt/torch/opt/dynamic.py +++ b/modelopt/torch/opt/dynamic.py @@ -1204,11 +1204,15 @@ class DynamicSpace: """ return any(True for _ in self.named_dynamic_modules()) - def named_hparams(self, configurable: bool | None = None) -> Iterator[tuple[str, Hparam]]: + def named_hparams( + self, configurable: bool | None = None, unique: bool | None = None + ) -> Iterator[tuple[str, Hparam]]: """Recursively yield the name and instance of *all* hparams. Args: configurable: Whether to include configurable hparams. + unique: Whether to include unique hparams. If ``configurable`` is set to ``True``, + then ``unique`` will be set to ``True`` by default. Yields: (name, Hparam): tuple containing the name and hparam. @@ -1218,9 +1222,13 @@ class DynamicSpace: """ _memo = set() assert configurable in [None, True], "Only all or configurable hparams are supported!" + assert unique in [None, True], "Only all or unique hparams are supported!" + if configurable: + unique = True + for mod_name, mod in self.named_dynamic_modules(): for hp_name, hp in mod.named_hparams(configurable=configurable): - if configurable is None or hp not in _memo: + if unique is None or hp not in _memo: yield mod_name + ("." if mod_name else "") + hp_name, hp _memo.add(hp) diff --git a/modelopt/torch/opt/hparam.py b/modelopt/torch/opt/hparam.py index d7457b25f..c05046308 100644 --- a/modelopt/torch/opt/hparam.py +++ b/modelopt/torch/opt/hparam.py @@ -137,6 +137,9 @@ class Hparam: else: assert val_set == curr, f"Cannot update choices: current {curr}, new: {val_set}" + def reset_choices(self) -> None: + """Reset the choices of the hparam.""" + @property def min(self) -> HPType: """Return min value from among choices.""" diff --git a/modelopt/torch/opt/plugins/__init__.py b/modelopt/torch/opt/plugins/__init__.py index b86ef1eb7..79c4367fb 100644 --- a/modelopt/torch/opt/plugins/__init__.py +++ b/modelopt/torch/opt/plugins/__init__.py @@ -19,6 +19,9 @@ from modelopt.torch.utils import import_plugin from .huggingface import * +with import_plugin("megatron core model config"): + from .megatron_model_config import * + with import_plugin("megatron core dist checkpointing"): from .mcore_dist_checkpointing import * diff --git a/modelopt/torch/opt/plugins/huggingface.py b/modelopt/torch/opt/plugins/huggingface.py index 4a911e7a9..5f8c16e71 100644 --- a/modelopt/torch/opt/plugins/huggingface.py +++ b/modelopt/torch/opt/plugins/huggingface.py @@ -64,7 +64,15 @@ def _get_modelopt_state_path(model_name_or_path: str) -> str: def _patch_model_init_for_modelopt(cls, model_path, extra_context=None): """Patch for `cls.init` method to restore ModelOpt state after `init`.""" # Note: Keeping original config in local as the package will be shared among threads - _original__init__ = cls.__init__ + added_original_init = False + if hasattr(cls, "original_init"): + _original__init__ = cls.original_init + else: + _original__init__ = cls.__init__ + cls.original_init = _original__init__ + # Avoid patching the init method twice, which can happen if one model is wrapped in another + # e.g. in the case of distillation + added_original_init = True @functools.wraps(_original__init__) def new_init_fn(self, *args, **kwargs): @@ -81,6 +89,8 @@ def _patch_model_init_for_modelopt(cls, model_path, extra_context=None): try: yield finally: + if added_original_init: + delattr(cls, "original_init") cls.__init__ = _original__init__ diff --git a/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py b/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py index 69e331f60..10c2eff05 100644 --- a/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py +++ b/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py @@ -27,7 +27,6 @@ from megatron.core.dist_checkpointing.serialization import get_default_load_shar from megatron.core.dist_checkpointing.strategies.common import COMMON_STATE_FNAME from megatron.core.dist_checkpointing.validation import StrictHandling from megatron.core.transformer.module import Float16Module -from packaging.version import Version import modelopt import modelopt.torch.opt as mto @@ -140,6 +139,9 @@ def _load_extra_state_from_sharded_checkpoint( ) -> None: """Load extra state from sharded checkpoint. + Note: since extra_state is a subset of full the sharded_state_dict, we use + strict=StrictHandling.LOG_UNEXPECTED instead of LOG_ALL. + Args: model: the model to load extra state into checkpoint_name: the checkpoint folder path @@ -151,7 +153,7 @@ def _load_extra_state_from_sharded_checkpoint( extra_sharded_state_dict, checkpoint_name, get_default_load_sharded_strategy(checkpoint_name), - strict=StrictHandling.LOG_ALL, + strict=StrictHandling.LOG_UNEXPECTED, ) extra_state_dict_no_prefix = {} @@ -193,35 +195,21 @@ def restore_sharded_modelopt_state( print(f"nvidia-modelopt ckpt/inst version: {modelopt_load_version}/{modelopt.__version__}") - is_dev_load_ver = "dev" in modelopt_load_version - is_legacy_load_ver = Version(modelopt_load_version) <= Version("0.30.0") + # After 0.29, we no longer store (or shard) any quantizer_state in the modelopt_state. + # quantizer_state (or other per-module state) is stored with the main distributed + # checkpoint as extra_state at the QuantModule level. + # + # The process of resuming modelopt_state becomes 2-phase: + # 1. Load the global modelopt_state and call mto.restore_from_modelopt_state. + # Modes are restored in order. Modes with per-module state stored as + # extra_state are partially restored (stop at DynamicModule replacement) + # + model[0] = mto.restore_from_modelopt_state(model[0], common_modelopt_state) - if is_legacy_load_ver and not is_dev_load_ver: - raise ValueError( - "nvidia-modelopt>0.29 updated how model_state is stored NeMo-MCore " - "distributed checkpoint (`torch-dist`). Newly generated checkpoints " - "can no longer be loaded with older nvidia-modelopt<=0.29." - "Old checkpoints also cannot be loaded. To convert the checkpoint, use" - "legacy modelopt to load and store in `torch` instead of `torch-dist`." - "Load the `torch` checkpoint with nvidia-modelopt>0.29 and store" - "`torch-dist` again to complete." - ) - else: - # After 0.29, we no longer store (or shard) any quantizer_state in the modelopt_state. - # quantizer_state (or other per-module state) is stored with the main distributed - # checkpoint as extra_state at the QuantModule level. - # - # The process of resuming modelopt_state becomes 2-phase: - # 1. Load the global modelopt_state and call mto.restore_from_modelopt_state. - # Modes are restored in order. Modes with per-module state stored as - # extra_state are partially restored (stop at DynamicModule replacement) - # - model[0] = mto.restore_from_modelopt_state(model[0], common_modelopt_state) - - try: - _load_extra_state_from_sharded_checkpoint(model[0], checkpoint_name, prefix) - except: # noqa: E722 - # [WAR]: nemo2 is calling this function with an empty prefix. - # The prefix however should be `module.` instead. This should be fixed - # from the NeMo side. This is just a WAR. - _load_extra_state_from_sharded_checkpoint(model[0], checkpoint_name, "module.") + try: + _load_extra_state_from_sharded_checkpoint(model[0], checkpoint_name, prefix) + except: # noqa: E722 + # [WAR]: nemo2 is calling this function with an empty prefix. + # The prefix however should be `module.` instead. This should be fixed + # from the NeMo side. This is just a WAR. + _load_extra_state_from_sharded_checkpoint(model[0], checkpoint_name, "module.") diff --git a/modelopt/torch/opt/plugins/megatron_model_config.py b/modelopt/torch/opt/plugins/megatron_model_config.py new file mode 100644 index 000000000..c54202915 --- /dev/null +++ b/modelopt/torch/opt/plugins/megatron_model_config.py @@ -0,0 +1,62 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Megatron-Core model config (TransformerConfig+).""" + +from collections.abc import Callable +from dataclasses import dataclass + +import torch.nn.functional as F +from megatron.core.transformer.transformer_config import TransformerConfig + + +@dataclass +class Llama31Config8B(TransformerConfig): + """Configuration class for GPT models. + + Extends TransformerConfig with additional parameters specific to GPT models + and provides utility methods for model configuration. + """ + + # From megatron.core.models.gpt.gpt_model.GPTModel + transformer_layer_spec = None + vocab_size: int = None + max_sequence_length: int = 8192 + position_embedding_type = "rope" + rotary_percent: float = 1.0 + rotary_base: int = 500000 + rope_scaling: bool = True + rope_scaling_factor: float = 8.0 + + # Specific TransformerConfig + seq_length: int = 8192 + num_layers: int = 32 + hidden_size: int = 4096 + ffn_hidden_size: int = 14336 + kv_channels: int = 128 + num_attention_heads: int = 32 + num_query_groups: int = 8 + init_method_std: float = 0.01 + normalization: str = "RMSNorm" + layernorm_epsilon: float = 1.0e-05 + activation_func: Callable = F.silu + gated_linear_unit: bool = True + add_bias_linear: bool = False + attention_dropout: float = 0.0 + hidden_dropout: float = 0.0 + + # Different from the default values in TransformerConfig + attention_softmax_in_fp32: bool = False + gradient_accumulation_fusion: bool = False diff --git a/modelopt/torch/opt/utils.py b/modelopt/torch/opt/utils.py index e893ad58c..88db2a958 100644 --- a/modelopt/torch/opt/utils.py +++ b/modelopt/torch/opt/utils.py @@ -58,10 +58,10 @@ def is_dynamic(model: nn.Module) -> bool: def named_hparams( - model: nn.Module, configurable: bool | None = None + model: nn.Module, configurable: bool | None = None, unique: bool | None = None ) -> Generator[tuple[str, Hparam], None, None]: """Recursively yield the name and instance of *all* hparams.""" - yield from _DynamicSpaceUnwrapped(model).named_hparams(configurable) + yield from _DynamicSpaceUnwrapped(model).named_hparams(configurable, unique) def named_dynamic_modules(model: nn.Module) -> Generator[tuple[str, DynamicModule], None, None]: diff --git a/modelopt/torch/prune/plugins/mcore_gpt_minitron.py b/modelopt/torch/prune/plugins/mcore_gpt_minitron.py index de2cf9917..b24a53b17 100644 --- a/modelopt/torch/prune/plugins/mcore_gpt_minitron.py +++ b/modelopt/torch/prune/plugins/mcore_gpt_minitron.py @@ -35,6 +35,7 @@ from modelopt.torch.nas.utils import sort_parameters from modelopt.torch.opt.config import ModeloptBaseConfig, get_kwargs_for_create_model_with_rules from modelopt.torch.opt.searcher import BaseSearcher, SearchConfig, SearchStateDict from modelopt.torch.opt.utils import named_hparams +from modelopt.torch.utils import print_rank_0 from ..fastnas import FastNASModeDescriptor from ..pruning import PruneModeRegistry @@ -60,6 +61,13 @@ def get_supported_model_config_map() -> dict[type, str]: except Exception: pass + try: + from megatron.core.models.mamba import MambaModel + + supported_model_config_map[MambaModel] = "config" + except Exception: + pass + try: from nemo.collections import llm from nemo.collections.nlp.models.language_modeling.megatron_gpt_model import ( @@ -67,6 +75,7 @@ def get_supported_model_config_map() -> dict[type, str]: ) supported_model_config_map[MegatronGPTModel] = "cfg" + # NOTE: llm.MambaModel is a subclass of llm.GPTModel supported_model_config_map[llm.GPTModel] = "config" except Exception: pass @@ -126,12 +135,13 @@ class MCoreGPTMinitronSearcher(BaseSearcher): ) self.hps_to_sort.add("num_heads_per_group") - for n, hp in named_hparams(self.model, configurable=True): + for n, hp in named_hparams(self.model, unique=True): hp_name = n.split(".")[-1] - if hp_name in export_config: + if hp.is_configurable and hp_name in export_config: assert export_config[hp_name] in hp.choices, ( - f"Invalid choice for {hp_name}! Available choices: {hp.choices}" + f"Invalid choice {export_config[hp_name]} for {n}! Available choices: {hp.choices}" ) + hp.reset_choices() # Make sure ConcatHparam choices are updated after modify() def run_search(self) -> None: """Run actual search.""" @@ -150,9 +160,10 @@ class MCoreGPTMinitronSearcher(BaseSearcher): assert self.forward_loop is not None is_training = self.model.training self.model.eval() + print_rank_0("Running forward loop...") with torch.no_grad(): self.forward_loop(self.model) - sort_parameters(self.model, self.hps_to_sort) + sort_parameters(self.model, self.hps_to_sort, verbose=True) self.model.train(is_training) # Prune homogeneously @@ -186,6 +197,12 @@ MCoreGPTMinitronConfig: type[ModeloptBaseConfig] = create_model( "num_query_groups_divisor": 1, "ffn_hidden_size_divisor": 64, }, + "megatron.core.models.mamba.MambaModel": { + "hidden_size_divisor": 64, + "num_heads_per_group_divisor": 1, + "num_query_groups_divisor": 1, + "ffn_hidden_size_divisor": 64, + }, }, doc='Configuration for the ``"mcore_gpt_minitron"`` mode.', ), diff --git a/modelopt/torch/quantization/algorithms.py b/modelopt/torch/quantization/algorithms.py new file mode 100644 index 000000000..b99311e6f --- /dev/null +++ b/modelopt/torch/quantization/algorithms.py @@ -0,0 +1,671 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Module for advanced quantization algorithms.""" + +import gc +import hashlib +import json +import types +import warnings +from collections import defaultdict +from collections.abc import Callable, Sequence +from typing import Any + +import regex as re +import torch +import torch.distributed +import torch.nn as nn +from tqdm import tqdm + +from modelopt.torch.opt.conversion import ModeloptStateManager +from modelopt.torch.opt.hparam import CustomHPType, Hparam, HPType +from modelopt.torch.opt.searcher import LPS, BaseSearcher, SearchConfig, SearchStateDict +from modelopt.torch.opt.utils import get_hparam, named_hparams +from modelopt.torch.utils import create_param_grad_clear_hook, print_rank_0, report_memory +from modelopt.torch.utils.distributed import DistributedProcessGroup, is_master + +from . import model_calib +from .config import FP8_DEFAULT_CFG, NVFP4_DEFAULT_CFG, QuantizeConfig, QuantizerAttributeConfig +from .conversion import set_quantizer_by_cfg +from .nn import QuantLinearConvBase, SequentialQuantizer, TensorQuantizer +from .utils import is_quantized_linear, multi_context + + +def estimate_quant_compression(quant_cfg: QuantizeConfig) -> float: + """Estimate the compression ratio of a quantization configuration. + + Right now, we find the minimum compression ratio across all quantizer attribute configs. + This is not perfect but is a good proxy for the overall compression ratio. We will improve + this in future releases. + + Args: + quant_cfg: The quantization configuration to estimate compression for. + + Returns: + float: The estimated compression ratio (0.0 to 1.0). + """ + + def estimate_quant_compression_for_quantizer(quantizer_attr_cfg): + if isinstance(quantizer_attr_cfg, list): + return min(estimate_quant_compression_for_quantizer(q) for q in quantizer_attr_cfg) + if isinstance(quantizer_attr_cfg, dict): + return estimate_quant_compression_for_quantizer(list(quantizer_attr_cfg.values())) + + if isinstance(quantizer_attr_cfg, QuantizerAttributeConfig): + if not quantizer_attr_cfg.enable: + return 1.0 + if not hasattr(quantizer_attr_cfg, "num_bits"): + return 1.0 + if isinstance(quantizer_attr_cfg.num_bits, tuple): + return (sum(quantizer_attr_cfg.num_bits) + 1) / 16 + elif isinstance(quantizer_attr_cfg.num_bits, int): + return quantizer_attr_cfg.num_bits / 16 + else: + raise ValueError(f"Unknown quantization config {quantizer_attr_cfg.num_bits}") + + raise ValueError(f"Unknown type {type(quantizer_attr_cfg)}, {quantizer_attr_cfg}") + + return estimate_quant_compression_for_quantizer(list(quant_cfg.quant_cfg.values())) + + +class QuantRecipe(CustomHPType): + """A subclass of QuantizeConfig enabling auto_quantize specific configurations.""" + + def __init__( + self, quant_cfg: dict[str, Any] | None = None, quant_format_idx: int | None = None + ): + """Initialize the QuantRecipe with the quantization configuration.""" + if quant_cfg is None: + self.config = QuantizeConfig(quant_cfg={"*": {"enable": False}}, algorithm="max") + else: + self.config = QuantizeConfig(**quant_cfg) + + # Disable KV Cache quantization + # Currently KV Cache quantization is enabled for some quantization formats and disabled for others + # This breaks the monotonicity of the quantization formats in terms of weight compression Vs accuracy + self.config.quant_cfg["*output_quantizer"] = QuantizerAttributeConfig(enable=False) + + self.compression = estimate_quant_compression(self.config) + + self.str_repr = ( + f"quantization_formats[{quant_format_idx}]:effective-bits-{self.compression * 16}" + ) + + @property + def num_bits(self) -> int: + """Get the number of bits for the quantization format.""" + return int(self.compression * 16) + + def __str__(self) -> str: + return f"{self.str_repr}" + + def __repr__(self) -> str: + return f"{self.config}" + + def __lt__(self, other: "QuantRecipe"): + return self.compression < other.compression + + def __eq__(self, other: object): + assert isinstance(other, QuantRecipe) + return self.config == other.config + + def __hash__(self) -> int: + sorted_json = json.dumps(json.loads(self.config.model_dump_json()), sort_keys=True) + return int(hashlib.md5(sorted_json.encode("utf-8"), usedforsecurity=False).hexdigest(), 16) + + @staticmethod + def disable_folding_pqs_to_weights(): + """Disable the folding of pre_quant_scale to weights.""" + model_calib._ENABLE_FOLDING_PQS_TO_WEIGHTS = False + + @staticmethod + def fold_pqs_to_weights(model): + """Fold the pre_quant_scale in weight_quantizers to weights.""" + model_calib._ENABLE_FOLDING_PQS_TO_WEIGHTS = True + for name, module in model.named_modules(): + if is_quantized_linear(module): + with SequentialQuantizer.convert_to_single_quantizer(model): + if module.weight_quantizer.pre_quant_scale is not None: + weight_pqs = module.weight_quantizer.pre_quant_scale + delattr(module.weight_quantizer, "_pre_quant_scale") + model_calib._apply_weight_pre_quant_scale(module, weight_pqs) + + +class QuantRecipeHparam(Hparam): + """An Hparam for quantization recipes. + + In addition, this Hparam also: + 1. Keeps a link to its modules and sets the quantizers for the module based on the active recipe. + 2. Keeps track of the importance of each recipe in a dict instead of a tensor + """ + + def __init__( + self, + choices: Sequence[QuantRecipe], + original: QuantRecipe | None = None, + nn_modules: list[nn.Module] | None = None, + ) -> None: + """Initializes Hparam with original value and choices.""" + choices = sorted(set(choices) | {QuantRecipe(quant_cfg=None)}) + super().__init__(choices, original) + self.nn_modules = nn_modules if nn_modules else [] + + # This is a hack; We dont want to make the input_quantizer, weight_quantizer, output_quantizer + # a dynamic attribute for backward compatibility with the model_calib.py + # TODO: Make input_quantizer, weight_quantizer, output_quantizer a dynamic attribute and get rid of this hack + self._all_quantizer_choices = {quant_recipe: {} for quant_recipe in self.choices} + + quant_recipe: QuantRecipe + for quant_recipe in self.choices: + for nn_module in self.nn_modules: + for quantizer_attr_name in [ + "input_quantizer", + "weight_quantizer", + "output_quantizer", + ]: + setattr(nn_module, quantizer_attr_name, TensorQuantizer()) + + set_quantizer_by_cfg(nn_module, quant_recipe.config.quant_cfg) + self._all_quantizer_choices[quant_recipe][nn_module] = { + quantizer_attr_name: getattr(nn_module, quantizer_attr_name) + for quantizer_attr_name in [ + "input_quantizer", + "weight_quantizer", + "output_quantizer", + ] + } + + self.active = self.original + + self._importance_dict = { + quant_recipe: dict.fromkeys(self.nn_modules, 0.0) for quant_recipe in self.choices + } + + @property + def active(self) -> HPType: + """Return the currently active value.""" + return self._active + + @active.setter + def active(self, val: HPType | None): + """Set the active value with a sanity check for choices and dynamic hparams.""" + val = self.original if val is None else val + assert val in self._choices, f"val = {val}, choices = {self.choices}" + if self.is_configurable: + self._active = val + else: + assert self._active == val + + for nn_module, quantizer_choices in self._all_quantizer_choices[val].items(): + for quantizer_attr_name, quantizer in quantizer_choices.items(): + setattr(nn_module, quantizer_attr_name, quantizer) + + @property + def importance(self) -> dict: + """Return the importance dict mapping recipe and importance.""" + return { + quant_recipe: sum(importance_dict.values()) + for quant_recipe, importance_dict in self._importance_dict.items() + } + + +class AutoQuantizeSearcher(BaseSearcher): + """A searcher for AutoQuantize algorithm. + + In AutoQuantize, we search for the best per-layer quantization configuration that minimizes the sum of per-layer + scores while meeting the specified constraint. AutoQuantize uses Linear Programming Solver to find the + optimal quantization configuration. + + The auto_quantize score for a layer quantization configuration is an approximation of model loss change change due + to quantizing the particular layer with the particular configuration. + The approximation is based on taylor expansion of the loss function wrt to the quantized output of the layer and + substitution of Fisher information for Hessian. + This approximation is mathematically correct for models where the loss + is a log likelihood loss such as BERT, GPT, etc. However, the auto_quantize score can still be used as a proxy + for other models such as ResNet. + """ + + candidate_stats: dict[str, dict[str, list[float]]] + best: dict[str, Any] + gradient_checkpointing_enable_contexts: list[tuple[Callable, Callable]] = [] + + rules = [ + r"^(.*?)\.(q_proj|k_proj|v_proj)$", # q_proj, k_proj, v_proj for llama like models + r"^(.*?)\.(gate_proj|up_proj)$", # gate_proj, up_proj for llama like models + r"^(.*?)\.(\d+\.(w1|w2|w3))$", # mixtral experts + r"^(.*?)\.((w1_linear|w2_linear|w3_linear)\.\d+)$", # dbrx experts + ] + + @property + def default_search_config(self): + """Get the default config for the searcher.""" + return { + "quantization_formats": [NVFP4_DEFAULT_CFG, FP8_DEFAULT_CFG], + "data_loader": None, + "forward_step": None, + "loss_func": None, + "forward_backward_step": None, + "num_calib_steps": 512, + "num_score_steps": 128, + "deployment": None, + "verbose": is_master(), + "checkpoint": None, + } + + @property + def default_state_dict(self) -> SearchStateDict: + """Get the default state dict for AutoQuantize.""" + return { + "candidate_stats": defaultdict(dict), + "best": {"recipe": {}, "constraints": {}, "score": float("inf"), "is_satisfied": False}, + } + + def sanitize_search_config(self, config: SearchConfig | None) -> SearchConfig: + """Sanitize the search config dict.""" + config = config or {} + if "score_func" in config: + warnings.warn("`score_func` is ignored for `auto_quantize`.") + config.pop("score_func") + config = super().sanitize_search_config(config) + assert config["data_loader"] is not None, ( + "`data_loader` must be provided for `auto_quantize`." + ) + assert config["forward_step"] is not None, ( + "`forward_step` must be provided for `auto_quantize`." + ) + + if config["forward_backward_step"] is None: + assert config["loss_func"] is not None, ( + "`loss_func` or `forward_backward_step` must be provided for `auto_quantize`." + ) + config["forward_backward_step"] = self._get_default_forward_backward_step() + + return config + + @staticmethod + def _is_auto_quantize_module(module): + return is_quantized_linear(module) or isinstance(module, QuantLinearConvBase) + + @staticmethod + def _get_search_recipes(quantization_formats): + return sorted( + [ + QuantRecipe(quant_cfg=q, quant_format_idx=i) + for i, q in enumerate(quantization_formats) + ] + ) + + @classmethod + def register_gradient_checkpointing_enable_context( + cls, is_supported_checker: Callable, context: Callable + ): + """Register a gradient checkpointing enable context for `AutoQuantize` score estimation. + + If the `is_supported_checker(model)` returns True, the `context(model)` will be used to enable gradient + checkpointing. + """ + cls.gradient_checkpointing_enable_contexts.append((is_supported_checker, context)) + + def _get_default_forward_backward_step(self): + def forward_backward_step(model, data): + output = self.config["forward_step"](model, data) + loss = self.config["loss_func"](output, data) + try: + loss.backward() + except RuntimeError as e: + raise RuntimeError( + "AutoQuantize: Error while calling `backward()` on the loss returned by `loss_func`. " + "Please fix this!" + ) from e + + return forward_backward_step + + @torch.enable_grad() + def _estimate_auto_quantize_scores(self): + # TODO: remove the no-quant recipe + def auto_quantize_score_estimate_forward(module, input, *args, **kwargs): + module.quant_recipe = QuantRecipe(quant_cfg=None, quant_format_idx=None) + output = module._forward_original(input, *args, **kwargs) + + # If gradient checkpointing is enabled, gradient will not be enabled in the global forward pass. + # With gradient checkpointing, gradients are computed in the local forward pass during backward pass + + # Lets compute the output_diff and save it in memory only if gradient is enabled to be memory efficient + if not torch.is_grad_enabled(): + return output + + module.output_diff_dict = {} + with torch.no_grad(): + for recipe in module.get_hparam("quant_recipe").choices: + if recipe.compression >= 1.0: + continue + module.quant_recipe = recipe + output_diff = module._forward_original(input, *args, **kwargs) + + if isinstance(output_diff, tuple): + output_diff = output_diff[0] - output[0] + else: + output_diff -= output + module.output_diff_dict[recipe] = output_diff + + return output + + def backward_hook(module, grad_input, grad_output): + for recipe, output_diff in module.output_diff_dict.items(): + score = ((grad_output[0].float() ** 2) * (output_diff.float() ** 2)).sum() + module.get_hparam("quant_recipe")._importance_dict[recipe][module] += score.item() + module.output_diff_dict[recipe] = None + + del module.output_diff_dict + + def setup_params_for_score_estimation(name, param, params_metadata): + params_metadata[name] = {"requires_grad": param.requires_grad} + param.requires_grad = True + accum_grad, handle = create_param_grad_clear_hook(param) + params_metadata[name]["accum_grad"] = accum_grad # We need to keep the accum_grad alive + params_metadata[name]["handle"] = handle + + def setup_module_for_score_estimation(module): + module._forward_original = module.forward + module.forward = types.MethodType(auto_quantize_score_estimate_forward, module) + module._backward_hook_handle = module.register_full_backward_hook(backward_hook) + + def cleanup_module_after_score_estimation(module): + module.forward = module._forward_original + del module._forward_original + + module._backward_hook_handle.remove() + + def cleanup_params_after_score_estimation(name, param, params_metadata): + param.requires_grad = params_metadata[name]["requires_grad"] + params_metadata[name]["handle"].remove() + + for name, module in self.model.named_modules(): + if self._is_auto_quantize_module(module): + # Monkey patch the forward methods to cache Y(Q(W), Q(X)) - Y(W,X) + setup_module_for_score_estimation(module) + + params_metadata = {} + for name, param in self.model.named_parameters(): + # Let us delete the gradient as soon as they are computed to save memory + # In addition, this method enables gradient for all parameters + # This is needed to make sure the re-entrant activation checkpointing works + setup_params_for_score_estimation(name, param, params_metadata) + + gc.collect() + if torch.cuda.is_available(): + torch.cuda.reset_peak_memory_stats() + report_memory("AutoQuantize: starting score estimation, ") + + self._run_func( + self.config["forward_backward_step"], + num_iters=self.config["num_score_steps"], + desc="Estimating auto_quantize scores", + ) + + if torch.cuda.is_available(): + report_memory("AutoQuantize: After score estimation") + + for name, module in self.model.named_modules(): + if self._is_auto_quantize_module(module): + cleanup_module_after_score_estimation(module) + + for name, param in self.model.named_parameters(): + cleanup_params_after_score_estimation(name, param, params_metadata) + + # Delete the params_metadata + del params_metadata + gc.collect() + + @classmethod + def insert_hparams_after_merge_rules(cls, model, quant_recipes): + """Restrict the search space using the merge rules and insert the hparams for the model.""" + # TRTLLM fuses linear layers such as q_proj, k_proj, v_proj into same layer + # Hence we need to restrict the search space so that all these layers share the same recipe + # Lets group the modules based on the rules and insert the same hparam for all the modules in the group + search_map: dict[str, list[nn.Module]] = {} + for name, module in model.named_modules(): + if not cls._is_auto_quantize_module(module): + continue + prefix = name + for rule in cls.rules: + pattern = re.compile(rule) + match = pattern.match(name) + if match: + prefix = match.group(1) + # We support only one rule for matching per module + break + if prefix not in search_map: + search_map[prefix] = [module] + else: + search_map[prefix].append(module) + + for prefix, modules in search_map.items(): + hparam = QuantRecipeHparam( + quant_recipes, + original=quant_recipes[0], + nn_modules=modules, + ) + for module in modules: + module._register_hparam("quant_recipe", hparam) + + def _get_formatted_weight_compression_constraint(self): + effective_bits = self.constraints["effective_bits"] + assert effective_bits > 0 and effective_bits <= 16, ( + "effective_bits should be between 0 and 16." + ) + weight_compression = self.constraints["effective_bits"] / 16.0 + + return weight_compression + + def _verify_constraint(self, search_recipes): + assert self.constraints["effective_bits"] >= search_recipes[0].num_bits, ( + f"The effective_bits {self.constraints['effective_bits']} constraint cannot be lower than the " + f"num_bits of most aggressive quantization format for this search which is " + f"{search_recipes[0]} whose num_bits = {search_recipes[0].num_bits}." + ) + + def _run_func(self, func, num_iters=1, desc=""): + for i, data in tqdm( + zip(range(num_iters), self.config["data_loader"]), + desc=desc, + total=num_iters, + ): + func(self.model, data) + + def before_search(self): + """Prepare the model for search by calibrating the quantizers and collecting ``AutoQuantize`` score.""" + # Import here to avoid circular import + from modelopt.torch.quantization.model_quant import calibrate + + super().before_search() + + search_recipes = self._get_search_recipes(self.config["quantization_formats"]) + self._verify_constraint(search_recipes) + self.insert_hparams_after_merge_rules(self.model, search_recipes) + + QuantRecipe.disable_folding_pqs_to_weights() + + # Iterate over the search recipes and calibrate the quantizers for each recipe + for recipe in search_recipes: + if recipe.compression >= 1.0: + continue + + # Lets reduce the number of calibration steps for AWQ since it takes longer + num_calib_steps = ( + self.config["num_calib_steps"] + if "awq" not in str(recipe.config.algorithm) + else max(1, self.config["num_calib_steps"] // 4) + ) + + def forward_loop(model): + self._run_func( + self.config["forward_step"], + num_iters=num_calib_steps, + desc=f"Calibrating for {recipe}", + ) + + for name, hparam in named_hparams(self.model, configurable=True): + if not isinstance(hparam, QuantRecipeHparam): + continue + hparam.active = recipe + + # Now calibrate the quantizers for the recipe + calibrate( + self.model, + algorithm=recipe.config.algorithm, + forward_loop=forward_loop, + ) + # Calibrate adds a new mode to the model. Since auto_quantize mixes the quantization recipes + # across layers, lets not save this new mode in the modelopt state. + # TODO: This is a hack. We need to create a mode for auto_quantize to handle this in a clean way. + ModeloptStateManager(self.model).state_dict().pop() + + self.model.eval() + with multi_context( + *( + context(self.model) + for is_supported_checker, context in self.gradient_checkpointing_enable_contexts + if is_supported_checker(self.model) + ) + ): + self._estimate_auto_quantize_scores() + + def run_search(self): + """Search for the best per-layer quantization configuration and return the best model and configuration. + + AutoQuantize uses Linear Programming Solver to find the optimal quantization configuration which + minimizes the sum of per-layer auto_quantize scores while meeting the specified constraint. + """ + + def get_total_weight_size(modules): + return sum( + (module.weight.numel() if self._is_auto_quantize_module(module) else 0) + for module in modules + ) + + def _get_constraints_for_search(lower_bound=None): + total_model_weight_size = get_total_weight_size(self.model.modules()) + + upper_bound = self._get_formatted_weight_compression_constraint() + + if lower_bound: + lower_bound = lower_bound * upper_bound + + constraints = { + "weight_size_after_compression": ( + lower_bound * total_model_weight_size if lower_bound else lower_bound, + upper_bound * total_model_weight_size, + ) + } + return constraints, "weight_size_after_compression" + + verbose = self.config["verbose"] + assert len(self.constraints) == 1 and "effective_bits" in self.constraints, ( + f"`constraints` must contain only 'effective_bits' constraint. " + f"Got {self.constraints.keys()}" + ) + + search_recipes = self._get_search_recipes(self.config["quantization_formats"]) + for name, hparam in named_hparams(self.model, configurable=True): + if not isinstance(hparam, QuantRecipeHparam): + continue + formats, scores, costs = [], [], [] + prev_score = float("inf") + for recipe in search_recipes: + formats.append(recipe) + score = hparam.importance[recipe] + cost = get_total_weight_size(hparam.nn_modules) * recipe.compression + + # Lets get the score across Data Parallel (DP) and Tensor Parallel (TP) groups + # This way we constraint the same quantization format for the same layer across the DP/TP groups + # The cost we use here is weight size. They are the same across DP/TP groups. + _ps = self.model.get_submodule(name.split(".quant_recipe")[0]).parallel_state + # The score is the sum of the scores across DP and TP groups + score = DistributedProcessGroup.get_dist_syncd_obj( + score, [_ps.data_parallel_group, _ps.tensor_parallel_group], sum + ) + + scores.append(min(score, prev_score)) + costs.append(cost) + prev_score = score + self.candidate_stats[name]["formats"] = formats + self.candidate_stats[name]["scores"] = scores + self.candidate_stats[name]["costs"] = costs + + for lower_bound in [None, 0.99, 0.90]: + # The LP solver for auto_quantize sometimes fails to find a solution if a lower bound is not + # specified. I dont know why this happens. + # As a workaround, lets specify a lower bound for the weight compression if previous + # search without lower bound fails. + constraints, constraint_name = _get_constraints_for_search(lower_bound) + + lps = LPS( + name="AutoQuantize", + constraints=constraints, + constraints_to_candidate_costs={ + constraint_name: [ + candidate_stat["costs"] for candidate_stat in self.candidate_stats.values() + ] + }, + candidate_scores=[ + candidate_stat["scores"] for candidate_stat in self.candidate_stats.values() + ], + objective_type="minimize", + verbose=verbose, + ) + selections, self.status = lps() + if self.status == "Optimal": + break + + self.best = {} + + if self.status != "Optimal": + warnings.warn( + "AutoQuantize FAILED to find a solution! The searched model might not meet all constraints. " + ) + self.best["is_satisfied"] = False + else: + self.best["is_satisfied"] = True + + best_recipe = {} + best_constraints, best_scores = 0, 0 + for name, selected_idx in zip(self.candidate_stats.keys(), selections): + best_recipe_for_name = self.candidate_stats[name]["formats"][selected_idx] + + # LP solver could give different solutions for the same layer across DP/TP groups even though + # the scores and costs are the same. Lets make sure the same quantization format is selected across DP/TP + _ps = self.model.get_submodule(name.split(".quant_recipe")[0]).parallel_state + best_recipe_for_name = DistributedProcessGroup.get_dist_syncd_obj( + best_recipe_for_name, + [_ps.data_parallel_group, _ps.tensor_parallel_group], + lambda a: a[0], + ) + + best_recipe[name] = best_recipe_for_name + get_hparam(self.model, name).active = best_recipe_for_name + best_constraints += self.candidate_stats[name]["costs"][selected_idx] + best_scores += self.candidate_stats[name]["scores"][selected_idx] + if verbose: + print_rank_0( + f"AutoQuantize best recipe for {name.replace('.quant_recipe', '')}: {best_recipe[name]}" + ) + + self.best["recipe"] = best_recipe + self.best["constraints"] = {constraint_name: best_constraints} + self.best["score"] = best_scores + + QuantRecipe.fold_pqs_to_weights(self.model) diff --git a/modelopt/torch/quantization/compress.py b/modelopt/torch/quantization/compress.py index 1109a5877..ccc947212 100644 --- a/modelopt/torch/quantization/compress.py +++ b/modelopt/torch/quantization/compress.py @@ -89,7 +89,6 @@ def compress_convert( raise ValueError( f"Invalid compression configuration: {to_compress}, expected a boolean as value." ) - # If real quant quantizer is present, real quantize the weights. pack_real_quantize_weight(model) diff --git a/modelopt/torch/quantization/config.py b/modelopt/torch/quantization/config.py index bc4fd1ae7..5d000ecfc 100644 --- a/modelopt/torch/quantization/config.py +++ b/modelopt/torch/quantization/config.py @@ -141,18 +141,6 @@ from typing import Literal from pydantic import ValidationInfo, field_validator, model_validator -# use the form "from x import y as y" syntax to direct mypy that these configs are re-exported -from modelopt.core.torch.quantization.config import NVFP4_AFFINE_KV_CFG as NVFP4_AFFINE_KV_CFG -from modelopt.core.torch.quantization.config import NVFP4_AWQ_CLIP_CFG as NVFP4_AWQ_CLIP_CFG -from modelopt.core.torch.quantization.config import NVFP4_AWQ_FULL_CFG as NVFP4_AWQ_FULL_CFG -from modelopt.core.torch.quantization.config import NVFP4_AWQ_LITE_CFG as NVFP4_AWQ_LITE_CFG -from modelopt.core.torch.quantization.config import NVFP4_DEFAULT_CFG as NVFP4_DEFAULT_CFG -from modelopt.core.torch.quantization.config import NVFP4_KV_CFG as NVFP4_KV_CFG -from modelopt.core.torch.quantization.config import NVFP4_KV_ROTATE_CFG as NVFP4_KV_ROTATE_CFG -from modelopt.core.torch.quantization.config import NVFP4_MXFP8_CFG as NVFP4_MXFP8_CFG -from modelopt.core.torch.quantization.config import ( - NVFP4_SVDQUANT_DEFAULT_CFG as NVFP4_SVDQUANT_DEFAULT_CFG, -) from modelopt.torch.opt.config import ModeloptBaseConfig, ModeloptField from modelopt.torch.utils.network import ConstructorLike @@ -369,25 +357,237 @@ FP8_AFFINE_KV_CFG = { "algorithm": "max", } +NVFP4_DEFAULT_CFG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + **_default_disabled_quantizer_cfg, + }, + "algorithm": "max", +} + + +NVFP4_AWQ_LITE_CFG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + **_default_disabled_quantizer_cfg, + }, + "algorithm": "awq_lite", +} + +NVFP4_AWQ_CLIP_CFG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + **_default_disabled_quantizer_cfg, + }, + "algorithm": {"method": "awq_clip"}, +} + +NVFP4_AWQ_FULL_CFG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + **_default_disabled_quantizer_cfg, + }, + "algorithm": {"method": "awq_full", "alpha_step": 0.1}, +} + + +NVFP4_AFFINE_KV_CFG = { + "quant_cfg": { + "*[kv]_bmm_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + "bias": {-2: None, -4: None, "type": "static"}, + }, + "default": {"enable": False}, + }, + "algorithm": "max", +} + +NVFP4_KV_CFG = { + "quant_cfg": { + "*[kv]_bmm_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "default": {"enable": False}, + }, + "algorithm": "max", +} + +# Moved from examples/diffusers/quantization/config.py to here +NVFP4_FP8_MHA_CONFIG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*output_quantizer": {"enable": False}, + "*q_bmm_quantizer": { + "num_bits": (4, 3), + "axis": None, + }, + "*k_bmm_quantizer": { + "num_bits": (4, 3), + "axis": None, + }, + "*v_bmm_quantizer": { + "num_bits": (4, 3), + "axis": None, + }, + "*softmax_quantizer": { + "num_bits": (4, 3), + "axis": None, + }, + "transformer_blocks*bmm2_output_quantizer": { + "num_bits": (4, 3), + "axis": None, + }, + "default": {"enable": False}, + }, + "algorithm": "max", +} + +NVFP4_KV_ROTATE_CFG = { + "quant_cfg": { + "*q_bmm_quantizer": { + "enable": False, + "rotate": True, + }, + "*k_bmm_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + "rotate": True, + }, + "*v_bmm_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + }, + "algorithm": "max", +} + +NVFP4_SVDQUANT_DEFAULT_CFG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + **_default_disabled_quantizer_cfg, + }, + "algorithm": {"method": "svdquant", "lowrank": 32}, +} + +NVFP4_FP8_CFG = { + "quant_cfg": { + "*weight_quantizer": { + "num_bits": (2, 1), + "block_sizes": {-1: 32, "type": "dynamic", "scale_bits": (4, 3)}, + "axis": None, + "enable": True, + }, + "*input_quantizer": { + "num_bits": (4, 3), + "axis": None, + "enable": True, + }, + **_default_disabled_quantizer_cfg, + }, + "algorithm": "max", +} + + choices: set[str] = { + "FP8_2D_BLOCKWISE_WEIGHT_ONLY_CFG", + "FP8_AFFINE_KV_CFG", + "FP8_DEFAULT_CFG", + "FP8_KV_CFG", + "FP8_PER_CHANNEL_PER_TOKEN_CFG", + "INT4_AWQ_CFG", + "INT4_BLOCKWISE_WEIGHT_ONLY_CFG", "INT8_DEFAULT_CFG", "INT8_SMOOTHQUANT_CFG", - "FP8_DEFAULT_CFG", - "FP8_2D_BLOCKWISE_WEIGHT_ONLY_CFG", - "FP8_PER_CHANNEL_PER_TOKEN_CFG", - "INT4_BLOCKWISE_WEIGHT_ONLY_CFG", - "INT4_AWQ_CFG", - "W4A8_AWQ_BETA_CFG", - "NVFP4_DEFAULT_CFG", - "NVFP4_AWQ_LITE_CFG", + "MXFP4_DEFAULT_CFG", + "MXFP8_DEFAULT_CFG", + "MXINT8_DEFAULT_CFG", + "NVFP4_AFFINE_KV_CFG", "NVFP4_AWQ_CLIP_CFG", "NVFP4_AWQ_FULL_CFG", - "NVFP4_KV_ROTATE_CFG", - "FP8_KV_CFG", - "FP8_AFFINE_KV_CFG", + "NVFP4_AWQ_LITE_CFG", + "NVFP4_DEFAULT_CFG", + "NVFP4_FP8_MHA_CONFIG", "NVFP4_KV_CFG", - "NVFP4_AFFINE_KV_CFG", - "MXFP8_DEFAULT_CFG", + "NVFP4_KV_ROTATE_CFG", + "NVFP4_FP8_CFG", + "NVFP4_SVDQUANT_DEFAULT_CFG", + "W4A8_AWQ_BETA_CFG", "W4A8_MXFP4_FP8_CFG", } diff --git a/modelopt/torch/quantization/export_onnx.py b/modelopt/torch/quantization/export_onnx.py index e49331901..965da8ae7 100644 --- a/modelopt/torch/quantization/export_onnx.py +++ b/modelopt/torch/quantization/export_onnx.py @@ -592,7 +592,7 @@ def export_fp4( sx_f32_per_tensor = 1.0 else: assert num_bits == (2, 1) - sx_f32_per_tensor = float(amax) / 6.0 + sx_f32_per_tensor = float(amax) / 6.0 / 448.0 x_f4, sx_f8 = _fp4_dynamic_quantize( g, inputs, sx_f32_per_tensor, trt_high_precision_dtype, block_size diff --git a/modelopt/torch/quantization/mode.py b/modelopt/torch/quantization/mode.py index 197d2b1b2..a7a7f66bc 100644 --- a/modelopt/torch/quantization/mode.py +++ b/modelopt/torch/quantization/mode.py @@ -158,7 +158,7 @@ class RealQuantizeModeDescriptor(ModeDescriptor): def next_modes(self) -> set[str] | None: """Real quantization should be the last mode in the chain.""" # TODO: update this to support QLoRA - return {"max_calibrate"} + return {"max_calibrate", "eagle"} @property def config_class(self) -> type[ModeloptBaseConfig]: diff --git a/modelopt/torch/quantization/model_quant.py b/modelopt/torch/quantization/model_quant.py index 4c936e5f4..b8e02726f 100644 --- a/modelopt/torch/quantization/model_quant.py +++ b/modelopt/torch/quantization/model_quant.py @@ -25,7 +25,6 @@ import torch import torch.nn as nn import modelopt.torch.quantization as mtq -from modelopt.core.torch.quantization.algorithms import AutoQuantizeSearcher from modelopt.torch.opt import apply_mode from modelopt.torch.opt.searcher import ForwardLoop from modelopt.torch.opt.utils import forward_with_reshard @@ -33,6 +32,7 @@ from modelopt.torch.quantization.config import QuantizeConfig from modelopt.torch.quantization.conversion import set_quantizer_by_cfg from . import config +from .algorithms import AutoQuantizeSearcher from .config import QuantizeAlgoCfgType from .conversion import set_quantizer_attribute from .mode import QuantizeModeRegistry, get_modelike_from_algo_cfg diff --git a/modelopt/torch/quantization/nn/modules/quant_linear.py b/modelopt/torch/quantization/nn/modules/quant_linear.py index 033f0b008..8616f130e 100644 --- a/modelopt/torch/quantization/nn/modules/quant_linear.py +++ b/modelopt/torch/quantization/nn/modules/quant_linear.py @@ -15,6 +15,8 @@ """Quantized Linear.""" +import warnings + import torch import torch.nn as nn import torch.nn.functional as F @@ -135,6 +137,7 @@ class RealQuantLinear(QuantModule): """Quantized version of nn.Linear with real quantization.""" list_of_scale_tensors = ["_scale", "double_scale", "_scale_zeros"] + allow_real_quant_gemm = True def forward(self, input, *args, **kwargs): """RealQuant layer forward function.""" @@ -149,6 +152,8 @@ class RealQuantLinear(QuantModule): and input.numel() > 1 # If we need to calibrate the input, we fallback to fake quant and not (self.input_quantizer.is_enabled and self.input_quantizer._if_calib) + # Our forward might not work for every implementation, so we allow user to disable it + and self.allow_real_quant_gemm ): # If the input is not quantized, we use the default GEMM. real_quant_gemm = ( @@ -204,3 +209,13 @@ class RealQuantLinear(QuantModule): # Function to dynamically override load_state_dict dynamically_update_state_methods(self) + + def _apply(self, fn): + """Override the _apply method to ensure that the weight is real-quantized.""" + # Check if fn is a tensor_cast_fun and print warning if so + if hasattr(fn, "__name__") and "tensor_cast" in fn.__name__.lower(): + warnings.warn("RealQuantLinear does not support tensor_cast_fun.") + return self + else: + # Process the function normally + return super()._apply(fn) diff --git a/modelopt/torch/quantization/nn/modules/quant_module.py b/modelopt/torch/quantization/nn/modules/quant_module.py index b98caca05..cd10eb9a0 100644 --- a/modelopt/torch/quantization/nn/modules/quant_module.py +++ b/modelopt/torch/quantization/nn/modules/quant_module.py @@ -88,20 +88,30 @@ class QuantModule(DynamicModule): def fold_weight(self): """Fold the weight for faster eval.""" - if ( - hasattr(self, "weight_quantizer") - and hasattr(self, "weight") - and self.weight_quantizer.fake_quant - ): - self.weight.data.copy_(self.weight_quantizer(self.weight.float()).to(self.weight.dtype)) - self.weight_quantizer.disable() - _attrs = [ - "_pre_quant_scale", - "_amax", - ] - for attr in _attrs: - if hasattr(self.weight_quantizer, attr): - delattr(self.weight_quantizer, attr) + # Handle all attributes that end with _weight_quantizer + for name in dir(self): + attr = getattr(self, name) + if ( + name.endswith("weight_quantizer") + and isinstance(attr, TensorQuantizer) + and attr.fake_quant + ): + # Get the corresponding weight name by removing _weight_quantizer suffix + weight_name = name[:-10] + + assert hasattr(self, weight_name), ( + f"{name} doesn't have a corresponding {weight_name} in {self.__class__.__name__}" + ) + weight = getattr(self, weight_name) + weight.data.copy_(attr(weight.float()).to(weight.dtype)) + attr.disable() + _attrs = [ + "_pre_quant_scale", + "_amax", + ] + for attr_name in _attrs: + if hasattr(attr, attr_name): + delattr(attr, attr_name) QuantModuleRegistry = _DMRegistryCls("Quant", QuantModule) diff --git a/modelopt/torch/quantization/nn/modules/tensor_quantizer.py b/modelopt/torch/quantization/nn/modules/tensor_quantizer.py index ea0228e1d..4ffad541d 100644 --- a/modelopt/torch/quantization/nn/modules/tensor_quantizer.py +++ b/modelopt/torch/quantization/nn/modules/tensor_quantizer.py @@ -309,7 +309,10 @@ class TensorQuantizer(nn.Module): @bias_type.setter def bias_type(self, value): - assert value in ["static", "dynamic"], "bias_type must be either 'static' or 'dynamic'." + assert value in [ + "static", + "dynamic", + ], "bias_type must be either 'static' or 'dynamic'." self._bias["type"] = value @property @@ -525,7 +528,7 @@ class TensorQuantizer(nn.Module): def _validate_amax(self, amax): # Dynamic control flow is not supported by torch dynamo - if not is_torch_export_mode(): + if not is_torch_export_mode() and not torch._dynamo.is_compiling(): assert torch.all(amax >= 0) and not torch.any(torch.isinf(amax)), ( f"Got invalid amax: {amax}" ) @@ -663,7 +666,12 @@ class TensorQuantizer(nn.Module): ) outputs = scaled_e4m3( - inputs, amax, self.bias_value, E, M, self._trt_high_precision_dtype + inputs, + amax, + self.bias_value, + E, + M, + self._trt_high_precision_dtype, ) else: # Integer static quantization, e.g., INT4_BLOCKWISE @@ -908,8 +916,12 @@ class TensorQuantizer(nn.Module): self._input_dtype = inputs.dtype if hasattr(inputs, "dtype") else None return inputs - # GLOBALS could break TorchDynamo for some Pytorch versions (i.e., 2.3.0) - if not is_torch_export_mode() and GLOBALS.in_onnx_export: + if ( + not is_torch_export_mode() + and not torch._dynamo.is_compiling() + and GLOBALS.in_onnx_export + ): + # GLOBALS could break TorchDynamo for some Pytorch versions (i.e., 2.3.0) self._check_onnx_readiness(inputs) if self.block_sizes is not None and self._fake_quant: @@ -1288,6 +1300,9 @@ class SequentialQuantizer(nn.Sequential): yield - for parent_module, sequential_quantizers_list in original_sequential_quantizers.items(): + for ( + parent_module, + sequential_quantizers_list, + ) in original_sequential_quantizers.items(): for name, sequential_quantizer in sequential_quantizers_list: setattr(parent_module, name, sequential_quantizer) diff --git a/modelopt/torch/quantization/plugins/__init__.py b/modelopt/torch/quantization/plugins/__init__.py index d52b5dd79..bf930e07b 100644 --- a/modelopt/torch/quantization/plugins/__init__.py +++ b/modelopt/torch/quantization/plugins/__init__.py @@ -63,3 +63,6 @@ with import_plugin("transformer_engine"): with import_plugin("transformers trainer"): from .transformers_trainer import * + +with import_plugin("vllm"): + from .vllm import * diff --git a/modelopt/torch/quantization/plugins/huggingface.py b/modelopt/torch/quantization/plugins/huggingface.py index bcda6bc08..ce1ceb8af 100644 --- a/modelopt/torch/quantization/plugins/huggingface.py +++ b/modelopt/torch/quantization/plugins/huggingface.py @@ -32,10 +32,10 @@ import torch.nn as nn import transformers from transformers.models.t5.modeling_t5 import T5Attention -from modelopt.core.torch.quantization.algorithms import AutoQuantizeSearcher from modelopt.torch.opt.dynamic import DynamicModule from modelopt.torch.utils.distributed import ParallelState +from ..algorithms import AutoQuantizeSearcher from ..conversion import register from ..nn import QuantInputBase, QuantModule, QuantModuleRegistry, TensorQuantizer from ..nn.modules.quant_linear import _QuantLinear diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index d8db6ee83..159788eb8 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -287,6 +287,8 @@ class _QuantMegatronMLP(_MegatronMLP): class _RealQuantMegatronColumnParallelLinear(RealQuantLinear, _MegatronColumnParallelLinear): + allow_real_quant_gemm = False # We don't support real quant gemm for ColumnParallelLinear + def _parameter_to_keep_in_quantizer_state_dict(self, key): return any(k in key for k in self.list_of_scale_tensors) @@ -310,6 +312,8 @@ class _RealQuantMegatronColumnParallelLinear(RealQuantLinear, _MegatronColumnPar class _RealQuantMegatronRowParallelLinear(RealQuantLinear, _MegatronRowParallelLinear): + allow_real_quant_gemm = False # We don't support real quant gemm for RowParallelLinear + def _parameter_to_keep_in_quantizer_state_dict(self, key): return any(k in key for k in self.list_of_scale_tensors) diff --git a/modelopt/torch/quantization/plugins/peft.py b/modelopt/torch/quantization/plugins/peft.py index 27630b555..144a0345b 100644 --- a/modelopt/torch/quantization/plugins/peft.py +++ b/modelopt/torch/quantization/plugins/peft.py @@ -18,6 +18,8 @@ import torch.nn.functional as F from peft.tuners.lora.layer import Linear as LoraLinear +from modelopt.torch.quantization.qtensor.base_qtensor import QTensorWrapper + from ..nn import QuantModule, QuantModuleRegistry, TensorQuantizer __all__ = [] @@ -35,20 +37,52 @@ class _QuantLoraLinear(QuantModule): if self.disable_adapters or adapter_names is not None or self.merged: return super().forward(x, args, kwargs) - weight = self.base_layer.weight - for active_adapter in self.active_adapters: - if active_adapter not in self.lora_A.keys(): # noqa: SIM118 - continue - lora_a = self.lora_A[active_adapter] - lora_b = self.lora_B[active_adapter] - scaling = self.scaling[active_adapter] - - if not self.use_dora[active_adapter]: - weight = weight + scaling * lora_b.weight @ lora_a.weight - else: - raise NotImplementedError("dora not implemented") - x = self.input_quantizer(x) - weight = self.weight_quantizer(weight) - output = self.output_quantizer(F.linear(x, weight, self.base_layer.bias)) + weight = self.base_layer.weight + is_compressed = isinstance(weight, QTensorWrapper) + if not is_compressed: + for active_adapter in self.active_adapters: + if active_adapter not in self.lora_A.keys(): # noqa: SIM118 + continue + lora_a = self.lora_A[active_adapter] + lora_b = self.lora_B[active_adapter] + scaling = self.scaling[active_adapter] + + if not self.use_dora[active_adapter]: + weight = weight + scaling * lora_b.weight @ lora_a.weight + else: + raise NotImplementedError("dora not implemented") + weight = self.weight_quantizer(weight) + output = F.linear(x, weight, self.base_layer.bias) + else: + # For compressed weights, compute base output and LoRA outputs separately + base_output = self.base_layer(x) + + # Only compute LoRA outputs if there are active adapters + if self.active_adapters: + # Start with zero LoRA output + lora_output = None + + for active_adapter in self.active_adapters: + if active_adapter not in self.lora_A.keys(): # noqa: SIM118 + continue + lora_a = self.lora_A[active_adapter] + lora_b = self.lora_B[active_adapter] + scaling = self.scaling[active_adapter] + + if not self.use_dora[active_adapter]: + # Compute LoRA output step by step to maintain gradient flow + lora_a_output = F.linear(x, lora_a.weight) + lora_b_output = F.linear(lora_a_output, lora_b.weight) + adapter_output = scaling * lora_b_output + + lora_output = ( + adapter_output if lora_output is None else lora_output + adapter_output + ) + + output = base_output + lora_output + else: + output = base_output + + output = self.output_quantizer(output) return output diff --git a/modelopt/torch/quantization/plugins/transformers_trainer.py b/modelopt/torch/quantization/plugins/transformers_trainer.py index 2a2fb3edc..9125df1de 100644 --- a/modelopt/torch/quantization/plugins/transformers_trainer.py +++ b/modelopt/torch/quantization/plugins/transformers_trainer.py @@ -205,6 +205,9 @@ class QADTrainer(QATTrainer, KDTrainer): self.model.cuda() if self.quant_cfg is not None and not is_quantized(self.model): self._quantize_model(use_eval_loop=False) + if getattr(self.args, "lora_config", None) is not None: + self.model.add_adapter(self.args.lora_config, adapter_name="adapter") + print_rank_0("Lora adapter added.") self._convert_to_distillation_model() def _convert_to_distillation_model(self): diff --git a/modelopt/torch/quantization/plugins/vllm.py b/modelopt/torch/quantization/plugins/vllm.py new file mode 100644 index 000000000..fb606b0d5 --- /dev/null +++ b/modelopt/torch/quantization/plugins/vllm.py @@ -0,0 +1,199 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Support quantization for VLLM layers.""" + +import importlib + +import torch +import vllm.model_executor.layers.fused_moe.layer as vllm_fused_moe_layer +import vllm.model_executor.layers.linear as vllm_linear + +from ...utils.distributed import ParallelState +from ..nn import QuantLinearConvBase, QuantModule, QuantModuleRegistry, TensorQuantizer + +vllm_fused_moe_package = importlib.import_module("vllm.model_executor.layers.fused_moe.fused_moe") + + +class FakeQuantMethod: + """A class that implements fake quantization methods for vLLM models. + + This class provides functionality to apply quantization methods to model layers + in a way that's compatible with vLLM's architecture. + """ + + def __init__(self, quant_method): + """Initialize the FakeQuantMethod. + + Args: + quant_method: The quantization method to be applied to the model layers. + """ + self.quant_method = quant_method + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + """Apply the quantization method to a given layer. + + Args: + layer (torch.nn.Module): The neural network layer to be quantized. + x (torch.Tensor): The input tensor to the layer. + bias (torch.Tensor | None, optional): The bias tensor to the layer. Defaults to None. + + Returns: + torch.Tensor: The quantized output tensor. + """ + x = layer.input_quantizer(x) + if layer.weight_quantizer.is_enabled: + original_weight = layer.weight + layer.weight = layer.weight_quantizer(layer.weight) + output = self.quant_method.apply(layer, x, bias) + layer.weight = original_weight + else: + output = self.quant_method.apply(layer, x, bias) + output = layer.output_quantizer(output) + return output + + +class _VLLMParallelLinear(QuantModule): + def _setup(self): + self.input_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_input) + self.weight_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_weight) + self.output_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_output) + self.output_quantizer.disable() + assert type(self.quant_method) is vllm_linear.UnquantizedLinearMethod, ( + f"quant_method is {type(self.quant_method)}" + ) + self.fake_quant_method = FakeQuantMethod(self.quant_method) + self.parallel_state = ParallelState(-1, -1) + + def forward(self, input_): + # This context manager will conflict with torch.compile + # with replace_function(self, "quant_method", self.fake_quant_method): + # Manually replace quant_method instead + self._quant_method = self.quant_method + self.quant_method = self.fake_quant_method + output = super().forward(input_) + self.quant_method = self._quant_method + return output + + +@QuantModuleRegistry.register({vllm_linear.RowParallelLinear: "vllm_RowParallelLinear"}) +class _QuantVLLMRowParallelLinear(_VLLMParallelLinear): + pass + + +@QuantModuleRegistry.register({vllm_linear.ColumnParallelLinear: "vllm_ColumnParallelLinear"}) +class _QuantVLLMColumnParallelLinear(_VLLMParallelLinear): + pass + + +@QuantModuleRegistry.register( + {vllm_linear.MergedColumnParallelLinear: "vllm_MergedColumnParallelLinear"} +) +class _QuantVLLMMergedColumnParallelLinear(_VLLMParallelLinear): + pass + + +@QuantModuleRegistry.register({vllm_linear.QKVParallelLinear: "vllm_QKVParallelLinear"}) +class _QuantVLLMQKVParallelLinear(_VLLMParallelLinear): + pass + + +# ReplicatedLinear is for MoE router and should not be quantized + + +# FusedMoE layer requires handling for UnquantizedFusedMoEMethod +@QuantModuleRegistry.register({vllm_fused_moe_layer.FusedMoE: "vllm_FusedMoE"}) +class _QuantVLLMFusedMoE(QuantModule): + def _setup(self): + self.w13_input_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_input) + self.w2_input_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_input) + self.w13_weight_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_weight) + self.w2_weight_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_weight) + self.w13_output_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_output) + self.w2_output_quantizer = TensorQuantizer(QuantLinearConvBase.default_quant_desc_output) + self.w13_output_quantizer.disable() + self.w2_output_quantizer.disable() + assert type(self.quant_method) is vllm_fused_moe_layer.UnquantizedFusedMoEMethod, ( + f"quant_method is {type(self.quant_method)}" + ) + self.parallel_state = ParallelState(-1, -1) + + def invoke_fused_moe_quantized( + self, + A: torch.Tensor, # noqa: N803 + B: torch.Tensor, # noqa: N803 + C: torch.Tensor, # noqa: N803 + *args, + **kwargs, + ): + if B is self.w13_weight: + # First layer of expert + A = self.w13_input_quantizer(A) # noqa: N806 + if self.w13_weight_quantizer.is_enabled: + original_weight = self.w13_weight + self.w13_weight = self.w13_weight_quantizer(self.w13_weight) + vllm_fused_moe_package._invoke_fused_moe_kernel(A, B, C, *args, **kwargs) + self.w13_weight = original_weight + else: + vllm_fused_moe_package._invoke_fused_moe_kernel(A, B, C, *args, **kwargs) + if self.w13_output_quantizer.is_enabled: + C[:] = self.w13_output_quantizer(C) + elif B is self.w2_weight: + A = self.w2_input_quantizer(A) # noqa: N806 + if self.w2_weight_quantizer.is_enabled: + original_weight = self.w2_weight + self.w2_weight = self.w2_weight_quantizer(self.w2_weight) + vllm_fused_moe_package._invoke_fused_moe_kernel(A, B, C, *args, **kwargs) + self.w2_weight = original_weight + else: + vllm_fused_moe_package._invoke_fused_moe_kernel(A, B, C, *args, **kwargs) + if self.w2_output_quantizer.is_enabled: + C[:] = self.w2_output_quantizer(C) + else: + raise ValueError("Cannot determine first or second layer of expert") + + def forward(self, hidden_states: torch.Tensor, router_logits: torch.Tensor): + # This is again due to the bad coding of vLLM + # fused_moe submodule is overwritten by the fused_moe function + # so we need to import the fused_moe module explicitly + assert vllm_fused_moe_package.invoke_fused_moe_kernel is not None + # This context manager will conflict with torch.compile + # with replace_function( + # vllm_fused_moe_package, + # "invoke_fused_moe_kernel", + # self.invoke_fused_moe_quantized, + # ): + self._invoke_fused_moe_quantized = self.invoke_fused_moe_quantized + self.invoke_fused_moe_quantized = self.invoke_fused_moe_quantized + output = super().forward(hidden_states, router_logits) + self.invoke_fused_moe_quantized = self._invoke_fused_moe_quantized + return output + + @property + def mopt_ckpt_versn(self): + """Checkpoint version of the modelopt.""" + return None + + @mopt_ckpt_versn.setter + def mopt_ckpt_versn(self, version: str): + """Set the checkpoint version for the TensorQuantizer states.""" + # vLLM defined an apply method that overwrites nn.Module.apply + # To avoid conflicting, disable the apply call here + # self.apply(_set_ckpt_version) diff --git a/modelopt/torch/quantization/qtensor/__init__.py b/modelopt/torch/quantization/qtensor/__init__.py index b3531ec85..c4ed88f87 100644 --- a/modelopt/torch/quantization/qtensor/__init__.py +++ b/modelopt/torch/quantization/qtensor/__init__.py @@ -15,11 +15,10 @@ """Tensor Class for Real Quantization.""" -from modelopt.core.torch.quantization.qtensor.nvfp4_tensor import * - from .base_qtensor import * from .fp8_tensor import * from .int4_tensor import * from .int8_tensor import * from .mxfp4_tensor import * from .nf4_tensor import * +from .nvfp4_tensor import * diff --git a/modelopt/torch/quantization/qtensor/base_qtensor.py b/modelopt/torch/quantization/qtensor/base_qtensor.py index 38dbe2590..ee01167a8 100644 --- a/modelopt/torch/quantization/qtensor/base_qtensor.py +++ b/modelopt/torch/quantization/qtensor/base_qtensor.py @@ -173,7 +173,7 @@ def pack_real_quantize_weight(module, force_quantize: bool = False): with SequentialQuantizer.convert_to_single_quantizer(module), torch.no_grad(): for _, m in module.named_modules(): - if hasattr(m, "weight") and m.weight.is_meta: + if hasattr(m, "weight") and (m.weight is None or m.weight.is_meta): continue if ( hasattr(m, "weight_quantizer") diff --git a/modelopt/torch/quantization/qtensor/nvfp4_tensor.py b/modelopt/torch/quantization/qtensor/nvfp4_tensor.py new file mode 100644 index 000000000..b67f90950 --- /dev/null +++ b/modelopt/torch/quantization/qtensor/nvfp4_tensor.py @@ -0,0 +1,282 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Implements NVFP4 quantization for efficient tensor storage and computation.""" + +import numpy as np +import torch + +from ..backends.utils import fp4_compatible +from ..qtensor.base_qtensor import BaseQuantizedTensor +from ..utils import reduce_block_padding + +# Define conversion tables +e2m1_bounds = torch.tensor([0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5]) +e2m1_values = torch.tensor([0, 0.5, 1, 1.5, 2, 3, 4, 6, 0, -0.5, -1, -1.5, -2, -3, -4, -6]) + +__all__ = ["NVFP4QTensor"] + + +class NVFP4QTensor(BaseQuantizedTensor): + """Implements the INT4 quantization on tensors for more efficient storage or computation. + + Attributes: + quantized_data (torch.Tensor): The quantized data stored as a packed uint8 tensor. + """ + + e2m1_values_on_device = {} + + @classmethod + def get_e2m1_values(cls, device): + """Returns the e2m1 values on the device.""" + if device not in cls.e2m1_values_on_device: + cls.e2m1_values_on_device[device] = e2m1_values.to(device) + return cls.e2m1_values_on_device[device] + + @classmethod + def get_weights_scaling_factor_2_from_quantizer(cls, weight_quantizer): + """Returns per tensor weight scaling factor from the weight_quantizer amax.""" + # Assert that weight_quantizer has attribute amax + assert hasattr(weight_quantizer, "_amax"), "Weight quantizer does not have attribute amax" + return weight_quantizer._amax.float() / 6.0 / 448.0 + + @classmethod + def get_weights_scaling_factor( + cls, + input: torch.Tensor, + block_size: int, + weights_scaling_factor_2: torch.Tensor | None = None, + keep_high_precision: bool = False, + ): + """Returns quantized per block weight scaling factor.""" + if weights_scaling_factor_2 is None: + weights_scaling_factor_2 = cls.get_weights_scaling_factor_2(input) + + # Get per_block amax + [n, k] = input.shape[-2:] + assert block_size != 0, "Block size is zero. Cannot return per_block amax for given input." + + assert k % block_size == 0, ( + "Weight shape is not divisible for block size for block quantiation." + ) + + input = input.reshape((*tuple(input.shape[:-2]), n, k // block_size, block_size)) + # Get per block amax + per_block_amax = input.abs().amax(dim=-1).float() + # Get per-block-scale + per_block_scale = per_block_amax / 6.0 + # Quantize per_block_scale to FP8 + q_per_block_scale = per_block_scale / weights_scaling_factor_2 + # Set all zero values in scale to 1.0 + q_per_block_scale[per_block_scale == 0] = 1.0 + # Convert to torch.float8_e4m3fn + if not keep_high_precision: + q_per_block_scale = q_per_block_scale.to(torch.float8_e4m3fn) + return q_per_block_scale, weights_scaling_factor_2 + + @classmethod + def get_weights_scaling_factor_2(cls, input: torch.Tensor): + """Returns per tensor weight scaling factor.""" + return input.abs().amax().float() / 6.0 / 448.0 + + @classmethod + def get_activation_scaling_factor(cls, quantizer): + """Returns the activation scaling factor for export.""" + # TODO: Update to use module and not quantizer + if not quantizer.is_enabled: + return None + + amax = quantizer.export_amax() + + if amax is None: + return None + + activation_scaling_factor = amax.float() / (quantizer.maxbound) + activation_scaling_factor = activation_scaling_factor / 448.0 + + assert torch.all(activation_scaling_factor > 0), ( + f" activation scaling factor {activation_scaling_factor} not positive." + ) + + return activation_scaling_factor + + @staticmethod + def _cast_fp4(weight: torch.Tensor): + """Converts tensor to uint4.""" + # Get device + device = weight.device + + # Define mask to perform rounding + mask = torch.tensor([0, 1, 0, 1, 0, 1, 0], dtype=torch.uint8).to(device) + mask_shape = list(weight.shape) + mask = mask.expand([*mask_shape, 7]) + + sign_bit = (weight < 0).to(torch.uint8) + + weight_abs = weight.abs_() + # Calculate the ordinal value based on the bounds + ord = torch.searchsorted(e2m1_bounds.to(device), weight_abs, out_int32=True).to(torch.uint8) + # All values equal to e2m1_bounds at odd indices are rounded up and even indices are rounded down + round = torch.any((weight_abs.unsqueeze(-1) == e2m1_bounds.to(device)) * mask, dim=-1) + fp4_val = (sign_bit * 0b1000 + ord + round).to(torch.uint8) + return fp4_val + + @classmethod + def quantize( + cls, + input: torch.Tensor, + block_size: int, + weights_scaling_factor: torch.Tensor | None = None, + weights_scaling_factor_2: torch.Tensor | None = None, + keep_high_precision: bool = False, + try_tensorrt: bool = False, + ): + """Converting a tensor to a quantized format based on NVFP4 quantization. + + Args: + input (torch.Tensor): The input tensor to be quantized. + block_size (int): The size of each block for quantization. + weights_scaling_factor (torch.Tensor): The scaling factor for the weights. + weights_scaling_factor_2 (torch.Tensor): The scaling factor for the weights. + keep_high_precision (bool): Whether to keep output scales at high precision. + + Returns: + tuple: Contains quantized data, quantized per block scaling factor, and per tensor scaling factor. + """ + # Get original input shape + input_shape = input.shape + input_dtype = input.dtype + + # pad the input if needed + input = reduce_block_padding(input, block_sizes={-1: block_size}) + + if weights_scaling_factor_2 is None: + weights_scaling_factor_2 = cls.get_weights_scaling_factor_2(input) + + # try call trtllm fp4 quantization if possible + if ( + fp4_compatible() + and weights_scaling_factor is None + and try_tensorrt + and block_size == 16 + ): + try: + import tensorrt_llm # noqa: F401 + + # Make sure this utils is available for dequantize + from tensorrt_llm._torch.auto_deploy.utils.quantization_utils import ( + cutlass_fp4_scale_to_modelopt_fp4_scale, # noqa: F401 + ) + + packed_weight, weights_scaling_factor = torch.ops.trtllm.fp4_quantize( + input, 1.0 / weights_scaling_factor_2, block_size, False + ) + # weights_scaling_factor is ready for nvfp4_gemm to use; + # however, it is different from the non trtllm version, so when dequantize, + # it will be converted. + return ( + cls(input_shape, input_dtype, packed_weight), + weights_scaling_factor, + weights_scaling_factor_2, + ) + except ImportError: + pass + + if weights_scaling_factor is None: + weights_scaling_factor, _ = cls.get_weights_scaling_factor( + input, block_size, weights_scaling_factor_2 + ) + + # Reshape the weight and scale factors + input = input.view((*tuple(input.shape[:-1]), -1, block_size)) + + # Scale weights + scaled_weight = input / ( + (weights_scaling_factor.to(torch.float32) * weights_scaling_factor_2).unsqueeze(-1) + ) + + # Reshape weights to original + scaled_weight = scaled_weight.view((*tuple(scaled_weight.shape[:-2]), -1)) + + if keep_high_precision: + return scaled_weight + # Cast weights to fp4 + q_weight = cls._cast_fp4(scaled_weight) + # Pack weights + packed_weight = (q_weight[..., 1::2] << 4) | q_weight[..., 0::2] + return ( + cls(input_shape, input_dtype, packed_weight), + weights_scaling_factor, + weights_scaling_factor_2, + ) + + def dequantize(self, dtype: torch.dtype = None, **kwarg): + """Dequantze NVFP4 packed tensor to a target dtype.""" + if dtype is None: + dtype = self.metadata["dtype"] + + def _unpack_tensor(input: torch.Tensor): + # Initalize storage for unpacked tensor + unpacked = torch.empty( + [input.shape[0], input.shape[1] * 2], dtype=dtype, device=input.device + ) + unpacked_shape = unpacked.shape + + unpacked[..., 1::2] = input >> 4 + unpacked[..., 0::2] = input & 0x0F + + unpacked = unpacked.reshape(-1) + unpacked = self.get_e2m1_values(input.device)[unpacked.long()] + + return unpacked.reshape(unpacked_shape) + + # Get scales from kwargs + if kwarg["scale"].dtype == torch.uint8 and kwarg["scale"].ndim == 1: + # If quantization is done by trtllm, convert cutlass fp4 scale to modelopt fp4 scale + try: + from tensorrt_llm._torch.auto_deploy.utils.quantization_utils import ( + cutlass_fp4_scale_to_modelopt_fp4_scale, + ) + + kwarg["scale"] = cutlass_fp4_scale_to_modelopt_fp4_scale( + kwarg["scale"], self.metadata["shape"][-2:] + ) + except ImportError as e: + raise ImportError( + "This tensor is quantized by trtllm, but tensorrt_llm cannot be imported." + ) from e + q_per_block_scale = ( + kwarg["scale"].to(torch.float32) + if kwarg["scale"].dtype == torch.float8_e4m3fn + else kwarg["scale"] + ) + block_sizes = kwarg["block_sizes"][-1] + per_block_quant_scale = kwarg["double_scale"] + + # Dequantize scales + per_block_scale = q_per_block_scale * per_block_quant_scale + + # Unpack and unscale weights + deq_data = _unpack_tensor(self._quantized_data) + + deq_data = deq_data.view( + deq_data.shape[0], deq_data.shape[1] // block_sizes, -1 + ) * per_block_scale.unsqueeze(-1) + + return ( + deq_data.view(-1)[: np.prod(self.metadata["shape"])] + .reshape(self.metadata["shape"]) + .to(dtype) + ) diff --git a/modelopt/torch/quantization/tensor_quant.py b/modelopt/torch/quantization/tensor_quant.py index dd0b56b51..ab8d3cf6f 100644 --- a/modelopt/torch/quantization/tensor_quant.py +++ b/modelopt/torch/quantization/tensor_quant.py @@ -25,7 +25,6 @@ from torch.onnx import symbolic_helper import modelopt.torch.quantization.triton as triton_kernel from .config import QuantizerAttributeConfig -from .export_onnx import export_fp4, export_fp8, export_int8, export_mxfp8 from .extensions import get_cuda_ext, get_cuda_ext_fp8, get_cuda_ext_mx mx_format_map = { @@ -191,7 +190,7 @@ def _dynamic_block_quantize_impl( and not DISABLE_TRITON_KERNEL and amax is not None ): - return triton_kernel.fp4_fake_quant_block(inputs, amax.item()) + return triton_kernel.fp4_fake_quant_block(inputs, amax) cuda_ext_mx = get_cuda_ext_mx(raise_if_failed=True) return cuda_ext_mx.fused_amax_convert( inputs, @@ -325,6 +324,8 @@ class FakeTensorQuantFunction(Function): trt_high_precision_dtype=None, ): """ONNX symbolic function.""" + from .export_onnx import export_int8 + return export_int8( g, inputs, amax, num_bits, unsigned, narrow_range, trt_high_precision_dtype ) @@ -395,6 +396,8 @@ class ScaledE4M3Function(Function): @symbolic_helper.parse_args("v", "t", "t", "i", "i", "s") def symbolic(g, inputs, amax=None, bias=None, E=4, M=3, trt_high_precision_dtype=None): # noqa: N803 """ONNX symbolic function.""" + from .export_onnx import export_fp8 + return export_fp8(g, inputs, amax, trt_high_precision_dtype) @staticmethod @@ -411,7 +414,12 @@ class ScaledE4M3Function(Function): ctx.amax = amax outputs = quantize_op( - inputs, amax, num_bits=8, exponent_bits=4, unsigned=False, narrow_range=False + inputs, + amax, + num_bits=8, + exponent_bits=4, + unsigned=False, + narrow_range=False, ) if bias is not None: @@ -424,7 +432,9 @@ class ScaledE4M3Function(Function): """Implements straight through estimation with clipping.""" (inputs,) = ctx.saved_tensors amax = torch.tensor( - ctx.amax if ctx.amax is not None else 448.0, dtype=torch.float32, device=inputs.device + ctx.amax if ctx.amax is not None else 448.0, + dtype=torch.float32, + device=inputs.device, ) grad_inputs = _fake_tensor_quant_backward(inputs, amax, grad_outputs) return grad_inputs, None, None, None, None, None @@ -453,7 +463,13 @@ def _dynamic_block_quantize_forward( scale_exponent_bits = scale_bits[0] scale_num_bits = scale_bits[0] + scale_bits[1] + 1 outputs = dynamic_block_quantize_op( - inputs, block_size, amax, num_bits, exponent_bits, scale_num_bits, scale_exponent_bits + inputs, + block_size, + amax, + num_bits, + exponent_bits, + scale_num_bits, + scale_exponent_bits, ) return outputs @@ -475,6 +491,8 @@ class DynamicBlockQuantizationFunction(Function): onnx_quantizer_type="dynamic", ): """ONNX symbolic function.""" + from .export_onnx import export_fp4, export_mxfp8 + if num_bits == (2, 1) and scale_bits == (4, 3): return export_fp4( g, @@ -643,6 +661,8 @@ class TensorQuantFunction(Function): trt_high_precision_dtype=None, ): """ONNX symbolic function.""" + from .export_onnx import export_int8 + return export_int8( g, inputs, amax, num_bits, unsigned, narrow_range, trt_high_precision_dtype ) diff --git a/modelopt/torch/quantization/triton/fp4_kernel.py b/modelopt/torch/quantization/triton/fp4_kernel.py index 6d4df33c7..79a7b259c 100644 --- a/modelopt/torch/quantization/triton/fp4_kernel.py +++ b/modelopt/torch/quantization/triton/fp4_kernel.py @@ -33,7 +33,7 @@ def fp4_fake_quant_kernel( y_ptr, M, N, - global_scale, + global_scale_ptr, BLOCK_SIZE: tl.constexpr, TILE_SIZE: tl.constexpr, NUM_FP4_BLOCKS: tl.constexpr, @@ -45,7 +45,7 @@ def fp4_fake_quant_kernel( y_ptr (tl.pointer): Pointer to the output buffer M (int): Number of rows in the matrix N (int): Number of columns in the matrix - global_scale (float): Global scaling factor + global_scale_ptr (tl.pointer): Pointer to the global scaling factor tensor BLOCK_SIZE (tl.constexpr): Size of each FP4 quantization block TILE_SIZE (tl.constexpr): Size of the processing block NUM_FP4_BLOCKS (tl.constexpr): Number of FP4 blocks within TILE_SIZE @@ -53,6 +53,9 @@ def fp4_fake_quant_kernel( pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) + # Load global scale from tensor + global_scale = tl.load(global_scale_ptr) + # Calculate offsets offs_m = pid_m * TILE_SIZE + tl.arange(0, TILE_SIZE) offs_n = pid_n * TILE_SIZE + tl.arange(0, TILE_SIZE) @@ -117,7 +120,7 @@ def fp4_fake_quant_kernel( def fp4_fake_quant_block( x: torch.Tensor, - global_amax: float, + global_amax: torch.Tensor, block_size: int = 16, tile_size: int = 128, ) -> torch.Tensor: @@ -125,7 +128,8 @@ def fp4_fake_quant_block( Args: x (torch.Tensor): Input tensor of shape (M, N) - global_scale (float): Global scaling factor + global_amax (torch.Tensor): Global max value of the input tensor + This needs to be a tensor to be cuda-graph compatible block_size (int): Size of FP4 quantization blocks tile_size (int): Size of processing blocks @@ -139,7 +143,10 @@ def fp4_fake_quant_block( M, N = x.size() y = torch.empty_like(x, dtype=torch.get_default_dtype()) - grid = lambda meta: (triton.cdiv(M, meta["TILE_SIZE"]), triton.cdiv(N, meta["TILE_SIZE"])) + grid = lambda meta: ( + triton.cdiv(M, meta["TILE_SIZE"]), + triton.cdiv(N, meta["TILE_SIZE"]), + ) global_scale = (global_amax / 6.0) / 448.0 num_fp4_blocks = tile_size // block_size fp4_fake_quant_kernel[grid]( diff --git a/modelopt/torch/speculative/config.py b/modelopt/torch/speculative/config.py index a95658b4c..dc1277359 100644 --- a/modelopt/torch/speculative/config.py +++ b/modelopt/torch/speculative/config.py @@ -136,6 +136,13 @@ class EagleConfig(ModeloptBaseConfig): ), ) + parallel_draft_step: int = ModeloptField( + default=1, + description=( + "The number of tokens generated in parallel draft. If set to 1, draft is not in parallel mode." + ), + ) + class MTPConfig(ModeloptBaseConfig): """MTP config.""" diff --git a/modelopt/torch/speculative/eagle/conversion.py b/modelopt/torch/speculative/eagle/conversion.py index 571a57727..32efb9886 100644 --- a/modelopt/torch/speculative/eagle/conversion.py +++ b/modelopt/torch/speculative/eagle/conversion.py @@ -50,6 +50,7 @@ def convert_to_eagle_model(model: nn.Module, config: EagleConfig) -> ConvertRetu draft_vocab_size=config.draft_vocab_size, use_mtp_layernorm=config.use_mtp_layernorm, ffn_hidden_size=config.ffn_hidden_size, + parallel_draft_step=config.parallel_draft_step, ) # no metadata, all specifed via config. diff --git a/modelopt/torch/speculative/eagle/eagle_model.py b/modelopt/torch/speculative/eagle/eagle_model.py index 86cd3e55b..a13f0ed02 100644 --- a/modelopt/torch/speculative/eagle/eagle_model.py +++ b/modelopt/torch/speculative/eagle/eagle_model.py @@ -15,6 +15,8 @@ """Eagle model to support eagle decoding.""" +import torch + from modelopt.torch.opt.dynamic import DynamicModule @@ -36,6 +38,7 @@ class EagleModel(DynamicModule): eagle_disable_moe, draft_vocab_size, use_mtp_layernorm, + parallel_draft_step, ): """Base Eagle Model modify function. Child class should implement the details.""" self.eagle_num_layers = eagle_num_layers @@ -47,6 +50,7 @@ class EagleModel(DynamicModule): self.eagle_disable_moe = eagle_disable_moe self.draft_vocab_size = draft_vocab_size self.use_mtp_layernorm = use_mtp_layernorm + self.parallel_draft_step = parallel_draft_step # Use default aux_hidden_state layers if use_aux_hidden_state is True # but no layer id is given @@ -57,3 +61,7 @@ class EagleModel(DynamicModule): assert not self.eagle_hidden_state_distillation, ( "EAGLE-3 does not support hidden state distillation!" ) + + if self.parallel_draft_step > 1: + for i in range(self.parallel_draft_step - 1): + self.register_buffer(f"mask_token_{i}", torch.tensor(-1)) diff --git a/modelopt/torch/speculative/plugins/__init__.py b/modelopt/torch/speculative/plugins/__init__.py index 245338239..5e3f4bff2 100644 --- a/modelopt/torch/speculative/plugins/__init__.py +++ b/modelopt/torch/speculative/plugins/__init__.py @@ -23,8 +23,11 @@ write your own one. Currently, we support plugins for from modelopt.torch.utils import import_plugin -with import_plugin("megatron"): - from .megatron import * +with import_plugin("megatron_eagle"): + from .megatron_eagle import * + +with import_plugin("megatron_medusa"): + from .megatron_medusa import * with import_plugin("transformers"): from .transformers import * diff --git a/modelopt/torch/speculative/plugins/megatron.py b/modelopt/torch/speculative/plugins/megatron_eagle.py similarity index 52% rename from modelopt/torch/speculative/plugins/megatron.py rename to modelopt/torch/speculative/plugins/megatron_eagle.py index 5892c096d..d5c3ff264 100644 --- a/modelopt/torch/speculative/plugins/megatron.py +++ b/modelopt/torch/speculative/plugins/megatron_eagle.py @@ -13,14 +13,14 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Plugin to add Medusa support for megatron-core GPT model.""" +"""Plugin to add EAGLE support for Megatron-Core GPT model.""" import copy import warnings +from collections import deque import megatron.core import torch -import torch.nn.functional as F from megatron.core import InferenceParams, tensor_parallel from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding @@ -40,6 +40,7 @@ from megatron.core.tensor_parallel.mappings import ( scatter_to_sequence_parallel_region, ) from megatron.core.transformer.attention import SelfAttention +from megatron.core.transformer.enums import AttnMaskType from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_block import TransformerBlock @@ -48,13 +49,15 @@ from megatron.core.transformer.utils import sharded_state_dict_default from megatron.core.utils import make_tp_sharded_tensor_for_checkpoint from packaging.version import Version +from ...opt.plugins.megatron_model_config import Llama31Config8B from ..eagle.conversion import EagleDMRegistry from ..eagle.eagle_model import EagleModel -from ..medusa.conversion import MedusaDMRegistry -from ..medusa.medusa_model import MedusaModel -from ..mtp.conversion import MTPDMRegistry -from ..mtp.mtp_model import MTPModel -from ..utils import AcceptanceRateValidation, get_default_attention_mask_and_position_ids +from ..utils import ( + AcceptanceRateValidation, + Tree, + TreeNode, + get_default_attention_mask_and_position_ids, +) try: from megatron.core.post_training.modelopt.gpt.model_specs import get_gpt_modelopt_spec @@ -111,286 +114,185 @@ def right_padding(input_ids: torch.Tensor, hidden_states: torch.Tensor = None): return padded_input_ids, seq_len -class MedusaLayer(MegatronModule): - """MedusaLayer impl following TensorRT-LLM's model definition. +def set_multi_step_attention_mask(attn_mask, step): + """Given an original attention_mask, construct a multi-step attention_mask. - Medusa layer consists of a column parallel linear following a silu. - """ + i0 i1 i2 i3 i4 i5 i6 i7 (base input_ids) + ======================= + h0 h1 h2 h3 h4 h5 h6 h7 (base hidden_states) + l0 l1 l2 l3 l4 l5 l6 l7 (base labels) - def __init__(self, config): - """Constructor. - Args: - config: MCore transformer config - """ - super().__init__(config=config) + (1st) | i1 i2 i3 i4 i5 i6 i7 -- | + (out) | h0 h1 h2 h3 h4 h5 h6 h7 | + ========================================= + f1 l1 | i1 h0 | x | + f2 l2 | i2 h1 | x x | + f3 l3 | i3 h2 | x x x | + f4 l4 | i4 h3 | x x x x | + f5 l5 | i5 h4 | x x x x x | + f6 l6 | i6 h5 | x x x x x x | + f7 l7 | i7 h6 | x x x x x x x | + -- -- | -- h7 | o o o o o o o o | + ========================================= - device = ( - torch.device("cpu") if config.use_cpu_initialization else torch.cuda.current_device() + + (2nd) | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | + (out) | h0 h1 h2 h3 h4 h5 h6 h7 | -- F1 F2 F3 F4 F5 F6 F7 | + =================================================================== + F1 l1 | i1 h0 | x | | + F2 l2 | i2 h1 | x x | | + F3 l3 | i3 h2 | x x x | | + F4 l4 | i4 h3 | x x x x | | + F5 l5 | i5 h4 | x x x x x | | + F6 l6 | i6 h5 | x x x x x x | | + F7 l7 | i7 h6 | x x x x x x x | | + -- -- | -- h7 | o o o o o o o o | | + =================================================================== + -- -- | i1 -- | | | + G2 l2 | i2 F1 | x o | x | + G3 l3 | i3 F2 | x x o | x | + G4 l4 | i4 F3 | x x x o | x | + G5 l5 | i5 F4 | x x x x o | x | + G6 l6 | i6 F5 | x x x x x o | x | + G7 l7 | i7 F6 | x x x x x x o | x | + -- -- | -- F7 | | | + =================================================================== + + + (3rd) | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | + (out) | h0 h1 h2 h3 h4 h5 h6 h7 | -- F1 F2 F3 F4 F5 F6 F7 | -- -- G2 G3 G4 G5 G6 G7 | + ============================================================================================= + F1 l1 | i1 h0 | x | | | + F2 l2 | i2 h1 | x x | | | + F3 l3 | i3 h2 | x x x | | | + F4 l4 | i4 h3 | x x x x | | | + F5 l5 | i5 h4 | x x x x x | | | + F6 l6 | i6 h5 | x x x x x x | | | + F7 l7 | i7 h6 | x x x x x x x | | | + -- -- | -- h7 | o o o o o o o o | | | + ============================================================================================= + -- -- | i1 -- | | | | + G2 l2 | i2 F1 | x o | x | | + G3 l3 | i3 F2 | x x o | x | | + G4 l4 | i4 F3 | x x x o | x | | + G5 l5 | i5 F4 | x x x x o | x | | + G6 l6 | i6 F5 | x x x x x o | x | | + G7 l7 | i7 F6 | x x x x x x o | x | | + -- -- | -- F7 | | | | + ============================================================================================= + -- -- | i1 -- | | | | + -- -- | i2 -- | | | | + H3 l3 | i3 G2 | x o o | x o | x | + H4 l4 | i4 G3 | x x o o | x o | x | + H5 l5 | i5 G4 | x x x o o | x o | x | + H6 l6 | i6 G5 | x x x x o o | x o | x | + H7 l7 | i7 G6 | x x x x x o o | x o | x | + -- -- | -- G7 | | | | + ============================================================================================= + + + (4th) | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | + (out) | h0 h1 h2 h3 h4 h5 h6 h7 | -- F1 F2 F3 F4 F5 F6 F7 | -- -- G2 G3 G4 G5 G6 G7 | -- -- -- H3 H4 H5 H6 H7 | + ======================================================================================================================= + F1 l1 | i1 h0 | x | | | | + F2 l2 | i2 h1 | x x | | | | + F3 l3 | i3 h2 | x x x | | | | + F4 l4 | i4 h3 | x x x x | | | | + F5 l5 | i5 h4 | x x x x x | | | | + F6 l6 | i6 h5 | x x x x x x | | | | + F7 l7 | i7 h6 | x x x x x x x | | | | + -- -- | -- h7 | o o o o o o o o | | | | + ======================================================================================================================= + -- -- | i1 -- | | | | | + G2 l2 | i2 F1 | x o | x | | | + G3 l3 | i3 F2 | x x o | x | | | + G4 l4 | i4 F3 | x x x o | x | | | + G5 l5 | i5 F4 | x x x x o | x | | | + G6 l6 | i6 F5 | x x x x x o | x | | | + G7 l7 | i7 F6 | x x x x x x o | x | | | + -- -- | -- F7 | | | | | + ======================================================================================================================= + -- -- | i1 -- | | | | | + -- -- | i2 -- | | | | | + H3 l3 | i3 G2 | x o o | x o | x | | + H4 l4 | i4 G3 | x x o o | x o | x | | + H5 l5 | i5 G4 | x x x o o | x o | x | | + H6 l6 | i6 G5 | x x x x o o | x o | x | | + H7 l7 | i7 G6 | x x x x x o o | x o | x | | + -- -- | -- G7 | | | | | + ======================================================================================================================= + -- -- | i1 -- | | | | | + -- -- | i2 -- | | | | | + -- -- | i3 -- | | | | | + K4 l4 | i4 H3 | x | x | x | x | + K5 l5 | i5 H4 | x x | x | x | x | + K6 l6 | i6 H5 | x x x | x | x | x | + K7 l7 | i7 H6 | x x x x | x | x | x | + -- -- | -- H7 | | | | | + ======================================================================================================================= + """ # noqa: E501 + assert step > 1, "step should be larger than 1 in multi-step attention mask." + assert step <= 4, "Currently only a step of 4 or smaller is supported!" + + s = attn_mask.shape[-1] + zero_mask = torch.ones_like(attn_mask).bool() + mask_2_1 = attn_mask.clone().detach() + mask_2_1[:, :, :, :-1] = mask_2_1[:, :, :, 1:] + mask_2_2 = torch.ones_like(attn_mask).bool() + for i in range(1, s - 1): + mask_2_2[:, :, i, i] = False + + if step == 2: + attn_mask = torch.cat( + ( + torch.cat((attn_mask, zero_mask), dim=-1), + torch.cat((mask_2_1, mask_2_2), dim=-1), + ), + dim=-2, ) + return attn_mask - self.activation_func = F.silu + mask_3_1 = mask_2_1.clone().detach() + mask_3_1[:, :, :, :-1] = mask_3_1[:, :, :, 1:] + mask_3_2 = mask_2_2.clone().detach() + mask_3_2[:, :, :, :-1] = mask_3_2[:, :, :, 1:] + mask_3_2[:, :, 1, 0] = True + mask_3_3 = mask_2_2.clone().detach() + mask_3_3[:, :, 1, 1] = True - self.linear = torch.nn.Linear( - config.hidden_size, - config.hidden_size, - dtype=config.params_dtype, - device=device, + if step == 3: + attn_mask = torch.cat( + ( + torch.cat((attn_mask, zero_mask, zero_mask), dim=-1), + torch.cat((mask_2_1, mask_2_2, zero_mask), dim=-1), + torch.cat((mask_3_1, mask_3_2, mask_3_3), dim=-1), + ), + dim=-2, ) + return attn_mask - def forward(self, x): - """Forward function.""" - y = self.linear(x) - return x + self.activation_func(y), None + mask_4_1 = mask_3_1.clone().detach() + mask_4_1[:, :, :, :-1] = mask_4_1[:, :, :, 1:] + mask_4_2 = mask_3_2.clone().detach() + mask_4_2[:, :, :, :-1] = mask_4_2[:, :, :, 1:] + mask_4_2[:, :, 2, 0] = True + mask_4_3 = mask_3_3.clone().detach() + mask_4_3[:, :, :, :-1] = mask_4_3[:, :, :, 1:] + mask_4_3[:, :, 2, 1] = True + mask_4_4 = mask_3_3.clone().detach() + mask_4_4[:, :, 2, 2] = True - -class MedusaHead(MegatronModule): - """MedusaHead impl following TensorRT-LLM's model definition. - - https://github.com/NVIDIA/TensorRT-LLM/blob/main/tensorrt_llm/models/medusa/model.py - Medusa head consists of several MedusaLayers and an lm_head. - """ - - def __init__(self, config, vocab_size: int, num_layers: int = 1, parallel_output: bool = True): - """Constructor. - - Args: - config: MCore transformer config - vocab_size: vocabulary size - num_layers: number of Medusa layers - parallel_output: if False, then all_gather the logits - """ - super().__init__(config=config) - - self.medusa_layers = torch.nn.ModuleList([MedusaLayer(config) for _ in range(num_layers)]) - - self.lm_head = tensor_parallel.ColumnParallelLinear( - config.hidden_size, - vocab_size, - config=config, - init_method=config.init_method, - bias=False, - skip_bias_add=False, - gather_output=not parallel_output, - skip_weight_param_allocation=False, - ) - - def load_state_dict_post_hook(module, incompatible_keys): - incompatible_keys.missing_keys.clear() - incompatible_keys.unexpected_keys.clear() - - self.register_load_state_dict_post_hook(load_state_dict_post_hook) - - def forward(self, x): - """Forward function.""" - for layer in self.medusa_layers: - x, _ = layer(x) - return self.lm_head(x) - - def sharded_state_dict( - self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None - ) -> ShardedStateDict: - """Return MCore sharded_state_dict.""" - assert not sharded_offsets, "Unexpected sharded offsets" - sharded_state_dict = {} - layer_prefix = f"{prefix}medusa_layers." - for i, layer in enumerate(self.medusa_layers): - state_dict_prefix = f"{layer_prefix}{i}." - sharded_pp_offset = [] - layer_sharded_state_dict = layer.sharded_state_dict( - state_dict_prefix, sharded_pp_offset, metadata - ) - sharded_state_dict.update(layer_sharded_state_dict) - sharded_state_dict.update( - self.lm_head.sharded_state_dict(f"{prefix}lm_head.", sharded_offsets, metadata) - ) - return sharded_state_dict - - -@MedusaDMRegistry.register({GPTModel: "megatron.core.models.gpt.GPTModel"}) -class _DynamicMedusaGPTModel(MedusaModel): - """A ``megatron.core.models.gpt.GPTModel`` model with dynamic hyperparams.""" - - def _setup(self): - super()._setup() - self._register_temp_attribute("medusa_report_acc", True) - self._register_temp_attribute("medusa_freeze_base_model", True) - self._register_temp_attribute("calibration_mode", False) - - def modify( - self, - medusa_num_heads=0, - medusa_num_layers=0, - medusa_freeze_base_model=True, - medusa_report_acc=True, - ): - """Constructor. - - Args: - config: MedusaConfig that specifies the medusa head configuration as well as - weights of base model and medusa head. - """ - if self.config.pipeline_model_parallel_size > 1: - warnings.warn( - "Pipeline parallelism detected! _DynamicMedusaGPTModel only supports " - "pipeline parallelism during TensorRT-LLM checkpoint export." - ) - super().modify(medusa_num_heads=medusa_num_heads, medusa_num_layers=medusa_num_layers) - - self.medusa_report_acc = medusa_report_acc - self.medusa_freeze_base_model = medusa_freeze_base_model - - # Freeze all parameters - if self.medusa_freeze_base_model: - for name, param in self.named_parameters(): - param.requires_grad = False - - if self.post_process: - self.medusa_heads = torch.nn.ModuleList( - [ - MedusaHead(self.config, self.vocab_size, num_layers=self.medusa_num_layers) - for _ in range(self.medusa_num_heads) - ] - ) - - def _base_model_forward(self, *args, labels: torch.Tensor = None, **kwargs): - if self.post_process: - # Set the post_process to False such that the forward will return the hidden_state. - self.post_process = False - # Calling parent's forward to get hidden_states - hidden_states = GPTModel.forward(self, *args, labels=labels, **kwargs) - # Reset the post_process to True - self.post_process = True - else: - hidden_states = GPTModel.forward(self, *args, labels=None, **kwargs) - - return hidden_states - - def _medusa_forward(self, hidden_states): - draft_logits = [] - # Medusa heads forward. We want to run through all the heads just to make sure all modules - # are exercised during calibration. - for i, head in enumerate(self.medusa_heads): - new_logits, _ = head(hidden_states) - - draft_logits.append(new_logits) - - return draft_logits - - def forward(self, *args, labels: torch.Tensor = None, **kwargs): - """Forward pass of the Medusa GPTModel. - - Returns: - torch.Tensor: If labels are provided, then return lm_loss of all heads. Otherwise, - return the original logits. - """ - hidden_states = self._base_model_forward(*args, labels=labels, **kwargs) - - if not self.post_process: - return hidden_states - - output_weight = None - if self.share_embeddings_and_output_weights: - output_weight = self.shared_embedding_or_output_weight() - # Original output logits - logits, _ = self.output_layer(hidden_states, weight=output_weight) - - draft_logits = self._medusa_forward(hidden_states) - - if self.medusa_report_acc and labels is not None: - acc = [] - for i, _ in enumerate(self.medusa_heads): - gathered_logits = gather_from_tensor_model_parallel_region(draft_logits[i]) - medusa_top1 = gathered_logits.transpose(0, 1).argmax(dim=-1)[:, : -(1 + i)] - medusa_labels = labels[:, 1 + i :] - top1_p = torch.eq(medusa_labels, medusa_top1).sum() / medusa_top1.numel() - acc.append(top1_p) - - if get_tensor_model_parallel_rank() == 0: - print(f"Medusa Training Accuracy: {acc}") - - # Return the original logits untouched. - if labels is None: - # [s b h] => [b s h] - return logits.transpose(0, 1).contiguous() - - # Base model loss - # If medusa_freeze_base_model is set to True, - # the base model is frozen . - loss = self.compute_language_model_loss(labels, logits) - # Medusa loss - for i, _ in enumerate(self.medusa_heads): - medusa_labels = labels[:, 1 + i :] - medusa_loss = self.compute_language_model_loss( - medusa_labels, draft_logits[i][: -(1 + i), :] - ) - loss[:, 1 + i :] += medusa_loss - - return loss - - def sharded_state_dict( - self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None - ) -> ShardedStateDict: - """Override the shared_state_dict to take care medusa_heads.""" - assert not sharded_offsets, "Unexpected sharded offsets" - - sharded_state_dict = GPTModel.sharded_state_dict(self, prefix, sharded_offsets, metadata) - - if not hasattr(self, "medusa_heads") or self.medusa_heads is None: - return sharded_state_dict - - # This is a remedy for nn.ModuleList. GPTModel.sharded_state_dict() is calling into - # MegatronModule.sharded_state_dict() which requires all children to implement - # sharded_state_dict(). medusa_heads is an nn.ModuleList which only has state_dict() - # implemented. As a result, all the submodules will not be sharded. - # - # The remedy is to pop all medusa_heads* out and call the MedusaHead sharded_state_dict() - # again to populate the correct sharded_staet_dict. - extra_keys = [] - for key in sharded_state_dict: - if "medusa_heads" in key: - extra_keys += [key] - for key in extra_keys: - sharded_state_dict.pop(key, None) - - layer_prefix = f"{prefix}medusa_heads." - for i, layer in enumerate(self.medusa_heads): - layer_sharded_state_dict = layer.sharded_state_dict(f"{layer_prefix}{i}.", [], metadata) - sharded_state_dict.update(layer_sharded_state_dict) - return sharded_state_dict - - def pseudo_speculative_generate(self, *args, steps=1, **kwargs): - """Pseudo generate of the Medusa GPTModel. - - Returns: - base_token (torch.Tensor): token from base model - draft_tokens (torch.Tensor): draft tokens from medusa heads - """ - hidden_states = self._base_model_forward(*args, labels=None, **kwargs) - - if not self.post_process: - return hidden_states - - output_weight = None - if self.share_embeddings_and_output_weights: - output_weight = self.shared_embedding_or_output_weight() - # Original output logits - logits, _ = self.output_layer(hidden_states, weight=output_weight) - - draft_logits = self._medusa_forward(hidden_states) - - logits = gather_from_tensor_model_parallel_region(logits.transpose(0, 1).contiguous()) - draft_logits = [ - gather_from_tensor_model_parallel_region(logit.transpose(0, 1).contiguous()) - for logit in draft_logits - ] - - # [b, s] - base_token = logits[:, -1:].argmax(dim=-1) - draft_tokens = [logit[:, -1:].argmax(dim=-1) for logit in draft_logits] - draft_tokens = torch.cat(draft_tokens, dim=-1) - - return base_token, draft_tokens + attn_mask = torch.cat( + ( + torch.cat((attn_mask, zero_mask, zero_mask, zero_mask), dim=-1), + torch.cat((mask_2_1, mask_2_2, zero_mask, zero_mask), dim=-1), + torch.cat((mask_3_1, mask_3_2, mask_3_3, zero_mask), dim=-1), + torch.cat((mask_4_1, mask_4_2, mask_4_3, mask_4_4), dim=-1), + ), + dim=-2, + ) + return attn_mask class EagleLanguageModelEmbedding(LanguageModelEmbedding): @@ -499,14 +401,14 @@ class EagleModule(MegatronModule): num_layers: number of Eagle layers rotary_pos_emb: If None, use the default Llama-3.1 rope (GPT-NeoX). """ - super().__init__(config=config) - + # Override transformer_config before superclass initialization self._num_eagle_layers = num_layers self._use_input_layernorm_in_first_layer = use_input_layernorm_in_first_layer self._use_mtp_layernorm = use_mtp_layernorm self._num_aux_hidden_states = num_aux_hidden_states - eagle_config = self._get_eagle_transformer_config(config) + super().__init__(config=eagle_config) + eagle_transformer_layer_spec = self._get_eagle_transformer_layer_spec(eagle_config) if self._num_aux_hidden_states > 0: @@ -540,7 +442,6 @@ class EagleModule(MegatronModule): bias=bias, ).to(device) - # [TODO]: chenhany separate the rope from the base model is necessary self.rotary_pos_emb = rotary_pos_emb # Eagle does not use the final_layernorm in decoder. @@ -576,17 +477,32 @@ class EagleModule(MegatronModule): tp_comm_buffer_name="qkv", ) + # Sanity check + if self.decoder.layers[0].self_attention.attention_type == AttnMaskType.arbitrary: + raise ValueError("EAGLE-3 must use arbitrary attention mask.") + def _get_eagle_transformer_config(self, base_model_config): eagle_config = copy.deepcopy(base_model_config) eagle_config.num_layers = self._num_eagle_layers + # Unset the PP config. eagle_config.pipeline_model_parallel_size = 1 - # Unset the uneven PP config. + eagle_config.virtual_pipeline_model_parallel_size = None eagle_config.num_layers_in_first_pipeline_stage = None eagle_config.num_layers_in_last_pipeline_stage = None return eagle_config def _get_eagle_transformer_layer_spec(self, eagle_config): - transformer_layer_spec = get_gpt_modelopt_spec(eagle_config, remap_te_layernorm=True) + """Get the TransformerLayer implementation spec. + + IMPORTANT: EagleModule must use arbitrary_attention_mask since we need to + manupulate the mask to compute the correct loss. The default + causal mask will result in leaking. + """ + transformer_layer_spec = get_gpt_modelopt_spec( + eagle_config, + remap_te_layernorm=True, + use_arbitrary_attention_mask=True, + ) # If heterogenous layers (e.g. DeepSeek), transformer_layer_spec is a # TransformerBlockSubmodules instead. We use the last layer_specs. if "TransformerBlockSubmodules" in str(type(transformer_layer_spec)): @@ -622,7 +538,6 @@ class EagleModule(MegatronModule): def forward( self, - input_ids: torch.Tensor, embeddings: torch.Tensor, hidden_states: torch.Tensor, attention_mask: torch.Tensor, @@ -632,8 +547,12 @@ class EagleModule(MegatronModule): extra_block_kwargs: dict | None = None, ) -> torch.Tensor: """Forward function.""" - # input_ids [b, s] - seq_len = input_ids.shape[-1] + # NOTE: Even if sequence_parallel is used, the rotary_seq_len must be in the original + # length. Since we get the seq_len from hidden_states.shape[0], we need to + # multiply the the tp back. + rotary_seq_len = hidden_states.shape[0] + if self.config.sequence_parallel: + rotary_seq_len *= self.config.tensor_model_parallel_size if self._use_mtp_layernorm: embeddings = self.enorm(embeddings) @@ -651,10 +570,8 @@ class EagleModule(MegatronModule): decoder_input = hidden_states if rotary_pos_emb is None: - # For MLA, rotary_pos_emb is computed per attention. - # [TODO] (chenhany): multi_latent_attention case seems wrong when training the 2nd loss rotary_pos_emb = ( - None if self.config.multi_latent_attention else self.rotary_pos_emb(seq_len) + None if self.config.multi_latent_attention else self.rotary_pos_emb(rotary_seq_len) ) self._next_hidden_states_input = None @@ -695,18 +612,34 @@ class EagleLlama3Module(EagleModule): use_mtp_layernorm: bool = False, bias: bool = False, num_aux_hidden_states: int = 0, - ffn_hidden_size: int = 0, + ffn_hidden_size: int | None = 0, ): """Constructor.""" - eagle_config = copy.deepcopy(config) + eagle_config = Llama31Config8B( + # Getting ModelParallelConfig from the base model + tensor_model_parallel_size=config.tensor_model_parallel_size, + sequence_parallel=config.sequence_parallel, + expert_tensor_parallel_size=config.expert_tensor_parallel_size, + use_cpu_initialization=config.use_cpu_initialization, + fp16=config.fp16, + bf16=config.bf16, + params_dtype=config.params_dtype, + # Override hidden_size and ffn_hidden_size from the base model + hidden_size=config.hidden_size, + ffn_hidden_size=config.ffn_hidden_size, + ) - # Make sure the transformer is Llama3 style - eagle_config.kv_channels = 128 - eagle_config.num_attention_heads = eagle_config.hidden_size // 128 - eagle_config.num_query_groups = 8 - eagle_config.num_moe_experts = None - eagle_config.multi_latent_attention = False + # If base model is using MHA/GQA, then use the same config to simply KV-cache impl. + if config.kv_channels is not None: + eagle_config.kv_channels = config.kv_channels + if config.num_attention_heads > 0: + eagle_config.num_attention_heads = config.num_attention_heads + else: + eagle_config.num_attention_heads = eagle_config.hidden_size // eagle_config.kv_channels + if config.num_query_groups is not None: + eagle_config.num_query_groups = config.num_query_groups + # Override ffn_hidden_size if provided to widen the transformer. if ffn_hidden_size > 0: eagle_config.ffn_hidden_size = ffn_hidden_size @@ -771,6 +704,7 @@ class _DynamicEagleGPTModel(EagleModel): eagle_disable_moe=False, draft_vocab_size=0, use_mtp_layernorm=False, + parallel_draft_step=1, eagle_self_logit_distillation=True, eagle_freeze_base_model=True, eagle_report_acc=True, @@ -781,6 +715,12 @@ class _DynamicEagleGPTModel(EagleModel): "Pipeline parallelism detected! _DynamicEagleGPTModel only supports " "pipeline parallelism during TensorRT-LLM checkpoint export." ) + + # Since there is a chance that EAGLE3 can have heterogenous layers (1st layer + # qkv is 2x large than the rest), we enable MCore hetereogeneous checkpoint. + if hasattr(self.config, "hetereogenous_dist_checkpoint"): + self.config.hetereogenous_dist_checkpoint = True + super().modify( eagle_num_layers=eagle_num_layers, use_input_layernorm_in_first_layer=use_input_layernorm_in_first_layer, @@ -791,6 +731,7 @@ class _DynamicEagleGPTModel(EagleModel): eagle_disable_moe=eagle_disable_moe, draft_vocab_size=draft_vocab_size, use_mtp_layernorm=use_mtp_layernorm, + parallel_draft_step=parallel_draft_step, ) self.eagle_report_acc = eagle_report_acc self.eagle_self_logit_distillation = eagle_self_logit_distillation @@ -799,8 +740,8 @@ class _DynamicEagleGPTModel(EagleModel): # EAGLE-3 auxiluary hidden_states (only work for TP+EP, does not work for PP) self._aux_hidden_states = [] - if self.position_embedding_type != "rope": - raise ValueError("For EAGLE, only rotary embedding is supported") + if self.position_embedding_type not in ["rope", "yarn"]: + raise ValueError("For EAGLE, only RoPE or YaRN embedding are supported") if not self.pre_process and self.post_process: self.embedding = EagleLanguageModelEmbedding( @@ -872,8 +813,13 @@ class _DynamicEagleGPTModel(EagleModel): skip_weight_param_allocation=False, ) - def _get_eagle_input_hidden_states(self, hidden_states): - """When _aux_hidden_states is not empty, then this is EAGLE-3.""" + def _get_eagle_input_hidden_states(self, hidden_states: torch.Tensor, apply_fc: bool = True): + """When _aux_hidden_states is not empty, then this is EAGLE-3. + + Args: + hidden_states: last hidden_states + apply_fc: whether to apply EAGLE3 fc + """ if len(self._aux_hidden_states) == 0: return hidden_states @@ -881,8 +827,11 @@ class _DynamicEagleGPTModel(EagleModel): aux_hidden_states = torch.cat(self._aux_hidden_states, dim=-1) self._aux_hidden_states.clear() - # [s / TP, b, 3h] -> [s / TP, b, h] - return self.eagle_module.fc(aux_hidden_states)[0] + if apply_fc: + # [s / TP, b, 3h] -> [s / TP, b, h] + return self.eagle_module.fc(aux_hidden_states)[0] + else: + return aux_hidden_states def _get_eagle_module_inputs( self, @@ -892,124 +841,7 @@ class _DynamicEagleGPTModel(EagleModel): position_ids: torch.Tensor, features: torch.Tensor | None = None, ): - """Getting EAGLE module inputs. - - i0 i1 i2 i3 i4 i5 i6 i7 (base input_ids) - ======================= - h0 h1 h2 h3 h4 h5 h6 h7 (base hidden_states) - l0 l1 l2 l3 l4 l5 l6 l7 (base labels) - - - (1st) | i1 i2 i3 i4 i5 i6 i7 -- | - (out) | h0 h1 h2 h3 h4 h5 h6 h7 | - ========================================= - f1 l1 | i1 h0 | x | - f2 l2 | i2 h1 | x x | - f3 l3 | i3 h2 | x x x | - f4 l4 | i4 h3 | x x x x | - f5 l5 | i5 h4 | x x x x x | - f6 l6 | i6 h5 | x x x x x x | - f7 l7 | i7 h6 | x x x x x x x | - -- -- | -- h7 | o o o o o o o o | - ========================================= - - - (2nd) | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | - (out) | h0 h1 h2 h3 h4 h5 h6 h7 | -- F1 F2 F3 F4 F5 F6 F7 | - =================================================================== - F1 l1 | i1 h0 | x | | - F2 l2 | i2 h1 | x x | | - F3 l3 | i3 h2 | x x x | | - F4 l4 | i4 h3 | x x x x | | - F5 l5 | i5 h4 | x x x x x | | - F6 l6 | i6 h5 | x x x x x x | | - F7 l7 | i7 h6 | x x x x x x x | | - -- -- | -- h7 | o o o o o o o o | | - =================================================================== - -- -- | i1 -- | | | - G2 l2 | i2 F1 | x o | x | - G3 l3 | i3 F2 | x x o | x | - G4 l4 | i4 F3 | x x x o | x | - G5 l5 | i5 F4 | x x x x o | x | - G6 l6 | i6 F5 | x x x x x o | x | - G7 l7 | i7 F6 | x x x x x x o | x | - -- -- | -- F7 | | | - =================================================================== - - - (3rd) | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | - (out) | h0 h1 h2 h3 h4 h5 h6 h7 | -- F1 F2 F3 F4 F5 F6 F7 | -- -- G2 G3 G4 G5 G6 G7 | - ============================================================================================= - F1 l1 | i1 h0 | x | | | - F2 l2 | i2 h1 | x x | | | - F3 l3 | i3 h2 | x x x | | | - F4 l4 | i4 h3 | x x x x | | | - F5 l5 | i5 h4 | x x x x x | | | - F6 l6 | i6 h5 | x x x x x x | | | - F7 l7 | i7 h6 | x x x x x x x | | | - -- -- | -- h7 | o o o o o o o o | | | - ============================================================================================= - -- -- | i1 -- | | | | - G2 l2 | i2 F1 | x o | x | | - G3 l3 | i3 F2 | x x o | x | | - G4 l4 | i4 F3 | x x x o | x | | - G5 l5 | i5 F4 | x x x x o | x | | - G6 l6 | i6 F5 | x x x x x o | x | | - G7 l7 | i7 F6 | x x x x x x o | x | | - -- -- | -- F7 | | | | - ============================================================================================= - -- -- | i1 -- | | | | - -- -- | i2 -- | | | | - H3 l3 | i3 G2 | x o o | x o | x | - H4 l4 | i4 G3 | x x o o | x o | x | - H5 l5 | i5 G4 | x x x o o | x o | x | - H6 l6 | i6 G5 | x x x x o o | x o | x | - H7 l7 | i7 G6 | x x x x x o o | x o | x | - -- -- | -- G7 | | | | - ============================================================================================= - - - (4th) | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | i1 i2 i3 i4 i5 i6 i7 -- | - (out) | h0 h1 h2 h3 h4 h5 h6 h7 | -- F1 F2 F3 F4 F5 F6 F7 | -- -- G2 G3 G4 G5 G6 G7 | -- -- -- H3 H4 H5 H6 H7 | - ======================================================================================================================= - F1 l1 | i1 h0 | x | | | | - F2 l2 | i2 h1 | x x | | | | - F3 l3 | i3 h2 | x x x | | | | - F4 l4 | i4 h3 | x x x x | | | | - F5 l5 | i5 h4 | x x x x x | | | | - F6 l6 | i6 h5 | x x x x x x | | | | - F7 l7 | i7 h6 | x x x x x x x | | | | - -- -- | -- h7 | o o o o o o o o | | | | - ======================================================================================================================= - -- -- | i1 -- | | | | | - G2 l2 | i2 F1 | x o | x | | | - G3 l3 | i3 F2 | x x o | x | | | - G4 l4 | i4 F3 | x x x o | x | | | - G5 l5 | i5 F4 | x x x x o | x | | | - G6 l6 | i6 F5 | x x x x x o | x | | | - G7 l7 | i7 F6 | x x x x x x o | x | | | - -- -- | -- F7 | | | | | - ======================================================================================================================= - -- -- | i1 -- | | | | | - -- -- | i2 -- | | | | | - H3 l3 | i3 G2 | x o o | x o | x | | - H4 l4 | i4 G3 | x x o o | x o | x | | - H5 l5 | i5 G4 | x x x o o | x o | x | | - H6 l6 | i6 G5 | x x x x o o | x o | x | | - H7 l7 | i7 G6 | x x x x x o o | x o | x | | - -- -- | -- G7 | | | | | - ======================================================================================================================= - -- -- | i1 -- | | | | | - -- -- | i2 -- | | | | | - -- -- | i3 -- | | | | | - K4 l4 | i4 H3 | x | x | x | x | - K5 l5 | i5 H4 | x x | x | x | x | - K6 l6 | i6 H5 | x x x | x | x | x | - K7 l7 | i7 H6 | x x x x | x | x | x | - -- -- | -- H7 | | | | | - ======================================================================================================================= - """ # noqa: E501 - s = input_ids.shape[1] + """Getting EAGLE module inputs.""" b = hidden_states.shape[1] h = hidden_states.shape[2] @@ -1026,7 +858,64 @@ class _DynamicEagleGPTModel(EagleModel): eagle_inputs = {} - if features is None: + if self.parallel_draft_step > 1: + eagle_inputs["input_ids"] = padded_input_ids + eagle_inputs["position_ids"] = position_ids + if rotary_pos_emb is not None: + eagle_inputs["rotary_pos_emb"] = rotary_pos_emb + else: + # [TODO] (yeyu): there will be problem here with MLA + eagle_inputs["rotary_pos_emb"] = None + + if self.config.sequence_parallel: + gathered_hidden_states = gather_from_sequence_parallel_region(hidden_states) + else: + gathered_hidden_states = hidden_states + eagle_inputs["hidden_states"] = gathered_hidden_states + + for i in range(self.parallel_draft_step - 1): + eagle_inputs["input_ids"] = torch.cat( + ( + eagle_inputs["input_ids"], + torch.full( + padded_input_ids.shape, + getattr(self, f"mask_token_{i}"), + device=padded_input_ids.device, + dtype=padded_input_ids.dtype, + ), + ), + dim=-1, + ) + + eagle_inputs["hidden_states"] = torch.cat( + ( + eagle_inputs["hidden_states"], + torch.zeros( + (1 + i, b, h), dtype=hidden_states.dtype, device=hidden_states.device + ), + gathered_hidden_states[: -(1 + i)], + ), + dim=0, + ) + + eagle_inputs["position_ids"] = torch.cat( + (eagle_inputs["position_ids"], position_ids), dim=-1 + ) + + if rotary_pos_emb is not None: + eagle_inputs["rotary_pos_emb"] = torch.cat( + (eagle_inputs["rotary_pos_emb"], rotary_pos_emb), dim=0 + ) + + if self.config.sequence_parallel: + eagle_inputs["hidden_states"] = scatter_to_sequence_parallel_region( + eagle_inputs["hidden_states"] + ) + + eagle_inputs["attention_mask"] = set_multi_step_attention_mask( + attn_mask, self.parallel_draft_step + ) + elif features is None: eagle_inputs["input_ids"] = padded_input_ids eagle_inputs["hidden_states"] = hidden_states eagle_inputs["attention_mask"] = attn_mask @@ -1057,21 +946,7 @@ class _DynamicEagleGPTModel(EagleModel): eagle_inputs["hidden_states"] ) - zero_mask = torch.ones_like(attn_mask).bool() - mask_2_1 = attn_mask.clone().detach() - mask_2_1[:, :, :, :-1] = mask_2_1[:, :, :, 1:] - mask_2_2 = torch.ones_like(attn_mask).bool() - for i in range(1, s - 1): - mask_2_2[:, :, i, i] = False - - attn_mask = torch.cat( - ( - torch.cat((attn_mask, zero_mask), dim=-1), - torch.cat((mask_2_1, mask_2_2), dim=-1), - ), - dim=-2, - ) - eagle_inputs["attention_mask"] = attn_mask + eagle_inputs["attention_mask"] = set_multi_step_attention_mask(attn_mask, 2) eagle_inputs["position_ids"] = torch.cat((position_ids, position_ids), dim=-1) if rotary_pos_emb is not None: @@ -1104,31 +979,7 @@ class _DynamicEagleGPTModel(EagleModel): eagle_inputs["hidden_states"] ) - zero_mask = torch.ones_like(attn_mask).bool() - mask_2_1 = attn_mask.clone().detach() - mask_2_1[:, :, :, :-1] = mask_2_1[:, :, :, 1:] - mask_2_2 = torch.ones_like(attn_mask).bool() - for i in range(1, s - 1): - mask_2_2[:, :, i, i] = False - - mask_3_1 = mask_2_1.clone().detach() - mask_3_1[:, :, :, :-1] = mask_3_1[:, :, :, 1:] - mask_3_2 = mask_2_2.clone().detach() - mask_3_2[:, :, :, :-1] = mask_3_2[:, :, :, 1:] - mask_3_2[:, :, 1, 0] = True - mask_3_3 = mask_2_2.clone().detach() - mask_3_3[:, :, 1, 1] = True - - attn_mask = torch.cat( - ( - torch.cat((attn_mask, zero_mask, zero_mask), dim=-1), - torch.cat((mask_2_1, mask_2_2, zero_mask), dim=-1), - torch.cat((mask_3_1, mask_3_2, mask_3_3), dim=-1), - ), - dim=-2, - ) - - eagle_inputs["attention_mask"] = attn_mask + eagle_inputs["attention_mask"] = set_multi_step_attention_mask(attn_mask, 3) eagle_inputs["position_ids"] = torch.cat( (position_ids, position_ids, position_ids), dim=-1 ) @@ -1166,43 +1017,7 @@ class _DynamicEagleGPTModel(EagleModel): eagle_inputs["hidden_states"] ) - zero_mask = torch.ones_like(attn_mask).bool() - mask_2_1 = attn_mask.clone().detach() - mask_2_1[:, :, :, :-1] = mask_2_1[:, :, :, 1:] - mask_2_2 = torch.ones_like(attn_mask).bool() - for i in range(1, s - 1): - mask_2_2[:, :, i, i] = False - - mask_3_1 = mask_2_1.clone().detach() - mask_3_1[:, :, :, :-1] = mask_3_1[:, :, :, 1:] - mask_3_2 = mask_2_2.clone().detach() - mask_3_2[:, :, :, :-1] = mask_3_2[:, :, :, 1:] - mask_3_2[:, :, 1, 0] = True - mask_3_3 = mask_2_2.clone().detach() - mask_3_3[:, :, 1, 1] = True - - mask_4_1 = mask_3_1.clone().detach() - mask_4_1[:, :, :, :-1] = mask_4_1[:, :, :, 1:] - mask_4_2 = mask_3_2.clone().detach() - mask_4_2[:, :, :, :-1] = mask_4_2[:, :, :, 1:] - mask_4_2[:, :, 2, 0] = True - mask_4_3 = mask_3_3.clone().detach() - mask_4_3[:, :, :, :-1] = mask_4_3[:, :, :, 1:] - mask_4_3[:, :, 2, 1] = True - mask_4_4 = mask_3_3.clone().detach() - mask_4_4[:, :, 2, 2] = True - - attn_mask = torch.cat( - ( - torch.cat((attn_mask, zero_mask, zero_mask, zero_mask), dim=-1), - torch.cat((mask_2_1, mask_2_2, zero_mask, zero_mask), dim=-1), - torch.cat((mask_3_1, mask_3_2, mask_3_3, zero_mask), dim=-1), - torch.cat((mask_4_1, mask_4_2, mask_4_3, mask_4_4), dim=-1), - ), - dim=-2, - ) - - eagle_inputs["attention_mask"] = attn_mask + eagle_inputs["attention_mask"] = set_multi_step_attention_mask(attn_mask, 4) eagle_inputs["position_ids"] = torch.cat( (position_ids, position_ids, position_ids, position_ids), dim=-1 ) @@ -1249,6 +1064,7 @@ class _DynamicEagleGPTModel(EagleModel): inference_params: InferenceParams = None, packed_seq_params: PackedSeqParams = None, extra_block_kwargs: dict | None = None, + return_eagle_inputs: bool = False, ): # Word and rotary positional embeddings if decoder_input is not None: @@ -1262,10 +1078,12 @@ class _DynamicEagleGPTModel(EagleModel): extra_kwargs = {"packed_seq_params": None} if mcore_version_higher_than("0.9.0") else {} + rotary_pos_emb = None + yarn_mscale = 1.0 if self.config.multi_latent_attention: # For MLA, rotary_pos_emb is computed per attention. rotary_pos_emb = None - else: + elif self.position_embedding_type == "rope": rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( inference_params, self.decoder, @@ -1274,8 +1092,22 @@ class _DynamicEagleGPTModel(EagleModel): **extra_kwargs, ) rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len) + elif self.position_embedding_type == "yarn": + rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( + inference_params, + None, + decoder_input, + self.config, + **extra_kwargs, + ) + rotary_pos_emb, yarn_mscale = self.rotary_pos_emb(rotary_seq_len) + else: + raise ValueError( + f"Only RoPE or YaRN are supported but got {self.position_embedding_type}" + ) - # Run base modeld decoder forward. + # [TODO]: yarn_mscale needs to be passed into TransformerBlock forward when supported. + # Now the default value for yarn_mscale = 1.0 hidden_states = self.decoder( hidden_states=decoder_input, attention_mask=attention_mask, @@ -1285,7 +1117,10 @@ class _DynamicEagleGPTModel(EagleModel): **(extra_block_kwargs or {}), ) - return hidden_states + if return_eagle_inputs: + return hidden_states, decoder_input + else: + return hidden_states, None def _eagle_forward( self, @@ -1296,7 +1131,6 @@ class _DynamicEagleGPTModel(EagleModel): extra_block_kwargs: dict | None = None, ): eagle_hidden_states, eagle_hidden_states_pre_final_layernorm = self.eagle_module( - eagle_inputs["input_ids"], eagle_inputs["embedding"], eagle_inputs["hidden_states"], eagle_inputs["attention_mask"], @@ -1323,12 +1157,17 @@ class _DynamicEagleGPTModel(EagleModel): inference_params: InferenceParams = None, packed_seq_params: PackedSeqParams = None, extra_block_kwargs: dict | None = None, + return_eagle_inputs: bool = False, + loss_decay_factor: float = 0.9, **kwargs, ) -> torch.Tensor: if input_ids is not None and (position_ids is None or attention_mask is None): attention_mask, position_ids = get_default_attention_mask_and_position_ids(input_ids) - hidden_states = self._base_model_forward( + # When return_eagle_inputs is True, return decoder_input_for_eagle. + # When LLM, decoder_input_for_eagle is just the text embeddings. However, when VLM + # decoder_input_for_eagle will also contain projected image/video embeddings. + hidden_states, decoder_input_for_eagle = self._base_model_forward( input_ids, position_ids, attention_mask, @@ -1336,20 +1175,35 @@ class _DynamicEagleGPTModel(EagleModel): inference_params, packed_seq_params, extra_block_kwargs, + return_eagle_inputs=return_eagle_inputs, ) # Typically, this is only the case when PP > 1. if not self.post_process: return hidden_states - # If EAGLE-3, aux_hidden_states are gathered by the forward_hook - eagle_module_input_hidden_states = self._get_eagle_input_hidden_states(hidden_states) - output_weight = None if self.share_embeddings_and_output_weights: output_weight = self.shared_embedding_or_output_weight() logits_sbh, _ = self.output_layer(hidden_states, weight=output_weight) + # If EAGLE-3, aux_hidden_states are gathered by the forward_hook + if return_eagle_inputs: + eagle_module_input_hidden_states = self._get_eagle_input_hidden_states( + hidden_states, apply_fc=False + ) + # In case of VLM, there will be other fields for pixels. + return { + "input_ids": input_ids, + "decoder_input": decoder_input_for_eagle, + "hidden_states": eagle_module_input_hidden_states, + "logits": logits_sbh, + } + else: + eagle_module_input_hidden_states = self._get_eagle_input_hidden_states( + hidden_states, apply_fc=True + ) + # Either inference or calibration mode, we want to make sure all weights have been exercised. # This makes sure all quantized weights have amax calibrated if inference_params is None or self.calibration_mode: @@ -1373,13 +1227,21 @@ class _DynamicEagleGPTModel(EagleModel): if labels is None: return logits_sbh.transpose(0, 1).contiguous() - # Base model loss # If eagle_freeze_base_model is set to True, # the base model is frozen . loss = self.compute_language_model_loss(labels, logits_sbh) - loss_0 = self._compute_eagle_loss(logits_sbh, labels, eagle_logits_0) loss = 0.0 * loss - loss[:, 1:] += 1.0 * loss_0 + + if self.parallel_draft_step > 1: + for i in range(self.parallel_draft_step): + eagle_logits = eagle_logits_0[i * labels.shape[1] : (i + 1) * labels.shape[1]] + loss_ = self._compute_eagle_loss(logits_sbh, labels, eagle_logits) + loss_ = loss_[:, i:] + loss[:, i + 1 :] += 1.0 * loss_ + return loss + + loss_0 = self._compute_eagle_loss(logits_sbh, labels, eagle_logits_0) + loss[:, 1:] += loss_decay_factor * loss_0 if self.eagle_report_acc and not self.training: acc = [] @@ -1420,7 +1282,7 @@ class _DynamicEagleGPTModel(EagleModel): loss_1 = self._compute_eagle_loss(logits_sbh, labels, eagle_logits_1) # [b, s - 2] loss_1 = loss_1[:, 1:] - loss[:, 2:] += 1.0 * loss_1 + loss[:, 2:] += loss_decay_factor**2 * loss_1 if self.eagle_report_acc and not self.training: acc = [] @@ -1462,7 +1324,7 @@ class _DynamicEagleGPTModel(EagleModel): loss_2 = self._compute_eagle_loss(logits_sbh, labels, eagle_logits_2) # [b, s - 3] loss_2 = loss_2[:, 2:] - loss[:, 3:] += 1.0 * loss_2 + loss[:, 3:] += loss_decay_factor**3 * loss_2 if self.eagle_report_acc and not self.training: acc = [] @@ -1504,7 +1366,7 @@ class _DynamicEagleGPTModel(EagleModel): loss_3 = self._compute_eagle_loss(logits_sbh, labels, eagle_logits_3) # [b, s - 4] loss_3 = loss_3[:, 3:] - loss[:, 4:] += 1.0 * loss_3 + loss[:, 4:] += loss_decay_factor**4 * loss_3 if self.eagle_report_acc and not self.training: acc = [] @@ -1526,121 +1388,176 @@ class _DynamicEagleGPTModel(EagleModel): return loss - def tree_decode(self, input_ids: torch.Tensor, tree=None): - self.eval() + def tree_decode(self, input_ids: torch.Tensor, tree: Tree): + """Tree-based decoding for EAGLE model using a mask-based approach. - choices = [] if tree is None else tree.choices - draft_len = len(choices) - num_additional_tokens = 0 if tree is None else tree.depth + This function implements a tree-based decoding strategy where each path of the tree + represents potential token sequences. The function uses attention masks to control + token dependencies and generate multiple candidate sequences in parallel. - micro_batch_size, seq_len = input_ids.size() - att_mask_len = seq_len + 1 + draft_len + Args: + input_ids (torch.Tensor): Input token IDs of shape [batch_size, seq_len] + treepaths (list[list[int]]): List of treepaths to decode - # Prepare attention masks - attention_mask = torch.tril( - torch.ones( - (1, att_mask_len, att_mask_len), - device=input_ids.device, - ), - ).view(1, 1, att_mask_len, att_mask_len) - attention_mask = attention_mask < 0.5 + Returns: + tuple: (base_token, base_draft_node, draft_tokens) + - base_token: The next token predicted by the base model + - base_draft_node: A TreeNode containing the base token prediction with a + hierarchical structure of child nodes, where each child node + represents a draft token generated by EAGLE + - draft_tokens: all the draft tokens generated by EAGLE + """ + # Initial setup and base model forward pass + padded_input_ids, seq_len = right_padding(input_ids) + attention_mask, position_ids = get_default_attention_mask_and_position_ids(padded_input_ids) - if draft_len > 0: - draft_attention_mask = torch.tensor( - tree.attn_mask, - device=input_ids.device, - ).view(1, 1, draft_len, draft_len) - attention_mask[0, 0, -draft_len:, -draft_len:] = draft_attention_mask - - # Prepare position ids and rope - position_ids = torch.arange( - att_mask_len, - dtype=torch.long, - device=input_ids.device, + # Get base model hidden states + hidden_states, _ = self._base_model_forward( + padded_input_ids, + position_ids, + attention_mask, ) - rotary_pos_emb = self.rotary_pos_emb(seq_len + 1 + num_additional_tokens) - if num_additional_tokens > 0: - rotary_pos_emb, additional_rotary_pos_emb = torch.split( - rotary_pos_emb, - (seq_len + 1, num_additional_tokens), - dim=0, + if not self.post_process: + return hidden_states + + # Generate base token prediction + output_weight = ( + self.shared_embedding_or_output_weight() + if self.share_embeddings_and_output_weights + else None + ) + logits_sbh, _ = self.output_layer(hidden_states, weight=output_weight) + logits_sbh = logits_sbh[:seq_len, :, :] + + base_token = ( + gather_from_tensor_model_parallel_region(logits_sbh)[-1:, :, :] + .argmax(dim=-1) + .transpose(0, 1) + ) + + # Early return if no steps needed + if not tree.root.children: + self._aux_hidden_states.clear() + return base_token, None, None + + # Prepare for tree decoding + eagle_ids = torch.cat((input_ids[:, 1:], base_token), dim=-1) + # EAGLE-3 + # Only the first iteration input_hidden_states are from aux_hidden_state layers + hidden_states = self._get_eagle_input_hidden_states(hidden_states) + + if self.config.sequence_parallel: + hidden_states = gather_from_sequence_parallel_region(hidden_states) + hidden_states = hidden_states[:seq_len, :, :] + + # relative id from [seq_len-1, seq_len] contains draft token position + # [seq_len, seq_len + num_child_level_1] contains the number of children for level 1 and so on + relative_ids = torch.tensor( + [seq_len - 1, *list(tree.num_children.values())], + device=input_ids.device, + ).cumsum(dim=0) + + draft_position_ids = torch.arange(relative_ids[-1], device=input_ids.device) + cur_pos = seq_len - 1 + for idx in range(len(relative_ids) - 1): + draft_position_ids[relative_ids[idx] : relative_ids[idx + 1]] = cur_pos + cur_pos += 1 + + draft_attention_mask = torch.full( + (1, 1, relative_ids[-1], relative_ids[-1]), True, device=input_ids.device + ).triu_(1) + draft_attention_mask[:, :, :, seq_len:] = True + draft_attention_mask[:, :, seq_len - 1 :, seq_len - 1 :] = tree.attention_mask + + draft_rotary_pos_emb = self.eagle_module.rotary_pos_emb(seq_len + tree.max_depth) + draft_rotary_pos_emb = torch.cat( + [draft_rotary_pos_emb[index : index + 1] for index in draft_position_ids], dim=0 + ) + + base_draft_node = TreeNode(base_token) + queue = deque([(base_draft_node, tree.root)]) + draft_tokens = [] + # Tree decoding loop + for step in range(tree.max_depth): + # Prepare inputs for EAGLE forward pass + padded_eagle_ids, seq_len, padded_hidden_states = right_padding( + eagle_ids, hidden_states ) - for i, c in enumerate(choices): - position_ids[seq_len + 1 + i] = seq_len + len(c) - rotary_pos_emb = torch.cat( - (rotary_pos_emb, additional_rotary_pos_emb[[len(c) - 1], :, :, :]), dim=0 + + if self.config.sequence_parallel: + padded_hidden_states = scatter_to_sequence_parallel_region(padded_hidden_states) + + eagle_attention_mask, eagle_position_ids = get_default_attention_mask_and_position_ids( + padded_eagle_ids + ) + length = eagle_ids.shape[-1] + eagle_attention_mask[:, :, :length, :length] = draft_attention_mask[ + :, :, :length, :length + ] + eagle_attention_mask[:, :, length:, length:] = True + eagle_position_ids[:length] = draft_position_ids[:length] + padded_rotary_pos_emb = self.eagle_module.rotary_pos_emb(padded_eagle_ids.shape[-1]) + padded_rotary_pos_emb[:length] = draft_rotary_pos_emb[:length] + + eagle_inputs = { + "input_ids": padded_eagle_ids, + "embedding": self.embedding( + input_ids=padded_eagle_ids, + position_ids=eagle_position_ids, + ), + "hidden_states": padded_hidden_states, + "attention_mask": eagle_attention_mask, + "rotary_pos_emb": padded_rotary_pos_emb, + } + + # Forward pass through EAGLE + _, eagle_logits, eagle_next_hidden_states_input = self._eagle_forward( + eagle_inputs, + output_weight, + ) + # Process EAGLE outputs + eagle_logits = eagle_logits[:seq_len, :, :] + if self.config.sequence_parallel: + eagle_next_hidden_states_input = gather_from_sequence_parallel_region( + eagle_next_hidden_states_input + ) + eagle_next_hidden_states_input = eagle_next_hidden_states_input[:seq_len, :, :] + # Generate and store top-k tokens for each tree node + for rel_idx in range(relative_ids[step], relative_ids[step + 1]): + draft_node, tree_node = queue.popleft() + n_topk = max(tree_node.children.keys()) + 1 if tree_node.children else 0 + # Get top-k tokens for current position + new_ids = ( + gather_from_tensor_model_parallel_region(eagle_logits)[ + rel_idx : rel_idx + 1, :, : + ] + .topk(n_topk, dim=-1)[1] + .squeeze(0) ) - position_ids = position_ids.unsqueeze(0) - - # Forward - decoder_input = self.embedding( - input_ids=input_ids, - position_ids=position_ids[:, :seq_len], - ) - hidden_states = self.decoder( - hidden_states=decoder_input, - attention_mask=attention_mask[:, :, :seq_len, :seq_len], - inference_params=None, - rotary_pos_emb=rotary_pos_emb[:seq_len, :, :, :], - packed_seq_params=None, - ) - logits, _ = self.output_layer(hidden_states) - - # [s b h] => [b s h] - all_logprob = gather_from_tensor_model_parallel_region(logits[[-1], :, :].transpose(0, 1)) - all_logprob = torch.softmax(all_logprob, dim=-1) - top_vals, top_ids = all_logprob[:, -1, :].topk(1, dim=-1) - - if num_additional_tokens == 0: - return top_ids - - new_tokens = top_ids - eagle_ids = torch.cat((input_ids[:, 1:], top_ids), dim=-1) - eagle_rotary_pos_emb = rotary_pos_emb[1 : 1 + eagle_ids.shape[-1], :, :, :] - eagle_hidden_states = hidden_states - - for i in range(num_additional_tokens): - eagle_position_ids = position_ids[:, 1 : 1 + eagle_ids.shape[-1]] - eagle_attn_mask = attention_mask[ - :, :, 1 : 1 + eagle_ids.shape[-1], 1 : 1 + eagle_ids.shape[-1] - ] - eagle_rotary_pos_emb = rotary_pos_emb[1 : 1 + eagle_ids.shape[-1], :, :, :] - - eagle_embeddings = self.embedding( - input_ids=eagle_ids, - position_ids=eagle_position_ids, - ) - - new_hidden_states = self.eagle_module( - eagle_embeddings, - eagle_hidden_states, - eagle_attn_mask, - rotary_pos_emb=eagle_rotary_pos_emb, - ) - - eagle_logits, _ = self.output_layer(new_hidden_states) - - all_logprob = gather_from_tensor_model_parallel_region( - eagle_logits[tree.relative_ids[i], :, :].transpose(0, 1) - ) - all_logprob = torch.softmax(all_logprob, dim=-1) - - for idx, tk in zip(tree.relative_ids[i], tree.top_k[i]): - if tk == 0: - continue - top_vals, top_ids = all_logprob[:, idx, :].topk(tk, dim=-1) - - new_tokens = torch.cat((new_tokens, top_ids), dim=-1) - eagle_ids = torch.cat((eagle_ids, top_ids), dim=-1) - - for _ in range(tk): - eagle_hidden_states = torch.cat( - (eagle_hidden_states, new_hidden_states[[idx], :, :]), dim=0 + for child_idx, child_node in tree_node.children.items(): + eagle_ids = torch.cat( + (eagle_ids, new_ids[:, child_idx : child_idx + 1]), dim=-1 ) + # value of the node is token id + new_draft_node = TreeNode(new_ids[:, child_idx]) + draft_tokens.append(new_ids[:, child_idx]) + draft_node.children[child_idx] = new_draft_node + queue.append((new_draft_node, child_node)) - return new_tokens + # Update hidden states for each branch + hidden_states = torch.cat( + ( + hidden_states, + eagle_next_hidden_states_input[rel_idx : rel_idx + 1].repeat( + len(tree_node.children), 1, 1 + ), + ), + dim=0, + ) + draft_tokens = torch.cat(draft_tokens, dim=-1) + return base_token, base_draft_node, draft_tokens def pseudo_speculative_generate( self, @@ -1657,7 +1574,7 @@ class _DynamicEagleGPTModel(EagleModel): attention_mask, position_ids = get_default_attention_mask_and_position_ids(padded_input_ids) - hidden_states = self._base_model_forward( + hidden_states, _ = self._base_model_forward( padded_input_ids, position_ids, attention_mask, @@ -1697,6 +1614,12 @@ class _DynamicEagleGPTModel(EagleModel): draft_tokens = [] for _ in range(steps): + if self.parallel_draft_step > 1: + for i in range(self.parallel_draft_step - 1): + eagle_ids = torch.cat( + (eagle_ids, getattr(self, f"mask_token_{i}").view((1, 1))), dim=-1 + ) + hidden_states = torch.cat((hidden_states, hidden_states[-1:]), dim=0) padded_eagle_ids, seq_len, padded_hidden_states = right_padding( eagle_ids, hidden_states ) @@ -1715,12 +1638,6 @@ class _DynamicEagleGPTModel(EagleModel): eagle_inputs["hidden_states"] = padded_hidden_states eagle_inputs["attention_mask"] = eagle_attention_mask - # if self.config.multi_latent_attention: - # # For MLA, rotary_pos_emb is computed per attention. - # rotary_pos_emb = None - # else: - # rotary_pos_emb = self.rotary_pos_emb(padded_eagle_ids.shape[-1]) - # [TODO] (chenhany): let the module compute itself eagle_inputs["rotary_pos_emb"] = None @@ -1736,13 +1653,26 @@ class _DynamicEagleGPTModel(EagleModel): ) eagle_next_hidden_states_input = eagle_next_hidden_states_input[:seq_len, :, :] - draft_token = ( - gather_from_tensor_model_parallel_region(eagle_logits)[-1:, :, :] - .argmax(dim=-1) - .transpose(0, 1) - ) + if self.parallel_draft_step > 1: + draft_token = ( + gather_from_tensor_model_parallel_region(eagle_logits)[ + -self.parallel_draft_step :, :, : + ] + .argmax(dim=-1) + .transpose(0, 1) + ) + else: + draft_token = ( + gather_from_tensor_model_parallel_region(eagle_logits)[-1:, :, :] + .argmax(dim=-1) + .transpose(0, 1) + ) if self.draft_vocab_size > 0: draft_token += self.eagle_module.d2t[draft_token] + + if self.parallel_draft_step > 1: + return base_token, draft_token + draft_tokens.append(draft_token) eagle_ids = torch.cat((eagle_ids, draft_token), dim=-1) @@ -1755,386 +1685,6 @@ class _DynamicEagleGPTModel(EagleModel): return base_token, draft_tokens -@MTPDMRegistry.register({GPTModel: "megatron.core.models.gpt.GPTModel"}) -class _DynamicMTPGPTModel(MTPModel): - """A ``megatron.core.models.gpt.GPTModel`` model with dynamic hyperparams.""" - - def _setup(self): - super()._setup() - self._register_temp_attribute("mtp_self_logit_distillation", True) - self._register_temp_attribute("mtp_freeze_base_model", True) - self._register_temp_attribute("calibration_mode", False) - - def modify( - self, - mtp_num_layers=0, - mtp_num_module=0, - mtp_freeze_list=[], - use_last_layernorm=False, - mtp_self_logit_distillation=True, - mtp_freeze_base_model=True, - mtp_report_acc=True, - ): - if self.config.pipeline_model_parallel_size > 1: - warnings.warn( - "Pipeline parallelism detected! _DynamicMTPGPTModel only supports " - "pipeline parallelism during TensorRT-LLM checkpoint export." - ) - super().modify( - mtp_num_layers=mtp_num_layers, - mtp_num_module=mtp_num_module, - mtp_freeze_list=mtp_freeze_list, - use_last_layernorm=use_last_layernorm, - ) - self.mtp_report_acc = mtp_report_acc - self.mtp_self_logit_distillation = mtp_self_logit_distillation - self.mtp_freeze_base_model = mtp_freeze_base_model - - if self.position_embedding_type != "rope": - raise ValueError("For MTP, only rotary embedding is supported") - - if not self.pre_process and self.post_process: - self.embedding = EagleLanguageModelEmbedding( - config=self.config, - vocab_size=self.vocab_size, - max_sequence_length=self.max_sequence_length, - position_embedding_type=self.position_embedding_type, - ) - - # Freeze all parameters - if self.mtp_freeze_base_model: - for name, param in self.named_parameters(): - param.requires_grad = False - - # Only the last PP stage has the additional projection and decoder layer. - # This is to simplify the export. - if self.post_process: - self.mtp = torch.nn.ModuleList() - for i in range(self.mtp_num_module): - mtp = EagleModule( - self.config, - self.rotary_pos_emb, - self.mtp_num_layers, - self.use_last_layernorm, - use_input_layernorm_in_first_layer=True, - use_mtp_layernorm=True, - bias=False, - ) - if i in self.mtp_freeze_list: - for name, param in mtp.named_parameters(): - param.requires_grad = False - self.mtp.append(mtp) - - self.kld = logits_kld_loss - - def _base_model_forward( - self, - input_ids: torch.Tensor, - position_ids: torch.Tensor, - attention_mask: torch.Tensor, - decoder_input: torch.Tensor = None, - inference_params: InferenceParams = None, - packed_seq_params: PackedSeqParams = None, - extra_block_kwargs: dict | None = None, - ): - # Word and rotary positional embeddings - if decoder_input is not None: - pass - elif self.pre_process: - decoder_input = self.embedding(input_ids=input_ids, position_ids=position_ids) - else: - # intermediate stage of pipeline - # decoder will get hidden_states from decoder.input_tensor - decoder_input = None - - extra_kwargs = {"packed_seq_params": None} if mcore_version_higher_than("0.9.0") else {} - - if self.config.multi_latent_attention: - # For MLA, rotary_pos_emb is computed per attention. - rotary_pos_emb = None - else: - rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len( - inference_params, - self.decoder, - decoder_input, - self.config, - **extra_kwargs, - ) - rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len) - - # Run base modeld decoder forward. - hidden_states = self.decoder( - hidden_states=decoder_input, - attention_mask=attention_mask, - inference_params=inference_params, - rotary_pos_emb=rotary_pos_emb, - packed_seq_params=packed_seq_params, - **(extra_block_kwargs or {}), - ) - - return hidden_states - - def _mtp_forward( - self, - index, - input_ids, - hidden_states, - position_ids, - attention_mask, - rotary_pos_emb, - output_weight, - inference_params: InferenceParams = None, - packed_seq_params: PackedSeqParams = None, - extra_block_kwargs: dict | None = None, - ): - mtp_embeddings = self.embedding( - input_ids=input_ids, - position_ids=position_ids, - ) - hidden_states = self.mtp[index]( - mtp_embeddings, - hidden_states, - attention_mask, # TODO (chenhany): this may needs some fix - rotary_pos_emb=rotary_pos_emb, - inference_params=inference_params, - packed_seq_params=packed_seq_params, - **(extra_block_kwargs or {}), - ) - mtp_logits, _ = self.output_layer(hidden_states, weight=output_weight) - - return hidden_states, mtp_logits - - def forward( - self, - input_ids: torch.Tensor, - position_ids: torch.Tensor = None, - attention_mask: torch.Tensor = None, - decoder_input: torch.Tensor = None, - labels: torch.Tensor = None, - inference_params: InferenceParams = None, - packed_seq_params: PackedSeqParams = None, - extra_block_kwargs: dict | None = None, - **kwargs, - ) -> torch.Tensor: - if position_ids is None or attention_mask is None: - attention_mask, position_ids = get_default_attention_mask_and_position_ids(input_ids) - - hidden_states = self._base_model_forward( - input_ids, - position_ids, - attention_mask, - decoder_input, - inference_params, - packed_seq_params, - extra_block_kwargs, - ) - - if not self.post_process: - return hidden_states - - output_weight = None - if self.share_embeddings_and_output_weights: - output_weight = self.shared_embedding_or_output_weight() - logits_sbh, _ = self.output_layer(hidden_states, weight=output_weight) - - if inference_params is None or self.calibration_mode: - draft_logits = [] - for i in range(self.mtp_num_module): - mtp_ids = torch.cat( - ( - input_ids[:, 1 + i :], - torch.zeros( - input_ids.shape[0], - 1 + i, - dtype=input_ids.dtype, - device=input_ids.device, - ), - ), - dim=-1, - ) - padding_zeros = torch.zeros( - position_ids.shape[0], - 1 + i, - dtype=position_ids.dtype, - device=position_ids.device, - ) - mtp_position_ids = torch.cat((position_ids[:, 1 + i :], padding_zeros), dim=1) - - mtp_attention_mask = attention_mask.clone().detach() - mtp_attention_mask[:, :, : -(1 + i), : -(1 + i)] = attention_mask[ - :, :, 1 + i :, 1 + i : - ] - mtp_attention_mask[:, :, -(1 + i) :, :] = True - mtp_attention_mask[:, :, :, -(1 + i) :] = True - - # For MLA, rotary_pos_emb is computed per attention. - rotary_pos_emb = ( - None - if self.config.multi_latent_attention - else self.rotary_pos_emb(mtp_ids.shape[-1]) - ) - - hidden_states, mtp_logits = self._mtp_forward( - i, - mtp_ids, - hidden_states, - mtp_position_ids, - mtp_attention_mask, - rotary_pos_emb, - output_weight, - inference_params, - packed_seq_params, - extra_block_kwargs, - ) - draft_logits.append(mtp_logits) - - # If labels are not provided, return the original logits. We only return after - # all mtp weights have been exercised for quantization calibration purpose. - if labels is None: - return logits_sbh.transpose(0, 1).contiguous() - - # Base model loss - # If mtp_freeze_base_model is set to True, - # the base model is frozen . - loss = self.compute_language_model_loss(labels, logits_sbh) - - for i, mtp_logits in enumerate(draft_logits): - # Compute lm loss (classification loss) or KLDivergence - if self.mtp_self_logit_distillation: - mtp_loss = self.kld( - mtp_logits[: -(1 + i), :, :], - logits_sbh[1 + i :, :, :], - ) - else: - mtp_loss = self.compute_language_model_loss( - labels[:, 1 + i :], mtp_logits[: -(1 + i), :, :] - ) - - loss[:, 1 + i] += mtp_loss - - acc = [] - if self.mtp_report_acc: - with torch.no_grad(): - gathered_logits = gather_from_tensor_model_parallel_region(mtp_logits) - mtp_top1 = gathered_logits.transpose(0, 1).argmax(dim=-1) - mtp_top1 = mtp_top1[:, : -(1 + i)] - top1_p = torch.eq(labels[:, 1 + i :], mtp_top1).sum() / mtp_top1.numel() - acc.append(top1_p) - - if get_tensor_model_parallel_rank() == 0: - print(f"MTP_{i} Training Accuracy: {acc}") - - return loss - - def sharded_state_dict( - self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None - ) -> ShardedStateDict: - """Override the shared_state_dict to take care mtp.""" - assert not sharded_offsets, "Unexpected sharded offsets" - - sharded_state_dict = GPTModel.sharded_state_dict(self, prefix, sharded_offsets, metadata) - - if not hasattr(self, "mtp") or self.mtp is None: - return sharded_state_dict - - # This is a remedy for nn.ModuleList. GPTModel.sharded_state_dict() is calling into - # MegatronModule.sharded_state_dict() which requires all children to implement - # sharded_state_dict(). mtp is an nn.ModuleList which only has state_dict() - # implemented. As a result, all the submodules will not be sharded. - # - # The remedy is to pop all mtp* out and call the EagleModule sharded_state_dict() - # again to populate the correct sharded_staet_dict. - extra_keys = [] - for key in sharded_state_dict: - if "mtp" in key: - extra_keys += [key] - for key in extra_keys: - sharded_state_dict.pop(key, None) - - layer_prefix = f"{prefix}mtp." - for i, layer in enumerate(self.mtp): - layer_sharded_state_dict = layer.sharded_state_dict(f"{layer_prefix}{i}.", [], metadata) - sharded_state_dict.update(layer_sharded_state_dict) - return sharded_state_dict - - def pseudo_speculative_generate( - self, - input_ids: torch.Tensor, - steps: int = 1, - ): - """Pseudo generate of the MTP GPTModel. - - Returns: - base_token (torch.Tensor): token from base model - draft_tokens (torch.Tensor): draft tokens from MTP - """ - padded_input_ids, seq_len = right_padding(input_ids) - - attention_mask, position_ids = get_default_attention_mask_and_position_ids(padded_input_ids) - - hidden_states, logits_sbh, output_weight = self._base_model_forward( - padded_input_ids, - position_ids, - attention_mask, - ) - - # Removing the padding - logits_sbh = logits_sbh[:seq_len, :, :] - if self.config.sequence_parallel: - hidden_states = gather_from_sequence_parallel_region(hidden_states) - hidden_states = hidden_states[:seq_len, :, :] - - base_token = ( - gather_from_tensor_model_parallel_region(logits_sbh)[-1:, :, :] - .argmax(dim=-1) - .transpose(0, 1) - ) - mtp_ids = torch.cat((input_ids[:, 1:], base_token), dim=-1) - - draft_tokens = [] - for i in range(self.mtp_num_module): - padded_mtp_ids, seq_len, padded_hidden_states = right_padding(mtp_ids, hidden_states) - if self.config.sequence_parallel: - padded_hidden_states = scatter_to_sequence_parallel_region(padded_hidden_states) - mtp_attention_mask, mtp_position_ids = get_default_attention_mask_and_position_ids( - padded_mtp_ids - ) - if self.config.multi_latent_attention: - # For MLA, rotary_pos_emb is computed per attention. - rotary_pos_emb = None - else: - rotary_pos_emb = self.rotary_pos_emb(padded_mtp_ids.shape[-1]) - - mtp_hidden_states, mtp_logits = self._mtp_forward( - i, - padded_mtp_ids, - padded_hidden_states, - mtp_position_ids, - mtp_attention_mask, - rotary_pos_emb, - output_weight, - ) - - mtp_logits = mtp_logits[:seq_len, :, :] - if self.config.sequence_parallel: - mtp_hidden_states = gather_from_sequence_parallel_region(mtp_hidden_states) - mtp_hidden_states = mtp_hidden_states[:seq_len, :, :] - - draft_token = ( - gather_from_tensor_model_parallel_region(mtp_logits)[-1:, :, :] - .argmax(dim=-1) - .transpose(0, 1) - ) - draft_tokens.append(draft_token) - - mtp_ids = torch.cat((mtp_ids, draft_token), dim=-1) - hidden_states = torch.cat((hidden_states, mtp_hidden_states[-1:, :, :]), dim=0) - - draft_tokens = torch.cat(draft_tokens, dim=-1) - - return base_token, draft_tokens - - class MegatronARValidation(AcceptanceRateValidation): """This is the subclass for megatron model AR validation.""" diff --git a/modelopt/torch/speculative/plugins/megatron_medusa.py b/modelopt/torch/speculative/plugins/megatron_medusa.py new file mode 100644 index 000000000..21f8b51ab --- /dev/null +++ b/modelopt/torch/speculative/plugins/megatron_medusa.py @@ -0,0 +1,312 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Plugin to add Medusa support for Megatron-Core GPT model.""" + +import warnings + +import torch +import torch.nn.functional as F +from megatron.core import tensor_parallel +from megatron.core.dist_checkpointing.mapping import ShardedStateDict +from megatron.core.models.gpt import GPTModel +from megatron.core.parallel_state import get_tensor_model_parallel_rank +from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region +from megatron.core.transformer.module import MegatronModule + +from ..medusa.conversion import MedusaDMRegistry +from ..medusa.medusa_model import MedusaModel + + +class MedusaLayer(MegatronModule): + """MedusaLayer impl following TensorRT-LLM's model definition. + + Medusa layer consists of a column parallel linear following a silu. + """ + + def __init__(self, config): + """Constructor. + + Args: + config: MCore transformer config + """ + super().__init__(config=config) + + device = ( + torch.device("cpu") if config.use_cpu_initialization else torch.cuda.current_device() + ) + + self.activation_func = F.silu + + self.linear = torch.nn.Linear( + config.hidden_size, + config.hidden_size, + dtype=config.params_dtype, + device=device, + ) + + def forward(self, x): + """Forward function.""" + y = self.linear(x) + return x + self.activation_func(y), None + + +class MedusaHead(MegatronModule): + """MedusaHead impl following TensorRT-LLM's model definition. + + https://github.com/NVIDIA/TensorRT-LLM/blob/main/tensorrt_llm/models/medusa/model.py + Medusa head consists of several MedusaLayers and an lm_head. + """ + + def __init__(self, config, vocab_size: int, num_layers: int = 1, parallel_output: bool = True): + """Constructor. + + Args: + config: MCore transformer config + vocab_size: vocabulary size + num_layers: number of Medusa layers + parallel_output: if False, then all_gather the logits + """ + super().__init__(config=config) + + self.medusa_layers = torch.nn.ModuleList([MedusaLayer(config) for _ in range(num_layers)]) + + self.lm_head = tensor_parallel.ColumnParallelLinear( + config.hidden_size, + vocab_size, + config=config, + init_method=config.init_method, + bias=False, + skip_bias_add=False, + gather_output=not parallel_output, + skip_weight_param_allocation=False, + ) + + def load_state_dict_post_hook(module, incompatible_keys): + incompatible_keys.missing_keys.clear() + incompatible_keys.unexpected_keys.clear() + + self.register_load_state_dict_post_hook(load_state_dict_post_hook) + + def forward(self, x): + """Forward function.""" + for layer in self.medusa_layers: + x, _ = layer(x) + return self.lm_head(x) + + def sharded_state_dict( + self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None + ) -> ShardedStateDict: + """Return MCore sharded_state_dict.""" + assert not sharded_offsets, "Unexpected sharded offsets" + sharded_state_dict = {} + layer_prefix = f"{prefix}medusa_layers." + for i, layer in enumerate(self.medusa_layers): + state_dict_prefix = f"{layer_prefix}{i}." + sharded_pp_offset = [] + layer_sharded_state_dict = layer.sharded_state_dict( + state_dict_prefix, sharded_pp_offset, metadata + ) + sharded_state_dict.update(layer_sharded_state_dict) + sharded_state_dict.update( + self.lm_head.sharded_state_dict(f"{prefix}lm_head.", sharded_offsets, metadata) + ) + return sharded_state_dict + + +@MedusaDMRegistry.register({GPTModel: "megatron.core.models.gpt.GPTModel"}) +class _DynamicMedusaGPTModel(MedusaModel): + """A ``megatron.core.models.gpt.GPTModel`` model with dynamic hyperparams.""" + + def _setup(self): + super()._setup() + self._register_temp_attribute("medusa_report_acc", True) + self._register_temp_attribute("medusa_freeze_base_model", True) + self._register_temp_attribute("calibration_mode", False) + + def modify( + self, + medusa_num_heads=0, + medusa_num_layers=0, + medusa_freeze_base_model=True, + medusa_report_acc=True, + ): + """Constructor. + + Args: + config: MedusaConfig that specifies the medusa head configuration as well as + weights of base model and medusa head. + """ + if self.config.pipeline_model_parallel_size > 1: + warnings.warn( + "Pipeline parallelism detected! _DynamicMedusaGPTModel only supports " + "pipeline parallelism during TensorRT-LLM checkpoint export." + ) + super().modify(medusa_num_heads=medusa_num_heads, medusa_num_layers=medusa_num_layers) + + self.medusa_report_acc = medusa_report_acc + self.medusa_freeze_base_model = medusa_freeze_base_model + + # Freeze all parameters + if self.medusa_freeze_base_model: + for name, param in self.named_parameters(): + param.requires_grad = False + + if self.post_process: + self.medusa_heads = torch.nn.ModuleList( + [ + MedusaHead(self.config, self.vocab_size, num_layers=self.medusa_num_layers) + for _ in range(self.medusa_num_heads) + ] + ) + + def _base_model_forward(self, *args, labels: torch.Tensor = None, **kwargs): + if self.post_process: + # Set the post_process to False such that the forward will return the hidden_state. + self.post_process = False + # Calling parent's forward to get hidden_states + hidden_states = GPTModel.forward(self, *args, labels=labels, **kwargs) + # Reset the post_process to True + self.post_process = True + else: + hidden_states = GPTModel.forward(self, *args, labels=None, **kwargs) + + return hidden_states + + def _medusa_forward(self, hidden_states): + draft_logits = [] + # Medusa heads forward. We want to run through all the heads just to make sure all modules + # are exercised during calibration. + for i, head in enumerate(self.medusa_heads): + new_logits, _ = head(hidden_states) + + draft_logits.append(new_logits) + + return draft_logits + + def forward(self, *args, labels: torch.Tensor = None, **kwargs): + """Forward pass of the Medusa GPTModel. + + Returns: + torch.Tensor: If labels are provided, then return lm_loss of all heads. Otherwise, + return the original logits. + """ + hidden_states = self._base_model_forward(*args, labels=labels, **kwargs) + + if not self.post_process: + return hidden_states + + output_weight = None + if self.share_embeddings_and_output_weights: + output_weight = self.shared_embedding_or_output_weight() + # Original output logits + logits, _ = self.output_layer(hidden_states, weight=output_weight) + + draft_logits = self._medusa_forward(hidden_states) + + if self.medusa_report_acc and labels is not None: + acc = [] + for i, _ in enumerate(self.medusa_heads): + gathered_logits = gather_from_tensor_model_parallel_region(draft_logits[i]) + medusa_top1 = gathered_logits.transpose(0, 1).argmax(dim=-1)[:, : -(1 + i)] + medusa_labels = labels[:, 1 + i :] + top1_p = torch.eq(medusa_labels, medusa_top1).sum() / medusa_top1.numel() + acc.append(top1_p) + + if get_tensor_model_parallel_rank() == 0: + print(f"Medusa Training Accuracy: {acc}") + + # Return the original logits untouched. + if labels is None: + # [s b h] => [b s h] + return logits.transpose(0, 1).contiguous() + + # Base model loss + # If medusa_freeze_base_model is set to True, + # the base model is frozen . + loss = self.compute_language_model_loss(labels, logits) + # Medusa loss + for i, _ in enumerate(self.medusa_heads): + medusa_labels = labels[:, 1 + i :] + medusa_loss = self.compute_language_model_loss( + medusa_labels, draft_logits[i][: -(1 + i), :] + ) + loss[:, 1 + i :] += medusa_loss + + return loss + + def sharded_state_dict( + self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None + ) -> ShardedStateDict: + """Override the shared_state_dict to take care medusa_heads.""" + assert not sharded_offsets, "Unexpected sharded offsets" + + sharded_state_dict = GPTModel.sharded_state_dict(self, prefix, sharded_offsets, metadata) + + if not hasattr(self, "medusa_heads") or self.medusa_heads is None: + return sharded_state_dict + + # This is a remedy for nn.ModuleList. GPTModel.sharded_state_dict() is calling into + # MegatronModule.sharded_state_dict() which requires all children to implement + # sharded_state_dict(). medusa_heads is an nn.ModuleList which only has state_dict() + # implemented. As a result, all the submodules will not be sharded. + # + # The remedy is to pop all medusa_heads* out and call the MedusaHead sharded_state_dict() + # again to populate the correct sharded_staet_dict. + extra_keys = [] + for key in sharded_state_dict: + if "medusa_heads" in key: + extra_keys += [key] + for key in extra_keys: + sharded_state_dict.pop(key, None) + + layer_prefix = f"{prefix}medusa_heads." + for i, layer in enumerate(self.medusa_heads): + layer_sharded_state_dict = layer.sharded_state_dict(f"{layer_prefix}{i}.", [], metadata) + sharded_state_dict.update(layer_sharded_state_dict) + return sharded_state_dict + + def pseudo_speculative_generate(self, *args, steps=1, **kwargs): + """Pseudo generate of the Medusa GPTModel. + + Returns: + base_token (torch.Tensor): token from base model + draft_tokens (torch.Tensor): draft tokens from medusa heads + """ + hidden_states = self._base_model_forward(*args, labels=None, **kwargs) + + if not self.post_process: + return hidden_states + + output_weight = None + if self.share_embeddings_and_output_weights: + output_weight = self.shared_embedding_or_output_weight() + # Original output logits + logits, _ = self.output_layer(hidden_states, weight=output_weight) + + draft_logits = self._medusa_forward(hidden_states) + + logits = gather_from_tensor_model_parallel_region(logits.transpose(0, 1).contiguous()) + draft_logits = [ + gather_from_tensor_model_parallel_region(logit.transpose(0, 1).contiguous()) + for logit in draft_logits + ] + + # [b, s] + base_token = logits[:, -1:].argmax(dim=-1) + draft_tokens = [logit[:, -1:].argmax(dim=-1) for logit in draft_logits] + draft_tokens = torch.cat(draft_tokens, dim=-1) + + return base_token, draft_tokens diff --git a/modelopt/torch/speculative/plugins/transformers.py b/modelopt/torch/speculative/plugins/transformers.py index ad93d2d45..11440f5aa 100644 --- a/modelopt/torch/speculative/plugins/transformers.py +++ b/modelopt/torch/speculative/plugins/transformers.py @@ -294,6 +294,7 @@ class HFEagleModel(EagleModel): eagle_disable_moe, # Not used in HFEagleModel draft_vocab_size, use_mtp_layernorm, + parallel_draft_step=1, ffn_hidden_size=0, ): """Constructor. @@ -311,6 +312,7 @@ class HFEagleModel(EagleModel): eagle_disable_moe=eagle_disable_moe, draft_vocab_size=draft_vocab_size, use_mtp_layernorm=use_mtp_layernorm, + parallel_draft_step=parallel_draft_step, ) self.config.eagle = { diff --git a/modelopt/torch/speculative/utils.py b/modelopt/torch/speculative/utils.py index 233bbf695..81b9e5bc8 100644 --- a/modelopt/torch/speculative/utils.py +++ b/modelopt/torch/speculative/utils.py @@ -16,7 +16,8 @@ """Utils for speculative decoding.""" import copy -from collections import Counter +import warnings +from collections import Counter, defaultdict, deque import torch import torch.distributed @@ -61,22 +62,98 @@ def get_default_attention_mask_and_position_ids(input_ids: torch.Tensor): return attention_mask, position_ids -def tree_decode(draft_logits: list[torch.Tensor], tree: list[list[int]]): - """Decode tokens using the tree. +class TreeNode: + """A node in the speculative decoding tree structure. - Args: - draft_logits: a list of logits. Each logit represent a future position. - tree: a tree for decoding. Each sublist is a branch from root where the number - represents the topk index. + Each node represents a token position in the sequence and maintains a dictionary of child nodes, """ - draft_tokens = [] - for seq in tree: - tokens = [] - for i, index in enumerate(seq): - token = draft_logits[i][:, -1].topk(index + 1, dim=-1).indices[:, -1:] - tokens.append(token) - draft_tokens.append(torch.cat(tokens, dim=-1)) - return draft_tokens + + def __init__(self, value: int, children: dict | None = None): + """Initialize a TreeNode. + + Args: + value (int): the value of the node + children (dict): a dictionary of children nodes + """ + self.value = value + self.children = children if children is not None else {} + + +class Tree: + """A tree structure for speculative decoding that defines valid token prediction paths. + + This class implements a tree-based structure used in speculative decoding to represent + multiple possible token prediction paths. The tree is constructed from a list of paths, + where each path is a sequence of token positions. + + """ + + def __init__(self, tree_paths: list[list[int]]): + """Initialize a Tree. + + Args: + tree_paths (list[list[int]]): a list of tree paths + """ + self.total_nodes = 1 + self.root = TreeNode(0) + self.num_children = defaultdict(int) + self.max_depth = 0 + self.create_tree(tree_paths) + self.create_attention_mask() + + def create_tree(self, tree_paths): + """Create the tree structure from the list of tree paths. + + This function builds the tree by iterating through each path in the tree_paths list. + For each path, it traverses the tree, creating nodes and updating the number of children + at each level. + """ + tree_paths.sort() + self.num_children[0] = 1 + for node_path in tree_paths: + parent_node = self.root + for i, node in enumerate(node_path): + # if node is not a child of parent_node, add it + if node not in parent_node.children: + if i != len(node_path) - 1: + raise ValueError( + f"Incomplete tree path found at {node_path}, {i}th (non-leaf) node doesn't exist" + ) + # value of the node is position id + child_node = TreeNode(node) + parent_node.children[child_node.value] = child_node + # keep track of the number of children of per level + self.num_children[i + 1] += 1 + parent_node = parent_node.children[node] + + self.total_nodes += 1 + # update max depth + self.max_depth = max(self.max_depth, len(node_path)) + + def create_attention_mask(self): + """Create the attention mask for the tree. + + This function constructs the attention mask for the tree based on the tree structure. + It ensures that each token can only attend to its valid predecessors according to the tree. + """ + queue = deque([[node, 0] for node in self.root.children.values()]) + self.attention_mask = torch.full( + (self.total_nodes, self.total_nodes), True, device=torch.cuda.current_device() + ) + # Base token (in the first column) is attended by all draft tokens + self.attention_mask[:, 0] = False + cur_idx = 1 + while queue: + # iterate over all nodes at current level and update attention mask + for _ in range(len(queue)): + node, node_idx = queue.popleft() + self.attention_mask[cur_idx, : node_idx + 1] = self.attention_mask[ + node_idx, : node_idx + 1 + ] + self.attention_mask[cur_idx, cur_idx] = False + for child in node.children.values(): + queue.append([child, cur_idx]) + cur_idx += 1 class ResBlock(nn.Module): @@ -158,15 +235,45 @@ class AcceptanceRateValidation: osl: output sequence length """ - def check_draft(self, ground_truth, input_ids, draft_tokens, tree=None): + def check_draft(self, ground_truth, input_ids, draft_tokens): """This function checks if the draft tokens should be accepted (same as ground truth). - If tree is None, it is eager mode. + Args: + ground_truth: the ground truth token ids + input_ids: the input token ids + draft_tokens: the draft tokens + + Returns: + input_ids: the updated input token ids """ if draft_tokens is None: return input_ids - if tree is None: + if isinstance(draft_tokens, TreeNode): + # Initialize tracking variables + token_matched = False # Flag to track if current token matches ground truth + # Iterate through each step/level in the tree + while draft_tokens.children: + # Check each candidate token at current level + for child in draft_tokens.children.values(): + # Check if draft token matches ground truth token + if child.value == ground_truth[:, input_ids.shape[1]]: + # Accept matching token and update sequence + input_id = child.value.unsqueeze(0) + input_ids = torch.cat((input_ids, input_id), dim=-1) + # Update position for next level traversal + draft_tokens = child + token_matched = True + break + else: + token_matched = False + + # Stop if either: + # 1. No match found at current level + # 2. We've reached the end of ground truth sequence + if (not token_matched) or (input_ids.shape[1] == ground_truth.shape[1]): + break + else: # eager mode for i in range(draft_tokens.shape[-1]): input_id = draft_tokens[:, i : i + 1] @@ -176,25 +283,30 @@ class AcceptanceRateValidation: break else: break - else: - # tree decoding - pass return input_ids - def check_data_consistancy_across_ranks(self, data, group=None): + def check_data_consistancy_across_ranks(self, data, group=None, fail_when_mismatch=True): """This function checks the data consistancy across all ranks in the group. Use rank 0 data as the golden set to broadcast to all ranks. Each rank will then compare to this data and through error if different. """ + if data is None: + return golden_set = copy.deepcopy(data) - torch.distributed.broadcast(data, src=0, group=group) + torch.distributed.broadcast(golden_set, src=0, group=group) if not torch.equal(data, golden_set): - raise ValueError( - "Data diverges across ranks. For Megatron, 'moe-token-dispatcher-type'" - "should set to 'alltoall'." - ) + if fail_when_mismatch: + raise ValueError( + "Data diverges across ranks. For Megatron, 'moe-token-dispatcher-type'" + "should set to 'alltoall'." + ) + else: + warnings.warn( + "Data diverges across ranks. Forcing all ranks' data equal to rank 0." + ) + return golden_set def validate( self, @@ -202,8 +314,8 @@ class AcceptanceRateValidation: prompt=None, input_ids=None, ground_truth=None, - tree=None, steps=1, + tree_paths=None, ): """This function validate the AR of the model given the input sequence.""" if input_ids is None: @@ -213,18 +325,33 @@ class AcceptanceRateValidation: if ground_truth is None: ground_truth = self.get_ground_truth(input_ids, osl) - self.check_data_consistancy_across_ranks(ground_truth) + ground_truth = self.check_data_consistancy_across_ranks(ground_truth) cnt = 0 draft_tokens = None + if tree_paths: + tree = Tree(tree_paths) + while input_ids.shape[1] < ground_truth.shape[1]: cnt += 1 - input_ids = self.check_draft(ground_truth, input_ids, draft_tokens, tree) + input_ids = self.check_draft(ground_truth, input_ids, draft_tokens) if input_ids.shape[1] == ground_truth.shape[1]: break - input_id, draft_tokens = self.model.pseudo_speculative_generate(input_ids, steps=steps) - self.check_data_consistancy_across_ranks(input_id) - self.check_data_consistancy_across_ranks(draft_tokens) + + if tree_paths: + input_id, draft_tokens, pred_tokens = self.model.tree_decode(input_ids, tree=tree) + pred_tokens = self.check_data_consistancy_across_ranks( + pred_tokens, fail_when_mismatch=False + ) + else: + input_id, draft_tokens = self.model.pseudo_speculative_generate( + input_ids, steps=steps + ) + draft_tokens = self.check_data_consistancy_across_ranks( + draft_tokens, fail_when_mismatch=False + ) + + input_id = self.check_data_consistancy_across_ranks(input_id) input_ids = torch.cat((input_ids, input_id), dim=-1) ar = (ground_truth.shape[1] - isl) / cnt diff --git a/modelopt/torch/trace/plugins/megatron.py b/modelopt/torch/trace/plugins/megatron.py index ec9d51ad8..b385e5426 100644 --- a/modelopt/torch/trace/plugins/megatron.py +++ b/modelopt/torch/trace/plugins/megatron.py @@ -19,11 +19,18 @@ from megatron.core.models.gpt import GPTModel from ..symbols import Symbol, SymInfo, SymMap +try: + from megatron.core.models.mamba import MambaModel + + HAS_MAMBA = True +except ImportError: + HAS_MAMBA = False + # NOTE: No need to register symbols for VocabParallelEmbedding, SelfAttention, MLP, LayerNorm, Row/Col Parallel Linear, -# etc. as they are not traced and manually handled in the _DynamicGPTModel class -@SymMap.register(GPTModel) -def get_megatron_gpt_model_sym_info(mod: GPTModel) -> SymInfo: - """Get symbol information for ``GPTModel`` layers.""" +# etc. as they are not traced and manually handled in the _DynamicMCoreLanguageModel class +@SymMap.register([GPTModel] + ([MambaModel] if HAS_MAMBA else [])) +def get_megatron_language_model_sym_info(mod) -> SymInfo: + """Get symbol information for ``GPTModel`` and ``MambaModel`` layers.""" hidden_size = Symbol(is_searchable=True) return SymInfo(is_shape_preserving=True, hidden_size=hidden_size) diff --git a/modelopt/torch/utils/dataset_utils.py b/modelopt/torch/utils/dataset_utils.py index 350331b85..b775d9c5a 100644 --- a/modelopt/torch/utils/dataset_utils.py +++ b/modelopt/torch/utils/dataset_utils.py @@ -213,6 +213,7 @@ def get_max_batch_size( max_sample_length: int = 512, sample_memory_usage_ratio: float = 1.0, sample_input_single_batch: torch.Tensor = None, + enable_grad: bool = False, ): """Get the maximum batch size that can be used for the model.""" @@ -238,7 +239,7 @@ def get_max_batch_size( ) # Calculate single batch inference with dummy input. - with torch.no_grad(): + with torch.set_grad_enabled(enable_grad): infer_method(sample_input_single_batch) free_mem_after, max_allocated_after = _get_free_gpu_mem() @@ -267,7 +268,7 @@ def get_max_batch_size( # For some models on multi GPU, we observe the memory per batch is not a constant. # So we just test the target batch size and make sure we do not go OOM. while target_data_batch > 1: - with torch.no_grad(): + with torch.set_grad_enabled(enable_grad): try: infer_method(target_input) break diff --git a/tests/examples/diffusers/conftest.py b/modelopt/torch/utils/plugins/__init__.py similarity index 51% rename from tests/examples/diffusers/conftest.py rename to modelopt/torch/utils/plugins/__init__.py index 4974e3f3e..dbcc48663 100644 --- a/tests/examples/diffusers/conftest.py +++ b/modelopt/torch/utils/plugins/__init__.py @@ -13,34 +13,12 @@ # See the License for the specific language governing permissions and # limitations under the License. -import pytest +"""Handles utility plugins for third-party modules.""" +from modelopt.torch.utils import import_plugin -@pytest.fixture(scope="session") -def int8_args(): - # fmt: off - return [ - "--format", "int8", - "--calib-size", "8", - "--collect-method", "min-mean", - "--percentile", "1.0", - "--alpha", "0.8", - "--quant-level", "3.0", - "--n-steps", "20", - "--batch-size", "2", - "--quant-algo", "smoothquant", - ] - # fmt: on +with import_plugin("megatron_generate"): + from .megatron_generate import * - -@pytest.fixture(scope="session") -def fp8_args(): - # fmt: off - return [ - "--format", "fp8", - "--calib-size", "8", - "--quant-level", "3.0", - "--n-steps", "20", - "--batch-size", "2", - ] - # fmt: on +with import_plugin("megatron_mmlu"): + from .megatron_mmlu import * diff --git a/modelopt/torch/utils/plugins/megatron_generate.py b/modelopt/torch/utils/plugins/megatron_generate.py new file mode 100644 index 000000000..893fb993d --- /dev/null +++ b/modelopt/torch/utils/plugins/megatron_generate.py @@ -0,0 +1,207 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""A simple generate Megatron (V)LM models.""" + +import torch +from megatron.core import mpu +from megatron.core.inference.communication_utils import broadcast_from_last_pipeline_stage +from megatron.core.inference.contexts import StaticInferenceContext +from megatron.core.pipeline_parallel import get_forward_backward_func +from megatron.core.timers import Timer +from megatron.core.transformer import MegatronModule +from tqdm import tqdm + + +def get_current_memory_info(): + """Get current memory usage.""" + remaining_mem, total_mem = torch.cuda.mem_get_info() + info = "rank {:3}/{:3} memory remaining {:03}% ({:d}/{:d} MB) ".format( + torch.distributed.get_rank(), + torch.distributed.get_world_size(), + int(remaining_mem * 100 / total_mem), + remaining_mem // 1048576, + total_mem // 1048576, + ) + return info + + +def megatron_generate( + model: MegatronModule, + input_ids: torch.LongTensor, + pixel_values: torch.FloatTensor | None = None, + image_grid_thw: torch.LongTensor | None = None, + image_sizes: torch.LongTensor | None = None, + osl: int = 32, + eos_token_id: list[int] = [], + enable_kv_cache: bool = True, + disable_tqdm: bool = False, + return_dict: bool = False, +) -> torch.Tensor | dict: + """A simple generate function for Megatron Core V(LM) models. + + This function supports TP, PP, EP, and ETP. Sequence parallelism is only supported without KV-cache + decoding (automatically turned off if KV-cache is enabled). Context parallelism is not tested. + For MHA and GQA, both native DotProductAttention and TEDotProductAttention are supported. For MLA, + only TEDotProductAttention is supported. + + When PP>1, all input args must be provided by all PP ranks. Similarly, outputs are broadcasted to + all PP ranks (from the last pipeline stage). + + Args: + model: The model to generate from. + input_ids: The sequence used as a prompt to generate. + pixel_values: (`torch.FloatTensor` of shape `(batch_size, num_channels, image_size, image_size)`): + The tensors corresponding to the input images. + image_grid_thw: (`torch.LongTensor` of shape `(num_images, 3)`, *optional*): + The temporal, height and width of feature shape of each image in LLM. + image_sizes: The image sizes. + osl: The maximum sequence length to generate. + eos_token_id: The end of sequence token id. + enable_kv_cache: Whether to enable KV-cache decoding. + disable_tqdm: Whether to disable the tqdm progress bar. + return_dict: Whether to return a dictionary that includes other metrics. + """ + if not isinstance(model, MegatronModule): + raise ValueError("megatron_generate only supports Megatron Core models.") + + if model.config.sequence_parallel and enable_kv_cache: + enable_kv_cache = False + print("Turing off kv-cache decoding since is not implemented for sequence parallelism!") + + model.eval() + + # Create a static inference context if KV-cache is enabled. + max_batch_size = input_ids.shape[0] + max_seq_len = input_ids.shape[-1] + osl + inference_context = ( + StaticInferenceContext(max_batch_size, max_seq_len) if enable_kv_cache else None + ) + + def _dummy_loss_func(output_tensor, non_loss_data=True): + """Need a dummy loss function.""" + return output_tensor + + def _forward_step_func(data, model): + """Forward step function.""" + batch_size = data["tokens"].shape[0] + seq_len = data["tokens"].shape[-1] + device = data["tokens"].device + + # ModelOpt transoformer_spec by default use arbitrary attention mask type; hence we need to + # compute the attention_mask for prefilling. Alternatively, if "causal" attention mask type + # is used, the attention_mask is not needed. During generation, the attn_mask_type is overridden + # to "no_mask" by SelfAttention.forward() if inference_context is provided. + if seq_len > 1: + attention_mask = ( + torch.triu(torch.ones((batch_size, seq_len, seq_len), device=device), diagonal=1) + .bool() + .view(batch_size, 1, seq_len, seq_len) + ) + else: + attention_mask = None + + # NOTE: we don't support traditional positional embedding. Only RoPE or YaRN are supported. + position_ids = None + + output_tensor = model( + data["tokens"], + position_ids, + attention_mask, + inference_context=inference_context, + runtime_gather_output=True, + ) + return output_tensor, _dummy_loss_func + + disable_tqdm = disable_tqdm or torch.distributed.get_rank() > 0 + + output_ids = torch.tensor([]) + step_pbar = tqdm(range(osl), disable=disable_tqdm, leave=False) + + time_ttft = 0 + time_remaining_outputs = 0 + timer = Timer("generate") + timer.start(barrier=True) + + for step in step_pbar: + step_pbar.set_description(get_current_memory_info()) + + if model.config.sequence_parallel: + tp = model.config.tensor_model_parallel_size + num_pad_tokens = (tp - input_ids.shape[-1] % tp) % tp + else: + num_pad_tokens = 0 + + if inference_context is not None and step > 0: + tokens = input_ids[:, -1:] + inference_context.enable_decode_mode() + elif num_pad_tokens > 0: + padding_shape = (input_ids.shape[0], num_pad_tokens) + padded_tokens = torch.full( + padding_shape, 0, dtype=input_ids.dtype, device=input_ids.device + ) + tokens = torch.cat((input_ids, padded_tokens), dim=-1) + else: + tokens = input_ids + + list_of_logits = get_forward_backward_func()( + forward_step_func=_forward_step_func, + data_iterator=[{"tokens": tokens}], + model=model, + num_microbatches=1, + seq_length=tokens.shape[-1], + micro_batch_size=max_batch_size, + decoder_seq_length=tokens.shape[-1], + forward_only=True, + collect_non_loss_data=True, + ) + + if inference_context is not None: + inference_context.sequence_len_offset += tokens.shape[-1] + + if mpu.is_pipeline_last_stage(): + eager_ids = ( + list_of_logits[0][:, -(num_pad_tokens + 1), :].argmax(dim=-1, keepdim=True).detach() + ) + else: + eager_ids = None + + eager_ids = broadcast_from_last_pipeline_stage( + [max_batch_size, 1], input_ids.dtype, eager_ids + ) + + if step > 0: + output_ids = torch.cat([output_ids, eager_ids], dim=-1) + else: + time_ttft = timer.elapsed(barrier=True) + output_ids = eager_ids + + input_ids = torch.cat([input_ids, eager_ids], dim=-1) + + if eager_ids.item() in eos_token_id: + break + + time_remaining_outputs = timer.elapsed(barrier=True) + + # print(f"time_ttft: {time_ttft}, time_remaining_outputs: {time_remaining_outputs}") + + if return_dict: + return { + "output_ids": output_ids, + "ttft": time_ttft, + "tps": time_remaining_outputs / (output_ids.shape[-1] - 1), + } + else: + return output_ids diff --git a/modelopt/torch/utils/plugins/megatron_mmlu.py b/modelopt/torch/utils/plugins/megatron_mmlu.py new file mode 100644 index 000000000..0b9614739 --- /dev/null +++ b/modelopt/torch/utils/plugins/megatron_mmlu.py @@ -0,0 +1,150 @@ +# Adapted from https://github.com/declare-lab/instruct-eval/blob/720e66f627369266ed1cfd74426666ec37e524bc/mmlu.py + +# MIT License +# +# Copyright (c) 2020 Dan Hendrycks +# Copyright (c) 2023 Deep Cognition and Language Research (DeCLaRe) Lab +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. + +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""A simple MMLU evaluation for Megatron LM models.""" + +import requests +import torch +import transformers +from datasets import load_dataset + +from .megatron_generate import megatron_generate + + +def _get_all_subjects(): + """All subjects (anatomy, ...) can be acquired from querying all subsets and splits.""" + response = requests.get( + "https://datasets-server.huggingface.co/splits?dataset=cais/mmlu", timeout=10 + ) + data = response.json() + all_subjects = set() + for split in data["splits"]: + all_subjects.add(split["config"]) + for name in ["all", "auxiliary_train"]: + all_subjects.discard(name) + return sorted(all_subjects) + + +def megatron_mmlu( + model, + tokenizer: transformers.PreTrainedTokenizer, + few_shots: int = 0, + percentage: float = 0.05, + enable_kv_cache: bool = False, +) -> None: + """Evaluate the model on MMLU. + + Args: + model: The model to evaluate. + tokenizer: The tokenizer to use. + few_shots: The number of few-shot examples to use. + percentage: The percentage of the test set to evaluate on. + enable_kv_cache: Whether to disable KV-cache. + """ + all_correct = {} + all_subjects = _get_all_subjects() + + def _format_example(example, include_answer: bool = True): + """Format an example into a multi-choices problem.""" + prompt = example["question"] + for choice, answer in zip(["A", "B", "C", "D"], example["choices"]): + prompt += f"\n{choice}. {answer}" + if include_answer: + prompt += "Answer: {}\n\n".format(example["answer"]) + else: + prompt += "\nAnswer:" + return prompt + + def _generate_prompt(test_example, dev_examples, few_shots=0): + """Generating few-shot prompts.""" + prompt = "The following are multiple choice questions (with answers) about {}.\n\n".format( + " ".join(test_example["subject"].split("_")) + ) + for i in range(few_shots): + prompt += _format_example(dev_examples[i]) + prompt += _format_example(test_example, include_answer=False) + return prompt + + if torch.distributed.get_rank() == 0: + print(f"\nMMLU ({percentage * 100}%, {few_shots}-shot) evaluation started...\n", flush=True) + print("{:48} | (ACC) | Count/Total".format("Subject"), flush=True) + print("{:48} | {:5} | {:11}".format("-" * 48, "-" * 5, "-" * 11), flush=True) + + for subject in all_subjects: + test_data = load_dataset("cais/mmlu", subject, split="test") + dev_data = load_dataset("cais/mmlu", subject, split="dev") + + correct = [] + for idx, test_example in enumerate(test_data): + if idx > percentage * len(test_data): + break + prompt = _generate_prompt(test_example, dev_data, few_shots=few_shots) + label = ["A", "B", "C", "D"][test_example["answer"]] + tokens = tokenizer(prompt, return_tensors="pt") + generated_ids = megatron_generate( + model, + tokens.input_ids.cuda(), + osl=2, + disable_tqdm=True, + enable_kv_cache=enable_kv_cache, + ) + predict = tokenizer.batch_decode(generated_ids)[0].strip() + correct += [True] if predict.startswith(label) else [False] + all_correct[subject] = correct + + if torch.distributed.get_rank() == 0: + print( + f"{subject:48} | {sum(correct) / len(correct):.3f} | {sum(correct):5}/{len(correct):5}", + flush=True, + ) + + avg_correct = [] + + for subject, correct in all_correct.items(): + avg_correct += correct + + if torch.distributed.get_rank() == 0: + print("{:48} | {:5} | {:11}".format("-" * 48, "-" * 5, "-" * 11), flush=True) + print( + "{:48} | {:.3f} | {:5}/{:5}".format( + "average", sum(avg_correct) / len(avg_correct), sum(avg_correct), len(avg_correct) + ), + flush=True, + ) diff --git a/pyproject.toml b/pyproject.toml index d3f17b6cc..86dc0ceba 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,7 +52,8 @@ extend-ignore = [ "PT011", "PT018", "RUF002", "RUF012", "SIM115", - "UP038" + "UP038", "UP032", + "SIM108", "SIM102", ] diff --git a/setup.py b/setup.py index a766894ed..a1c4dcf46 100644 --- a/setup.py +++ b/setup.py @@ -15,16 +15,14 @@ """The package setup script for modelopt customizing certain aspects of the installation process.""" -import os -import platform - import setuptools +from setuptools_scm import get_version # Package configuration ############################################################################ name = "nvidia-modelopt" -version = os.environ.get( - "SETUPTOOLS_SCM_PRETEND_VERSION", "0.33.0" if platform.system() == "Linux" else "0.27.0" -) +# TODO: Set version to static stable release version when creating the release branch +# version = os.environ.get("SETUPTOOLS_SCM_PRETEND_VERSION", "X.Y.Z") +version = get_version(root=".", fallback_version="0.0.0") packages = setuptools.find_namespace_packages(include=["modelopt*"]) package_dir = {"": "."} package_data = {"modelopt": ["**/*.h", "**/*.cpp", "**/*.cu"]} @@ -33,7 +31,6 @@ setup_kwargs = {} # Required and optional dependencies ############################################################### required_deps = [ # Common - f"nvidia-modelopt-core=={version}", "ninja", # for faster building of C++ / CUDA extensions "numpy", "packaging", @@ -56,11 +53,13 @@ optional_deps = { "cppimport", "cupy-cuda12x; platform_machine != 'aarch64' and platform_system != 'Darwin'", "ml_dtypes", # for bfloat16 conversion - "onnx>=1.18.0", "onnx-graphsurgeon", + "onnx>=1.18.0", + "onnxconverter-common", "onnxruntime~=1.22.0 ; platform_machine == 'aarch64' or platform_system == 'Darwin'", "onnxruntime-gpu~=1.22.0 ; platform_machine != 'aarch64' and platform_system != 'Darwin' and platform_system != 'Windows'", # noqa: E501 "onnxruntime-gpu==1.20.0; platform_system == 'Windows'", + "onnxscript", # For test_onnx_dynamo_export unit test "onnxsim ; python_version < '3.12' and platform_machine != 'aarch64'", "polygraphy>=0.49.22", ], @@ -70,7 +69,7 @@ optional_deps = { "diffusers>=0.32.2", "huggingface_hub>=0.24.0", "peft>=0.12.0", - "transformers>=4.48,<4.54", # Version match done in modelopt/torch/__init__.py as well + "transformers>=4.48,<5.0", # Version match done in modelopt/torch/__init__.py as well ], # linter tools "dev-lint": [ @@ -82,13 +81,12 @@ optional_deps = { # testing "dev-test": [ "coverage", - "onnxscript", # For test_onnx_dynamo_export unit test "pytest", "pytest-cov", "pytest-timeout", "timm", - "tox", - "tox-current-env>=0.0.12", # Incompatible with tox==4.18.0 + "tox>4.18", + "tox-current-env>=0.0.12", ], # docs "dev-docs": [ diff --git a/tests/_test_utils/import_helper.py b/tests/_test_utils/import_helper.py index 6d9fa860d..03d3c8f24 100644 --- a/tests/_test_utils/import_helper.py +++ b/tests/_test_utils/import_helper.py @@ -41,7 +41,7 @@ def skip_if_no_libcudnn(): pytest.skip(f"{e}!", allow_module_level=True) -def skip_if_no_megatron(apex_or_te_required: bool = False): +def skip_if_no_megatron(apex_or_te_required: bool = False, mamba_required: bool = False): try: import megatron # noqa: F401 except ImportError: @@ -50,16 +50,26 @@ def skip_if_no_megatron(apex_or_te_required: bool = False): try: import apex # noqa: F401 - HAS_APEX = True # noqa: N806 + has_apex = True except ImportError: - HAS_APEX = False # noqa: N806 + has_apex = False try: import transformer_engine # noqa: F401 - HAS_TE = True # noqa: N806 + has_te = True except ImportError: - HAS_TE = False # noqa: N806 + has_te = False - if apex_or_te_required and not HAS_APEX and not HAS_TE: + try: + import mamba_ssm # noqa: F401 + + has_mamba = True + except ImportError: + has_mamba = False + + if apex_or_te_required and not has_apex and not has_te: pytest.skip("Apex or TE required for Megatron test", allow_module_level=True) + + if mamba_required and not has_mamba: + pytest.skip("Mamba required for Megatron test", allow_module_level=True) diff --git a/tests/_test_utils/model.py b/tests/_test_utils/model.py index 41f91f702..6e2fe17f7 100644 --- a/tests/_test_utils/model.py +++ b/tests/_test_utils/model.py @@ -62,3 +62,19 @@ LLAVA_PATH = _select_path( remote_id="llava-hf/llava-1.5-7b-hf", local_id="llava-1.5-7b-hf", ) + +# Diffusers +FLUX_SCHNELL_PATH = _select_path( + remote_id="hf-internal-testing/tiny-flux-pipe", + local_id="black-forest-labs/FLUX.1-schnell", +) + +SDXL_1_0_PATH = _select_path( + remote_id="hf-internal-testing/tiny-sdxl-pipe", + local_id="stabilityai/stable-diffusion-xl-base-1.0", +) + +SD3_PATH = _select_path( + remote_id="hf-internal-testing/tiny-sd3-pipe", + local_id="stabilityai/stable-diffusion-3-medium-diffusers", +) diff --git a/tests/_test_utils/ptq_utils.py b/tests/_test_utils/ptq_utils.py new file mode 100644 index 000000000..f943faadb --- /dev/null +++ b/tests/_test_utils/ptq_utils.py @@ -0,0 +1,105 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import importlib.metadata as metadata +import subprocess +from dataclasses import asdict, dataclass +from pathlib import Path + +import pytest +import torch + +PTQ_EXAMPLE_DIR = Path(__file__).parents[2] / "examples" / "llm_ptq" + + +@dataclass +class PTQCommand: + quant: str + export_fmt: str = "tensorrt_llm" + tasks: str = "build" + calib: int = 16 + sparsity: str | None = None + kv_cache_quant: str | None = None + trust_remote_code: bool = False + calib_batch_size: int | None = None + auto_quantize_bits: float | None = None + tp: int | None = None + pp: int | None = None + min_sm: int | None = None + min_gpu: int | None = None + + def run(self, model_path: str): + if self.min_sm and torch.cuda.get_device_capability() < ( + self.min_sm // 10, + self.min_sm % 10, + ): + pytest.skip(reason=f"Requires sm{self.min_sm} or higher") + return + + if self.min_gpu and torch.cuda.device_count() < self.min_gpu: + pytest.skip(reason=f"Requires at least {self.min_gpu} GPUs") + return + + param_dict = asdict(self) + + param_dict.pop("min_sm", None) + param_dict.pop("min_gpu", None) + + trust_remote_code = param_dict.pop("trust_remote_code", False) + + args = ["--model", model_path] + for key, value in param_dict.items(): + if value is not None: + args.append(f"--{key}") + args.append(f"{value}") + + if trust_remote_code: + args.append("--trust_remote_code") + + self.command = ["scripts/huggingface_example.sh", "--no-verbose", *args] + subprocess.run(self.command, cwd=PTQ_EXAMPLE_DIR, check=True) + + def param_str(self): + param_dict = asdict(self) + param_dict.pop("trust_remote_code", False) + return "_".join(str(value) for value in param_dict.values() if value is not None).replace( + ",", "_" + ) + + +class WithRequirements: + requirements = [] + + @pytest.fixture(scope="class", autouse=True) + def install(self): + save_deps = [] + for mod, ver in self.requirements: + try: + save_ver = metadata.version(mod) + except metadata.PackageNotFoundError: + save_ver = None + + save_deps.append((mod, save_ver)) + + spec = f"{mod}=={ver}" if ver else mod + subprocess.run(["pip", "install", spec], check=True) + + yield + + for mod, ver in save_deps: + if ver: + subprocess.run(["pip", "install", f"{mod}=={ver}"], check=True) + else: + subprocess.run(["pip", "uninstall", "--yes", mod], check=True) diff --git a/tests/_test_utils/torch_dist/plugins/megatron_common.py b/tests/_test_utils/torch_dist/plugins/megatron_common.py index 3bfea254c..95a1464d1 100644 --- a/tests/_test_utils/torch_dist/plugins/megatron_common.py +++ b/tests/_test_utils/torch_dist/plugins/megatron_common.py @@ -13,16 +13,15 @@ # See the License for the specific language governing permissions and # limitations under the License. import copy +from warnings import warn import torch import torch.nn as nn import torch.nn.functional as F from _test_utils.import_helper import skip_if_no_megatron -from packaging.version import Version skip_if_no_megatron() -from megatron.core import __version__ as mcore_version from megatron.core import dist_checkpointing from megatron.core.inference.communication_utils import broadcast_from_last_pipeline_stage from megatron.core.inference.model_inference_wrappers.gpt.gpt_inference_wrapper import ( @@ -36,6 +35,7 @@ from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_local_spec, get_gpt_layer_with_transformer_engine_spec, ) +from megatron.core.models.mamba import MambaModel from megatron.core.parallel_state import ( initialize_model_parallel, is_pipeline_first_stage, @@ -43,6 +43,8 @@ from megatron.core.parallel_state import ( ) from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.attention import SelfAttention +from megatron.core.transformer.mlp import MLP from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_config import TransformerConfig @@ -56,14 +58,25 @@ try: from megatron.core.post_training.modelopt.gpt.model_specs import get_gpt_modelopt_spec HAS_TE = True -except ImportError: +except ImportError as e: + warn(f"Transformer Engine not installed: {e}") HAS_TE = False +try: + from megatron.core.post_training.modelopt.mamba.model_specs import get_mamba_stack_modelopt_spec + from megatron.core.ssm.mamba_layer import MambaLayer + + HAS_MAMBA = True +except ImportError as e: + warn(f"Mamba not installed: {e}") + HAS_MAMBA = False + try: import apex # noqa: F401 HAS_APEX = True -except ImportError: +except ImportError as e: + warn(f"Apex not installed: {e}") HAS_APEX = False @@ -118,24 +131,29 @@ class MegatronModel(MegatronModule): def get_mcore_gpt_model( tensor_model_parallel_size: int = 1, pipeline_model_parallel_size: int = 1, + initialize_megatron: bool = False, *, num_layers: int = 2, - num_layers_in_first_pipeline_stage: int | None = None, - num_layers_in_last_pipeline_stage: int | None = None, hidden_size: int = 64, num_attention_heads: int = 8, num_query_groups: int | None = None, ffn_hidden_size: int | None = 128, - max_sequence_length: int = 32, + max_sequence_length: int = 16, vocab_size: int = 64, activation_func: str = "swiglu", normalization: str = "LayerNorm", transformer_impl: str = "modelopt" if HAS_TE else "local", + # Uneven PP + num_layers_in_first_pipeline_stage: int | None = None, + num_layers_in_last_pipeline_stage: int | None = None, ) -> GPTModel: assert activation_func in ["swiglu", "squared_relu"] assert normalization in ["LayerNorm", "RMSNorm"] assert transformer_impl in ["local", "transformer_engine", "modelopt"] - print(f"Using `transformer_impl={transformer_impl}` model spec for building GPT Model.") + print(f"Using `{transformer_impl=}` model spec for building GPT Model.") + + if initialize_megatron: + initialize_for_megatron(tensor_model_parallel_size, pipeline_model_parallel_size) def squared_relu(x): return torch.pow(F.relu(x), 2) @@ -154,18 +172,17 @@ def get_mcore_gpt_model( gated_linear_unit=(activation_func == "swiglu"), pipeline_dtype=torch.float32, add_bias_linear=False, - # Uneven PP num_layers_in_first_pipeline_stage=num_layers_in_first_pipeline_stage, num_layers_in_last_pipeline_stage=num_layers_in_last_pipeline_stage, ) if transformer_impl == "local": assert HAS_APEX, "Apex not installed" - transformer_layer_spec = get_gpt_layer_local_spec() + transformer_layer_spec = get_gpt_layer_local_spec(normalization=normalization) else: assert HAS_TE, "Transformer Engine not installed" transformer_layer_spec = ( - get_gpt_modelopt_spec(config) + get_gpt_modelopt_spec(config, remap_te_layernorm=True) if transformer_impl == "modelopt" else get_gpt_layer_with_transformer_engine_spec() ) @@ -184,26 +201,99 @@ def get_mcore_gpt_model( return model +def get_mcore_mamba_model( + tensor_model_parallel_size: int = 1, + pipeline_model_parallel_size: int = 1, + initialize_megatron: bool = False, + *, + num_layers: int = 3, + hybrid_override_pattern: str | None = None, + hidden_size: int = 64, + num_attention_heads: int = 8, + num_query_groups: int | None = None, + ffn_hidden_size: int | None = 128, + max_sequence_length: int = 4, + vocab_size: int = 64, + # Mamba-specific parameters + mamba_state_dim: int = 32, + mamba_head_dim: int = 16, + mamba_num_groups: int = 2, + # Uneven PP + num_layers_in_first_pipeline_stage: int | None = None, + num_layers_in_last_pipeline_stage: int | None = None, +) -> MambaModel: + assert HAS_MAMBA, "Mamba not installed" + + if initialize_megatron: + initialize_for_megatron(tensor_model_parallel_size, pipeline_model_parallel_size) + + config = TransformerConfig( + tensor_model_parallel_size=tensor_model_parallel_size, + pipeline_model_parallel_size=pipeline_model_parallel_size, + sequence_parallel=False, + num_layers=num_layers, + hidden_size=hidden_size, + num_attention_heads=num_attention_heads, + num_query_groups=num_query_groups, + ffn_hidden_size=ffn_hidden_size, + pipeline_dtype=torch.float32, + mamba_state_dim=mamba_state_dim, + mamba_head_dim=mamba_head_dim, + mamba_num_groups=mamba_num_groups, + num_layers_in_first_pipeline_stage=num_layers_in_first_pipeline_stage, + num_layers_in_last_pipeline_stage=num_layers_in_last_pipeline_stage, + ) + + if hybrid_override_pattern is None: + # Generate pattern by repeating "M*-" and trimming to match num_layers + # For num_layers=3, return "M*-" (Mamba -> Attention -> MLP) + # For num_layers=5, return "M*-M*" (Mamba -> Attention -> MLP -> Mamba -> Attention) + hybrid_override_pattern = ("M*-" * num_layers)[:num_layers] + else: + assert len(hybrid_override_pattern) == num_layers + print(f"Using `{hybrid_override_pattern=}` for building Mamba Model.") + + model = MambaModel( + config=config, + mamba_stack_spec=get_mamba_stack_modelopt_spec(remap_te_layernorm=True), + vocab_size=vocab_size, + max_sequence_length=max_sequence_length, + hybrid_override_pattern=hybrid_override_pattern, + pre_process=is_pipeline_first_stage(), + post_process=is_pipeline_last_stage(), + share_embeddings_and_output_weights=False, + position_embedding_type="rope", + ) + return model + + @torch.no_grad() -def run_mcore_gpt_inference( - model: GPTModel, prompt_tokens: torch.Tensor, active_hidden_size: int | None = None +def run_mcore_inference( + model: GPTModel | MambaModel, + prompt_tokens: torch.Tensor, + active_hidden_size: int | None = None, ) -> torch.Tensor: - """Run inference on a wrapped Megatron GPT model. + """Run inference on a wrapped Megatron GPT or Mamba model. Args: - model: Megatron GPT model. + model: Megatron GPT or Mamba model. prompt_tokens: Input tokens for inference. active_hidden_size: Hidden size to use for inference. If not provided, infer the hidden_size - from `model.decoder.layers[0].self_attention.linear_qkv.input_size`. NOTE: `model.config.hidden_size` may not be the same as the active hidden size for the model since for a NAS search space-converted model, the hidden size may be different until the model is exported. NOTE: If depth pruned model and some PP have 0 layers, this would not work. """ batch_size = prompt_tokens.shape[0] - active_hidden_size = ( - active_hidden_size or model.decoder.layers[0].self_attention.linear_qkv.input_size - ) + if active_hidden_size is None: + if HAS_MAMBA and isinstance(model.decoder.layers[0], MambaLayer): + active_hidden_size = model.decoder.layers[0].mixer.d_model + elif isinstance(model.decoder.layers[0].self_attention, SelfAttention): + active_hidden_size = model.decoder.layers[0].self_attention.linear_qkv.input_size + elif isinstance(model.decoder.layers[0].mlp, MLP): + active_hidden_size = model.decoder.layers[0].mlp.linear_fc1.input_size + else: + raise ValueError(f"Cannot infer hidden size from {type(model.decoder.layers[0])=}") inference_wrapper_config = InferenceWrapperConfig( hidden_size=active_hidden_size, inference_batch_times_seqlen_threshold=batch_size * model.max_sequence_length, @@ -213,13 +303,10 @@ def run_mcore_gpt_inference( ) wrapped_model = GPTInferenceWrapper(model, inference_wrapper_config) wrapped_model.prep_model_for_inference(prompt_tokens) - if Version(mcore_version) >= Version("0.11"): - inference_input = wrapped_model.prep_inference_input(prompt_tokens) - inference_input = wrapped_model.get_batch_for_context_window( - inference_input, 0, model.max_sequence_length - ) - else: - inference_input = wrapped_model.get_batch_for_context_window(0, model.max_sequence_length) + inference_input = wrapped_model.prep_inference_input(prompt_tokens) + inference_input = wrapped_model.get_batch_for_context_window( + inference_input, 0, model.max_sequence_length + ) # Note: This is returned in all TP ranks or last PP stage in PP models logits = wrapped_model.run_one_forward_step(inference_input) @@ -231,14 +318,14 @@ def run_mcore_gpt_inference( return logits # shape: (batch_size, max_sequence_length, vocab_size) -def run_mcore_gpt_inference_with_dummy_input( - model: GPTModel, batch_size: int = 2, hidden_size: int | None = None +def run_mcore_inference_with_dummy_input( + model: GPTModel | MambaModel, batch_size: int = 2, hidden_size: int | None = None ) -> torch.Tensor: - """Run inference on a wrapped Megatron GPT model.""" + """Run inference on a wrapped Megatron GPT or Mamba model.""" prompt_tokens = torch.randint( 0, model.vocab_size, (batch_size, model.max_sequence_length) ).cuda() - return run_mcore_gpt_inference(model, prompt_tokens, hidden_size) + return run_mcore_inference(model, prompt_tokens, hidden_size) def initialize_for_megatron( diff --git a/tests/_test_utils/torch_model/transformers_models.py b/tests/_test_utils/torch_model/transformers_models.py index 5c38b33eb..6c304014d 100644 --- a/tests/_test_utils/torch_model/transformers_models.py +++ b/tests/_test_utils/torch_model/transformers_models.py @@ -26,6 +26,8 @@ from transformers import ( BertForQuestionAnswering, LlamaConfig, LlamaForCausalLM, + Qwen3Config, + Qwen3ForCausalLM, T5Config, T5Model, T5Tokenizer, @@ -34,6 +36,22 @@ from transformers import ( import modelopt.torch.opt as mto +def get_tiny_qwen3(**config_kwargs) -> Qwen3ForCausalLM: + kwargs = { + "hidden_size": 32, + "intermediate_size": 32, + "num_hidden_layers": 2, + "num_attention_heads": 16, + "num_key_value_heads": 2, + "max_position_embeddings": 32, + "vocab_size": 32, + } + kwargs.update(**config_kwargs) + tiny_qwen3 = Qwen3ForCausalLM(Qwen3Config(**kwargs)) + + return tiny_qwen3 + + def get_tiny_llama(**config_kwargs) -> LlamaForCausalLM: kwargs = { "hidden_size": 32, diff --git a/tests/examples/cnn_qat/test_resnet50.py b/tests/examples/cnn_qat/test_resnet50.py new file mode 100644 index 000000000..77da56c1f --- /dev/null +++ b/tests/examples/cnn_qat/test_resnet50.py @@ -0,0 +1,77 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os + +import pytest +from _test_utils.examples.run_command import run_example_command +from _test_utils.torch_misc import minimum_gpu + +imagenet_path = os.getenv("IMAGENET_PATH") +if not imagenet_path or not os.path.isdir(imagenet_path): + pytest.skip( + "IMAGENET_PATH environment variable is not set or does not point to a valid directory", + allow_module_level=True, + ) + + +def _build_common_command(): + """Build common command arguments for CNN QAT training.""" + train_data_path = os.path.join(imagenet_path, "train") + val_data_path = os.path.join(imagenet_path, "val") + for p in (train_data_path, val_data_path): + if not os.path.isdir(p): + pytest.skip(f"Expected dataset folder '{p}' not found", allow_module_level=True) + + return [ + "--train-data-path", + train_data_path, + "--val-data-path", + val_data_path, + "--batch-size", + "64", + "--num-workers", + "8", + "--epochs", + "5", + "--lr", + "1e-4", + "--print-freq", + "50", + ] + + +def _run_qat_command(base_cmd, common_args, output_dir, example_dir="cnn_qat"): + """Helper function to run QAT command with common arguments.""" + full_command = base_cmd + common_args + ["--output-dir", str(output_dir)] + run_example_command(full_command, example_dir) + + +@minimum_gpu(1) +def test_cnn_qat_single_gpu(tmp_path): + """Test CNN QAT on single GPU.""" + common_args = _build_common_command() + base_command = ["python", "torchvision_qat.py", "--gpu", "0"] + + _run_qat_command(base_command, common_args, tmp_path) + + +@minimum_gpu(2) +def test_cnn_qat_multi_gpu(tmp_path): + """Test CNN QAT on multiple GPUs.""" + common_args = _build_common_command() + base_command = ["torchrun", "--nproc_per_node=2", "torchvision_qat.py"] + + _run_qat_command(base_command, common_args, tmp_path) diff --git a/tests/examples/diffusers/test_diffusers.py b/tests/examples/diffusers/test_diffusers.py new file mode 100644 index 000000000..4900bc508 --- /dev/null +++ b/tests/examples/diffusers/test_diffusers.py @@ -0,0 +1,159 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from pathlib import Path +from typing import NamedTuple + +import pytest +from _test_utils.examples.run_command import run_diffusers_cmd +from _test_utils.model import FLUX_SCHNELL_PATH, SD3_PATH, SDXL_1_0_PATH +from _test_utils.torch_misc import minimum_sm + + +class DiffuserModel(NamedTuple): + dtype: str + name: str + path: str + format_type: str + quant_algo: str + quant_level: str + collect_method: str + + def _run_cmd(self, script: str, *args: str) -> None: + cmd_args = [ + "python", + script, + "--model", + self.name, + "--override-model-path", + self.path, + "--model-dtype", + self.dtype, + ] + cmd_args.extend(args) + run_diffusers_cmd(cmd_args) + + def _format_args(self) -> list[str]: + return [ + "--calib-size", + "8", + "--percentile", + "1.0", + "--alpha", + "0.8", + "--n-steps", + "20", + "--batch-size", + "2", + "--format", + self.format_type, + "--collect-method", + self.collect_method, + "--quant-level", + self.quant_level, + "--quant-algo", + self.quant_algo, + ] + + def quantize(self, tmp_path: Path) -> None: + self._run_cmd( + "quantize.py", + *self._format_args(), + "--trt-high-precision-dtype", + self.dtype, + "--quantized-torch-ckpt-save-path", + str(tmp_path / f"{self.name}_{self.format_type}_{self.quant_level}.pt"), + "--onnx-dir", + str(tmp_path / f"{self.name}_{self.format_type}_{self.quant_level}_onnx"), + ) + + def restore(self, tmp_path: Path) -> None: + self._run_cmd( + "quantize.py", + *self._format_args(), + "--trt-high-precision-dtype", + self.dtype, + "--restore-from", + str(tmp_path / f"{self.name}_{self.format_type}_{self.quant_level}.pt"), + "--onnx-dir", + str(tmp_path / f"{self.name}_{self.format_type}_{self.quant_level}_onnx"), + ) + + def inference(self, tmp_path: Path) -> None: + self._run_cmd( + "diffusion_trt.py", + "--onnx-load-path", + str(tmp_path / f"{self.name}_{self.format_type}_{self.quant_level}_onnx/model.onnx"), + "--dq-only", + ) + + +@pytest.mark.parametrize( + "model", + [ + DiffuserModel( + name="flux-schnell", + path=FLUX_SCHNELL_PATH, + dtype="BFloat16", + format_type="int8", + quant_algo="smoothquant", + quant_level="3.0", + collect_method="min-mean", + ), + DiffuserModel( + name="sd3-medium", + path=SD3_PATH, + dtype="Half", + format_type="int8", + quant_algo="smoothquant", + quant_level="3.0", + collect_method="min-mean", + ), + pytest.param( + DiffuserModel( + name="sdxl-1.0", + path=SDXL_1_0_PATH, + dtype="Half", + format_type="fp8", + quant_algo="max", + quant_level="3.0", + collect_method="default", + ), + marks=minimum_sm(89), + ), + DiffuserModel( + name="sdxl-1.0", + path=SDXL_1_0_PATH, + dtype="Half", + format_type="int8", + quant_algo="smoothquant", + quant_level="3.0", + collect_method="min-mean", + ), + ], + ids=[ + "flux_schnell_bf16_int8_smoothquant_3.0_min_mean", + "sd3_medium_fp16_int8_smoothquant_3.0_min_mean", + "sdxl_1.0_fp16_fp8_max_3.0_default", + "sdxl_1.0_fp16_int8_smoothquant_3.0_min_mean", + ], +) +def test_diffusers_quantization( + model: DiffuserModel, + tmp_path: Path, +) -> None: + model.quantize(tmp_path) + model.restore(tmp_path) + model.inference(tmp_path) diff --git a/tests/examples/diffusers/test_flux_quantization.py b/tests/examples/diffusers/test_flux_quantization.py deleted file mode 100644 index 6c8208b93..000000000 --- a/tests/examples/diffusers/test_flux_quantization.py +++ /dev/null @@ -1,56 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from _test_utils.examples.run_command import run_diffusers_cmd - -# Use the tiny models for faster testing -FLUX_ARGS = [ - "--model", - "flux-schnell", - "--override-model-path", - "hf-internal-testing/tiny-flux-pipe", -] -FLUX_DTYPE = "BFloat16" - - -# fmt: off -def test_flux_int8_level3(int8_args, tmp_path): - torch_ckpt_path = tmp_path / "flux-schnell.int8.level3.pt" - onnx_dir = tmp_path / "flux-schnell.int8.level3.onnx" - - # INT8 Level 3 - run_diffusers_cmd( - [ - "python", "quantize.py", - *FLUX_ARGS, *int8_args, - "--model-dtype", FLUX_DTYPE, - "--trt-high-precision-dtype", FLUX_DTYPE, - "--batch-size", "1", - "--quantized-torch-ckpt-save-path", torch_ckpt_path, - "--onnx-dir", onnx_dir, - ], - ) - - # Inference - DQ only - run_diffusers_cmd( - [ - "python", "diffusion_trt.py", - *FLUX_ARGS, - "--model-dtype", FLUX_DTYPE, - "--onnx-load-path", onnx_dir / "model.onnx", - "--dq-only", - ], - ) diff --git a/tests/examples/diffusers/test_sd3_quantization.py b/tests/examples/diffusers/test_sd3_quantization.py deleted file mode 100644 index 32e94ad1c..000000000 --- a/tests/examples/diffusers/test_sd3_quantization.py +++ /dev/null @@ -1,65 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from _test_utils.examples.run_command import run_diffusers_cmd - -# Use the tiny models for faster testing -# NOTE: For tiny sdxl and sd3, we have to use Float dtype instead of Half from transformers v4.49.0 -SD3_ARGS = ["--model", "sd3-medium", "--override-model-path", "hf-internal-testing/tiny-sd3-pipe"] -SD3_DTYPE = "Float" - - -# fmt: off -def test_sd3_int8_level3(int8_args, tmp_path): - torch_ckpt_path = tmp_path / "sd3-medium.int8.level3.pt" - onnx_dir = tmp_path / "sd3-medium.int8.level3.onnx" - restore_onnx_dir = tmp_path / "sd3-medium.int8.level3.restore.onnx" - - # INT8 Level 3 - run_diffusers_cmd( - [ - "python", "quantize.py", - *SD3_ARGS, *int8_args, - "--model-dtype", SD3_DTYPE, - "--trt-high-precision-dtype", SD3_DTYPE, - "--batch-size", "2", - "--quantized-torch-ckpt-save-path", torch_ckpt_path, - "--onnx-dir", onnx_dir, - ], - ) - - # INT8 Level 3 Restore - run_diffusers_cmd( - [ - "python", "quantize.py", - *SD3_ARGS, *int8_args, - "--model-dtype", SD3_DTYPE, - "--trt-high-precision-dtype", SD3_DTYPE, - "--restore-from", torch_ckpt_path, - "--onnx-dir", restore_onnx_dir, - ], - ) - - # Inference - DQ only - run_diffusers_cmd( - [ - "python", "diffusion_trt.py", - *SD3_ARGS, - "--model-dtype", SD3_DTYPE, - "--onnx-load-path", onnx_dir / "model.onnx", - "--dq-only", - ], - ) diff --git a/tests/examples/diffusers/test_sdxl_quantization.py b/tests/examples/diffusers/test_sdxl_quantization.py deleted file mode 100644 index 463ef4e2f..000000000 --- a/tests/examples/diffusers/test_sdxl_quantization.py +++ /dev/null @@ -1,107 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from warnings import warn - -from _test_utils.examples.run_command import run_diffusers_cmd - -# Use the tiny models for faster testing -# NOTE: For tiny sdxl and sd3, we have to use Float dtype instead of Half from transformers v4.49.0 -SDXL_ARGS = ["--model", "sdxl-1.0", "--override-model-path", "hf-internal-testing/tiny-sdxl-pipe"] -SDXL_DTYPE = "Float" - - -# fmt: off -def test_sdxl_fp8_level3(fp8_args, tmp_path, cuda_capability): - torch_ckpt_path = tmp_path / "sdxl.fp8.level3.pt" - onnx_dir = tmp_path / "sdxl.fp8.level3.onnx" - restore_onnx_dir = tmp_path / "sdxl.fp8.level3.restore.onnx" - - # FP8 Level 3 - run_diffusers_cmd( - [ - "python", "quantize.py", - *SDXL_ARGS, *fp8_args, - "--model-dtype", SDXL_DTYPE, - "--trt-high-precision-dtype", SDXL_DTYPE, - "--quantized-torch-ckpt-save-path", torch_ckpt_path, - "--onnx-dir", onnx_dir, - ], - ) - - # FP8 Level 3 Restore - run_diffusers_cmd( - [ - "python", "quantize.py", - *SDXL_ARGS, *fp8_args, - "--model-dtype", SDXL_DTYPE, - "--trt-high-precision-dtype", SDXL_DTYPE, - "--restore-from", torch_ckpt_path, - "--onnx-dir", restore_onnx_dir, - ], - ) - - # Inference - DQ only - if cuda_capability >= (8, 9): - run_diffusers_cmd( - [ - "python", "diffusion_trt.py", - *SDXL_ARGS, - "--model-dtype", SDXL_DTYPE, - "--onnx-load-path", onnx_dir / "model.onnx", - "--dq-only", - ], - ) - else: - warn("CUDA capability >= 8.9 is required for FP8 inference!") - - -def test_sdxl_int8_level3(int8_args, tmp_path): - torch_ckpt_path = tmp_path / "sdxl.int8.level3.pt" - onnx_dir = tmp_path / "sdxl.int8.level3.onnx" - - # INT8 Level 3 - run_diffusers_cmd( - [ - "python", "quantize.py", - *SDXL_ARGS, *int8_args, - "--model-dtype", SDXL_DTYPE, - "--trt-high-precision-dtype", SDXL_DTYPE, - "--batch-size", "2", - "--quantized-torch-ckpt-save-path", torch_ckpt_path, - "--onnx-dir", onnx_dir, - ], - ) - - # Inference - run_diffusers_cmd( - [ - "python", "diffusion_trt.py", - *SDXL_ARGS, - "--model-dtype", SDXL_DTYPE, - "--restore-from", torch_ckpt_path, - ], - ) - - # Inference - DQ only - run_diffusers_cmd( - [ - "python", "diffusion_trt.py", - *SDXL_ARGS, - "--model-dtype", SDXL_DTYPE, - "--onnx-load-path", onnx_dir / "model.onnx", - "--dq-only", - ], - ) diff --git a/tests/examples/llm_ptq/test_bart.py b/tests/examples/llm_ptq/test_bart.py deleted file mode 100644 index 3e20d7aa2..000000000 --- a/tests/examples/llm_ptq/test_bart.py +++ /dev/null @@ -1,31 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import pytest -from _test_utils.examples.run_command import run_llm_ptq_command -from _test_utils.model import BART_PATH -from _test_utils.torch_misc import minimum_sm - - -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp16", "tensorrt_llm")]) -def test_bart(quant, export_fmt): - run_llm_ptq_command(model=BART_PATH, quant=quant, export_fmt=export_fmt) - - -@minimum_sm(89) -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp8", "tensorrt_llm")]) -def test_bart_sm89(quant, export_fmt): - run_llm_ptq_command(model=BART_PATH, quant=quant, export_fmt=export_fmt) diff --git a/tests/examples/llm_ptq/test_llama.py b/tests/examples/llm_ptq/test_llama.py deleted file mode 100644 index e3ca4c73c..000000000 --- a/tests/examples/llm_ptq/test_llama.py +++ /dev/null @@ -1,145 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os - -import pytest -from _test_utils.examples.run_command import run_llm_ptq_command -from _test_utils.model import TINY_LLAMA_PATH -from _test_utils.torch_misc import minimum_gpu, minimum_sm - -# TODO: Enable export Affine NVFP4 KV cache tests when supported by TRTLLM -# TODO: sparsegpt test is disable due to some bug when sparsifying models with Pytorch 2.4.0a0+3bcc3cddb5.nv24.7 - - -# Use actual TinyLlama-1.1B for nightly tests -@pytest.fixture(scope="module") -def llama_path(tiny_llama_path): - fast_tests = os.getenv("MODELOPT_FAST_TESTS", "true").lower() == "true" - if fast_tests: - return tiny_llama_path - return TINY_LLAMA_PATH - - -@pytest.mark.parametrize( - ("quant", "export_fmt", "sparsity"), - [ - ("fp16", "tensorrt_llm", None), - ("bf16", "tensorrt_llm", None), - ("int8_sq", "tensorrt_llm", None), - # ("int8_sq", "tensorrt_llm", "sparsegpt"), - ("int4_awq", "tensorrt_llm", None), - ("int4_awq", "hf", None), - ("nvfp4", "tensorrt_llm", None), - ("nvfp4", "hf", None), - ("nvfp4_awq", "tensorrt_llm", None), - ("nvfp4_awq", "hf", None), - ], -) -def test_llama(llama_path, quant, export_fmt, sparsity): - run_llm_ptq_command(model=llama_path, quant=quant, export_fmt=export_fmt, sparsity=sparsity) - - -@pytest.mark.parametrize( - ("quant", "export_fmt"), - [ - ("int4_awq,nvfp4,fp8,w4a8_awq", "tensorrt_llm"), - ("int4_awq,nvfp4,fp8", "hf"), - ], -) -def test_llama_autoquant(llama_path, quant, export_fmt): - run_llm_ptq_command( - model=llama_path, - quant=quant, - export_fmt=export_fmt, - calib_batch_size=4, - auto_quantize_bits=6.4, - ) - - -@pytest.mark.parametrize( - ("quant", "export_fmt", "kv_cache_quant"), - [ - ("nvfp4_awq", "tensorrt_llm", "nvfp4"), - ("nvfp4_awq", "hf", "nvfp4"), - # ("nvfp4_awq", "tensorrt_llm", "nvfp4_affine"), - # ("nvfp4_awq", "hf", "nvfp4_affine"), - ], -) -def test_llama_kv_cache(llama_path, quant, export_fmt, kv_cache_quant): - run_llm_ptq_command( - model=llama_path, quant=quant, export_fmt=export_fmt, kv_cache_quant=kv_cache_quant - ) - - -@pytest.mark.parametrize( - ("quant", "export_fmt", "kv_cache_quant"), - [ - ("int4_awq,nvfp4,fp8,w4a8_awq", "tensorrt_llm", "nvfp4"), - ("int4_awq,nvfp4,fp8,w4a8_awq", "hf", "nvfp4"), - # ("int4_awq,nvfp4,fp8,w4a8_awq", "tensorrt_llm", "nvfp4_affine"), - # ("int4_awq,nvfp4,fp8,w4a8_awq", "hf", "nvfp4_affine"), - ], -) -def test_llama_autoquant_kv_cache(llama_path, quant, export_fmt, kv_cache_quant): - run_llm_ptq_command( - model=llama_path, - quant=quant, - export_fmt=export_fmt, - calib_batch_size=4, - auto_quantize_bits=6.4, - kv_cache_quant=kv_cache_quant, - ) - - -@pytest.mark.parametrize( - ("quant", "export_fmt", "sparsity", "kv_cache_quant"), - [ - ("fp8", "tensorrt_llm", None, None), - ("fp8", "tensorrt_llm", None, "none"), # disable kv cache quantization - # ("fp8", "tensorrt_llm", "sparsegpt", None), - ("fp8", "hf", None, None), - ("w4a8_awq", "tensorrt_llm", None, None), - ], -) -@minimum_sm(89) -def test_llama_sm89(llama_path, quant, export_fmt, sparsity, kv_cache_quant): - run_llm_ptq_command( - model=llama_path, - quant=quant, - export_fmt=export_fmt, - sparsity=sparsity, - kv_cache_quant=kv_cache_quant, - ) - - -@pytest.mark.parametrize( - ("quant", "tasks", "sparsity", "tp", "pp"), - [ - # TP - ("fp16", "build", None, 2, 1), - # ("fp16", "build", "sparsegpt", 1), - ("nvfp4", "build", None, 2, 1), - ("fp16", "benchmark", None, 2, 1), - # ("fp16", "benchmark", "sparsegpt", 2, 1), - # PP - # ("nvfp4", "build", None, 1, 2), - # ("fp16", "build", None, 1, 2), - # ("fp16", "build", "sparsegpt", 1, 2), - ], -) -@minimum_gpu(2) -def test_llama_multi_gpu(llama_path, quant, tasks, sparsity, tp, pp): - run_llm_ptq_command(model=llama_path, quant=quant, tasks=tasks, sparsity=sparsity, tp=tp, pp=pp) diff --git a/tests/examples/llm_ptq/test_llm_ptq.py b/tests/examples/llm_ptq/test_llm_ptq.py new file mode 100644 index 000000000..b3eccf2b6 --- /dev/null +++ b/tests/examples/llm_ptq/test_llm_ptq.py @@ -0,0 +1,163 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +import os + +import pytest +from _test_utils.model import BART_PATH, MIXTRAL_PATH, T5_PATH, TINY_LLAMA_PATH, WHISPER_PATH +from _test_utils.ptq_utils import PTQCommand, WithRequirements + + +@pytest.mark.parametrize( + "command", + [ + PTQCommand(quant="fp16"), + PTQCommand(quant="fp8", min_sm=89), + ], + ids=PTQCommand.param_str, +) +def test_ptq_bart(command): + command.run(BART_PATH) + + +class TestT5(WithRequirements): + requirements = [("transformers", "4.48.0")] + + @pytest.mark.parametrize( + "command", + [ + PTQCommand(quant="fp16"), + PTQCommand(quant="fp8", min_sm=89), + ], + ids=PTQCommand.param_str, + ) + def test_ptq_t5(self, command): + command.run(T5_PATH) + + +@pytest.mark.parametrize( + "command", + [ + PTQCommand(quant="fp16"), + PTQCommand(quant="fp8", min_sm=89), + PTQCommand(quant="fp8", export_fmt="hf", min_sm=89), + ], + ids=PTQCommand.param_str, +) +def test_ptq_mixtral(command): + command.run(MIXTRAL_PATH) + + +class TestWhisper(WithRequirements): + requirements = [ + ("librosa", None), + ("soundfile", None), + ] + + @pytest.mark.parametrize( + "command", + [ + # Auto-batch-size computation seems to take >10mins for Whisper hence using a fixed batch size + PTQCommand(quant="fp16", calib_batch_size=16), + PTQCommand(quant="fp8", calib_batch_size=16, min_sm=89), + ], + ids=PTQCommand.param_str, + ) + def test_ptq_whisper(self, command): + command.run(WHISPER_PATH) + + +@pytest.fixture(scope="module") +def llama_path(tiny_llama_path): + fast_tests = os.getenv("MODELOPT_FAST_TESTS", "true").lower() == "true" + if fast_tests: + return tiny_llama_path + return TINY_LLAMA_PATH + + +@pytest.mark.parametrize( + "command", + [ + PTQCommand(quant="fp16"), + PTQCommand(quant="bf16"), + PTQCommand(quant="int8_sq"), + # ("int8_sq", "tensorrt_llm", "sparsegpt"), + PTQCommand(quant="int4_awq"), + PTQCommand(quant="int4_awq", export_fmt="hf"), + PTQCommand(quant="nvfp4"), + PTQCommand(quant="nvfp4", export_fmt="hf"), + PTQCommand(quant="nvfp4_awq"), + PTQCommand(quant="nvfp4_awq", export_fmt="hf"), + # + # autoquant + PTQCommand( + quant="int4_awq,nvfp4,fp8,w4a8_awq", + calib_batch_size=4, + auto_quantize_bits=6.4, + ), + PTQCommand( + quant="int4_awq,nvfp4,fp8", + export_fmt="hf", + calib_batch_size=4, + auto_quantize_bits=6.4, + ), + # + # kv_cache + PTQCommand(quant="nvfp4_awq", kv_cache_quant="nvfp4"), + PTQCommand(quant="nvfp4_awq", export_fmt="hf", kv_cache_quant="nvfp4"), + # ("nvfp4_awq", "tensorrt_llm", "nvfp4_affine"), + # ("nvfp4_awq", "hf", "nvfp4_affine"), + # + # autoquant_kv_cache + PTQCommand( + quant="int4_awq,nvfp4,fp8,w4a8_awq", + kv_cache_quant="nvfp4", + calib_batch_size=4, + auto_quantize_bits=6.4, + ), + PTQCommand( + quant="int4_awq,nvfp4,fp8,w4a8_awq", + export_fmt="hf", + kv_cache_quant="nvfp4", + calib_batch_size=4, + auto_quantize_bits=6.4, + ), + # ("int4_awq,nvfp4,fp8,w4a8_awq", "tensorrt_llm", "nvfp4_affine"), + # ("int4_awq,nvfp4,fp8,w4a8_awq", "hf", "nvfp4_affine"), + # + # sm89 + PTQCommand(quant="fp8", min_sm=89), + PTQCommand(quant="fp8", kv_cache_quant="none", min_sm=89), + # ("fp8", "tensorrt_llm", "sparsegpt", None), + PTQCommand(quant="fp8", export_fmt="hf", min_sm=89), + PTQCommand(quant="w4a8_awq", min_sm=89), + # + # multi_gpu + # TP + PTQCommand(quant="fp16", tp=2, pp=1, min_gpu=2), + # ("fp16", "build", "sparsegpt", 1), + PTQCommand(quant="nvfp4", tp=2, pp=1, min_gpu=2), + PTQCommand(quant="fp16", tasks="benchmark", tp=2, pp=1, min_gpu=2), + # ("fp16", "benchmark", "sparsegpt", 2, 1), + # PP + # ("nvfp4", "build", None, 1, 2), + # ("fp16", "build", None, 1, 2), + # ("fp16", "build", "sparsegpt", 1, 2), + ], + ids=PTQCommand.param_str, +) +def test_ptq_llama(command, llama_path): + command.run(llama_path) diff --git a/tests/examples/llm_ptq/test_mixtral.py b/tests/examples/llm_ptq/test_mixtral.py deleted file mode 100644 index 0fea7f266..000000000 --- a/tests/examples/llm_ptq/test_mixtral.py +++ /dev/null @@ -1,31 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -import pytest -from _test_utils.examples.run_command import run_llm_ptq_command -from _test_utils.model import MIXTRAL_PATH -from _test_utils.torch_misc import minimum_sm - - -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp16", "tensorrt_llm")]) -def test_mixtral(quant, export_fmt): - run_llm_ptq_command(model=MIXTRAL_PATH, quant=quant, export_fmt=export_fmt) - - -@minimum_sm(89) -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp8", "tensorrt_llm"), ("fp8", "hf")]) -def test_mixtral_sm89(quant, export_fmt): - run_llm_ptq_command(model=MIXTRAL_PATH, quant=quant, export_fmt=export_fmt) diff --git a/tests/examples/llm_ptq/test_t5.py b/tests/examples/llm_ptq/test_t5.py deleted file mode 100644 index b81040f50..000000000 --- a/tests/examples/llm_ptq/test_t5.py +++ /dev/null @@ -1,42 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import subprocess -from pathlib import Path - -import pytest -from _test_utils.examples.run_command import run_llm_ptq_command -from _test_utils.model import T5_PATH -from _test_utils.torch_misc import minimum_sm - - -@pytest.fixture(scope="session", autouse=True) -def install_t5_requirements(): - subprocess.run( - ["pip", "install", "-r", "requirements-t5.txt"], - cwd=Path(__file__).parent.parent.parent.parent / "examples/llm_ptq", - check=True, - ) - - -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp16", "tensorrt_llm")]) -def test_t5(quant, export_fmt): - run_llm_ptq_command(model=T5_PATH, quant=quant, export_fmt=export_fmt) - - -@minimum_sm(89) -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp8", "tensorrt_llm")]) -def test_t5_sm89(quant, export_fmt): - run_llm_ptq_command(model=T5_PATH, quant=quant, export_fmt=export_fmt) diff --git a/tests/examples/llm_ptq/test_whisper.py b/tests/examples/llm_ptq/test_whisper.py deleted file mode 100644 index e52e1f921..000000000 --- a/tests/examples/llm_ptq/test_whisper.py +++ /dev/null @@ -1,44 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import subprocess -from pathlib import Path - -import pytest -from _test_utils.examples.run_command import run_llm_ptq_command -from _test_utils.torch_misc import minimum_sm - -WHISPER_PATH = "openai/whisper-tiny" - - -@pytest.fixture(scope="session", autouse=True) -def install_whisper_requirements(): - subprocess.run( - ["pip", "install", "-r", "requirements-whisper.txt"], - cwd=Path(__file__).parent.parent.parent.parent / "examples/llm_ptq", - check=True, - ) - - -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp16", "tensorrt_llm")]) -def test_whisper(quant, export_fmt): - run_llm_ptq_command(model=WHISPER_PATH, quant=quant, export_fmt=export_fmt) - - -@minimum_sm(89) -@pytest.mark.parametrize(("quant", "export_fmt"), [("fp8", "tensorrt_llm")]) -def test_whisper_sm89(quant, export_fmt): - # Auto-batch-size computation seems to take >10mins for Whisper hence using a fixed batch size - run_llm_ptq_command(model=WHISPER_PATH, quant=quant, export_fmt=export_fmt, calib_batch_size=16) diff --git a/tests/gpu/onnx/test_onnx_torch_int4_awq.py b/tests/gpu/onnx/test_quantize_onnx_torch_int4_awq.py similarity index 94% rename from tests/gpu/onnx/test_onnx_torch_int4_awq.py rename to tests/gpu/onnx/test_quantize_onnx_torch_int4_awq.py index b8b413ccb..826d90155 100644 --- a/tests/gpu/onnx/test_onnx_torch_int4_awq.py +++ b/tests/gpu/onnx/test_quantize_onnx_torch_int4_awq.py @@ -34,6 +34,10 @@ if int4.has_cupy: else: import numpy as np +# TODO: Rename this script to test_onnx_torch_int4_awq.py +# For that, we need to investigate failure in 'pytest tests/gpu/onnx'. +# test_qdq_utils_fp8.py::test_fused_q[bf16,fp16] fails if this script runs after the int4 test, but not before. + def test_int4_awq(tmp_path): def _forward_loop(model, dataloader): diff --git a/tests/gpu/torch/export/test_unified_export_megatron.py b/tests/gpu/torch/export/test_unified_export_megatron.py index 7476271fb..2d21da56a 100644 --- a/tests/gpu/torch/export/test_unified_export_megatron.py +++ b/tests/gpu/torch/export/test_unified_export_megatron.py @@ -21,20 +21,15 @@ import torch import transformers from _test_utils.import_helper import skip_if_no_megatron from _test_utils.torch_dist.dist_utils import spawn_multiprocess_job -from _test_utils.torch_dist.plugins.megatron_common import ( - get_mcore_gpt_model, - initialize_for_megatron, -) +from _test_utils.torch_dist.plugins.megatron_common import get_mcore_gpt_model from _test_utils.torch_model.transformers_models import create_tiny_llama_dir skip_if_no_megatron(apex_or_te_required=True) import modelopt.torch.speculative as mtsp from modelopt.torch.export import export_mcore_gpt_to_hf, import_mcore_gpt_from_hf -from modelopt.torch.speculative.plugins.megatron import ( - _DynamicEagleGPTModel, - _DynamicMedusaGPTModel, -) +from modelopt.torch.speculative.plugins.megatron_eagle import _DynamicEagleGPTModel +from modelopt.torch.speculative.plugins.megatron_medusa import _DynamicMedusaGPTModel def _test_unified_export_megatron(tmp_path, model_type, arch, algo, rank, size): @@ -50,11 +45,10 @@ def _test_unified_export_megatron(tmp_path, model_type, arch, algo, rank, size): activation_func = "squared_relu" if model_type == "nemotron" else "swiglu" normalization = "LayerNorm" if model_type == "nemotron" else "RMSNorm" - initialize_for_megatron(tensor_model_parallel_size=size, pipeline_model_parallel_size=1) - model = get_mcore_gpt_model( tensor_model_parallel_size=size, pipeline_model_parallel_size=1, + initialize_megatron=True, num_layers=num_layers, hidden_size=hidden_size, num_attention_heads=num_attention_heads, @@ -147,11 +141,10 @@ def _test_unified_import_megatron(tiny_llama_dir, rank, size): activation_func = "swiglu" normalization = "RMSNorm" - initialize_for_megatron(tensor_model_parallel_size=size, pipeline_model_parallel_size=1) - model = get_mcore_gpt_model( tensor_model_parallel_size=size, pipeline_model_parallel_size=1, + initialize_megatron=True, num_layers=num_layers, hidden_size=hidden_size, num_attention_heads=num_attention_heads, diff --git a/tests/gpu/torch/nas/plugins/test_megatron_dynamic_modules.py b/tests/gpu/torch/nas/plugins/test_megatron_gpt_dynamic_modules.py similarity index 88% rename from tests/gpu/torch/nas/plugins/test_megatron_dynamic_modules.py rename to tests/gpu/torch/nas/plugins/test_megatron_gpt_dynamic_modules.py index 8b25e00c2..00f924346 100644 --- a/tests/gpu/torch/nas/plugins/test_megatron_dynamic_modules.py +++ b/tests/gpu/torch/nas/plugins/test_megatron_gpt_dynamic_modules.py @@ -24,9 +24,8 @@ skip_if_no_megatron(apex_or_te_required=True) from _test_utils.torch_dist.dist_utils import spawn_multiprocess_job from _test_utils.torch_dist.plugins.megatron_common import ( get_mcore_gpt_model, - initialize_for_megatron, - run_mcore_gpt_inference, - run_mcore_gpt_inference_with_dummy_input, + run_mcore_inference, + run_mcore_inference_with_dummy_input, ) from _test_utils.torch_misc import set_seed from megatron.core.parallel_state import destroy_model_parallel @@ -39,7 +38,7 @@ from megatron.core.transformer.transformer_layer import TransformerLayer import modelopt.torch.nas as mtn from modelopt.torch.nas.plugins.megatron import ( _DynamicColumnParallelLinear, - _DynamicGPTModel, + _DynamicMCoreLanguageModel, _DynamicMLP, _DynamicProjRowParallelLinear, _DynamicQKVColumnParallelLinear, @@ -65,15 +64,14 @@ def _test_gpt_search_space( num_layers = min(size * 2, 8) hidden_size = 256 ffn_hidden_size = 128 - max_sequence_length = 32 + max_sequence_length = 16 vocab_size = 64 batch_size = 2 - initialize_for_megatron(tensor_model_parallel_size=1, pipeline_model_parallel_size=size) - model = get_mcore_gpt_model( tensor_model_parallel_size=1, pipeline_model_parallel_size=size, + initialize_megatron=True, num_layers=num_layers, hidden_size=hidden_size, num_attention_heads=num_attention_heads, @@ -87,7 +85,7 @@ def _test_gpt_search_space( model = mtn.convert(model, "mcore_gpt_minitron") - assert isinstance(model, _DynamicGPTModel) + assert isinstance(model, _DynamicMCoreLanguageModel) for m in model.modules(): if isinstance(m, VocabParallelEmbedding): assert isinstance(m, _DynamicVocabParallelEmbedding) @@ -104,24 +102,26 @@ def _test_gpt_search_space( # NOTE: `search_space_size` does not reduce across TP/PP groups ss_size_per_pp = search_space_size(model) + ffn_hidden_size_choices = ffn_hidden_size // channel_divisor + hidden_size_choices = hidden_size // channel_divisor + num_layers_per_pp = num_layers // size assert ( ss_size_per_pp - == (num_attention_heads * ffn_hidden_size // channel_divisor) ** (num_layers / size) - * hidden_size + == (num_attention_heads * ffn_hidden_size_choices) ** num_layers_per_pp * num_layers - // channel_divisor + * hidden_size_choices ) # Make sure forward pass works on min and centroid subnets prompt_tokens = torch.randint(0, vocab_size, (batch_size, max_sequence_length)).cuda() for sample_func in [min, max, centroid]: mtn.sample(model, sample_func) - output = run_mcore_gpt_inference(model, prompt_tokens) + output = run_mcore_inference(model, prompt_tokens) assert output.shape == (batch_size, max_sequence_length, vocab_size) # Make sure export and forward pass works on centroid model mtn.export(model) - _ = run_mcore_gpt_inference(model, prompt_tokens, model.hidden_size) + _ = run_mcore_inference(model, prompt_tokens, model.hidden_size) assert not any(named_dynamic_modules(model)) @@ -157,13 +157,10 @@ def _test_gpt_parameter_sorting(activation_func, rank, size): vocab_size = 128 batch_size = 2 - initialize_for_megatron( - tensor_model_parallel_size=1, pipeline_model_parallel_size=size, seed=SEED - ) - model = get_mcore_gpt_model( tensor_model_parallel_size=1, pipeline_model_parallel_size=size, + initialize_megatron=True, num_layers=num_layers, hidden_size=hidden_size, num_attention_heads=num_attention_heads, @@ -184,27 +181,27 @@ def _test_gpt_parameter_sorting(activation_func, rank, size): # Compute activations for sorting for _ in range(5): - run_mcore_gpt_inference_with_dummy_input(model, batch_size) + run_mcore_inference_with_dummy_input(model, batch_size) # Get the output of the original model prompt_tokens = torch.randint(0, vocab_size, (batch_size, max_sequence_length)).cuda() - y1 = run_mcore_gpt_inference(model, prompt_tokens) + y1 = run_mcore_inference(model, prompt_tokens) search_space.sort_parameters() # check if all ffn_hidden_size, num_heads_per_group, num_query_groups, hidden_size have been sorted - num_sortable_per_pp = sum( - 1 for _, hp in search_space.named_hparams(configurable=True) if hp.importance is not None - ) - expected_num_sortable_hps_per_layer = 4 - assert num_sortable_per_pp == expected_num_sortable_hps_per_layer * num_layers // size + sortable_per_pp = [ + n for n, hp in search_space.named_hparams(configurable=True) if hp.importance is not None + ] + # 3 hps per layer + 1 for hidden_size (num_layers is not sorted!) + assert len(sortable_per_pp) == 3 * num_layers // size + 1 # Export since sorting force reassigns SelfAttention weights which we dont want to re-sort! # TODO: ideally we shouldn't need this search_space.export() # sanity check if the model functionality is preserved after sorting - y2 = run_mcore_gpt_inference(model, prompt_tokens) + y2 = run_mcore_inference(model, prompt_tokens) # # check if the inference results after sorting is the same assert all( @@ -230,18 +227,15 @@ def test_expand_head_indices(): def test_megatron_self_attention_head_sorting(distributed_setup_size_1): - initialize_for_megatron(tensor_model_parallel_size=1, pipeline_model_parallel_size=1, seed=SEED) - model = get_mcore_gpt_model( tensor_model_parallel_size=1, pipeline_model_parallel_size=1, + initialize_megatron=True, num_layers=1, hidden_size=16, num_attention_heads=8, num_query_groups=2, ffn_hidden_size=16, - max_sequence_length=32, - vocab_size=32, activation_func="squared_relu", ) diff --git a/tests/gpu/torch/prune/plugins/test_mcore_gpt_minitron_pruning.py b/tests/gpu/torch/prune/plugins/test_mcore_gpt_minitron_pruning.py index 396871459..b07a71964 100644 --- a/tests/gpu/torch/prune/plugins/test_mcore_gpt_minitron_pruning.py +++ b/tests/gpu/torch/prune/plugins/test_mcore_gpt_minitron_pruning.py @@ -24,8 +24,7 @@ skip_if_no_megatron(apex_or_te_required=True) from _test_utils.torch_dist.dist_utils import spawn_multiprocess_job from _test_utils.torch_dist.plugins.megatron_common import ( get_mcore_gpt_model, - initialize_for_megatron, - run_mcore_gpt_inference_with_dummy_input, + run_mcore_inference_with_dummy_input, ) import modelopt.torch.prune as mtp @@ -47,7 +46,7 @@ def _test_mcore_gpt_pruning( ): hidden_size = 256 ffn_hidden_size = 256 - max_sequence_length = 32 + max_sequence_length = 16 vocab_size = 64 batch_size = 2 @@ -67,11 +66,10 @@ def _test_mcore_gpt_pruning( else: raise ValueError(f"Unsupported size {size}") - initialize_for_megatron(tensor_model_parallel_size=1, pipeline_model_parallel_size=size) - model = get_mcore_gpt_model( tensor_model_parallel_size=1, pipeline_model_parallel_size=size, + initialize_megatron=True, num_layers=num_layers, hidden_size=hidden_size, num_attention_heads=num_attention_heads, @@ -87,7 +85,7 @@ def _test_mcore_gpt_pruning( def forward_loop(m): for _ in range(5): - run_mcore_gpt_inference_with_dummy_input(m, batch_size, hidden_size) + run_mcore_inference_with_dummy_input(m, batch_size, hidden_size) pruned_ffn = ffn_hidden_size // pruned_ffn_div pruned_num_attention_heads = num_attention_heads // pruned_num_attention_heads_div @@ -132,7 +130,7 @@ def _test_mcore_gpt_pruning( ) # Assert forward pass works on the pruned model - run_mcore_gpt_inference_with_dummy_input(model, batch_size, pruned_hidden_size) + run_mcore_inference_with_dummy_input(model, batch_size, pruned_hidden_size) # Assert model.config is updated for correct save/restoring assert model.config.ffn_hidden_size == pruned_ffn diff --git a/tests/gpu/torch/quantization/plugins/test_megatron.py b/tests/gpu/torch/quantization/plugins/test_megatron.py index 56137fef3..5c69ffc08 100644 --- a/tests/gpu/torch/quantization/plugins/test_megatron.py +++ b/tests/gpu/torch/quantization/plugins/test_megatron.py @@ -23,7 +23,7 @@ from _test_utils.torch_dist.plugins.megatron_common import ( MegatronModel, get_mcore_gpt_model, initialize_for_megatron, - run_mcore_gpt_inference, + run_mcore_inference, sharded_state_dict_test_helper, ) from _test_utils.torch_misc import set_seed @@ -151,7 +151,7 @@ def _test_sharded_state_dict(tmp_path, config, hidden_size, modelopt_version, co ).cuda() def forward_fn(model): - return run_mcore_gpt_inference(model, prompt_tokens) + return run_mcore_inference(model, prompt_tokens) model_ref = mtq.quantize(model_ref, config, forward_fn) if compress: @@ -272,7 +272,7 @@ def test_regular_state_dict(distributed_setup_size_1, hidden_size): ).cuda() def forward_fn(model): - return run_mcore_gpt_inference(model, prompt_tokens) + return run_mcore_inference(model, prompt_tokens) model_ref = mtq.quantize(model_ref, mixed_precision_config, forward_fn) @@ -316,7 +316,7 @@ def _test_fp8_real_quantize_helper(rank, size): prompt_tokens = torch.randint(0, model.vocab_size, (2, model.max_sequence_length)).cuda() def forward_fn(model): - return run_mcore_gpt_inference(model, prompt_tokens) + return run_mcore_inference(model, prompt_tokens) forward_fn(model) diff --git a/tests/gpu/torch/speculative/plugins/test_speculative_megatron_modules.py b/tests/gpu/torch/speculative/plugins/test_speculative_megatron_modules.py index ec180db61..ddaa75273 100644 --- a/tests/gpu/torch/speculative/plugins/test_speculative_megatron_modules.py +++ b/tests/gpu/torch/speculative/plugins/test_speculative_megatron_modules.py @@ -12,7 +12,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - +from collections import deque from functools import partial import pytest @@ -22,16 +22,13 @@ from _test_utils.import_helper import skip_if_no_megatron skip_if_no_megatron(apex_or_te_required=True) from _test_utils.torch_dist.dist_utils import spawn_multiprocess_job -from _test_utils.torch_dist.plugins.megatron_common import ( - get_mcore_gpt_model, - initialize_for_megatron, -) +from _test_utils.torch_dist.plugins.megatron_common import get_mcore_gpt_model +from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region import modelopt.torch.speculative as mtsp -from modelopt.torch.speculative.plugins.megatron import ( - _DynamicEagleGPTModel, - _DynamicMedusaGPTModel, -) +from modelopt.torch.speculative.plugins.megatron_eagle import _DynamicEagleGPTModel, right_padding +from modelopt.torch.speculative.plugins.megatron_medusa import _DynamicMedusaGPTModel +from modelopt.torch.speculative.utils import Tree, get_default_attention_mask_and_position_ids def _test_speculative_gpt_model( @@ -43,11 +40,10 @@ def _test_speculative_gpt_model( vocab_size = 64 batch_size = 2 - initialize_for_megatron(tensor_model_parallel_size=size, pipeline_model_parallel_size=1) - model = get_mcore_gpt_model( tensor_model_parallel_size=size, pipeline_model_parallel_size=1, + initialize_megatron=True, num_attention_heads=num_attention_heads, num_query_groups=num_query_groups, max_sequence_length=max_sequence_length, @@ -104,7 +100,6 @@ def _test_speculative_gpt_model( labels = torch.randint(0, vocab_size, (batch_size, max_sequence_length)).cuda() eagle_loss = model(prompt_tokens, position_ids, attention_mask, labels=labels) - print(eagle_loss.shape) assert eagle_loss.shape[0] == batch_size assert eagle_loss.shape[1] == max_sequence_length @@ -138,3 +133,171 @@ def test_speculative_gpt_model( ), backend="nccl", ) + + +def generate_next_tokens(model, eagle_ids, hidden_states, topk=1): + padded_eagle_ids, seq_len, padded_hidden_states = right_padding(eagle_ids, hidden_states) + eagle_attention_mask, eagle_position_ids = get_default_attention_mask_and_position_ids( + padded_eagle_ids + ) + + eagle_inputs = {} + eagle_inputs["input_ids"] = padded_eagle_ids + eagle_inputs["embedding"] = model.embedding( + input_ids=padded_eagle_ids, + position_ids=eagle_position_ids, + ) + eagle_inputs["hidden_states"] = padded_hidden_states + eagle_inputs["attention_mask"] = eagle_attention_mask + + eagle_inputs["rotary_pos_emb"] = None + + _, eagle_logits, eagle_next_hidden_states_input = model._eagle_forward(eagle_inputs, None) + + eagle_logits = eagle_logits[seq_len - 1 : seq_len, :, :] + eagle_next_hidden_states_input = eagle_next_hidden_states_input[seq_len - 1 : seq_len, :, :] + + draft_token = ( + gather_from_tensor_model_parallel_region(eagle_logits).topk(topk, dim=-1)[1].transpose(0, 1) + ) + return draft_token, eagle_next_hidden_states_input + + +def _test_tree_decode(tree_paths, greedy_steps, rank, size): + activation_func = "squared_relu" + normalization = "RMSNorm" + + num_attention_heads = 8 + num_query_groups = size + max_sequence_length = 32 + vocab_size = 64 + batch_size = 1 + + config = {"eagle_num_layers": 1} + + model = get_mcore_gpt_model( + tensor_model_parallel_size=size, + pipeline_model_parallel_size=1, + initialize_megatron=True, + num_attention_heads=num_attention_heads, + num_query_groups=num_query_groups, + max_sequence_length=max_sequence_length, + vocab_size=vocab_size, + activation_func=activation_func, + normalization=normalization, + ).cuda() + + model = mtsp.convert(model, [("eagle", config)]) + + # Prepare inputs for forward. + prompt_tokens = torch.randint(0, vocab_size, (batch_size, max_sequence_length)).cuda() + attention_mask = torch.tril(torch.ones((1, 1, max_sequence_length, max_sequence_length))).cuda() + position_ids = torch.arange(max_sequence_length, dtype=torch.long).unsqueeze(0).cuda() + attention_mask = attention_mask < 0.5 + + model.eval() + tree = Tree(tree_paths) + + input_id, draft_tokens, pred_tokens = model.tree_decode(prompt_tokens, tree=tree) + + # check for empty tree paths + if not tree_paths: + assert draft_tokens is None, "draft_tokens should be None for empty tree paths" + return + + # check when tree decode is same as greedy decode + if greedy_steps: + spec_input_id, spec_draft_tokens = model.pseudo_speculative_generate( + prompt_tokens, steps=greedy_steps + ) + assert (pred_tokens == spec_draft_tokens[0]).all(), ( + f"pred_tokens should be equal to spec_draft_tokens, {pred_tokens} != {spec_draft_tokens[0]}" + ) + assert input_id == spec_input_id[0], ( + f"spec_input_id should be equal to input_id, {input_id} != {spec_input_id[0]}" + ) + return + + orig_hidden_states, _ = model._base_model_forward( + prompt_tokens, + position_ids, + attention_mask, + ) + + # Get Eagle-specific input hidden states + eagle_hidden_states = model._get_eagle_input_hidden_states(orig_hidden_states) + # Extract tokens for Eagle processing (excluding first token) + eagle_tokens = prompt_tokens[:, 1:] + + # Initialize lists to store draft tokens and hidden states + draft_tokens_list = [input_id] + eagle_hidden_states_list = [eagle_hidden_states] + # Track indices for token and hidden state mapping + index_list = [[[0, 0]]] + + # Initialize queue for breadth-first tree traversal + queue = deque([(draft_tokens, 0)]) + + # Process tree nodes in breadth-first order + while queue: + tree_token, index = queue.popleft() + if not tree_token.children: + continue + # Collect tokens and hidden states for current node + tokens = [] + hidden_states = [] + for token_idx, state_idx in index_list[index]: + tokens.append(draft_tokens_list[token_idx]) + hidden_states.append(eagle_hidden_states_list[state_idx]) + # Concatenate tokens and hidden states for processing + tokens = torch.cat([eagle_tokens, torch.cat(tokens, dim=-1)], dim=-1) + hidden_states = torch.cat(hidden_states, dim=0) + # Generate next token and get updated hidden states + draft_token, eagle_next_hidden_states_input = generate_next_tokens( + model, tokens, hidden_states, topk=len(tree_token.children) + ) + # Verify generated tokens match expected tree structure + for child_idx, tree_node in enumerate(tree_token.children.values()): + assert tree_node.value[0] == draft_token[0, 0, child_idx], ( + f"token mismatch at {tree_node.value[0]} != {draft_token[0, 0, child_idx]}" + ) + # Update tracking variables + cur_len = len(draft_tokens_list) + eagle_hidden_states_list.append(eagle_next_hidden_states_input) + # Process children and add them to the queue + for child_idx, child_tree_token in enumerate(tree_token.children.values()): + queue.append([child_tree_token, len(index_list)]) + draft_tokens_list.append(draft_token[:, :, child_idx]) + index_list.append( + index_list[index][:] + [[cur_len + child_idx, len(eagle_hidden_states_list) - 1]] + ) + + +@pytest.mark.parametrize( + ("greedy_steps", "tree_paths"), + [ + (None, []), + (3, [[0], [0, 0], [0, 0, 0]]), + ( + None, + [ + [0], + [1], + [0, 0], + [0, 1], + [1, 1], + [0, 0, 0], + [0, 0, 1], + [1, 0, 0], + [1, 0], + [0, 0, 1, 0], + ], + ), + ], +) +def test_tree_decode_model(greedy_steps, tree_paths): + spawn_multiprocess_job( + size=torch.cuda.device_count(), + job=partial(_test_tree_decode, tree_paths, greedy_steps), + backend="nccl", + ) diff --git a/tests/unit/onnx/test_quantize_int4.py b/tests/unit/onnx/test_quantize_zint4.py similarity index 95% rename from tests/unit/onnx/test_quantize_int4.py rename to tests/unit/onnx/test_quantize_zint4.py index f2e094109..4249f1fe1 100644 --- a/tests/unit/onnx/test_quantize_int4.py +++ b/tests/unit/onnx/test_quantize_zint4.py @@ -25,6 +25,10 @@ import modelopt.onnx.quantization as moq from modelopt.onnx.quantization.int4 import quantize as quantize_int4 from modelopt.onnx.utils import save_onnx +# TODO: Rename this script to *_int4.py +# For that, we need to investigate failure in 'pytest tests/unit/onnx'. +# test_quantize_int8.py::test_int8[bf16] fails if this script runs after the int4 test, but not before. + def _matmul_model(w: np.ndarray, in_shape: Sequence[int], out_shape: Sequence[int], tmp_path): # Assumes diff --git a/tests/unit/torch/distill/test_distill.py b/tests/unit/torch/distill/test_distill.py index d6579beaf..c989f1448 100644 --- a/tests/unit/torch/distill/test_distill.py +++ b/tests/unit/torch/distill/test_distill.py @@ -111,6 +111,23 @@ def test_distillation_model_multiloss_balancer(): distillation_model.compute_kd_loss(student_loss=output.mean()) +def test_distillation_model_mft(): + student = tiny_mobilenet().train() + config = { + "teacher_model": tiny_alexnet, + "criterion": mtd.MFTLoss(threshold=0.2), + "loss_balancer": None, + } + + distillation_model = mtd.convert(student, mode=[("kd_loss", config)]) + + input_tensor = get_input_tensor() + labels = torch.randint(0, 10, (input_tensor.size(0),)) # Dummy labels for MFT + distillation_model(input_tensor) + loss = distillation_model.compute_kd_loss(labels=labels) + assert isinstance(loss, torch.Tensor) and loss.numel() == 1 + + def test_distillation_mode_default_config(): student = tiny_mobilenet() with pytest.raises(AssertionError): diff --git a/tests/unit/torch/opt/plugins/test_hf_patching.py b/tests/unit/torch/opt/plugins/test_hf_patching.py new file mode 100644 index 000000000..4d795c8f1 --- /dev/null +++ b/tests/unit/torch/opt/plugins/test_hf_patching.py @@ -0,0 +1,66 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +from _test_utils.torch_model.transformers_models import ( + create_tiny_llama_dir, + get_tiny_qwen3, + tf_output_tester, +) +from transformers import AutoModelForCausalLM + +import modelopt.torch.distill as mtd +import modelopt.torch.opt as mto + + +def _teacher_factory(model_name_or_path, teacher_model_type): + if teacher_model_type == "qwen3": + return get_tiny_qwen3() + else: + return AutoModelForCausalLM.from_pretrained( + model_name_or_path, + ) + + +@pytest.mark.parametrize( + ("model_cls", "teacher_model_type"), + [ + (AutoModelForCausalLM, "llama"), + (AutoModelForCausalLM, "qwen3"), + ], +) +def test_nested_model_save_restore(tmp_path, model_cls, teacher_model_type): + tiny_llama_dir = create_tiny_llama_dir(tmp_path) + + model_ref = model_cls.from_pretrained(tiny_llama_dir) + + kd_config = { + "teacher_model": ( + _teacher_factory, + (tiny_llama_dir, teacher_model_type), + {}, + ), + "criterion": mtd.LogitsDistillationLoss(), + "expose_minimal_state_dict": False, + } + model = mtd.convert(model_ref, mode=[("kd_loss", kd_config)]) + model.save_pretrained(tiny_llama_dir / "modelopt_model") + + model_test = model_cls.from_pretrained(tiny_llama_dir / "modelopt_model") + + tf_output_tester(model, model_test) + # since distill model contains loss function, we compare state of model and teacher model manually + assert mto.modelopt_state(model.model) == mto.modelopt_state(model_test.model) + assert mto.modelopt_state(model._teacher_model) == mto.modelopt_state(model_test._teacher_model) diff --git a/tests/unit/torch/quantization/test_autoquant.py b/tests/unit/torch/quantization/test_autoquant.py index 0ff2cb223..6c7ef17a3 100644 --- a/tests/unit/torch/quantization/test_autoquant.py +++ b/tests/unit/torch/quantization/test_autoquant.py @@ -23,7 +23,7 @@ from _test_utils.torch_quantization.models import SimpleConv, SimpleConvLinear, import modelopt.torch.opt as mto import modelopt.torch.quantization as mtq -from modelopt.core.torch.quantization.algorithms import ( +from modelopt.torch.quantization.algorithms import ( QuantRecipe, QuantRecipeHparam, estimate_quant_compression, diff --git a/tests/unit/torch/trace/plugins/test_transformers_attention_symbols.py b/tests/unit/torch/trace/plugins/test_transformers_attention_symbols.py index a74d67de0..ef4289ba3 100644 --- a/tests/unit/torch/trace/plugins/test_transformers_attention_symbols.py +++ b/tests/unit/torch/trace/plugins/test_transformers_attention_symbols.py @@ -26,7 +26,10 @@ from modelopt.torch.trace.plugins.transformers import get_hf_attn_sym_info @pytest.mark.parametrize( ("cls_type", "config"), [ - (BertAttention, BertConfig(hidden_size=8, num_attention_heads=2)), + ( + BertAttention, + BertConfig(hidden_size=8, num_attention_heads=2, attn_implementation="eager"), + ), (GPTJAttention, GPTJConfig(n_embd=12, n_head=2, rotary_dim=4)), ], ) diff --git a/tests/unit/torch/trace/test_symbol.py b/tests/unit/torch/trace/test_symbol.py index 5f6d9eb2d..1396912fc 100644 --- a/tests/unit/torch/trace/test_symbol.py +++ b/tests/unit/torch/trace/test_symbol.py @@ -19,14 +19,6 @@ import torch.nn as nn from modelopt.torch.trace import RobustTracer, Symbol, SymMap from modelopt.torch.trace.modules.nn import get_conv_sym_info, get_linear_sym_info -try: - import megatron # noqa: F401 - import transformer_engine # noqa: F401 - - SKIP = True -except ImportError: - SKIP = False - def test_symbol_cls(): sym = Symbol(elastic_dims={1, 2}, cl_type=Symbol.CLType.INCOMING) @@ -117,11 +109,7 @@ def test_sym_map(model): assert_num_symbols() -@pytest.mark.skipif(SKIP, reason="This cpu unit test will fail on GPU with Megatron/TE installed!") def test_sym_map_registry(): - # NOTE: If running with transformer_engine or megatron-core installed, this test will fail. - # Ignoring this error for now, as it will only be there if running CPU tests on a GPU machine - # with the above packages installed. mods_in_registry = { nn.Linear, nn.BatchNorm1d, @@ -151,6 +139,20 @@ def test_sym_map_registry(): except ImportError: pass + try: + from megatron.core.models.gpt import GPTModel + + mods_in_registry.add(GPTModel) + except ImportError: + pass + + try: + from megatron.core.models.mamba import MambaModel + + mods_in_registry.add(MambaModel) + except ImportError: + pass + not_a_leaf = {nn.Sequential} dependent_registry = set()