| """Step 3d (clean metric) — end-to-end re-route validated by NLL, not generation. |
| |
| Teacher-force the gold answer; at each answer position choose the ACTIVE block set, build the |
| attention mask, and measure the gold token's NLL. No generation loop -> immune to the rambly-output / |
| memory-fallback noise that broke the accuracy metric. Modes (gold NLL, lower = better): |
| |
| dense : all blocks active (ceiling) |
| static : top-k blocks frozen from the first answer position's real attention |
| landmark-reroute : top-k by attention to each block's resident LANDMARK token, re-picked per position |
| oracle-reroute : top-k by full-block attention mass, re-picked per position (best-case selection) |
| |
| If landmark-reroute NLL ~= dense ~= oracle-reroute, and static NLL is high, on-demand re-route works. |
| Counterfactual facts (fictional capitals) force reading the docs; gold blocks are randomly placed. |
| |
| Run: source env.sh && python scripts/e2e_nll.py --model Qwen/Qwen3-8B --n-hops 3 |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import random |
|
|
| import torch |
| import torch.nn.functional as F |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
|
|
| NEG = -1e9 |
| SYS = "You are a helpful assistant. Use the documents to answer. Give only the capitals, comma-separated, in order." |
| PAIRS = [("France", "Paris"), ("Japan", "Tokyo"), ("Egypt", "Cairo"), ("Brazil", "Brasilia"), |
| ("Canada", "Ottawa"), ("Kenya", "Nairobi"), ("Norway", "Oslo"), ("Peru", "Lima"), |
| ("India", "Delhi"), ("Spain", "Madrid"), ("Greece", "Athens"), ("Cuba", "Havana")] |
| FICT = ["Zandar", "Qeropolis", "Vunbry", "Marnex", "Trellin", "Osquith", "Balmoor", "Kestol", |
| "Yarrow", "Drennik", "Fenwick", "Lorvath", "Sundeep", "Ashgard", "Perrin", "Wexley"] |
|
|
|
|
| def build(tok, rng, n_blocks, n_hops, sink_ids): |
| docs = rng.sample(range(len(PAIRS)), n_blocks) |
| asked = rng.sample(docs, n_hops) |
| slot = {d: bi for bi, d in enumerate(docs)} |
| fic = FICT[:]; rng.shuffle(fic) |
| cap = {d: fic[i] for i, d in enumerate(docs)} |
|
|
| ids, blk, sink, gold = [], [], [], [] |
|
|
| def add(text, b, is_sink=False, special=False, gblock=None, is_ans=False): |
| t = tok(text, add_special_tokens=special)["input_ids"] |
| ids.extend(t); blk.extend([b] * len(t)); sink.extend([is_sink] * len(t)) |
| if is_ans: |
| gold.extend([gblock] * len(t)) |
|
|
| add(SYS + "\n", -1, special=True) |
| for d in docs: |
| for s in sink_ids: |
| ids.append(s); blk.append(slot[d]); sink.append(True) |
| add(f"[Document {slot[d]+1}] The capital of {PAIRS[d][0]} is {cap[d]}.\n", slot[d]) |
| add(f"List the capitals of: {', '.join(PAIRS[d][0] for d in asked)}.\nAnswer:", -2) |
| a0 = len(ids) |
| for j, d in enumerate(asked): |
| prefix = " " + cap[d] if j == 0 else ", " + cap[d] |
| add(prefix, -2, gblock=slot[d], is_ans=True) |
| return ids, blk, sink, a0, gold |
|
|
|
|
| def landmarks_content(blk, sink, n_blocks): |
| lm, isc = {}, [False] * len(blk) |
| for j, (b, s) in enumerate(zip(blk, sink)): |
| if b >= 0 and not s: |
| isc[j] = True; lm[b] = j |
| return [lm[b] for b in range(n_blocks)], isc |
|
|
|
|
| def mask_for(blk, sink, a0, n_blocks, active_rows, dev): |
| n = len(blk) |
| b = torch.tensor(blk, device=dev) |
| lm, isc = landmarks_content(blk, sink, n_blocks) |
| causal = torch.tril(torch.ones(n, n, dtype=torch.bool, device=dev)) |
| is_static = (b == -1) | (b == -2) |
| is_lm = torch.zeros(n, dtype=torch.bool, device=dev); is_lm[torch.tensor(lm, device=dev)] = True |
| static_key = is_static | is_lm |
| isc_t = torch.tensor(isc, device=dev) |
| allowed = torch.zeros(n, n, dtype=torch.bool, device=dev) |
| bi, bj = b.view(n, 1), b.view(1, n) |
| allowed[:a0] = ((bj == -1) | (bi == bj) | (bi == -2))[:a0] |
| for i in range(a0, n): |
| act = active_rows[i - a0] |
| vis = static_key.clone() |
| if act: |
| vis = vis | (isc_t & torch.isin(bj.view(n)[:], torch.tensor(sorted(act), device=dev))) |
| allowed[i] = vis |
| allowed &= causal |
| return torch.where(allowed, 0.0, NEG).view(1, 1, n, n).float() |
|
|
|
|
| @torch.no_grad() |
| def gold_nll(model, ids, mask, a0, dev): |
| t = torch.tensor([ids], device=dev) |
| lp = F.log_softmax(model(input_ids=t, attention_mask=mask).logits[0].float(), -1) |
| tgt = t[0, a0:] |
| return -lp[a0 - 1:-1].gather(-1, tgt.view(-1, 1)).squeeze(-1) |
|
|
|
|
| @torch.no_grad() |
| def signals(model, ids, blk, sink, a0, n_blocks, dev): |
| """Dense pass -> per-answer-position full-block mass and landmark attention.""" |
| n = len(ids) |
| dense_rows = {i: set(range(n_blocks)) for i in range(n - a0)} |
| m = mask_for(blk, sink, a0, n_blocks, dense_rows, dev) |
| out = model(input_ids=torch.tensor([ids], device=dev), attention_mask=m, output_attentions=True) |
| a = torch.stack(out.attentions, 0)[:, 0].mean(dim=(0, 1)).float() |
| lm, isc = landmarks_content(blk, sink, n_blocks) |
| content = torch.zeros(n_blocks, n, device=dev) |
| for j, c in enumerate(isc): |
| if c: |
| content[blk[j], j] = 1.0 |
| full_mass = a[a0:n] @ content.T |
| lm_mass = a[a0:n][:, torch.tensor(lm, device=dev)] |
| return full_mass.cpu(), lm_mass.cpu(), m |
|
|
|
|
| def topk_rows(scores, k, A, static=False): |
| if static: |
| sel = set(scores[0].argsort(descending=True)[:k].tolist()) |
| return {t: set(sel) for t in range(A)} |
| return {t: set(scores[t].argsort(descending=True)[:k].tolist()) for t in range(A)} |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--model", default="Qwen/Qwen3-8B") |
| ap.add_argument("--n-blocks", type=int, default=8) |
| ap.add_argument("--n-hops", type=int, default=3) |
| ap.add_argument("--k", type=int, default=1) |
| ap.add_argument("--n-examples", type=int, default=30) |
| ap.add_argument("--n-sinks", type=int, default=4) |
| ap.add_argument("--seed", type=int, default=1) |
| args = ap.parse_args() |
| dev = "cuda"; rng = random.Random(args.seed) |
|
|
| print(f"loading {args.model} ...") |
| tok = AutoTokenizer.from_pretrained(args.model) |
| model = AutoModelForCausalLM.from_pretrained( |
| args.model, dtype=torch.bfloat16, attn_implementation="eager").to(dev).eval() |
| sink_ids = tok("\n", add_special_tokens=False)["input_ids"] * args.n_sinks |
|
|
| agg = {m: [] for m in ["dense", "static", "landmark", "oracle"]} |
| cov = {m: [] for m in ["static", "landmark", "oracle"]} |
| loads = {m: [] for m in ["dense", "static", "landmark", "oracle"]} |
| for _ in range(args.n_examples): |
| ids, blk, sink, a0, gold = build(tok, rng, args.n_blocks, args.n_hops, sink_ids) |
| A = len(ids) - a0 |
| full_mass, lm_mass, dense_mask = signals(model, ids, blk, sink, a0, args.n_blocks, dev) |
| sel = { |
| "dense": {t: set(range(args.n_blocks)) for t in range(A)}, |
| "static": topk_rows(full_mass, args.k, A, static=True), |
| "landmark": topk_rows(lm_mass, args.k, A), |
| "oracle": topk_rows(full_mass, args.k, A), |
| } |
| for mode, rows in sel.items(): |
| mask = dense_mask if mode == "dense" else mask_for(blk, sink, a0, args.n_blocks, rows, dev) |
| nll = gold_nll(model, ids, mask, a0, dev) |
| agg[mode].append(nll.mean().item()) |
| loads[mode].append(sum(len(rows[t]) for t in range(A)) / A) |
| if mode != "dense": |
| cov[mode].append(sum(gold[t] in rows[t] for t in range(A)) / A) |
|
|
| def mean(x): |
| return sum(x) / len(x) |
| print(f"\nmodel={args.model} n_blocks={args.n_blocks} n_hops={args.n_hops} k={args.k} n={args.n_examples}") |
| print(f"{'mode':>10} | {'gold NLL':>9} | {'avg active':>10} | {'gold-block coverage':>19}") |
| for mode in ["dense", "static", "landmark", "oracle"]: |
| c = f"{mean(cov[mode]):.2f}" if mode in cov else "-" |
| print(f"{mode:>10} | {mean(agg[mode]):>9.3f} | {mean(loads[mode]):>10.2f} | {c:>19}") |
| d, s, l = mean(agg['dense']), mean(agg['static']), mean(agg['landmark']) |
| print(f"\nlandmark-reroute closes {(s-l)/(s-d)*100 if s>d else 0:.0f}% of the static->dense NLL gap " |
| f"at {mean(loads['landmark']):.1f}/{args.n_blocks} blocks active.") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|