mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Add support for offline speculative decoding model PTQ (#883)
## What does this PR do? **Type of change:** new feature **Overview:** This PR enables loading in a ModelOpt pretrained offline speculative decoding model (e.g., EAGLE3) and performs PTQ on it and export. ## Usage Follow the speculative_decoding examples to train an offline speculative decoding model first. Then follow the command below to quantize and export it: ```bash python hf_ptq.py --pyt_ckpt_path <dir_of_offline_specdec_model> --specdec_offline_dataset <dir_of_dataset> ``` ## Testing <!-- Mention how have you tested your change if applicable. --> ## Before your PR is "*Ready for review*" <!-- If you haven't finished some of the above items you can still open `Draft` PR. --> - **Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)** and your commits are signed. - **Is this change backward compatible?**: Yes/No <!--- If No, explain why. --> - **Did you write any new necessary tests?**: Yes/No - **Did you add or update any necessary documentation?**: Yes/No - **Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**: Yes/No <!--- Only for new features, API changes, critical bug fixes or bw breaking changes. --> ## Additional Information <!-- E.g. related issue. --> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Offline speculative decoding workflow: support loading a local dataset for calibration, generation, and export; new CLI option to specify the offline dataset. * **Improvements** * Export and quantization paths now accept and propagate offline speculative-decoding inputs. * Offline data loading honors a sample-size limit and enforces safe batch sizing for calibration. * **Bug Fixes** * Better handling of model/config mismatches and varied batch types in offline flows. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Ye Yu <yeyu@nvidia.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
8428f0610a
commit
cccfded8a9
+146
-29
@@ -18,6 +18,7 @@ import copy
|
||||
import random
|
||||
import time
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -64,6 +65,10 @@ from modelopt.torch.export.model_utils import get_language_model_from_vl, is_mul
|
||||
from modelopt.torch.quantization.config import _default_disabled_quantizer_cfg, need_calibration
|
||||
from modelopt.torch.quantization.plugins.accelerate import init_quantized_weights
|
||||
from modelopt.torch.quantization.utils import is_quantized
|
||||
from modelopt.torch.speculative.eagle.utils import (
|
||||
EagleOfflineDataCollator,
|
||||
OfflineSupervisedDataset,
|
||||
)
|
||||
from modelopt.torch.utils.dataset_utils import (
|
||||
create_forward_loop,
|
||||
get_dataset_dataloader,
|
||||
@@ -163,6 +168,34 @@ def extract_and_prepare_language_model_from_vl(full_model):
|
||||
return None, None
|
||||
|
||||
|
||||
class _DeviceDataLoader:
|
||||
"""Wrapper around a DataLoader that moves each batch to a target device."""
|
||||
|
||||
def __init__(self, dataloader: DataLoader, device: torch.device):
|
||||
self.dataloader = dataloader
|
||||
self.device = device
|
||||
|
||||
def __iter__(self):
|
||||
for batch in self.dataloader:
|
||||
yield _move_batch_to_device(batch, self.device)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dataloader)
|
||||
|
||||
|
||||
def _move_batch_to_device(batch: dict, device: torch.device) -> dict:
|
||||
"""Recursively move all tensors in a batch dict to the given device."""
|
||||
|
||||
def _to_device(value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
return value.to(device)
|
||||
if isinstance(value, dict):
|
||||
return {k: _to_device(v) for k, v in value.items()}
|
||||
return value
|
||||
|
||||
return {k: _to_device(v) for k, v in batch.items()}
|
||||
|
||||
|
||||
def make_calib_dataloader(
|
||||
args: argparse.Namespace,
|
||||
language_model: torch.nn.Module,
|
||||
@@ -170,10 +203,28 @@ def make_calib_dataloader(
|
||||
tokenizer: PreTrainedTokenizerBase | None,
|
||||
device: torch.device,
|
||||
model_type: str | None,
|
||||
) -> tuple[DataLoader, str | None]:
|
||||
) -> tuple[DataLoader | _DeviceDataLoader, str | None]:
|
||||
calib_dataloader = None
|
||||
first_text_speech_dataset = None
|
||||
if args.calib_with_images:
|
||||
if args.specdec_offline_dataset is not None:
|
||||
offline_data_path = Path(args.specdec_offline_dataset)
|
||||
dumped_files = sorted(str(p) for p in offline_data_path.glob("*.pt"))
|
||||
if not dumped_files:
|
||||
raise ValueError(f"No .pt files found in {args.specdec_offline_dataset}")
|
||||
if args.calib_size[0] > 0:
|
||||
dumped_files = dumped_files[: args.calib_size[0]]
|
||||
dataset = OfflineSupervisedDataset(dumped_files)
|
||||
collator = EagleOfflineDataCollator(train_len=args.calib_seq)
|
||||
raw_loader = DataLoader(
|
||||
dataset,
|
||||
batch_size=args.batch_size,
|
||||
shuffle=False,
|
||||
collate_fn=collator,
|
||||
)
|
||||
# Wrap to move batches to the target device; device-transfer logic is kept
|
||||
# out of the data collator to avoid interference with dataloader prefetching.
|
||||
calib_dataloader = _DeviceDataLoader(raw_loader, device)
|
||||
elif args.calib_with_images:
|
||||
# VLM image-text calibration path: assume Nemotron VLM dataset by default.
|
||||
assert processor is not None, (
|
||||
"Please provide a processor (e.g., AutoProcessor) for image calibration."
|
||||
@@ -358,7 +409,7 @@ def auto_quantize(
|
||||
def load_model(args: argparse.Namespace):
|
||||
# If low memory mode is enabled, we compress the model while loading the HF checkpoint.
|
||||
calibration_only = False
|
||||
if not args.low_memory_mode:
|
||||
if args.specdec_offline_dataset is not None or not args.low_memory_mode:
|
||||
full_model = get_model(
|
||||
args.pyt_ckpt_path,
|
||||
args.device,
|
||||
@@ -459,15 +510,27 @@ def load_model(args: argparse.Namespace):
|
||||
language_model = extracted_lm
|
||||
model_type = extracted_model_type
|
||||
else:
|
||||
if args.dataset is None:
|
||||
args.dataset = ["cnn_dailymail", "nemotron-post-training-dataset-v2"]
|
||||
warnings.warn(
|
||||
"No dataset specified. Defaulting to cnn_dailymail and nemotron-post-training-dataset-v2."
|
||||
if args.specdec_offline_dataset is not None:
|
||||
language_model = full_model
|
||||
else:
|
||||
if args.dataset is None:
|
||||
args.dataset = ["cnn_dailymail", "nemotron-post-training-dataset-v2"]
|
||||
warnings.warn(
|
||||
"No dataset specified. Defaulting to cnn_dailymail and nemotron-post-training-dataset-v2."
|
||||
)
|
||||
# Adjust calib_size to match dataset length by extending or truncating as needed
|
||||
args.calib_size = (args.calib_size + [args.calib_size[-1]] * len(args.dataset))[
|
||||
: len(args.dataset)
|
||||
]
|
||||
|
||||
# We only quantize the language model for VLMs other than the type supported above.
|
||||
extracted_lm, extracted_model_type = extract_and_prepare_language_model_from_vl(
|
||||
full_model
|
||||
)
|
||||
# Adjust calib_size to match dataset length by extending or truncating as needed
|
||||
args.calib_size = (args.calib_size + [args.calib_size[-1]] * len(args.dataset))[
|
||||
: len(args.dataset)
|
||||
]
|
||||
if extracted_lm is not None:
|
||||
language_model = extracted_lm
|
||||
model_type = extracted_model_type
|
||||
|
||||
tokenizer = get_tokenizer(args.pyt_ckpt_path, trust_remote_code=args.trust_remote_code)
|
||||
|
||||
default_padding_side = tokenizer.padding_side
|
||||
@@ -475,12 +538,6 @@ def load_model(args: argparse.Namespace):
|
||||
# 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.
|
||||
extracted_lm, extracted_model_type = extract_and_prepare_language_model_from_vl(full_model)
|
||||
if extracted_lm is not None:
|
||||
language_model = extracted_lm
|
||||
model_type = extracted_model_type
|
||||
|
||||
if model_type == "phi4mm":
|
||||
warnings.warn("Please set the default input_mode to InputMode.LANGUAGE before quantizing.")
|
||||
|
||||
@@ -581,7 +638,12 @@ def mono_quantize(
|
||||
if args.calib_with_images and is_nemotron_vl_model:
|
||||
calibrate_loop = create_vlm_calibration_loop(full_model, calib_dataloader)
|
||||
else:
|
||||
calibrate_loop = create_forward_loop(dataloader=calib_dataloader)
|
||||
calibrate_loop = create_forward_loop(
|
||||
dataloader=calib_dataloader,
|
||||
allowed_non_tensor_keys={"base_model_outputs"}
|
||||
if args.specdec_offline_dataset is not None
|
||||
else None,
|
||||
)
|
||||
|
||||
if calibration_only:
|
||||
language_model = mtq.calibrate(
|
||||
@@ -736,7 +798,7 @@ def pre_quantize(
|
||||
full_model: torch.nn.Module,
|
||||
model_type: str | None,
|
||||
tokenizer: PreTrainedTokenizerBase | None,
|
||||
calib_dataloader: DataLoader,
|
||||
calib_dataloader: DataLoader | None,
|
||||
is_nemotron_vl_model: bool,
|
||||
):
|
||||
"""
|
||||
@@ -746,7 +808,12 @@ def pre_quantize(
|
||||
post-quantize generation.
|
||||
|
||||
"""
|
||||
# Offline specdec models skip pre-quantize preview (no tokenizer or standard dataloader)
|
||||
if args.specdec_offline_dataset is not None:
|
||||
return None, None
|
||||
|
||||
# Only run single sample for preview
|
||||
assert calib_dataloader is not None, "calib_dataloader is required for pre-quantize preview"
|
||||
preview_input_ids = next(iter(calib_dataloader))[
|
||||
"input_features" if model_type == "whisper" else "input_ids"
|
||||
][0:1]
|
||||
@@ -781,6 +848,7 @@ def pre_quantize(
|
||||
def post_quantize(
|
||||
args: argparse.Namespace,
|
||||
full_model: torch.nn.Module,
|
||||
language_model: torch.nn.Module,
|
||||
model_type: str | None,
|
||||
tokenizer: PreTrainedTokenizerBase | None,
|
||||
processor: BaseImageProcessor | ProcessorMixin | None,
|
||||
@@ -788,14 +856,31 @@ def post_quantize(
|
||||
generated_ids_before_ptq,
|
||||
is_nemotron_vl_model,
|
||||
first_text_speech_dataset,
|
||||
default_padding_side,
|
||||
default_pad_token,
|
||||
calib_dataloader: DataLoader,
|
||||
):
|
||||
"""
|
||||
Processing after the quantization.
|
||||
Processing after the quantization, then export.
|
||||
|
||||
Currently we run one round of generation using the quantized model for a sample prompt,
|
||||
and compare it with pre-quantize generation.
|
||||
For offline speculative decoding models, skip generation comparison and proceed
|
||||
directly to export. For standard models, run one round of generation using the
|
||||
quantized model for a sample prompt and compare it with pre-quantize generation.
|
||||
|
||||
"""
|
||||
# Early exit for offline speculative decoding: skip generation comparison and export directly.
|
||||
# The model's get_dummy_inputs() provides the right input format for the export forward pass.
|
||||
if args.specdec_offline_dataset is not None:
|
||||
export_quantized(
|
||||
args,
|
||||
full_model,
|
||||
language_model,
|
||||
model_type,
|
||||
tokenizer,
|
||||
default_padding_side,
|
||||
default_pad_token,
|
||||
)
|
||||
return
|
||||
|
||||
if args.verbose:
|
||||
try:
|
||||
@@ -873,6 +958,16 @@ def post_quantize(
|
||||
f"example outputs after ptq: {output_decode(generated_ids_after_ptq, preview_input_ids.shape[1])}"
|
||||
)
|
||||
|
||||
export_quantized(
|
||||
args,
|
||||
full_model,
|
||||
language_model,
|
||||
model_type,
|
||||
tokenizer,
|
||||
default_padding_side,
|
||||
default_pad_token,
|
||||
)
|
||||
|
||||
|
||||
def quantize_main(
|
||||
args: argparse.Namespace,
|
||||
@@ -892,6 +987,13 @@ def quantize_main(
|
||||
if args.calib_with_images:
|
||||
print("Image-text calibration enabled. Using default batch_size=1 for calibration.")
|
||||
args.batch_size = 1
|
||||
# Speculative decoding offline model dost not support get_max_batch_size() because of
|
||||
# the customized dataloader, so we set batch_size to 1 to avoid OOM.
|
||||
elif args.specdec_offline_dataset is not None:
|
||||
print(
|
||||
"Offline speculative decoding calibration enabled. Using default batch_size=1 for calibration."
|
||||
)
|
||||
args.batch_size = 1
|
||||
else:
|
||||
# Calibration/sparsification will actually take much more memory than regular inference
|
||||
# due to intermediate tensors for fake quantization. Setting sample_memory_usage_ratio
|
||||
@@ -1020,6 +1122,7 @@ def quantize_main(
|
||||
post_quantize(
|
||||
args,
|
||||
full_model,
|
||||
language_model,
|
||||
model_type,
|
||||
tokenizer,
|
||||
processor,
|
||||
@@ -1027,15 +1130,9 @@ def quantize_main(
|
||||
generated_ids_before_ptq,
|
||||
is_nemotron_vl_model,
|
||||
first_text_speech_dataset,
|
||||
)
|
||||
export_quantized(
|
||||
args,
|
||||
full_model,
|
||||
language_model,
|
||||
model_type,
|
||||
tokenizer,
|
||||
default_padding_side,
|
||||
default_pad_token,
|
||||
calib_dataloader,
|
||||
)
|
||||
|
||||
|
||||
@@ -1099,6 +1196,14 @@ def parse_args() -> argparse.Namespace:
|
||||
type=str,
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--specdec_offline_dataset",
|
||||
help=(
|
||||
"If set, the model is a speculative decoding model,"
|
||||
"which uses offline dataset for calibration. "
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib_with_images",
|
||||
action="store_true",
|
||||
@@ -1256,6 +1361,12 @@ def parse_args() -> argparse.Namespace:
|
||||
if args.moe_calib_experts_ratio is not None and not (0.0 < args.moe_calib_experts_ratio <= 1.0):
|
||||
parser.error("--moe_calib_experts_ratio must be in the range (0.0, 1.0].")
|
||||
|
||||
if args.specdec_offline_dataset is not None and args.sparsity_fmt != "dense":
|
||||
parser.error("--specdec_offline_dataset is only supported with --sparsity_fmt dense (PTQ).")
|
||||
|
||||
if args.specdec_offline_dataset is not None and args.low_memory_mode:
|
||||
parser.error("--specdec_offline_dataset is not compatible with --low_memory_mode.")
|
||||
|
||||
return args
|
||||
|
||||
|
||||
@@ -1311,4 +1422,10 @@ if __name__ == "__main__":
|
||||
|
||||
args.dataset = args.dataset.split(",") if isinstance(args.dataset, str) else args.dataset
|
||||
args.calib_size = [int(num_sample) for num_sample in args.calib_size.split(",")]
|
||||
|
||||
if args.specdec_offline_dataset is not None and len(args.calib_size) != 1:
|
||||
raise ValueError(
|
||||
"--specdec_offline_dataset expects a single --calib value, not a comma-separated list."
|
||||
)
|
||||
|
||||
main(args)
|
||||
|
||||
@@ -206,9 +206,10 @@ def main(args: argparse.Namespace) -> None:
|
||||
continue
|
||||
|
||||
# Tokenize and check length
|
||||
input_ids = tokenizer.apply_chat_template(
|
||||
tokenized = tokenizer.apply_chat_template(
|
||||
conversations, return_tensors="pt", add_generation_template=False
|
||||
)["input_ids"]
|
||||
)
|
||||
input_ids = tokenized["input_ids"] if isinstance(tokenized, dict) else tokenized
|
||||
num_input_tokens = input_ids.shape[1]
|
||||
if num_input_tokens <= 10 or num_input_tokens > args.max_seq_len:
|
||||
num_skipped_too_long += 1
|
||||
|
||||
@@ -20,19 +20,19 @@ from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from types import FrameType
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import transformers
|
||||
from datasets import load_dataset
|
||||
from packaging.version import Version
|
||||
from scripts.ar_validate import validate_ar
|
||||
from torch.utils.data import Dataset
|
||||
from transformers import Trainer, TrainerCallback
|
||||
from transformers.trainer_pt_utils import LabelSmoother
|
||||
|
||||
import modelopt
|
||||
from modelopt.torch.speculative.eagle.utils import (
|
||||
EagleOfflineDataCollator,
|
||||
OfflineSupervisedDataset,
|
||||
)
|
||||
from modelopt.torch.speculative.utils import get_ttt_msk_func
|
||||
from modelopt.torch.utils import print_rank_0
|
||||
from modelopt.torch.utils.distributed import is_master
|
||||
@@ -49,92 +49,8 @@ try:
|
||||
except (ImportError, AttributeError):
|
||||
wandb = None
|
||||
|
||||
IGNORE_TOKEN_ID = LabelSmoother.ignore_index
|
||||
|
||||
|
||||
class OfflineSupervisedDataset(Dataset):
|
||||
"""Offline dataset for supervised fine-tuning.
|
||||
|
||||
This dataset loads data on-the-fly from pre-processed .pt data files.
|
||||
|
||||
Args:
|
||||
dumped_files (list): A list of file paths to the dumped .pt files.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dumped_files,
|
||||
):
|
||||
super().__init__()
|
||||
self.dumped_files = dumped_files
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dumped_files)
|
||||
|
||||
def __getitem__(self, i) -> dict[str, torch.Tensor]:
|
||||
try:
|
||||
offline_data = torch.load(self.dumped_files[i])
|
||||
except Exception as e:
|
||||
print(
|
||||
f"[ERROR] Failed to load file at index={i}, "
|
||||
f"path='{self.dumped_files[i]}', error={e}. "
|
||||
"Reusing data from previous index (i-1)."
|
||||
)
|
||||
return self.__getitem__(i - 1)
|
||||
|
||||
labels = torch.full_like(offline_data["input_ids"], IGNORE_TOKEN_ID)
|
||||
labels[..., :-1] = offline_data["input_ids"][..., 1:]
|
||||
|
||||
ret = {
|
||||
"input_ids": offline_data["input_ids"],
|
||||
"base_model_hidden_states": offline_data["hidden_states"],
|
||||
"aux_hidden_states": offline_data["aux_hidden_states"],
|
||||
"attention_mask": torch.ones_like(offline_data["input_ids"]),
|
||||
"loss_mask": torch.ones_like(offline_data["input_ids"]),
|
||||
"labels": labels,
|
||||
}
|
||||
return ret
|
||||
|
||||
|
||||
class EagleOfflineDataCollator:
|
||||
"""Data collator that truncate or pads data for offline training."""
|
||||
|
||||
def __init__(self, train_len):
|
||||
self.train_len = train_len
|
||||
|
||||
def _pad_or_truncate(self, x: torch.Tensor, length: int, dim: int = 0):
|
||||
"""Pad or truncate a tensor to length along a given dimension."""
|
||||
dim = dim % x.ndim # support negative dimension
|
||||
|
||||
# allocate output tensor
|
||||
out_shape = list(x.shape)
|
||||
out_shape[dim] = length
|
||||
out = x.new_zeros(out_shape)
|
||||
|
||||
# consturct copy slice
|
||||
slc = [slice(None)] * x.ndim
|
||||
slc[dim] = slice(0, min(length, x.size(dim)))
|
||||
|
||||
# populate output tensor
|
||||
out[tuple(slc)] = x[tuple(slc)]
|
||||
return out
|
||||
|
||||
def __call__(self, features: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
base_batch = {
|
||||
k: torch.stack([self._pad_or_truncate(item[k], self.train_len) for item in features])
|
||||
for k in ["input_ids", "attention_mask", "loss_mask", "labels"]
|
||||
}
|
||||
|
||||
base_model_outputs = {
|
||||
k: torch.stack([self._pad_or_truncate(item[k], self.train_len) for item in features])
|
||||
for k in ["base_model_hidden_states", "aux_hidden_states"]
|
||||
}
|
||||
|
||||
batch = {
|
||||
**base_batch,
|
||||
"base_model_outputs": base_model_outputs,
|
||||
}
|
||||
return batch
|
||||
# Re-export for backward compatibility
|
||||
__all__ = ["EagleOfflineDataCollator", "OfflineSupervisedDataset"]
|
||||
|
||||
|
||||
def make_eagle_supervised_data_module(
|
||||
@@ -168,6 +84,11 @@ def make_eagle_supervised_data_module(
|
||||
if not dumped_files:
|
||||
raise ValueError(f"No .pt files found in {data_args.offline_data_path}")
|
||||
|
||||
# sample_size=-1 means use all samples; positive integer selects that many
|
||||
if data_args.sample_size == 0 or data_args.sample_size < -1:
|
||||
raise ValueError("sample_size must be -1 (use all samples) or a positive integer")
|
||||
if data_args.sample_size > 0:
|
||||
dumped_files = dumped_files[: data_args.sample_size]
|
||||
train_dataset = OfflineSupervisedDataset(dumped_files)
|
||||
data_collator = EagleOfflineDataCollator(train_len=train_len)
|
||||
|
||||
|
||||
@@ -91,6 +91,14 @@ class DataArguments:
|
||||
)
|
||||
vlm_img_dir: str = field(default=None, metadata={"help": "Path to the VLM image directory."})
|
||||
vlm_processor: str = field(default=None, metadata={"help": "Path to the VLM processor."})
|
||||
sample_size: int = field(
|
||||
default=-1,
|
||||
metadata={"help": "Number of samples to use for training. Use -1 to use all samples."},
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
if self.sample_size == 0 or self.sample_size < -1:
|
||||
raise ValueError("sample_size must be -1 (use all samples) or a positive integer")
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -219,6 +227,13 @@ def train():
|
||||
else:
|
||||
# To avoid OOM for large models, we load and convert model on CPU first.
|
||||
# Model will be moved to GPU during HF trainer.init().
|
||||
if use_offline_training:
|
||||
# Load config first to preserve original num_hidden_layers before
|
||||
# load_vlm_or_llm may reduce layers for offline space savings.
|
||||
model_config = transformers.AutoConfig.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
)
|
||||
model = load_vlm_or_llm(
|
||||
model_args.model_name_or_path,
|
||||
use_fake_base=model_args.use_fake_base_for_offline,
|
||||
@@ -227,6 +242,18 @@ def train():
|
||||
device_map="cpu",
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
)
|
||||
if use_offline_training:
|
||||
# When doing offline training, we need to set num_hidden_layers
|
||||
# since we override it when loading the model for space savings.
|
||||
# Some models (e.g. Kimi-K2.5) use non-standard config attributes,
|
||||
# so fall back to the model's own config if the attribute is missing.
|
||||
model.config.num_orig_hidden_layers = getattr(
|
||||
model_config, "num_hidden_layers", model.config.num_hidden_layers
|
||||
)
|
||||
if hasattr(model.config, "layer_types"):
|
||||
del (
|
||||
model.config.layer_types
|
||||
) # remove layer_types to avoid mismatch with the modified model
|
||||
tokenizer = transformers.AutoTokenizer.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
model_max_length=training_args.training_seq_len,
|
||||
|
||||
@@ -92,7 +92,11 @@ class SpeculativeDecodingExporter(ABC):
|
||||
self.model = model
|
||||
|
||||
@abstractmethod
|
||||
def export(self, export_dir: Path | str, dtype: torch.dtype | None = None):
|
||||
def export(
|
||||
self,
|
||||
export_dir: Path | str,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
"""Export the model to the deployment format."""
|
||||
raise NotImplementedError("Subclasses must implement this method.")
|
||||
|
||||
@@ -185,7 +189,11 @@ class EagleExporter(SpeculativeDecodingExporter):
|
||||
|
||||
return template_config
|
||||
|
||||
def export(self, export_dir: Path | str, dtype: torch.dtype | None = None):
|
||||
def export(
|
||||
self,
|
||||
export_dir: Path | str,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
"""Export the model to the deployment format."""
|
||||
# Make export dir
|
||||
export_dir = Path(export_dir)
|
||||
|
||||
@@ -381,6 +381,9 @@ def requantize_resmooth_fused_llm_layers(model: torch.nn.Module):
|
||||
elif getattr(model.config, "is_encoder_decoder", False):
|
||||
# For other encoder-decoder models (non-VL), pass both encoder and decoder input ids
|
||||
model(fake_input, decoder_input_ids=decoder_fake_input)
|
||||
elif hasattr(model, "get_dummy_inputs"):
|
||||
# For speculative decoding models (EAGLE, etc.), use model-provided dummy inputs
|
||||
model(**model.get_dummy_inputs())
|
||||
else:
|
||||
model(fake_input)
|
||||
|
||||
|
||||
@@ -35,7 +35,13 @@
|
||||
|
||||
"""Eagle model utils."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
from transformers.trainer_pt_utils import LabelSmoother
|
||||
|
||||
IGNORE_TOKEN_ID = LabelSmoother.ignore_index
|
||||
|
||||
|
||||
# Copied from transformers.models.bart.modeling_bart._make_causal_mask
|
||||
@@ -70,3 +76,90 @@ def expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: int | None = No
|
||||
inverted_mask = 1.0 - expanded_mask
|
||||
|
||||
return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)
|
||||
|
||||
|
||||
class OfflineSupervisedDataset(Dataset):
|
||||
"""Offline dataset for supervised fine-tuning with pre-dumped hidden states.
|
||||
|
||||
This dataset loads data on-the-fly from pre-processed .pt data files generated by
|
||||
``examples/speculative_decoding/main.py --mode dump_offline_data``. Each .pt file
|
||||
contains a dict with the following keys:
|
||||
|
||||
- ``input_ids``: token IDs of shape ``(seq_len,)``
|
||||
- ``hidden_states``: base model last hidden states of shape ``(seq_len, hidden_size)``
|
||||
- ``aux_hidden_states``: auxiliary hidden states of shape ``(seq_len, hidden_size)``
|
||||
- ``base_model_input_embeds``: input embeddings of shape ``(seq_len, hidden_size)``
|
||||
|
||||
Args:
|
||||
dumped_files (list): A list of file paths to the dumped .pt files.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dumped_files,
|
||||
):
|
||||
"""Initialize with a list of .pt file paths."""
|
||||
super().__init__()
|
||||
self.dumped_files = dumped_files
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dumped_files)
|
||||
|
||||
def __getitem__(self, i) -> dict[str, torch.Tensor]:
|
||||
offline_data = torch.load(self.dumped_files[i], weights_only=True)
|
||||
|
||||
labels = torch.full_like(offline_data["input_ids"], IGNORE_TOKEN_ID)
|
||||
labels[..., :-1] = offline_data["input_ids"][..., 1:]
|
||||
|
||||
ret = {
|
||||
"input_ids": offline_data["input_ids"],
|
||||
"base_model_hidden_states": offline_data["hidden_states"],
|
||||
"aux_hidden_states": offline_data["aux_hidden_states"],
|
||||
"attention_mask": torch.ones_like(offline_data["input_ids"]),
|
||||
"loss_mask": torch.ones_like(offline_data["input_ids"]),
|
||||
"labels": labels,
|
||||
}
|
||||
return ret
|
||||
|
||||
|
||||
class EagleOfflineDataCollator:
|
||||
"""Data collator that truncates or pads data for offline training."""
|
||||
|
||||
def __init__(self, train_len):
|
||||
"""Initialize with the target sequence length for truncation/padding."""
|
||||
self.train_len = train_len
|
||||
|
||||
def _pad_or_truncate(self, x: torch.Tensor, length: int, dim: int = 0):
|
||||
"""Pad or truncate a tensor to length along a given dimension."""
|
||||
dim = dim % x.ndim # support negative dimension
|
||||
|
||||
# allocate output tensor
|
||||
out_shape = list(x.shape)
|
||||
out_shape[dim] = length
|
||||
out = x.new_zeros(out_shape)
|
||||
|
||||
# construct copy slice
|
||||
slc = [slice(None)] * x.ndim
|
||||
slc[dim] = slice(0, min(length, x.size(dim)))
|
||||
|
||||
# populate output tensor
|
||||
out[tuple(slc)] = x[tuple(slc)]
|
||||
return out
|
||||
|
||||
def __call__(self, features: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
"""Collate a list of feature dicts into a single padded/truncated batch."""
|
||||
base_batch = {
|
||||
k: torch.stack([self._pad_or_truncate(item[k], self.train_len) for item in features])
|
||||
for k in ["input_ids", "attention_mask", "loss_mask", "labels"]
|
||||
}
|
||||
|
||||
base_model_outputs = {
|
||||
k: torch.stack([self._pad_or_truncate(item[k], self.train_len) for item in features])
|
||||
for k in ["base_model_hidden_states", "aux_hidden_states"]
|
||||
}
|
||||
|
||||
batch = {
|
||||
**base_batch,
|
||||
"base_model_outputs": base_model_outputs,
|
||||
}
|
||||
return batch
|
||||
|
||||
@@ -460,6 +460,36 @@ class HFEagleModel(EagleModel):
|
||||
print(f"Failed to create NVTX range {name}: {e}")
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def get_dummy_inputs(self) -> dict:
|
||||
"""Construct dummy inputs for export forward pass.
|
||||
|
||||
Returns a dict of kwargs that can be passed to forward(). For offline EAGLE models,
|
||||
this includes dummy base_model_outputs with the right tensor shapes so the export
|
||||
pipeline doesn't need to thread real calibration data through multiple layers.
|
||||
"""
|
||||
device = self.device
|
||||
dtype = next(self.parameters()).dtype
|
||||
hidden_size = self._base_llm_config.hidden_size
|
||||
dummy_inputs = {
|
||||
"input_ids": torch.ones(1, 2, dtype=torch.long, device=device),
|
||||
}
|
||||
if self.eagle_offline:
|
||||
base_model_outputs = {
|
||||
"base_model_hidden_states": torch.zeros(
|
||||
1, 2, hidden_size, dtype=dtype, device=device
|
||||
),
|
||||
"base_model_input_embeds": torch.zeros(
|
||||
1, 2, hidden_size, dtype=dtype, device=device
|
||||
),
|
||||
}
|
||||
if self.eagle_config.use_aux_hidden_state:
|
||||
num_aux = len(self.eagle_config.eagle_aux_hidden_state_layer_ids)
|
||||
base_model_outputs["aux_hidden_states"] = torch.zeros(
|
||||
1, 2, hidden_size * num_aux, dtype=dtype, device=device
|
||||
)
|
||||
dummy_inputs["base_model_outputs"] = base_model_outputs
|
||||
return dummy_inputs
|
||||
|
||||
def get_exporter(self) -> SpeculativeDecodingExporter:
|
||||
"""Get the exporter for the draft model."""
|
||||
exporter_cls = (
|
||||
|
||||
@@ -508,20 +508,25 @@ def get_max_batch_size(
|
||||
return 512
|
||||
|
||||
|
||||
def _process_batch(batch_data, infer_method, max_working_batch_size=None):
|
||||
def _process_batch(
|
||||
batch_data, infer_method, max_working_batch_size=None, allowed_non_tensor_keys=None
|
||||
):
|
||||
"""Process a batch of data through the model's inference method.
|
||||
|
||||
Args:
|
||||
batch_data: Dictionary containing the batch data
|
||||
infer_method: Model's inference method (either forward or generate)
|
||||
max_working_batch_size: Maximum batch size known to work without OOM
|
||||
allowed_non_tensor_keys: Set of key names whose values may be non-tensor types
|
||||
|
||||
Returns:
|
||||
The maximum batch size that worked successfully
|
||||
"""
|
||||
assert all(torch.is_tensor(data) or data is None for data in batch_data.values()), (
|
||||
"batch_data values must be tensors"
|
||||
)
|
||||
allowed_non_tensor_keys = allowed_non_tensor_keys or set()
|
||||
assert all(
|
||||
torch.is_tensor(data) or data is None or key in allowed_non_tensor_keys
|
||||
for key, data in batch_data.items()
|
||||
), f"batch_data values must be tensors or None, except for keys: {allowed_non_tensor_keys}."
|
||||
# Get the batch size of current data
|
||||
batch_size = batch_data[next(iter(batch_data.keys()))].shape[0]
|
||||
|
||||
@@ -538,7 +543,7 @@ def _process_batch(batch_data, infer_method, max_working_batch_size=None):
|
||||
split_data[key] = batch_data[key][i:end_idx, ...]
|
||||
|
||||
max_working_batch_size = _process_batch(
|
||||
split_data, infer_method, max_working_batch_size
|
||||
split_data, infer_method, max_working_batch_size, allowed_non_tensor_keys
|
||||
)
|
||||
|
||||
return max_working_batch_size
|
||||
@@ -566,19 +571,28 @@ def _process_batch(batch_data, infer_method, max_working_batch_size=None):
|
||||
split_data_2 = {key: batch_data[key][mid:, ...] for key in batch_data}
|
||||
|
||||
# Recursively process each half and track max working batch size
|
||||
max_working_batch_size = _process_batch(split_data_1, infer_method)
|
||||
max_working_batch_size = _process_batch(split_data_2, infer_method, max_working_batch_size)
|
||||
max_working_batch_size = _process_batch(
|
||||
split_data_1, infer_method, allowed_non_tensor_keys=allowed_non_tensor_keys
|
||||
)
|
||||
max_working_batch_size = _process_batch(
|
||||
split_data_2, infer_method, max_working_batch_size, allowed_non_tensor_keys
|
||||
)
|
||||
|
||||
# Return the minimum of the two (to be conservative)
|
||||
return max_working_batch_size
|
||||
|
||||
|
||||
def _forward_loop(model: torch.nn.Module, dataloader: DataLoader) -> None:
|
||||
def _forward_loop(
|
||||
model: torch.nn.Module,
|
||||
dataloader: DataLoader,
|
||||
allowed_non_tensor_keys: set | None = None,
|
||||
) -> None:
|
||||
"""Runs forward passes through the model using data from the dataloader.
|
||||
|
||||
Args:
|
||||
model: The PyTorch model to run inference on
|
||||
dataloader: DataLoader containing the batched input data
|
||||
allowed_non_tensor_keys: Set of key names whose values may be non-tensor types
|
||||
"""
|
||||
with torch.no_grad():
|
||||
is_enc_dec = model_type_is_enc_dec(model)
|
||||
@@ -587,7 +601,9 @@ def _forward_loop(model: torch.nn.Module, dataloader: DataLoader) -> None:
|
||||
|
||||
for _, data in enumerate(tqdm(dataloader)):
|
||||
# Process batch and update max working batch size
|
||||
max_working_batch_size = _process_batch(data, infer_method, max_working_batch_size)
|
||||
max_working_batch_size = _process_batch(
|
||||
data, infer_method, max_working_batch_size, allowed_non_tensor_keys
|
||||
)
|
||||
|
||||
|
||||
def create_forward_loop(
|
||||
@@ -600,6 +616,7 @@ def create_forward_loop(
|
||||
device: str | None = None,
|
||||
include_labels: bool = False,
|
||||
dataloader: DataLoader | None = None,
|
||||
allowed_non_tensor_keys: set | None = None,
|
||||
) -> Callable:
|
||||
"""Creates and returns a forward loop function configured for a specific model, dataset, and tokenizer.
|
||||
|
||||
@@ -618,6 +635,9 @@ def create_forward_loop(
|
||||
device: Target device for the returned dataloader.
|
||||
include_labels: Whether to include labels in the dataloader.
|
||||
dataloader: If provided, use the provided dataloader instead.
|
||||
allowed_non_tensor_keys: Set of key names whose batch values may be non-tensor types.
|
||||
Useful when the dataloader yields batches with non-standard fields (e.g., nested
|
||||
model outputs).
|
||||
|
||||
Example usage for quantization:
|
||||
|
||||
@@ -657,7 +677,7 @@ def create_forward_loop(
|
||||
include_labels=include_labels,
|
||||
)
|
||||
|
||||
return lambda model: _forward_loop(model, dataloader)
|
||||
return lambda model: _forward_loop(model, dataloader, allowed_non_tensor_keys)
|
||||
|
||||
|
||||
def model_type_is_enc_dec(model):
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
# 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.
|
||||
|
||||
"""End-to-end CI test for the offline speculative decoding PTQ workflow.
|
||||
|
||||
Covers the three-stage pipeline:
|
||||
1. Collect hidden states from the base model → .pt files
|
||||
2. Train an offline EAGLE draft model → ModelOpt checkpoint
|
||||
3. PTQ the offline checkpoint → quantized export
|
||||
|
||||
Running all three stages in sequence validates that the data format produced
|
||||
by stage 1 is correctly consumed by stage 2 and that the checkpoint produced
|
||||
by stage 2 is correctly quantized in stage 3.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import safetensors.torch
|
||||
import torch
|
||||
from _test_utils.examples.run_command import MODELOPT_ROOT, run_example_command
|
||||
|
||||
EAGLE3_YAML = str(
|
||||
MODELOPT_ROOT / "modelopt_recipes" / "general" / "speculative_decoding" / "eagle3.yaml"
|
||||
)
|
||||
|
||||
# Tiny EAGLE architecture overrides (dotlist entries)
|
||||
_TINY_EAGLE_ARCH = [
|
||||
"eagle.eagle_architecture_config.max_position_embeddings=128",
|
||||
"eagle.eagle_architecture_config.num_hidden_layers=1",
|
||||
"eagle.eagle_architecture_config.intermediate_size=64",
|
||||
"eagle.eagle_architecture_config.num_attention_heads=2",
|
||||
"eagle.eagle_architecture_config.num_key_value_heads=2",
|
||||
"eagle.eagle_architecture_config.head_dim=64",
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def offline_ptq_dirs(tmp_path_factory):
|
||||
"""Shared output directories for all stages."""
|
||||
return {
|
||||
"hidden_states": tmp_path_factory.mktemp("hidden_states"),
|
||||
"eagle_ckpt": tmp_path_factory.mktemp("eagle_ckpt"),
|
||||
"ptq_export": tmp_path_factory.mktemp("ptq_export"),
|
||||
}
|
||||
|
||||
|
||||
def test_collect_hidden_states(tiny_llama_path, tiny_daring_anteater_path, offline_ptq_dirs):
|
||||
"""Stage 1: generate .pt hidden state files from the base model."""
|
||||
run_example_command(
|
||||
[
|
||||
"python",
|
||||
"collect_hidden_states/compute_hidden_states_hf.py",
|
||||
"--model",
|
||||
tiny_llama_path,
|
||||
"--input-data",
|
||||
str(tiny_daring_anteater_path),
|
||||
"--output-dir",
|
||||
str(offline_ptq_dirs["hidden_states"]),
|
||||
"--debug-max-num-conversations",
|
||||
"2",
|
||||
],
|
||||
"speculative_decoding",
|
||||
)
|
||||
|
||||
pt_files = list(offline_ptq_dirs["hidden_states"].glob("*.pt"))
|
||||
assert len(pt_files) > 0, "No .pt files generated by compute_hidden_states_hf.py"
|
||||
|
||||
# Validate the format expected by OfflineSupervisedDataset
|
||||
sample = torch.load(str(pt_files[0]))
|
||||
assert "input_ids" in sample, "Missing 'input_ids' in .pt file"
|
||||
assert "hidden_states" in sample, "Missing 'hidden_states' in .pt file"
|
||||
|
||||
|
||||
def test_offline_eagle_training(tiny_llama_path, tiny_daring_anteater_path, offline_ptq_dirs):
|
||||
"""Stage 2: train an EAGLE3 draft model using the offline hidden states."""
|
||||
output_dir = offline_ptq_dirs["eagle_ckpt"] / "trained"
|
||||
|
||||
overrides = [
|
||||
f"model.model_name_or_path={tiny_llama_path}",
|
||||
f"data.data_path={tiny_daring_anteater_path}",
|
||||
f"data.offline_data_path={offline_ptq_dirs['hidden_states']}",
|
||||
f"training.output_dir={output_dir}",
|
||||
"training.num_train_epochs=1",
|
||||
"training.learning_rate=1e-5",
|
||||
"training.training_seq_len=64",
|
||||
"training.save_steps=1",
|
||||
*_TINY_EAGLE_ARCH,
|
||||
]
|
||||
|
||||
run_example_command(
|
||||
["./launch_train.sh", "--config", EAGLE3_YAML, *overrides],
|
||||
"speculative_decoding",
|
||||
setup_free_port=True,
|
||||
)
|
||||
|
||||
assert output_dir.exists(), "EAGLE training did not produce an output directory"
|
||||
|
||||
|
||||
def test_offline_ptq(offline_ptq_dirs):
|
||||
"""Stage 3: run PTQ on the offline EAGLE checkpoint using the hidden state dataset."""
|
||||
run_example_command(
|
||||
[
|
||||
"python",
|
||||
"hf_ptq.py",
|
||||
"--pyt_ckpt_path",
|
||||
str(offline_ptq_dirs["eagle_ckpt"] / "trained"),
|
||||
"--qformat",
|
||||
"fp8",
|
||||
"--calib_size",
|
||||
"2",
|
||||
"--batch_size",
|
||||
"1",
|
||||
"--specdec_offline_dataset",
|
||||
str(offline_ptq_dirs["hidden_states"]),
|
||||
"--export_path",
|
||||
str(offline_ptq_dirs["ptq_export"]),
|
||||
],
|
||||
"llm_ptq",
|
||||
)
|
||||
|
||||
# Verify the exported checkpoint exists and has the expected EAGLE keys
|
||||
export_dir = offline_ptq_dirs["ptq_export"]
|
||||
assert (export_dir / "model.safetensors").exists(), "PTQ export missing model.safetensors"
|
||||
assert (export_dir / "config.json").exists(), "PTQ export missing config.json"
|
||||
|
||||
from modelopt.torch.export.plugins.hf_spec_export import LLAMA_EAGLE_SINGLE_LAYER
|
||||
|
||||
state_dict = safetensors.torch.load_file(export_dir / "model.safetensors")
|
||||
for key in LLAMA_EAGLE_SINGLE_LAYER["required"] - {"fc", "layers.0.hidden_norm"}:
|
||||
assert f"{key}.weight" in state_dict, f"Missing key '{key}.weight' in exported state dict"
|
||||
@@ -0,0 +1,272 @@
|
||||
# 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.
|
||||
|
||||
"""Unit tests for offline speculative decoding PTQ support."""
|
||||
|
||||
import argparse
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Load eagle_utils from examples/ via importlib (not a package, so no import).
|
||||
# eagle_utils has a top-level `from scripts.ar_validate import validate_ar` that
|
||||
# only resolves when run from examples/speculative_decoding/. We stub it out here.
|
||||
# ---------------------------------------------------------------------------
|
||||
import sys
|
||||
import types
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from _test_utils.torch.transformers_models import get_tiny_llama
|
||||
|
||||
import modelopt.torch.speculative as mtsp
|
||||
from modelopt.torch.speculative.eagle.default_config import default_eagle_config
|
||||
from modelopt.torch.speculative.eagle.utils import (
|
||||
EagleOfflineDataCollator,
|
||||
OfflineSupervisedDataset,
|
||||
)
|
||||
|
||||
_mock_scripts = types.ModuleType("scripts")
|
||||
_mock_ar = types.ModuleType("scripts.ar_validate")
|
||||
_mock_ar.validate_ar = lambda *args, **kwargs: None # type: ignore[attr-defined]
|
||||
sys.modules.setdefault("scripts", _mock_scripts)
|
||||
sys.modules.setdefault("scripts.ar_validate", _mock_ar)
|
||||
|
||||
_EAGLE_UTILS_PATH = os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"../../../../..",
|
||||
"examples/speculative_decoding/eagle_utils.py",
|
||||
)
|
||||
_spec = importlib.util.spec_from_file_location("eagle_utils", _EAGLE_UTILS_PATH)
|
||||
assert _spec is not None and _spec.loader is not None
|
||||
_eagle_utils = importlib.util.module_from_spec(_spec)
|
||||
_spec.loader.exec_module(_eagle_utils)
|
||||
make_eagle_supervised_data_module = _eagle_utils.make_eagle_supervised_data_module
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sample_size truncation tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_data_args(sample_size, tmp_path, n_files=5):
|
||||
"""Create a temp dir with n_files dummy .pt files and an argparse.Namespace."""
|
||||
for i in range(n_files):
|
||||
torch.save({}, tmp_path / f"sample_{i}.pt")
|
||||
return argparse.Namespace(
|
||||
vlm_processor=None,
|
||||
vlm_img_dir=None,
|
||||
offline_data_path=str(tmp_path),
|
||||
lazy_preprocess=True,
|
||||
sample_size=sample_size,
|
||||
)
|
||||
|
||||
|
||||
def test_sample_size_positive_truncates(tmp_path):
|
||||
"""sample_size > 0 should truncate the dataset to that many samples."""
|
||||
data_args = _make_data_args(sample_size=3, tmp_path=tmp_path, n_files=5)
|
||||
tokenizer = MagicMock()
|
||||
module = make_eagle_supervised_data_module(tokenizer, data_args, train_len=8)
|
||||
assert len(module["train_dataset"]) == 3
|
||||
|
||||
|
||||
def test_sample_size_minus_one_uses_all(tmp_path):
|
||||
"""sample_size=-1 should use all samples."""
|
||||
data_args = _make_data_args(sample_size=-1, tmp_path=tmp_path, n_files=5)
|
||||
tokenizer = MagicMock()
|
||||
module = make_eagle_supervised_data_module(tokenizer, data_args, train_len=8)
|
||||
assert len(module["train_dataset"]) == 5
|
||||
|
||||
|
||||
def test_sample_size_zero_raises(tmp_path):
|
||||
"""sample_size=0 should raise ValueError."""
|
||||
data_args = _make_data_args(sample_size=0, tmp_path=tmp_path, n_files=5)
|
||||
tokenizer = MagicMock()
|
||||
with pytest.raises(ValueError, match="sample_size must be -1"):
|
||||
make_eagle_supervised_data_module(tokenizer, data_args, train_len=8)
|
||||
|
||||
|
||||
def test_sample_size_larger_than_dataset_uses_all(tmp_path):
|
||||
"""sample_size > number of files should use all samples without error."""
|
||||
data_args = _make_data_args(sample_size=100, tmp_path=tmp_path, n_files=5)
|
||||
tokenizer = MagicMock()
|
||||
module = make_eagle_supervised_data_module(tokenizer, data_args, train_len=8)
|
||||
assert len(module["train_dataset"]) == 5
|
||||
|
||||
|
||||
def test_sample_size_no_pt_files_raises(tmp_path):
|
||||
"""Empty directory should raise ValueError."""
|
||||
data_args = argparse.Namespace(
|
||||
vlm_processor=None,
|
||||
vlm_img_dir=None,
|
||||
offline_data_path=str(tmp_path),
|
||||
lazy_preprocess=True,
|
||||
sample_size=-1,
|
||||
)
|
||||
tokenizer = MagicMock()
|
||||
with pytest.raises(ValueError, match="No .pt files found"):
|
||||
make_eagle_supervised_data_module(tokenizer, data_args, train_len=8)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_dummy_inputs() for export forward pass
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
TINY_EAGLE_ARCH_CFG = {
|
||||
"num_hidden_layers": 1,
|
||||
"intermediate_size": 32,
|
||||
"num_attention_heads": 16,
|
||||
"num_key_value_heads": 16,
|
||||
"head_dim": 2,
|
||||
"use_last_layernorm": True,
|
||||
"use_aux_hidden_state": False,
|
||||
"eagle_aux_hidden_state_layer_ids": [],
|
||||
}
|
||||
|
||||
TINY_EAGLE_MODE_CFG = {
|
||||
"eagle_architecture_config": {**default_eagle_config, **TINY_EAGLE_ARCH_CFG},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def eagle_model():
|
||||
model = get_tiny_llama(num_hidden_layers=4)
|
||||
mtsp.convert(model, mode=[("eagle", TINY_EAGLE_MODE_CFG)])
|
||||
return model
|
||||
|
||||
|
||||
def test_get_dummy_inputs_online(eagle_model):
|
||||
"""Online EAGLE model returns input_ids only (no base_model_outputs)."""
|
||||
eagle_model.eagle_offline = False
|
||||
dummy = eagle_model.get_dummy_inputs()
|
||||
assert "input_ids" in dummy
|
||||
assert "base_model_outputs" not in dummy
|
||||
|
||||
|
||||
def test_get_dummy_inputs_offline(eagle_model):
|
||||
"""Offline EAGLE model returns input_ids and base_model_outputs with correct shapes."""
|
||||
eagle_model.eagle_offline = True
|
||||
dummy = eagle_model.get_dummy_inputs()
|
||||
assert "input_ids" in dummy
|
||||
assert "base_model_outputs" in dummy
|
||||
hidden_size = eagle_model.config.hidden_size
|
||||
assert dummy["base_model_outputs"]["base_model_hidden_states"].shape[-1] == hidden_size
|
||||
assert dummy["base_model_outputs"]["base_model_input_embeds"].shape[-1] == hidden_size
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OfflineSupervisedDataset tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SEQ_LEN = 16
|
||||
HIDDEN_SIZE = 8
|
||||
|
||||
|
||||
def _make_offline_pt(path, seq_len=SEQ_LEN, hidden_size=HIDDEN_SIZE):
|
||||
"""Write a realistic .pt file matching the format expected by OfflineSupervisedDataset."""
|
||||
data = {
|
||||
"input_ids": torch.randint(0, 100, (seq_len,)),
|
||||
"hidden_states": torch.randn(seq_len, hidden_size),
|
||||
"aux_hidden_states": torch.randn(seq_len, hidden_size),
|
||||
"base_model_input_embeds": torch.randn(seq_len, hidden_size),
|
||||
}
|
||||
torch.save(data, path)
|
||||
return data
|
||||
|
||||
|
||||
def test_offline_dataset_len_and_getitem(tmp_path):
|
||||
"""OfflineSupervisedDataset should load .pt files and return proper keys."""
|
||||
n = 3
|
||||
files = []
|
||||
for i in range(n):
|
||||
p = tmp_path / f"sample_{i}.pt"
|
||||
_make_offline_pt(p)
|
||||
files.append(str(p))
|
||||
|
||||
ds = OfflineSupervisedDataset(files)
|
||||
assert len(ds) == n
|
||||
|
||||
item = ds[0]
|
||||
assert set(item.keys()) == {
|
||||
"input_ids",
|
||||
"base_model_hidden_states",
|
||||
"aux_hidden_states",
|
||||
"attention_mask",
|
||||
"loss_mask",
|
||||
"labels",
|
||||
}
|
||||
assert item["input_ids"].shape == (SEQ_LEN,)
|
||||
assert item["attention_mask"].shape == (SEQ_LEN,)
|
||||
assert item["labels"].shape == (SEQ_LEN,)
|
||||
|
||||
|
||||
def test_offline_dataset_labels_shift(tmp_path):
|
||||
"""Labels should be input_ids shifted left by 1."""
|
||||
p = tmp_path / "sample.pt"
|
||||
orig = _make_offline_pt(p)
|
||||
ds = OfflineSupervisedDataset([str(p)])
|
||||
item = ds[0]
|
||||
# labels[:-1] should equal input_ids[1:]
|
||||
assert torch.equal(item["labels"][:-1], orig["input_ids"][1:])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EagleOfflineDataCollator tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_collator_truncates(tmp_path):
|
||||
"""Collator should truncate sequences longer than train_len."""
|
||||
train_len = 8
|
||||
p = tmp_path / "sample.pt"
|
||||
_make_offline_pt(p, seq_len=SEQ_LEN) # SEQ_LEN > train_len
|
||||
ds = OfflineSupervisedDataset([str(p)])
|
||||
collator = EagleOfflineDataCollator(train_len=train_len)
|
||||
batch = collator([ds[0]])
|
||||
assert batch["input_ids"].shape == (1, train_len)
|
||||
assert batch["base_model_outputs"]["base_model_hidden_states"].shape[1] == train_len
|
||||
|
||||
|
||||
def test_collator_pads(tmp_path):
|
||||
"""Collator should pad sequences shorter than train_len."""
|
||||
train_len = 32
|
||||
p = tmp_path / "sample.pt"
|
||||
_make_offline_pt(p, seq_len=SEQ_LEN) # SEQ_LEN < train_len
|
||||
ds = OfflineSupervisedDataset([str(p)])
|
||||
collator = EagleOfflineDataCollator(train_len=train_len)
|
||||
batch = collator([ds[0]])
|
||||
assert batch["input_ids"].shape == (1, train_len)
|
||||
# Padded region should be zeros
|
||||
assert (batch["input_ids"][0, SEQ_LEN:] == 0).all()
|
||||
|
||||
|
||||
def test_collator_batches_multiple(tmp_path):
|
||||
"""Collator should stack multiple samples into a batch."""
|
||||
train_len = SEQ_LEN
|
||||
files = []
|
||||
for i in range(4):
|
||||
p = tmp_path / f"sample_{i}.pt"
|
||||
_make_offline_pt(p)
|
||||
files.append(str(p))
|
||||
ds = OfflineSupervisedDataset(files)
|
||||
collator = EagleOfflineDataCollator(train_len=train_len)
|
||||
batch = collator([ds[i] for i in range(4)])
|
||||
assert batch["input_ids"].shape == (4, train_len)
|
||||
assert batch["base_model_outputs"]["base_model_hidden_states"].shape == (
|
||||
4,
|
||||
train_len,
|
||||
HIDDEN_SIZE,
|
||||
)
|
||||
@@ -103,6 +103,48 @@ def test_batch_contents_preserved():
|
||||
assert processed_values == [0, 1, 2, 3]
|
||||
|
||||
|
||||
def test_process_batch_allowed_non_tensor_keys_accepted():
|
||||
"""Non-tensor values under allowed_non_tensor_keys should not raise."""
|
||||
batch_data = {
|
||||
"input_ids": torch.ones((2, 8), dtype=torch.long),
|
||||
"base_model_outputs": [{"hidden_states": torch.zeros(2, 8, 16)}], # non-tensor
|
||||
}
|
||||
|
||||
def mock_infer(**kwargs):
|
||||
pass
|
||||
|
||||
# Should not raise
|
||||
_process_batch(batch_data, mock_infer, allowed_non_tensor_keys={"base_model_outputs"})
|
||||
|
||||
|
||||
def test_process_batch_non_tensor_without_allowlist_raises():
|
||||
"""Non-tensor values without allowlist should raise AssertionError."""
|
||||
batch_data = {
|
||||
"input_ids": torch.ones((2, 8), dtype=torch.long),
|
||||
"base_model_outputs": [{"hidden_states": torch.zeros(2, 8, 16)}],
|
||||
}
|
||||
|
||||
def mock_infer(**kwargs):
|
||||
pass
|
||||
|
||||
with pytest.raises(AssertionError):
|
||||
_process_batch(batch_data, mock_infer)
|
||||
|
||||
|
||||
def test_process_batch_other_keys_still_validated():
|
||||
"""Non-tensor values under non-allowed keys should still raise even with allowlist set."""
|
||||
batch_data = {
|
||||
"input_ids": torch.ones((2, 8), dtype=torch.long),
|
||||
"unexpected_key": "some_string", # not in allowed list
|
||||
}
|
||||
|
||||
def mock_infer(**kwargs):
|
||||
pass
|
||||
|
||||
with pytest.raises(AssertionError):
|
||||
_process_batch(batch_data, mock_infer, allowed_non_tensor_keys={"base_model_outputs"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("test_local_path", [True, False])
|
||||
def test_get_dataset_samples_with_unsupported_minipile_dataset(tmp_path, test_local_path):
|
||||
pytest.importorskip("datasets")
|
||||
|
||||
Executable
+27
@@ -0,0 +1,27 @@
|
||||
#!/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.
|
||||
|
||||
SCRIPT_DIR="$(dirname "$(readlink -f "$0")")"
|
||||
source ${SCRIPT_DIR}/../service_utils.sh
|
||||
|
||||
trap 'error_handler $0 $LINENO' ERR
|
||||
|
||||
###################################################################################################
|
||||
|
||||
python modules/Model-Optimizer/examples/llm_ptq/hf_ptq.py \
|
||||
--model ${HF_MODEL_CKPT} \
|
||||
${@}
|
||||
@@ -0,0 +1,128 @@
|
||||
# EAGLE3 offline speculative decoding pipeline with PTQ for Qwen3-8B.
|
||||
#
|
||||
# 5-step pipeline:
|
||||
# task_0: Data synthesis — query TRT-LLM server to generate prompt samples
|
||||
# task_1: Dump hidden states — run target model to capture hidden states
|
||||
# task_2: Offline training — train the EAGLE3 draft head and export checkpoint
|
||||
# task_3: PTQ — quantize the EAGLE3 model using offline hidden states
|
||||
# task_4: Benchmark — evaluate speculative decoding speedup via VLLM
|
||||
#
|
||||
# All tasks share /scratchspace to pass artifacts between steps.
|
||||
#
|
||||
# Usage:
|
||||
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_offline_eagle3_ptq.yaml --yes
|
||||
# uv run slurm.py --yaml modules/Model-Optimizer/tools/launcher/examples/Qwen/Qwen3-8B/hf_offline_eagle3_ptq.yaml --yes
|
||||
|
||||
job_name: Qwen3-8B_EAGLE3_offline_PTQ
|
||||
pipeline:
|
||||
allow_to_fail: false
|
||||
skip: false
|
||||
note:
|
||||
|
||||
global_vars:
|
||||
hf_model: /hf-local/Qwen/Qwen3-8B
|
||||
|
||||
# Step 1: Data synthesis via TRT-LLM server
|
||||
# Args before "--" go to trtllm-serve; args after "--" go to tools/query.py.
|
||||
task_0:
|
||||
script: common/tensorrt_llm/query.sh
|
||||
args:
|
||||
- --model <<global_vars.hf_model>>
|
||||
- --tp_size 4
|
||||
- --ep_size 4
|
||||
- --max_num_tokens 32000
|
||||
- --port 8000
|
||||
- --host 0.0.0.0
|
||||
- --trust_remote_code
|
||||
- --
|
||||
- --data /hf-local/modelopt/Speculative-Decoding-Prompt-Samples
|
||||
- --save /scratchspace/data
|
||||
environment:
|
||||
- HF_LOCAL: /hf-local
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
container: nvcr.io/nvidia/tensorrt-llm/release:1.2.0
|
||||
|
||||
# Step 2: Dump hidden states from target model
|
||||
task_1:
|
||||
script: common/eagle3/dump_offline_data.sh
|
||||
args:
|
||||
- --input-data /scratchspace/data
|
||||
- --output-dir /scratchspace/offline_hidden_states
|
||||
- --max-seq-len 8192
|
||||
- --tp 4
|
||||
- --moe-ep 4
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 4
|
||||
gpus_per_node: 4
|
||||
container: nvcr.io/nvidia/tensorrt-llm/release:1.2.0
|
||||
|
||||
# Step 3: Train EAGLE3 draft head (offline, single task) and export checkpoint
|
||||
task_2:
|
||||
script: common/eagle3/offline_training.sh
|
||||
args:
|
||||
- --offline-data /scratchspace/offline_hidden_states
|
||||
- --data_path None
|
||||
- --mode eagle3
|
||||
- --num_epochs 1
|
||||
- --lr 3e-4
|
||||
- --save_steps 500000
|
||||
- --output_dir /scratchspace/eagle3
|
||||
- --train_bs 8
|
||||
- --training_seq_len 4096
|
||||
- --eagle_config modules/Model-Optimizer/examples/speculative_decoding/eagle_config.json
|
||||
- --disable_tqdm True
|
||||
- --ar_validate_steps 500000
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 1
|
||||
gpus_per_node: 4
|
||||
container: nvcr.io/nvidia/tensorrt-llm/release:1.2.0
|
||||
|
||||
# Step 4: PTQ — quantize the EAGLE3 model using offline hidden states
|
||||
task_3:
|
||||
script: common/eagle3/hf_ptq.sh
|
||||
args:
|
||||
- --specdec_offline_dataset /scratchspace/offline_hidden_states
|
||||
- --qformat fp8
|
||||
- --calib_size 512
|
||||
- --export_path /scratchspace/export_quantized
|
||||
environment:
|
||||
- HF_MODEL_CKPT: /scratchspace/export
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 1
|
||||
gpus_per_node: 1
|
||||
|
||||
# Step 5: Benchmark speculative decoding (VLLM backend)
|
||||
task_4:
|
||||
script: common/specdec_bench/quick_check.sh
|
||||
args:
|
||||
- --draft_model_dir /scratchspace/export_quantized
|
||||
- --draft_length 3
|
||||
- --output_length 4096
|
||||
- --engine VLLM
|
||||
- --tp_size 4
|
||||
- --ep_size 1
|
||||
- --speculative_algorithm EAGLE3
|
||||
- --mtbench /hf-local/HuggingFaceH4/mt_bench_prompts/raw/question.jsonl
|
||||
- --concurrency 1
|
||||
environment:
|
||||
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
||||
slurm_config:
|
||||
_factory_: "slurm_factory"
|
||||
nodes: 1
|
||||
ntasks_per_node: 1
|
||||
gpus_per_node: 4
|
||||
container: vllm/vllm-openai:latest
|
||||
Reference in New Issue
Block a user