diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 6ecb660b73..6b89232fbe 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -5,6 +5,7 @@ repos: hooks: # Run the linter. - id: ruff - args: [ --fix ] + args: [ --fix, --config=pyproject.toml ] # Run the formatter. - id: ruff-format + args: [ --config=pyproject.toml ] diff --git a/configs/test.toml b/configs/test.toml deleted file mode 100644 index 5209e96ba7..0000000000 --- a/configs/test.toml +++ /dev/null @@ -1,19 +0,0 @@ -name_model = "debugmodel" -project = "debug_150m_zero_band" - -[train] -micro_bs = 4 # change this base on the gpu - -[data] -seq_length = 8192 -dataset_name_or_paths = "/data/datasets/open-web-math" -dataset_ratio = "100" -num_workers = 1 - -[optim] -batch_size = 128 -warmup_steps = 1000 -total_steps = 88_000 - -[optim.optim] -lr = 4e-4 \ No newline at end of file diff --git a/install.sh b/install.sh index be5e65d9a9..23e8a8506b 100755 --- a/install.sh +++ b/install.sh @@ -57,11 +57,6 @@ main() { log_info "Updating git submodules..." git submodule update --init --recursive - log_info "Downloading data..." - mkdir -p datasets - uv run python scripts/subset_data.py --dataset_name PrimeIntellect/fineweb-edu --data_world_size 1 --data_rank 0 --max_shards 128 - mv fineweb-edu/ datasets/fineweb-edu/ - log_info "Installation completed! You can double check that everything is install correctly by running" } diff --git a/pyproject.toml b/pyproject.toml index 5dcac41736..934deed7f5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,6 +17,7 @@ dependencies = [ "pyarrow", "wandb", "vllm", + "jaxtyping" ] @@ -29,6 +30,7 @@ allow-direct-references = true # allow direct references to git repos in depende [tool.ruff] line-length = 140 +ignore = ["F722", "F821"] [tool.uv] dev-dependencies = ["ruff>=0.5.0", "pre-commit>=3.0.0","pytest>=7.0.0", "faker"] diff --git a/scripts/subset_data.py b/scripts/subset_data.py deleted file mode 100644 index 2e6c648475..0000000000 --- a/scripts/subset_data.py +++ /dev/null @@ -1,120 +0,0 @@ -#!/usr/bin/env python -# coding: utf-8 -# Usage: -# python scripts/subset_data.py --dataset_name PrimeIntellect/fineweb-edu --data_world_size 12 --data_rank 1 - -import argparse -import subprocess -from typing import Dict, List, Optional -import functools -from datasets import load_dataset_builder, BuilderConfig -import logging -from huggingface_hub import get_token -import os -import multiprocessing as mp -from tqdm import tqdm - -logger = logging.getLogger(__name__) -logger.setLevel(logging.DEBUG) -ch = logging.StreamHandler() -ch.setLevel(logging.DEBUG) -formatter = logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s") -ch.setFormatter(formatter) -logger.addHandler(ch) - - -@functools.lru_cache(maxsize=None) -def _get_ds_config_dict(path: str, name: Optional[str] = None) -> Dict[str, BuilderConfig]: - ds_builder = load_dataset_builder(path=path, name=name) - return ds_builder.builder_configs - - -def _get_datafiles(path: str, name: Optional[str] = None, split: str = "train") -> List[str]: - builder_config = _get_ds_config_dict(path=path, name=name) - if name is None: - if "default" not in builder_config: - name = next(iter(builder_config.keys())) - else: - name = "default" - return builder_config[name].data_files[split] - - -def _download_file(data_file: str, save_path: str) -> None: - """Download a file from huggingface.co - - Args: - data_file (str): The file to download. e.g. 'hf://datasets/PrimeIntellect/fineweb-edu@14efaa24d7dff8a745bf4918e415878546542346/data1/train-00450.parquet' - save_path (str): The path to save the file. e.g. 'data1/train-00450.parquet' - """ - assert data_file.startswith("hf://") - data_file = data_file.replace("hf://", "").replace("@", "/resolve/") - - if "/" in save_path: - parent = "/".join(save_path.split("/")[:-1]) - if not os.path.exists(parent): - logger.debug(f"Creating directory: {parent}") - os.makedirs(parent, exist_ok=True) - - cmd = [ - "wget", - f'--header="Authorization: Bearer {get_token()}"', - f"https://huggingface.co/{data_file}?download=true", - f"-O {save_path}", - ] - result = subprocess.run(" ".join(cmd), shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) - if result.returncode != 0: - logger.error(f"Error downloading file: {data_file}") - logger.error(result.stderr.decode("utf-8")) - - -def _download_file_wrapper(args): - return _download_file(*args) - - -def _get_save_path(data_file: str) -> str: - ret_list = data_file.split("@")[-1].split("/")[1:] - return args.dataset_name.split("/")[-1] + "/" + "/".join(ret_list) - - -def main(args): - g_data_files = _get_datafiles(args.dataset_name) - logger.debug(f"Length of data_files: {len(g_data_files)}") - if len(args.filter) > 0: - args.filter = args.filter.split(",") - data_files = [] - for _filter in args.filter: - data_files.extend([f for f in g_data_files if _filter in f]) - else: - data_files = g_data_files - - logger.debug(f"Length of data_files: {len(data_files)}") - data_files = data_files[args.data_rank :: args.data_world_size][: args.max_shards] - logger.debug(f"Data files: {data_files}") - logger.debug(f"Length of data_files processing: {len(data_files)}") - - if args.dry_run: - return - - with mp.Pool(args.num_workers) as pool: - save_paths = list(pool.imap(_get_save_path, tqdm(data_files, desc="Getting save paths"))) - _ = list( - tqdm( - pool.imap(_download_file_wrapper, zip(data_files, save_paths)), - desc="Downloading files", - total=len(data_files), - bar_format="{l_bar}{bar:10}{r_bar}", - ) - ) - - -if __name__ == "__main__": - parser = argparse.ArgumentParser(description="Download and process data from a HF dataset") - parser.add_argument("--dataset_name", type=str, default="PrimeIntellect/fineweb-edu", help="dataset name") - parser.add_argument("--dry_run", action="store_true", help="do not download data") - parser.add_argument("--filter", type=str, default="", help="search shards by the filter") - parser.add_argument("--data_rank", type=int, default=0, help="start index") - parser.add_argument("--data_world_size", type=int, default=4, help="world size") - parser.add_argument("--max_shards", type=int, default=1000) - parser.add_argument("--num_workers", type=int, default=12) - args = parser.parse_args() - main(args) diff --git a/src/zeroband/train.py b/src/zeroband/train.py index 2892ed9c74..a904e67fd6 100644 --- a/src/zeroband/train.py +++ b/src/zeroband/train.py @@ -10,6 +10,7 @@ from zeroband.models import ModelName, get_model_and_tokenizer from zeroband.training.checkpoint import TrainingProgress, load_checkpoint_fsdp_state, save_checkpoint_fsdp_state from zeroband.training.data import DataConfig, get_dataloader +from zeroband.training.loss import grpo_loss from zeroband.training.lr_scheduler import get_scheduler from zeroband.training.utils import ( PerfCounter, @@ -21,7 +22,7 @@ from zeroband.logger import get_logger from pydantic_config import BaseConfig, parse_argv -import torch.nn.functional as F +from jaxtyping import Float, Int from zeroband.training.world_info import get_world_info @@ -96,8 +97,7 @@ def train(config: Config): model, model_config, tokenizer = get_model_and_tokenizer(config.name_model) - # fmt: skip - train_dataloader = get_dataloader(tokenizer=tokenizer,world_size=world_info.world_size,rank=world_info.rank,batch_size=config.train.micro_bs,data_config=config.data) # fmt: skip + train_dataloader = get_dataloader(tokenizer=tokenizer, batch_size=config.train.micro_bs, data_config=config.data) train_dataloader_iterator = iter(train_dataloader) @@ -140,20 +140,15 @@ def train(config: Config): model.set_requires_gradient_sync(not is_accumulating) batch = next(train_dataloader_iterator) - input_ids = batch["input_ids"].to("cuda") - labels = batch["labels"].to("cuda") - logits = model(input_ids=input_ids).logits.contiguous() - flatten_logits = logits.reshape(-1, logits.size(-1)) # b seq vocab -> (b * seq) vocab - flatten_labels = labels.reshape(-1) # b seq -> (b * seq) + input_ids: Int[torch.Tensor, "batch seq"] = batch["input_ids"].to("cuda") + advantages: Float[torch.Tensor, "batch"] = batch["advantages"].to("cuda") + ref_logprobs: Float[torch.Tensor, "batch seq"] = batch["ref_logprobs"].to("cuda") - ce_loss = F.cross_entropy(flatten_logits, flatten_labels) + policy_logprobs = model(input_ids=input_ids).logits.contiguous() - del logits - del flatten_logits - del flatten_labels + loss = grpo_loss(policy_logprobs, ref_logprobs, advantages) / gradient_accumulation_steps - loss = ce_loss / gradient_accumulation_steps loss.backward() loss_batch += loss.detach().clone() diff --git a/src/zeroband/training/data.py b/src/zeroband/training/data.py index 5eb9efe2c9..dbac73ffc1 100644 --- a/src/zeroband/training/data.py +++ b/src/zeroband/training/data.py @@ -1,38 +1,20 @@ -from dataclasses import dataclass, asdict -import random -from typing import Any, Generator, Optional, List, Dict, TypedDict, Union -import functools +from typing import Any, Generator, TypedDict from pydantic_config import BaseConfig -from zeroband.logger import get_logger import torch -from torch.utils.data import IterableDataset, Dataset +from torch.utils.data import IterableDataset from torchdata.stateful_dataloader import StatefulDataLoader -from torch.distributed.checkpoint.stateful import Stateful -from datasets import load_dataset_builder, BuilderConfig -from pyarrow import parquet as pq -from transformers import PreTrainedTokenizer - - -TEST_VOCAB_SIZE = 1024 +from jaxtyping import Float, Int class DataConfig(BaseConfig): dataset_name_or_paths: str = "datasets/fineweb-edu" - val_dataset_name_or_paths: str | None = None seq_length: int = 1024 fake: bool = False num_workers: int = 4 - max_train_samples: int | None = None - max_eval_samples: int | None = None - dataset_ratio: str | None = None - data_rank: int | None = None - data_world_size: int | None = None - reverse_data_files: bool = False - split_by_data_rank: bool = True class FakeTokenizedDataset(IterableDataset): @@ -46,10 +28,13 @@ def __init__(self, seq_len: int, vocab_size: int): def __iter__(self) -> Generator[dict[str, Any], Any, None]: while True: - len_ = random.randint(1, self.seq_len) - input_ids = torch.randint(3, self.vocab_size, (len_,)).tolist() + len_ = self.seq_len + input_ids = torch.randint(3, self.vocab_size, (len_,)) + advantages = torch.randn(len_) + ref_logprobs = torch.randn(len_, self.vocab_size) self.step += 1 - yield {"input_ids": input_ids} + yield {"input_ids": input_ids, "advantages": advantages, "ref_logprobs": ref_logprobs} + # yield {"input_ids": input_ids} def state_dict(self): return {"step": self.step} @@ -62,364 +47,16 @@ def load_state_dict(self, state_dict): class BatchOutput(TypedDict): - input_ids: torch.IntTensor - labels: torch.IntTensor - seqlens: list[int] - - -@dataclass -class SequencePackingDataSetState: - inputs_ids: list[int] - labels: list[int] - seqlens: list[int] - - -class SequencePackingDataSet(IterableDataset, Stateful): - """ - This class wrap a dataset and wrap it into an iterable that return sequence of max_seq_length - packed - """ - - def __init__(self, dataset: Dataset, max_seq_length: int, eos_token: int): - self.dataset = dataset - self.max_seq_length = max_seq_length - self.eos_token = eos_token - - self.state = SequencePackingDataSetState(inputs_ids=[], labels=[], seqlens=[]) - - def __iter__(self) -> Generator[BatchOutput, Any, None]: - for og_sample in self.dataset: - og_sample: list[int] = og_sample["input_ids"] - - og_sample = og_sample + [self.eos_token] - sample_inputs_ids = og_sample[:-1] - sample_labels = og_sample[1:] - - token_remaining = self.max_seq_length - len(self.state.inputs_ids) - - if len(sample_inputs_ids) < token_remaining: - self.state.inputs_ids.extend(sample_inputs_ids) - self.state.labels.extend(sample_labels) - self.state.seqlens.append(len(sample_inputs_ids)) - - else: - self.state.inputs_ids.extend(sample_inputs_ids[:token_remaining]) - self.state.labels.extend(sample_labels[:token_remaining]) - self.state.seqlens.append(token_remaining) - - data = { - "input_ids": torch.Tensor(self.state.inputs_ids).to(dtype=torch.long), - "labels": torch.Tensor(self.state.labels).to(dtype=torch.long), - "seqlens": self.state.seqlens, - } - self.state.inputs_ids = [] - self.state.labels = [] - self.state.seqlens = [] - - yield data - - def state_dict(self): - return {"dataset": self.dataset.state_dict(), "state": asdict(self.state)} - - def load_state_dict(self, state_dict): - self.dataset.load_state_dict(state_dict["dataset"]) - self.state = SequencePackingDataSetState(**state_dict["state"]) - - -def collate_fn(samples: list[dict[str, torch.LongTensor]]) -> dict[str, torch.LongTensor | list[torch.LongTensor]]: - assert samples[0].keys() == {"input_ids", "labels", "seqlens"} - - inputs_ids = [] - labels = [] - seqlens = [] - - for sample in samples: - inputs_ids.append(sample["input_ids"]) - labels.append(sample["labels"]) - - seqlens.append(torch.Tensor(sample["seqlens"]).long()) - - return { - "input_ids": torch.stack(inputs_ids, dim=0), - "labels": torch.stack(labels, dim=0), - "seqlens": seqlens, - } - - -@dataclass -class PQDatasetState: - files: List[str] - file_index: int - row_index: int - increment: int - init_row_index: int - - -class ParquetDataset(IterableDataset, Stateful): - """ - this class is a wrapper around a parquet dataset compatible with datasets and statefull compatible. The dataset is infinite and will restart from the last state if the iterator is exhausted. - TODO: - * [ ] handle mutli proc dataloader pytorch - """ - - def __init__(self, files: List[str], tokenizer: PreTrainedTokenizer): - self.arg_files = files - self.tokenizer = tokenizer - - self.state = None - - def _lazy_init(self): - worker_info = torch.utils.data.get_worker_info() - if worker_info is not None: - if worker_info.num_workers > len(self.arg_files): - get_logger().warning( - f"dataloader rank {worker_info.id} Number of workers {worker_info.num_workers} is greater than the number of files {len(self.arg_files)}" - ) - self.state = PQDatasetState( - files=self.arg_files, - file_index=0, - row_index=worker_info.id, - increment=worker_info.num_workers, - init_row_index=worker_info.id, - ) - return - - files = self.arg_files[worker_info.id :: worker_info.num_workers] - else: - files = self.arg_files - - self.state = PQDatasetState(files=files, file_index=0, row_index=0, increment=1, init_row_index=0) - - def __iter__(self): - # we lazy init the parquet dataset to get the worker info from dataloader multi process - if self.state is None: - self._lazy_init() - - while True: - file = self.state.files[self.state.file_index] - - parquet_file = pq.ParquetFile(file) - table = parquet_file.read()["text"] - - while True: - row = table[self.state.row_index] - - self.state.row_index += self.state.increment - if self.state.row_index >= len(table): - self.state.row_index = self.state.init_row_index - self.state.file_index += 1 - if self.state.file_index >= len(self.state.files): # infinite datasets - self.state.file_index = 0 - - yield {"input_ids": self.tokenizer.encode(str(row))} - - @property - def is_empty(self): - return len(self.arg_files) == 0 - - def state_dict(self) -> dict[str, Any]: - return asdict(self.state) if self.state is not None else {} - - def load_state_dict(self, state_dict): - self.state = PQDatasetState(**state_dict) - + input_ids: Int[torch.Tensor, "batch seq"] + advantages: Float[torch.Tensor, "batch"] + ref_logprobs: Float[torch.Tensor, "batch seq vocab"] -@dataclass -class InterleaveDatasetState: - current_index: int - seed: int - -class InterleaveDataset(IterableDataset, Stateful): - """This class take a list of datasets and interleave them. It is stateful and can be used with pytorch dataloader. - - It draw a sample from each dataset with a probability given by the probabilities list. - - The state can be saved and restored. Under the hood we just fast forward the random generator to the current position. - """ - - def __init__(self, datasets: List[ParquetDataset], probabilities: List[float], seed: int = 42): - assert len(datasets) > 0, "At least one dataset is required" - assert len(datasets) == len(probabilities), "The number of datasets and probabilities must be the same" - - self.probabilities = [] - self.datasets = [] - - for dataset, prob in zip(datasets, probabilities): - if not dataset.is_empty: - self.datasets.append(dataset) - self.probabilities.append(prob) - else: - get_logger().warning(f"Dataset {dataset} is empty. Skipping.") - - self.state = InterleaveDatasetState(current_index=0, seed=seed) - self._init_random_state() - - def _init_random_state(self): - """Initialize random generator and advance to current position""" - ... - self.random_generator = random.Random(self.state.seed) - # Advance the RNG to the current position - for _ in range(self.state.current_index): - self._get_dataset_to_yield_from() - - def _get_dataset_to_yield_from(self) -> int: - return self.random_generator.choices(range(len(self.datasets)), weights=self.probabilities, k=1)[0] - - def __iter__(self): - data_iters = [iter(dataset) for dataset in self.datasets] - while True: - dataset_to_yield_from = self._get_dataset_to_yield_from() - - sample = next(data_iters[dataset_to_yield_from]) - self.state.current_index += 1 - - yield sample - - def state_dict(self): - state = {"interleave_state": asdict(self.state)} - - for i, dataset in enumerate(self.datasets): - state[f"dataset_{i}"] = dataset.state_dict() - return state - - def load_state_dict(self, state_dict): - self.state = InterleaveDatasetState(**state_dict["interleave_state"]) - for i, dataset in enumerate(self.datasets): - dataset.load_state_dict(state_dict[f"dataset_{i}"]) - self._init_random_state() - - -def get_dataloader( - tokenizer, - world_size: int, - rank: int, - batch_size: int, - data_config: DataConfig, -) -> StatefulDataLoader: +def get_dataloader(tokenizer, batch_size: int, data_config: DataConfig) -> StatefulDataLoader[BatchOutput]: + """Get a dataloader for the training dataset""" if data_config.fake: - train_dataset = FakeTokenizedDataset(data_config.seq_length, TEST_VOCAB_SIZE) + train_dataset = FakeTokenizedDataset(data_config.seq_length, len(tokenizer)) else: - train_dataset = load_all_datasets(data_config=data_config, split="train", tokenizer=tokenizer, rank=rank, world_size=world_size) - - dataset = SequencePackingDataSet(train_dataset, data_config.seq_length, eos_token=tokenizer.eos_token_id) - - return StatefulDataLoader( - dataset, - batch_size=batch_size, - collate_fn=collate_fn, - num_workers=data_config.num_workers, - ) - - -@functools.lru_cache(maxsize=None) -def _get_ds_config_dict(path: str, name: Optional[str] = None) -> Dict[str, BuilderConfig]: - ds_builder = load_dataset_builder(path=path, name=name) - return ds_builder.builder_configs - - -def _get_datafiles(path: str, name: Optional[str] = None, split: str = "train") -> List[str]: - builder_config = _get_ds_config_dict(path=path, name=name) - if name is None or len(name) == 0: - if "default" not in builder_config: - get_logger().warning(f"Default config not found for {path}. Using first config.") - name = next(iter(builder_config.keys())) - else: - name = "default" - return builder_config[name].data_files[split] - - -def _nice_print(kwargs: Dict[str, Union[str, List[str]]]) -> str: - def _foo(a): - if isinstance(a, list): - return str(a[:5]) + "..." + str(a[-5:]) if len(a) > 10 else str(a) - return str(a) - - return str({k: _foo(v) for k, v in kwargs.items()}) - - -def _load_datasets( - dataset_names: str, - split: str, - tokenizer: PreTrainedTokenizer, - data_rank: Optional[int] = None, - data_world_size: Optional[int] = None, - streaming: bool = True, - probabilities: Optional[List[float]] = None, - reverse_data_files: bool = False, -) -> InterleaveDataset: - get_logger().debug(dataset_names) - ds_args = [] - for _ds in dataset_names.split(","): - _ds_name, _, _ds_config = _ds.partition(":") - _ds_args: dict[str, Any] = {"path": _ds_name} - if _ds_config: - _ds_args["name"] = _ds_config - _data_files = _get_datafiles(_ds_name, _ds_config, split) - if reverse_data_files: - _data_files = _data_files[::-1] - _ds_args["data_files"] = _data_files - if data_rank is not None and data_world_size is not None: - _ds_args["data_files"] = _data_files[data_rank::data_world_size] - - ds_args.append(_ds_args) - - # logger.debug(f"Datasets ({split}):\n" + "\n".join(map(_nice_print, ds_args))) - # logger.debug(f"Probabilities: {probabilities}") - get_logger().debug(f"Loading datasets{' in streaming mode' if streaming else ''}") - datasets = [] - for ds_arg in ds_args: - # logger.debug(f"Loading dataset: {ds_arg['data_files']}") - _ds = ParquetDataset(files=ds_arg["data_files"], tokenizer=tokenizer) - datasets.append(_ds) - - if len(datasets) > 1: - ds = InterleaveDataset(datasets=datasets, probabilities=probabilities) - else: - ds = datasets[0] - - get_logger().info(f"Loaded datasets ({split})") - return ds - - -def _get_probabilities(data_config: DataConfig) -> Optional[List[float]]: - if data_config.dataset_ratio is None: - return None - if len(data_config.dataset_name_or_paths.split(",")) != len(data_config.dataset_ratio.split(":")): - raise ValueError("Number of datasets and dataset ratios must be the same") - nums = [float(i) for i in data_config.dataset_ratio.split(":")] - denom = sum(nums) - return [i / denom for i in nums] - - -def load_all_datasets( - data_config: DataConfig, - split: str, - tokenizer: PreTrainedTokenizer, - rank: int, - world_size: int, -) -> InterleaveDataset: - """Load all datasets and interleave them""" - - if data_config.split_by_data_rank and (data_config.data_rank is not None and data_config.data_world_size is not None): - split_rank = data_config.data_rank * world_size + rank - split_world_size = data_config.data_world_size * world_size - else: - split_rank = rank - split_world_size = world_size - - get_logger().info("Loading Train dataset(s)") - - ds = _load_datasets( - dataset_names=data_config.dataset_name_or_paths, - split=split, - data_rank=split_rank, - data_world_size=split_world_size, - probabilities=_get_probabilities(data_config), - reverse_data_files=data_config.reverse_data_files, - tokenizer=tokenizer, - ) - - get_logger().info(f"Train dataset: {ds}") + train_dataset = FakeTokenizedDataset(data_config.seq_length, len(tokenizer)) - return ds + return StatefulDataLoader(train_dataset, batch_size=batch_size, num_workers=data_config.num_workers) diff --git a/src/zeroband/training/loss.py b/src/zeroband/training/loss.py new file mode 100644 index 0000000000..f30f018c58 --- /dev/null +++ b/src/zeroband/training/loss.py @@ -0,0 +1,38 @@ +import torch +from torch import Tensor +from jaxtyping import Float + + +@torch.compile +def grpo_loss( + policy_logprobs: Float[Tensor, "batch seq"], + ref_logprobs: Float[Tensor, "batch seq"], + advantages: Float[Tensor, "batch"], + beta: float = 0.04, + epsilon: float = 0.2, +): + """ + DeepSeek Math Loss: https://arxiv.org/abs/2402.03300 + """ + + # Expand advantages to match sequence dimension + advantages = advantages.unsqueeze(-1) # [batch_size, 1] + + # Policy ratio + ratio = torch.exp(policy_logprobs - ref_logprobs) + clipped_ratio = torch.clamp(ratio, 1 - epsilon, 1 + epsilon) + + # Policy loss + policy_loss = -torch.min(ratio * advantages, clipped_ratio * advantages) + + # KL penalty (unbiased estimator) + kl_div = ref_logprobs / policy_logprobs - torch.log(ref_logprobs / policy_logprobs) - 1 + + # Reduce across sequence length + policy_loss = policy_loss.mean(dim=-1) + kl_penalty = kl_div.mean(dim=-1) + + # Final loss (mean across batch) + loss = (policy_loss + beta * kl_penalty).mean() + + return loss diff --git a/tests/units/test_loss.py b/tests/units/test_loss.py new file mode 100644 index 0000000000..16cc3e1ce8 --- /dev/null +++ b/tests/units/test_loss.py @@ -0,0 +1,11 @@ +from zeroband.training.loss import grpo_loss +import torch + + +def test_grpo_loss(): + policy_logprobs = torch.randn(10, 10).cuda() + ref_logprobs = torch.randn(10, 10).cuda() + advantages = torch.randn(10).cuda() + loss = grpo_loss(policy_logprobs, ref_logprobs, advantages) + assert loss.shape == () + assert loss.item() is not None diff --git a/uv.lock b/uv.lock index a4d38a3be7..b94b85d3af 100644 --- a/uv.lock +++ b/uv.lock @@ -623,6 +623,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ef/a6/62565a6e1cf69e10f5727360368e451d4b7f58beeac6173dc9db836a5b46/iniconfig-2.0.0-py3-none-any.whl", hash = "sha256:b6a85871a79d2e3b22d2d1b94ac2824226a63c6b741c88f7ae975f18b6778374", size = 5892 }, ] +[[package]] +name = "jaxtyping" +version = "0.2.38" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "wadler-lindig" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/34/a5/83fbf2ed24f8bd9af80536b3139e9c9cb8fb096d6ceeb28965b847fae9ae/jaxtyping-0.2.38.tar.gz", hash = "sha256:84d509341437189e82d7dbb59a2970435724851ca79fd8550e886cd37c048333", size = 45785 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/db/7e/da7b57a1f3af7303a0f3c8594d820fc0d3a9bbe3810a357eb21eb166e76b/jaxtyping-0.2.38-py3-none-any.whl", hash = "sha256:bc209ab8ec29917b6f0c7dec4a8ea1fc276f7d94f25b71c01d1243ec2b21ae12", size = 56375 }, +] + [[package]] name = "jinja2" version = "3.1.5" @@ -2354,6 +2366,15 @@ dependencies = [ ] sdist = { url = "https://files.pythonhosted.org/packages/f9/ed/1e4625f290e24c1381aeeb66620b32755b17f2894a4227fb49fb20fa0ec9/vllm-0.1.2.tar.gz", hash = "sha256:6fff59a918ec822c3d7a9060c3af5604a1bbf26ec4db15656678768019aefad7", size = 93849 } +[[package]] +name = "wadler-lindig" +version = "0.1.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/7a/a2/e8fc843e14aec55d18572a93c5443cc89a6b7d537b90d804f7b373301d9f/wadler_lindig-0.1.3.tar.gz", hash = "sha256:476fb7015135f714cef8f8eac7c44b164c8b993345e651a9b6f25b7b112440c9", size = 15197 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/39/3b/5b918a0da0d6920e7f7328cf0ab00df31b905d709f458596304f09096785/wadler_lindig-0.1.3-py3-none-any.whl", hash = "sha256:3018e4e6b115a7ef21c77414a41cbe7e03e83f6b5e25004958e33432a17f3c94", size = 20140 }, +] + [[package]] name = "wandb" version = "0.19.6" @@ -2582,6 +2603,7 @@ version = "0.1.0" source = { editable = "." } dependencies = [ { name = "datasets" }, + { name = "jaxtyping" }, { name = "ninja" }, { name = "numpy" }, { name = "pyarrow" }, @@ -2606,6 +2628,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "datasets", specifier = ">=3.0.0" }, + { name = "jaxtyping" }, { name = "ninja" }, { name = "numpy" }, { name = "pyarrow" },