320 lines
10 KiB
Python
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()
|