commit d0ca860f8549d17ae795f548ba6b342d45b2cc18 Author: owenqwenstarsky Date: Fri Mar 13 12:41:08 2026 -0500 Add mini-10m training and inference pipeline diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..eae09af --- /dev/null +++ b/.gitignore @@ -0,0 +1,10 @@ +# Python +__pycache__/ +*.pyc + +# Local caches and outputs +.cache/ +runs/ + +# Virtualenv +.venv/ diff --git a/README.md b/README.md new file mode 100644 index 0000000..2191d1b --- /dev/null +++ b/README.md @@ -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" +``` diff --git a/infer.py b/infer.py new file mode 100644 index 0000000..3d1ec80 --- /dev/null +++ b/infer.py @@ -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() diff --git a/model.py b/model.py new file mode 100644 index 0000000..28108e1 --- /dev/null +++ b/model.py @@ -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()) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..ee29f40 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +torch>=2.2.0 +datasets>=2.18.0 +tqdm>=4.66.0 +numpy>=1.26.0 diff --git a/train.py b/train.py new file mode 100644 index 0000000..a1838a5 --- /dev/null +++ b/train.py @@ -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()