Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 6 additions & 2 deletions src/zeroband/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
import concurrent.futures
import time
from toploc.utils import sha256sum
from safetensors import safe_open

# from vllm.model_executor.model_loader
from vllm.model_executor.model_loader.loader import _process_weights_after_loading
Expand Down Expand Up @@ -173,7 +174,10 @@ def reload_model_weights(llm: LLM, ckpt_path: str):
# Access the internal model from vLLM
model = llm.llm_engine.model_executor.driver_worker.model_runner.model
# Load state dict
state_dict = torch.load(ckpt_path, map_location="cpu", weights_only=True)
state_dict = {}
with safe_open(ckpt_path, framework="pt", device="cpu") as f:
for key in f.keys():
state_dict[key] = f.get_tensor(key)

# Create a better weight iterator that filters out empty keys and handles prefixes
def weights_iterator():
Expand Down Expand Up @@ -291,7 +295,7 @@ def logits_processor_hook(module, input):
stable_file = last_step / "stable"
if stable_file.exists():
logger.info(f"Reloading model weights from {config.rollout_path} step {maybe_new_step}")
llm = reload_model_weights(llm, Path(config.rollout_path) / f"step_{maybe_new_step}/model.pt")
llm = reload_model_weights(llm, Path(config.rollout_path) / f"step_{maybe_new_step}/model.safetensors")
ckpt_step = maybe_new_step
total_problems = 0
total_tokens = 0
Expand Down
93 changes: 93 additions & 0 deletions src/zeroband/shardcast_downloader.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
import argparse
from shardcast import ClientNode
import time
import logging
from pathlib import Path
import multiprocessing as mp

POLL_INTERVAL = 10
logger = logging.getLogger(__name__)


def main(servers: list[str], output_dir: Path, versions_to_keep: int = -1, backlog_version: int = -1):
"""
Download the latest version of the model from the servers and delete expired versions.
Versions will be saved as v{version}.safetensors in the output directory.

Args:
servers: list of servers to download from
output_dir: directory to save the downloaded files
versions_to_keep: number of versions to keep
backlog_version: version to attempt to get first
"""
client = ClientNode(servers, str(output_dir))

while True:
# 1. Pick the version to download
available_versions = sorted([int(x[1:]) for x in client.list_available_versions().keys()])
version = available_versions[-1]
if version <= backlog_version:
backlog_version = -1
if backlog_version != -1:
version = backlog_version
backlog_version += 1
safetensors_filepath = output_dir / f"step_{version}/model.safetensors"

# 2. Check if the version exists
if version not in available_versions:
logger.warning(f"Version {version} not found")
time.sleep(POLL_INTERVAL)
continue
if safetensors_filepath.exists():
logger.info(f"Version {version} already exists")
time.sleep(POLL_INTERVAL)
continue

# 3. Download the version
logger.info(f"Downloading version {version}")
start = time.time()
filepath = client.download_version(f"v{version}", str(safetensors_filepath))
logger.info(f"Downloaded in {time.time() - start} seconds")

# 4. Make stable file if successful and delete expired versions
if filepath is not None:
(output_dir / f"step_{version}/stable").touch()

if versions_to_keep != -1:
try:
logger.info(f"Deleting expired version {version - versions_to_keep}")
expired_version = output_dir / f"step_{version - versions_to_keep}/model.safetensors"
expired_version.unlink()
(output_dir / f"step_{version - versions_to_keep}/stable").unlink()
except FileNotFoundError:
logger.warning(f"Expired version {version - versions_to_keep} not found")
except Exception as e:
logger.warning(f"Error deleting expired version {version - versions_to_keep}: {e}")


def run_main_bg(servers: list[str], output_dir: Path, versions_to_keep: int = -1, backlog_version: int = -1) -> mp.Process:
"""
Run the main function in a background process.

Args:
servers: list of servers to download from
output_dir: directory to save the downloaded files
versions_to_keep: number of versions to keep
backlog_version: version to attempt to get first

Returns:
mp.Process: The created process running the main function
"""
process = mp.Process(target=main, args=(servers, output_dir, versions_to_keep, backlog_version))
process.start()
return process


