Files
jspace/scripts/loop_common.py
T
2026-07-14 01:14:57 +02:00

313 lines
13 KiB
Python

"""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:
unk = tok.unk_token_id
eos = [t for t in (tok.eos_token_id,
tok.convert_tokens_to_ids("<end_of_turn>"),
tok.convert_tokens_to_ids("<turn|>"))
if t is not None and t >= 0 and t != unk]
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)