"""Probe: cross-token residual-stream injection (stateful workspace). Instead of looping the band k times within each token step, carry the band output S across token steps: at step t, one band pass with input merge(e_t, S_{t-1}) (positions aligned; the new position inherits the last position's state). Every position deepens by one iteration per emitted token -- the amortized loop. Untrained probe: is this stable? coherent? does the J-lens track? """ import argparse import sys from pathlib import Path import torch from loop_common import BandLooper, MergeAdapter, chat_prompt sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from jlens.core import JLens, load_model # noqa: E402 ROOT = Path(__file__).resolve().parent.parent def seed_state(S, mode, gamma=0.9): """State the new position inherits: last / mean / recency-weighted mean.""" if mode == "last": return S[:, -1:] if mode == "mean": return S.mean(dim=1, keepdim=True) if mode == "ema": T = S.shape[1] w = gamma ** torch.arange(T - 1, -1, -1, device=S.device, dtype=torch.float32) w = (w / w.sum()).view(1, T, 1) return (S.float() * w).sum(dim=1, keepdim=True).to(S.dtype) raise ValueError(mode) @torch.no_grad() def generate_carry(looper, adapter, tok, ids, max_new=60, carry=True, lens=None, concept_id=None, mode="last"): """Greedy decode; band output carried across token steps via the merge.""" S = None trace = [] eos = {tok.eos_token_id, tok.convert_tokens_to_ids("")} for step in range(max_new): calls, _ = looper.capture(ids, logits_to_keep=1) e = looper._hin[looper.l0] if S is None or not carry: s = looper.band(e, calls) # plain first pass else: S_pad = torch.cat([S, seed_state(S, mode)], dim=1) s = looper.band(adapter(e, S_pad), calls) # ONE pass, carried S = s logits = looper.suffix_logits(s, calls, last_only=True) nxt = logits[0, -1].argmax().item() if lens is not None and concept_id is not None: probs = lens._readout(s[0].float() @ lens.Jbar[looper.l1].T) probs = torch.softmax(probs.float(), -1) trace.append({"step": step, "P_concept": probs[:, concept_id].max().item(), "norm": (s.norm() / e.norm()).item()}) ids = torch.cat([ids, torch.tensor([[nxt]], device=ids.device)], 1) if nxt in eos: break return ids, trace def main(): ap = argparse.ArgumentParser() ap.add_argument("--adapter", default=None) ap.add_argument("--max-new", type=int, default=60) ap.add_argument("--mode", default="last", choices=["last", "mean", "ema"]) args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) looper = BandLooper(model) adapter = MergeAdapter().cuda() tag = "untrained" if args.adapter: adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) tag = Path(args.adapter).stem jbar = torch.load(ROOT / "results" / "jbar.pt", map_location="cuda") lens = JLens(model, tok, jbar["Jbar"].float()) prompts = [ ("spider", "The animal that spins webs has how many legs? " "Answer with just the number.", " spider"), ("story", "In one sentence, why is the sky blue?", " blue"), ("math", "Tom has 4 boxes of 12 pens and gives away 9 pens. " "How many pens does he have left? Think step by step briefly, " "then give the number.", " pens"), ] for name, q, concept in prompts: cid = tok.encode(concept, add_special_tokens=False)[0] ids = tok(chat_prompt(tok, q, ""), return_tensors="pt", add_special_tokens=False)["input_ids"].cuda() base, _ = generate_carry(looper, adapter, tok, ids, max_new=args.max_new, carry=False) carr, tr = generate_carry(looper, adapter, tok, ids, max_new=args.max_new, carry=True, lens=lens, concept_id=cid, mode=args.mode) n = ids.shape[1] print(f"\n=== [{tag}] {name} ===") print("baseline :", tok.decode(base[0, n:], skip_special_tokens=True)) print("carried :", tok.decode(carr[0, n:], skip_special_tokens=True)) ks = [0, 2, 5, 10, 20, len(tr) - 1] print("trace :", " ".join( f"t={tr[i]['step']}: P={tr[i]['P_concept']:.3f} " f"|s|/|e|={tr[i]['norm']:.2f}" for i in sorted(set(k for k in ks if 0 <= k < len(tr))))) if __name__ == "__main__": main()