Inital commit
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user