"""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_ (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 # ACTLossHead.model = FixedPointReasoningModel_ACTV1 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() # ACTV1.set_num_iters() reads config.max_iter_eval in eval mode # accumulate per-position correct/valid + exact-sequence over the test set pos_cor = pos_val = None exact_n = exact_tot = 0 # exact-sequence (all valid positions correct) — Sudoku/Maze metric 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"] # [B, L] lab = batch["labels"] # [B, L] valid = lab != IGNORE_LABEL_ID cor = valid & (p == lab) c = cor.float().sum(0) # [L] v = valid.float().sum(0) # [L] 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() # per position (final-state acc at length pos+1) L = len(perpos) # cumulative per-k (matches length_gen accuracy@k) 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()