| """Eval-only probe: load a trained checkpoint and measure per-position accuracy |
| on the test set at one or more eval loop budgets (test-time loop scaling). |
| |
| Usage: |
| DISABLE_COMPILE=1 python3 eval_probe.py CKPT_DIR STEP [--iters 16,32,64] [--data DIR] |
| |
| CKPT_DIR has all_config.yaml + step_<STEP> (a plain state_dict). Prints, per eval |
| budget: per-position accuracy (pos 0..), final-state accuracy vs k, and the |
| cumulative per-k curve (same metric as length_gen_eval). |
| """ |
| import os, sys, argparse, json |
| os.environ.setdefault("DISABLE_COMPILE", "1") |
| os.environ.setdefault("WANDB_MODE", "offline") |
|
|
| import torch |
| from omegaconf import OmegaConf |
|
|
| from pretrain_config import PretrainConfig |
| from create_model import create_model |
| from pretrain import create_dataloader, autocast_ctx |
| from models.losses import IGNORE_LABEL_ID |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("ckpt_dir") |
| ap.add_argument("step") |
| ap.add_argument("--iters", default="", help="comma list of max_iter_eval; default = config value") |
| ap.add_argument("--data", default="", help="override test data dir") |
| ap.add_argument("--maxk", type=int, default=128) |
| ap.add_argument("--cur_bias", default="", help="force LoopAttn.cur_bias (use 0 for old uniform-init checkpoints)") |
| ap.add_argument("--max_batches", type=int, default=0, help="cap eval batches (0=all); test sets can be huge (Sudoku 422k)") |
| ap.add_argument("--stepsize_decay", default="", help="override FPRM damping decay γ (paper Sudoku eval=0.997)") |
| ap.add_argument("--decay_patience", default="", help="override FPRM patience P (paper Sudoku eval=10)") |
| ap.add_argument("--init_std", default="", help="override FP-init std (official singlez uses scale_init=2.0; our model_utils reads init_std)") |
| ap.add_argument("--out", default="", help="optional json out") |
| args = ap.parse_args() |
|
|
| cfg_d = OmegaConf.to_container(OmegaConf.load(os.path.join(args.ckpt_dir, "all_config.yaml")), resolve=True) |
| cfg_d["load_checkpoint"] = os.path.join(args.ckpt_dir, f"step_{args.step}") |
| cfg_d["resume_from"] = None |
| cfg_d["metrics_out"] = None |
| if args.data: |
| cfg_d["data_paths"] = [args.data] |
| cfg_d["data_paths_test"] = [] |
| cfg_d["k_eval_min"] = 2 |
| cfg_d["k_eval_max"] = args.maxk |
| if args.stepsize_decay != "": |
| cfg_d["arch"]["stepsize_decay"] = float(args.stepsize_decay) |
| if args.decay_patience != "": |
| cfg_d["arch"]["decay_patience"] = int(args.decay_patience) |
| if args.init_std != "": |
| cfg_d["arch"]["init_std"] = float(args.init_std) |
|
|
| iters_list = [int(x) for x in args.iters.split(",") if x] or [cfg_d["arch"].get("max_iter_eval") or cfg_d["arch"]["max_iter"]] |
|
|
| base = PretrainConfig(**cfg_d) |
| eval_loader, eval_metadata = create_dataloader( |
| base, "test", test_set_mode=True, epochs_per_iter=1, |
| global_batch_size=base.global_batch_size, rank=0, world_size=1, |
| ) |
| model, _, _ = create_model(base, eval_metadata, rank=0, world_size=1, strict_load=False) |
| if args.cur_bias != "": |
| n = 0 |
| for m in model.modules(): |
| if hasattr(m, "cur_bias"): |
| m.cur_bias.data.fill_(float(args.cur_bias)); n += 1 |
| print(f"forced cur_bias={args.cur_bias} on {n} LoopAttn module(s)") |
| model.eval() |
|
|
| results = {} |
| for it in iters_list: |
| act = model.model |
| act.config.max_iter_eval = it |
| act.config.max_iter = it |
| act.inner.config.max_iter_eval = it |
| act.inner.config.max_iter = it |
| model.set_num_iters() |
|
|
| |
| pos_cor = pos_val = None |
| exact_n = exact_tot = 0 |
| nsteps_seen = None |
| with torch.no_grad(): |
| for bi, (set_name, batch, _) in enumerate(eval_loader): |
| if args.max_batches and bi >= args.max_batches: |
| break |
| batch = {k: v.cuda() for k, v in batch.items()} |
| with torch.device("cuda"): |
| carry = model.initial_carry(batch) |
| steps = 0 |
| while True: |
| with autocast_ctx(base): |
| carry, _, _, _, preds, all_finish = model(carry=carry, batch=batch, return_keys=["preds"]) |
| steps += 1 |
| if all_finish: |
| break |
| nsteps_seen = steps |
| p = preds["preds"] |
| lab = batch["labels"] |
| valid = lab != IGNORE_LABEL_ID |
| cor = valid & (p == lab) |
| c = cor.float().sum(0) |
| v = valid.float().sum(0) |
| if pos_cor is None: |
| pos_cor, pos_val = c, v |
| else: |
| pos_cor += c; pos_val += v |
| vs = valid.sum(1); cs = cor.sum(1); sv = vs > 0 |
| exact_n += int((sv & (cs == vs)).sum().item()) |
| exact_tot += int(sv.sum().item()) |
|
|
| perpos = (pos_cor / pos_val.clamp_min(1)).cpu().tolist() |
| L = len(perpos) |
| |
| import numpy as np |
| cc = np.cumsum(pos_cor.cpu().numpy()); vv = np.cumsum(pos_val.cpu().numpy()) |
| cumk = (cc / np.clip(vv, 1, None)).tolist() |
| exact_acc = exact_n / max(exact_tot, 1) |
| results[it] = {"perpos": perpos, "cumk": cumk, "infer_steps": nsteps_seen, "exact_seq_acc": exact_acc} |
| print(f"\n>>> max_iter_eval={it}: EXACT-sequence acc = {exact_acc:.4f} (per-position mean = {sum(perpos)/len(perpos):.4f})") |
|
|
| show = [0,1,2,3,4,5,6,7,8,9,11,15,19,23,31,63,127] |
| print(f"\n=== max_iter_eval={it} (ran {nsteps_seen} eval loops) ===") |
| print(" per-position acc (final-state @ length pos+1):") |
| print(" " + " ".join(f"p{j}={perpos[j]:.3f}" for j in show if j < L)) |
| print(" cumulative acc@k:") |
| print(" " + " ".join(f"k{j+1}={cumk[j]:.3f}" for j in show if j < L)) |
|
|
| if args.out: |
| json.dump(results, open(args.out, "w")) |
| print(f"\nwrote {args.out}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|