diff --git a/results-loop/probe_discount.pt b/results-loop/probe_discount.pt new file mode 100644 index 0000000..4943954 Binary files /dev/null and b/results-loop/probe_discount.pt differ diff --git a/results-loop/probe_discount_ff.pt b/results-loop/probe_discount_ff.pt new file mode 100644 index 0000000..a4997f4 Binary files /dev/null and b/results-loop/probe_discount_ff.pt differ diff --git a/scripts/probe_discount.py b/scripts/probe_discount.py new file mode 100644 index 0000000..c355db6 --- /dev/null +++ b/scripts/probe_discount.py @@ -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"))