| """Depth-scan harvest policies for one deterministic model. |
| |
| The flip analysis (d24 vs d96, warm-s2 probe) showed: solved puzzles freeze |
| (202/211 identical preds), unsolved ones keep wandering (0/45 frozen), and the |
| d48+ plateau is right->wrong (-9) cancelling wrong->right (+14). This script |
| harvests that flux with depth-scan policies, all deterministic: |
| |
| fixed-D : exact/perpos at each depth (reference) |
| stable-halt : scan depths ascending; lock a puzzle at the first depth whose |
| prediction equals the previous depth's (converged), else last. |
| Deployable: per-puzzle early exit on stability. |
| cell-vote : per-cell majority across depths (temporal self-ensemble). |
| oracle : right at ANY depth (ceiling for per-puzzle depth selection). |
| |
| Usage: DISABLE_COMPILE=1 python eval_stability.py checkpoints/fl-b01-L4-warm-s2 \ |
| --step 78120 --depths 6,12,24,36,48,72,96 --max_batches 8 |
| """ |
| import os, argparse, glob |
| os.environ.setdefault("DISABLE_COMPILE", "1"); os.environ.setdefault("WANDB_MODE", "offline") |
| import numpy as np |
| 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 |
|
|
| ap = argparse.ArgumentParser() |
| ap.add_argument("ckpt_dir"); ap.add_argument("--step", default="") |
| ap.add_argument("--depths", default="6,12,24,36,48,72,96"); ap.add_argument("--max_batches", type=int, default=8) |
| a = ap.parse_args() |
|
|
| cfg = OmegaConf.to_container(OmegaConf.load(os.path.join(a.ckpt_dir, "all_config.yaml")), resolve=True) |
| cfg.update(load_checkpoint=None, resume_from=None, metrics_out=None) |
| base = PretrainConfig(**cfg) |
| loader, meta = 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, meta, rank=0, world_size=1, strict_load=False) |
| strip = lambda sd: {k.replace("_orig_mod.", ""): v for k, v in sd.items()} |
| bl = sorted(glob.glob(os.path.join(a.ckpt_dir, "step_*_train_state.pt")), key=lambda p: int(p.split("step_")[1].split("_")[0])) |
| b = os.path.join(a.ckpt_dir, f"step_{a.step}_train_state.pt") if a.step else bl[-1] |
| d = torch.load(b, map_location="cpu", weights_only=False) |
| model.load_state_dict(strip(d["model"]), strict=False) |
| model.load_state_dict(strip(d["ema"]), strict=False) |
| model.eval() |
| icfg = model.model.inner.config |
| depths = [int(x) for x in a.depths.split(",") if x] |
| print(f"{os.path.basename(b)} (EMA) depths={depths} batches={a.max_batches}", flush=True) |
|
|
| preds_by_d, labels = {}, None |
| for D in depths: |
| icfg.L_cycles = D; icfg.n_backwards_L = D |
| ps, ls = [], [] |
| with torch.no_grad(): |
| for bi, (_, batch, _) in enumerate(loader): |
| if bi >= a.max_batches: break |
| batch = {k: v.cuda() for k, v in batch.items()} |
| with torch.device("cuda"): carry = model.initial_carry(batch) |
| while True: |
| with autocast_ctx(base): |
| carry, _, _, _, preds, fin = model(carry=carry, batch=batch, return_keys=["preds"]) |
| if fin: break |
| ps.append(preds["preds"].short().cpu()); ls.append(batch["labels"].short().cpu()) |
| preds_by_d[D] = torch.cat(ps).numpy() |
| if labels is None: labels = torch.cat(ls).numpy() |
| v = labels != IGNORE_LABEL_ID |
| ex = (((preds_by_d[D] == labels) & v).sum(1) == v.sum(1)) |
| pp = ((preds_by_d[D] == labels) & v).sum() / v.sum() |
| print(f" fixed d{D:<3d}: exact {ex.mean():.4f} perpos {pp:.4f}", flush=True) |
|
|
| v = labels != IGNORE_LABEL_ID |
| n = labels.shape[0] |
| right = {D: (((preds_by_d[D] == labels) & v).sum(1) == v.sum(1)) for D in depths} |
|
|
| |
| oracle = np.zeros(n, bool) |
| for D in depths: oracle |= right[D] |
|
|
| |
| final = np.array([preds_by_d[depths[-1]][i] for i in range(n)]) |
| lockD = np.full(n, depths[-1]) |
| locked = final.copy() |
| done = np.zeros(n, bool) |
| for i in range(1, len(depths)): |
| Dp, Dc = depths[i-1], depths[i] |
| same = ((preds_by_d[Dp] == preds_by_d[Dc]) | ~v).all(1) |
| take = same & ~done |
| locked[take] = preds_by_d[Dc][take]; lockD[take] = Dc; done |= take |
| ex_halt = (((locked == labels) & v).sum(1) == v.sum(1)) |
| pp_halt = ((locked == labels) & v).sum() / v.sum() |
|
|
| |
| stack = np.stack([preds_by_d[D] for D in depths]) |
| K = stack.shape[0] |
| vote = np.zeros_like(stack[0]) |
| for i in range(n): |
| col = stack[:, i, :] |
| for t in np.where(v[i])[0]: |
| vals, cnt = np.unique(col[:, t], return_counts=True) |
| vote[i, t] = vals[np.argmax(cnt)] |
| ex_vote = (((vote == labels) & v).sum(1) == v.sum(1)) |
| pp_vote = ((vote == labels) & v).sum() / v.sum() |
|
|
| print(f"\n policy | exact | perpos | note") |
| print(f" stable-halt | {ex_halt.mean():.4f} | {pp_halt:.4f} | mean lock depth {lockD.mean():.1f} (compute<=d{depths[-1]})") |
| print(f" cell-vote | {ex_vote.mean():.4f} | {pp_vote:.4f} | majority over {K} depths") |
| print(f" oracle | {oracle.mean():.4f} | — | right at ANY depth (ceiling)") |
| print(f" (n={n})", flush=True) |
|
|