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
+90
View File
@@ -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
```
+68
View File
@@ -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()
+265
View File
@@ -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
+6
View File
@@ -0,0 +1,6 @@
torch>=2.2
transformers>=4.40
datasets>=2.19
tokenizers>=0.19
numpy>=1.24
tqdm>=4.66
+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()
+71
View File
@@ -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()