171 lines
6.9 KiB
Python
171 lines
6.9 KiB
Python
"""Learnable soft path through ALL layers (pre-registration item 14).
|
|
|
|
The frozen model = 35 same-typed functions on one residual bus. Instead of
|
|
a hand-fixed loop over L14-30, learn a gate matrix g[t, l] in [0,1]: on
|
|
loop iteration t, layer l's residual delta is scaled by g[t, l] (prompt
|
|
positions only; generated/answer positions always run ungated). Gates are
|
|
initialized as a Gaussian bump over depth centered mid-band, so at init a
|
|
loop iteration is approximately the hand band pass — then SGD may move the
|
|
compute envelope anywhere in [0, n_layers). The anchor merge adapter is
|
|
kept at each iteration boundary for stability (rho < 1).
|
|
|
|
Reading the result: if the learned envelope concentrates on the lens band,
|
|
gradient descent independently rediscovers the workspace; if it wins with
|
|
mass elsewhere, the lens placement story needs revision.
|
|
|
|
E2B caveat (stated in advance): KV sharing makes attention deltas of
|
|
layers >= 15 loop-inert; gate mass there is interpretable for MLP deltas
|
|
only.
|
|
"""
|
|
|
|
import math
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
from torch.utils.checkpoint import checkpoint
|
|
|
|
from loop_common import BAND, BandLooper, _text_model
|
|
|
|
|
|
class PathGates(nn.Module):
|
|
"""g[t, l] = sigmoid(logits[t, l]); row 0 = warm sweep, rows 1..k = loops."""
|
|
|
|
def __init__(self, n_layers, k_max, mu=None, sigma=6.0, band=BAND):
|
|
super().__init__()
|
|
mu = (band[0] + band[1]) / 2 if mu is None else mu
|
|
init = torch.empty(k_max + 1, n_layers)
|
|
for l in range(n_layers):
|
|
p = math.exp(-((l - mu) ** 2) / (2 * sigma ** 2))
|
|
p = min(max(p, 1e-3), 1 - 1e-3)
|
|
init[:, l] = math.log(p / (1 - p))
|
|
self.logits = nn.Parameter(init)
|
|
|
|
def g(self, t):
|
|
return torch.sigmoid(
|
|
self.logits[min(t, self.logits.shape[0] - 1)].float())
|
|
|
|
def envelope(self):
|
|
with torch.no_grad():
|
|
return torch.sigmoid(self.logits.float()).tolist()
|
|
|
|
|
|
class GatedLooper(BandLooper):
|
|
"""BandLooper over the FULL depth with per-iteration per-layer gates."""
|
|
|
|
def __init__(self, model, gates):
|
|
tm = _text_model(model)
|
|
super().__init__(model, band=(0, len(tm.layers) - 1))
|
|
self.gates = gates
|
|
|
|
def _gated(self, h, calls, t, loop_mask=None):
|
|
g = self.gates.g(t).to(h.device)
|
|
x = h
|
|
for i in range(self.l0, self.l1 + 1):
|
|
args, kwargs = calls[i]
|
|
out = self.tm.layers[i](x, *args, **kwargs)
|
|
if isinstance(out, tuple):
|
|
out = out[0]
|
|
gi = g[i].to(x.dtype)
|
|
if loop_mask is not None: # ungated (g=1) off the prompt span
|
|
gi = torch.where(loop_mask[..., None], gi,
|
|
torch.ones_like(loop_mask[..., None],
|
|
dtype=x.dtype))
|
|
x = x + gi * (out - x)
|
|
return x
|
|
|
|
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):
|
|
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
|
|
e = self._hin[self.l0].detach()
|
|
|
|
def sweep(x, t):
|
|
if use_checkpoint:
|
|
return checkpoint(
|
|
lambda x_: self._gated(x_, calls, t, loop_mask), x,
|
|
use_reentrant=False)
|
|
return self._gated(x, calls, t, loop_mask)
|
|
|
|
s = sweep(e, 0) # warm sweep, t=0 (gates trainable here too)
|
|
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._gated(x, calls, i + 1, loop_mask)
|
|
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)
|
|
s = sweep(x, i + 1)
|
|
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 generate_frozen_prompt(self, adapter, tok, input_ids, k,
|
|
max_new_tokens=220, attention_mask=None,
|
|
stop_strs=(), feedforward=False,
|
|
conv_out=None):
|
|
"""Per-layer KV write-in: run the final gated sweep recording every
|
|
layer's INPUT, then one native prefill with pre-forward hooks
|
|
swapping each layer's hidden_states to the recorded stream — the
|
|
cache then holds exactly the gated states; decode is native."""
|
|
if k == 0:
|
|
return super().generate_frozen_prompt(
|
|
adapter, tok, input_ids, 0, max_new_tokens=max_new_tokens,
|
|
attention_mask=attention_mask, stop_strs=stop_strs)
|
|
calls, _ = self.capture(input_ids, attention_mask, logits_to_keep=1)
|
|
e = self._hin[self.l0]
|
|
s = self._gated(e, calls, 0)
|
|
for i in range(k):
|
|
x = adapter(e, s)
|
|
if i < k - 1:
|
|
s = self._gated(x, calls, i + 1)
|
|
# final sweep: record per-layer inputs of the gated stream
|
|
xs = {}
|
|
g = self.gates.g(k).to(x.device)
|
|
h = x
|
|
for i in range(self.l0, self.l1 + 1):
|
|
xs[i] = h
|
|
args, kwargs = calls[i]
|
|
out = self.tm.layers[i](h, *args, **kwargs)
|
|
if isinstance(out, tuple):
|
|
out = out[0]
|
|
h = h + g[i].to(h.dtype) * (out - h)
|
|
del calls
|
|
|
|
P = input_ids.shape[1]
|
|
handles = []
|
|
for i in range(self.l0, self.l1 + 1):
|
|
def pre(mod, args, kwargs, i=i):
|
|
hh = kwargs.get("hidden_states",
|
|
args[0] if args else None)
|
|
if hh is not None and hh.shape[1] == P: # prefill only
|
|
if "hidden_states" in kwargs:
|
|
kwargs["hidden_states"] = xs[i].to(hh.dtype)
|
|
return args, kwargs
|
|
return (xs[i].to(hh.dtype),) + args[1:], kwargs
|
|
return None
|
|
handles.append(self.tm.layers[i].register_forward_pre_hook(
|
|
pre, with_kwargs=True))
|
|
try:
|
|
gen = self.model.generate(
|
|
input_ids=input_ids, attention_mask=attention_mask,
|
|
max_new_tokens=max_new_tokens, do_sample=False,
|
|
pad_token_id=tok.pad_token_id or 0)
|
|
finally:
|
|
for hd in handles:
|
|
hd.remove()
|
|
return gen
|