From 09ccc0aa32565df9ba22b799d86a49f9b721b465 Mon Sep 17 00:00:00 2001 From: Owen Qwen Date: Sat, 6 Jun 2026 17:06:08 -0500 Subject: [PATCH] Inital commit --- README.md | 90 +++++++++ generate.py | 68 +++++++ model.py | 265 +++++++++++++++++++++++++++ requirements.txt | 6 + train.py | 444 +++++++++++++++++++++++++++++++++++++++++++++ train_tokenizer.py | 71 ++++++++ 6 files changed, 944 insertions(+) create mode 100644 README.md create mode 100644 generate.py create mode 100644 model.py create mode 100644 requirements.txt create mode 100644 train.py create mode 100644 train_tokenizer.py diff --git a/README.md b/README.md new file mode 100644 index 0000000..f2ddcdc --- /dev/null +++ b/README.md @@ -0,0 +1,90 @@ +# 1B Dense GPT Pretraining Skeleton + +This repo contains: + +- `model.py` — GPT-style dense decoder-only Transformer (~1.02B params with a 128k vocab) +- `train_tokenizer.py` — optional 128k byte-level BPE tokenizer training on FineWeb-EDU +- `train.py` — DDP pretraining on FineWeb-EDU with gradient accumulation and checkpoint resume/init +- `generate.py` — quick text generation from a checkpoint +- `requirements.txt` + +## Model defaults + +The default architecture targets a custom 128k tokenizer: + +- vocab: tokenizer length, expected `128000` +- layers: `16` +- hidden size: `2048` +- query heads: `16` +- KV heads: `8` (GQA) +- SwiGLU intermediate size: `5632` +- tied input/output embeddings + +This is about **1.02B trainable parameters**. + +## Install + +```bash +pip install -r requirements.txt +``` + +## Train tokenizer (optional) + +If you do not already have a custom 128k tokenizer: + +```bash +python train_tokenizer.py --out_dir tokenizers/fwe_128k --vocab_size 128000 +``` + +Then pass `--tokenizer_path tokenizers/fwe_128k` to training/generation. + +## Train on all GPUs + +On the training box: + +```bash +torchrun --standalone --nproc_per_node=$(nvidia-smi -L | wc -l) train.py \ + --tokenizer_path /path/to/custom-128k-tokenizer \ + --out_dir runs/one_b_fineweb_edu \ + --total_tokens 10000000000 \ + --block_size 4096 \ + --micro_batch_size 1 \ + --grad_accum_steps 64 \ + --precision bf16 +``` + +The script streams `HuggingFaceFW/fineweb-edu` with config `sample-10BT` by default. + +## Resume same run + +Restores model, optimizer, scaler, RNG, step, and token count: + +```bash +torchrun --standalone --nproc_per_node=$(nvidia-smi -L | wc -l) train.py \ + --tokenizer_path /path/to/custom-128k-tokenizer \ + --out_dir runs/one_b_fineweb_edu \ + --resume runs/one_b_fineweb_edu/ckpt_last.pt +``` + +For exact streamed data-position resume, keep `--num_workers 0`. + +## Start a new run from checkpoint weights + +Loads model weights/config only and starts a fresh optimizer/schedule in a new output dir: + +```bash +torchrun --standalone --nproc_per_node=$(nvidia-smi -L | wc -l) train.py \ + --tokenizer_path /path/to/custom-128k-tokenizer \ + --out_dir runs/continued \ + --init_from runs/one_b_fineweb_edu/ckpt_last.pt +``` + +## Generate + +```bash +python generate.py \ + --checkpoint runs/one_b_fineweb_edu/ckpt_last.pt \ + --tokenizer_path /path/to/custom-128k-tokenizer \ + --prompt "The purpose of education is" \ + --max_new_tokens 128 +``` diff --git a/generate.py b/generate.py new file mode 100644 index 0000000..a6ff914 --- /dev/null +++ b/generate.py @@ -0,0 +1,68 @@ +#!/usr/bin/env python3 +"""Generate text from a trained checkpoint.""" + +from __future__ import annotations + +import argparse + +import torch +from transformers import AutoTokenizer + +from model import GPT, GPTConfig + + +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: + return torch.load(path, map_location=map_location) + + +def main() -> None: + p = argparse.ArgumentParser() + p.add_argument("--checkpoint", type=str, required=True) + p.add_argument("--tokenizer_path", type=str, required=True) + p.add_argument("--prompt", type=str, default="The purpose of education is") + p.add_argument("--max_new_tokens", type=int, default=128) + p.add_argument("--temperature", type=float, default=0.8) + p.add_argument("--top_k", type=int, default=50) + p.add_argument("--top_p", type=float, default=0.95) + p.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") + p.add_argument("--precision", choices=["bf16", "fp16", "fp32"], default="bf16") + p.add_argument("--compile", action="store_true") + args = p.parse_args() + + tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_path, use_fast=True) + ckpt = load_checkpoint(args.checkpoint, map_location="cpu") + config = GPTConfig(**ckpt["model_config"]) + if config.vocab_size != len(tokenizer): + raise ValueError(f"checkpoint vocab_size={config.vocab_size}, tokenizer vocab_size={len(tokenizer)}") + + device = torch.device(args.device) + model = GPT(config) + model.load_state_dict(ckpt["model"], strict=True) + model.to(device) + model.eval() + if args.compile: + model = torch.compile(model) + + ids = tokenizer.encode(args.prompt, add_special_tokens=False) + x = torch.tensor([ids], dtype=torch.long, device=device) + dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[args.precision] + use_amp = args.precision != "fp32" and device.type == "cuda" + + with torch.no_grad(): + with torch.autocast(device_type=device.type, dtype=dtype, enabled=use_amp): + y = model.generate( + x, + max_new_tokens=args.max_new_tokens, + temperature=args.temperature, + top_k=args.top_k, + top_p=args.top_p, + eos_token_id=tokenizer.eos_token_id, + ) + print(tokenizer.decode(y[0].tolist(), skip_special_tokens=True)) + + +if __name__ == "__main__": + main() diff --git a/model.py b/model.py new file mode 100644 index 0000000..b589cd8 --- /dev/null +++ b/model.py @@ -0,0 +1,265 @@ +"""A compact GPT-style dense decoder-only Transformer (~1B params with defaults). + +Default shape is intended for a custom 128k vocabulary tokenizer: +- vocab_size=128000 +- n_layer=16, n_embd=2048, n_head=16, n_kv_head=8, intermediate_size=5632 +- tied input/output embeddings + +This lands at about 1.02B trainable parameters. +""" + +from __future__ import annotations + +import json +import math +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Optional + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +@dataclass +class GPTConfig: + vocab_size: int = 128_000 + block_size: int = 4096 + n_layer: int = 16 + n_head: int = 16 + n_kv_head: int = 8 + n_embd: int = 2048 + intermediate_size: int = 5632 + dropout: float = 0.0 + bias: bool = False + rope_theta: float = 10_000.0 + norm_eps: float = 1e-5 + tie_embeddings: bool = True + + @classmethod + def from_json(cls, path: str | Path) -> "GPTConfig": + with open(path, "r", encoding="utf-8") as f: + return cls(**json.load(f)) + + def to_json(self, path: str | Path) -> None: + with open(path, "w", encoding="utf-8") as f: + json.dump(asdict(self), f, indent=2) + + +class RMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-5): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x = x.float() + x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) + return (self.weight * x).to(dtype) + + +def rotate_half(x: torch.Tensor) -> torch.Tensor: + x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + +def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + # x: [B, H, T, D], cos/sin: [1, 1, T, D] + return (x * cos) + (rotate_half(x) * sin) + + +class CausalSelfAttention(nn.Module): + def __init__(self, config: GPTConfig): + super().__init__() + assert config.n_embd % config.n_head == 0 + assert config.n_head % config.n_kv_head == 0 + self.n_head = config.n_head + self.n_kv_head = config.n_kv_head + self.head_dim = config.n_embd // config.n_head + self.kv_repeat = config.n_head // config.n_kv_head + self.dropout = config.dropout + + self.q_proj = nn.Linear(config.n_embd, config.n_head * self.head_dim, bias=config.bias) + self.k_proj = nn.Linear(config.n_embd, config.n_kv_head * self.head_dim, bias=config.bias) + self.v_proj = nn.Linear(config.n_embd, config.n_kv_head * self.head_dim, bias=config.bias) + self.o_proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias) + self.resid_dropout = nn.Dropout(config.dropout) + + inv_freq = 1.0 / ( + config.rope_theta ** (torch.arange(0, self.head_dim, 2).float() / self.head_dim) + ) + self.register_buffer("inv_freq", inv_freq, persistent=False) + + def _rope_cache(self, seq_len: int, device: torch.device, dtype: torch.dtype): + t = torch.arange(seq_len, device=device, dtype=self.inv_freq.dtype) + freqs = torch.outer(t, self.inv_freq.to(device)) + emb = torch.cat((freqs, freqs), dim=-1) + cos = emb.cos()[None, None, :, :].to(dtype) + sin = emb.sin()[None, None, :, :].to(dtype) + return cos, sin + + def forward(self, x: torch.Tensor) -> torch.Tensor: + bsz, seq_len, embd = x.size() + + q = self.q_proj(x).view(bsz, seq_len, self.n_head, self.head_dim).transpose(1, 2) + k = self.k_proj(x).view(bsz, seq_len, self.n_kv_head, self.head_dim).transpose(1, 2) + v = self.v_proj(x).view(bsz, seq_len, self.n_kv_head, self.head_dim).transpose(1, 2) + + cos, sin = self._rope_cache(seq_len, x.device, q.dtype) + q = apply_rope(q, cos, sin) + k = apply_rope(k, cos, sin) + + if self.kv_repeat != 1: + k = k.repeat_interleave(self.kv_repeat, dim=1) + v = v.repeat_interleave(self.kv_repeat, dim=1) + + y = F.scaled_dot_product_attention( + q, + k, + v, + attn_mask=None, + dropout_p=self.dropout if self.training else 0.0, + is_causal=True, + ) + y = y.transpose(1, 2).contiguous().view(bsz, seq_len, embd) + return self.resid_dropout(self.o_proj(y)) + + +class SwiGLU(nn.Module): + def __init__(self, config: GPTConfig): + super().__init__() + self.gate_proj = nn.Linear(config.n_embd, config.intermediate_size, bias=config.bias) + self.up_proj = nn.Linear(config.n_embd, config.intermediate_size, bias=config.bias) + self.down_proj = nn.Linear(config.intermediate_size, config.n_embd, bias=config.bias) + self.dropout = nn.Dropout(config.dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = F.silu(self.gate_proj(x)) * self.up_proj(x) + return self.dropout(self.down_proj(x)) + + +class Block(nn.Module): + def __init__(self, config: GPTConfig): + super().__init__() + self.ln_1 = RMSNorm(config.n_embd, eps=config.norm_eps) + self.attn = CausalSelfAttention(config) + self.ln_2 = RMSNorm(config.n_embd, eps=config.norm_eps) + self.mlp = SwiGLU(config) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x + self.attn(self.ln_1(x)) + x = x + self.mlp(self.ln_2(x)) + return x + + +class GPT(nn.Module): + def __init__(self, config: GPTConfig): + super().__init__() + self.config = config + self.transformer = nn.ModuleDict( + dict( + wte=nn.Embedding(config.vocab_size, config.n_embd), + drop=nn.Dropout(config.dropout), + h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]), + ln_f=RMSNorm(config.n_embd, eps=config.norm_eps), + ) + ) + self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) + if config.tie_embeddings: + self.lm_head.weight = self.transformer.wte.weight + + self.apply(self._init_weights) + # Slightly scale residual projections as in GPT-2 for stability. + for name, param in self.named_parameters(): + if name.endswith("o_proj.weight") or name.endswith("down_proj.weight"): + nn.init.normal_(param, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer)) + + def _init_weights(self, module: nn.Module) -> None: + if isinstance(module, nn.Linear): + nn.init.normal_(module.weight, mean=0.0, std=0.02) + if module.bias is not None: + nn.init.zeros_(module.bias) + elif isinstance(module, nn.Embedding): + nn.init.normal_(module.weight, mean=0.0, std=0.02) + + def forward( + self, idx: torch.Tensor, targets: Optional[torch.Tensor] = None + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + bsz, seq_len = idx.shape + if seq_len > self.config.block_size: + raise ValueError(f"sequence length {seq_len} exceeds block_size {self.config.block_size}") + + x = self.transformer.wte(idx) + x = self.transformer.drop(x) + for block in self.transformer.h: + x = block(x) + x = self.transformer.ln_f(x) + + if targets is not None: + logits = self.lm_head(x) + loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=-1) + else: + logits = self.lm_head(x[:, [-1], :]) + loss = None + return logits, loss + + @torch.no_grad() + def generate( + self, + idx: torch.Tensor, + max_new_tokens: int, + temperature: float = 1.0, + top_k: Optional[int] = None, + top_p: Optional[float] = None, + eos_token_id: Optional[int] = None, + ) -> torch.Tensor: + for _ in range(max_new_tokens): + idx_cond = idx[:, -self.config.block_size :] + logits, _ = self(idx_cond) + logits = logits[:, -1, :] + if temperature <= 0: + next_id = torch.argmax(logits, dim=-1, keepdim=True) + else: + logits = logits / temperature + if top_k is not None and top_k > 0: + v, _ = torch.topk(logits, min(top_k, logits.size(-1))) + logits[logits < v[:, [-1]]] = -float("inf") + if top_p is not None and 0 < top_p < 1: + sorted_logits, sorted_indices = torch.sort(logits, descending=True) + probs = torch.softmax(sorted_logits, dim=-1) + cumprobs = torch.cumsum(probs, dim=-1) + mask = cumprobs > top_p + mask[..., 1:] = mask[..., :-1].clone() + mask[..., 0] = False + sorted_logits[mask] = -float("inf") + logits = torch.full_like(logits, -float("inf")) + logits.scatter_(dim=-1, index=sorted_indices, src=sorted_logits) + probs = torch.softmax(logits, dim=-1) + next_id = torch.multinomial(probs, num_samples=1) + idx = torch.cat((idx, next_id), dim=1) + if eos_token_id is not None and torch.all(next_id.squeeze(-1) == eos_token_id): + break + return idx + + def configure_optimizers(self, weight_decay: float, learning_rate: float, betas: tuple[float, float]): + decay_params = [] + nodecay_params = [] + for name, p in self.named_parameters(): + if not p.requires_grad: + continue + if p.dim() >= 2: + decay_params.append(p) + else: + nodecay_params.append(p) + optim_groups = [ + {"params": decay_params, "weight_decay": weight_decay}, + {"params": nodecay_params, "weight_decay": 0.0}, + ] + return torch.optim.AdamW(optim_groups, lr=learning_rate, betas=betas, fused=torch.cuda.is_available()) + + def num_parameters(self, non_embedding: bool = False) -> int: + n = sum(p.numel() for p in self.parameters()) + if non_embedding: + n -= self.transformer.wte.weight.numel() + return n diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..9ba4aa9 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,6 @@ +torch>=2.2 +transformers>=4.40 +datasets>=2.19 +tokenizers>=0.19 +numpy>=1.24 +tqdm>=4.66 diff --git a/train.py b/train.py new file mode 100644 index 0000000..cb59d6d --- /dev/null +++ b/train.py @@ -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() diff --git a/train_tokenizer.py b/train_tokenizer.py new file mode 100644 index 0000000..776aacc --- /dev/null +++ b/train_tokenizer.py @@ -0,0 +1,71 @@ +#!/usr/bin/env python3 +"""Train a custom 128k byte-level BPE tokenizer from FineWeb-EDU.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +from datasets import load_dataset +from tokenizers import Tokenizer, decoders, models, pre_tokenizers, processors, trainers +from transformers import PreTrainedTokenizerFast + + +def text_iterator(dataset_name: str, dataset_config: str | None, split: str, text_column: str, max_docs: int | None): + ds = load_dataset(dataset_name, dataset_config, split=split, streaming=True) + n = 0 + for row in ds: + text = row.get(text_column) + if text: + yield text + n += 1 + if max_docs is not None and n >= max_docs: + break + + +def main() -> None: + p = argparse.ArgumentParser() + p.add_argument("--out_dir", type=str, required=True) + p.add_argument("--vocab_size", type=int, default=128_000) + 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("--max_docs", type=int, default=2_000_000, help="Cap docs used for tokenizer training; set 0 for no cap") + args = p.parse_args() + + out_dir = Path(args.out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + max_docs = None if args.max_docs == 0 else args.max_docs + + eos = "<|endoftext|>" + unk = "<|unk|>" + tokenizer = Tokenizer(models.BPE(unk_token=unk)) + tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False) + tokenizer.decoder = decoders.ByteLevel() + tokenizer.post_processor = processors.ByteLevel(trim_offsets=False) + + trainer = trainers.BpeTrainer( + vocab_size=args.vocab_size, + min_frequency=2, + show_progress=True, + special_tokens=[unk, eos], + ) + tokenizer.train_from_iterator( + text_iterator(args.dataset_name, args.dataset_config, args.split, args.text_column, max_docs), + trainer=trainer, + ) + + fast = PreTrainedTokenizerFast( + tokenizer_object=tokenizer, + unk_token=unk, + bos_token=eos, + eos_token=eos, + pad_token=eos, + ) + fast.save_pretrained(out_dir) + print(f"saved tokenizer with vocab size {len(fast):,} to {out_dir}") + + +if __name__ == "__main__": + main()