Files

320 lines
10 KiB
Python

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()