| |
| """Which exact positions receive the attention mass? Per-head top-k diagnostic. |
| |
| Disambiguates "the model attends to separator tokens as a class" from "the model |
| parks on one specific position that happens to be a separator" (a positional sink). |
| Also reports norm-weighted attention (||alpha_i * v_i|| in the style of Kobayashi et |
| al. 2020), since a raw-attention sink can be a no-op if its value vector is small. |
| """ |
| import argparse |
| import json |
| import os |
| import sys |
|
|
| import numpy as np |
| import torch |
|
|
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) |
| from attention_atlas import build_prompt, load_model, graph_roles, arm_labels |
|
|
|
|
| def role_str(roles, p, node_role): |
| kind, payload = roles[p] |
| if payload is None: |
| return kind |
| comp, depth, _ = node_role.get(payload, ("off", -99, -1)) |
| if kind == "latent": |
| return f"latent{payload}" |
| return f"{kind}[{payload}]{comp}{depth}" |
|
|
|
|
| @torch.no_grad() |
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--ckpt", default=None) |
| 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("--n", type=int, default=64) |
| ap.add_argument("--hops", type=int, nargs="+", default=[1, 5, 10]) |
| ap.add_argument("--topk", type=int, default=5) |
| ap.add_argument("--device", default="cuda:0") |
| args = ap.parse_args() |
|
|
| model, tok = load_model(args.ckpt, args.model_id, args.device) |
| data = json.load(open(args.val))[: args.n] |
| L = args.L |
|
|
| seqs, metas = [], [] |
| order = np.arange(len(data[0]["edges"])) |
| for s in data: |
| ids, roles, root_pos = build_prompt(s, tok, L, order, False) |
| seqs.append(ids) |
| metas.append((s, roles, root_pos)) |
| input_ids = torch.tensor(seqs, device=args.device) |
| B, T = input_ids.shape |
| attn_mask = torch.ones_like(input_ids) |
| pos = torch.arange(T, device=args.device).unsqueeze(0).expand(B, -1) |
|
|
| out = model.forward(input_ids, attn_mask, input_ids.clone(), pos) |
| rep = model.base_causallm(inputs_embeds=out.inputs_embeds, attention_mask=attn_mask, |
| position_ids=pos, output_attentions=True) |
| att = torch.stack(rep.attentions, 0).float().cpu().numpy() |
|
|
| |
| vnorm = value_norms(model, out.inputs_embeds, attn_mask, pos) |
|
|
| n_layer, _, n_head = att.shape[0], att.shape[1], att.shape[2] |
| print(f"\n########## {args.tag} ##########") |
| for m in args.hops: |
| print(f"\n--- hop {m} (query = root_pos+{m-1}) ---") |
| for l in range(n_layer): |
| for h in range(n_head): |
| |
| |
| |
| rows = np.stack([att[l, bi, h, mp[2] + m - 1, : mp[2] + m] |
| for bi, mp in enumerate(metas)]) |
| mean_row = rows.mean(0) |
| top = np.argsort(-mean_row)[: args.topk] |
| s0, roles0, _ = metas[0] |
| nr0 = graph_roles(s0, L) |
| desc = " ".join(f"{p}:{mean_row[p]:.2f}({role_str(roles0, p, nr0)})" |
| for p in top) |
| |
| nw = mean_row * np.stack([vnorm[l, bi, h, : mp[2] + m] |
| for bi, mp in enumerate(metas)]).mean(0) |
| nw = nw / max(nw.sum(), 1e-9) |
| topn = np.argsort(-nw)[: args.topk] |
| descn = " ".join(f"{p}:{nw[p]:.2f}({role_str(roles0, p, nr0)})" |
| for p in topn) |
| print(f" L{l}H{h} raw | {desc}") |
| print(f" nw | {descn}") |
|
|
|
|
| def value_norms(model, embeds, attn_mask, pos): |
| """||v_i|| per layer/head, captured by hooking each block's attention module.""" |
| store = {} |
|
|
| def mk_hook(idx): |
| def hook(mod, inp, out): |
| hidden = inp[0] |
| _, _, v = mod.c_attn(hidden).split(mod.split_size, dim=2) |
| B, T, _ = v.shape |
| v = v.view(B, T, mod.num_heads, -1) |
| store[idx] = v.norm(dim=-1).permute(0, 2, 1).float().cpu().numpy() |
| return hook |
|
|
| handles = [blk.attn.register_forward_hook(mk_hook(i)) |
| for i, blk in enumerate(model.base_causallm.transformer.h)] |
| model.base_causallm(inputs_embeds=embeds, attention_mask=attn_mask, position_ids=pos) |
| for hd in handles: |
| hd.remove() |
| return np.stack([store[i] for i in sorted(store)], 0) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|