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