#!/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()