"""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 AdaptiveMergeAdapter(nn.Module): """Merge with state-dependent anchor coefficient (Lys-inspired, trained). alpha(e, s) = sigmoid(w·[e;ŝ] + b), per position; w zero-init and b = logit(0.3), so at init this is exactly the fixed alpha=0.3 merge.""" def __init__(self, d=1536, hidden=512, alpha0=0.3): super().__init__() 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) self.alpha_head = nn.Linear(2 * d, 1) nn.init.zeros_(self.alpha_head.weight) import math as _m nn.init.constant_(self.alpha_head.bias, _m.log(alpha0 / (1 - alpha0))) 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) ) cat = torch.cat([e32, s_hat], dim=-1) a = torch.sigmoid(self.alpha_head(cat)) out = (1 - a) * e32 + a * s_hat + self.mlp(cat) 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): try: from lora_band import loop_active loop_active(True) out = self._run(h, calls, self.l0, self.l1) loop_active(False) return out except ImportError: 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, bptt=None): """bptt: backprop only through the last `bptt` iterations (McLeish- style truncated BPTT); earlier iterations run under no_grad.""" """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] n_nograd = max(0, k - bptt) if bptt else 0 for i in range(k): if i < n_nograd: with torch.no_grad(): x = adapter(e, s) if loop_mask is not None: x = torch.where(loop_mask[..., None], x, e) s = self.band(x, calls) s = s.detach() states.append(s) continue 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("") 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(""), tok.convert_tokens_to_ids("")) 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: #### ") 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"] + "", 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)