Skip to content
Merged
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
3 changes: 2 additions & 1 deletion .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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 ]
19 changes: 0 additions & 19 deletions configs/test.toml

This file was deleted.

5 changes: 0 additions & 5 deletions install.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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"
}

Expand Down
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ dependencies = [
"pyarrow",
"wandb",
"vllm",
"jaxtyping"
]


Expand All @@ -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"]
120 changes: 0 additions & 120 deletions scripts/subset_data.py

This file was deleted.

21 changes: 8 additions & 13 deletions src/zeroband/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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()

Expand Down
Loading