| |
| """Attention atlas for coconut latent CoT on the 2-arm star task. |
| |
| Descriptive, assumption-free map of *where in the prompt* the model looks when it |
| produces each hop of the latent chain. |
| |
| Method |
| ------ |
| `Coconut.forward` fills the `<|latent|>` slots one pass at a time, each pass using a |
| sliced KV cache, so the per-pass attention blocks are ragged and awkward to stitch. |
| But the forward also returns the fully-populated `inputs_embeds`. Because attention is |
| causal, the fill at position t only ever depended on positions < t, so replaying those |
| final embeddings through the base LM in a single uncached pass reproduces the model |
| exactly while giving one clean (T, T) attention matrix per head. The script asserts |
| logit equivalence between the two paths. |
| |
| Indexing |
| -------- |
| Hop m is read at query position q_m = root_pos + (m - 1), matching |
| `graph_metrics.perhop_categorize` and `scripts/probe_latents.py`. So q_1 is the root |
| token (no latents yet) and for m >= 2 the query *is* latent slot m-1. Attention to the |
| "latest thought" is therefore self-attention; attention to latents m-2, m-3, ... is the |
| part that is actually optional. Both are reported separately. |
| |
| Because a single full-L forward has the same causal prefix at q_m as a short stage-m |
| prompt would, one forward per graph yields every stage's attention at once. |
| |
| Position-vs-content control |
| --------------------------- |
| Training shuffles edge order (`dataset.get_prefix`), so a trained model should key on |
| node identity rather than slot index. Each graph is run under several edge permutations |
| (with candidate order swapped too). We then correlate, across permutations, the |
| attention indexed by node id versus indexed by absolute slot. Content-driven attention |
| gives high node-id correlation and low slot correlation. |
| |
| Usage |
| ----- |
| python scripts/attention_atlas.py \ |
| --ckpt ckpts/star-coconut-L10-bfs-backtrack-ce095/checkpoint_2000 \ |
| --tag ce095 --L 10 --device cuda:0 |
| """ |
| import argparse |
| import json |
| import os |
| import sys |
| from collections import defaultdict |
|
|
| import numpy as np |
| import torch |
| from transformers import AutoModelForCausalLM, AutoConfig |
|
|
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) |
| from stokenizer import STokenizer |
| from coconut import Coconut |
|
|
|
|
| |
|
|
| def graph_roles(sample, L): |
| """node id -> (component, depth, arm_id). arm_id is the depth-1 ancestor.""" |
| parent = {int(d): int(s) for s, d in sample["edges"]} |
| role = {} |
| for comp, rootkey, nbrkey in (("pos", "root", "neighbor_k"), |
| ("neg", "neg_root", "neg_neighbor_k")): |
| r = int(sample[rootkey]) |
| role[r] = (comp, 0, -1) |
| for k in range(1, L + 1): |
| for n in sample[nbrkey].get(str(k), []): |
| n = int(n) |
| anc = n |
| for _ in range(k - 1): |
| anc = parent[anc] |
| role[n] = (comp, k, anc) |
| return role |
|
|
|
|
| def arm_labels(sample, role): |
| """Which arm id is the shortest-path arm, and which neg arm holds neg_target.""" |
| tgt_arm = role[int(sample["steps"][0])][2] if sample["steps"] else -1 |
| neg_arm = role[int(sample["neg_target"])][2] |
| return tgt_arm, neg_arm |
|
|
|
|
| |
|
|
| def build_prompt(sample, tok, L, edge_order, swap_cands): |
| """Return (token_ids, roles, root_pos). roles[i] describes what position i is.""" |
| edges = sample["edges"] |
| toks, roles = ["<eos>"], [("bos", None)] |
| for slot, ei in enumerate(edge_order): |
| s, d = int(edges[ei][0]), int(edges[ei][1]) |
| if slot > 0: |
| |
| |
| |
| toks.append("|") |
| roles.append(("sep", int(edges[edge_order[slot - 1]][1]))) |
| toks.append(str(s)) |
| roles.append(("edge_src", s)) |
| toks.append(str(d)) |
| roles.append(("edge_dst", d)) |
|
|
| tgt, neg = int(sample["target"]), int(sample["neg_target"]) |
| cands = [(neg, "cand_decoy"), (tgt, "cand_target")] if swap_cands \ |
| else [(tgt, "cand_target"), (neg, "cand_decoy")] |
| toks.append("[Q]") |
| roles.append(("Q", None)) |
| for nid, lbl in cands: |
| toks.append(str(nid)) |
| roles.append((lbl, nid)) |
| toks.append("[R]") |
| roles.append(("R", None)) |
| toks.append(str(int(sample["root"]))) |
| roles.append(("root", int(sample["root"]))) |
|
|
| root_pos = len(toks) - 1 |
| for i in range(1, L + 1): |
| toks.append("<|latent|>") |
| roles.append(("latent", i)) |
|
|
| return [tok.convert_tokens_to_ids(t) for t in toks], roles, root_pos |
|
|
|
|
| |
|
|
| REL_KEYS = ["<=-3", "-2", "-1", "0", "+1", "+2", ">=+3"] |
|
|
|
|
| def rel_key(delta): |
| if delta <= -3: |
| return "<=-3" |
| if delta >= 3: |
| return ">=+3" |
| return f"{delta:+d}" if delta != 0 else "0" |
|
|
|
|
| COARSE = ["self", "bos", "sep", "Q", "R", "root", "cand_target", "cand_decoy", |
| "latent_prev", "node_pos", "node_neg"] |
|
|
|
|
| def bucket_positions(roles, root_pos, m, node_role, tgt_arm, neg_arm, L): |
| """Partition key positions 0..q_m into named buckets. |
| |
| Returns (coarse, graph, latent_rec, srcdst) where each maps name -> list of |
| positions. `coarse` is an exact partition of 0..q_m; the others are marginals. |
| """ |
| q = root_pos + (m - 1) |
| coarse = defaultdict(list) |
| graph = defaultdict(list) |
| latent_rec = defaultdict(list) |
| srcdst = defaultdict(list) |
|
|
| for p in range(q + 1): |
| kind, payload = roles[p] |
| if p == q: |
| coarse["self"].append(p) |
| continue |
|
|
| if kind in ("bos", "Q", "R"): |
| coarse[kind].append(p) |
| continue |
| if kind == "latent": |
| coarse["latent_prev"].append(p) |
| latent_rec[f"d{(m - 1) - payload}"].append(p) |
| continue |
|
|
| |
| |
| nid = payload |
| comp, depth, arm = node_role.get(nid, ("off", -99, -1)) |
| pfx = "sep" if kind == "sep" else "" |
| if kind == "sep": |
| coarse["sep"].append(p) |
| elif kind == "root": |
| coarse["root"].append(p) |
| elif kind in ("cand_target", "cand_decoy"): |
| coarse[kind].append(p) |
| else: |
| coarse["node_pos" if comp == "pos" else "node_neg"].append(p) |
| srcdst[kind].append(p) |
|
|
| if comp == "off": |
| continue |
| rk = rel_key(depth - m) |
| graph[f"{pfx}{comp}:{rk}"].append(p) |
| if comp == "pos": |
| graph[f"{pfx}posarm:{'tgt' if arm == tgt_arm else 'sib'}:{rk}"].append(p) |
| else: |
| graph[f"{pfx}negarm:{'q' if arm == neg_arm else 'o'}:{rk}"].append(p) |
|
|
| return coarse, graph, latent_rec, srcdst |
|
|
|
|
| |
|
|
| def load_model(ckpt, model_id, device, num_nodes=100): |
| tok = STokenizer(num_nodes=num_nodes) |
| cfg = AutoConfig.from_pretrained(model_id) |
| base = AutoModelForCausalLM.from_config(cfg, attn_implementation="eager") |
| model = Coconut(base, |
| tok.convert_tokens_to_ids("<|latent|>"), |
| tok.convert_tokens_to_ids("<|start-latent|>"), |
| tok.convert_tokens_to_ids("<|end-latent|>"), |
| tok.eos_token_id) |
| if ckpt: |
| sd = torch.load(ckpt, map_location="cpu") |
| missing, unexpected = model.load_state_dict(sd, strict=False) |
| print(f" loaded {len(sd)} tensors (missing={len(missing)} unexpected={len(unexpected)})") |
| else: |
| print(" RANDOM INIT (no checkpoint)") |
| model.to(device).eval() |
| return model, tok |
|
|
|
|
| @torch.no_grad() |
| def attentions_for_batch(model, input_ids, device): |
| """Coconut forward, then replay the filled embeddings for clean attention maps.""" |
| B, T = input_ids.shape |
| attn_mask = torch.ones_like(input_ids) |
| pos = torch.arange(T, device=device).unsqueeze(0).expand(B, -1) |
|
|
| out = model.forward(input_ids, attn_mask, input_ids.clone(), pos) |
| replay = model.base_causallm(inputs_embeds=out.inputs_embeds, |
| attention_mask=attn_mask, |
| position_ids=pos, |
| output_attentions=True) |
| max_dev = (replay.logits - out.logits).abs().max().item() |
| |
| att = torch.stack(replay.attentions, dim=0) |
| return att, max_dev |
|
|
|
|
| |
|
|
| @torch.no_grad() |
| def run(args): |
| device = args.device |
| model, tok = load_model(args.ckpt, args.model_id, device) |
| data = json.load(open(args.val))[: args.max_samples] |
| L = args.L |
|
|
| n_layer = model.base_causallm.config.n_layer |
| n_head = model.base_causallm.config.n_head |
| HK = [(l, h) for l in range(n_layer) for h in range(n_head)] |
|
|
| |
| acc = defaultdict(lambda: np.zeros((args.n_perm, n_layer, n_head, L))) |
| entropy = np.zeros((args.n_perm, n_layer, n_head, L)) |
| n_seen = 0 |
| worst_dev = 0.0 |
|
|
| |
| corr_nid, corr_slot = [], [] |
|
|
| rng = np.random.default_rng(args.seed) |
| n_edges = len(data[0]["edges"]) |
|
|
| for start in range(0, len(data), args.batch_size): |
| batch = data[start:start + args.batch_size] |
| |
| perms = [] |
| for pi in range(args.n_perm): |
| if pi == 0: |
| perms.append((np.arange(n_edges), False)) |
| else: |
| perms.append((rng.permutation(n_edges), bool(rng.integers(2)))) |
|
|
| per_perm_nid, per_perm_slot = [], [] |
|
|
| for pi, (order, swap) in enumerate(perms): |
| seqs, metas = [], [] |
| for s in batch: |
| ids, roles, root_pos = build_prompt(s, tok, L, order, swap) |
| seqs.append(ids) |
| metas.append((s, roles, root_pos)) |
| input_ids = torch.tensor(seqs, device=device) |
| att, dev = attentions_for_batch(model, input_ids, device) |
| worst_dev = max(worst_dev, dev) |
| att = att.float().cpu().numpy() |
|
|
| nid_vecs = np.zeros((len(batch), n_layer, n_head, L, args.num_nodes)) |
| slot_vecs = np.zeros((len(batch), n_layer, n_head, L, input_ids.shape[1])) |
|
|
| for bi, (s, roles, root_pos) in enumerate(metas): |
| node_role = graph_roles(s, L) |
| tgt_arm, neg_arm = arm_labels(s, node_role) |
| for m in range(1, L + 1): |
| q = root_pos + (m - 1) |
| row = att[:, bi, :, q, :q + 1] |
| coarse, graph, lat, sd_ = bucket_positions( |
| roles, root_pos, m, node_role, tgt_arm, neg_arm, L) |
|
|
| for name, ps in list(coarse.items()) + list(graph.items()) \ |
| + [("lat:" + k, v) for k, v in lat.items()] \ |
| + [("tok:" + k, v) for k, v in sd_.items()]: |
| if ps: |
| acc[name][pi, :, :, m - 1] += row[:, :, ps].sum(axis=2) |
|
|
| p = np.clip(row, 1e-12, None) |
| entropy[pi, :, :, m - 1] += -(p * np.log(p)).sum(axis=2) |
|
|
| |
| |
| |
| |
| for pp in range(q + 1): |
| kind, payload = roles[pp] |
| if payload is not None and kind != "latent": |
| nid_vecs[bi, :, :, m - 1, payload] += row[:, :, pp] |
| slot_vecs[bi, :, :, m - 1, pp] = row[:, :, pp] |
|
|
| per_perm_nid.append(nid_vecs) |
| per_perm_slot.append(slot_vecs) |
|
|
| n_seen += len(batch) |
|
|
| if args.n_perm >= 2: |
| a_n, b_n = per_perm_nid[0], per_perm_nid[1] |
| a_s, b_s = per_perm_slot[0], per_perm_slot[1] |
| corr_nid.append(_rowwise_corr(a_n, b_n)) |
| corr_slot.append(_rowwise_corr(a_s, b_s)) |
|
|
| print(f" {n_seen}/{len(data)} graphs", flush=True) |
|
|
| for k in acc: |
| acc[k] /= n_seen |
| entropy /= n_seen |
|
|
| result = { |
| "tag": args.tag, |
| "ckpt": args.ckpt, |
| "n_graphs": n_seen, |
| "n_perm": args.n_perm, |
| "L": L, |
| "n_layer": n_layer, |
| "n_head": n_head, |
| "max_logit_dev": worst_dev, |
| "buckets": {k: v.tolist() for k, v in acc.items()}, |
| "entropy": entropy.tolist(), |
| "control": { |
| "node_id_corr": float(np.nanmean(np.concatenate(corr_nid))) if corr_nid else None, |
| "slot_corr": float(np.nanmean(np.concatenate(corr_slot))) if corr_slot else None, |
| }, |
| } |
| os.makedirs(args.outdir, exist_ok=True) |
| path = os.path.join(args.outdir, f"atlas_{args.tag}.json") |
| with open(path, "w") as f: |
| json.dump(result, f) |
| print(f"wrote {path}") |
| summarize(result) |
| return result |
|
|
|
|
| def _rowwise_corr(a, b): |
| """Pearson r between matching rows of two (..., N) arrays, flattened.""" |
| a = a.reshape(-1, a.shape[-1]) |
| b = b.reshape(-1, b.shape[-1]) |
| a = a - a.mean(axis=1, keepdims=True) |
| b = b - b.mean(axis=1, keepdims=True) |
| num = (a * b).sum(axis=1) |
| den = np.sqrt((a * a).sum(axis=1) * (b * b).sum(axis=1)) |
| with np.errstate(invalid="ignore", divide="ignore"): |
| return np.where(den > 0, num / den, np.nan) |
|
|
|
|
| def summarize(res): |
| L, buckets = res["L"], res["buckets"] |
| ent = np.array(res["entropy"])[0].mean(axis=(0, 1)) |
| print(f"\n=== {res['tag']} | {res['n_graphs']} graphs | " |
| f"max|Δlogit|={res['max_logit_dev']:.2e} ===") |
| c = res["control"] |
| if c["node_id_corr"] is not None: |
| print(f"permutation control (node tokens only): node-id r={c['node_id_corr']:.3f}" |
| f" slot r={c['slot_corr']:.3f}") |
|
|
| def get(name): |
| return np.array(buckets[name])[0].mean(axis=(0, 1)) if name in buckets \ |
| else np.zeros(L) |
|
|
| sink = get("bos") + get("sep") |
| print("\n-- raw attention mass --") |
| print("hop | bos sep self latPrev nodePOS nodeNEG cand QRroot | entropy") |
| for m in range(L): |
| print(f"{m+1:3d} | {get('bos')[m]:.3f} {get('sep')[m]:.3f} {get('self')[m]:.3f} " |
| f"{get('latent_prev')[m]:.3f} {get('node_pos')[m]:.3f} " |
| f"{get('node_neg')[m]:.3f} " |
| f"{(get('cand_target') + get('cand_decoy'))[m]:.3f} " |
| f"{(get('Q') + get('R') + get('root'))[m]:.3f} | {ent[m]:.3f}") |
|
|
| |
| |
| |
| sepm = np.clip(get("sep"), 1e-9, None) |
| print("\n-- separator attention by edge depth (frac of sep mass; " |
| "chance .051/comp-rel, .026/arm) --") |
| print("hop | POS rel-2 rel-1 rel0 rel+1 | NEG rel0 | rel0 tgtArm sibArm") |
| for m in range(L): |
| s = sepm[m] |
| print(f"{m+1:3d} | {get('seppos:-2')[m]/s:.3f} {get('seppos:-1')[m]/s:.3f} " |
| f"{get('seppos:0')[m]/s:.3f} {get('seppos:+1')[m]/s:.3f} | " |
| f"{get('sepneg:0')[m]/s:.3f} | " |
| f"{get('sepposarm:tgt:0')[m]/s:.3f} " |
| f"{get('sepposarm:sib:0')[m]/s:.3f}") |
|
|
| denom = np.clip(1.0 - sink, 1e-9, None) |
| print("\n-- renormalized excluding bos/sep sink (content-bearing attention) --") |
| print("hop | self latPrev | nodePOS nodeNEG | rel-2 rel-1 rel0 rel+1 | " |
| "tgtArm sibArm") |
| for m in range(L): |
| d = denom[m] |
| print(f"{m+1:3d} | {get('self')[m]/d:.3f} {get('latent_prev')[m]/d:.3f} | " |
| f"{get('node_pos')[m]/d:.3f} {get('node_neg')[m]/d:.3f} | " |
| f"{get('pos:-2')[m]/d:.3f} {get('pos:-1')[m]/d:.3f} " |
| f"{get('pos:0')[m]/d:.3f} {get('pos:+1')[m]/d:.3f} | " |
| f"{get('posarm:tgt:0')[m]/d:.3f} {get('posarm:sib:0')[m]/d:.3f}") |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--ckpt", default=None, help="omit for random-init floor") |
| ap.add_argument("--tag", required=True) |
| ap.add_argument("--val", default="data/star_2arm_L10_valid_fo_bfs.json") |
| ap.add_argument("--model_id", default="configs/symbol-2layer-8head-768dim-L20.json") |
| ap.add_argument("--L", type=int, default=10) |
| ap.add_argument("--num_nodes", type=int, default=100) |
| ap.add_argument("--max_samples", type=int, default=256) |
| ap.add_argument("--batch_size", type=int, default=32) |
| ap.add_argument("--n_perm", type=int, default=3) |
| ap.add_argument("--seed", type=int, default=0) |
| ap.add_argument("--device", default="cuda:0") |
| ap.add_argument("--outdir", default="figs/attention_atlas") |
| run(ap.parse_args()) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|