if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--servers", type=str, nargs="+", required=True)
parser.add_argument("--output_dir", type=str, required=True)
parser.add_argument("--versions-to-keep", type=int, default=-1)
parser.add_argument("--backlog-version", type=int, default=-1)
args = parser.parse_args()
main(args.servers, Path(args.output_dir), args.versions_to_keep, args.backlog_version)
49 changes: 23 additions & 26 deletions src/zeroband/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
from zeroband.models import AttnImpl, ModelName, ModelType, get_model_and_tokenizer
from zeroband.training.checkpoint import TrainingProgress, load_checkpoint_fsdp_state, save_checkpoint_fsdp_state, save_ckpt_for_rollout
from zeroband.training.data import DataConfig, get_dataloader
from zeroband.training.loss import grpo_loss, selective_log_softmax, entropy_loss
from zeroband.training.loss import entropy_loss, grpo_loss, log_prob_from_logits
from zeroband.training.lr_scheduler import get_scheduler
from zeroband.training.utils import PerfCounter, apply_ac_ckpt

Expand Down Expand Up @@ -90,6 +90,8 @@ class Config(BaseConfig):
on_policy_log_prob: bool = False
max_async_level: int = 2 # the amount of rollout checkpoints to keep

entropy_loss_coeff: float = 0.001

@model_validator(mode="after")
def check_liger(self):
if self.train.liger_qwen:
Expand Down Expand Up @@ -214,28 +216,22 @@ def train(config: Config):

# here we want to pre-compute the logprobs with the model before update
with torch.no_grad():
if config.on_policy_log_prob:
data = []

for rollout_step in range(config.optim.step_per_rollout):
for grad_acc_step in range(gradient_accumulation_steps):
batch = next(train_dataloader_iterator)
input_ids = batch["input_ids"].to("cuda")
data = []

logits: Float[torch.Tensor, "batch seq vocab"] = model(input_ids=input_ids).logits.contiguous()
for rollout_step in range(config.optim.step_per_rollout):
for grad_acc_step in range(gradient_accumulation_steps):
batch = next(train_dataloader_iterator)
input_ids = batch["input_ids"].to("cuda")

input_ids = input_ids[:, 1:]
logits = logits[:, :-1, :] / config.temperature
logits: Float[torch.Tensor, "batch seq vocab"] = model(input_ids=input_ids).logits.contiguous()

per_token_logps = selective_log_softmax(logits, input_ids)
batch["logprobs"] = per_token_logps.to("cpu")
log_probs = log_prob_from_logits(logits, input_ids, config.temperature)
batch["logprobs"] = log_probs.to("cpu")

del logits, per_token_logps
data.append(batch)
del logits, log_probs
data.append(batch)

logprobs_aware_iterator = iter(data)
else:
logprobs_aware_iterator = train_dataloader_iterator
logprobs_aware_iterator = iter(data)

for rollout_step in range(config.optim.step_per_rollout):
loss_batch = 0
Expand Down Expand Up @@ -272,8 +268,6 @@ def train(config: Config):
advantages = batch["advantages"].to("cuda")
loss_mask = loss_mask.to("cuda")
original_logprobs = batch["logprobs"].to("cuda")
if not config.on_policy_log_prob:
original_logprobs = original_logprobs[:, 1:]

