mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Update files on GitHub
Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
This commit is contained in:
@@ -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
@@ -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.
|
||||
|
||||
@@ -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>`_.
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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}",
|
||||
],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Executable
+131
@@ -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"
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
#### A Library to Quantize and Compress Deep Learning Models for Optimized Inference on Native Windows RTX GPUs
|
||||
|
||||
[](https://nvidia.github.io/TensorRT-Model-Optimizer/)
|
||||
[](https://pypi.org/project/nvidia-modelopt/0.27.0/)
|
||||
[](https://pypi.org/project/nvidia-modelopt/)
|
||||
[](../../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.
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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."
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]}")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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
@@ -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]:
|
||||
|
||||
@@ -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.",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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."),
|
||||
}
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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 *
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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.',
|
||||
),
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,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](
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user