File size: 6,686 Bytes
fedf514 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 | """Audit the sparse block decode: (A) generation vs source docs, (B) NO data leakage from masked-out
blocks, (C) how much active context is actually saved.
(B) is the key correctness check: a "canary" secret sits in one block. We show that when that block is
NOT in the active set, the model provably cannot reproduce the secret (masked KV has zero influence);
when it IS active, it can. So sparsity is real, not a leaky mask.
"""
from __future__ import annotations
import argparse, torch, torch.nn as nn
import scripts.demo_infer as DM
CANARY = "Zephyr-Quokka-8827"
DOCS = [
f"CONFIDENTIAL: the master password for the vault is {CANARY}.", # block 0 = canary
"The Ostara Festival is held every spring in the town of Wrenfield.",
"Nordwind Labs is headquartered in Bergen, Norway.",
"The lantern parade of the Ostara Festival follows the River Coll.", # block 3 = answer to Q1
"Photosynthesis converts carbon dioxide and water into glucose.",
"Mount Kilimanjaro is the highest mountain in Africa.",
"The gcd function uses Euclid's algorithm.",
"Canberra is the capital city of Australia.",
]
Q_MAIN = "Along which river does the Ostara Festival's lantern parade take place?"
Q_CANARY = "What is the master password for the vault?"
def content_lens(blk, snk, nb):
L = {b: 0 for b in range(nb)}
for b, s in zip(blk, snk):
if b >= 0 and not s: L[b] += 1
return L
@torch.no_grad()
def gen_fixed(model, tok, ids, blk, snk, q0, active_set, summ, max_new=16, trace=False):
"""Greedy sparse decode with a FIXED active set (no re-routing). Optionally trace savings."""
dev = model.device; nb = max(blk) + 1
seq, bl, sk = list(ids), list(blk), list(snk)
clen = content_lens(blk, snk, nb); total_content = sum(clen.values())
ar = {}; att_frac = []
for step in range(max_new):
n = len(seq)
ar[n - q0] = set(range(nb)) if step == 0 else set(active_set)
m = DM.mask(bl, sk, q0, nb, ar, n, dev)
nxt = int(model(input_ids=torch.tensor([seq], device=dev), attention_mask=m).logits[0, -1].argmax())
if trace and step > 0:
act_content = sum(clen[b] for b in active_set)
resident = nb * summ # all blocks' summaries stay resident
att_frac.append((act_content + resident) / total_content)
seq.append(nxt); bl.append(-2); sk.append(False)
if nxt == tok.eos_token_id: break
out = tok.decode(seq[q0:], skip_special_tokens=True)
return out, (sum(att_frac) / len(att_frac) if att_frac else 1.0)
@torch.no_grad()
def route_trace(model, tok, ids, blk, snk, q0, router, k, summ, max_new=16):
"""Real router decode; return generation + per-step active blocks + avg attended-content fraction."""
dev = model.device; nb = max(blk) + 1
seq, bl, sk = list(ids), list(blk), list(snk)
clen = content_lens(blk, snk, nb); total = sum(clen.values())
ar = {}; active = set(range(nb)); steps = []; fracs = []
for step in range(max_new):
n = len(seq)
ar[n - q0] = set(range(nb)) if step == 0 else set(active)
m = DM.mask(bl, sk, q0, nb, ar, n, dev)
out = model(input_ids=torch.tensor([seq], device=dev), attention_mask=m, output_attentions=True)
nxt = int(out.logits[0, -1].argmax())
att = torch.stack(out.attentions, 0)[:, 0, :, -1, :]
bp = DM.block_content_pos(bl, sk, nb)
feat = torch.zeros(nb, att.shape[0]*att.shape[1], device=dev)
for b in range(nb):
feat[b] = att[:, :, torch.tensor(bp[b][-summ:], device=dev)].mean(-1).reshape(-1).float()
if step > 0:
steps.append(sorted(active)); fracs.append((sum(clen[b] for b in active) + nb*summ) / total)
active = set(router(feat).squeeze(-1).argsort(descending=True)[:k].tolist())
seq.append(nxt); bl.append(-2); sk.append(False)
if nxt == tok.eos_token_id: break
return tok.decode(seq[q0:], skip_special_tokens=True), steps, (sum(fracs)/len(fracs) if fracs else 1.0), total, nb
def main():
ap = argparse.ArgumentParser(); ap.add_argument("--model", required=True)
ap.add_argument("--router", required=True); ap.add_argument("--k", type=int, default=2)
args = ap.parse_args(); dev = "cuda"
from transformers import AutoModelForCausalLM, AutoTokenizer
tok = AutoTokenizer.from_pretrained(args.model)
model = AutoModelForCausalLM.from_pretrained(args.model, dtype=torch.bfloat16,
attn_implementation="eager", device_map="cuda").eval()
ck = torch.load(args.router, map_location=dev); summ = ck["summary_tokens"]
router = nn.Linear(ck["in_dim"], 1).to(dev); router.load_state_dict(ck["state_dict"]); router.eval()
print("="*90 + "\nSOURCE DOCUMENTS (blocks):")
for i, d in enumerate(DOCS): print(f" [block {i}] {d}")
# (A) generation vs source
ids, blk, snk, q0 = DM.build(tok, DOCS, Q_MAIN)
g, steps, frac, total, nb = route_trace(model, tok, ids, blk, snk, q0, router, args.k, summ, 20)
print(f"\n(A) Q: {Q_MAIN}")
print(f" sparse gen: {g.split('Answer:')[-1].strip()[:120]!r}")
print(f" -> grounded in block 3 ('River Coll'). canary block 0 NOT needed.")
# (C) active-context savings
print(f"\n(C) active-context savings (k={args.k} of {nb} blocks):")
print(f" per-step active blocks (first few): {steps[:6]}")
print(f" avg block-content tokens attended per decode step: {frac*100:.0f}% "
f"(vs 100% dense) -> ~{1/max(frac,1e-9):.1f}x less context read")
print(f" (total block-content tokens = {total}; canary block 0 in active steps: "
f"{sum(0 in s for s in steps)}/{len(steps)})")
# (B) leakage: ask the CANARY question under active sets that EXCLUDE vs INCLUDE block 0
idc, blc, snc, q0c = DM.build(tok, DOCS, Q_CANARY)
excl, _ = gen_fixed(model, tok, idc, blc, snc, q0c, active_set={3, 5}, summ=summ, max_new=16) # block 0 masked
incl, _ = gen_fixed(model, tok, idc, blc, snc, q0c, active_set={0, 3}, summ=summ, max_new=16) # block 0 active
print(f"\n(B) DATA-LEAKAGE test — canary secret = {CANARY!r} lives in block 0")
print(f" Q: {Q_CANARY}")
print(f" active={{3,5}} (block 0 MASKED): {excl.split('Answer:')[-1].strip()[:80]!r}")
print(f" -> canary leaked? {CANARY in excl}")
print(f" active={{0,3}} (block 0 ACTIVE): {incl.split('Answer:')[-1].strip()[:80]!r}")
print(f" -> canary present? {CANARY in incl}")
print(" PASS: masked block cannot leak" if (CANARY not in excl and CANARY in incl)
else " CHECK: unexpected leakage behaviour")
if __name__ == "__main__":
main()
|