summaryrefslogtreecommitdiff
path: root/src/train.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/train.py')
-rw-r--r--src/train.py154
1 files changed, 154 insertions, 0 deletions
diff --git a/src/train.py b/src/train.py
new file mode 100644
index 0000000..f0c63e9
--- /dev/null
+++ b/src/train.py
@@ -0,0 +1,154 @@
+"""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)
+ lam = cfg.halt_penalty if (cfg.halting and step >= cfg.halt_warmup_steps) else 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()