Add mini-10m training and inference pipeline

This commit is contained in:
owenqwenstarsky
2026-03-13 12:41:08 -05:00
commit d0ca860f85
6 changed files with 774 additions and 0 deletions
+10
View File
@@ -0,0 +1,10 @@
# Python
__pycache__/
*.pyc
# Local caches and outputs
.cache/
runs/
# Virtualenv
.venv/
+73
View File
@@ -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"
```
+198
View File
@@ -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()
+170
View File
@@ -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())
+4
View File
@@ -0,0 +1,4 @@
torch>=2.2.0
datasets>=2.18.0
tqdm>=4.66.0
numpy>=1.26.0
+319
View File
@@ -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()