116 lines
4.7 KiB
Python
116 lines
4.7 KiB
Python
"""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"))
|