diff --git a/examples/ONLINE_TRAINING.md b/examples/ONLINE_TRAINING.md new file mode 100644 index 000000000..1014dd23b --- /dev/null +++ b/examples/ONLINE_TRAINING.md @@ -0,0 +1,62 @@ +# Online training + +This readme walks through the process of online training an Eagle3 draft model. + +## Prepare data + +In a python environment with `speculators` installed, prepare the training dataset. Pass in the target model name/path, dataset name/path (you can pass in multiple datasets), and the output directory. + +``` +python scripts/prepare_data.py --model Qwen/Qwen3-8B --data sharegpt --output ./output +``` + +**Produces:** + +``` +./output/ + data-00000-of-00002.arrow # ⎤ + data-00001-of-00002.arrow # | Processed dataset on disk + dataset_info.json # | + state.json # ⎦ + + token_freq.pt # Token frequencies for vocab mapping +``` + +## Launch vLLM + +In a python environment with `vllm` installed, launch a vllm server configured for hidden states extraction. We provide a wrapper script (`scripts/launch_vllm.py`) to make this easier. + +``` +CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/launch_vllm.py Qwen/Qwen3-8B -- --data-parallel-size 4 --port 8000 +``` + +Note: anything that comes after the `--` will be passed directly to vllm. The `--data-parallel-size` and `--port` are examples of optional arguments for configuring vLLM. `--tensor-parallel-size` also works as expected. + +**Produces:** Model ready to serve requests on on port 8000 + +## Run training + +In a python environment with `speculators` installed, launch the training process. `torchrun` (and the arguments to it) are used to launch a multi-gpu training job. These can be omitted if training on a single gpu. + +``` +CUDA_VISIBLE_DEVICES=4,5,6,7 torchrun --standalone --nproc_per_node 4 scripts/train.py --verifier-name-or-path Qwen/Qwen3-8B --data-path ./output --vllm-endpoint http://localhost:8000/v1 --save-path ./output/checkpoint --draft-model-size 32000 +``` + +**Produces:** If `--draft-model-size` is set, vocab mappings will be generated and cached to the `--data-path` directory. + +``` +./output/ + data-00000-of-00002.arrow # ⎤ + data-00001-of-00002.arro # | + dataset_info.json # | From `scripts/prepare_data.py` step + state.json # | + token_freq.pt # ⎦ + + td2.npy # ⎤ Vocab mappings + d2t.npy # ⎦ + + checkpoints/ # Training checkpoints (loadable by vLLM) + 0/ + 1/ + ... +``` diff --git a/pyproject.toml b/pyproject.toml index d6d641159..54c2f710a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -41,6 +41,7 @@ dependencies = [ "huggingface-hub", "loguru>=0.7.2,<=0.7.3", "numpy>=2.0.0,<=2.4.2", + "openai>=2.0.0", "protobuf", "psutil", "pydantic>=2.0.0", diff --git a/scripts/build_vocab_mapping.py b/scripts/build_vocab_mapping.py index 93baf0f1e..ff134334e 100644 --- a/scripts/build_vocab_mapping.py +++ b/scripts/build_vocab_mapping.py @@ -20,10 +20,10 @@ import numpy as np import torch -from transformers import AutoConfig from speculators.train.vocab_mapping import ( build_vocab_mappings_from_distribution, + get_target_vocab_size, ) logging.basicConfig( @@ -71,29 +71,6 @@ def parse_args(): return parser.parse_args() -def get_target_vocab_size(args): - has_vocab = args.target_vocab_size is not None - has_model = args.target_model_path is not None - - if has_vocab and has_model: - raise ValueError("Cannot specify both target-vocab-size and target-model-path") - - if not has_vocab and not has_model: - raise ValueError("Must specify either target-vocab-size or target-model-path") - - if has_vocab: - return args.target_vocab_size - - logger.info(f"Loading target model config from {args.target_model_path}") - config = AutoConfig.from_pretrained(args.target_model_path) - - # For multimodal models (Qwen3VL, etc.), extract text_config - if hasattr(config, "text_config"): - config = config.text_config - - return config.vocab_size - - def main(): args = parse_args() @@ -103,7 +80,9 @@ def main(): token_freq_dict = torch.load(token_freq_path, weights_only=True) - target_vocab_size = get_target_vocab_size(args) + target_vocab_size = get_target_vocab_size( + args.target_vocab_size, args.target_model_path + ) d2t, t2d = build_vocab_mappings_from_distribution( token_freq_dict=token_freq_dict, diff --git a/scripts/data_generation_offline.py b/scripts/data_generation_offline.py index 756d49fad..1fe7ec35a 100644 --- a/scripts/data_generation_offline.py +++ b/scripts/data_generation_offline.py @@ -82,6 +82,7 @@ def parse_args(): parser.add_argument( "--train-data-path", type=str, + action="append", required=True, help="Path to training data (same as used in preprocessing)", ) @@ -340,13 +341,12 @@ def main(): dataset, _ = load_and_preprocess_dataset( target_model_path=args.target_model_path, - train_data_path=args.train_data_path, + train_data_paths=args.train_data_path, seq_length=args.seq_length, build_dataset_num_proc=args.num_preprocessing_workers, seed=args.seed, max_samples=args.max_samples, token_freq_path=args.token_freq_path, - cache_dir=args.hf_cache_dir, assistant_pattern=args.assistant_pattern, turn_dropout=args.turn_dropout, ) diff --git a/scripts/data_generation_offline2.py b/scripts/data_generation_offline2.py new file mode 100644 index 000000000..11eaf785a --- /dev/null +++ b/scripts/data_generation_offline2.py @@ -0,0 +1,321 @@ +#!/usr/bin/env python3 +""" +Offline Hidden States Generation Pipeline + +This script generates hidden states and saves them to disk for offline training. + +Usage: + python data_generation_offline.py \ + --model meta-llama/Llama-3.1-8B-Instruct \ + --preprocessed-data sharegpt \ + --output ./training_data \ + --max-samples 5000 +""" + +import argparse +import asyncio +import logging +import shutil +import sys +from pathlib import Path +from typing import Any + +import openai +from datasets import load_from_disk +from safetensors import safe_open +from tqdm import tqdm + +from speculators.data_generation.vllm_client import generate_hidden_states_async +from speculators.train.logger import setup_root_logger + +logger = logging.getLogger(__name__) + + +def parse_args(): + parser = argparse.ArgumentParser(description="Generate EAGLE training data offline") + + # Model arguments + parser.add_argument( + "--model", + type=str, + default=None, + help=( + "HuggingFace model ID or local path for target model (default auto select)." + "For verification purposes only." + ), + ) + parser.add_argument( + "--endpoint", + type=str, + default="http://localhost:8000/v1", + help=( + "The address of the vLLM instance to use for hidden states generation " + "(default: 'http://localhost:8000/v1'). " + "Note: the vLLM instance must be configured for hidden states extraction." + ), + ) + + # Data arguments + parser.add_argument( + "--preprocessed-data", + type=str, + required=True, + help="Path to preprocessed dataset (dataset produced by prepare_data.py)", + ) + parser.add_argument( + "--max-samples", + type=int, + default=None, + help="Maximum number of samples to process (default: None, process all)", + ) + + # Output arguments + parser.add_argument( + "--output", + type=str, + default=None, + help=( + "Directory to generated hidden states files " + "(default args.preprocessed_data / 'hidden_states')" + ), + ) + + # Hidden states generation arguments + parser.add_argument( + "--layer-ids", + type=int, + nargs="+", + default=None, + help=( + "List of layer IDs from which to capture hidden states " + "(default: auto-select)" + ), + ) + parser.add_argument( + "--concurrency", + type=int, + default=32, + help=( + "Number of active vLLM requests at a time." + "Note: number of async workers set to 2*concurrency" + ), + ) + parser.add_argument( + "--validate-outputs", + action="store_true", + help=( + "Load generated safetensor files and check output token ids match prompt" + " tokens and hidden states seq_len matches num tokens" + ), + ) + + # Processing arguments + parser.add_argument( + "--start-idx", + type=int, + default=0, + help="Starting index for output files (default: 0)", + ) + return parser.parse_args() + + +def get_existing_hidden_state_indices(output_path: Path) -> list[int]: + """Find existing `hs_i.safetensors` files (where i is the file index)""" + + existing_file_indices = [] + + if not output_path.exists(): + return existing_file_indices + + for file_path in output_path.iterdir(): + if file_path.name.startswith("hs_") and file_path.name.endswith(".safetensors"): + index_str = file_path.stem[3:] # Remove "hs_" prefix + try: + file_index = int(index_str) + existing_file_indices.append(file_index) + except ValueError: + continue + + return sorted(existing_file_indices) + + +def get_indices_to_process( + num_samples: int, max_samples: int | None, existing: list[int] +) -> list[int]: + """Determines which indices should be processed. If max_samples is None + returns all dataset indices not in existing. Otherwise gets the first + `max_samples - len(existing)` samples not already in existing. + + Args: + num_samples: Total size of preprocessed dataset + max_samples: (Optional) limit for number of samples to process + existing: list of ids that have already been processed + + Returns: + list of dataset indices to process + """ + + if len(existing) >= num_samples: + logger.info("All samples already processed!") + return [] + if max_samples and len(existing) >= max_samples: + logger.info("At least max_samples already processed!") + return [] + + if len(existing) > 0: + logger.info(f"Found {len(existing)} existing samples.") + + existing_s = set(existing) + if max_samples is None: + return [i for i in range(num_samples) if i not in existing_s] + + num_remaining = min(max_samples, num_samples) - len(existing) + to_process = [] + cur = 0 + while num_remaining > 0 and cur < num_samples: + if cur not in existing_s: + to_process.append(cur) + num_remaining -= 1 + + cur += 1 + + return to_process + + +def check_safetensors_file(path: Path, tokens: list[int]): + with safe_open(path, "pt") as f: + t_ids = f.get_tensor("token_ids").tolist() + if t_ids != tokens: + raise ValueError( + f"Token ids in {path} don't match expected token ids {tokens}" + ) + + hs_slice = f.get_slice("hidden_states") + hs_shape = list(hs_slice.get_shape()) + if len(tokens) != hs_shape[0]: + raise ValueError( + f"Sequence length of hidden states {hs_shape[0]} in {path}" + f" doesn't match num tokens {len(tokens)}" + ) + + +async def worker( + client, + model: str, + queue: "asyncio.Queue[dict[str, Any]]", + pbar: tqdm, + vllm_semaphore: asyncio.Semaphore, + write_semaphore: asyncio.Semaphore, + hidden_states_output_dir: Path, + validate_outputs: bool, +): + """Worker that pulls items from queue and sends them to the vLLM endpoint.""" + while True: + item = await queue.get() + if item is None: + queue.task_done() + return + + idx = item["idx"] + input_ids = item["input_ids"].tolist() + + target_hidden_states_path = hidden_states_output_dir / f"hs_{idx}.safetensors" + + try: + async with vllm_semaphore: # Limit number of active generate calls + hidden_states_path = await generate_hidden_states_async( + client, model, input_ids + ) + async with write_semaphore: # Limit number of active disk writes + await asyncio.to_thread( + shutil.move, hidden_states_path, target_hidden_states_path + ) + if validate_outputs: + await asyncio.to_thread( + check_safetensors_file, target_hidden_states_path, input_ids + ) + finally: + pbar.update(1) + queue.task_done() + + +async def generate_and_save_hidden_states(args, dataset): + if args.output is None: + hidden_states_dir = Path(args.preprocessed_data) / "hidden_states" + else: + hidden_states_dir = Path(args.output) + hidden_states_dir.mkdir(parents=True, exist_ok=True) + + existing_file_indices = get_existing_hidden_state_indices(hidden_states_dir) + num_samples = len(dataset) + + to_process = get_indices_to_process( + num_samples, args.max_samples, existing_file_indices + ) + if not to_process: + return + + logger.info(f"Processing {len(to_process)} samples") + + queue: asyncio.Queue = asyncio.Queue(maxsize=args.concurrency * 4) + vllm_semaphore = asyncio.Semaphore(args.concurrency) + write_semaphore = asyncio.Semaphore(args.concurrency) + + async with openai.AsyncOpenAI(base_url=args.endpoint, api_key="EMPTY") as client: + list_models = await client.models.list() + model_id = list_models.data[0].id + if args.model and args.model != model_id: + raise ValueError( + f"An explicit model name was passed ({args.model}) which doesn't match" + "found model_id {model_id}." + "Please make sure --endpoint is set to the correct vllm instance." + ) + + with tqdm(total=len(to_process)) as pbar: + workers = [ + asyncio.create_task( + worker( + client, + model_id, + queue, + pbar, + vllm_semaphore, + write_semaphore, + hidden_states_dir, + args.validate_outputs, + ) + ) + for _ in range(args.concurrency * 2) + ] + + for i in to_process: + item = dataset[i] + await queue.put({"idx": i, "input_ids": item["input_ids"]}) + + logger.info("Waiting for remaining file saves to complete...") + # Signale workers to stop + for _ in range(len(workers)): + await queue.put(None) + await asyncio.gather(*workers) + + logger.info(f"Saved {len(to_process)} new data points to {args.output}") + + +def main(): + args = parse_args() + setup_root_logger() + + logger.info("EAGLE Offline Data Generation") + + dataset = load_from_disk(args.preprocessed_data) + + try: + asyncio.run(generate_and_save_hidden_states(args, dataset)) + except KeyboardInterrupt: + sys.exit(130) + + logger.info("Data generation complete!") + + +if __name__ == "__main__": + main() diff --git a/scripts/launch_vllm.py b/scripts/launch_vllm.py new file mode 100644 index 000000000..1465bb77a --- /dev/null +++ b/scripts/launch_vllm.py @@ -0,0 +1,86 @@ +import argparse +import json +import os + + +def parse_args(): + parser = argparse.ArgumentParser( + description="Launch vLLM for hidden states extraction", + usage=( + "launch_vllm.py [-h] MODEL [--hidden-states-path HIDDEN_STATES_PATH] " + "[--layers LAYERS [LAYERS ...]] -- *VLLM_ARGS" + ), + ) + parser.add_argument( + "model", type=str, help="Model name or path to extract hidden states from" + ) + parser.add_argument( + "--hidden-states-path", + type=str, + default="/tmp/hidden_states", # noqa: S108 + help="The directory to save hidden states to. Default '/tmp/hidden_states'.", + ) + parser.add_argument( + "--layers", + type=int, + nargs="+", + help=( + "(Optional) A (space separated) list of integer layer ids. Default layers " + "[2, num_hidden_layers // 2, num_hidden_layers - 3, num_hidden_layers]." + ), + ) + parser.add_argument( + "vllm_args", nargs=argparse.REMAINDER, help="Arguments to be passed to vLLM" + ) + + return parser.parse_args() + + +def main(): + args = parse_args() + + if args.layers: + layers = args.layers + else: + # Import here so that it isn't required if layers passed explicitly + from transformers import AutoConfig # noqa: PLC0415 + + config = AutoConfig.from_pretrained(args.model) + if hasattr(config, "text_config"): + config = config.text_config + + num_hidden_layers = config.num_hidden_layers + layers = [2, num_hidden_layers // 2, num_hidden_layers - 3, num_hidden_layers] + + speculative_config = { + "method": "extract_hidden_states", + "num_speculative_tokens": 1, + "draft_model_config": { + "hf_config": {"eagle_aux_hidden_state_layer_ids": layers} + }, + } + kv_transfer_config = { + "kv_connector": "ExampleHiddenStatesConnector", + "kv_role": "kv_producer", + "kv_connector_extra_config": {"shared_storage_path": args.hidden_states_path}, + } + + cmd = [ + "vllm", + "serve", + args.model, + "--speculative_config", + json.dumps(speculative_config), + "--kv_transfer_config", + json.dumps(kv_transfer_config), + *args.vllm_args, + ] + + print("Running command:") + print(" ".join(cmd)) + + os.execvp(cmd[0], cmd) # noqa: S606 + + +if __name__ == "__main__": + main() diff --git a/scripts/prepare_data.py b/scripts/prepare_data.py new file mode 100644 index 000000000..c84b2c812 --- /dev/null +++ b/scripts/prepare_data.py @@ -0,0 +1,175 @@ +#!/usr/bin/env python3 +""" +Prepare data for speculator training + +This script processes an input dataset and: +1. Applies chat template + tokenizes each sample +2. Produces a loss/assistant mask for each sample +3. Records token frequency statistics + +The output of this script is: +1. Processed dataset ready for online training or offline datagen in output_dir +2. Token frequency statistics file at token_freq_path + +Preprocessing will be skipped if the dataset already exists at the output directory. +Token frequencies are saved in the output directory by default. + +Usage: + python prepare_data.py \ + --model meta-llama/Llama-3.1-8B-Instruct \ + --data sharegpt \ + --output ./training_data \ + --max-samples 5000 +""" + +import argparse +import glob +import logging +import sys +from pathlib import Path + +from speculators.data_generation.logging_utils import PipelineLogger # noqa: E402 +from speculators.data_generation.preprocessing import ( # noqa: E402 + load_and_preprocess_dataset, +) + +logging.basicConfig( + level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s" +) +log = PipelineLogger(__name__) + + +def parse_args(): + parser = argparse.ArgumentParser(description="Prepare data for speculator training") + + # Model arguments + parser.add_argument( + "--model", + type=str, + required=True, + help="HuggingFace model ID or local path for target model", + ) + + # Data arguments + parser.add_argument( + "--data", + type=str, + action="append", + required=True, + help="Path to training data (same as used in preprocessing)", + ) + parser.add_argument( + "--seq-length", + type=int, + default=8192, + help="Maximum sequence length for preprocessing and model (default: 8192)", + ) + parser.add_argument( + "--max-samples", + type=int, + default=None, + help="Maximum number of samples to process (default: None, process all)", + ) + parser.add_argument( + "--token-freq-path", + type=str, + default=None, + help=( + "Path to save token frequency distribution" + "(default: args.output / 'token_freq.pt')" + ), + ) + parser.add_argument( + "--assistant-pattern", + type=str, + default=None, + help=( + "Custom regex pattern for matching assistant responses. " + "If not provided, auto-detected from chat template." + ), + ) + parser.add_argument( + "--turn-dropout", + action="store_true", + help=( + "Enable turn dropout: randomly keeps first N consecutive turns " + "per conversation for data augmentation." + ), + ) + + # Output arguments + parser.add_argument( + "--output", type=str, required=True, help="Directory to save output dataset" + ) + parser.add_argument( + "--overwrite", + action="store_true", + help=( + "Forcibly rerun `prepare_data.py`.Deletes existing content in output dir" + ), + ) + + # Processing arguments + parser.add_argument( + "--seed", + type=int, + default=0, + help="Random seed (must match preprocessing seed, default: 0)", + ) + parser.add_argument( + "--num-preprocessing-workers", + type=int, + default=8, + help="Number of CPU processes for dataset preprocessing (default: 8)", + ) + return parser.parse_args() + + +def main(): + args = parse_args() + + log.section("Preparing data") + log.config( + { + "Target Model": args.model, + "Dataset": args.data, + "Output Dir": args.output, + } + ) + + output = Path(args.output) + if output.exists(): + if not args.overwrite and glob.glob(str(output / "*.arrow")): + log.warning( + "Dataset files already exists in output directory, skipping " + "preprocessing. To existing overwrite files use --overwrite." + ) + sys.exit(0) + else: + output.mkdir(parents=True) + + token_freq_path = ( + output / "token_freq.pt" + if args.token_freq_path is None + else Path(args.token_freq_path) + ) + + dataset, _ = load_and_preprocess_dataset( + target_model_path=args.model, + train_data_paths=args.data, + seq_length=args.seq_length, + build_dataset_num_proc=args.num_preprocessing_workers, + seed=args.seed, + max_samples=args.max_samples, + token_freq_path=token_freq_path, + assistant_pattern=args.assistant_pattern, + turn_dropout=args.turn_dropout, + ) + + log.info("Done preparing data") + log.section(f"Writing dataset to {args.output}") + dataset.save_to_disk(args.output) + + +if __name__ == "__main__": + main() diff --git a/scripts/train.py b/scripts/train.py index 0f688ba74..e00018eda 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -1,6 +1,8 @@ import argparse +import logging import random import warnings +from pathlib import Path import numpy as np import torch @@ -11,10 +13,11 @@ from speculators.model import SpeculatorModel from speculators.train.data import ( + BaseEagle3Dataset, + Eagle3ArrowDataset, Eagle3SampleFileDataset, create_collate_fn, split_files, - standardize_data_v1, ) from speculators.train.distributed_batch_sampler import ( MultipackDistributedBatchSamplerV2, @@ -23,6 +26,12 @@ from speculators.train.noise_transforms import AddUniformNoise from speculators.train.trainer import Trainer, TrainerConfig from speculators.train.utils import maybe_destroy_distributed, maybe_setup_distributed +from speculators.train.vocab_mapping import ( + build_vocab_mappings_from_distribution, + get_target_vocab_size, +) + +logger = logging.getLogger(__name__) DRAFT_ARCH_CONFIGS: dict[str, type] = { "llama": LlamaConfig, @@ -44,14 +53,12 @@ def set_seed(seed: int, deterministic: bool = False): def setup_dataloader( - file_list: list[str], + dataset: BaseEagle3Dataset, world_size: int, local_rank: int, - add_noise: bool = True, - noise_std: float = 0.05, + hidden_size: int, num_workers: int = 12, prefetch_factor: int = 4, - hidden_states_dtype: torch.dtype = torch.bfloat16, ) -> DataLoader: """Setup dataloader for training. Args: @@ -65,22 +72,6 @@ def setup_dataloader( Returns: DataLoader: Dataloader for training. """ - if add_noise: - noise_transform = AddUniformNoise( - std=noise_std, tensors=("hidden_states", "verifier_last_hidden_states") - ) - else: - noise_transform = None - - standardize_fn = standardize_data_v1 - - dataset = Eagle3SampleFileDataset( - file_list=file_list, - max_len=args.total_seq_len, - transform=noise_transform, - standardize_fn=standardize_fn, - hidden_states_dtype=hidden_states_dtype, - ) batch_sampler = MultipackDistributedBatchSamplerV2( batch_max_length=args.total_seq_len, lengths=dataset.approx_lengths, @@ -93,7 +84,7 @@ def setup_dataloader( num_workers=num_workers, prefetch_factor=prefetch_factor, pin_memory=True, - collate_fn=create_collate_fn(args.total_seq_len), + collate_fn=create_collate_fn(args.total_seq_len, hidden_size), persistent_workers=True, ) @@ -138,6 +129,78 @@ def create_transformer_layer_config( ) +def _load_mappings(d2t_path, t2d_path, expected_draft_vocab_size: int | None): + logger.info(f"Loading vocab mappings from '{d2t_path}' and '{t2d_path}'") + # Load d2t and t2d tensors if provided + d2t = torch.from_numpy(np.load(d2t_path)) + t2d = torch.from_numpy(np.load(t2d_path)) + draft_vocab_size = d2t.shape[0] + if expected_draft_vocab_size and expected_draft_vocab_size != draft_vocab_size: + raise ValueError( + f"Explicit vocab mapping (t2d & d2t) files were provided, but don't" + f"match the provided --draft-vocab-size {draft_vocab_size}." + f"d2t.shape={d2t.shape}, dim 0 should match provided value." + ) + return d2t, t2d, draft_vocab_size + + +def parse_vocab_mappings(args: argparse.Namespace): + if args.d2t_path or args.t2d_path: + if not (args.d2t_path and args.t2d_path): + raise ValueError( + "Both t2d and d2t must be provided together, or both must be omitted. " + f"Got t2d={'provided' if args.t2d_path is not None else 'not provided'}" + f"d2t={'provided' if args.d2t_path is not None else 'not provided'}" + ) + + return _load_mappings(args.d2t_path, args.t2d_path, args.draft_vocab_size) + + data_path = Path(args.data_path) + default_t2d_path = data_path / "t2d.npy" + default_d2t_path = data_path / "d2t.npy" + + if default_t2d_path.exists() and default_d2t_path.exists(): + return _load_mappings(default_d2t_path, default_t2d_path, args.draft_vocab_size) + + token_freq_path = args.token_freq_path or data_path / "token_freq.pt" + token_freq_path = Path(token_freq_path) + if token_freq_path.exists() and args.draft_vocab_size is not None: + logger.info("No vocab mappings provided. Regenerating from token frequencies") + token_freq_dict = torch.load(token_freq_path, weights_only=True) + + target_vocab_size = get_target_vocab_size(None, args.verifier_name_or_path) + + d2t, t2d = build_vocab_mappings_from_distribution( + token_freq_dict=token_freq_dict, + draft_vocab_size=args.draft_vocab_size, + target_vocab_size=target_vocab_size, + ) + draft_vocab_size = d2t.shape[0] + if args.draft_vocab_size and args.draft_vocab_size != draft_vocab_size: + raise ValueError( + f"Explicit vocab mapping (t2d & d2t) files were provided, but don't" + f"match the provided --draft-vocab-size {draft_vocab_size}." + f"d2t.shape={d2t.shape}, dim 0 should match provided value." + ) + + logger.info(f"Caching vocab mapping files to '{data_path}'") + np.save(data_path / "d2t.npy", d2t.cpu().numpy()) + np.save(data_path / "t2d.npy", t2d.cpu().numpy()) + + return d2t, t2d, draft_vocab_size + + logger.warning( + "No vocab mappings found, and can't generate new ones because either " + f"token_freq_path='{token_freq_path}' doesn't exist or --draft-vocab-size is " + "None. Using full verifier vocab" + ) + # When vocab mapping is not provided, use the full verifier vocab + verifier_config = AutoConfig.from_pretrained(args.verifier_name_or_path) + if hasattr(verifier_config, "text_config"): + verifier_config = verifier_config.text_config + return None, None, verifier_config.vocab_size + + def main(args: argparse.Namespace): # Set random seed for reproducibility set_seed(args.seed, args.deterministic_cuda) @@ -150,32 +213,13 @@ def main(args: argparse.Namespace): # Setup distributed training local_rank, world_size, rank, is_distributed = maybe_setup_distributed() - device = torch.device(local_rank) if not hasattr(torch, args.hidden_states_dtype): raise ValueError( "--hidden-states-dtype must be a dtype attribute of torch. e.g. `bfloat16`" ) hidden_states_dtype = getattr(torch, args.hidden_states_dtype) - # Load t2d and d2t tensors if provided - if args.d2t_path or args.t2d_path: - if not (args.d2t_path and args.t2d_path): - raise ValueError( - "Both t2d and d2t must be provided together, or both must be omitted. " - f"Got t2d={'provided' if args.t2d_path is not None else 'not provided'}" - f"d2t={'provided' if args.d2t_path is not None else 'not provided'}" - ) - d2t = torch.from_numpy(np.load(args.d2t_path)).to(device) - t2d = torch.from_numpy(np.load(args.t2d_path)).to(device) - draft_vocab_size = d2t.shape[0] - else: - d2t = None - t2d = None - # When vocab mapping is not provided, use the full verifier vocab - verifier_config = AutoConfig.from_pretrained(args.verifier_name_or_path) - if hasattr(verifier_config, "text_config"): - verifier_config = verifier_config.text_config - draft_vocab_size = verifier_config.vocab_size + d2t, t2d, draft_vocab_size = parse_vocab_mappings(args) # Setup speculator config transformer_layer_config = create_transformer_layer_config( @@ -194,35 +238,68 @@ def main(args: argparse.Namespace): args.from_pretrained, t2d=t2d, d2t=d2t ) else: + args_dict = vars(args) + args_dict["draft_vocab_size"] = draft_vocab_size draft_model = model_class.from_training_args( verifier_config=transformer_layer_config, t2d=t2d, d2t=d2t, - draft_vocab_size=draft_vocab_size, - **vars(args), + **args_dict, ) # Setup dataloaders - train_files, val_files = split_files(args.data_path, ratio=0.9) + noise_transform = AddUniformNoise(std=args.noise_std) + if args.legacy_data: + warnings.warn( + "Using '--legacy-data' is deprecated and will be removed soon.", + category=DeprecationWarning, + stacklevel=2, + ) + train_files, val_files = split_files(args.data_path, ratio=0.9) + train_dataset: BaseEagle3Dataset = Eagle3SampleFileDataset( + file_list=train_files, max_len=args.total_seq_len, transform=noise_transform + ) + val_dataset: BaseEagle3Dataset = Eagle3SampleFileDataset( + file_list=val_files, max_len=args.total_seq_len + ) + else: + train_dataset = Eagle3ArrowDataset( + datapath=args.data_path, + max_len=args.total_seq_len, + hidden_states_path=args.hidden_states_path, + vllm_endpoint=args.vllm_endpoint, + on_missing=args.on_missing, + on_generate=args.on_generate, + transform=noise_transform, + split_ratio=0.9, + model=args.verifier_name_or_path, + ) + val_dataset = Eagle3ArrowDataset( + datapath=args.data_path, + max_len=args.total_seq_len, + hidden_states_path=args.hidden_states_path, + vllm_endpoint=args.vllm_endpoint, + on_missing=args.on_missing, + on_generate=args.on_generate, + split_ratio=-0.1, + model=args.verifier_name_or_path, + ) + train_loader = setup_dataloader( - train_files, + train_dataset, world_size, local_rank, - add_noise=True, - noise_std=args.noise_std, + transformer_layer_config.hidden_size, num_workers=args.num_workers, prefetch_factor=args.prefetch_factor, - hidden_states_dtype=hidden_states_dtype, ) val_loader = setup_dataloader( - val_files, + val_dataset, world_size, local_rank, - add_noise=False, - noise_std=args.noise_std, + transformer_layer_config.hidden_size, num_workers=args.num_workers, prefetch_factor=args.prefetch_factor, - hidden_states_dtype=hidden_states_dtype, ) # Get trainer kwargs from model class @@ -277,6 +354,60 @@ def parse_args(): help="The pretrained draft model to finetune", ) parser.add_argument("--data-path", type=str, default="./data") + parser.add_argument( + "--hidden-states-path", + type=str, + default=None, + help=( + "The path where cached hidden states files are stored. (Default: " + "args.data_path / 'hidden_states')" + ), + ) + parser.add_argument( + "--vllm-endpoint", + type=str, + default="http://localhost:8000/v1", + help=( + "vLLM endpoint address to use if generating hidden states on-demand." + " Only required if `--on-missing=generate` and samples are missing." + " Note: the vLLM instance must be configured to cache hidden states" + " to a location that is accessible from the training instance. i.e." + " on the same node, or a shared network drive. (Default: 'http://localhost:8000/v1')" + ), + ) + parser.add_argument( + "--on-missing", + choices=["generate", "skip", "warn", "raise"], + default="generate", + help=( + "Dataloader behaviour when there are no cached hidden states for a sample." + "Default: 'generate', which attempts to generate the hidden states on-" + "demand using the provided vLLM endpoint. The other options skip the sample" + ", skip and warn, or raise an error respectively." + ), + ) + parser.add_argument( + "--on-generate", + choices=["cache", "delete"], + default="delete", + help=( + "Dataloader behaviour when a new hidden state has been generated" + " (only applies if args.on_missing=='generate'). Default: 'delete', " + "deletes hidden states once they are loaded. 'cache' will instead store" + "the hidden states in the args.hidden_states_path. This can be used to " + "enable hybrid online/offline training, with hidden states generated on the" + "first epoch, and reused on subsequent epochs." + ), + ) + parser.add_argument( + "--legacy-data", + action="store_true", + help=( + "DEPRECATED. Use the old data format which stores hidden states alongside " + "token_ids and assistant_masks, in data_i.pt files. This option will be " + "removed soon." + ), + ) parser.add_argument("--save-path", type=str, default="./checkpoints") parser.add_argument("--epochs", type=int, default=20) parser.add_argument("--lr", type=float, default=1e-4) @@ -299,6 +430,19 @@ def parse_args(): help="Architecture for draft decoder layers. Defaults to 'llama'. " "Note: only 'llama' is currently supported in vLLM for inference.", ) + + parser.add_argument( + "--token-freq-path", + type=str, + default=None, + help="Path to token frequency distribution file (.pt)", + ) + parser.add_argument( + "--draft-vocab-size", + type=int, + default=None, + help="Vocabulary size for the draft model", + ) parser.add_argument("--d2t-path", type=str, default=None) parser.add_argument("--t2d-path", type=str, default=None) parser.add_argument("--ttt-steps", type=int, default=3) diff --git a/src/speculators/data_generation/__init__.py b/src/speculators/data_generation/__init__.py index 8819704df..a4aa8645d 100644 --- a/src/speculators/data_generation/__init__.py +++ b/src/speculators/data_generation/__init__.py @@ -1,7 +1 @@ """Data generation utilities for EAGLE-style speculative decoding training.""" - -from speculators.data_generation.vllm_hidden_states_generator import ( - VllmHiddenStatesGenerator, -) - -__all__ = ["VllmHiddenStatesGenerator"] diff --git a/src/speculators/data_generation/preprocessing.py b/src/speculators/data_generation/preprocessing.py index 928e97f69..8a151bb60 100644 --- a/src/speculators/data_generation/preprocessing.py +++ b/src/speculators/data_generation/preprocessing.py @@ -1,12 +1,13 @@ import bisect import random import re +from pathlib import Path from re import Pattern from typing import Any, cast import torch from datasets import Dataset as HFDataset -from datasets import load_dataset +from datasets import concatenate_datasets, load_dataset from transformers import AutoTokenizer, PreTrainedTokenizerBase from speculators.data_generation.configs import DATASET_CONFIGS @@ -22,7 +23,7 @@ log = PipelineLogger(__name__) -def _visualize_sample(_dataset, preprocessed, tokenizer, idx: int = 0): +def _visualize_sample(preprocessed, tokenizer, idx: int = 0): """Visualize a single sample with color-coded trainable regions.""" # Get preprocessed sample prep_sample = preprocessed[idx] @@ -225,7 +226,7 @@ def _create_loss_mask_from_offsets( assistant_pattern: str | Pattern[str], ) -> torch.Tensor: """Create loss mask by finding assistant response spans in formatted text.""" - loss_mask = torch.zeros(len(offsets), dtype=torch.long) + loss_mask = torch.zeros(len(offsets), dtype=torch.bool) matches_found = 0 token_starts = [offset[0] for offset in offsets] @@ -263,7 +264,7 @@ def _preprocess_batch( ) -> dict[str, list]: """Process a batch of conversations into tokenized format with loss masks.""" - results: dict[str, list] = {"input_ids": [], "loss_mask": []} + results: dict[str, list] = {"input_ids": [], "loss_mask": [], "seq_len": []} conversations = examples.get("conversations", []) if not conversations: @@ -339,6 +340,7 @@ def _preprocess_batch( # Append to results results["input_ids"].append(torch.tensor(input_ids, dtype=torch.long)) results["loss_mask"].append(loss_mask) + results["seq_len"].append(len(input_ids)) except (TypeError, ValueError, KeyError, AttributeError, RuntimeError) as e: log.error( @@ -393,21 +395,17 @@ def build_eagle3_dataset( num_proc=num_proc, batch_size=1000, remove_columns=original_cols, - load_from_cache_file=True, + keep_in_memory=True, # skip caching ) dataset.set_format(type="torch") return dataset -def load_raw_dataset( - train_data_path: str, num_proc: int = 8, cache_dir: str | None = None -) -> HFDataset: +def load_raw_dataset(train_data_path: str, num_proc: int = 8) -> HFDataset: """Load raw dataset from local file or HuggingFace.""" if train_data_path.endswith((".jsonl", ".json")): - return load_dataset( - "json", data_files=train_data_path, split="train", cache_dir=cache_dir - ) + return load_dataset("json", data_files=train_data_path, split="train") if train_data_path not in DATASET_CONFIGS: raise ValueError( @@ -416,7 +414,7 @@ def load_raw_dataset( ) config = DATASET_CONFIGS[train_data_path] - raw_dataset = load_dataset(config.hf_path, split=config.split, cache_dir=cache_dir) + raw_dataset = load_dataset(config.hf_path, split=config.split) if config.normalize_fn is not None: raw_dataset = raw_dataset.map(config.normalize_fn, num_proc=num_proc) @@ -426,13 +424,12 @@ def load_raw_dataset( def load_and_preprocess_dataset( target_model_path: str, - train_data_path: str, + train_data_paths: list[str], seq_length: int, build_dataset_num_proc: int = 8, seed: int = 0, max_samples: int | None = None, - token_freq_path: str = "./token_freq.pt", # noqa: S107 - cache_dir: str | None = None, + token_freq_path: Path | str = "./token_freq.pt", # noqa: S107 assistant_pattern: str | None = None, turn_dropout: bool = False, ) -> tuple[HFDataset, PreTrainedTokenizerBase]: @@ -461,7 +458,7 @@ def load_and_preprocess_dataset( """ log.section("Starting dataset preprocessing") - log.subsection("Loading tokenizer and dataset") + log.subsection("Loading tokenizer") tokenizer = AutoTokenizer.from_pretrained(target_model_path, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token @@ -472,40 +469,47 @@ def load_and_preprocess_dataset( "Please use a model with a pre-configured chat template." ) - raw_dataset = load_raw_dataset( - train_data_path, num_proc=build_dataset_num_proc, cache_dir=cache_dir - ) - raw_dataset = raw_dataset.shuffle(seed=seed) + processed_datasets = [] + for train_data_path in train_data_paths: + log.subsection(f"Processing {train_data_path}") + raw_dataset = load_raw_dataset(train_data_path, num_proc=build_dataset_num_proc) + raw_dataset = raw_dataset.shuffle(seed=seed) + + if max_samples is not None and len(raw_dataset) > 3 * max_samples: + # Reduce size to 3 * max_samples to reduce processing + # This will then be reduced further to max_samples + # after combining datasets and shuffling + raw_dataset = raw_dataset.select(range(3 * max_samples)) + + log.info(f"Loaded {len(raw_dataset)} samples") + + if turn_dropout: + log.info("Turn dropout enabled: randomly keeping N consecutive turns") + + preprocessed_dataset = build_eagle3_dataset( + dataset=raw_dataset, + tokenizer=tokenizer, + max_length=seq_length, + num_proc=build_dataset_num_proc, + assistant_pattern=assistant_pattern, + turn_dropout=turn_dropout, + ) + processed_datasets.append(preprocessed_dataset) + combined_dataset = concatenate_datasets(processed_datasets) + combined_dataset.shuffle(seed=seed) if max_samples is not None and len(raw_dataset) > max_samples: - raw_dataset = raw_dataset.select(range(max_samples)) - - log.info(f"Loaded {len(raw_dataset)} samples") - - log.subsection("Tokenizing and building dataset") - if cache_dir: - log.info(f"Preprocessed data will be cached at: {cache_dir}") - if turn_dropout: - log.info("Turn dropout enabled: randomly keeping N consecutive turns") - - preprocessed_dataset = build_eagle3_dataset( - dataset=raw_dataset, - tokenizer=tokenizer, - max_length=seq_length, - num_proc=build_dataset_num_proc, - assistant_pattern=assistant_pattern, - turn_dropout=turn_dropout, - ) + combined_dataset = combined_dataset.select(range(max_samples)) log.subsection("Computing token frequency distribution") save_token_frequency_distribution( - dataset=preprocessed_dataset, + dataset=combined_dataset, output_path=token_freq_path, ) log.subsection("Visualizing sample") - _visualize_sample(raw_dataset, preprocessed_dataset, tokenizer, idx=0) + _visualize_sample(combined_dataset, tokenizer, idx=0) log.section("Dataset preprocessing complete") - return preprocessed_dataset, tokenizer + return combined_dataset, tokenizer diff --git a/src/speculators/data_generation/vllm_client.py b/src/speculators/data_generation/vllm_client.py new file mode 100644 index 000000000..b78a9d346 --- /dev/null +++ b/src/speculators/data_generation/vllm_client.py @@ -0,0 +1,52 @@ +import openai + + +class InvalidResponseError(Exception): + pass + + +def extract_output(completion, token_ids) -> str: + prompt_token_ids = getattr(completion.choices[0], "prompt_token_ids", None) + + if prompt_token_ids is None: + raise InvalidResponseError("Response missing prompt_token_ids") + + if prompt_token_ids != token_ids: + raise InvalidResponseError( + f"Prompt token IDs mismatch: expected {token_ids}, got {prompt_token_ids}" + ) + + if not hasattr(completion, "kv_transfer_params"): + raise InvalidResponseError("Response missing kv_transfer_params") + + return completion.kv_transfer_params.get("hidden_states_path") + + +async def generate_hidden_states_async( + client: openai.AsyncClient, model: str, token_ids: list[int] +) -> str: + """ + Runs decode w/ max_tokens 1 to generate hidden states and returns path to + hidden states file. + """ + + completion = await client.completions.create( + model=model, + prompt=token_ids, + max_tokens=1, + extra_body={"return_token_ids": True}, + ) + + return extract_output(completion, token_ids) + + +def generate_hidden_states( + client: openai.Client, model: str, token_ids: list[int] +) -> str: + completion = client.completions.create( + model=model, + prompt=token_ids, + max_tokens=1, + extra_body={"return_token_ids": True}, + ) + return extract_output(completion, token_ids) diff --git a/src/speculators/data_generation/vllm_hidden_states_generator.py b/src/speculators/data_generation/vllm_hidden_states_generator.py index eadd6032b..0c987aab7 100644 --- a/src/speculators/data_generation/vllm_hidden_states_generator.py +++ b/src/speculators/data_generation/vllm_hidden_states_generator.py @@ -1,5 +1,6 @@ """Extract hidden states from intermediate layers during prefill using vLLM.""" +import warnings from typing import Literal import torch @@ -84,6 +85,12 @@ def __init__( tensor_parallel_size: int = 1, max_num_batched_tokens: int | None = None, ): + warnings.warn( + "VllmHiddenStatesGenerator and the associated data_generation_offline.py" + " script are deprecatd and will be removed shortly.", + DeprecationWarning, + stacklevel=2, + ) self.model_path = model_path self.tensor_parallel_size = tensor_parallel_size self._request_counter = 0 diff --git a/src/speculators/py.typed b/src/speculators/py.typed new file mode 100644 index 000000000..e69de29bb diff --git a/src/speculators/train/data.py b/src/speculators/train/data.py index 2c0362953..f72f57270 100644 --- a/src/speculators/train/data.py +++ b/src/speculators/train/data.py @@ -3,14 +3,24 @@ import math import os import random +import shutil +import warnings from collections.abc import Callable +from os import PathLike from pathlib import Path -from typing import Any +from typing import Any, Literal +import openai import torch import torch.nn.functional as F # noqa: N812 +from datasets import load_from_disk +from safetensors.torch import load_file from torch.utils.data import Dataset +from speculators.data_generation.vllm_client import ( + InvalidResponseError, + generate_hidden_states, +) from speculators.train.noise_transforms import TransformTensors BatchType = dict[str, Any] @@ -92,6 +102,26 @@ def split_files(datapath: str, ratio: float = 0.9, seed: int = 0): StandardizeFnSig = Callable[[dict[str, Any]], dict[str, Any]] +def create_empty_sample(hidden_size: int): + # data structure: { + # "hidden_states": [seq_len, 3 * hidden_size], + # "input_ids": [seq_len], + # "verifier_last_hidden_states": [seq_len, hidden_size], + # "loss_mask": [seq_len], + # "lengths": [1], + # "position_ids": [seq_len], + # } + + return { + "hidden_states": torch.empty(0, 3 * hidden_size), + "input_ids": torch.empty(0), + "verifier_last_hidden_states": torch.empty(0, hidden_size), + "loss_mask": torch.empty(0), + "lengths": torch.tensor([0], dtype=torch.long), + "position_ids": torch.arange(0, dtype=torch.long), + } + + def standardize_data_v1(data: dict[str, Any]) -> dict[str, Any]: # v1 data format: # { @@ -113,7 +143,221 @@ def standardize_data_v1(data: dict[str, Any]) -> dict[str, Any]: } -class Eagle3SampleFileDataset(Dataset): +class BaseEagle3Dataset(Dataset): + def __init__( + self, + max_len: int, + transform: TransformTensors | None = None, + hidden_states_dtype=torch.float, + ): + self.max_len = max_len + self.transform = transform + self.hidden_states_dtype = hidden_states_dtype + self.approx_lengths = self._compute_approx_lengths() + + def _compute_approx_lengths(self): + raise NotImplementedError + + def _get_raw_data(self, index): + raise NotImplementedError + + def __getitem__(self, index) -> BatchType | None: + data = self._get_raw_data(index) + + if data is None: + return data + + # data structure: { + # "hidden_states": [seq_len, 3 * hidden_size], + # "input_ids": [seq_len], + # "verifier_last_hidden_states": [seq_len, hidden_size], + # "loss_mask": [seq_len], + # } + + # Convert hidden states to the correct dtype + data = { + k: v.to(self.hidden_states_dtype) if "hidden_states" in k else v + for k, v in data.items() + } + + # Add lengths tensor + seq_len = data["input_ids"].shape[0] + data["lengths"] = torch.tensor([seq_len], dtype=torch.long) + # shape: [1] + + data["position_ids"] = torch.arange(seq_len, dtype=torch.long) + # shape: [seq_len] + + # data structure: { + # "hidden_states": [seq_len, 3 * hidden_size], + # "input_ids": [seq_len], + # "verifier_last_hidden_states": [seq_len, hidden_size], + # "loss_mask": [seq_len], + # "lengths": [1], + # "position_ids": [seq_len], + # } + + # Apply transform + if self.transform: + data = self.transform(data) + + # Note: shift_batch will reduce seq_len by 1 + return shift_batch(data) + + +class Eagle3ArrowDataset(BaseEagle3Dataset): + def __init__( + self, + max_len: int, + datapath: str | PathLike, + hidden_states_path: str | PathLike | None = None, + vllm_endpoint: str = "http://localhost:8000/v1", + on_missing: Literal["generate", "skip", "warn", "raise"] = "generate", + on_generate: Literal["cache", "delete"] = "delete", + split_ratio: float = 1.0, + transform: TransformTensors | None = None, + hidden_states_dtype=torch.float, + model: str | None = None, + ): + """Initialize the Eagle3ArrowDataset. + Args: + max_len: The maximum length of the sequence. + datapath: The path to the data directory that contains the preprocessed + arrow dataset. + transform: The transform to apply to the data. + hidden_states_dtype: The dtype of the hidden states. + """ + self.data = load_from_disk(datapath) + if split_ratio == 1.0: + pass + elif 1.0 > split_ratio > 0: + self.start_file_idx = 0 + split_idx = int(len(self.data) * split_ratio) + self.data = self.data.select(range(split_idx)) + elif -1.0 < split_ratio < 0: + split_idx = int(len(self.data) * (1.0 + split_ratio)) + self.start_file_idx = split_idx + self.data = self.data.select(range(split_idx, len(self.data))) + else: + raise ValueError("split_ratio must be in range (-1.0, 1.0] excluding 0.0.") + + self.hidden_states_path: Path = ( + Path(datapath) / "hidden_states" + if hidden_states_path is None + else Path(hidden_states_path) + ) + self.vllm_endpoint = vllm_endpoint + self.on_missing = on_missing + self.on_generate = on_generate + self.client: openai.OpenAI | None = None + self.model = model + + # Delay super init so that `_compute_approx_lengths` has required data + super().__init__(max_len, transform, hidden_states_dtype) + + def _map_to_file_idx(self, index: int): + return index + self.start_file_idx + + def _setup_client(self): + # Delay client setup so it runs in dataloader thread if on_missing="generate" + self.client = openai.OpenAI(base_url=self.vllm_endpoint, api_key="EMPTY") + list_models = self.client.models.list() + model_id = list_models.data[0].id + if self.model and self.model != model_id: + raise ValueError( + f"An explicit model name was passed ({self.model}) which doesn't match" + "found model_id {model_id}." + "Please make sure --endpoint is set to the correct vllm instance." + ) + self.model = model_id + + def __len__(self): + return len(self.data) + + def _compute_approx_lengths(self) -> list[int]: + """Get lengths of the dataset samples.""" + return list(self.data.with_format(None)["seq_len"]) + + def _maybe_load_hs_file(self, index: int) -> dict[str, torch.Tensor] | None: + file_idx = self._map_to_file_idx(index) + candidate_path = self.hidden_states_path / f"hs_{file_idx}.safetensors" + if candidate_path.exists(): + return load_file(candidate_path) + + return None + + def _maybe_generate_hs(self, index: int) -> dict[str, torch.Tensor] | None: + if not self.client: + self._setup_client() + + input_ids = self.data[index]["input_ids"].tolist() + try: + hs_filepath = generate_hidden_states(self.client, self.model, input_ids) # type:ignore[arg-type] + except InvalidResponseError as e: + warnings.warn(str(e), stacklevel=1) + return None + + loaded_hs = load_file(hs_filepath) + + match self.on_generate: + case "cache": + file_idx = self._map_to_file_idx(index) + target_path = self.hidden_states_path / f"hs_{file_idx}.safetensors" + shutil.move(hs_filepath, target_path) + case "delete": + Path(hs_filepath).unlink() + + return loaded_hs + + def _get_raw_data(self, index): + loaded_hs = self._maybe_load_hs_file(index) + + if loaded_hs is None: + match self.on_missing: + case "generate": + loaded_hs = self._maybe_generate_hs(index) + case "skip": + return None + case "warn": + warnings.warn( + f"Failed to load hidden states for sample {index}. Skipping...", + stacklevel=1, + ) + return None + case "raise": + raise RuntimeError( + f"Failed to load hidden states for sample {index}." + ) + + if loaded_hs is None: + return loaded_hs + + # loaded_hs structure: { + # "hidden_states": [seq_len, 4, hidden_size] + # "token_ids": [seq_len] + # } + + if not torch.equal(loaded_hs["token_ids"], self.data[index]["input_ids"]): + warnings.warn( + f"Loaded token ids {loaded_hs['token_ids']} for index {index} don't" + f"match input ids {self.data[index]['input_ids']}", + stacklevel=1, + ) + return None + + return { + "hidden_states": loaded_hs["hidden_states"][:, :-1].flatten( + 1 + ), # [seq_len, 3 * hidden_size] + "input_ids": loaded_hs["token_ids"], # [seq_len] + "verifier_last_hidden_states": loaded_hs["hidden_states"][ + :, -1 + ], # [seq_len, hidden_size] + "loss_mask": self.data[index]["loss_mask"], # [seq_len] + } + + +class Eagle3SampleFileDataset(BaseEagle3Dataset): def __init__( self, max_len: int, @@ -121,7 +365,6 @@ def __init__( file_list: list[str] | None = None, transform: TransformTensors | None = None, hidden_states_dtype=None, - standardize_fn: StandardizeFnSig = standardize_data_v1, ): """Initialize the Eagle3SampleFileDataset. Args: @@ -139,6 +382,7 @@ def __init__( Note: datapath or file_list must be provided, but not both. """ + if datapath is not None and file_list is not None: raise ValueError( "Either `datapath` or `file_list` must be provided, but " @@ -157,11 +401,9 @@ def __init__( ) self.data: list[str] = file_list - self.max_len = max_len - self.transform = transform - self.standardize_fn = standardize_fn - self.hidden_states_dtype = hidden_states_dtype - self.approx_lengths = self._compute_approx_lengths() + + # Delay super init so that `_compute_approx_lengths` has required data + super().__init__(max_len, transform, hidden_states_dtype) def __len__(self): return len(self.data) @@ -189,7 +431,12 @@ def _compute_approx_lengths(self) -> list[int]: pass # Fallback: approximate lengths from file sizes - lengths_0 = self.__getitem__(0)["lengths"] + item_0 = self.__getitem__(0) + if item_0 is None: + raise ValueError( + "Failed to load first element of datasets for length approximation" + ) + lengths_0 = item_0["lengths"] # this is a single sample so there is only one length lengths_0 = lengths_0[0].item() size_0 = Path(self.data[0]).stat().st_size @@ -199,57 +446,28 @@ def _compute_approx_lengths(self) -> list[int]: for fname in self.data ] - def __getitem__(self, index) -> BatchType: - data = torch.load( - self.data[index], mmap=True, weights_only=True, map_location="cpu" + def _get_raw_data(self, index): + return standardize_data_v1( + torch.load( + self.data[index], mmap=True, weights_only=True, map_location="cpu" + ) ) - data = self.standardize_fn(data) - # data structure: { - # "hidden_states": [seq_len, 3 * hidden_size], - # "input_ids": [seq_len], - # "verifier_last_hidden_states": [seq_len, hidden_size], - # "loss_mask": [seq_len], - # } - - if self.hidden_states_dtype is not None: - # Convert hidden states to the correct dtype - data = { - k: v.to(self.hidden_states_dtype) if "hidden_states" in k else v - for k, v in data.items() - } - # Add lengths tensor - seq_len = data["input_ids"].shape[0] - data["lengths"] = torch.tensor([seq_len], dtype=torch.long) - # shape: [1] - - data["position_ids"] = torch.arange(seq_len, dtype=torch.long) - # shape: [seq_len] - - # data structure: { - # "hidden_states": [seq_len, 3 * hidden_size], - # "input_ids": [seq_len], - # "verifier_last_hidden_states": [seq_len, hidden_size], - # "loss_mask": [seq_len], - # "lengths": [1], - # "position_ids": [seq_len], - # } - - # Apply transform - if self.transform: - data = self.transform(data) - - # Note: shift_batch will reduce seq_len by 1 - return shift_batch(data) +def create_collate_fn(max_len: int, hidden_size: int): + def collate_fn(batch: list[BatchType | None]) -> BatchType: + # Filter failed samples + batch = [b for b in batch if b is not None] + if not batch: + # Create empty sample which then gets padded to full + # batch size if no valid samples are found + batch = [create_empty_sample(hidden_size)] -def create_collate_fn(max_len: int): - def collate_fn(batch: list[BatchType]) -> BatchType: collated_data = {} - for key in batch[0]: + for key in batch[0]: # type: ignore[union-attr] # Concatenate the tensors along the seq (0th) dimension - collated_data[key] = torch.cat([b[key] for b in batch], dim=0) + collated_data[key] = torch.cat([b[key] for b in batch], dim=0) # type: ignore[index] # shape: [total_seq_len, ...] if key != "lengths": diff --git a/src/speculators/train/logger.py b/src/speculators/train/logger.py index 4fe530a76..179c66fe4 100644 --- a/src/speculators/train/logger.py +++ b/src/speculators/train/logger.py @@ -496,6 +496,9 @@ def setup_root_logger(level="INFO"): level=level, format="%(message)s", datefmt="[%X]", handlers=[handler] ) + # Disable verbose HTTP response logs from httpx + logging.getLogger("httpx").propagate = False + def setup_metric_logger(loggers, run_name, output_dir): """Configure the metric logging system with specified backends. diff --git a/src/speculators/train/vocab_mapping.py b/src/speculators/train/vocab_mapping.py index a1641b811..90351c53a 100644 --- a/src/speculators/train/vocab_mapping.py +++ b/src/speculators/train/vocab_mapping.py @@ -6,6 +6,7 @@ import torch from datasets import Dataset as HFDataset from tqdm import tqdm # type: ignore[import-untyped] +from transformers import AutoConfig __all__ = [ "build_vocab_mappings_from_distribution", @@ -15,8 +16,8 @@ def save_token_frequency_distribution( dataset: HFDataset, - output_path: str = "./token_freq.pt", -) -> str: + output_path: Path | str = "./token_freq.pt", +): """Save token frequency distribution from the dataset. Args: @@ -28,7 +29,7 @@ def save_token_frequency_distribution( """ path = Path(output_path) if path.exists(): - return output_path + return token_freq: Counter[int] = Counter() for item in tqdm(dataset, desc="Counting token frequencies"): @@ -44,8 +45,6 @@ def save_token_frequency_distribution( Path(path).parent.mkdir(parents=True, exist_ok=True) torch.save(token_freq_dict, path) - return output_path - def combine_token_frequency_distributions( token_freq_paths: list[str | Path], @@ -95,3 +94,25 @@ def build_vocab_mappings_from_distribution( target_to_draft[selected_token_ids] = True return draft_to_target, target_to_draft + + +def get_target_vocab_size(target_vocab_size, target_model_path): + has_vocab = target_vocab_size is not None + has_model = target_model_path is not None + + if has_vocab and has_model: + raise ValueError("Cannot specify both target-vocab-size and target-model-path") + + if not has_vocab and not has_model: + raise ValueError("Must specify either target-vocab-size or target-model-path") + + if has_vocab: + return target_vocab_size + + config = AutoConfig.from_pretrained(target_model_path) + + # For multimodal models (Qwen3VL, etc.), extract text_config + if hasattr(config, "text_config"): + config = config.text_config + + return config.vocab_size diff --git a/tests/datagen/test_preprocessing.py b/tests/datagen/test_preprocessing.py index 62399414e..4e180fe12 100644 --- a/tests/datagen/test_preprocessing.py +++ b/tests/datagen/test_preprocessing.py @@ -214,7 +214,7 @@ def test_create_loss_mask_simple(): mask = _create_loss_mask_from_offsets(text, offsets, pattern) assert len(mask) == len(offsets) - assert mask.dtype == torch.long + assert mask.dtype == torch.bool # Tokens in assistant responses should have mask = 1 # "Hi there!" is at positions 6-8 (indices in offsets) diff --git a/tests/datagen/test_vllm_hidden_states.py b/tests/datagen/test_vllm_hidden_states.py index ddcccb3e7..4ff26f1ed 100644 --- a/tests/datagen/test_vllm_hidden_states.py +++ b/tests/datagen/test_vllm_hidden_states.py @@ -9,7 +9,9 @@ import torch from transformers import AutoModelForCausalLM, AutoTokenizer -from speculators.data_generation import VllmHiddenStatesGenerator +from speculators.data_generation.vllm_hidden_states_generator import ( + VllmHiddenStatesGenerator, +) logger = logging.getLogger(__name__) diff --git a/tests/unit/train/test_data.py b/tests/unit/train/test_data.py index 34bb3cf52..682096839 100644 --- a/tests/unit/train/test_data.py +++ b/tests/unit/train/test_data.py @@ -128,7 +128,8 @@ def test_standardize_data_v1(): def test_collate_fn_basic(): """Test basic collation functionality.""" max_len = 10 - collate_fn = create_collate_fn(max_len) + hidden_size = 1 + collate_fn = create_collate_fn(max_len, hidden_size) batch = [ { @@ -202,21 +203,22 @@ def test_collate_fn_basic(): def test_collate_fn_length_truncation(): """Test that lengths are truncated when they exceed max_len.""" max_len = 11 - collate_fn = create_collate_fn(max_len) + hidden_size = 8 + collate_fn = create_collate_fn(max_len, hidden_size) batch = [ { "input_ids": torch.arange(5, dtype=torch.long), - "hidden_states": torch.randn(5, 24), - "verifier_last_hidden_states": torch.randn(5, 8), + "hidden_states": torch.randn(5, 3 * hidden_size), + "verifier_last_hidden_states": torch.randn(5, hidden_size), "loss_mask": torch.ones(5, dtype=torch.long), "lengths": torch.tensor([5], dtype=torch.long), "position_ids": torch.arange(5, dtype=torch.long), }, { "input_ids": torch.arange(7, dtype=torch.long), - "hidden_states": torch.randn(7, 24), - "verifier_last_hidden_states": torch.randn(7, 8), + "hidden_states": torch.randn(7, 3 * hidden_size), + "verifier_last_hidden_states": torch.randn(7, hidden_size), "loss_mask": torch.ones(7, dtype=torch.long), "lengths": torch.tensor([7], dtype=torch.long), "position_ids": torch.arange(7, dtype=torch.long), @@ -353,13 +355,11 @@ def test_dataset_getitem_v1_format(tmp_path: Path): torch.save(data, file_path) dataset = Eagle3SampleFileDataset( - max_len=12, - file_list=[str(file_path)], - standardize_fn=standardize_data_v1, - hidden_states_dtype=output_dtype, + max_len=12, file_list=[str(file_path)], hidden_states_dtype=output_dtype ) item = dataset[0] + assert item is not None for key, value in item.items(): assert torch.allclose(value, expected_output[key]), ( @@ -386,11 +386,7 @@ def test_dataset_loads_lengths_from_sample_lengths_json(tmp_path: Path): json.dump(expected_lengths, f) file_list = sorted([str(f) for f in tmp_path.glob("data_*.pt")]) - dataset = Eagle3SampleFileDataset( - max_len=50, - file_list=file_list, - standardize_fn=standardize_data_v1, - ) + dataset = Eagle3SampleFileDataset(max_len=50, file_list=file_list) assert dataset.approx_lengths == [9, 14, 19], ( f"Expected [9, 14, 19], got {dataset.approx_lengths}" @@ -410,11 +406,7 @@ def test_dataset_fallback_when_sample_lengths_json_missing(tmp_path: Path): torch.save(data, tmp_path / "data_0.pt") file_list = [str(tmp_path / "data_0.pt")] - dataset = Eagle3SampleFileDataset( - max_len=50, - file_list=file_list, - standardize_fn=standardize_data_v1, - ) + dataset = Eagle3SampleFileDataset(max_len=50, file_list=file_list) # Should use fallback and return a list with one length assert len(dataset.approx_lengths) == 1 @@ -439,9 +431,5 @@ def test_dataset_fallback_when_sample_lengths_json_malformed(tmp_path: Path): json.dump({"0": 9}, f) file_list = sorted([str(f) for f in tmp_path.glob("data_*.pt")]) - dataset = Eagle3SampleFileDataset( - max_len=50, - file_list=file_list, - standardize_fn=standardize_data_v1, - ) + dataset = Eagle3SampleFileDataset(max_len=50, file_list=file_list) assert len(dataset.approx_lengths) == 2