"""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 TAG = "discount" if "--idx" in sys.argv: import json _i = int(sys.argv[sys.argv.index("--idx") + 1]) _star = {it["idx"]: it for it in json.load(open(OUT / "star_data.json")) if it["split"] == "test"} Q = _star[_i]["question"] TAG = f"gsm{_i}" 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=72, 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 / (f"probe_{TAG}_ff.pt" if FF else f"probe_{TAG}.pt")) print("\nsaved", OUT / (f"probe_{TAG}_ff.pt" if FF else f"probe_{TAG}.pt"))