"""Loop-attention utilization probe (TRM host). Loads a trained loop-attn checkpoint (EMA weights, the eval-time weights), runs a few test batches, and records the actual softmax attention distribution that `LoopAttn` places over the loop-state history {z_L_anchor, o_1..o_t}. Answers the core thesis question: does loop-attn route over the loop history, or collapse back to carry-last (all mass on the most-recent source)? Usage: DISABLE_COMPILE=1 python3 probe_loopattn.py checkpoints/trm-lar-s0 [--max_batches 4] """ import os, sys, 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 import models.loop_attnres as LA def main(): ap = argparse.ArgumentParser() ap.add_argument("ckpt_dir") ap.add_argument("--step", default="", help="train_state step to load EMA from; default=latest") ap.add_argument("--max_batches", type=int, default=4) ap.add_argument("--dump", default="", help="write per-S attention distribution to this json") 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"] = None cfg_d["resume_from"] = None cfg_d["metrics_out"] = None 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) # Load EMA weights (the eval-time weights). bundles = sorted(glob.glob(os.path.join(args.ckpt_dir, "step_*_train_state.pt")), key=lambda p: int(p.split("step_")[1].split("_")[0])) bundle = (os.path.join(args.ckpt_dir, f"step_{args.step}_train_state.pt") if args.step else bundles[-1]) d = torch.load(bundle, map_location="cpu", weights_only=False) # EMA tracks only dense params; puzzle_emb/H_init/L_init live in d["model"]. strip = lambda sd: {k.replace("_orig_mod.", ""): v for k, v in sd.items()} model.load_state_dict(strip(d["model"]), strict=False) missing, unexpected = model.load_state_dict(strip(d["ema"]), strict=False) print(f"loaded model+EMA from {os.path.basename(bundle)} (ema missing={len(missing)})") model.eval() rd = float(model.model.inner.loop_attn.recency_decay) print(f"learned recency_decay = {rd:+.4f} (init was {cfg_d['arch'].get('loop_attnres_recency_init')})") # Hook: monkeypatch LoopAttn.forward to record the softmax distribution per S. by_S = {} # S -> list of [S] mean-over-(B,T) attention vectors orig = LA.LoopAttn.forward def patched(self, sources): v = torch.stack(sources, dim=0) S = v.shape[0] dist = torch.arange(S - 1, -1, -1, device=v.device, dtype=torch.float32) rbias = (-self.recency_decay.float() * dist) logits = torch.einsum("d,sbtd->sbt", self.w.to(v.dtype), LA.rms_norm(v, self.eps)).float() logits = logits + rbias.view(S, 1, 1) a = logits.softmax(dim=0) # [S,B,T] flat = a.reshape(S, -1) # [S, B*T] by_S.setdefault(S, []).append((flat.mean(1).detach().cpu().numpy(), flat.std(1).detach().cpu().numpy())) return torch.einsum("sbt,sbtd->btd", a.to(v.dtype), v) LA.LoopAttn.forward = patched try: with torch.no_grad(): for bi, (set_name, batch, _) in enumerate(eval_loader): if bi >= args.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, _, _, _, _, all_finish = model(carry=carry, batch=batch, return_keys=[]) if all_finish: break finally: LA.LoopAttn.forward = orig print("\n=== mean attention over loop sources (source 0 = distant anchor z_L_in; " "last = most-recent loop output) ===") dump = {"recency_decay": rd, "by_S": {}} for S in sorted(by_S): means = np.stack([m for m, s in by_S[S]]).mean(0) # [S] stds = np.stack([s for m, s in by_S[S]]).mean(0) # [S] avg per-token std # recency-only prior for contrast dist = np.arange(S - 1, -1, -1) prior = np.exp(-rd * dist); prior /= prior.sum() newest = means[-1] hist = 1.0 - newest bars = " ".join(f"s{i}={means[i]:.3f}" for i in range(S)) print(f" S={S:2d} (n={len(by_S[S])} calls): {bars}") print(f" newest={newest:.3f} history={hist:.3f} " f"(recency-prior newest={prior[-1]:.3f}) " f"per-token std(newest)={stds[-1]:.4f}") dump["by_S"][S] = {"mean": means.tolist(), "std": stds.tolist()} if getattr(args, "dump", ""): import json json.dump(dump, open(args.dump, "w"), indent=1) print(f"\nwrote {args.dump}") if __name__ == "__main__": main()