Inital commit
This commit is contained in:
@@ -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
@@ -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()
|
||||
@@ -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
|
||||
@@ -0,0 +1,6 @@
|
||||
torch>=2.2
|
||||
transformers>=4.40
|
||||
datasets>=2.19
|
||||
tokenizers>=0.19
|
||||
numpy>=1.24
|
||||
tqdm>=4.66
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user