J-lens workspace reproduction + loop retrofit: lens, band looping, adapters, controls, multi-task evals

Reproduction of the 2026 workspace/J-lens paper on gemma-4 (E2B/12B/26B),
plus the workspace-loop retrofit line: merge adapter, prompt-only latent
planning (MBPP), carry variant, attribution controls (FF/pause/untrained),
band-location ablation, Blocksworld harness, 12B replication scripts.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-14 00:54:12 +02:00
co-authored by Claude Fable 5
commit ef9c08966c
42 changed files with 5068 additions and 0 deletions
+117
View File
@@ -0,0 +1,117 @@
"""Blocksworld: generator, verifier, prompts.
Pure planning domain — no code syntax, no arithmetic. Plans are symbolically
verifiable by simulation, so STaR bucketing works. Instances are generated
(contamination-free by construction) with difficulty = blocks + moves.
"""
import json
import random
import re
MOVE_RE = re.compile(
r"move\s+([A-Z])\s+(?:onto|on top of|on)\s+(?:the\s+)?(table|[A-Z])",
re.IGNORECASE)
def gen_instance(n_blocks, rng):
blocks = [chr(65 + i) for i in range(n_blocks)]
def random_stacks():
bs = blocks[:]
rng.shuffle(bs)
stacks, i = [], 0
while i < len(bs):
take = rng.randint(1, len(bs) - i)
stacks.append(bs[i : i + take])
i += take
return stacks
init = random_stacks()
goal = random_stacks()
while goal == init:
goal = random_stacks()
return {"blocks": blocks, "init": init, "goal": goal}
def fmt_state(stacks):
out = []
for st in stacks:
if len(st) == 1:
out.append(f"{st[0]} is on the table")
else:
out.append(f"{st[0]} is on the table with "
+ " on top, then ".join(
[f"{b}" for b in st[1:]]) + " on top")
# clearer explicit form
lines = []
for st in stacks:
lines.append(f"stack: {' -> '.join(st)} (bottom -> top)")
return "; ".join(lines)
def question(inst):
return (
"You are stacking blocks. Only the TOP block of a stack can be "
"moved, one block at a time.\n"
f"Blocks: {', '.join(inst['blocks'])}\n"
f"Initial state: {fmt_state(inst['init'])}\n"
f"Goal state: {fmt_state(inst['goal'])}\n"
"Give a plan as a numbered list of moves, each exactly of the form "
"'move X onto Y' or 'move X onto the table'.")
def verify_plan(inst, text, max_moves=40):
stacks = [st[:] for st in inst["init"]]
def top_of(b):
for st in stacks:
if st and st[-1] == b:
return st
return None
moves = MOVE_RE.findall(text)
if not moves or len(moves) > max_moves:
return False
for b, tgt in moves:
b = b.upper()
tgt = tgt if tgt.lower() == "table" else tgt.upper()
src = top_of(b)
if src is None:
return False # b not clear (or nonexistent)
if tgt == "table" or tgt.lower() == "table":
src.pop()
stacks.append([b])
else:
dst = top_of(tgt)
if dst is None or b == tgt:
return False
src.pop()
dst.append(b)
stacks = [st for st in stacks if st]
norm = sorted(tuple(st) for st in stacks)
return norm == sorted(tuple(st) for st in inst["goal"])
def make_dataset(n_train=400, n_test=200, seed=0):
rng = random.Random(seed)
items = []
for split, n in (("train", n_train), ("test", n_test)):
for i in range(n):
nb = rng.choice([3, 3, 4, 4, 5])
inst = gen_instance(nb, rng)
items.append({"split": split, "task_id": f"bw_{split}_{i}",
"n_blocks": nb, **inst,
"question": question(inst)})
return items
if __name__ == "__main__":
ds = make_dataset()
print(json.dumps(ds[0], indent=1))
# verifier self-test: identity plan on trivial instance
inst = {"blocks": ["A", "B"], "init": [["A"], ["B"]],
"goal": [["A", "B"]]}
assert verify_plan(inst, "1. move B onto A")
assert not verify_plan(inst, "1. move A onto A")
print("verifier self-test ok")
+66
View File
@@ -0,0 +1,66 @@
"""Blocksworld STaR labeling with the frozen model (direct vs CoT plan)."""
import json
import os
import sys
import time
from pathlib import Path
import torch
from bw_common import make_dataset, verify_plan
from prep_mbpp import batch_generate
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
DIRECT_SUFFIX = "\n\nGive only the numbered list of moves, nothing else."
COT_SUFFIX = ("\n\nFirst think step by step about which blocks must move "
"and in what order (briefly), then give the numbered list of "
"moves.")
def chat(tok, it, suffix):
return tok.apply_chat_template(
[{"role": "user", "content": it["question"] + suffix}],
tokenize=False, add_generation_prompt=True)
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
items = make_dataset(n_train=400, n_test=200, seed=0)
print(f"{len(items)} instances", flush=True)
for tag, suffix, mx in (("direct", DIRECT_SUFFIX, 200),
("cot", COT_SUFFIX, 500)):
t0 = time.time()
gens = batch_generate(model, tok,
[chat(tok, it, suffix) for it in items],
max_new_tokens=mx, batch_size=24)
oks = [verify_plan(it, g) for it, g in zip(items, gens)]
for it, g, ok in zip(items, gens, oks):
it[f"{tag}_ok"] = bool(ok)
it[f"{tag}_plan"] = g if ok else None
print(f"{tag}: acc={sum(oks)/len(items):.3f} "
f"({time.time()-t0:.0f}s)", flush=True)
for it in items:
it["label"] = ("easy" if it["direct_ok"]
else "hard" if it["cot_ok"] else "drop")
it["sol_plan"] = (it["direct_plan"] if it["direct_ok"]
else it["cot_plan"])
for split in ("train", "test"):
sub = [it for it in items if it["split"] == split]
print(f"{split}: easy={sum(i['label']=='easy' for i in sub)} "
f"hard={sum(i['label']=='hard' for i in sub)} "
f"drop={sum(i['label']=='drop' for i in sub)}", flush=True)
json.dump(items, open(OUT / "bw_data.json", "w"), indent=1)
print("wrote", OUT / "bw_data.json")
if __name__ == "__main__":
main()
+143
View File
@@ -0,0 +1,143 @@
"""Design C machinery: k-loop the prompt, then per-token carry generation.
Structure per sequence: [prompt] [p pause tokens] [answer]
1. prefill: prompt positions looped k times through the merge (settled plan)
2. carry scan: each pause/answer position's band input is
x_t = merge(e_t, s_{t-1})
where s_{t-1} is the *previous position's* final band output — the
latent thought chain runs through the pause positions' workspace states
(token content of pauses is inert: <unused0>).
3. suffix -> logits; CE on answer tokens only.
The scan is sequential per position (RNN-style through the band); each step
re-runs the band on the full sequence — causality keeps settled positions
stable, provisional later positions are recomputed next step anyway.
"""
import torch
from torch.utils.checkpoint import checkpoint
PAUSE_ID = 6 # <unused0>
def _pos_ids(mask):
return (mask.cumsum(-1) - 1).clamp(min=0)
def prompt_prefill(looper, adapter, e, calls, prompt_mask, k):
"""k merge->band loops over prompt positions. Returns (S, X)."""
with torch.no_grad():
s = looper.band(e, calls)
x = e
for _ in range(k):
x = torch.where(prompt_mask[..., None], adapter(e, s), e)
s = looper.band(x, calls)
return s, x
def carry_steps(looper, adapter, e, calls, S, X, step_updates,
use_checkpoint=False):
"""Sequential scan. step_updates: list of (row_idx, pos) index tensors."""
for rows, pos in step_updates:
if rows.numel() == 0:
continue
seed = S[rows, pos - 1]
x_new = adapter(e[rows, pos], seed)
X = X.clone()
X[rows, pos] = x_new.to(X.dtype)
S = (checkpoint(lambda X_: looper.band(X_, calls), X,
use_reentrant=False) if use_checkpoint
else looper.band(X, calls))
return S, X
def build_step_updates(prompt_lens, total_lens, device):
"""For right-padded batches: step j updates row b at prompt_lens[b]+j."""
max_span = int((total_lens - prompt_lens).max())
updates = []
for j in range(max_span):
pos = prompt_lens + j
sel = pos < total_lens
rows = torch.nonzero(sel, as_tuple=False).squeeze(-1)
updates.append((rows.to(device), pos[sel].to(device)))
return updates
def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
k, use_checkpoint=False, feedforward=False):
"""Teacher-forced design-C forward (right-padded batch).
feedforward=True: pause-token control — same positions get the adapter as
x_t = merge(e_t, e_t) in ONE band pass; no state propagates between
positions. Isolates 'pause compute + weights' from the recurrence."""
dev = input_ids.device
calls, _ = looper.capture(input_ids, attention_mask, logits_to_keep=1)
e = looper._hin[looper.l0].detach()
ar = torch.arange(input_ids.shape[1], device=dev)
prompt_mask = ar[None, :] < prompt_lens[:, None].to(dev)
if feedforward:
touched = attention_mask.bool()
x = torch.where(touched[..., None], adapter(e, e), e)
S = (checkpoint(lambda x_: looper.band(x_, calls), x,
use_reentrant=False) if use_checkpoint
else looper.band(x, calls))
return looper.suffix_logits(S, calls)
S, X = prompt_prefill(looper, adapter, e, calls, prompt_mask, k)
total_lens = attention_mask.sum(-1)
updates = build_step_updates(prompt_lens.to(dev), total_lens.to(dev), dev)
S, X = carry_steps(looper, adapter, e, calls, S, X, updates,
use_checkpoint=use_checkpoint)
return looper.suffix_logits(S, calls)
@torch.no_grad()
def generate_carry_c(looper, adapter, tok, input_ids, attention_mask,
k, p, max_new_tokens=10, feedforward=False):
"""Greedy design-C generation (left-padded batch, uniform positions).
Appends p pause tokens, prefill-loops the prompt, carries through the
pauses, then generates with per-token carry."""
B = input_ids.shape[0]
dev = input_ids.device
pauses = torch.full((B, p), PAUSE_ID, dtype=torch.long, device=dev)
ids = torch.cat([input_ids, pauses], 1)
mask = torch.cat([attention_mask,
torch.ones_like(pauses)], 1)
n_prompt = input_ids.shape[1]
eos = {tok.eos_token_id, tok.convert_tokens_to_ids("<end_of_turn>")}
done = torch.zeros(B, dtype=torch.bool, device=dev)
X_store = None # merged band inputs for settled positions
for step in range(max_new_tokens + 1): # step 0 = prefill + pause scan
calls, _ = looper.capture(ids, mask, logits_to_keep=1,
position_ids=_pos_ids(mask))
e = looper._hin[looper.l0]
if feedforward:
x = torch.where(mask.bool()[..., None], adapter(e, e), e)
S = looper.band(x, calls)
X = x
elif X_store is None:
ar = torch.arange(ids.shape[1], device=dev)
prompt_mask = (ar[None, :] < n_prompt) & mask.bool()
S, X = prompt_prefill(looper, adapter, e, calls, prompt_mask, k)
updates = [(torch.arange(B, device=dev),
torch.full((B,), n_prompt + j, device=dev,
dtype=torch.long)) for j in range(p)]
S, X = carry_steps(looper, adapter, e, calls, S, X, updates)
else:
X = torch.cat([X_store, e[:, X_store.shape[1]:]], 1)
rows = torch.arange(B, device=dev)
pos = torch.full((B,), ids.shape[1] - 1, device=dev,
dtype=torch.long)
S = looper.band(X, calls) # settled prefix + provisional new pos
S, X = carry_steps(looper, adapter, e, calls, S, X, [(rows, pos)])
X_store = X
logits = looper.suffix_logits(S, calls, last_only=True)
nxt = logits[:, -1].argmax(-1)
nxt = torch.where(done, torch.full_like(nxt, list(eos)[0]), nxt)
ids = torch.cat([ids, nxt[:, None]], 1)
mask = torch.cat([mask, (~done)[:, None].long()], 1)
done |= torch.tensor([t.item() in eos for t in nxt], device=dev)
if done.all():
break
return ids
+67
View File
@@ -0,0 +1,67 @@
"""Estimate J_l = E_{prompt, t, t'>=t}[dh_final,t'/dh_l,t] over a pretraining-like corpus."""
import argparse, json, sys, time
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model, prompt_jacobian_pairsum
def corpus_texts(n, min_chars=400):
from datasets import load_dataset
ds = load_dataset("HuggingFaceFW/fineweb-edu", name="sample-10BT",
split="train", streaming=True)
got = 0
for ex in ds:
t = ex["text"].strip()
if len(t) >= min_chars:
yield t
got += 1
if got >= n:
return
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--n-prompts", type=int, default=256)
ap.add_argument("--seq-len", type=int, default=64)
ap.add_argument("--chunk", type=int, default=128)
ap.add_argument("--out", default="results/jbar.pt")
ap.add_argument("--dtype", default="float32")
args = ap.parse_args()
model, tok = load_model(dtype=getattr(torch, args.dtype))
L = len(model.model.language_model.layers)
d = model.config.get_text_config().hidden_size
Jsum = torch.zeros(L, d, d, device="cuda", dtype=torch.float32)
pairs_total = 0
t0 = time.time()
for i, text in enumerate(corpus_texts(args.n_prompts)):
ids = tok(text, return_tensors="pt", truncation=True,
max_length=args.seq_len)["input_ids"].cuda()
if ids.shape[1] < args.seq_len:
continue
J, n_pairs = prompt_jacobian_pairsum(model, ids, chunk=args.chunk)
Jsum += J
pairs_total += n_pairs
if i % 5 == 0 or i == args.n_prompts - 1:
el = time.time() - t0
print(f"[{i+1}/{args.n_prompts}] {el:.0f}s ({el/(i+1):.1f}s/prompt)",
flush=True)
if i % 50 == 49: # checkpoint
torch.save({"Jbar": (Jsum / pairs_total).cpu(),
"n_prompts": i + 1, "seq_len": args.seq_len},
args.out + ".ckpt")
Jbar = (Jsum / pairs_total).cpu()
Path(args.out).parent.mkdir(parents=True, exist_ok=True)
torch.save({"Jbar": Jbar, "n_prompts": args.n_prompts,
"seq_len": args.seq_len, "pairs": pairs_total}, args.out)
print("saved", args.out)
if __name__ == "__main__":
main()
+119
View File
@@ -0,0 +1,119 @@
"""Schematic: the J-lens view of the residual stream + depth regimes + looping."""
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.patches import FancyBboxPatch, FancyArrowPatch, Rectangle
from matplotlib.patches import ConnectionPatch
fig = plt.figure(figsize=(13, 8.5))
gs = fig.add_gridspec(1, 2, width_ratios=[1.15, 1], wspace=0.05)
axL = fig.add_subplot(gs[0]); axR = fig.add_subplot(gs[1])
for a in (axL, axR):
a.set_xlim(0, 10); a.set_ylim(0, 10); a.axis("off")
# ---------- LEFT: residual stream with regime bands (gemma-4-E2B, 35 layers) ----------
axL.text(5, 9.6, "The residual stream under the J-lens", ha="center",
fontsize=15, fontweight="bold")
axL.text(5, 9.18, "gemma-4-E2B · 35 layers, d=1536", ha="center",
fontsize=9.5, color="#555")
# central stream bar
sx, sw = 4.15, 1.3
axL.add_patch(Rectangle((sx, 0.7), sw, 7.7, facecolor="#f2f2f2",
edgecolor="#888", lw=1.2, zorder=1))
axL.annotate("", xy=(sx+sw/2, 8.75), xytext=(sx+sw/2, 0.55),
arrowprops=dict(arrowstyle="-|>", color="#444", lw=2), zorder=0)
axL.text(sx+sw/2, 0.32, "token embedding (input)", ha="center", fontsize=8, color="#333")
axL.text(sx+sw/2, 8.95, "logits (output)", ha="center", fontsize=8, color="#333")
def y(layer): return 0.8 + (layer/34)*7.35
bands = [
("transduction", 0, 5, "#d9d9d9", "detokenization /\nlexical assembly",
["J-lens: noise / surface form"]),
("sensor", 6, 13, "#bcd4f0", "input written\ninto the stream",
["J-lens: current input token", "(echo ↑ 26%)"]),
("workspace", 14, 30, "#bfe3c0", "abstract concepts\nheld & broadcast",
["J-lens: unspoken content — 'spider', 'big'",
"≈ 1.08B params (58% of decoder)",
"swap-gates fire here · ignition threshold"]),
("motor", 31, 34, "#f2c4c4", "output staged\nfor unembedding",
["J-lens: next output token (→68%)"]),
]
for name, l0, l1, c, fn, lens_lines in bands:
yb, yt = y(l0)-0.10, y(l1)+0.10
ymid = (yb+yt)/2
axL.add_patch(Rectangle((sx, yb), sw, yt-yb, facecolor=c, edgecolor="#666",
lw=1, alpha=0.95, zorder=2))
axL.text(sx+sw/2, ymid, f"L{l0}{l1}", ha="center", va="center",
fontsize=8.5, fontweight="bold", zorder=3)
# band name above-left inside its band
axL.text(sx-0.2, ymid, name, ha="right", va="center",
fontsize=10.5, fontweight="bold", color="#222")
# function (right, top) + lens lines stacked below
axL.text(sx+sw+0.3, ymid+0.28, fn, ha="left", va="center", fontsize=8, color="#222")
for j, ln in enumerate(lens_lines):
axL.text(sx+sw+0.3, ymid-0.16-0.28*j, ln, ha="left", va="center",
fontsize=6.9, color="#556", style="italic")
# three verbs: write in (sensor), hold (workspace), read out (motor) — offset from
# the band-name row so arrows don't clip the labels
axL.annotate("", xy=(sx-0.02, y(12.3)), xytext=(sx-0.7, y(12.3)),
arrowprops=dict(arrowstyle="-|>", color="#2b6cb0", lw=1.8))
axL.text(sx-0.72, y(12.3)+0.2, "write in", ha="right", fontsize=7.5, color="#2b6cb0")
axL.add_patch(FancyArrowPatch((sx-0.08, y(28.5)), (sx-0.08, y(25.5)),
connectionstyle="arc3,rad=-1.0", arrowstyle="-|>",
color="#2f855a", lw=1.7, mutation_scale=11))
axL.text(sx-0.62, y(27), "hold", ha="right", va="center", fontsize=7.5, color="#2f855a")
axL.annotate("", xy=(sx-0.02, y(33.7)), xytext=(sx-0.7, y(33.7)),
arrowprops=dict(arrowstyle="-|>", color="#c53030", lw=1.8))
axL.text(sx-0.72, y(33.7)+0.2, "read out", ha="right", fontsize=7.5, color="#c53030")
# ---------- RIGHT: the looping result ----------
axR.text(5, 9.7, "Can we loop the workspace?", ha="center",
fontsize=15, fontweight="bold")
def box(ax, x, y0, w, h, fc, ec, txt, fs=8, fw="normal", tc="#111"):
ax.add_patch(FancyBboxPatch((x, y0), w, h, boxstyle="round,pad=0.06,rounding_size=0.12",
facecolor=fc, edgecolor=ec, lw=1.4))
ax.text(x+w/2, y0+h/2, txt, ha="center", va="center", fontsize=fs, fontweight=fw, color=tc)
# Case A: band loop (not a self-map) -- fails
axR.text(0.4, 8.9, "A. Loop the whole band (L14→30)", fontsize=10.5, fontweight="bold")
box(axR, 0.6, 7.3, 2.2, 1.0, "#bfe3c0", "#2f855a", "band\nL14 → L30")
axR.add_patch(FancyArrowPatch((2.9, 8.15), (2.9, 7.45),
connectionstyle="arc3,rad=1.3", arrowstyle="-|>", color="#c53030",
lw=2, mutation_scale=14))
axR.text(4.9, 7.8, "out lives 17 layers\ndownstream of in", fontsize=8, color="#333")
axR.text(0.6, 6.75, "✗ NOT a self-map (out-space ≠ in-space)",
fontsize=9, color="#c53030", fontweight="bold")
axR.text(0.6, 6.4, " → collapses in 1 iteration: '8'')'",
fontsize=8.3, color="#555")
axR.plot([0.4, 9.6], [5.95, 5.95], color="#ccc", lw=1)
# Case B: single tied layer (self-map)
axR.text(0.4, 5.55, "B. Loop one tied layer h ← h + β·Δ(h)", fontsize=10.5, fontweight="bold")
box(axR, 0.6, 4.0, 2.2, 1.0, "#bcd4f0", "#2b6cb0", "layer L\nin = out")
axR.add_patch(FancyArrowPatch((2.9, 4.85), (2.9, 4.15),
connectionstyle="arc3,rad=1.3", arrowstyle="-|>", color="#2b6cb0",
lw=2, mutation_scale=14))
axR.text(4.9, 4.5, "same residual point\nin and out", fontsize=8, color="#333")
axR.text(0.6, 3.45, "✓ self-map — stable; damping (low-pass) extends it",
fontsize=9, color="#2f855a", fontweight="bold")
axR.text(0.6, 3.05, " β=0.5 @ L26 keeps '8' through ~6 loops; |ΔH| ↓ (converges)",
fontsize=8.3, color="#555")
axR.text(0.6, 2.6, "✗ but frozen → degenerate fixed point; concept never sharpens",
fontsize=9, color="#c53030", fontweight="bold")
axR.text(0.6, 2.2, " extra loop compute ≠ extra reasoning",
fontsize=8.3, color="#555")
# verdict box
box(axR, 0.6, 0.5, 9.0, 1.35, "#fff7e6", "#d69e2e",
"Verdict: mechanically loopable at the layer level, and low-pass\n"
"damping is necessary for stability — but not sufficient. Turning loop-depth\n"
"into reasoning needs the loop trained in (LoRA-per-loop / Huginn / MoR).",
fs=8.6, fw="normal", tc="#7a5a00")
fig.savefig("results/jlens_diagram.png", dpi=140, bbox_inches="tight",
facecolor="white")
print("wrote results/jlens_diagram.png")
+83
View File
@@ -0,0 +1,83 @@
"""Blocksworld eval: pass@1 vs k (prompt-only loops, fast path)."""
import argparse
import json
import os
import sys
import time
from pathlib import Path
import torch
from bw_common import verify_plan
from bw_prep import DIRECT_SUFFIX, chat
from loop_common import BandLooper, MergeAdapter
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
@torch.no_grad()
def acc_at_k(looper, adapter, tok, items, k, batch=12, max_new=200):
oks, per_item = [], []
for i in range(0, len(items), batch):
torch.cuda.empty_cache()
chunk = items[i : i + batch]
enc = tok([chat(tok, it, DIRECT_SUFFIX) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
max_new_tokens=max_new,
attention_mask=enc["attention_mask"])
for j, it in enumerate(chunk):
txt = tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True)
ok = verify_plan(it, txt)
oks.append(ok)
per_item.append({"task_id": it["task_id"], "ok": bool(ok)})
by = {}
for it, ok in zip(items, oks):
d = by.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
return (sum(oks) / len(items),
{l: c / n for l, (c, n) in by.items()}, per_item)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="bw")
ap.add_argument("--ks", default="0,2,4")
args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")]
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
items = [it for it in json.load(open(OUT / "bw_data.json"))
if it["split"] == "test"]
print(f"[{args.tag}] Blocksworld eval on {len(items)} items", flush=True)
res = {"tag": args.tag, "ks": {}, "n": len(items)}
for k in ks:
t0 = time.time()
acc, by, per_item = acc_at_k(looper, adapter, tok, items, k)
res["ks"][k] = {"acc": acc, "by_label": by, "per_item": per_item}
print(f"k={k}: acc={acc:.3f} "
f"by_label={ {l: round(v,3) for l,v in by.items()} }"
f" ({time.time()-t0:.0f}s)", flush=True)
json.dump(res, open(OUT / f"eval_bw_{args.tag}.json", "w"), indent=1)
print("wrote", OUT / f"eval_bw_{args.tag}.json")
if __name__ == "__main__":
main()
+90
View File
@@ -0,0 +1,90 @@
"""Eval design C (prefill k + carry through p pauses) on GSM8K test.
Usage: eval_carry.py --adapter ../results-loop/adapter_carry_e600.pt \
--tag carry --grid "0:0,2:0,2:2,2:6"
Grid entries are k:p pairs; k=0,p=0 is the plain baseline (same harness).
"""
import argparse
import json
import sys
import time
from pathlib import Path
import torch
from carry_common import generate_carry_c
from loop_common import (DIRECT_SUFFIX, BandLooper, MergeAdapter, chat_prompt,
last_number, num_eq)
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
@torch.no_grad()
def acc_at(looper, adapter, tok, items, k, p, batch=16, feedforward=False):
hits, per_label = 0, {}
for i in range(0, len(items), batch):
torch.cuda.empty_cache()
chunk = items[i : i + batch]
enc = tok([chat_prompt(tok, it["question"], DIRECT_SUFFIX)
for it in chunk], return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = generate_carry_c(looper, adapter, tok, enc["input_ids"],
enc["attention_mask"], k=k, p=p,
max_new_tokens=10, feedforward=feedforward)
n0 = enc["input_ids"].shape[1] + p
for j, it in enumerate(chunk):
txt = tok.decode(gen[j, n0:], skip_special_tokens=True)
ok = num_eq(last_number(txt), it["gold"])
hits += ok
d = per_label.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
return hits / len(items), {l: c / n for l, (c, n) in per_label.items()}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="carry")
ap.add_argument("--grid", default="0:0,2:0,2:2,2:6")
ap.add_argument("--n", type=int, default=0)
ap.add_argument("--feedforward", action="store_true")
args = ap.parse_args()
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter().cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
items = [it for it in json.load(open(OUT / "star_data.json"))
if it["split"] == "test"]
if args.n:
items = items[: args.n]
print(f"[{args.tag}] GSM8K carry eval on {len(items)} items, "
f"grid={args.grid}", flush=True)
res = {"tag": args.tag, "grid": {}, "n": len(items)}
for kp in args.grid.split(","):
k, p = (int(x) for x in kp.split(":"))
t0 = time.time()
acc, by_label = acc_at(looper, adapter, tok, items, k, p,
feedforward=args.feedforward)
res["grid"][kp] = {"acc": acc, "by_label": by_label}
print(f"k={k} p={p}: acc={acc:.3f} "
f"by_label={ {l: round(v,3) for l,v in by_label.items()} }"
f" ({time.time()-t0:.0f}s)", flush=True)
with open(OUT / f"eval_{args.tag}.json", "w") as f:
json.dump(res, f, indent=1)
print("wrote", OUT / f"eval_{args.tag}.json")
if __name__ == "__main__":
main()
+143
View File
@@ -0,0 +1,143 @@
"""HumanEval transfer eval: does the MBPP-trained loop adapter generalize?
No HumanEval training exists (no train split) — this is pure distribution
transfer. Labels the 164 items with the frozen model (direct vs terse-plan,
greedy; descriptive buckets only), then evaluates the loop adapter at
k grid + the untrained control.
"""
import argparse
import json
import os
import subprocess
import sys
import tempfile
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from datasets import load_dataset
from loop_common import BandLooper, MergeAdapter
from prep_mbpp import batch_generate, extract_code
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
DIRECT = ("Complete the following Python function. Return the COMPLETE "
"function (signature included) in a ```python code block. "
"No explanation.\n\n```python\n{prompt}```")
PLAN = ("First write a very brief plan: at most 4 short bullet lines. Then "
"return the COMPLETE function (signature included) in a ```python "
"code block.\n\n```python\n{prompt}```")
def he_prompt(tok, item, tmpl):
return tok.apply_chat_template(
[{"role": "user", "content": tmpl.format(prompt=item["prompt"])}],
tokenize=False, add_generation_prompt=True)
def run_he_tests(code, item, timeout=10):
if not code:
return False
script = (code + "\n\n" + item["test"] +
f"\ncheck({item['entry_point']})\n")
try:
with tempfile.TemporaryDirectory() as td:
r = subprocess.run([sys.executable, "-c", script], cwd=td,
capture_output=True, timeout=timeout)
return r.returncode == 0
except (subprocess.TimeoutExpired, OSError):
return False
@torch.no_grad()
def loop_eval(looper, adapter, tok, items, k, batch=8, max_new=380):
codes = []
for i in range(0, len(items), batch):
torch.cuda.empty_cache()
chunk = items[i : i + batch]
enc = tok([he_prompt(tok, it, DIRECT) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
max_new_tokens=max_new,
attention_mask=enc["attention_mask"])
for j in range(len(chunk)):
codes.append(extract_code(
tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True)))
with ThreadPoolExecutor(8) as ex:
oks = list(ex.map(lambda ci: run_he_tests(ci[0], ci[1]),
zip(codes, items)))
return oks
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="he")
ap.add_argument("--ks", default="0,2,4")
args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")]
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
items = list(load_dataset("openai/openai_humaneval")["test"])
print(f"[{args.tag}] HumanEval: {len(items)} items", flush=True)
# labeling pass (descriptive buckets; greedy — outcome-selection caveat)
lab_path = OUT / "humaneval_labels.json"
if lab_path.exists():
labels = json.load(open(lab_path))
else:
plans = batch_generate(model, tok,
[he_prompt(tok, it, PLAN) for it in items],
max_new_tokens=800, batch_size=16)
with ThreadPoolExecutor(8) as ex:
plan_ok = list(ex.map(
lambda gi: run_he_tests(extract_code(gi[0]), gi[1]),
zip(plans, items)))
labels = {it["task_id"]: bool(ok) for it, ok in zip(items, plan_ok)}
json.dump(labels, open(lab_path, "w"), indent=1)
print(f"plan-reachable: {sum(labels.values())}/{len(items)}", flush=True)
res = {"tag": args.tag, "ks": {}, "n": len(items)}
k0_ok = None
for k in ks:
t0 = time.time()
oks = loop_eval(looper, adapter, tok, items, k)
if k == 0:
k0_ok = oks
hard = [i for i, it in enumerate(items)
if k0_ok and not k0_ok[i] and labels[it["task_id"]]]
acc = sum(oks) / len(items)
hard_acc = (sum(oks[i] for i in hard) / len(hard)) if hard else None
res["ks"][k] = {"acc": acc, "hard_n": len(hard),
"hard_acc": hard_acc,
"per_item": [{"task_id": it["task_id"],
"ok": bool(o)}
for it, o in zip(items, oks)]}
print(f"k={k}: pass@1={acc:.3f} hard({len(hard)})="
f"{hard_acc if hard_acc is None else round(hard_acc,3)} "
f"({time.time()-t0:.0f}s)", flush=True)
json.dump(res, open(OUT / f"eval_humaneval_{args.tag}.json", "w"),
indent=1)
print("wrote", OUT / f"eval_humaneval_{args.tag}.json")
if __name__ == "__main__":
main()
+143
View File
@@ -0,0 +1,143 @@
"""Eval the looped-band model: accuracy vs k, and J-lens concept sharpening.
Usage:
eval_loop.py # untrained adapter (alpha-merge only)
eval_loop.py --adapter results-loop/adapter.pt --tag trained
"""
import argparse
import json
import os
import time
from pathlib import Path
import torch
from loop_common import (DIRECT_SUFFIX, BandLooper, MergeAdapter, chat_prompt,
last_number, num_eq)
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import JLens, load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
ROOT = OUT.parent
@torch.no_grad()
def accuracy_at_k(looper, adapter, tok, items, k, batch=16, prompt_only=False,
feedforward=False):
hits, per_label, per_item = 0, {}, []
for i in range(0, len(items), batch):
torch.cuda.empty_cache()
chunk = items[i : i + batch]
enc = tok([chat_prompt(tok, it["question"], DIRECT_SUFFIX) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
if prompt_only:
gen = looper.generate_frozen_prompt(
adapter, tok, enc["input_ids"], k, max_new_tokens=10,
attention_mask=enc["attention_mask"], feedforward=feedforward)
else:
gen = looper.loop_generate(adapter, tok, enc["input_ids"], k,
max_new_tokens=10,
attention_mask=enc["attention_mask"])
for j, it in enumerate(chunk):
txt = tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True)
ok = num_eq(last_number(txt), it["gold"])
per_item.append({"idx": it["idx"], "ok": bool(ok)})
hits += ok
d = per_label.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
return (hits / len(items),
{l: c / n for l, (c, n) in per_label.items()}, per_item)
@torch.no_grad()
def spider_sharpening(looper, adapter, model, tok, ks):
"""P('spider') under the J-lens at L30, and P('8') at the output, vs k."""
jbar = torch.load(ROOT / "results" / "jbar.pt", map_location="cuda")
Jbar = (jbar["Jbar"] if isinstance(jbar, dict) else jbar).float()
lens = JLens(model, tok, Jbar)
q = ("The animal that spins webs has how many legs? "
"Answer with just the number.")
ids = tok.apply_chat_template([{"role": "user", "content": q}],
return_tensors="pt", add_generation_prompt=True,
return_dict=True)["input_ids"].cuda()
spider = tok.encode(" spider", add_special_tokens=False)[0]
eight = [tok.encode(t, add_special_tokens=False)[0] for t in (" 8", "8")]
calls, _ = looper.capture(ids)
e = looper._hin[looper.l0]
s = looper.band(e, calls)
rows = []
kmax = max(ks)
for k in range(0, kmax + 1):
if k > 0:
s = looper.band(adapter(e, s), calls)
if k in ks:
_, probs = lens.read(s[0], looper.l1, topk=1)
logits = lens._readout((s[0].float() @ Jbar[looper.l1].T))
p_spider = torch.softmax(logits.float(), -1)[:, spider].max().item()
out = looper.suffix_logits(s, calls)
p8 = torch.softmax(out[0, -1].float(), -1)[eight].max().item()
rows.append({"k": k, "P_spider_lens": p_spider, "P_8_out": p8})
return rows
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="untrained")
ap.add_argument("--ks", default="0,1,2,4,8")
ap.add_argument("--n", type=int, default=0, help="cap test items (0=all)")
ap.add_argument("--prompt-only", action="store_true",
help="loop the prompt span only (unified regime, fast path)")
ap.add_argument("--no-spider", action="store_true")
ap.add_argument("--feedforward", action="store_true",
help="no-recurrence control arm: adapter(e,e) once")
args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")]
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
items = [it for it in json.load(open(OUT / "star_data.json"))
if it["split"] == "test"]
if args.n:
items = items[: args.n]
print(f"[{args.tag}] eval on {len(items)} test items, ks={ks}", flush=True)
res = {"tag": args.tag, "ks": {}, "n": len(items)}
for k in ks:
t0 = time.time()
acc, by_label, per_item = accuracy_at_k(looper, adapter, tok, items, k,
prompt_only=args.prompt_only,
feedforward=args.feedforward)
res["ks"][k] = {"acc": acc, "by_label": by_label,
"per_item": per_item}
print(f"k={k}: acc={acc:.3f} by_label={ {l: round(v,3) for l,v in by_label.items()} }"
f" ({time.time()-t0:.0f}s)", flush=True)
res["spider"] = ([] if args.no_spider else
spider_sharpening(looper, adapter, model, tok, set(ks)))
for r in res["spider"]:
print(f"spider k={r['k']}: P_lens={r['P_spider_lens']:.3f} "
f"P(8)={r['P_8_out']:.3f}", flush=True)
OUT.mkdir(exist_ok=True)
with open(OUT / f"eval_{args.tag}.json", "w") as f:
json.dump(res, f, indent=1)
print("wrote", OUT / f"eval_{args.tag}.json")
if __name__ == "__main__":
main()
+112
View File
@@ -0,0 +1,112 @@
"""Eval MBPP pass@1 vs loop depth k with prompt-only ("latent planning") loops.
Usage:
eval_loop_code.py --tag untrained
eval_loop_code.py --adapter ../results-loop/adapter_code.pt --tag trained
"""
import argparse
import json
import os
import sys
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from loop_common import BandLooper, MergeAdapter
from prep_mbpp import DIRECT_SUFFIX, extract_code, mbpp_prompt, run_tests
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
@torch.no_grad()
def pass1_at_k(looper, adapter, tok, items, k, batch=8, max_new=220,
feedforward=False, pause=0):
codes = []
for i in range(0, len(items), batch):
torch.cuda.empty_cache()
chunk = items[i : i + batch]
enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
if pause:
B = enc["input_ids"].shape[0]
pcol = torch.full((B, pause), 6, dtype=torch.long, device="cuda")
enc["input_ids"] = torch.cat([enc["input_ids"], pcol], 1)
enc["attention_mask"] = torch.cat(
[enc["attention_mask"], torch.ones_like(pcol)], 1)
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
max_new_tokens=max_new,
attention_mask=enc["attention_mask"],
feedforward=feedforward)
for j in range(len(chunk)):
txt = tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True)
codes.append(extract_code(txt))
with ThreadPoolExecutor(8) as ex:
oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]),
zip(codes, items)))
hits, per_label = 0, {}
for it, ok in zip(items, oks):
hits += ok
d = per_label.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
per_item = [{"task_id": it["task_id"], "ok": bool(ok)}
for it, ok in zip(items, oks)]
return (hits / len(items),
{l: c / n for l, (c, n) in per_label.items()}, per_item)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="untrained")
ap.add_argument("--ks", default="0,1,2,4")
ap.add_argument("--n", type=int, default=250)
ap.add_argument("--feedforward", action="store_true",
help="no-recurrence control arm: adapter(e,e) once")
ap.add_argument("--pause", type=int, default=0,
help="append p pause tokens to each prompt")
args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")]
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
items = [it for it in json.load(open(OUT / "mbpp_data.json"))
if it["split"] == "test"][: args.n]
print(f"[{args.tag}] MBPP eval on {len(items)} test items, ks={ks}",
flush=True)
res = {"tag": args.tag, "ks": {}, "n": len(items)}
for k in ks:
t0 = time.time()
acc, by_label, per_item = pass1_at_k(looper, adapter, tok, items, k,
feedforward=args.feedforward,
pause=args.pause)
res["ks"][k] = {"acc": acc, "by_label": by_label,
"per_item": per_item}
print(f"k={k}: pass@1={acc:.3f} "
f"by_label={ {l: round(v,3) for l,v in by_label.items()} }"
f" ({time.time()-t0:.0f}s)", flush=True)
with open(OUT / f"eval_code_{args.tag}.json", "w") as f:
json.dump(res, f, indent=1)
print("wrote", OUT / f"eval_code_{args.tag}.json")
if __name__ == "__main__":
main()
+54
View File
@@ -0,0 +1,54 @@
"""Explicit-planning reference: frozen model, plan-first prompt, MBPP test.
The number latent looping is compared against (equal-or-more FLOPs: ~200-400
visible plan tokens vs k band passes over the prompt)."""
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from prep_mbpp import batch_generate, extract_code, mbpp_prompt, run_tests
from prep_mbpp_fix import TERSE_PLAN_SUFFIX
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
items = [it for it in json.load(open(OUT / "mbpp_data.json"))
if it["split"] == "test"]
t0 = time.time()
gens = batch_generate(model, tok,
[mbpp_prompt(tok, it, TERSE_PLAN_SUFFIX)
for it in items], max_new_tokens=700, batch_size=16)
codes = [extract_code(g) for g in gens]
with ThreadPoolExecutor(8) as ex:
oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]),
zip(codes, items)))
by = {}
for it, ok in zip(items, oks):
d = by.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
per_item = [{"task_id": it["task_id"], "ok": bool(ok)}
for it, ok in zip(items, oks)]
res = {"acc": sum(oks) / len(items),
"by_label": {l: c / n for l, (c, n) in by.items()},
"per_item": per_item}
print(f"plan-first pass@1={res['acc']:.3f} "
f"by_label={ {l: round(v,3) for l,v in res['by_label'].items()} }"
f" ({time.time()-t0:.0f}s)", flush=True)
json.dump(res, open(OUT / "eval_plan_baseline.json", "w"), indent=1)
print("wrote", OUT / "eval_plan_baseline.json")
if __name__ == "__main__":
main()
+130
View File
@@ -0,0 +1,130 @@
"""Experiment 1: J-lens readouts.
(a) Two-hop reasoning: unspoken intermediate 'spider' visible mid-network.
(b) Multilingual: English intermediates during a Chinese task.
(c) Directed modulation: 'hold citrus in mind' while copying unrelated text.
(d) Layer profile: where in depth the lens carries abstract content.
"""
import os, sys
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import (JLens, chat_ids, collect_residuals,
generate_with_residuals, load_model)
RES = Path(os.environ.get("JLENS_RESULTS", "results"))
def show_table(jl, tok, ids, layers, positions, topk=6, header=""):
"""Print top J-lens tokens at (layer, position) grid."""
hs = collect_residuals(jl.model, ids)
toks = [tok.decode([t]) for t in ids[0].tolist()]
print(f"\n=== {header} ===")
print("positions:", {p: repr(toks[p]) for p in positions})
for l in layers:
row = []
for p in positions:
idx, prob = jl.read(hs[l, p], l, topk=topk)
row.append(" ".join(
(tok.decode([i]).strip() or "·") for i in idx.tolist()[:topk]))
print(f"L{l:>2} | " + " || ".join(row))
def concept_heatmap(jl, tok, ids, words, tag, hs=None):
"""P(concept tokens) by (layer, position); save tensor + print peak."""
if hs is None:
hs = collect_residuals(jl.model, ids)
tids = []
for w in words:
for v in (w, " " + w, w.capitalize(), " " + w.capitalize()):
e = tok.encode(v, add_special_tokens=False)
if len(e) == 1:
tids.append(e[0])
tids = sorted(set(tids))
L, T, _ = hs.shape
out = torch.zeros(L, T)
for l in range(L):
logits = jl._readout(hs[l].float() @ jl.Jbar[l].T)
probs = torch.softmax(logits.float(), dim=-1)
out[l] = probs[:, tids].sum(-1).cpu()
peak = out.max().item()
lmax, pmax = divmod(out.argmax().item(), T)
toks = [tok.decode([t]) for t in ids[0].tolist()] if ids is not None else None
print(f"[{tag}] words={words} peak P={peak:.3f} at layer {lmax}, "
f"pos {pmax}" + (f" ({toks[pmax]!r})" if toks and pmax < len(toks) else ""))
RES.mkdir(exist_ok=True)
torch.save({"map": out, "words": words,
"tokens": [tok.decode([t]) for t in ids[0].tolist()] if ids is not None else None},
RES / f"heat_{tag}.pt")
return out
def main():
model, tok = load_model(dtype=torch.bfloat16)
ck = torch.load(sys.argv[1] if len(sys.argv) > 1 else "results/jbar.pt",
map_location="cuda")
jl = JLens(model, tok, ck["Jbar"].cuda())
print(f"Jbar from {ck.get('n_prompts')} prompts")
L = len(model.model.language_model.layers)
mid = range(max(2, L // 6), L - 4, max(2, L // 16))
# (a) two-hop: spider
ids = chat_ids(tok, "The animal that spins webs has how many legs? "
"Answer with just a number.")
print("answer:", jl.generate(ids, 5))
show_table(jl, tok, ids, mid, list(range(ids.shape[1] - 10, ids.shape[1])),
header="two-hop spider (last 10 positions)")
concept_heatmap(jl, tok, ids, ["spider", "spiders"], "spider")
concept_heatmap(jl, tok, ids, ["eight", "8"], "eight")
# (b) multilingual: Chinese antonym of small -> big
ids = chat_ids(tok, "小的反义词是什么?只用一个字回答。")
print("\nanswer:", jl.generate(ids, 5))
show_table(jl, tok, ids, mid, list(range(ids.shape[1] - 8, ids.shape[1])),
header="Chinese antonym of 小")
concept_heatmap(jl, tok, ids, ["big", "large", "bigger"], "big_en")
# (c) directed modulation: think of citrus while copying text
copy_text = "The committee will meet on Thursday to review the budget."
for cond, instr in [
("citrus", "While you copy the text, silently think about citrus "
"fruits the entire time. Copy this text exactly, output "
f"nothing else: \"{copy_text}\""),
("control", f"Copy this text exactly, output nothing else: \"{copy_text}\""),
]:
ids = chat_ids(tok, instr)
text, all_ids, hs = generate_with_residuals(model, tok, ids, 30)
print(f"\n[{cond}] output: {text!r}")
# only look at generated positions
gen0 = ids.shape[1]
heat = concept_heatmap(jl, tok, all_ids.unsqueeze(0),
["citrus", "lemon", "orange", "lime"],
f"citrus_{cond}", hs=hs)
print(f" citrus P over generated positions: mean "
f"{heat[:, gen0:].mean():.4f} max {heat[:, gen0:].max():.4f}")
# (d) layer profile: lens entropy + top-token agreement with input/output
ids = chat_ids(tok, "Write one sentence about the ocean.")
text, all_ids, hs = generate_with_residuals(model, tok, ids, 20)
L, T, _ = hs.shape
ent, agree_in, agree_next = [], [], []
for l in range(L):
logits = jl._readout(hs[l].float() @ jl.Jbar[l].T)
probs = torch.softmax(logits.float(), dim=-1)
e = -(probs * probs.clamp_min(1e-12).log()).sum(-1).mean()
top = probs.argmax(-1)
ent.append(e.item())
agree_in.append((top[:-1] == all_ids[:len(top) - 1].cuda()).float().mean().item())
agree_next.append((top[:-1] == all_ids[1:len(top)].cuda()).float().mean().item())
print("\nlayer | lens entropy | top==current tok | top==next tok")
for l in range(0, L, 2):
print(f"L{l:>2} | {ent[l]:8.2f} | {agree_in[l]:.2f} | {agree_next[l]:.2f}")
torch.save({"entropy": ent, "agree_in": agree_in, "agree_next": agree_next},
RES / "layer_profile.pt")
if __name__ == "__main__":
main()
+103
View File
@@ -0,0 +1,103 @@
"""Experiment 2: causal swap interventions gated by the J-lens.
At each (layer, position) where the J-lens reads the source concept, transfer
the activation content source -> target (embedding write basis; see core.py).
(a) Two-hop: swap spider<->ant in the workspace -> answer flips 8 -> 6.
(b) Broadcast: swap France->China under different templates.
(c) Capital-question grid over country pairs.
"""
import itertools, os, sys
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import JLens, chat_ids, load_model
RES = Path(os.environ.get("JLENS_RESULTS", "results"))
def main():
model, tok = load_model(dtype=torch.bfloat16)
ck = torch.load(sys.argv[1] if len(sys.argv) > 1 else "results/jbar.pt",
map_location="cuda")
jl = JLens(model, tok, ck["Jbar"].cuda())
print(f"Jbar from {ck.get('n_prompts')} prompts")
# (a) spider -> ant
SP = [(" spider", " ant"), (" spiders", " ants"), ("spider", "ant")]
q = "The animal that spins webs has how many legs? Answer with just a number."
ids = chat_ids(tok, q)
print("\n--- two-hop swap: spider -> ant (baseline:",
repr(jl.generate(ids, 6)), ") ---")
for thr in (0.002, 0.005, 0.01):
for alpha in (1.0, 1.5, 2.0):
out, fired = jl.generate_swapped(ids, SP, thr=thr, alpha=alpha,
max_new_tokens=6)
print(f"thr={thr} a={alpha}: {out!r} (fired {len(fired)} slots)")
# paper write-basis for comparison
out, _ = jl.generate_swapped(ids, SP, thr=0.005, alpha=1.0, write="jlens",
max_new_tokens=6)
print(f"paper jlens-write thr=0.005 a=1.0: {out!r}")
# control: unrelated swap
out, _ = jl.generate_swapped(ids, [(" spider", " piano")], thr=0.005,
alpha=1.0, max_new_tokens=6)
print(f"control spider->piano: {out!r}")
# reverse: ant prompt -> spider
q2 = ("The insect that builds colonies and lifts many times its own "
"weight has how many legs? Answer with just a number.")
ids2 = chat_ids(tok, q2)
print("\nant prompt baseline:", repr(jl.generate(ids2, 6)))
for alpha in (1.0, 1.5, 2.0):
out, fired = jl.generate_swapped(
ids2, [(" ant", " spider"), (" ants", " spiders"), ("ant", "spider")],
thr=0.005, alpha=alpha, max_new_tokens=6)
print(f"swap ant->spider a={alpha}: {out!r} (fired {len(fired)})")
# (b) broadcast France -> China across templates
print("\n--- broadcast: France -> China (thr=0.01, a=1.0) ---")
FR = [(" France", " China"), ("France", "China"), (" French", " Chinese")]
templates = [
("capital", "What is the capital of France? Answer with just the city name."),
("language", "What language is spoken in France? Answer with one word."),
("continent", "Which continent is France on? Answer with one word."),
("currency", "What currency is used in France? Answer with one word."),
("river", "Name a famous river in France. Answer with one word."),
("food", "Name a famous dish from France. Answer with a short phrase."),
]
for tag, qq in templates:
ids = chat_ids(tok, qq)
base = jl.generate(ids, 8)
sw, fired = jl.generate_swapped(ids, FR, thr=0.01, alpha=1.0,
max_new_tokens=8)
print(f"[{tag:9}] base={base!r:30} swapped={sw!r} ({len(fired)} slots)")
# (c) capital grid over country pairs
print("\n--- capital-question swap grid (thr=0.01, a=1.0) ---")
countries = {"France": ("Paris", "French"), "China": ("Beijing", "Chinese"),
"Japan": ("Tokyo", "Japanese"), "Egypt": ("Cairo", "Egyptian"),
"Brazil": ("Brasília", "Brazilian"), "Canada": ("Ottawa", "Canadian")}
hits = tries = 0
rows = []
for (src, (scap, sadj)), (tgt, (tcap, tadj)) in itertools.permutations(
countries.items(), 2):
qq = f"What is the capital of {src}? Answer with just the city name."
ids = chat_ids(tok, qq)
pairs = [(" " + src, " " + tgt), (src, tgt), (" " + sadj, " " + tadj)]
sw, fired = jl.generate_swapped(ids, pairs, thr=0.01, alpha=1.0,
max_new_tokens=8)
ok = tcap.lower().replace("í", "i").split()[0][:5] in \
sw.lower().replace("í", "i")
hits += ok
tries += 1
rows.append((src, tgt, tcap, sw, ok))
print(f"{src:>7}->{tgt:<7} expect {tcap:<9} got {sw!r} {'OK' if ok else ''}")
print(f"\nswap grid success: {hits}/{tries}")
torch.save(rows, RES / "swap_grid.pt")
if __name__ == "__main__":
main()
+86
View File
@@ -0,0 +1,86 @@
"""Experiment 3: J-space occupancy.
Decompose residual activations as sparse non-negative combinations of k<=25
J-lens vectors (matching pursuit + NNLS refit) and measure the fraction of
activation variance the J-space carries per layer (paper: ~10%, intermediate
layers only).
"""
import os, sys
from pathlib import Path
import numpy as np
import torch
from scipy.optimize import nnls
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import JLens, chat_ids, collect_residuals, load_model
RES = Path(os.environ.get("JLENS_RESULTS", "results"))
K = 25
PROMPTS = [
"The animal that spins webs has how many legs? Answer with just a number.",
"Write one sentence about the ocean.",
"What is the capital of France? Answer with just the city name.",
"Explain photosynthesis in one sentence.",
]
def pursuit_r2(D, h, k=K):
"""Non-negative matching pursuit of h onto rows of D. Returns R^2, support."""
r = h.clone()
sel = []
for _ in range(k):
scores = D @ r
if sel:
scores[torch.tensor(sel, device=D.device)] = -1e30
j = int(scores.argmax())
if scores[j] <= 0:
break
sel.append(j)
A = D[sel].T # (d, |sel|)
x, _ = nnls(A.detach().cpu().numpy().astype(np.float64),
h.detach().cpu().numpy().astype(np.float64))
approx = A @ torch.tensor(x, device=D.device, dtype=D.dtype)
r = h - approx
r2 = 1 - (r.norm() / h.norm()) ** 2
return float(r2), sel
def main():
model, tok = load_model(dtype=torch.bfloat16)
ck = torch.load(sys.argv[1] if len(sys.argv) > 1 else "results/jbar.pt",
map_location="cuda")
jl = JLens(model, tok, ck["Jbar"].cuda())
WU = model.lm_head.weight.detach().float() # (V, d)
L = len(model.model.language_model.layers)
layers = list(range(1, L, 3))
r2_by_layer = {l: [] for l in layers}
for p in PROMPTS:
ids = chat_ids(tok, p)
hs = collect_residuals(model, ids)
T = ids.shape[1]
positions = list(range(max(1, T - 12), T)) # skip BOS region
for l in layers:
D = WU @ jl.Jbar[l] # (V, d) J-lens dictionary at layer l
D = D / D.norm(dim=1, keepdim=True).clamp_min(1e-8)
for t in positions:
r2, _ = pursuit_r2(D, hs[l, t].float())
r2_by_layer[l].append(r2)
del D
torch.cuda.empty_cache()
print(f"done prompt: {p[:40]}...", flush=True)
print("\nlayer | mean R^2 of k<=25 non-negative J-lens pursuit")
means = {}
for l in layers:
means[l] = float(np.mean(r2_by_layer[l]))
print(f"L{l:>2} | {means[l]:.3f}")
torch.save(means, RES / "jspace_r2.pt")
if __name__ == "__main__":
main()
+142
View File
@@ -0,0 +1,142 @@
"""Experiment 4: locate sensor / workspace / motor regimes by depth.
Per-layer diagnostics averaged over a prompt set:
sensor = frac. of positions where top J-lens token == current input token
motor = frac. where top J-lens token == NEXT input token (teacher-forced)
persist = mean Jaccard overlap of top-10 lens tokens at adjacent positions
(workspace content should persist across positions)
content = frac. of positions whose top lens token is a content word
(alphabetic, len>=3, not in a junk list)
Ignition test (paper: ambiguous inputs produce sharp binary commitment at
workspace onset): replace one token's embedding (main + per-layer) with a
w-mixture of two concept embeddings and track lens commitment
C = (P1-P2)/(P1+P2) by layer at a downstream position.
"""
import os, sys
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import JLens, chat_ids, collect_residuals, load_model
RES = Path(os.environ.get("JLENS_RESULTS", "results"))
PROMPTS = [
"The animal that spins webs has how many legs? Answer with just a number.",
"What is the capital of France? Answer with just the city name.",
"Write one sentence about the ocean.",
"Explain photosynthesis in one sentence.",
"Name a famous river in Egypt.",
"What language is spoken in Brazil? Answer with one word.",
"Summarize the plot of Romeo and Juliet in one sentence.",
"If I have 3 apples and eat one, how many are left?",
]
JUNK = set("·.<>|/\\()[]{}:;,!?\"'`~*#-_=+ \n\t")
def regime_profile(jl, tok, model):
L = len(jl.tm.layers)
sensor = torch.zeros(L)
motor = torch.zeros(L)
persist = torch.zeros(L)
content = torch.zeros(L)
n_pos = 0
n_adj = 0
for p in PROMPTS:
ids = chat_ids(tok, p)
hs = collect_residuals(model, ids)
T = ids.shape[1]
sl = slice(4, T - 5) # skip bos/turn tokens and trailing template
cur = ids[0, sl].cuda()
nxt = ids[0, 4 + 1:T - 4].cuda()
for l in range(L):
idx, _ = jl.read(hs[l, sl], l, topk=10) # (T', 10)
top1 = idx[:, 0]
sensor[l] += (top1 == cur).sum().item()
motor[l] += (top1 == nxt).sum().item()
sets = [set(r.tolist()) for r in idx]
for a, b in zip(sets, sets[1:]):
persist[l] += len(a & b) / len(a | b)
for t in top1.tolist():
s = tok.decode([t]).strip()
content[l] += (len(s) >= 3 and s.isalpha())
n_pos += cur.numel()
n_adj += cur.numel() - 1
return sensor / n_pos, motor / n_pos, persist / n_adj, content / n_pos
def ignition(jl, tok, model, pairs, ws=(0.0, 0.25, 0.5, 0.75, 1.0)):
"""Mix two concept embeddings at one position; commitment by layer."""
tm = jl.tm
L = len(tm.layers)
template = ("My favorite thing in the world is the X ."
" I think about it every single day because")
out = {}
for w1, w2 in pairs:
id1 = tok.encode(" " + w1, add_special_tokens=False)[0]
id2 = tok.encode(" " + w2, add_special_tokens=False)[0]
ids = tok(template, return_tensors="pt")["input_ids"].cuda()
pos = (ids[0] == tok.encode(" X", add_special_tokens=False)[0]) \
.nonzero()[0].item()
has_ple = getattr(tm, "embed_tokens_per_layer", None) is not None
with torch.no_grad():
pair_ids = torch.tensor([[id1, id2]], device="cuda")
rows_main = tm.embed_tokens(pair_ids)[0].detach()
rows_ple = tm.embed_tokens_per_layer(pair_ids)[0].detach() if has_ple else None
curves = torch.zeros(len(ws), L)
for wi, w in enumerate(ws):
def make_hook(rows, w=w, pos=pos):
def mix_hook(mod, inp, out_e):
e = out_e.clone()
e[0, pos] = w * rows[0] + (1 - w) * rows[1]
return e
return mix_hook
h1 = tm.embed_tokens.register_forward_hook(make_hook(rows_main))
h2 = tm.embed_tokens_per_layer.register_forward_hook(make_hook(rows_ple)) if has_ple else None
try:
hs = collect_residuals(model, ids)
finally:
h1.remove()
if h2: h2.remove()
rpos = pos + 3 # downstream read position
for l in range(L):
logits = jl._readout(hs[l, rpos].float() @ jl.Jbar[l].T)
probs = torch.softmax(logits.float(), -1)
p1, p2 = probs[id1].item(), probs[id2].item()
curves[wi, l] = (p1 - p2) / (p1 + p2 + 1e-12)
out[(w1, w2)] = curves
return out
def main():
model, tok = load_model(dtype=torch.bfloat16)
ck = torch.load(sys.argv[1] if len(sys.argv) > 1 else "results/jbar.pt",
map_location="cuda")
jl = JLens(model, tok, ck["Jbar"].cuda())
sensor, motor, persist, content = regime_profile(jl, tok, model)
print("layer | sensor(top==cur) motor(top==next) persist(top10 Jaccard) content-word")
for l in range(len(sensor)):
bars = lambda x: "#" * int(20 * x)
print(f"L{l:>2} | {sensor[l]:.2f} {bars(sensor[l]):<20} | "
f"{motor[l]:.2f} {bars(motor[l]):<20} | "
f"{persist[l]:.2f} | {content[l]:.2f}")
pairs = [("dog", "piano"), ("ocean", "violin"), ("dragon", "bicycle")]
ign = ignition(jl, tok, model, pairs)
print("\nignition: commitment C=(P1-P2)/(P1+P2) at read pos, by layer")
print("pair | w: " + " ".join(f"{w:4.2f}" for w in (0.0, .25, .5, .75, 1.0)))
for (w1, w2), curves in ign.items():
for l in range(2, len(sensor), 4):
print(f"{w1}/{w2:<9} L{l:>2} | " +
" ".join(f"{curves[wi, l]:+.2f}" for wi in range(5)))
torch.save({"sensor": sensor, "motor": motor, "persist": persist,
"content": content, "ignition": ign}, RES / "regimes.pt")
if __name__ == "__main__":
main()
+126
View File
@@ -0,0 +1,126 @@
"""Per-prompt gate: probe on the k=0 workspace state predicts 'will looping help?'
1. Features: L30 residual at the last prompt position (plain forward), MBPP
train items; labels easy(0)/hard(1) from the STaR pass (free supervision).
2. Logistic probe (d=1536 -> 1), class-balanced.
3. Gated eval on MBPP test: predicted-easy -> k=0, predicted-hard -> k=4 with
the dedicated loop adapter. Reports gated pass@1 vs uniform-k references.
"""
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from loop_common import BandLooper, MergeAdapter
from prep_mbpp import DIRECT_SUFFIX, extract_code, mbpp_prompt, run_tests
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import ResidualCapture, load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
LAYER = 30
K_HARD = 4
@torch.no_grad()
def collect_states(model, tok, items, batch=16):
feats = []
for i in range(0, len(items), batch):
chunk = items[i : i + batch]
enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
with ResidualCapture(model, layers=[LAYER]) as cap:
model(**enc, use_cache=False, logits_to_keep=1)
h = cap.acts[LAYER] # (B, T, d); left padding -> last pos is prompt end
feats.append(h[:, -1].float().cpu())
return torch.cat(feats)
def fit_probe(X, y, epochs=300, lr=0.05):
mu, sd = X.mean(0), X.std(0) + 1e-6
Xn = (X - mu) / sd
w = torch.zeros(X.shape[1], requires_grad=True)
b = torch.zeros(1, requires_grad=True)
opt = torch.optim.Adam([w, b], lr=lr)
pos_w = (y == 0).sum() / max(1, (y == 1).sum())
for _ in range(epochs):
z = Xn @ w + b
loss = torch.nn.functional.binary_cross_entropy_with_logits(
z, y.float(), pos_weight=pos_w)
opt.zero_grad()
loss.backward()
opt.step()
return w.detach(), b.detach(), mu, sd
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter().cuda()
adapter.load_state_dict(torch.load(OUT / "adapter_code_e399.pt",
map_location="cuda"))
data = json.load(open(OUT / "mbpp_data.json"))
train = [it for it in data if it["split"] == "train"
and it["label"] in ("easy", "hard")]
test = [it for it in data if it["split"] == "test"][:250]
print(f"collecting states: {len(train)} train, {len(test)} test", flush=True)
Xtr = collect_states(model, tok, train)
ytr = torch.tensor([it["label"] == "hard" for it in train]).long()
Xte = collect_states(model, tok, test)
w, b, mu, sd = fit_probe(Xtr, ytr)
ptr = torch.sigmoid(((Xtr - mu) / sd) @ w + b)
acc_tr = ((ptr > 0.5).long() == ytr).float().mean()
pte = torch.sigmoid(((Xte - mu) / sd) @ w + b)
pred_hard = (pte > 0.5).tolist()
yte = [it["label"] == "hard" for it in test]
tp = sum(p and t for p, t in zip(pred_hard, yte))
print(f"probe: train_acc={acc_tr:.3f} test: pred_hard="
f"{sum(pred_hard)} (true hard={sum(yte)}, tp={tp})", flush=True)
# gated generation: k=0 for predicted-easy, K_HARD for predicted-hard
codes = [None] * len(test)
for k, sel in ((0, [i for i, ph in enumerate(pred_hard) if not ph]),
(K_HARD, [i for i, ph in enumerate(pred_hard) if ph])):
for i in range(0, len(sel), 8):
idxs = sel[i : i + 8]
chunk = [test[j] for j in idxs]
enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = looper.generate_frozen_prompt(
adapter, tok, enc["input_ids"], k, max_new_tokens=220,
attention_mask=enc["attention_mask"])
for jj, j in enumerate(idxs):
codes[j] = extract_code(
tok.decode(gen[jj, enc["input_ids"].shape[1]:],
skip_special_tokens=True))
with ThreadPoolExecutor(8) as ex:
oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]),
zip(codes, test)))
acc = sum(oks) / len(test)
by = {}
for it, ok in zip(test, oks):
d = by.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
print(f"GATED pass@1={acc:.3f} "
f"by_label={ {l: round(c/n,3) for l,(c,n) in by.items()} }", flush=True)
json.dump({"acc": acc,
"by_label": {l: c / n for l, (c, n) in by.items()},
"probe_train_acc": float(acc_tr),
"pred_hard_test": int(sum(pred_hard))},
open(OUT / "eval_gated.json", "w"), indent=1)
print("wrote", OUT / "eval_gated.json")
if __name__ == "__main__":
main()
+75
View File
@@ -0,0 +1,75 @@
"""J-lens sharpening battery on MBPP prompts: does the trained loop
concentrate the workspace readout while reading the problem?
For N test prompts, at each k: J-lens distribution at L30, last prompt
position -> entropy + top-1 probability. Trained vs untrained adapter.
"""
import json
import sys
from pathlib import Path
import torch
from loop_common import BandLooper, MergeAdapter
from prep_mbpp import DIRECT_SUFFIX, mbpp_prompt
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import JLens, load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
ROOT = OUT.parent
N_ITEMS = 20
KS = (0, 1, 2, 4)
@torch.no_grad()
def battery(looper, adapter, lens, tok, items):
rows = []
for it in items:
ids = tok(mbpp_prompt(tok, it, DIRECT_SUFFIX), return_tensors="pt",
add_special_tokens=False)["input_ids"].cuda()
calls, _ = looper.capture(ids, logits_to_keep=1)
e = looper._hin[looper.l0]
s = looper.band(e, calls)
state = {0: s}
for k in range(1, max(KS) + 1):
s = looper.band(adapter(e, s), calls)
state[k] = s
for k in KS:
probs = torch.softmax(lens._readout(
state[k][0, -1].float() @ lens.Jbar[looper.l1].T).float(), -1)
ent = -(probs * (probs + 1e-12).log()).sum().item()
rows.append({"task_id": it["task_id"], "k": k, "entropy": ent,
"top1": probs.max().item()})
return rows
def main():
model, tok = load_model(dtype=torch.bfloat16)
looper = BandLooper(model)
jbar = torch.load(ROOT / "results" / "jbar.pt", map_location="cuda")
lens = JLens(model, tok, jbar["Jbar"].float())
items = [it for it in json.load(open(OUT / "mbpp_data.json"))
if it["split"] == "test"][:N_ITEMS]
res = {}
for tag, path in (("untrained", None),
("trained", OUT / "adapter_code_e399.pt")):
adapter = MergeAdapter().cuda()
if path:
adapter.load_state_dict(torch.load(path, map_location="cuda"))
rows = battery(looper, adapter, lens, tok, items)
res[tag] = rows
for k in KS:
sel = [r for r in rows if r["k"] == k]
ent = sum(r["entropy"] for r in sel) / len(sel)
top = sum(r["top1"] for r in sel) / len(sel)
print(f"[{tag}] k={k}: mean_entropy={ent:.3f} mean_top1={top:.3f}",
flush=True)
json.dump(res, open(OUT / "lens_battery.json", "w"), indent=1)
print("wrote", OUT / "lens_battery.json")
if __name__ == "__main__":
main()
+310
View File
@@ -0,0 +1,310 @@
"""Shared machinery for workspace-band looping (WORKSPACE_LOOPING.md).
Codifies probe 2d: loop L14-30 through a merge layer at the L13->L14 boundary,
L14_in = (1-alpha)*e + alpha*(s renormalized to |e|) + MLP([e; s_hat])
with e = L13 output (fixed anchor) and s = looped-back band output.
The MLP is zero-initialized, so the untrained adapter reproduces the
hand-built alpha-merge exactly.
Band forward is done by re-calling the decoder layers with the per-layer
(args, kwargs) captured from a normal forward pass -- position embeddings and
attention masks do not depend on hidden states, so they are reusable across
loop iterations.
"""
import re
import sys
from pathlib import Path
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import _text_model, load_model # noqa: E402
import os as _os
BAND = tuple(int(x) for x in
_os.environ.get("JLENS_BAND", "14,30").split(","))
# default: workspace band on gemma-4-E2B (inclusive); 12B: JLENS_BAND=36,45
ALPHA = 0.3 # anchor-dominant merge weight from probe 2d
class MergeAdapter(nn.Module):
"""(1-a)*e + a*s_hat + MLP([e; s_hat]); MLP zero-init => starts at probe 2d."""
def __init__(self, d=1536, hidden=512, alpha=ALPHA):
super().__init__()
self.alpha = alpha
self.mlp = nn.Sequential(
nn.Linear(2 * d, hidden), nn.GELU(), nn.Linear(hidden, d)
)
nn.init.zeros_(self.mlp[2].weight)
nn.init.zeros_(self.mlp[2].bias)
def forward(self, e, s):
dt = e.dtype
e32, s32 = e.float(), s.float()
s_hat = s32 * (
e32.norm(dim=-1, keepdim=True) / (s32.norm(dim=-1, keepdim=True) + 1e-6)
)
base = (1 - self.alpha) * e32 + self.alpha * s_hat
out = base + self.mlp(torch.cat([e32, s_hat], dim=-1))
return out.to(dt)
class BandLooper:
"""Capture layer-call kwargs once per forward, then re-run L14-30 manually."""
def __init__(self, model, band=BAND):
self.model = model
self.tm = _text_model(model)
self.l0, self.l1 = band
self.n_layers = len(self.tm.layers)
def capture(self, input_ids, attention_mask=None, logits_to_keep=0,
position_ids=None):
"""Plain forward; returns (calls dict, logits). calls[i] = (args, kwargs)."""
calls, hin, handles = {}, {}, []
for i in range(self.l0, self.n_layers):
def pre(mod, args, kwargs, i=i):
# normalize: strip hidden_states, keep the rest for re-calls
drop = {"hidden_states", "past_key_value", "past_key_values",
"use_cache"}
kw = {k: v for k, v in kwargs.items() if k not in drop}
if "hidden_states" in kwargs:
hin[i] = kwargs["hidden_states"]
calls[i] = (args, kw)
else:
hin[i] = args[0]
calls[i] = (args[1:], kw)
handles.append(
self.tm.layers[i].register_forward_pre_hook(pre, with_kwargs=True)
)
try:
with torch.no_grad():
out = self.model(input_ids=input_ids, attention_mask=attention_mask,
use_cache=False, logits_to_keep=logits_to_keep,
position_ids=position_ids)
finally:
for h in handles:
h.remove()
self._hin = hin
return calls, out.logits
def _run(self, h, calls, lo, hi):
for i in range(lo, hi + 1):
args, kwargs = calls[i]
h = self.tm.layers[i](h, *args, **kwargs)
if isinstance(h, tuple):
h = h[0]
return h
def band(self, h, calls):
return self._run(h, calls, self.l0, self.l1)
def suffix_logits(self, h, calls, last_only=False):
h = self._run(h, calls, self.l1 + 1, self.n_layers - 1)
if last_only:
h = h[:, -1:]
h = self.tm.norm(h.to(self.tm.norm.weight.dtype))
head = self.model.lm_head if hasattr(self.model, "lm_head") else self.model.get_output_embeddings()
logits = head(h)
cap = getattr(self.model.config.get_text_config(), "final_logit_softcapping", None)
if cap:
logits = cap * torch.tanh(logits / cap)
return logits
def loop_logits(self, adapter, input_ids, k, attention_mask=None,
use_checkpoint=False, return_states=False, last_only=False,
loop_mask=None, feedforward=False):
"""Teacher-forced logits after k merge->band loops. k=0 = plain forward.
loop_mask (B, T) bool: positions where the merge applies; elsewhere the
band input stays the anchor e ("latent planning" over the prompt span —
unmasked positions still re-attend to the looped states each iteration).
"""
calls, base_logits = self.capture(input_ids, attention_mask,
logits_to_keep=1 if last_only else 0)
if k == 0:
return (base_logits, None) if return_states else base_logits
del base_logits # full-vocab logits — do not hold across the loop
e = self._hin[self.l0].detach()
if feedforward:
# no-recurrence control: adapter sees (e, e), applied exactly once
x = adapter(e, e)
if loop_mask is not None:
x = torch.where(loop_mask[..., None], x, e)
s = (checkpoint(lambda x_: self.band(x_, calls), x,
use_reentrant=False) if use_checkpoint
else self.band(x, calls))
logits = self.suffix_logits(s, calls, last_only=last_only)
return (logits, [s]) if return_states else logits
with torch.no_grad():
s = self.band(e, calls) # s_0: no trainable params upstream
states = [s]
for _ in range(k):
x = adapter(e, s)
if loop_mask is not None:
x = torch.where(loop_mask[..., None], x, e)
if use_checkpoint:
s = checkpoint(lambda x_: self.band(x_, calls), x, use_reentrant=False)
else:
s = self.band(x, calls)
states.append(s)
logits = self.suffix_logits(s, calls, last_only=last_only)
return (logits, states) if return_states else logits
@torch.no_grad()
def loop_generate(self, adapter, tok, input_ids, k, max_new_tokens=12,
attention_mask=None, loop_prompt_only=False,
stop_strs=()):
"""Batched greedy decode with the looped forward (no KV cache).
loop_prompt_only: merge applies only to the initial prompt span;
generated tokens go through the plain band (latent planning)."""
ids = input_ids
mask = attention_mask
lmask = None
if loop_prompt_only:
lmask = (mask if mask is not None
else torch.ones_like(ids)).bool().clone()
n_prompt = ids.shape[1]
texts = [""] * ids.shape[0]
eos = {tok.eos_token_id}
eot = tok.convert_tokens_to_ids("<end_of_turn>")
if eot is not None and eot >= 0:
eos.add(eot)
done = torch.zeros(ids.shape[0], dtype=torch.bool, device=ids.device)
for _ in range(max_new_tokens):
logits = self.loop_logits(adapter, ids, k, attention_mask=mask,
last_only=True, loop_mask=lmask)
nxt = logits[:, -1].argmax(-1)
nxt = torch.where(done, torch.full_like(nxt, list(eos)[0]), nxt)
ids = torch.cat([ids, nxt[:, None]], 1)
if mask is not None:
mask = torch.cat([mask, (~done)[:, None].long()], 1)
if lmask is not None:
lmask = torch.cat(
[lmask, torch.zeros_like(lmask[:, :1])], 1)
for b, t in enumerate(nxt.tolist()):
if not done[b] and stop_strs:
texts[b] += tok.decode([t])
done |= torch.tensor([t.item() in eos for t in nxt], device=ids.device)
if stop_strs:
done |= torch.tensor(
[any(ss in tx for ss in stop_strs) for tx in texts],
device=ids.device)
if done.all():
break
return ids
@torch.no_grad()
def generate_frozen_prompt(self, adapter, tok, input_ids, k,
max_new_tokens=220, attention_mask=None,
stop_strs=(), feedforward=False):
"""Fast equivalent of loop_generate(loop_prompt_only=True).
The looped prompt states are constant across token steps (causality),
so: loop the prompt once to the merged input x* = adapter(e, s*), then
run ONE cached prefill with a hook swapping the band input to x*, and
generate normally with the KV cache (native speed)."""
if k == 0:
x_star = None
elif feedforward:
calls, _ = self.capture(input_ids, attention_mask,
logits_to_keep=1)
x_star = adapter(self._hin[self.l0], self._hin[self.l0])
del calls
else:
calls, _ = self.capture(input_ids, attention_mask,
logits_to_keep=1)
e = self._hin[self.l0]
s = self.band(e, calls)
x_star = e
for _ in range(k):
x_star = adapter(e, s)
s = self.band(x_star, calls)
del calls # prefill re-runs band(x_star) -> same final s as slow path
hook = None
if x_star is not None:
def swap(mod, args, kwargs):
h = kwargs.get("hidden_states", args[0] if args else None)
if h is not None and h.shape[1] == x_star.shape[1]: # prefill
if "hidden_states" in kwargs:
kwargs["hidden_states"] = x_star.to(h.dtype)
return args, kwargs
return (x_star.to(h.dtype),) + args[1:], kwargs
return None
hook = self.tm.layers[self.l0].register_forward_pre_hook(
swap, with_kwargs=True)
try:
eos = [t for t in (tok.eos_token_id,
tok.convert_tokens_to_ids("<end_of_turn>"))
if t is not None and t >= 0]
out = self.model.generate(
input_ids=input_ids, attention_mask=attention_mask,
max_new_tokens=max_new_tokens, do_sample=False,
eos_token_id=eos, pad_token_id=tok.pad_token_id or 0,
stop_strings=list(stop_strs) or None, tokenizer=tok)
finally:
if hook:
hook.remove()
return out
# ---------- GSM8K helpers ----------
NUM_RE = re.compile(r"-?\$?[\d,]*\.?\d+")
def gold_answer(ans_field):
return ans_field.split("####")[-1].strip().replace(",", "").replace("$", "")
def last_number(text):
hits = NUM_RE.findall(text)
if not hits:
return None
x = hits[-1].replace(",", "").replace("$", "").rstrip(".")
return x
def num_eq(a, b):
try:
return a is not None and b is not None and abs(float(a) - float(b)) < 1e-4
except ValueError:
return False
DIRECT_SUFFIX = "\n\nGive only the final numeric answer, nothing else."
COT_SUFFIX = ("\n\nThink step by step, then give the final numeric answer "
"on the last line as: #### <number>")
def chat_prompt(tok, question, suffix=DIRECT_SUFFIX):
return tok.apply_chat_template(
[{"role": "user", "content": question + suffix}],
tokenize=False, add_generation_prompt=True,
)
def build_train_batch(tok, items, device="cuda"):
"""Right-padded (input_ids, attention_mask, labels); labels only on answer tokens."""
seqs, labs = [], []
for it in items:
p = tok(chat_prompt(tok, it["question"]), add_special_tokens=False)["input_ids"]
a = tok(it["gold"] + "<end_of_turn>", add_special_tokens=False)["input_ids"]
seqs.append(p + a)
labs.append([-100] * len(p) + a)
T = max(len(s) for s in seqs)
pad = tok.pad_token_id or 0
ids = torch.full((len(seqs), T), pad, dtype=torch.long)
lab = torch.full((len(seqs), T), -100, dtype=torch.long)
msk = torch.zeros((len(seqs), T), dtype=torch.long)
for i, (s, l) in enumerate(zip(seqs, labs)):
ids[i, : len(s)] = torch.tensor(s)
lab[i, : len(s)] = torch.tensor(l)
msk[i, : len(s)] = 1
return ids.to(device), msk.to(device), lab.to(device)
+51
View File
@@ -0,0 +1,51 @@
"""MBPP latent-planning results: pass@1 vs loop depth, trained vs untrained."""
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
OUT = Path(__file__).resolve().parent.parent / "results-loop"
tr = json.load(open(OUT / "eval_code_trained.json"))
un = json.load(open(OUT / "eval_code_untrained.json"))
ks = sorted(int(k) for k in tr["ks"])
BLUE, GRAY = "#2b6cb0", "#8a8f98"
def curve(res, sel):
return [res["ks"][str(k)]["acc"] if sel == "all"
else res["ks"][str(k)]["by_label"].get(sel, 0.0) for k in ks]
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4.2))
for ax in (ax1, ax2):
ax.grid(True, color="#e5e5e5", lw=0.7)
ax.set_axisbelow(True)
for s in ("top", "right"):
ax.spines[s].set_visible(False)
ax.set_xticks(ks)
ax.set_xlabel("loop depth k (prompt-only)")
ax1.plot(ks, curve(tr, "all"), "-o", color=BLUE, lw=2, ms=5, label="trained · all")
ax1.plot(ks, curve(un, "all"), "-o", color=GRAY, lw=2, ms=5, label="untrained · all")
ax1.axhline(curve(tr, "all")[0], color="#bbb", lw=1, ls=":")
ax1.set_ylabel("pass@1 (250 MBPP test items)")
ax1.set_title("Overall: trained k=2 beats no-loop baseline", fontsize=11)
ax1.legend(fontsize=8, frameon=False)
ax2.plot(ks, curve(tr, "hard"), "-s", color=BLUE, lw=2, ms=5,
label="trained · hard (plan-only)")
ax2.plot(ks, curve(un, "hard"), "-s", color=GRAY, lw=2, ms=5,
label="untrained · hard")
ax2.set_ylabel("pass@1, hard bucket (~56 items)")
ax2.set_title("Planning-dependent problems: 3.6% → 46.4%", fontsize=11)
ax2.legend(fontsize=8, frameon=False)
ax2.annotate("silent loops recover ~46%\nof explicit-planning gap",
xy=(4, curve(tr, "hard")[-1]), xytext=(1.8, 0.30), fontsize=8,
color=BLUE, arrowprops=dict(arrowstyle="-", color=BLUE, lw=0.8))
fig.suptitle("MBPP latent planning: loop the workspace over the prompt, "
"then write code normally", fontsize=12, y=1.02)
fig.tight_layout()
fig.savefig(OUT / "loop_eval_code.png", dpi=140, bbox_inches="tight",
facecolor="white")
print("wrote", OUT / "loop_eval_code.png")
+105
View File
@@ -0,0 +1,105 @@
"""Loop convergence dynamics: does the recurrence reach a fixed point?
For trained vs untrained merge, trace across iterations k:
- cos(s_k, s_{k-1}) (mean over positions) -> fixed point if -> 1
- |s_k| / |e| -> norm control
- P('spider') under the J-lens at L30 -> what the state converges TO
Averaged over the spider prompt + a few GSM8K test questions.
"""
import json
import sys
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import torch
from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
KMAX = 10
BLUE, GRAY = "#2b6cb0", "#8a8f98"
@torch.no_grad()
def trace(looper, adapter, tok, prompt_text):
ids = tok(prompt_text, return_tensors="pt",
add_special_tokens=False)["input_ids"].cuda()
calls, _ = looper.capture(ids)
e = looper._hin[looper.l0]
s = looper.band(e, calls)
rows = []
for k in range(1, KMAX + 1):
new = looper.band(adapter(e, s), calls)
cos = torch.nn.functional.cosine_similarity(
new[0].float(), s[0].float(), dim=-1).mean().item()
rows.append({"k": k, "cos": cos,
"norm": (new.norm() / e.norm()).item()})
s = new
return rows
def main():
model, tok = load_model(dtype=torch.bfloat16)
looper = BandLooper(model)
gsm = [it for it in json.load(open(OUT / "star_data.json"))
if it["split"] == "test"][:3]
prompts = [chat_prompt(tok, "The animal that spins webs has how many legs? "
"Answer with just the number.", "")]
prompts += [chat_prompt(tok, it["question"], DIRECT_SUFFIX) for it in gsm]
curves = {}
for tag, path in (("untrained", None), ("trained", OUT / "adapter.pt")):
adapter = MergeAdapter().cuda()
if path:
adapter.load_state_dict(torch.load(path, map_location="cuda"))
traces = [trace(looper, adapter, tok, p) for p in prompts]
curves[tag] = {
"cos": [sum(t[i]["cos"] for t in traces) / len(traces)
for i in range(KMAX)],
"norm": [sum(t[i]["norm"] for t in traces) / len(traces)
for i in range(KMAX)],
}
print(tag, "cos:", [round(c, 3) for c in curves[tag]["cos"]], flush=True)
print(tag, "norm:", [round(c, 3) for c in curves[tag]["norm"]], flush=True)
ks = list(range(1, KMAX + 1))
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4))
for ax in (ax1, ax2):
ax.grid(True, color="#e5e5e5", lw=0.7)
ax.set_axisbelow(True)
for sp in ("top", "right"):
ax.spines[sp].set_visible(False)
ax.set_xticks(ks)
ax.set_xlabel("loop iteration k")
ax.axvspan(2, 4, color="#f2e8cf", alpha=0.45, zorder=0)
for tag, c in (("trained", BLUE), ("untrained", GRAY)):
ax1.plot(ks, curves[tag]["cos"], "-o", color=c, lw=2, ms=5, label=tag)
ax2.plot(ks, curves[tag]["norm"], "-o", color=c, lw=2, ms=5, label=tag)
ax1.set_ylabel("cos(s_k, s_{k1}) (mean over positions)")
ax1.set_title("Successive-state similarity: fixed point?", fontsize=11)
ax1.legend(fontsize=8, frameon=False, loc="lower right")
ax1.text(3, ax1.get_ylim()[0] + 0.02 * (ax1.get_ylim()[1] - ax1.get_ylim()[0]),
"accuracy &\nsharpening plateau", fontsize=7.5, color="#7a5a00",
ha="center")
ax2.set_ylabel("|s_k| / |e|")
ax2.set_title("State norm across iterations", fontsize=11)
ax2.legend(fontsize=8, frameon=False)
fig.suptitle("Loop dynamics (spider + 3 GSM8K prompts, mean)", fontsize=12,
y=1.02)
fig.tight_layout()
fig.savefig(OUT / "loop_dynamics.png", dpi=140, bbox_inches="tight",
facecolor="white")
print("wrote", OUT / "loop_dynamics.png")
if __name__ == "__main__":
main()
+57
View File
@@ -0,0 +1,57 @@
"""Two-panel figure: accuracy vs loop depth k, and J-lens concept sharpening."""
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
OUT = Path(__file__).resolve().parent.parent / "results-loop"
tr = json.load(open(OUT / "eval_trained.json"))
un = json.load(open(OUT / "eval_untrained.json"))
ks = sorted(int(k) for k in tr["ks"])
BLUE, GRAY = "#2b6cb0", "#8a8f98"
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4.2))
def curve(res, sel):
return [res["ks"][str(k)]["acc"] if sel == "all"
else res["ks"][str(k)]["by_label"].get(sel, 0.0) for k in ks]
for ax in (ax1, ax2):
ax.grid(True, color="#e5e5e5", lw=0.7, zorder=0)
ax.set_axisbelow(True)
for s in ("top", "right"):
ax.spines[s].set_visible(False)
ax.set_xticks(ks)
ax.set_xlabel("loop depth k")
# --- panel 1: GSM8K accuracy vs k ---
ax1.plot(ks, curve(tr, "all"), "-o", color=BLUE, lw=2, ms=5, label="trained · all")
ax1.plot(ks, curve(un, "all"), "-o", color=GRAY, lw=2, ms=5, label="untrained · all")
ax1.plot(ks, curve(tr, "hard"), "--s", color=BLUE, lw=2, ms=5, label="trained · hard (CoT-only)")
ax1.plot(ks, curve(un, "hard"), "--s", color=GRAY, lw=2, ms=5, label="untrained · hard")
ax1.set_ylabel("accuracy (greedy, 256 held-out items)")
ax1.set_title("GSM8K accuracy vs loop depth", fontsize=11)
ax1.legend(fontsize=8, frameon=False)
ax1.annotate("hard items: 0.8% → 6.3%", xy=(2, 0.063), xytext=(3.2, 0.115),
fontsize=8, color=BLUE,
arrowprops=dict(arrowstyle="-", color=BLUE, lw=0.8))
# --- panel 2: J-lens sharpening ---
sp_tr = {r["k"]: r["P_spider_lens"] for r in tr["spider"]}
sp_un = {r["k"]: r["P_spider_lens"] for r in un["spider"]}
ax2.plot(ks, [sp_tr[k] for k in ks], "-o", color=BLUE, lw=2, ms=5, label="trained")
ax2.plot(ks, [sp_un[k] for k in ks], "-o", color=GRAY, lw=2, ms=5, label="untrained")
ax2.set_ylabel("P('spider') under J-lens at L30")
ax2.set_title("Latent concept sharpening across loops", fontsize=11)
ax2.legend(fontsize=8, frameon=False)
ax2.text(ks[-1], sp_tr[ks[-1]] + 0.006, "8× over control", fontsize=8,
color=BLUE, ha="right")
fig.suptitle("Trained merge adapter: loop depth buys latent sharpening and some "
"CoT-only answers, not overall accuracy", fontsize=12, y=1.02)
fig.tight_layout()
fig.savefig(OUT / "loop_eval.png", dpi=140, bbox_inches="tight", facecolor="white")
print("wrote", OUT / "loop_eval.png")
+59
View File
@@ -0,0 +1,59 @@
"""Render heatmaps (concept P by layer x position) and the layer profile."""
import os, sys
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import torch
RES = Path(__file__).resolve().parent.parent / os.environ.get("JLENS_RESULTS", "results")
def plot_heat(path):
d = torch.load(path)
m, words, toks = d["map"].detach(), d["words"], d["tokens"]
fig, ax = plt.subplots(figsize=(min(16, 0.28 * m.shape[1] + 2), 6))
im = ax.imshow(m.numpy(), aspect="auto", origin="lower", cmap="magma",
vmin=0)
ax.set_ylabel("layer")
ax.set_xlabel("position")
if toks and len(toks) <= 60:
ax.set_xticks(range(len(toks)))
ax.set_xticklabels([t.replace("\n", "\\n") for t in toks],
rotation=90, fontsize=6)
ax.set_title(f"J-lens P({'/'.join(words)})")
fig.colorbar(im)
out = path.with_suffix(".png")
fig.tight_layout()
fig.savefig(out, dpi=130)
plt.close(fig)
print("wrote", out)
def plot_profile(path):
d = torch.load(path)
fig, ax1 = plt.subplots(figsize=(8, 4.5))
L = len(d["entropy"])
ax1.plot(range(L), d["entropy"], "k-", label="lens entropy")
ax1.set_xlabel("layer")
ax1.set_ylabel("entropy (nats)")
ax2 = ax1.twinx()
ax2.plot(range(L), d["agree_in"], "b--", label="top tok == current input")
ax2.plot(range(L), d["agree_next"], "r--", label="top tok == next output")
ax2.set_ylabel("agreement")
fig.legend(loc="upper center", ncol=3, fontsize=8)
fig.tight_layout()
out = path.with_suffix(".png")
fig.savefig(out, dpi=130)
plt.close(fig)
print("wrote", out)
if __name__ == "__main__":
for p in sorted(RES.glob("heat_*.pt")):
plot_heat(p)
lp = RES / "layer_profile.pt"
if lp.exists():
plot_profile(lp)
+132
View File
@@ -0,0 +1,132 @@
"""STaR-style difficulty labeling of MBPP with the frozen base model.
Two passes per item: direct code generation vs plan-first-then-code, each
executed against MBPP's unit tests (sandboxed subprocess). Labels:
easy (direct passes), hard (plan-only passes), drop (neither). The model's
own passing code is kept as the training target (in-distribution supervision).
"""
import json
import os
import re
import subprocess
import sys
import tempfile
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from datasets import load_dataset
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
OUT.mkdir(exist_ok=True)
DIRECT_SUFFIX = ("\n\nWrite only the Python function in a ```python code "
"block. No explanation.")
PLAN_SUFFIX = ("\n\nFirst write a very brief plan: at most 4 short bullet "
"lines, no headings, no math notation. Then write the complete "
"Python function in a ```python code block.")
CODE_RE = re.compile(r"```(?:python)?\s*\n(.*?)```", re.S)
def mbpp_prompt(tok, item, suffix):
tests = "\n".join(item["test_list"])
msg = (f"{item['text']}\nYour code should pass these tests:\n\n{tests}"
f"{suffix}")
return tok.apply_chat_template([{"role": "user", "content": msg}],
tokenize=False, add_generation_prompt=True)
def extract_code(text):
m = CODE_RE.findall(text)
return m[-1].strip() if m else None
def run_tests(code, item, timeout=10):
if not code:
return False
script = (item.get("test_setup_code") or "") + "\n" + code + "\n" + \
"\n".join(item["test_list"])
try:
with tempfile.TemporaryDirectory() as td:
r = subprocess.run([sys.executable, "-c", script], cwd=td,
capture_output=True, timeout=timeout)
return r.returncode == 0
except (subprocess.TimeoutExpired, OSError):
return False
@torch.no_grad()
def batch_generate(model, tok, prompts, max_new_tokens, batch_size=24):
outs = []
for i in range(0, len(prompts), batch_size):
chunk = prompts[i : i + batch_size]
enc = tok(chunk, return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = model.generate(**enc, max_new_tokens=max_new_tokens,
do_sample=False,
pad_token_id=tok.pad_token_id or 0)
for j in range(len(chunk)):
outs.append(tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True))
print(f" {min(i+batch_size, len(prompts))}/{len(prompts)}", flush=True)
return outs
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
ds = load_dataset("google-research-datasets/mbpp", "full")
items = []
for split, tag in (("train", "train"), ("validation", "train"),
("test", "test")):
for row in ds[split]:
items.append({"split": tag, "task_id": row["task_id"],
"text": row["text"], "test_list": row["test_list"],
"test_setup_code": row["test_setup_code"]})
print(f"{len(items)} items", flush=True)
for tag, suffix, mx in (("direct", DIRECT_SUFFIX, 220),
("plan", PLAN_SUFFIX, 700)):
t0 = time.time()
print(f"{tag} pass", flush=True)
gens = batch_generate(model, tok,
[mbpp_prompt(tok, it, suffix) for it in items], mx)
codes = [extract_code(g) for g in gens]
with ThreadPoolExecutor(8) as ex:
oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]),
zip(codes, items)))
for it, c, ok in zip(items, codes, oks):
it[f"{tag}_ok"] = bool(ok)
it[f"{tag}_code"] = c if ok else None
print(f"{tag} pass done in {time.time()-t0:.0f}s "
f"pass@1={sum(oks)/len(items):.3f}", flush=True)
for it in items:
it["label"] = ("easy" if it["direct_ok"]
else "hard" if it["plan_ok"] else "drop")
it["sol_code"] = it["direct_code"] if it["direct_ok"] else it["plan_code"]
for split in ("train", "test"):
sub = [it for it in items if it["split"] == split]
n = len(sub)
print(f"{split}: n={n} direct={sum(i['direct_ok'] for i in sub)/n:.3f} "
f"plan={sum(i['plan_ok'] for i in sub)/n:.3f} "
f"easy={sum(i['label']=='easy' for i in sub)} "
f"hard={sum(i['label']=='hard' for i in sub)} "
f"drop={sum(i['label']=='drop' for i in sub)}", flush=True)
with open(OUT / "mbpp_data.json", "w") as f:
json.dump(items, f, indent=1)
print("wrote", OUT / "mbpp_data.json")
if __name__ == "__main__":
main()
+65
View File
@@ -0,0 +1,65 @@
"""Re-run the MBPP plan pass on direct-fail items with a non-truncating budget.
The first plan pass (max_new=380) truncated ~all outputs before the code:
E2B writes verbose plans. Fix: terse-plan prompt + 700-token budget. Only
items with label=='drop' need re-labeling (easy is decided by the direct pass).
"""
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from prep_mbpp import batch_generate, extract_code, mbpp_prompt, run_tests
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
TERSE_PLAN_SUFFIX = ("\n\nFirst write a very brief plan: at most 4 short "
"bullet lines, no headings, no math notation. Then write "
"the complete Python function in a ```python code block.")
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
data = json.load(open(OUT / "mbpp_data.json"))
redo = [it for it in data if it["label"] == "drop"]
print(f"re-running plan pass on {len(redo)} direct-fail items", flush=True)
t0 = time.time()
gens = batch_generate(model, tok,
[mbpp_prompt(tok, it, TERSE_PLAN_SUFFIX)
for it in redo], max_new_tokens=700, batch_size=16)
codes = [extract_code(g) for g in gens]
n_trunc = sum(1 for g in gens if len(tok(g)["input_ids"]) >= 695)
with ThreadPoolExecutor(8) as ex:
oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]),
zip(codes, redo)))
for it, c, ok in zip(redo, codes, oks):
it["plan_ok"] = bool(ok)
it["plan_code"] = c if ok else None
it["label"] = "hard" if ok else "drop"
it["sol_code"] = it["direct_code"] if it["direct_ok"] else it["plan_code"]
print(f"done in {time.time()-t0:.0f}s plan-pass on fails: "
f"{sum(oks)}/{len(redo)} still-truncated={n_trunc}", flush=True)
for split in ("train", "test"):
sub = [it for it in data if it["split"] == split]
n = len(sub)
print(f"{split}: n={n} easy={sum(i['label']=='easy' for i in sub)} "
f"hard={sum(i['label']=='hard' for i in sub)} "
f"drop={sum(i['label']=='drop' for i in sub)}", flush=True)
with open(OUT / "mbpp_data.json", "w") as f:
json.dump(data, f, indent=1)
print("wrote", OUT / "mbpp_data.json")
if __name__ == "__main__":
main()
+103
View File
@@ -0,0 +1,103 @@
"""STaR-style difficulty labeling of GSM8K with the frozen base model.
Two passes per item with frozen gemma-4-E2B-it:
direct: answer with no CoT -> correct = "easy"
cot: step-by-step -> correct (and direct wrong) = "hard"
Items the model cannot solve even with CoT are dropped from training
(unreachable supervision) but kept in the test split for eval.
Output: results-loop/star_data.json
"""
import json
import os
import random
import time
from pathlib import Path
import torch
from datasets import load_dataset
from loop_common import (COT_SUFFIX, DIRECT_SUFFIX, chat_prompt, gold_answer,
last_number, num_eq)
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
N_TRAIN, N_TEST = 1024, 256
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
OUT.mkdir(exist_ok=True)
@torch.no_grad()
def batch_generate(model, tok, prompts, max_new_tokens, batch_size):
outs = []
for i in range(0, len(prompts), batch_size):
chunk = prompts[i : i + batch_size]
enc = tok(chunk, return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = model.generate(**enc, max_new_tokens=max_new_tokens,
do_sample=False,
pad_token_id=tok.pad_token_id or 0)
for j in range(len(chunk)):
outs.append(tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True))
print(f" {min(i+batch_size, len(prompts))}/{len(prompts)}", flush=True)
return outs
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
ds = load_dataset("openai/gsm8k", "main")
rng = random.Random(0)
tr_idx = rng.sample(range(len(ds["train"])), N_TRAIN)
te_idx = rng.sample(range(len(ds["test"])), N_TEST)
items = []
for split, idxs in (("train", tr_idx), ("test", te_idx)):
for i in idxs:
row = ds[split][i]
items.append({"split": split, "idx": i, "question": row["question"],
"gold": gold_answer(row["answer"])})
t0 = time.time()
print(f"direct pass ({len(items)} items)", flush=True)
direct = batch_generate(
model, tok, [chat_prompt(tok, it["question"], DIRECT_SUFFIX) for it in items],
max_new_tokens=10, batch_size=64)
print(f"direct pass done in {time.time()-t0:.0f}s", flush=True)
t0 = time.time()
print("cot pass", flush=True)
cot = batch_generate(
model, tok, [chat_prompt(tok, it["question"], COT_SUFFIX) for it in items],
max_new_tokens=320, batch_size=32)
print(f"cot pass done in {time.time()-t0:.0f}s", flush=True)
for it, d, c in zip(items, direct, cot):
it["direct_pred"] = last_number(d)
it["direct_ok"] = num_eq(it["direct_pred"], it["gold"])
it["cot_pred"] = last_number(c)
it["cot_ok"] = num_eq(it["cot_pred"], it["gold"])
it["label"] = ("easy" if it["direct_ok"]
else "hard" if it["cot_ok"] else "drop")
for split in ("train", "test"):
sub = [it for it in items if it["split"] == split]
n = len(sub)
print(f"{split}: n={n} direct_acc={sum(i['direct_ok'] for i in sub)/n:.3f} "
f"cot_acc={sum(i['cot_ok'] for i in sub)/n:.3f} "
f"easy={sum(i['label']=='easy' for i in sub)} "
f"hard={sum(i['label']=='hard' for i in sub)} "
f"drop={sum(i['label']=='drop' for i in sub)}", flush=True)
with open(OUT / "star_data.json", "w") as f:
json.dump(items, f, indent=1)
print("wrote", OUT / "star_data.json")
if __name__ == "__main__":
main()
+116
View File
@@ -0,0 +1,116 @@
"""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("<end_of_turn>")}
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()
+104
View File
@@ -0,0 +1,104 @@
"""Sampled relabeling of test-split difficulty buckets (protocol item:
outcome-selection fix).
Greedy labeling couples bucket selection to the same coin flips as the k=0
eval (circular: a bucket defined by baseline failure shows baseline ~0%).
Fix: labels from 3 temperature-sampled direct attempts, independent of the
greedy eval — hard_s = 0/3 sampled direct correct AND (plan/CoT reachable);
easy_s = >=2/3 correct; else mid_s. Adds 'label_sampled' to the data JSONs.
"""
import json
import sys
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from loop_common import DIRECT_SUFFIX as GSM_SUFFIX
from loop_common import chat_prompt, last_number, num_eq
from prep_mbpp import (DIRECT_SUFFIX as MBPP_SUFFIX, extract_code,
mbpp_prompt, run_tests)
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
N_SAMPLES = 3
TEMP = 0.8
@torch.no_grad()
def sample_batch(model, tok, prompts, max_new, batch=24, seed=0):
outs = [[] for _ in prompts]
for s in range(N_SAMPLES):
torch.manual_seed(1000 + s + seed)
for i in range(0, len(prompts), batch):
chunk = prompts[i : i + batch]
enc = tok(chunk, return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
gen = model.generate(**enc, max_new_tokens=max_new, do_sample=True,
temperature=TEMP, top_p=0.95,
pad_token_id=tok.pad_token_id or 0)
for j in range(len(chunk)):
outs[i + j].append(
tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True))
print(f" sample {s+1}/{N_SAMPLES} done", flush=True)
return outs
def relabel(items, n_correct, reachable_key):
for it, nc in zip(items, n_correct):
if nc == 0:
it["label_sampled"] = ("hard" if it[reachable_key] else "drop")
elif nc >= 2:
it["label_sampled"] = "easy"
else:
it["label_sampled"] = "mid"
def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
# --- MBPP test ---
mbpp = json.load(open(OUT / "mbpp_data.json"))
mtest = [it for it in mbpp if it["split"] == "test"]
print(f"MBPP: sampling {len(mtest)} items x{N_SAMPLES}", flush=True)
gens = sample_batch(model, tok,
[mbpp_prompt(tok, it, MBPP_SUFFIX) for it in mtest],
max_new=220)
nc = []
with ThreadPoolExecutor(8) as ex:
for it, gs in zip(mtest, gens):
oks = list(ex.map(lambda g: run_tests(extract_code(g), it), gs))
nc.append(sum(oks))
relabel(mtest, nc, "plan_ok")
json.dump(mbpp, open(OUT / "mbpp_data.json", "w"), indent=1)
dist = {l: sum(it.get("label_sampled") == l for it in mtest)
for l in ("easy", "mid", "hard", "drop")}
agree = sum(it["label"] == it.get("label_sampled") for it in mtest
if it.get("label_sampled") != "mid")
print(f"MBPP sampled labels: {dist} (greedy-agreement excl. mid: "
f"{agree}/{sum(1 for it in mtest if it.get('label_sampled') != 'mid')})",
flush=True)
# --- GSM8K test ---
gsm = json.load(open(OUT / "star_data.json"))
gtest = [it for it in gsm if it["split"] == "test"]
print(f"GSM: sampling {len(gtest)} items x{N_SAMPLES}", flush=True)
gens = sample_batch(model, tok,
[chat_prompt(tok, it["question"], GSM_SUFFIX)
for it in gtest], max_new=10)
nc = [sum(num_eq(last_number(g), it["gold"]) for g in gs)
for it, gs in zip(gtest, gens)]
relabel(gtest, nc, "cot_ok")
json.dump(gsm, open(OUT / "star_data.json", "w"), indent=1)
dist = {l: sum(it.get("label_sampled") == l for it in gtest)
for l in ("easy", "mid", "hard", "drop")}
print(f"GSM sampled labels: {dist}", flush=True)
if __name__ == "__main__":
main()
+149
View File
@@ -0,0 +1,149 @@
"""Stage B / design C: train the merge adapter for prompt-prefill + carry.
GSM8K only (the task the prompt-only loop failed on). Sequence:
[prompt] [p x <unused0> pauses] [gold answer]; k=2 prefill loops on the
prompt; carry scan through pauses + answer; CE on answer tokens.
Curriculum: easy p=2, hard p=6 (hard needs the longer latent chain).
"""
import argparse
import json
import math
import random
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from carry_common import PAUSE_ID, carry_logits
from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
STEPS = 600
BATCH = 4
LR = 1e-3
WARMUP = 20
K_PREFILL = 2
P_BY_LABEL = {"easy": 2, "hard": 6}
ap = argparse.ArgumentParser()
ap.add_argument("--feedforward", action="store_true",
help="pause-token control: adapter(e,e), no carry")
ap.add_argument("--seed", type=int, default=0)
ARGS = ap.parse_args()
SEED = ARGS.seed
TAG = "pausectl" if ARGS.feedforward else "carry"
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
def build_batch(tok, items, p, device="cuda"):
seqs, labs, plens = [], [], []
for it in items:
pr = tok(chat_prompt(tok, it["question"], DIRECT_SUFFIX),
add_special_tokens=False)["input_ids"]
a = tok(it["gold"] + "<end_of_turn>",
add_special_tokens=False)["input_ids"]
seqs.append(pr + [PAUSE_ID] * p + a)
labs.append([-100] * (len(pr) + p) + a)
plens.append(len(pr))
T = max(len(s) for s in seqs)
pad = tok.pad_token_id or 0
ids = torch.full((len(seqs), T), pad, dtype=torch.long)
lab = torch.full((len(seqs), T), -100, dtype=torch.long)
msk = torch.zeros((len(seqs), T), dtype=torch.long)
for i, (s, l) in enumerate(zip(seqs, labs)):
ids[i, : len(s)] = torch.tensor(s)
lab[i, : len(s)] = torch.tensor(l)
msk[i, : len(s)] = 1
return (ids.to(device), msk.to(device), lab.to(device),
torch.tensor(plens, device=device))
@torch.no_grad()
def val_loss(looper, adapter, tok, items, p, k=K_PREFILL):
tot, n = 0.0, 0
for i in range(0, len(items), BATCH):
ids, msk, lab, plens = build_batch(tok, items[i : i + BATCH], p)
logits = carry_logits(looper, adapter, ids, msk, plens, k, feedforward=ARGS.feedforward)
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
tot += loss.item() * len(ids)
n += len(ids)
return tot / n
def main():
rng = random.Random(SEED)
torch.manual_seed(SEED)
data = [it for it in json.load(open(OUT / "star_data.json"))
if it["split"] == "train" and it["label"] != "drop"]
model, tok = load_model(dtype=torch.bfloat16)
for pp in model.parameters():
pp.requires_grad_(False)
looper = BandLooper(model)
adapter = MergeAdapter().cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
keep = [it for it in data
if len(tok(it["question"])["input_ids"]) + 30 <= 400]
rng.shuffle(keep)
pool = {l: [it for it in keep if it["label"] == l]
for l in ("easy", "hard")}
val = {l: pool[l][:16] for l in pool}
pool = {l: pool[l][16:] for l in pool}
print(f"pool: easy={len(pool['easy'])} hard={len(pool['hard'])}",
flush=True)
log = []
t0 = time.time()
for step in range(STEPS):
lbl = ("easy", "hard")[step % 2]
batch = rng.sample(pool[lbl], BATCH)
p = P_BY_LABEL[lbl]
ids, msk, lab, plens = build_batch(tok, batch, p)
for g in opt.param_groups:
g["lr"] = lr_at(step)
logits = carry_logits(looper, adapter, ids, msk, plens, K_PREFILL, feedforward=ARGS.feedforward,
use_checkpoint=True)
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
log.append({"step": step, "label": lbl, "p": p, "loss": loss.item()})
if step % 10 == 0:
print(f"step {step:4d} {lbl} p={p} loss={loss.item():.4f} "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
vals = {}
for l in ("easy", "hard"):
for pp_ in (0, 2, 6):
vals[f"{l}_p{pp_}"] = val_loss(looper, adapter, tok,
val[l], pp_)
print(f" val@{step}: " + " ".join(
f"{n}={v:.3f}" for n, v in sorted(vals.items())), flush=True)
log.append({"step": step, "val": vals})
torch.save(adapter.state_dict(),
OUT / f"adapter_{TAG}_e{step+1}.pt")
json.dump(log, open(OUT / f"train_{TAG}_log.json", "w"), indent=1)
print("done", flush=True)
if __name__ == "__main__":
main()
+175
View File
@@ -0,0 +1,175 @@
"""Plan-distillation baseline: same 1.6M budget, no recurrence.
Teacher: frozen model WITH its own terse plan in context (the plan text the
STaR pipeline validated). Student: feedforward adapter (adapter(e,e) at
prompt positions, applied once), NO plan in context. Loss: KL(teacher →
student) on code tokens + 0.5 CE on gold code. Bounds how much of the loop's
gain is available by compressing plan information into the adapter weights.
"""
import argparse
import json
import math
import os
import random
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from loop_common import BandLooper, MergeAdapter
from prep_mbpp import (DIRECT_SUFFIX, CODE_RE, batch_generate, mbpp_prompt)
from prep_mbpp_fix import TERSE_PLAN_SUFFIX
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
STEPS = 600
BATCH = 4
LR = 1e-3
WARMUP = 20
KL_T = 1.0
SEED = 0
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
def plan_prefix(gen_text):
"""Text before the final fenced code block (the plan)."""
m = list(CODE_RE.finditer(gen_text))
return gen_text[: m[-1].start()].strip() if m else None
def ensure_plans(model, tok, items):
path = OUT / "mbpp_plans.json"
if path.exists():
plans = json.load(open(path))
else:
print(f"generating plans for {len(items)} items", flush=True)
gens = batch_generate(model, tok,
[mbpp_prompt(tok, it, TERSE_PLAN_SUFFIX)
for it in items], max_new_tokens=700,
batch_size=16)
plans = {str(it["task_id"]): plan_prefix(g)
for it, g in zip(items, gens)}
json.dump(plans, open(path, "w"), indent=1)
return plans
def build_pair(tok, it, plan, device="cuda"):
"""(student ids/mask/lmask, teacher ids, code positions in each)."""
code = "```python\n" + it["sol_code"] + "\n```<end_of_turn>"
a = tok(code, add_special_tokens=False)["input_ids"]
sp = tok(mbpp_prompt(tok, it, DIRECT_SUFFIX),
add_special_tokens=False)["input_ids"]
tp = tok(mbpp_prompt(tok, it, TERSE_PLAN_SUFFIX) + plan + "\n",
add_special_tokens=False)["input_ids"]
return sp, tp, a
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
rng = random.Random(args.seed)
torch.manual_seed(args.seed)
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
for p in model.parameters():
p.requires_grad_(False)
looper = BandLooper(model)
d = model.config.get_text_config().hidden_size
adapter = MergeAdapter(d=d).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
data = json.load(open(OUT / "mbpp_data.json"))
train = [it for it in data if it["split"] == "train"
and it["label"] != "drop" and it.get("sol_code")]
plans = ensure_plans(model, tok, train)
train = [it for it in train if plans.get(str(it["task_id"]))]
tok.padding_side = "right"
train = [it for it in train
if len(tok(mbpp_prompt(tok, it, TERSE_PLAN_SUFFIX))["input_ids"])
+ len(tok(plans[str(it["task_id"])])["input_ids"])
+ len(tok(it["sol_code"])["input_ids"]) + 30 <= 900]
print(f"distill pool: {len(train)}", flush=True)
pad = tok.pad_token_id or 0
log = []
t0 = time.time()
for step in range(STEPS):
batch = rng.sample(train, BATCH)
pairs = [build_pair(tok, it, plans[str(it["task_id"])])
for it in batch]
# student tensors
Ts = max(len(sp) + len(a) for sp, _, a in pairs)
Tt = max(len(tp) + len(a) for _, tp, a in pairs)
s_ids = torch.full((BATCH, Ts), pad, dtype=torch.long)
s_msk = torch.zeros((BATCH, Ts), dtype=torch.long)
t_ids = torch.full((BATCH, Tt), pad, dtype=torch.long)
t_msk = torch.zeros((BATCH, Tt), dtype=torch.long)
lab = torch.full((BATCH, Ts), -100, dtype=torch.long)
spans = []
for i, (sp, tp, a) in enumerate(pairs):
s_ids[i, : len(sp) + len(a)] = torch.tensor(sp + a)
s_msk[i, : len(sp) + len(a)] = 1
lab[i, len(sp): len(sp) + len(a)] = torch.tensor(a)
t_ids[i, : len(tp) + len(a)] = torch.tensor(tp + a)
t_msk[i, : len(tp) + len(a)] = 1
spans.append((len(sp), len(tp), len(a)))
s_ids, s_msk, lab = s_ids.cuda(), s_msk.cuda(), lab.cuda()
t_ids, t_msk = t_ids.cuda(), t_msk.cuda()
lmask = (lab == -100) & (s_msk == 1)
with torch.no_grad():
t_logits = model(input_ids=t_ids, attention_mask=t_msk,
use_cache=False).logits
for g in opt.param_groups:
g["lr"] = lr_at(step)
s_logits = looper.loop_logits(adapter, s_ids, 1, attention_mask=s_msk,
loop_mask=lmask, feedforward=True,
use_checkpoint=True)
kl = torch.zeros((), device="cuda")
n_tok = 0
for i, (ls, lt, la) in enumerate(spans):
# predicting code token j uses position (prefix_len + j - 1)
sl = s_logits[i, ls - 1: ls + la - 1].float()
tl = t_logits[i, lt - 1: lt + la - 1].float()
kl = kl + F.kl_div(
F.log_softmax(sl / KL_T, -1), F.log_softmax(tl / KL_T, -1),
log_target=True, reduction="sum")
n_tok += la
kl = kl / n_tok
ce = F.cross_entropy(s_logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
loss = kl + 0.5 * ce
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
log.append({"step": step, "kl": kl.item(), "ce": ce.item()})
if step % 10 == 0:
print(f"step {step:4d} kl={kl.item():.4f} ce={ce.item():.4f} "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
torch.save(adapter.state_dict(),
OUT / f"adapter_distill_e{step+1}.pt")
json.dump(log, open(OUT / "train_distill_log.json", "w"),
indent=1)
print("done", flush=True)
if __name__ == "__main__":
main()
+126
View File
@@ -0,0 +1,126 @@
"""Train ONLY the merge adapter at the L13->L14 boundary (WORKSPACE_LOOPING.md §3).
Frozen base model; band L14-30 unrolled k times through the adapter; loss =
cross-entropy on answer tokens of the *direct* (no-CoT) prompt, supervised by
gold answers on items the frozen model can reach (STaR filter from
prep_star_data.py). Difficulty->depth curriculum:
k=1: easy items k=2: easy+hard k=4: hard items
so hard items are only ever seen at depth, forcing the loop to be used.
"""
import json
import random
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from loop_common import BandLooper, MergeAdapter, build_train_batch
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
STEPS = 800
BATCH = 8
LR = 1e-3
WARMUP = 20
MAX_TOK = 400 # skip overlong items
K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))]
SEED = 0
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
import math
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
@torch.no_grad()
def val_answer_acc(looper, adapter, tok, items, k):
"""Teacher-forced exact match: argmax at every answer position."""
ok = 0
for i in range(0, len(items), BATCH):
ids, msk, lab = build_train_batch(tok, items[i : i + BATCH])
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk)
pred = logits[:, :-1].argmax(-1)
tgt = lab[:, 1:]
m = tgt != -100
ok += (((pred == tgt) | ~m).all(-1)).sum().item()
return ok / len(items)
def main():
rng = random.Random(SEED)
torch.manual_seed(SEED)
data = json.load(open(OUT / "star_data.json"))
model, tok = load_model(dtype=torch.bfloat16)
for p in model.parameters():
p.requires_grad_(False)
looper = BandLooper(model)
adapter = MergeAdapter().cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
train = [it for it in data if it["split"] == "train" and it["label"] != "drop"]
# length filter
keep = []
for it in train:
n = len(tok(it["question"])["input_ids"])
if n + 40 <= MAX_TOK:
keep.append(it)
train = keep
pool = {"easy": [it for it in train if it["label"] == "easy"],
"hard": [it for it in train if it["label"] == "hard"]}
val = {lbl: rng.sample(lst, min(32, len(lst))) for lbl, lst in pool.items()}
for lbl in pool:
pool[lbl] = [it for it in pool[lbl] if it not in val[lbl]]
print(f"train pool: easy={len(pool['easy'])} hard={len(pool['hard'])}", flush=True)
log = []
t0 = time.time()
for step in range(STEPS):
k, labels = K_BUCKETS[step % len(K_BUCKETS)]
cand = [it for lbl in labels for it in pool[lbl]]
batch = rng.sample(cand, BATCH)
ids, msk, lab = build_train_batch(tok, batch)
for g in opt.param_groups:
g["lr"] = lr_at(step)
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
use_checkpoint=True)
loss = F.cross_entropy(
logits[:, :-1].flatten(0, 1).float(), lab[:, 1:].flatten(),
ignore_index=-100)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
log.append({"step": step, "k": k, "loss": loss.item()})
if step % 10 == 0:
print(f"step {step:4d} k={k} loss={loss.item():.4f} "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
accs = {}
for kk in (0, 1, 2, 4):
accs[f"easy_k{kk}"] = val_answer_acc(looper, adapter, tok,
val["easy"], kk)
accs[f"hard_k{kk}"] = val_answer_acc(looper, adapter, tok,
val["hard"], kk)
print(f" val@{step}: " +
" ".join(f"{n}={v:.2f}" for n, v in accs.items()), flush=True)
log.append({"step": step, "val": accs})
torch.save(adapter.state_dict(), OUT / "adapter.pt")
json.dump(log, open(OUT / "train_log.json", "w"), indent=1)
torch.save(adapter.state_dict(), OUT / "adapter.pt")
json.dump(log, open(OUT / "train_log.json", "w"), indent=1)
print("done; adapter ->", OUT / "adapter.pt")
if __name__ == "__main__":
main()
+129
View File
@@ -0,0 +1,129 @@
"""Train the merge adapter on Blocksworld (prompt-only latent planning).
Same recipe as train_merge_code.py; supervision = the model's own verified
passing plan (direct for easy, CoT-derived for hard — the plan section only,
CE on plan tokens)."""
import argparse
import json
import math
import os
import random
import re
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from bw_prep import DIRECT_SUFFIX, chat
from loop_common import BandLooper, MergeAdapter
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
STEPS = 600
BATCH = 4
LR = 1e-3
WARMUP = 20
K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))]
ap = argparse.ArgumentParser()
ap.add_argument("--seed", type=int, default=0)
ARGS = ap.parse_args()
MOVELIST_RE = re.compile(r"(?:^|\n)\s*1\s*[.)]", re.M)
def plan_only(text):
"""From a CoT output, keep from the numbered list onward."""
m = MOVELIST_RE.search(text)
return text[m.start():].strip() if m else text.strip()
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
def build_batch(tok, items, device="cuda"):
seqs, labs = [], []
for it in items:
p = tok(chat(tok, it, DIRECT_SUFFIX),
add_special_tokens=False)["input_ids"]
a = tok(plan_only(it["sol_plan"]) + "<end_of_turn>",
add_special_tokens=False)["input_ids"]
seqs.append(p + a)
labs.append([-100] * len(p) + a)
T = max(len(s) for s in seqs)
pad = tok.pad_token_id or 0
ids = torch.full((len(seqs), T), pad, dtype=torch.long)
lab = torch.full((len(seqs), T), -100, dtype=torch.long)
msk = torch.zeros((len(seqs), T), dtype=torch.long)
for i, (s, l) in enumerate(zip(seqs, labs)):
ids[i, : len(s)] = torch.tensor(s)
lab[i, : len(s)] = torch.tensor(l)
msk[i, : len(s)] = 1
lmask = (lab == -100) & (msk == 1)
return ids.to(device), msk.to(device), lab.to(device), lmask.to(device)
def main():
rng = random.Random(ARGS.seed)
torch.manual_seed(ARGS.seed)
data = json.load(open(OUT / "bw_data.json"))
model, tok = load_model(dtype=torch.bfloat16)
for p in model.parameters():
p.requires_grad_(False)
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
train = [it for it in data if it["split"] == "train"
and it["label"] != "drop" and it.get("sol_plan")]
pool = {l: [it for it in train if it["label"] == l]
for l in ("easy", "hard")}
val = {l: pool[l][:12] for l in pool}
pool = {l: pool[l][12:] for l in pool}
print(f"pool: easy={len(pool['easy'])} hard={len(pool['hard'])}",
flush=True)
if min(len(pool["easy"]), len(pool["hard"])) < BATCH:
print("INSUFFICIENT POOL — aborting", flush=True)
return
log = []
t0 = time.time()
for step in range(STEPS):
k, labels = K_BUCKETS[step % len(K_BUCKETS)]
cand = [it for lbl in labels for it in pool[lbl]]
batch = rng.sample(cand, BATCH)
ids, msk, lab, lmask = build_batch(tok, batch)
for g in opt.param_groups:
g["lr"] = lr_at(step)
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
use_checkpoint=True, loop_mask=lmask)
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
log.append({"step": step, "k": k, "loss": loss.item()})
if step % 10 == 0:
print(f"step {step:4d} k={k} loss={loss.item():.4f} "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
torch.save(adapter.state_dict(),
OUT / f"adapter_bw_e{step+1}.pt")
json.dump(log, open(OUT / "train_bw_log.json", "w"), indent=1)
print("done", flush=True)
if __name__ == "__main__":
main()
+164
View File
@@ -0,0 +1,164 @@
"""Train the merge adapter on MBPP with prompt-only ("latent planning") looping.
Same recipe as train_merge.py, with two changes:
- loop_mask: the merge applies only to prompt positions; the code tokens are
teacher-forced through the plain band (they still attend to looped prompt
states each iteration) no exposure-bias gap by construction.
- supervision: the model's OWN passing code from prep_mbpp.py (easy: direct
pass; hard: code extracted from the plan-first pass), CE on code tokens.
"""
import argparse
import json
import os
import math
import random
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from loop_common import BandLooper, MergeAdapter
from prep_mbpp import DIRECT_SUFFIX, mbpp_prompt
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
STEPS = 600
BATCH = 4
LR = 1e-3
WARMUP = 20
MAX_TOK = 512
K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))]
ap = argparse.ArgumentParser()
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--pause", type=int, default=0,
help="pause-token control: p inert tokens after prompt, "
"feedforward adapter, no recurrence")
ARGS = ap.parse_args()
SEED = ARGS.seed
SUFFIX = (f"_s{SEED}" if SEED else "") + (f"_p{ARGS.pause}" if ARGS.pause else "")
PAUSE_ID = 6 # <unused0>
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
def build_code_batch(tok, items, device="cuda"):
"""Right-padded batch; labels on code tokens; loop_mask on prompt span."""
seqs, labs = [], []
for it in items:
p = tok(mbpp_prompt(tok, it, DIRECT_SUFFIX),
add_special_tokens=False)["input_ids"]
p = p + [PAUSE_ID] * ARGS.pause
a = tok("```python\n" + it["sol_code"] + "\n```<end_of_turn>",
add_special_tokens=False)["input_ids"]
seqs.append(p + a)
labs.append([-100] * len(p) + a)
T = max(len(s) for s in seqs)
pad = tok.pad_token_id or 0
ids = torch.full((len(seqs), T), pad, dtype=torch.long)
lab = torch.full((len(seqs), T), -100, dtype=torch.long)
msk = torch.zeros((len(seqs), T), dtype=torch.long)
for i, (s, l) in enumerate(zip(seqs, labs)):
ids[i, : len(s)] = torch.tensor(s)
lab[i, : len(s)] = torch.tensor(l)
msk[i, : len(s)] = 1
lmask = (lab == -100) & (msk == 1)
return (ids.to(device), msk.to(device), lab.to(device), lmask.to(device))
@torch.no_grad()
def val_loss(looper, adapter, tok, items, k):
tot, n = 0.0, 0
for i in range(0, len(items), BATCH):
ids, msk, lab, lmask = build_code_batch(tok, items[i : i + BATCH])
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
loop_mask=lmask, feedforward=bool(ARGS.pause))
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
tot += loss.item() * len(ids)
n += len(ids)
return tot / n
def main():
rng = random.Random(SEED)
torch.manual_seed(SEED)
data = json.load(open(OUT / "mbpp_data.json"))
model, tok = load_model(dtype=torch.bfloat16)
for p in model.parameters():
p.requires_grad_(False)
looper = BandLooper(model)
d = model.config.get_text_config().hidden_size
adapter = MergeAdapter(d=d).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
train = [it for it in data if it["split"] == "train"
and it["label"] != "drop" and it["sol_code"]]
train = [it for it in train
if len(tok(mbpp_prompt(tok, it, DIRECT_SUFFIX))["input_ids"])
+ len(tok(it["sol_code"])["input_ids"]) + 12 <= MAX_TOK]
pool = {"easy": [it for it in train if it["label"] == "easy"],
"hard": [it for it in train if it["label"] == "hard"]}
val = {lbl: lst[:16] for lbl, lst in pool.items()}
for lbl in pool:
pool[lbl] = pool[lbl][16:]
print(f"train pool: easy={len(pool['easy'])} hard={len(pool['hard'])}",
flush=True)
log = []
t0 = time.time()
for step in range(STEPS):
k, labels = K_BUCKETS[step % len(K_BUCKETS)]
cand = [it for lbl in labels for it in pool[lbl]]
batch = rng.sample(cand, min(BATCH, len(cand)))
ids, msk, lab, lmask = build_code_batch(tok, batch)
for g in opt.param_groups:
g["lr"] = lr_at(step)
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
use_checkpoint=True, loop_mask=lmask,
feedforward=bool(ARGS.pause))
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
log.append({"step": step, "k": k, "loss": loss.item()})
if step % 10 == 0:
print(f"step {step:4d} k={k} loss={loss.item():.4f} "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
vals = {}
for kk in (0, 1, 2, 4):
vals[f"easy_k{kk}"] = val_loss(looper, adapter, tok,
val["easy"], kk)
vals[f"hard_k{kk}"] = val_loss(looper, adapter, tok,
val["hard"], kk)
print(f" val@{step}: " +
" ".join(f"{n}={v:.3f}" for n, v in vals.items()), flush=True)
log.append({"step": step, "val": vals})
torch.save(adapter.state_dict(),
OUT / f"adapter_code{SUFFIX}_e{step+1}.pt")
json.dump(log, open(OUT / f"train_code_log{SUFFIX}.json", "w"), indent=1)
torch.save(adapter.state_dict(), OUT / f"adapter_code{SUFFIX}.pt")
json.dump(log, open(OUT / f"train_code_log{SUFFIX}.json", "w"), indent=1)
print("done; adapter ->", OUT / f"adapter_code{SUFFIX}.pt")
if __name__ == "__main__":
main()
+184
View File
@@ -0,0 +1,184 @@
"""Unified merge adapter: GSM8K + MBPP, prompt-only looping, fresh init.
One adapter, one regime (loop_mask = prompt span; generated/answer tokens
never loop), mixed-task batches. Hardened protocol elements: per-task val
holdouts, checkpoint every 100 steps kept separately (best-val selection and
k chosen on val, never test).
"""
import argparse
import json
import os
import math
import random
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX
from prep_mbpp import DIRECT_SUFFIX as MBPP_SUFFIX
from prep_mbpp import mbpp_prompt
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
STEPS = 800
BATCH = 6
LR = 1e-3
WARMUP = 20
MAX_TOK = 512
K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))]
SEED = 0
VAL_N = 16 # per task per bucket
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
def item_texts(tok, it):
if it["task"] == "gsm":
return (chat_prompt(tok, it["question"], DIRECT_SUFFIX),
it["gold"] + "<end_of_turn>")
return (mbpp_prompt(tok, it, MBPP_SUFFIX),
"```python\n" + it["sol_code"] + "\n```<end_of_turn>")
def build_batch(tok, items, device="cuda"):
seqs, labs = [], []
for it in items:
ptxt, atxt = item_texts(tok, it)
p = tok(ptxt, add_special_tokens=False)["input_ids"]
a = tok(atxt, add_special_tokens=False)["input_ids"]
seqs.append(p + a)
labs.append([-100] * len(p) + a)
T = max(len(s) for s in seqs)
pad = tok.pad_token_id or 0
ids = torch.full((len(seqs), T), pad, dtype=torch.long)
lab = torch.full((len(seqs), T), -100, dtype=torch.long)
msk = torch.zeros((len(seqs), T), dtype=torch.long)
for i, (s, l) in enumerate(zip(seqs, labs)):
ids[i, : len(s)] = torch.tensor(s)
lab[i, : len(s)] = torch.tensor(l)
msk[i, : len(s)] = 1
lmask = (lab == -100) & (msk == 1)
return ids.to(device), msk.to(device), lab.to(device), lmask.to(device)
@torch.no_grad()
def val_loss(looper, adapter, tok, items, k, feedforward=False):
tot, n = 0.0, 0
for i in range(0, len(items), BATCH):
ids, msk, lab, lmask = build_batch(tok, items[i : i + BATCH])
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
loop_mask=lmask, feedforward=feedforward)
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
tot += loss.item() * len(ids)
n += len(ids)
return tot / n
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--noloop", action="store_true",
help="feedforward control: adapter(e,e), no recurrence")
ap.add_argument("--tasks", default="gsm,mbpp")
ap.add_argument("--tag", default=None)
args = ap.parse_args()
tasks = args.tasks.split(",")
tag = args.tag or ("ff" if args.noloop else "uni")
rng = random.Random(SEED)
torch.manual_seed(SEED)
gsm = [dict(it, task="gsm") for it in json.load(open(OUT / "star_data.json"))
if it["split"] == "train" and it["label"] != "drop"]
mbpp = [dict(it, task="mbpp") for it in json.load(open(OUT / "mbpp_data.json"))
if it["split"] == "train" and it["label"] != "drop"
and it.get("sol_code")]
if "gsm" not in tasks:
gsm = []
if "mbpp" not in tasks:
mbpp = []
model, tok = load_model(dtype=torch.bfloat16)
for p in model.parameters():
p.requires_grad_(False)
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
def fits(it):
ptxt, atxt = item_texts(tok, it)
return (len(tok(ptxt)["input_ids"]) + len(tok(atxt)["input_ids"])
<= MAX_TOK)
items = [it for it in gsm + mbpp if fits(it)]
rng.shuffle(items)
pool, val = {"easy": [], "hard": []}, {}
for task in tasks:
for lbl in ("easy", "hard"):
sub = [it for it in items if it["task"] == task
and it["label"] == lbl]
val[(task, lbl)] = sub[:VAL_N]
pool.setdefault(lbl, []).extend(sub[VAL_N:])
print("pool: easy={} hard={} (gsm {} / mbpp {})".format(
len(pool["easy"]), len(pool["hard"]),
sum(it["task"] == "gsm" for l in pool.values() for it in l),
sum(it["task"] == "mbpp" for l in pool.values() for it in l)),
flush=True)
log = []
t0 = time.time()
for step in range(STEPS):
k, labels = K_BUCKETS[step % len(K_BUCKETS)]
cand = [it for lbl in labels for it in pool[lbl]]
batch = rng.sample(cand, BATCH)
ids, msk, lab, lmask = build_batch(tok, batch)
for g in opt.param_groups:
g["lr"] = lr_at(step)
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
use_checkpoint=True, loop_mask=lmask,
feedforward=args.noloop)
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
log.append({"step": step, "k": k, "loss": loss.item()})
if step % 10 == 0:
print(f"step {step:4d} k={k} loss={loss.item():.4f} "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
vals = {}
kk_grid = (0, 1) if args.noloop else (0, 1, 2, 4)
for (task, lbl), vitems in val.items():
for kk in kk_grid:
vals[f"{task}_{lbl}_k{kk}"] = val_loss(
looper, adapter, tok, vitems, kk,
feedforward=args.noloop)
print(f" val@{step}: " + " ".join(
f"{n}={v:.3f}" for n, v in sorted(vals.items())), flush=True)
log.append({"step": step, "val": vals})
torch.save(adapter.state_dict(),
OUT / f"adapter_{tag}_e{step+1}.pt")
json.dump(log, open(OUT / f"train_{tag}_log.json", "w"), indent=1)
print("done", flush=True)
if __name__ == "__main__":
main()