probe_discount: single-prompt whiteboard microscopy (lens table + digit probs, carry vs FF control) on the 80/15% problem
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,115 @@
|
||||
"""Single-prompt whiteboard microscopy (Nils, 2026-07-16 evening).
|
||||
|
||||
Run one discount problem through the item-21 arm-A carry loop (k=2, p=6),
|
||||
then lens-read every workspace-band layer at every post-prompt position:
|
||||
does anything token-aligned (complement 85, intermediate 12, answer 68,
|
||||
schema words) ever appear on the carried state, and in what order?
|
||||
|
||||
Prints: generated text; top-3 lens tokens per (pause/gen position x band
|
||||
layer subset); digit-probability matrix at the band exit (L30) — the
|
||||
vector that actually seeds the next position's carry. Saves the full
|
||||
table to results-loop/probe_discount.pt.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from carry_common import (PAUSE_ID, _pos_ids, build_step_updates,
|
||||
carry_steps, prompt_prefill, generate_carry_c)
|
||||
from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
from jlens.core import JLens, ResidualCapture, load_model # noqa: E402
|
||||
|
||||
OUT = Path(os.environ.get("LOOP_OUT",
|
||||
Path(__file__).resolve().parent.parent / "results-loop"))
|
||||
Q = "A shirt costs $80 and is on sale for 15% off. How much does it cost?"
|
||||
K, P = 2, 6
|
||||
FF = "--ff" in sys.argv
|
||||
ADAPTER = OUT / ("adapter_carrycot_ff_e400.pt" if FF else "adapter_carrycot_e400.pt")
|
||||
|
||||
model, tok = load_model(dtype=torch.bfloat16)
|
||||
looper = BandLooper(model)
|
||||
band = list(range(looper.l0, looper.l1 + 1))
|
||||
adapter = MergeAdapter(d=model.config.get_text_config().hidden_size).cuda()
|
||||
adapter.load_state_dict(torch.load(ADAPTER, map_location="cuda"))
|
||||
jbar = torch.load(Path(__file__).resolve().parent.parent / "results/jbar.pt",
|
||||
map_location="cuda")["Jbar"]
|
||||
jl = JLens(model, tok, jbar)
|
||||
|
||||
prompt = chat_prompt(tok, Q, DIRECT_SUFFIX)
|
||||
ids = tok(prompt, add_special_tokens=False, return_tensors="pt")["input_ids"].cuda()
|
||||
mask = torch.ones_like(ids)
|
||||
|
||||
with torch.no_grad():
|
||||
gen = generate_carry_c(looper, adapter, tok, ids, mask, K, P,
|
||||
max_new_tokens=48, feedforward=FF)
|
||||
text = tok.decode(gen[0, ids.shape[1] + P:], skip_special_tokens=False)
|
||||
print("GENERATED:", repr(text), flush=True)
|
||||
|
||||
# teacher-forced replay of the full sequence, capture band residuals
|
||||
full = gen[0][None, :].cuda()
|
||||
fmask = torch.ones_like(full)
|
||||
n_prompt = ids.shape[1]
|
||||
T = full.shape[1]
|
||||
with torch.no_grad():
|
||||
calls, _ = looper.capture(full, fmask, logits_to_keep=1,
|
||||
position_ids=_pos_ids(fmask))
|
||||
e = looper._hin[looper.l0]
|
||||
ar = torch.arange(T, device="cuda")
|
||||
pmask = ar[None, :] < n_prompt
|
||||
if FF:
|
||||
X = torch.where(fmask.bool()[..., None], adapter(e, e), e)
|
||||
else:
|
||||
S, X = prompt_prefill(looper, adapter, e, calls, pmask, K)
|
||||
updates = build_step_updates(torch.tensor([n_prompt], device="cuda"),
|
||||
torch.tensor([T], device="cuda"), "cuda")
|
||||
S, X = carry_steps(looper, adapter, e, calls, S, X, updates)
|
||||
with ResidualCapture(model, layers=band) as rc:
|
||||
looper.band(X, calls)
|
||||
acts = {l: rc.acts[l][0] for l in band} # (T, d) each
|
||||
|
||||
# lens-read every band layer at every post-prompt position
|
||||
table = {}
|
||||
for l in band:
|
||||
idx, p = jl.read(acts[l][n_prompt:], l, topk=5)
|
||||
table[l] = (idx.cpu(), p.cpu())
|
||||
|
||||
poslab = [f"pause{j}" for j in range(P)] + [
|
||||
repr(tok.decode([t])) for t in full[0, n_prompt + P:].tolist()]
|
||||
|
||||
print("\n=== top-3 lens tokens (position x layer) ===")
|
||||
show_layers = [14, 18, 22, 26, 30]
|
||||
for i, lab in enumerate(poslab):
|
||||
row = []
|
||||
for l in show_layers:
|
||||
idx, p = table[l]
|
||||
toks = "/".join(repr(tok.decode([t]))[1:-1] for t in idx[i, :3].tolist())
|
||||
row.append(f"L{l}:{toks}")
|
||||
print(f"{i:3d} {lab:12s} " + " ".join(row), flush=True)
|
||||
|
||||
print("\n=== digit probs at band exit L30 (the carried state) ===")
|
||||
digit_ids = [tok(str(d), add_special_tokens=False)["input_ids"][0]
|
||||
for d in range(10)]
|
||||
idx30, p30 = table[looper.l1]
|
||||
probs = torch.zeros(len(poslab), 10)
|
||||
full_read_idx, full_read_p = jl.read(acts[looper.l1][n_prompt:], looper.l1,
|
||||
topk=64)
|
||||
for i in range(len(poslab)):
|
||||
for d, tid in enumerate(digit_ids):
|
||||
hit = (full_read_idx[i] == tid).nonzero()
|
||||
if hit.numel():
|
||||
probs[i, d] = full_read_p[i, hit[0, 0]]
|
||||
hdr = " " + " " * 12 + " ".join(f"{d} " for d in range(10))
|
||||
print(hdr)
|
||||
for i, lab in enumerate(poslab):
|
||||
print(f"{i:3d} {lab:12s} " +
|
||||
" ".join(f"{probs[i, d]:.3f}" for d in range(10)))
|
||||
|
||||
torch.save({"question": Q, "k": K, "p": P, "generated": text,
|
||||
"band": band, "poslab": poslab, "table": table,
|
||||
"digit_probs": probs},
|
||||
OUT / ("probe_discount_ff.pt" if FF else "probe_discount.pt"))
|
||||
print("\nsaved", OUT / ("probe_discount_ff.pt" if FF else "probe_discount.pt"))
|
||||
Reference in New Issue
Block a user