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:
@@ -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)
|
||||
Reference in New Issue
Block a user