Add mini-10m training and inference pipeline
This commit is contained in:
+10
@@ -0,0 +1,10 @@
|
|||||||
|
# Python
|
||||||
|
__pycache__/
|
||||||
|
*.pyc
|
||||||
|
|
||||||
|
# Local caches and outputs
|
||||||
|
.cache/
|
||||||
|
runs/
|
||||||
|
|
||||||
|
# Virtualenv
|
||||||
|
.venv/
|
||||||
@@ -0,0 +1,73 @@
|
|||||||
|
# mini-10m
|
||||||
|
|
||||||
|
Simple decoder-only Transformer (~10.5M params) with a **16k max context window**, trained on the first **10,000 rows** of:
|
||||||
|
|
||||||
|
- `karpathy/tinystories-gpt4-clean`
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
|
`model.py` uses:
|
||||||
|
|
||||||
|
- Byte-level vocabulary (`vocab_size=256`)
|
||||||
|
- 7 Transformer blocks
|
||||||
|
- `d_model=352`, `n_heads=8`
|
||||||
|
- RoPE positional encoding (`max_seq_len=16384`)
|
||||||
|
- RMSNorm + GELU MLP
|
||||||
|
- Tied input/output embeddings
|
||||||
|
|
||||||
|
Parameter count is approximately **10.5M**.
|
||||||
|
|
||||||
|
## Setup
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m venv .venv
|
||||||
|
source .venv/bin/activate
|
||||||
|
pip install -r requirements.txt
|
||||||
|
```
|
||||||
|
|
||||||
|
## Quick dry run
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python train.py --dry-run
|
||||||
|
```
|
||||||
|
|
||||||
|
## Train on 10k TinyStories rows
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python train.py \
|
||||||
|
--dataset karpathy/tinystories-gpt4-clean \
|
||||||
|
--num-rows 10000 \
|
||||||
|
--cache-dir .cache/huggingface \
|
||||||
|
--max-seq-len 16384 \
|
||||||
|
--train-seq-len 2048 \
|
||||||
|
--batch-size 2 \
|
||||||
|
--grad-accum 8 \
|
||||||
|
--max-steps 500
|
||||||
|
```
|
||||||
|
|
||||||
|
Notes:
|
||||||
|
|
||||||
|
- `--max-seq-len 16384` sets model context capability.
|
||||||
|
- `--train-seq-len` can be smaller (for memory); this is common for small runs.
|
||||||
|
- Data loading uses Hugging Face streaming, so only the requested rows are consumed.
|
||||||
|
- Checkpoints and config are written to `runs/tiny10m/`.
|
||||||
|
|
||||||
|
## Inference
|
||||||
|
|
||||||
|
Use the latest checkpoint in a run directory:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python infer.py \
|
||||||
|
--run-dir runs/tiny10m \
|
||||||
|
--prompt "Once upon a time" \
|
||||||
|
--max-new-tokens 200 \
|
||||||
|
--temperature 0.8 \
|
||||||
|
--top-k 40 \
|
||||||
|
--top-p 0.95
|
||||||
|
```
|
||||||
|
|
||||||
|
Or target an exact checkpoint:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python infer.py --checkpoint runs/tiny10m/step_000500.pt --prompt "The little robot"
|
||||||
|
```
|
||||||
@@ -0,0 +1,198 @@
|
|||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from model import MiniLM, ModelConfig
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
p = argparse.ArgumentParser(description="Run inference with a trained mini-10m checkpoint")
|
||||||
|
p.add_argument("--checkpoint", type=str, default=None, help="Path to checkpoint .pt file")
|
||||||
|
p.add_argument("--run-dir", type=str, default="runs/tiny10m", help="Run dir containing step_*.pt")
|
||||||
|
p.add_argument("--prompt", type=str, default="Once upon a time")
|
||||||
|
p.add_argument("--max-new-tokens", type=int, default=200)
|
||||||
|
p.add_argument("--temperature", type=float, default=0.8)
|
||||||
|
p.add_argument("--top-k", type=int, default=40)
|
||||||
|
p.add_argument("--top-p", type=float, default=0.95)
|
||||||
|
p.add_argument("--seed", type=int, default=1337)
|
||||||
|
p.add_argument("--device", type=str, default="auto")
|
||||||
|
p.add_argument("--dtype", type=str, default="auto", choices=["auto", "float32", "bfloat16"])
|
||||||
|
p.add_argument(
|
||||||
|
"--max-seq-len",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="Override model max context window. If omitted, tries run config then defaults to 16384.",
|
||||||
|
)
|
||||||
|
return p.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def detect_device(user_value: str) -> torch.device:
|
||||||
|
if user_value != "auto":
|
||||||
|
return torch.device(user_value)
|
||||||
|
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def detect_dtype(user_value: str, device: torch.device) -> torch.dtype:
|
||||||
|
if user_value == "float32":
|
||||||
|
return torch.float32
|
||||||
|
if user_value == "bfloat16":
|
||||||
|
return torch.bfloat16
|
||||||
|
if device.type == "cuda" and torch.cuda.is_bf16_supported():
|
||||||
|
return torch.bfloat16
|
||||||
|
return torch.float32
|
||||||
|
|
||||||
|
|
||||||
|
def encode_byte_level(text: str) -> list[int]:
|
||||||
|
return list(text.encode("utf-8", errors="ignore"))
|
||||||
|
|
||||||
|
|
||||||
|
def decode_byte_level(tokens: list[int]) -> str:
|
||||||
|
return bytes(tokens).decode("utf-8", errors="replace")
|
||||||
|
|
||||||
|
|
||||||
|
def find_latest_checkpoint(run_dir: Path) -> Path:
|
||||||
|
ckpts = list(run_dir.glob("step_*.pt"))
|
||||||
|
if not ckpts:
|
||||||
|
raise FileNotFoundError(f"No checkpoints found in {run_dir}")
|
||||||
|
|
||||||
|
def step_num(path: Path) -> int:
|
||||||
|
m = re.search(r"step_(\d+)\.pt$", path.name)
|
||||||
|
return int(m.group(1)) if m else -1
|
||||||
|
|
||||||
|
return max(ckpts, key=step_num)
|
||||||
|
|
||||||
|
|
||||||
|
def maybe_load_run_config(run_dir: Path) -> dict:
|
||||||
|
cfg_path = run_dir / "config.json"
|
||||||
|
if not cfg_path.exists():
|
||||||
|
return {}
|
||||||
|
with open(cfg_path, "r", encoding="utf-8") as f:
|
||||||
|
return json.load(f)
|
||||||
|
|
||||||
|
|
||||||
|
def build_model(max_seq_len: int) -> MiniLM:
|
||||||
|
cfg = ModelConfig(
|
||||||
|
vocab_size=256,
|
||||||
|
d_model=352,
|
||||||
|
n_heads=8,
|
||||||
|
n_layers=7,
|
||||||
|
ffn_mult=4,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
dropout=0.0,
|
||||||
|
)
|
||||||
|
return MiniLM(cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def sample_next_token(
|
||||||
|
logits: torch.Tensor,
|
||||||
|
temperature: float,
|
||||||
|
top_k: int,
|
||||||
|
top_p: float,
|
||||||
|
) -> int:
|
||||||
|
if temperature <= 0:
|
||||||
|
return int(torch.argmax(logits).item())
|
||||||
|
|
||||||
|
logits = logits / temperature
|
||||||
|
|
||||||
|
if top_k > 0:
|
||||||
|
k = min(top_k, logits.size(-1))
|
||||||
|
values, _ = torch.topk(logits, k)
|
||||||
|
cutoff = values[..., -1]
|
||||||
|
logits = torch.where(logits < cutoff, torch.full_like(logits, float("-inf")), logits)
|
||||||
|
|
||||||
|
if 0.0 < top_p < 1.0:
|
||||||
|
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
||||||
|
sorted_probs = torch.softmax(sorted_logits, dim=-1)
|
||||||
|
cumulative_probs = torch.cumsum(sorted_probs, dim=-1)
|
||||||
|
|
||||||
|
sorted_mask = cumulative_probs > top_p
|
||||||
|
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
|
||||||
|
sorted_mask[..., 0] = False
|
||||||
|
|
||||||
|
mask = torch.zeros_like(sorted_mask, dtype=torch.bool)
|
||||||
|
mask.scatter_(0, sorted_indices, sorted_mask)
|
||||||
|
logits = torch.where(mask, torch.full_like(logits, float("-inf")), logits)
|
||||||
|
|
||||||
|
probs = torch.softmax(logits, dim=-1)
|
||||||
|
next_token = torch.multinomial(probs, num_samples=1)
|
||||||
|
return int(next_token.item())
|
||||||
|
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def generate(
|
||||||
|
model: MiniLM,
|
||||||
|
prompt_ids: list[int],
|
||||||
|
max_new_tokens: int,
|
||||||
|
temperature: float,
|
||||||
|
top_k: int,
|
||||||
|
top_p: float,
|
||||||
|
device: torch.device,
|
||||||
|
amp_dtype: torch.dtype,
|
||||||
|
) -> list[int]:
|
||||||
|
out_ids = list(prompt_ids)
|
||||||
|
|
||||||
|
for _ in range(max_new_tokens):
|
||||||
|
context = out_ids[-model.cfg.max_seq_len :]
|
||||||
|
x = torch.tensor(context, dtype=torch.long, device=device).unsqueeze(0)
|
||||||
|
with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"):
|
||||||
|
logits, _ = model(x)
|
||||||
|
next_logits = logits[0, -1]
|
||||||
|
next_id = sample_next_token(next_logits, temperature, top_k, top_p)
|
||||||
|
out_ids.append(next_id)
|
||||||
|
|
||||||
|
return out_ids
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
args = parse_args()
|
||||||
|
torch.manual_seed(args.seed)
|
||||||
|
torch.cuda.manual_seed_all(args.seed)
|
||||||
|
|
||||||
|
device = detect_device(args.device)
|
||||||
|
amp_dtype = detect_dtype(args.dtype, device)
|
||||||
|
|
||||||
|
run_dir = Path(args.run_dir)
|
||||||
|
ckpt_path = Path(args.checkpoint) if args.checkpoint else find_latest_checkpoint(run_dir)
|
||||||
|
run_cfg = maybe_load_run_config(ckpt_path.parent)
|
||||||
|
|
||||||
|
max_seq_len = args.max_seq_len
|
||||||
|
if max_seq_len is None:
|
||||||
|
max_seq_len = int(run_cfg.get("max_seq_len", 16_384))
|
||||||
|
|
||||||
|
model = build_model(max_seq_len=max_seq_len).to(device)
|
||||||
|
|
||||||
|
ckpt = torch.load(ckpt_path, map_location=device)
|
||||||
|
state_dict = ckpt["model"] if isinstance(ckpt, dict) and "model" in ckpt else ckpt
|
||||||
|
model.load_state_dict(state_dict, strict=True)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
prompt_ids = encode_byte_level(args.prompt)
|
||||||
|
if not prompt_ids:
|
||||||
|
raise ValueError("Prompt produced no byte tokens. Provide a non-empty prompt.")
|
||||||
|
|
||||||
|
out_ids = generate(
|
||||||
|
model=model,
|
||||||
|
prompt_ids=prompt_ids,
|
||||||
|
max_new_tokens=args.max_new_tokens,
|
||||||
|
temperature=args.temperature,
|
||||||
|
top_k=args.top_k,
|
||||||
|
top_p=args.top_p,
|
||||||
|
device=device,
|
||||||
|
amp_dtype=amp_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
generated_suffix = out_ids[len(prompt_ids) :]
|
||||||
|
print(f"Device: {device}")
|
||||||
|
print(f"Checkpoint: {ckpt_path}")
|
||||||
|
print(f"Max context: {max_seq_len}")
|
||||||
|
print("--- PROMPT ---")
|
||||||
|
print(args.prompt)
|
||||||
|
print("--- GENERATED ---")
|
||||||
|
print(decode_byte_level(generated_suffix))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ModelConfig:
|
||||||
|
vocab_size: int = 256
|
||||||
|
d_model: int = 352
|
||||||
|
n_heads: int = 8
|
||||||
|
n_layers: int = 7
|
||||||
|
ffn_mult: int = 4
|
||||||
|
max_seq_len: int = 16_384
|
||||||
|
dropout: float = 0.0
|
||||||
|
rope_theta: float = 10_000.0
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim: int, eps: float = 1e-6):
|
||||||
|
super().__init__()
|
||||||
|
self.eps = eps
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
norm = x.pow(2).mean(dim=-1, keepdim=True)
|
||||||
|
x = x * torch.rsqrt(norm + self.eps)
|
||||||
|
return self.weight * x
|
||||||
|
|
||||||
|
|
||||||
|
def precompute_rope_cache(max_seq_len: int, head_dim: int, theta: float, device: torch.device):
|
||||||
|
if head_dim % 2 != 0:
|
||||||
|
raise ValueError(f"head_dim must be even for RoPE, got {head_dim}")
|
||||||
|
inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
|
||||||
|
t = torch.arange(max_seq_len, device=device).float()
|
||||||
|
freqs = torch.outer(t, inv_freq)
|
||||||
|
return freqs.cos(), freqs.sin()
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor):
|
||||||
|
q1, q2 = q[..., ::2], q[..., 1::2]
|
||||||
|
k1, k2 = k[..., ::2], k[..., 1::2]
|
||||||
|
|
||||||
|
q_rot_1 = q1 * cos - q2 * sin
|
||||||
|
q_rot_2 = q1 * sin + q2 * cos
|
||||||
|
k_rot_1 = k1 * cos - k2 * sin
|
||||||
|
k_rot_2 = k1 * sin + k2 * cos
|
||||||
|
|
||||||
|
q_out = torch.stack((q_rot_1, q_rot_2), dim=-1).flatten(-2)
|
||||||
|
k_out = torch.stack((k_rot_1, k_rot_2), dim=-1).flatten(-2)
|
||||||
|
return q_out, k_out
|
||||||
|
|
||||||
|
|
||||||
|
class CausalSelfAttention(nn.Module):
|
||||||
|
def __init__(self, cfg: ModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
if cfg.d_model % cfg.n_heads != 0:
|
||||||
|
raise ValueError("d_model must be divisible by n_heads")
|
||||||
|
self.n_heads = cfg.n_heads
|
||||||
|
self.head_dim = cfg.d_model // cfg.n_heads
|
||||||
|
self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False)
|
||||||
|
self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
|
||||||
|
self.dropout = nn.Dropout(cfg.dropout)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
|
||||||
|
bsz, seq_len, dim = x.size()
|
||||||
|
qkv = self.qkv(x)
|
||||||
|
q, k, v = qkv.chunk(3, dim=-1)
|
||||||
|
|
||||||
|
q = q.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
|
||||||
|
k = k.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
|
||||||
|
v = v.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
|
||||||
|
|
||||||
|
# RoPE cache is [T, head_dim/2], expand to [1, 1, T, head_dim/2].
|
||||||
|
cos = cos.unsqueeze(0).unsqueeze(0)
|
||||||
|
sin = sin.unsqueeze(0).unsqueeze(0)
|
||||||
|
q, k = apply_rotary(q, k, cos, sin)
|
||||||
|
|
||||||
|
out = F.scaled_dot_product_attention(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
attn_mask=None,
|
||||||
|
dropout_p=self.dropout.p if self.training else 0.0,
|
||||||
|
is_causal=True,
|
||||||
|
)
|
||||||
|
out = out.transpose(1, 2).contiguous().view(bsz, seq_len, dim)
|
||||||
|
out = self.proj(out)
|
||||||
|
return self.dropout(out)
|
||||||
|
|
||||||
|
|
||||||
|
class MLP(nn.Module):
|
||||||
|
def __init__(self, cfg: ModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
hidden = cfg.ffn_mult * cfg.d_model
|
||||||
|
self.fc = nn.Linear(cfg.d_model, hidden, bias=False)
|
||||||
|
self.proj = nn.Linear(hidden, cfg.d_model, bias=False)
|
||||||
|
self.dropout = nn.Dropout(cfg.dropout)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
x = self.fc(x)
|
||||||
|
x = F.gelu(x, approximate="tanh")
|
||||||
|
x = self.proj(x)
|
||||||
|
return self.dropout(x)
|
||||||
|
|
||||||
|
|
||||||
|
class Block(nn.Module):
|
||||||
|
def __init__(self, cfg: ModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
self.norm1 = RMSNorm(cfg.d_model)
|
||||||
|
self.attn = CausalSelfAttention(cfg)
|
||||||
|
self.norm2 = RMSNorm(cfg.d_model)
|
||||||
|
self.mlp = MLP(cfg)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
|
||||||
|
x = x + self.attn(self.norm1(x), cos, sin)
|
||||||
|
x = x + self.mlp(self.norm2(x))
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class MiniLM(nn.Module):
|
||||||
|
def __init__(self, cfg: ModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
self.cfg = cfg
|
||||||
|
self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
|
||||||
|
self.drop = nn.Dropout(cfg.dropout)
|
||||||
|
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layers)])
|
||||||
|
self.norm_f = RMSNorm(cfg.d_model)
|
||||||
|
self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
|
||||||
|
|
||||||
|
# Weight tying keeps the model near 10M params with this config.
|
||||||
|
self.lm_head.weight = self.tok_emb.weight
|
||||||
|
|
||||||
|
cos, sin = precompute_rope_cache(
|
||||||
|
max_seq_len=cfg.max_seq_len,
|
||||||
|
head_dim=cfg.d_model // cfg.n_heads,
|
||||||
|
theta=cfg.rope_theta,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
)
|
||||||
|
self.register_buffer("rope_cos", cos, persistent=False)
|
||||||
|
self.register_buffer("rope_sin", sin, persistent=False)
|
||||||
|
|
||||||
|
def forward(self, idx: torch.Tensor, targets: torch.Tensor | None = None):
|
||||||
|
bsz, seq_len = idx.size()
|
||||||
|
if seq_len > self.cfg.max_seq_len:
|
||||||
|
raise ValueError(
|
||||||
|
f"Input sequence length {seq_len} exceeds max_seq_len {self.cfg.max_seq_len}"
|
||||||
|
)
|
||||||
|
|
||||||
|
x = self.tok_emb(idx)
|
||||||
|
x = self.drop(x)
|
||||||
|
|
||||||
|
cos = self.rope_cos[:seq_len].to(x.device)
|
||||||
|
sin = self.rope_sin[:seq_len].to(x.device)
|
||||||
|
|
||||||
|
for block in self.blocks:
|
||||||
|
x = block(x, cos, sin)
|
||||||
|
|
||||||
|
x = self.norm_f(x)
|
||||||
|
logits = self.lm_head(x)
|
||||||
|
|
||||||
|
loss = None
|
||||||
|
if targets is not None:
|
||||||
|
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
|
||||||
|
return logits, loss
|
||||||
|
|
||||||
|
|
||||||
|
def count_parameters(model: nn.Module) -> int:
|
||||||
|
return sum(p.numel() for p in model.parameters())
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
torch>=2.2.0
|
||||||
|
datasets>=2.18.0
|
||||||
|
tqdm>=4.66.0
|
||||||
|
numpy>=1.26.0
|
||||||
@@ -0,0 +1,319 @@
|
|||||||
|
import argparse
|
||||||
|
import itertools
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from datasets import load_dataset
|
||||||
|
from tqdm import trange
|
||||||
|
|
||||||
|
from model import MiniLM, ModelConfig, count_parameters
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
p = argparse.ArgumentParser(description="Train a ~10M parameter MiniLM on TinyStories")
|
||||||
|
p.add_argument("--dataset", type=str, default="karpathy/tinystories-gpt4-clean")
|
||||||
|
p.add_argument("--split", type=str, default="train")
|
||||||
|
p.add_argument("--text-column", type=str, default="text")
|
||||||
|
p.add_argument("--num-rows", type=int, default=10_000)
|
||||||
|
p.add_argument("--cache-dir", type=str, default=".cache/huggingface")
|
||||||
|
|
||||||
|
p.add_argument("--max-seq-len", type=int, default=16_384, help="Model context window")
|
||||||
|
p.add_argument(
|
||||||
|
"--train-seq-len",
|
||||||
|
type=int,
|
||||||
|
default=2_048,
|
||||||
|
help="Actual training sequence length (<= max-seq-len)",
|
||||||
|
)
|
||||||
|
p.add_argument("--batch-size", type=int, default=2)
|
||||||
|
p.add_argument("--grad-accum", type=int, default=8)
|
||||||
|
p.add_argument("--max-steps", type=int, default=500)
|
||||||
|
p.add_argument("--eval-every", type=int, default=50)
|
||||||
|
p.add_argument("--eval-batches", type=int, default=20)
|
||||||
|
|
||||||
|
p.add_argument("--lr", type=float, default=3e-4)
|
||||||
|
p.add_argument("--weight-decay", type=float, default=0.1)
|
||||||
|
p.add_argument("--warmup-steps", type=int, default=50)
|
||||||
|
p.add_argument("--grad-clip", type=float, default=1.0)
|
||||||
|
p.add_argument("--dropout", type=float, default=0.0)
|
||||||
|
|
||||||
|
p.add_argument("--seed", type=int, default=1337)
|
||||||
|
p.add_argument("--out-dir", type=str, default="runs/tiny10m")
|
||||||
|
p.add_argument("--save-every", type=int, default=100)
|
||||||
|
p.add_argument("--device", type=str, default="auto")
|
||||||
|
p.add_argument("--dtype", type=str, default="auto", choices=["auto", "float32", "bfloat16"])
|
||||||
|
p.add_argument("--compile", action="store_true", help="torch.compile model")
|
||||||
|
p.add_argument(
|
||||||
|
"--dry-run",
|
||||||
|
action="store_true",
|
||||||
|
help="Build model and run one forward pass on random data, then exit",
|
||||||
|
)
|
||||||
|
return p.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def set_seed(seed: int):
|
||||||
|
random.seed(seed)
|
||||||
|
np.random.seed(seed)
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
torch.cuda.manual_seed_all(seed)
|
||||||
|
|
||||||
|
|
||||||
|
def detect_device(user_value: str) -> torch.device:
|
||||||
|
if user_value != "auto":
|
||||||
|
return torch.device(user_value)
|
||||||
|
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def detect_dtype(user_value: str, device: torch.device) -> torch.dtype:
|
||||||
|
if user_value == "float32":
|
||||||
|
return torch.float32
|
||||||
|
if user_value == "bfloat16":
|
||||||
|
return torch.bfloat16
|
||||||
|
if device.type == "cuda" and torch.cuda.is_bf16_supported():
|
||||||
|
return torch.bfloat16
|
||||||
|
return torch.float32
|
||||||
|
|
||||||
|
|
||||||
|
def encode_byte_level(text: str) -> list[int]:
|
||||||
|
return list(text.encode("utf-8", errors="ignore"))
|
||||||
|
|
||||||
|
|
||||||
|
def load_tokens(
|
||||||
|
dataset_name: str,
|
||||||
|
split: str,
|
||||||
|
text_column: str,
|
||||||
|
num_rows: int,
|
||||||
|
cache_dir: str,
|
||||||
|
) -> np.ndarray:
|
||||||
|
ds = load_dataset(dataset_name, split=split, streaming=True, cache_dir=cache_dir)
|
||||||
|
|
||||||
|
all_tokens: list[int] = []
|
||||||
|
sep = [10, 10] # two newlines
|
||||||
|
rows_seen = 0
|
||||||
|
for row in itertools.islice(ds, num_rows):
|
||||||
|
if text_column not in row:
|
||||||
|
raise KeyError(f"Column '{text_column}' not found in dataset row keys={list(row.keys())}")
|
||||||
|
text = row[text_column]
|
||||||
|
all_tokens.extend(encode_byte_level(text))
|
||||||
|
all_tokens.extend(sep)
|
||||||
|
rows_seen += 1
|
||||||
|
|
||||||
|
if len(all_tokens) < 2048:
|
||||||
|
raise ValueError(f"Not enough tokens to train: got {len(all_tokens)}")
|
||||||
|
if rows_seen == 0:
|
||||||
|
raise ValueError("No rows were streamed from the dataset.")
|
||||||
|
return np.asarray(all_tokens, dtype=np.uint16)
|
||||||
|
|
||||||
|
|
||||||
|
def build_model(args: argparse.Namespace) -> MiniLM:
|
||||||
|
cfg = ModelConfig(
|
||||||
|
vocab_size=256,
|
||||||
|
d_model=352,
|
||||||
|
n_heads=8,
|
||||||
|
n_layers=7,
|
||||||
|
ffn_mult=4,
|
||||||
|
max_seq_len=args.max_seq_len,
|
||||||
|
dropout=args.dropout,
|
||||||
|
)
|
||||||
|
model = MiniLM(cfg)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def get_batch(tokens: torch.Tensor, seq_len: int, batch_size: int, device: torch.device):
|
||||||
|
max_start = tokens.size(0) - seq_len - 1
|
||||||
|
idx = torch.randint(0, max_start, (batch_size,), device=tokens.device)
|
||||||
|
x = torch.stack([tokens[i : i + seq_len] for i in idx])
|
||||||
|
y = torch.stack([tokens[i + 1 : i + seq_len + 1] for i in idx])
|
||||||
|
return x.to(device, non_blocking=True), y.to(device, non_blocking=True)
|
||||||
|
|
||||||
|
|
||||||
|
def cosine_lr(step: int, max_steps: int, warmup_steps: int, base_lr: float) -> float:
|
||||||
|
if step < warmup_steps:
|
||||||
|
return base_lr * (step + 1) / max(1, warmup_steps)
|
||||||
|
progress = (step - warmup_steps) / max(1, max_steps - warmup_steps)
|
||||||
|
return base_lr * 0.5 * (1.0 + math.cos(math.pi * progress))
|
||||||
|
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def evaluate(
|
||||||
|
model: nn.Module,
|
||||||
|
val_tokens: torch.Tensor,
|
||||||
|
seq_len: int,
|
||||||
|
eval_batches: int,
|
||||||
|
batch_size: int,
|
||||||
|
device: torch.device,
|
||||||
|
amp_dtype: torch.dtype,
|
||||||
|
) -> float:
|
||||||
|
model.eval()
|
||||||
|
losses = []
|
||||||
|
for _ in range(eval_batches):
|
||||||
|
xb, yb = get_batch(val_tokens, seq_len, batch_size, device)
|
||||||
|
with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"):
|
||||||
|
_, loss = model(xb, yb)
|
||||||
|
losses.append(loss.item())
|
||||||
|
model.train()
|
||||||
|
return float(np.mean(losses))
|
||||||
|
|
||||||
|
|
||||||
|
def save_checkpoint(
|
||||||
|
out_dir: Path,
|
||||||
|
model: nn.Module,
|
||||||
|
optimizer: torch.optim.Optimizer,
|
||||||
|
step: int,
|
||||||
|
train_loss: float,
|
||||||
|
val_loss: float | None,
|
||||||
|
):
|
||||||
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
ckpt_path = out_dir / f"step_{step:06d}.pt"
|
||||||
|
torch.save(
|
||||||
|
{
|
||||||
|
"model": model.state_dict(),
|
||||||
|
"optimizer": optimizer.state_dict(),
|
||||||
|
"step": step,
|
||||||
|
"train_loss": train_loss,
|
||||||
|
"val_loss": val_loss,
|
||||||
|
},
|
||||||
|
ckpt_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
args = parse_args()
|
||||||
|
if args.train_seq_len > args.max_seq_len:
|
||||||
|
raise ValueError("--train-seq-len must be <= --max-seq-len")
|
||||||
|
|
||||||
|
set_seed(args.seed)
|
||||||
|
device = detect_device(args.device)
|
||||||
|
amp_dtype = detect_dtype(args.dtype, device)
|
||||||
|
|
||||||
|
model = build_model(args).to(device)
|
||||||
|
n_params = count_parameters(model)
|
||||||
|
|
||||||
|
print(f"Device: {device}")
|
||||||
|
print(f"AMP dtype: {amp_dtype}")
|
||||||
|
print(f"Model parameters: {n_params:,}")
|
||||||
|
print(f"Max context window: {args.max_seq_len}")
|
||||||
|
|
||||||
|
if args.compile:
|
||||||
|
model = torch.compile(model)
|
||||||
|
|
||||||
|
if args.dry_run:
|
||||||
|
x = torch.randint(0, 256, (2, min(128, args.train_seq_len)), device=device)
|
||||||
|
y = torch.randint(0, 256, (2, min(128, args.train_seq_len)), device=device)
|
||||||
|
_, loss = model(x, y)
|
||||||
|
print(f"Dry-run loss: {loss.item():.4f}")
|
||||||
|
return
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Streaming dataset {args.dataset} ({args.split}) and consuming first {args.num_rows} rows..."
|
||||||
|
)
|
||||||
|
t0 = time.time()
|
||||||
|
os.makedirs(args.cache_dir, exist_ok=True)
|
||||||
|
tokens = load_tokens(
|
||||||
|
args.dataset,
|
||||||
|
args.split,
|
||||||
|
args.text_column,
|
||||||
|
args.num_rows,
|
||||||
|
args.cache_dir,
|
||||||
|
)
|
||||||
|
elapsed = time.time() - t0
|
||||||
|
print(f"Loaded {len(tokens):,} byte tokens in {elapsed:.1f}s")
|
||||||
|
|
||||||
|
split_idx = int(0.95 * len(tokens))
|
||||||
|
train_tokens = torch.from_numpy(tokens[:split_idx]).long()
|
||||||
|
val_tokens = torch.from_numpy(tokens[split_idx:]).long()
|
||||||
|
if device.type == "cuda":
|
||||||
|
train_tokens = train_tokens.pin_memory()
|
||||||
|
val_tokens = val_tokens.pin_memory()
|
||||||
|
|
||||||
|
min_needed = args.train_seq_len + 2
|
||||||
|
if len(train_tokens) < min_needed or len(val_tokens) < min_needed:
|
||||||
|
raise ValueError(
|
||||||
|
f"Tokenized dataset is too small for train_seq_len={args.train_seq_len}. "
|
||||||
|
f"Need at least {min_needed} tokens in each split."
|
||||||
|
)
|
||||||
|
|
||||||
|
optimizer = torch.optim.AdamW(
|
||||||
|
model.parameters(),
|
||||||
|
lr=args.lr,
|
||||||
|
betas=(0.9, 0.95),
|
||||||
|
weight_decay=args.weight_decay,
|
||||||
|
fused=(device.type == "cuda"),
|
||||||
|
)
|
||||||
|
|
||||||
|
out_dir = Path(args.out_dir)
|
||||||
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
with open(out_dir / "config.json", "w", encoding="utf-8") as f:
|
||||||
|
json.dump(vars(args), f, indent=2)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Starting training: steps={args.max_steps}, batch={args.batch_size}, "
|
||||||
|
f"grad_accum={args.grad_accum}, seq_len={args.train_seq_len}"
|
||||||
|
)
|
||||||
|
|
||||||
|
model.train()
|
||||||
|
running_loss = 0.0
|
||||||
|
|
||||||
|
for step in trange(1, args.max_steps + 1):
|
||||||
|
step_loss = 0.0
|
||||||
|
lr = cosine_lr(step - 1, args.max_steps, args.warmup_steps, args.lr)
|
||||||
|
for pg in optimizer.param_groups:
|
||||||
|
pg["lr"] = lr
|
||||||
|
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
|
for _ in range(args.grad_accum):
|
||||||
|
xb, yb = get_batch(train_tokens, args.train_seq_len, args.batch_size, device)
|
||||||
|
with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"):
|
||||||
|
_, loss = model(xb, yb)
|
||||||
|
loss = loss / args.grad_accum
|
||||||
|
loss.backward()
|
||||||
|
step_loss += loss.item()
|
||||||
|
|
||||||
|
if args.grad_clip > 0:
|
||||||
|
torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
|
||||||
|
optimizer.step()
|
||||||
|
|
||||||
|
running_loss += step_loss
|
||||||
|
|
||||||
|
if step % args.eval_every == 0 or step == 1:
|
||||||
|
train_loss = running_loss / args.eval_every if step > 1 else running_loss
|
||||||
|
running_loss = 0.0
|
||||||
|
val_loss = evaluate(
|
||||||
|
model=model,
|
||||||
|
val_tokens=val_tokens,
|
||||||
|
seq_len=args.train_seq_len,
|
||||||
|
eval_batches=args.eval_batches,
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
device=device,
|
||||||
|
amp_dtype=amp_dtype,
|
||||||
|
)
|
||||||
|
ppl = math.exp(min(20.0, val_loss))
|
||||||
|
print(
|
||||||
|
f"step={step:05d} lr={lr:.2e} train_loss={train_loss:.4f} "
|
||||||
|
f"val_loss={val_loss:.4f} val_ppl={ppl:.2f}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if step % args.save_every == 0 or step == args.max_steps:
|
||||||
|
latest_val = evaluate(
|
||||||
|
model=model,
|
||||||
|
val_tokens=val_tokens,
|
||||||
|
seq_len=args.train_seq_len,
|
||||||
|
eval_batches=max(1, args.eval_batches // 2),
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
device=device,
|
||||||
|
amp_dtype=amp_dtype,
|
||||||
|
)
|
||||||
|
save_checkpoint(out_dir, model, optimizer, step, step_loss, latest_val)
|
||||||
|
|
||||||
|
print(f"Training complete. Checkpoints saved in: {out_dir}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Reference in New Issue
Block a user