Inital commit

This commit is contained in:
2026-06-06 17:06:08 -05:00
commit 09ccc0aa32
6 changed files with 944 additions and 0 deletions
+444
View File
@@ -0,0 +1,444 @@
#!/usr/bin/env python3
"""Distributed pretraining on FineWeb-EDU with checkpoint resume/init support.
Launch on a training machine with all GPUs, for example:
torchrun --standalone --nproc_per_node=$(nvidia-smi -L | wc -l) train.py \
--tokenizer_path /path/to/custom-128k-tokenizer \
--out_dir runs/one_b_fwe \
--total_tokens 10000000000 \
--block_size 4096 \
--micro_batch_size 1 \
--grad_accum_steps 64
Resume the same run, including optimizer/scaler/RNG state:
torchrun ... train.py --tokenizer_path ... --out_dir runs/one_b_fwe --resume runs/one_b_fwe/ckpt_last.pt
Start a new run from model weights only:
torchrun ... train.py --tokenizer_path ... --out_dir runs/continued --init_from runs/one_b_fwe/ckpt_last.pt
"""
from __future__ import annotations
import argparse
import json
import math
import os
import random
import time
from dataclasses import asdict
from pathlib import Path
from typing import Iterable
import numpy as np
import torch
import torch.distributed as dist
from datasets import load_dataset
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
from transformers import AutoTokenizer
from model import GPT, GPTConfig
class FineWebEduTokenDataset(IterableDataset):
"""Streams FineWeb-EDU text, tokenizes it, and yields contiguous LM blocks.
For deterministic DDP resume, keep num_workers=0. The dataset will skip the
number of per-rank sequences already consumed when --resume is used.
"""
def __init__(
self,
tokenizer_path: str,
dataset_name: str,
dataset_config: str | None,
split: str,
text_column: str,
block_size: int,
seed: int,
rank: int,
world_size: int,
shuffle_buffer: int,
skip_sequences: int = 0,
add_eos: bool = True,
):
self.tokenizer_path = tokenizer_path
self.dataset_name = dataset_name
self.dataset_config = dataset_config
self.split = split
self.text_column = text_column
self.block_size = block_size
self.seed = seed
self.rank = rank
self.world_size = world_size
self.shuffle_buffer = shuffle_buffer
self.skip_sequences = skip_sequences
self.add_eos = add_eos
def _iter_token_blocks(self) -> Iterable[torch.Tensor]:
worker = get_worker_info()
if worker is None:
shard_rank = self.rank
num_shards = self.world_size
worker_seed = self.seed
else:
# Additional sharding across DataLoader workers. Exact resume is only
# guaranteed for num_workers=0, but this prevents duplicate examples.
shard_rank = self.rank * worker.num_workers + worker.id
num_shards = self.world_size * worker.num_workers
worker_seed = self.seed + worker.id
tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path, use_fast=True)
if tokenizer.eos_token_id is None:
raise ValueError("Tokenizer must define an eos_token_id for document boundaries.")
ds = load_dataset(
self.dataset_name,
self.dataset_config,
split=self.split,
streaming=True,
)
if self.shuffle_buffer > 0:
ds = ds.shuffle(seed=worker_seed, buffer_size=self.shuffle_buffer)
ds = ds.shard(num_shards=num_shards, index=shard_rank)
buf: list[int] = []
yielded = 0
for row in ds:
text = row.get(self.text_column)
if not text:
continue
ids = tokenizer.encode(text, add_special_tokens=False)
if self.add_eos:
ids.append(tokenizer.eos_token_id)
buf.extend(ids)
while len(buf) >= self.block_size + 1:
block = buf[: self.block_size + 1]
del buf[: self.block_size]
if yielded < self.skip_sequences:
yielded += 1
continue
yielded += 1
yield torch.tensor(block, dtype=torch.long)
def __iter__(self):
return self._iter_token_blocks()
def setup_distributed():
ddp = "RANK" in os.environ and "WORLD_SIZE" in os.environ
if ddp:
dist.init_process_group(backend="nccl")
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}")
master = rank == 0
else:
rank = 0
local_rank = 0
world_size = 1
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
master = True
return ddp, rank, local_rank, world_size, device, master
def cleanup_distributed(ddp: bool) -> None:
if ddp:
dist.destroy_process_group()
def set_seed(seed: int, rank: int) -> None:
seed = seed + rank
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def get_lr(tokens_seen: int, args: argparse.Namespace) -> float:
if tokens_seen < args.warmup_tokens:
return args.learning_rate * max(1, tokens_seen) / max(1, args.warmup_tokens)
if tokens_seen >= args.total_tokens:
return args.min_lr
progress = (tokens_seen - args.warmup_tokens) / max(1, args.total_tokens - args.warmup_tokens)
coeff = 0.5 * (1.0 + math.cos(math.pi * progress))
return args.min_lr + coeff * (args.learning_rate - args.min_lr)
def set_lr(optimizer: torch.optim.Optimizer, lr: float) -> None:
for group in optimizer.param_groups:
group["lr"] = lr
def unwrap_model(model: torch.nn.Module) -> GPT:
return model.module if isinstance(model, DDP) else model
def save_checkpoint(
path: Path,
raw_model: GPT,
optimizer: torch.optim.Optimizer,
scaler: torch.cuda.amp.GradScaler,
args: argparse.Namespace,
tokens_seen: int,
step: int,
) -> None:
ckpt = {
"model": raw_model.state_dict(),
"model_config": asdict(raw_model.config),
"optimizer": optimizer.state_dict(),
"scaler": scaler.state_dict(),
"tokens_seen": tokens_seen,
"step": step,
"args": vars(args),
"torch_rng_state": torch.get_rng_state(),
"cuda_rng_state_all": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None,
"numpy_rng_state": np.random.get_state(),
"python_rng_state": random.getstate(),
}
tmp = path.with_suffix(path.suffix + ".tmp")
torch.save(ckpt, tmp)
os.replace(tmp, path)
def load_checkpoint(path: str, map_location: str | torch.device = "cpu"):
try:
return torch.load(path, map_location=map_location, weights_only=False)
except TypeError: # older PyTorch
return torch.load(path, map_location=map_location)
def restore_rng(ckpt) -> None:
if "torch_rng_state" in ckpt:
torch.set_rng_state(ckpt["torch_rng_state"])
if torch.cuda.is_available() and ckpt.get("cuda_rng_state_all") is not None:
torch.cuda.set_rng_state_all(ckpt["cuda_rng_state_all"])
if "numpy_rng_state" in ckpt:
np.random.set_state(ckpt["numpy_rng_state"])
if "python_rng_state" in ckpt:
random.setstate(ckpt["python_rng_state"])
def build_arg_parser() -> argparse.ArgumentParser:
p = argparse.ArgumentParser()
p.add_argument("--tokenizer_path", type=str, required=True, help="Local path or HF id for your 128k tokenizer")
p.add_argument("--out_dir", type=str, default="runs/one_b_fineweb_edu")
p.add_argument("--resume", type=str, default=None, help="Resume full training state from checkpoint")
p.add_argument("--init_from", type=str, default=None, help="Initialize model weights from checkpoint, new optimizer/run")
p.add_argument("--dataset_name", type=str, default="HuggingFaceFW/fineweb-edu")
p.add_argument("--dataset_config", type=str, default="sample-10BT")
p.add_argument("--split", type=str, default="train")
p.add_argument("--text_column", type=str, default="text")
p.add_argument("--shuffle_buffer", type=int, default=10_000)
p.add_argument("--total_tokens", type=int, default=10_000_000_000)
p.add_argument("--block_size", type=int, default=4096)
p.add_argument("--micro_batch_size", type=int, default=1)
p.add_argument("--grad_accum_steps", type=int, default=64)
p.add_argument("--n_layer", type=int, default=16)
p.add_argument("--n_head", type=int, default=16)
p.add_argument("--n_kv_head", type=int, default=8)
p.add_argument("--n_embd", type=int, default=2048)
p.add_argument("--intermediate_size", type=int, default=5632)
p.add_argument("--dropout", type=float, default=0.0)
p.add_argument("--rope_theta", type=float, default=10_000.0)
p.add_argument("--untie_embeddings", action="store_true")
p.add_argument("--learning_rate", type=float, default=3e-4)
p.add_argument("--min_lr", type=float, default=3e-5)
p.add_argument("--warmup_tokens", type=int, default=200_000_000)
p.add_argument("--weight_decay", type=float, default=0.1)
p.add_argument("--beta1", type=float, default=0.9)
p.add_argument("--beta2", type=float, default=0.95)
p.add_argument("--grad_clip", type=float, default=1.0)
p.add_argument("--precision", choices=["bf16", "fp16", "fp32"], default="bf16")
p.add_argument("--compile", action="store_true")
p.add_argument("--seed", type=int, default=1337)
p.add_argument("--num_workers", type=int, default=0)
p.add_argument("--log_interval", type=int, default=10)
p.add_argument("--save_every_tokens", type=int, default=250_000_000)
return p
def main() -> None:
args = build_arg_parser().parse_args()
ddp, rank, local_rank, world_size, device, master = setup_distributed()
set_seed(args.seed, rank)
out_dir = Path(args.out_dir)
if master:
out_dir.mkdir(parents=True, exist_ok=True)
with open(out_dir / "args.json", "w", encoding="utf-8") as f:
json.dump(vars(args), f, indent=2)
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_path, use_fast=True)
vocab_size = len(tokenizer)
if master and vocab_size != 128_000:
print(f"[info] tokenizer vocab size is {vocab_size:,}; model will use that size (expected custom 128k).")
ckpt = None
tokens_seen = 0
step = 0
if args.resume and args.init_from:
raise ValueError("Use only one of --resume or --init_from.")
if args.resume:
ckpt = load_checkpoint(args.resume, map_location="cpu")
config = GPTConfig(**ckpt["model_config"])
tokens_seen = int(ckpt.get("tokens_seen", 0))
step = int(ckpt.get("step", 0))
elif args.init_from:
ckpt = load_checkpoint(args.init_from, map_location="cpu")
config = GPTConfig(**ckpt["model_config"])
config.block_size = args.block_size
else:
config = GPTConfig(
vocab_size=vocab_size,
block_size=args.block_size,
n_layer=args.n_layer,
n_head=args.n_head,
n_kv_head=args.n_kv_head,
n_embd=args.n_embd,
intermediate_size=args.intermediate_size,
dropout=args.dropout,
rope_theta=args.rope_theta,
tie_embeddings=not args.untie_embeddings,
)
if config.vocab_size != vocab_size:
raise ValueError(
f"Checkpoint/model vocab_size={config.vocab_size} but tokenizer has {vocab_size}. "
"Use the same tokenizer that was used for the checkpoint."
)
raw_model = GPT(config)
if ckpt is not None:
raw_model.load_state_dict(ckpt["model"], strict=True)
if master:
mode = "resumed full state from" if args.resume else "initialized weights from"
print(f"[info] {mode} {args.resume or args.init_from}")
raw_model.to(device)
optimizer = raw_model.configure_optimizers(
weight_decay=args.weight_decay,
learning_rate=args.learning_rate,
betas=(args.beta1, args.beta2),
)
model: torch.nn.Module = raw_model
if args.compile:
model = torch.compile(model)
use_fp16_scaler = args.precision == "fp16" and device.type == "cuda"
scaler = torch.cuda.amp.GradScaler(enabled=use_fp16_scaler)
if args.resume and ckpt is not None:
optimizer.load_state_dict(ckpt["optimizer"])
scaler.load_state_dict(ckpt["scaler"])
restore_rng(ckpt)
if ddp:
model = DDP(model, device_ids=[local_rank], output_device=local_rank, gradient_as_bucket_view=True)
dtype = {
"bf16": torch.bfloat16,
"fp16": torch.float16,
"fp32": torch.float32,
}[args.precision]
use_amp = args.precision != "fp32" and device.type == "cuda"
tokens_per_iter = args.micro_batch_size * config.block_size * args.grad_accum_steps * world_size
per_rank_sequences_seen = tokens_seen // max(1, world_size) // config.block_size
if args.num_workers > 0 and per_rank_sequences_seen > 0 and master:
print("[warn] exact dataset-position resume is only guaranteed with --num_workers 0")
dataset = FineWebEduTokenDataset(
tokenizer_path=args.tokenizer_path,
dataset_name=args.dataset_name,
dataset_config=args.dataset_config,
split=args.split,
text_column=args.text_column,
block_size=config.block_size,
seed=args.seed,
rank=rank,
world_size=world_size,
shuffle_buffer=args.shuffle_buffer,
skip_sequences=per_rank_sequences_seen,
)
loader = DataLoader(
dataset,
batch_size=args.micro_batch_size,
num_workers=args.num_workers,
pin_memory=device.type == "cuda",
)
data_iter = iter(loader)
if master:
print(f"[info] device={device}, world_size={world_size}, tokens/iter={tokens_per_iter:,}")
print(f"[info] parameters={raw_model.num_parameters():,} ({raw_model.num_parameters()/1e9:.3f}B)")
print(f"[info] starting at step={step:,}, tokens_seen={tokens_seen:,}")
config.to_json(out_dir / "model_config.json")
model.train()
last_log = time.time()
last_tokens = tokens_seen
next_save_at = ((tokens_seen // args.save_every_tokens) + 1) * args.save_every_tokens
while tokens_seen < args.total_tokens:
lr = get_lr(tokens_seen, args)
set_lr(optimizer, lr)
optimizer.zero_grad(set_to_none=True)
total_loss = 0.0
for micro_step in range(args.grad_accum_steps):
batch = next(data_iter).to(device, non_blocking=True)
x = batch[:, :-1].contiguous()
y = batch[:, 1:].contiguous()
sync_context = (
model.no_sync()
if ddp and micro_step < args.grad_accum_steps - 1
else torch.enable_grad()
)
with sync_context:
with torch.autocast(device_type=device.type, dtype=dtype, enabled=use_amp):
_, loss = model(x, y)
loss = loss / args.grad_accum_steps
total_loss += loss.detach().float().item()
scaler.scale(loss).backward()
if args.grad_clip > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(raw_model.parameters(), args.grad_clip)
scaler.step(optimizer)
scaler.update()
step += 1
tokens_seen += tokens_per_iter
if master and step % args.log_interval == 0:
now = time.time()
dt = max(now - last_log, 1e-9)
toks_per_s = (tokens_seen - last_tokens) / dt
print(
f"step {step:7d} | tokens {tokens_seen:13,d}/{args.total_tokens:,} | "
f"loss {total_loss:.4f} | lr {lr:.2e} | {toks_per_s:,.0f} tok/s"
)
last_log = now
last_tokens = tokens_seen
if master and (tokens_seen >= next_save_at or tokens_seen >= args.total_tokens):
save_checkpoint(out_dir / "ckpt_last.pt", raw_model, optimizer, scaler, args, tokens_seen, step)
save_checkpoint(out_dir / f"ckpt_{tokens_seen:013d}.pt", raw_model, optimizer, scaler, args, tokens_seen, step)
print(f"[info] saved checkpoint at {tokens_seen:,} tokens")
next_save_at += args.save_every_tokens
cleanup_distributed(ddp)
if __name__ == "__main__":
main()