# Loss
pg_loss, clip_ratio = grpo_loss(
Expand All @@ -283,17 +277,20 @@ def train(config: Config):

loss = pg_loss - config.entropy_loss_coeff * entropy
loss = loss / gradient_accumulation_steps
clip_ratio = clip_ratio / gradient_accumulation_steps

del batch, logits, input_ids, advantages, loss_mask, original_logprobs
clip_ratio = clip_ratio / gradient_accumulation_steps

# Backward
loss.backward()
loss_batch += loss.detach().clone()
pg_loss_batch += (pg_loss / gradient_accumulation_steps).detach().clone()
entropy_loss_batch += (entropy / gradient_accumulation_steps).detach().clone()
clip_ratio_batch += clip_ratio.detach().clone()
del loss, clip_ratio, pg_loss, entropy

del batch, logits, input_ids, advantages, loss_mask, original_logprobs, pg_loss, entropy, clip_ratio

# Backward
loss.backward()
loss_batch += loss.detach().clone()

del loss, clip_ratio, pg_loss, entropy

dist.all_reduce(tensor=loss_batch, op=dist.ReduceOp.AVG)
dist.all_reduce(tensor=pg_loss_batch, op=dist.ReduceOp.AVG)
Expand Down
6 changes: 4 additions & 2 deletions src/zeroband/training/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ def load_checkpoint_fsdp_state(
scheduler.load_state_dict(state["scheduler"])


def save_ckpt_for_rollout(model: ModelType, path: Path) -> Path:
def save_ckpt_for_rollout(model: ModelType, path: Path, dtype: torch.dtype = torch.bfloat16) -> Path:
"""
Save the checkpoint for rollout as one unified safetensors file.

Expand All @@ -112,7 +112,9 @@ def save_ckpt_for_rollout(model: ModelType, path: Path) -> Path:

# Only save on rank 0
if torch.distributed.get_rank() == 0:
save_file(state, path_file)
for key, value in state.items():
state[key] = value.to(dtype)
save_file(state, path_file, metadata={"format": "pt"})

stable_file = path / "stable"
stable_file.touch()
Expand Down
115 changes: 43 additions & 72 deletions src/zeroband/training/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,8 @@
from torch import Tensor
from jaxtyping import Float, Int, jaxtyped
from beartype import beartype as typechecker
import torch.nn.functional as F

from zeroband.training.verl_utils import logprobs_from_logits, masked_mean


# beartype here just make sure we have the correct shape
Expand All @@ -27,89 +28,59 @@ def grpo_loss(
epsilon: Clipping parameter for PPO
ignore_index: Specifies a target value that is ignored and does not contribute to the loss
"""
return _compile_grpo_loss(
logits=logits,
input_ids=input_ids,
advantages=advantages,
original_logprobs=original_logprobs,
loss_mask=loss_mask,
temperature=temperature,
epsilon=epsilon,
advantages = advantages[:, 1:]
loss_mask = loss_mask[:, 1:]

log_probs = log_prob_from_logits(logits, input_ids, temperature)

pg_loss, pg_clipfrac, _ = compute_policy_loss(
old_log_prob=original_logprobs, log_prob=log_probs, advantages=advantages, eos_mask=loss_mask, cliprange=epsilon
)
return pg_loss, pg_clipfrac


def selective_log_softmax(logits, index):
"""
credits to https://github.com/huggingface/trl/blob/07cfe1677e552b7d5c92b7740e5b2f0b057661d8/trl/trainer/utils.py#L1659
def log_prob_from_logits(logits, input_ids, temperature):
input_ids = input_ids[:, 1:]

logits.div_(temperature)
response_length = logits.shape[1]
logits = logits[:, -response_length - 1 : -1] # (bsz, response_length)
log_probs = logprobs_from_logits(logits, input_ids)
return log_probs

A memory-efficient implementation of the common `log_softmax -> gather` operation.

This function is equivalent to the following naive implementation:
```python
logps = torch.gather(logits.log_softmax(-1), dim=-1, index=index.unsqueeze(-1)).squeeze(-1)
```
def compute_policy_loss(old_log_prob, log_prob, advantages, eos_mask, cliprange):
"""Adapted from https://github.com/huggingface/trl/blob/main/trl/trainer/ppo_trainer.py#L1122

Args:
logits (`torch.Tensor`):
Logits tensor of shape `(..., num_classes)`.
index (`torch.Tensor`):
Index tensor of shape `(...)`, specifying the positions to gather from the log-softmax output.
old_log_prob: `(torch.Tensor)`
shape: (bs, response_length)
log_prob: `(torch.Tensor)`
shape: (bs, response_length)
advantages: `(torch.Tensor)`
shape: (bs, response_length)
eos_mask: `(torch.Tensor)`
shape: (bs, response_length)
cliprange: (float)
The clip range used in PPO. See https://arxiv.org/abs/1707.06347

Returns:
`torch.Tensor`:
Gathered log probabilities with the same shape as `index`.
"""
if logits.dtype in [torch.float32, torch.float64]:
selected_logits = torch.gather(logits, dim=-1, index=index.unsqueeze(-1)).squeeze(-1)
# loop to reduce peak mem consumption
logsumexp_values = torch.stack([torch.logsumexp(lg, dim=-1) for lg in logits])
per_token_logps = selected_logits - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x)
else:
# logsumexp approach is unstable with bfloat16, fall back to slightly less efficent approach
per_token_logps = []
for row_logits, row_labels in zip(logits, index): # loop to reduce peak mem consumption
row_logps = F.log_softmax(row_logits, dim=-1)
row_per_token_logps = row_logps.gather(dim=-1, index=row_labels.unsqueeze(-1)).squeeze(-1)
per_token_logps.append(row_per_token_logps)
per_token_logps = torch.stack(per_token_logps)
return per_token_logps


# @torch.compile
def _compile_grpo_loss(
logits: torch.Tensor,
input_ids: torch.Tensor,
advantages: torch.Tensor,
original_logprobs: torch.Tensor,
loss_mask: torch.Tensor,
temperature: float,
epsilon: float,
) -> tuple[Tensor, Tensor]:
# we start by dropping the bos token because it does not have a corresponding logit
input_ids = input_ids[:, 1:]
advantages = advantages[:, 1:]
# original_logprobs = original_logprobs[:, 1:] # no need to do it now
loss_mask = loss_mask[:, 1:]
pg_loss: `a scalar torch.Tensor`
policy gradient loss computed via PPO
pg_clipfrac: (float)
a float number indicating the fraction of policy gradient loss being clipped

# from the logits we drop the last logits because it corresponds to the next token that will be sample but is not here yet
logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token prediction

# Divide logits by sampling temperature.
# See https://huggingface.co/blog/the_n_implementation_details_of_rlhf_with_ppo#policy-training-implementation-details
logits = logits / temperature
per_token_logps = selective_log_softmax(logits, input_ids)

coef_1 = torch.exp(per_token_logps - original_logprobs)
coef_2 = torch.clamp(coef_1, 1 - epsilon, 1 + epsilon)
per_token_loss1 = -coef_1 * advantages
per_token_loss2 = -coef_2 * advantages
per_token_loss = torch.max(per_token_loss1, per_token_loss2)
"""
negative_approx_kl = log_prob - old_log_prob
ratio = torch.exp(negative_approx_kl)
ppo_kl = masked_mean(-negative_approx_kl, eos_mask)

loss = (per_token_loss * loss_mask).sum() / loss_mask.sum()
pg_losses = -advantages * ratio
pg_losses2 = -advantages * torch.clamp(ratio, 1.0 - cliprange, 1.0 + cliprange)

is_clipped = (per_token_loss1 < per_token_loss2).float()
clip_ratio = (is_clipped * loss_mask).sum() / loss_mask.sum()
return loss, clip_ratio
pg_loss = masked_mean(torch.max(pg_losses, pg_losses2), eos_mask)
pg_clipfrac = masked_mean(torch.gt(pg_losses2, pg_losses).float(), eos_mask)
return pg_loss, pg_clipfrac, ppo_kl


@jaxtyped(typechecker=typechecker)
Expand Down
Loading