"""Training loop. Usage: python -m src.train [model] [seed] [--flag ...]""" import csv import math import os import random import numpy as np import torch import torch.nn as nn from src.config import Config, parse_args from src.data import build_examples, decode_tokens, get_splits, make_batch from src.model_api import build_model, greedy_decode def set_seed(s: int) -> None: random.seed(s) np.random.seed(s) torch.manual_seed(s) @torch.no_grad() def evaluate(model, examples, cfg: Config): """Token accuracy (teacher-forced) + exact-match accuracy (greedy) + per-example detail.""" model.eval() token_correct = 0 token_total = 0 em_correct = 0 per_example = [] for x, target in examples: batch = make_batch([(x, target)], cfg) out = model(batch["x"], batch["y_in"]) logits = out["logits"][0] y = batch["y"][0] mask = batch["y_mask"][0] pred_tok = logits.argmax(-1) token_correct += int((pred_tok[mask] == y[mask]).sum()) token_total += int(mask.sum()) xi = torch.tensor(x, dtype=torch.long).unsqueeze(0) gen = greedy_decode(model, xi, cfg)[0].tolist() pred = decode_tokens(gen, cfg) target_n = decode_tokens(target, cfg) ok = pred == target_n em_correct += int(ok) per_example.append((tuple(x), target_n, pred, ok)) model.train() return token_correct / max(1, token_total), em_correct / len(examples), per_example def log_example_rows(log_examples, val_per_example): """Format the fixed 10 logged val inputs as '42:ok' strings, in fixed order.""" by_x = {x: (target, pred, ok) for x, target, pred, ok in val_per_example} rows = [] for x, _target in log_examples: t, p, ok = by_x[tuple(x)] label = "".join(map(str, x)) # "42" in both vocab modes rows.append(f"{label}:{'ok' if ok else f'{p}~{t}'}") return ";".join(rows) def main() -> None: cfg = parse_args() set_seed(cfg.seed) out_dir = os.path.join(cfg.out_dir, cfg.model, f"seed{cfg.seed}") os.makedirs(out_dir, exist_ok=True) cfg.save(os.path.join(out_dir, "config.json")) train_in, val_in = get_splits(cfg) train_ex = build_examples(train_in, cfg) val_ex = build_examples(val_in, cfg) model = build_model(cfg) print(f"model={cfg.model} params={model.param_count()} train={len(train_ex)} val={len(val_ex)} " f"vocab={cfg.vocab} eos={cfg.eos_id} pad={cfg.pad_id}") opt = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay) ce = nn.CrossEntropyLoss(reduction="none") log_examples = sorted(val_ex, key=lambda x: (len(str(x)), x))[: cfg.log_n_examples] csv_path = os.path.join(out_dir, "metrics.csv") fieldnames = ["step", "train_loss", "train_token_acc", "train_em", "val_token_acc", "val_em", "mean_halt_steps", "log_examples", "param_count"] best_val_em = -1.0 patience_left = cfg.early_stop_patience best_path = os.path.join(out_dir, "best.pt") n_batches = math.ceil(len(train_ex) / cfg.batch_size) rng = random.Random(cfg.seed) step = 0 done = False first_row = True def _eval_pass(cur_loss, halt_steps): nonlocal best_val_em, patience_left, done, first_row train_tok, train_em, _ = evaluate(model, train_ex, cfg) val_tok, val_em, val_per = evaluate(model, val_ex, cfg) mean_halt = float(halt_steps.mean()) if halt_steps is not None else float("nan") row = { "step": step, "train_loss": float(cur_loss), "train_token_acc": train_tok, "train_em": train_em, "val_token_acc": val_tok, "val_em": val_em, "mean_halt_steps": mean_halt, "log_examples": log_example_rows(log_examples, val_per), "param_count": model.param_count(), } with open(csv_path, "a", newline="") as fh: w = csv.DictWriter(fh, fieldnames=fieldnames) if first_row: w.writeheader() first_row = False w.writerow(row) if val_em > best_val_em: best_val_em = val_em torch.save(model.state_dict(), best_path) if val_em >= cfg.early_stop_em: patience_left -= 1 else: patience_left = cfg.early_stop_patience if patience_left <= 0: done = True print(f"step {step} loss {float(cur_loss):.4f} train_em {train_em:.3f} " f"val_em {val_em:.3f} val_tok {val_tok:.3f} halt {mean_halt:.2f}") while step < cfg.max_train_steps and not done: rng.shuffle(train_ex) for bi in range(n_batches): sl = train_ex[bi * cfg.batch_size: (bi + 1) * cfg.batch_size] if not sl: continue batch = make_batch(sl, cfg) out = model(batch["x"], batch["y_in"]) logits = out["logits"] y_safe = batch["y"].clamp(max=cfg.vocab - 1) # CE index guard for pad positions loss_tokens = ce(logits.reshape(-1, cfg.vocab), y_safe.reshape(-1)).reshape( logits.shape[0], -1) * batch["y_mask"].float() loss_tokens = loss_tokens.sum() / batch["y_mask"].sum().clamp(min=1) if cfg.halting and step >= cfg.halt_warmup_steps: if step >= cfg.halt_ramp_end_steps: lam = cfg.halt_penalty else: frac = (step - cfg.halt_warmup_steps) / max(1, cfg.halt_ramp_end_steps - cfg.halt_warmup_steps) lam = cfg.halt_penalty * frac else: lam = 0.0 halt = out["halt_steps"] penalty = halt.float().mean() * lam if halt is not None and lam > 0 else 0.0 loss = loss_tokens + penalty opt.zero_grad() loss.backward() opt.step() step += 1 if step % cfg.eval_every == 0 or step >= cfg.max_train_steps: _eval_pass(loss.detach(), halt.detach() if halt is not None else None) if done: break torch.save(model.state_dict(), os.path.join(out_dir, "last.pt")) print(f"DONE steps={step} best_val_em={best_val_em:.4f}") if __name__ == "__main__": main()