mylab-share-2 / eval_probe.py
pengxiang's picture
Add files using upload-large-folder tool
7b3a667 verified
Raw
History Blame Contribute Delete
6.47 kB
"""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 # 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()