Update files on GitHub

Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
This commit is contained in:
Keval Morabia
2025-07-31 22:48:09 +05:30
parent b40f4782d2
commit 4c611e47a6
139 changed files with 8483 additions and 3165 deletions
+1
View File
@@ -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|
+14 -1
View File
@@ -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)
^^^^^^^^^^^^^^^^^
@@ -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 <https://github.com/microsoft/Olive/tree/main/examples/phi3#quantize-models-with-nvidia-tensorrt-model-optimizer>`_ 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 <https://github.com/microsoft/Olive/tree/main/examples/phi3#quantize-models-with-nvidia-tensorrt-model-optimizer>`_ Olive example.
+19
View File
@@ -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 <modelopt.torch.opt.conversion.restore>` utility.
.. tip::
When training the student on a small corpus of ground truth data, consider using :class:`MFTLoss <modelopt.torch.distill.MFTLoss>` for to perform Minifinetuning in lieu of the standard
:class:`LogitsDistillationLoss <modelopt.torch.distill.losses.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 <modelopt.torch.distill.losses.MFTDistillationLoss>`, which can
be used in place of the standard :class:`LogitsDistillationLoss <modelopt.torch.distill.losses.LogitsDistillationLoss>`.
More information about the technique can be found in the original paper:
`Minifinetuning: Low-Data Generation Domain Adaptation through Corrective Self-Distillation <https://arxiv.org/abs/2506.15702>`_.
+27 -1
View File
@@ -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)
@@ -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/
+1 -1
View File
@@ -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": {
+8 -2
View File
@@ -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,
+20 -30
View File
@@ -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,
@@ -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
}
@@ -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
}
@@ -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
}
+1 -1
View File
@@ -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
+1 -8
View File
@@ -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)
+30
View File
@@ -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 <hf-model-name> --finetune-recipe <recipe-name>`. 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 <fp8, nvfp4>`.
#### QAD
In order to train using QAD, launch the example with `python qat/nemo_qat_flow.py --model-name <hf-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 <my_template.txt>` flag.
+17
View File
@@ -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}",
],
+9 -3
View File
@@ -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`)<br>
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
+37
View File
@@ -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
+131
View File
@@ -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"
+348
View File
@@ -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)
+2 -2
View File
@@ -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 <system_prompt_text>
python3 server_generate.py --data_path Daring-Anteater/train.jsonl --output_path finetune/data.jsonl --max_token 512 --chat --system_prompt <system_prompt_text>
```
#### SLURM Prepare Data
+15 -3
View File
@@ -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,
+10 -1
View File
@@ -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
@@ -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()
@@ -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"
+2 -2
View File
@@ -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
@@ -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.
+13 -4
View File
@@ -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
@@ -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",
@@ -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.
-22
View File
@@ -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."
)
+7 -3
View File
@@ -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]
@@ -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,
)
@@ -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
+7 -4
View File
@@ -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 <op_type>:<precision>, where precision can be fp32 (default) or fp16. "
"For example: op_type_1:fp16 op_type_2:fp32."
"Each item should have the format <op_type>:<precision> (all inputs and outputs have the same precision) "
"or <op_type>:[<inp1_precision>,<inp2_precision>,...]:[<out1_precision>,<out2_precision>,...] "
"(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(
@@ -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.
+44 -32
View File
@@ -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
+48 -35
View File
@@ -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],
+24 -25
View File
@@ -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
+7 -3
View File
@@ -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]}")
+25
View File
@@ -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)
@@ -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
+9 -3
View File
@@ -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.,
+12 -1
View File
@@ -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
+134 -2
View File
@@ -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:
+50 -34
View File
@@ -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>."
)
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
+100
View File
@@ -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 <op_type>:<precision> or"
" <op_type>:[<inp1_precision>,<inp2_precision>,...]:[<out1_precision>,<out2_precision>,...]."
)
# 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
+17 -1
View File
@@ -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]:
+1 -1
View File
@@ -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.",
+7 -1
View File
@@ -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):
+5 -1
View File
@@ -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.
+130 -5
View File
@@ -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(
+20 -4
View File
@@ -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
@@ -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 *
@@ -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,
+309 -6
View File
@@ -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()
+108 -258
View File
@@ -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),
}
+122 -57
View File
@@ -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),
}
+52 -59
View File
@@ -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."),
}
+41 -47
View File
@@ -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."),
}
@@ -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
)
+47
View File
@@ -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
+224 -520
View File
@@ -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(
+42 -4
View File
@@ -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
+136 -115
View File
@@ -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.
+13 -3
View File
@@ -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):
+6 -2
View File
@@ -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(
+3 -3
View File
@@ -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
+10 -2
View File
@@ -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)
+3
View File
@@ -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."""
+3
View File
@@ -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 *
+11 -1
View File
@@ -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__
@@ -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.")
@@ -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
+2 -2
View File
@@ -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]:
@@ -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.',
),
+671
View File
@@ -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)
-1
View File
@@ -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)
+225 -25
View File
@@ -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",
}
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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]:
+1 -1
View File
@@ -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
@@ -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)
@@ -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)
@@ -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)
@@ -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 *
@@ -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
@@ -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)
+49 -15
View File
@@ -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
@@ -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):
+199
View File
@@ -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)
@@ -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 *
@@ -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")
@@ -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)
)
+25 -5
View File
@@ -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
)
@@ -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](
+7
View File
@@ -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."""
@@ -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.
@@ -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))
@@ -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 *

Some files were not shown because too many files have changed in this diff Show More