diff options
Diffstat (limited to 'src/train.py')
| -rw-r--r-- | src/train.py | 154 |
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() |
