Chimera-64M / family.py
Quazim0t0's picture
Upload family.py with huggingface_hub
db5e2b4 verified
Raw History Blame Contribute Delete
22 kB
"""SpikeWhale / Byrne family traits, ported to Quazimoto-LM (attention backbone).
These are the *transformer-native* family blocks (their origin is the transformer
`modeling_byrne_embed.py`), so unlike the SNN port they operate directly on the
sequence hidden state [B,T,d] -- no per-step adaptation needed. Each keeps the
family's safe-at-init contract: a tanh/zero gate makes the block a no-op at start,
while the content (`up`/`down`) weights are NON-zero so the gate still receives
gradient (the double-zero saddle would deadlock it). DERF soft_clamp bounds any
new instability surface, in line with the family's stability discipline.
Included: HRMRefinementBlock (signature), MoESwiGLU, MTPHead, JEPAPredictorBlock.
Engram / ProgSem are bio/SNN-specific and SpikingLinearAttention is the SNN's
stand-in for the real attention this model already has, so they are omitted.
"""
from __future__ import annotations
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import instrument as _viz # live-visualizer capture hooks (no-op unless a recorder is active)
# sqrt(pi)/2: soft_clamp is the identity for small inputs and saturates smoothly to
# +/-bound with a non-zero gradient everywhere (no dead-gradient zones).
_ERF_K = math.sqrt(math.pi) / 2.0
def soft_clamp(x, bound):
return bound * torch.erf(x * (_ERF_K / bound))
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight
def sqrtsoftplus(x):
"""Family expert-scoring function: sqrt(softplus(x))."""
return torch.sqrt(F.softplus(x) + 1e-8)
class HRMRefinementBlock(nn.Module):
"""Signature family block: iterative gated refinement that, in the canonical
HRM spirit, starts the reasoning from a RANDOM initial state z0 and iterates
toward a solution conditioned on the input (`anchor`).
Unlike the other family traits this block does NOT start as a no-op: the gate
is initialised OPEN (gate_init_open). With an input-anchored no-op start the
gate received ~zero gradient and never woke up; a random z0 forces the block
to actively reconcile the random state against the input, so the open gate
carries real signal from step 0. To keep the deep trunk intact we contribute
only the reasoning DELTA (h - z0) as a residual -- z0 itself is never dumped
into the trunk, and a closed gate (h == z0) degrades cleanly to a no-op."""
def __init__(self, hidden_size, refine_dim, steps, eps=1e-3, gate_init_open=0.1):
super().__init__()
self.steps = steps
self.norm = RMSNorm(hidden_size, eps)
self.down = nn.Linear(hidden_size * 2, refine_dim, bias=False)
self.up = nn.Linear(refine_dim, hidden_size, bias=False)
# random initial reasoning state (learnable), broadcast over batch/time
self.z0 = nn.Parameter(torch.empty(hidden_size))
nn.init.trunc_normal_(self.z0, std=1.0, a=-2.0, b=2.0)
# gates start OPEN so the random-state reasoning reaches the output at init
go = math.atanh(min(gate_init_open, 0.9)) if gate_init_open > 0 else 0.0
self.gate = nn.Parameter(torch.full((steps,), go))
nn.init.normal_(self.down.weight, std=0.02)
nn.init.normal_(self.up.weight, std=0.02)
def forward(self, x): # x: [B,T,d]
B, T, _ = x.shape
anchor = x
h = self.z0.expand(B, T, -1) # random initial reasoning state
for t in range(self.steps):
inp = torch.cat([self.norm(h), anchor], dim=-1)
update = soft_clamp(self.up(F.silu(self.down(inp))), 10.0)
h = h + torch.tanh(self.gate[t]) * update
return x + (h - self.z0) # add reasoning delta, keep trunk
class ExpertFFN(nn.Module):
def __init__(self, hidden_size, intermediate_size):
super().__init__()
self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class MoESwiGLU(nn.Module):
"""Shared + top-k routed SwiGLU experts, sqrtsoftplus scoring, norm_topk_prob,
Switch-style load-balance aux (via `last_aux_loss`). down-projections zero-init
=> no-op at start."""
def __init__(self, hidden_size, intermediate_size, n_routed=4, n_shared=1,
top_k=2, aux_loss_coef=0.01):
super().__init__()
self.top_k = min(top_k, n_routed)
self.n_routed = n_routed
self.n_shared = n_shared
self.aux_loss_coef = aux_loss_coef
self.router = nn.Linear(hidden_size, n_routed, bias=False)
self.experts = nn.ModuleList([ExpertFFN(hidden_size, intermediate_size)
for _ in range(n_routed)])
self.shared = (ExpertFFN(hidden_size, intermediate_size * n_shared)
if n_shared > 0 else None)
for e in self.experts:
nn.init.zeros_(e.down_proj.weight)
if self.shared is not None:
nn.init.zeros_(self.shared.down_proj.weight)
self.last_aux_loss = None
def forward(self, x): # x: [B,T,d]
flat = x.reshape(-1, x.shape[-1])
out = torch.zeros_like(flat)
if self.shared is not None:
s = self.shared(flat)
out = out + (s / self.n_shared if self.n_shared > 1 else s)
logits = self.router(flat)
scores = sqrtsoftplus(logits)
topv, topi = scores.topk(self.top_k, dim=-1)
topv = topv / (topv.sum(-1, keepdim=True) + 1e-8)
for slot in range(self.top_k):
idx = topi[:, slot]
w = topv[:, slot].unsqueeze(-1)
for e_id, expert in enumerate(self.experts):
mask = idx == e_id
if mask.any():
out[mask] = out[mask] + w[mask] * expert(flat[mask])
probs = F.softmax(logits, dim=-1)
expert_mask = torch.zeros_like(probs)
expert_mask.scatter_(1, topi, 1.0)
self.last_aux_loss = (self.n_routed * (expert_mask.mean(0) * probs.mean(0)).sum()
* self.aux_loss_coef)
return out.view_as(x)
class MTPHead(nn.Module):
"""Multi-token-prediction head: zero-init d->d residual, reuses the tied readout."""
def __init__(self, hidden_size):
super().__init__()
self.proj = nn.Linear(hidden_size, hidden_size, bias=False)
nn.init.zeros_(self.proj.weight)
def forward(self, hidden):
return hidden + self.proj(hidden)
class TokenCompressor(nn.Module):
"""Frozen LSH-style projection (gradient never reaches it through the hash cast)."""
def __init__(self, hidden_size, compress_dim):
super().__init__()
self.proj = nn.Linear(hidden_size, compress_dim, bias=False)
nn.init.normal_(self.proj.weight, std=0.02)
self.proj.weight.requires_grad_(False)
def forward(self, x):
return self.proj(x)
class MultiHeadHashLookup(nn.Module):
"""N-gram hash memory: for n=1..max_ngram, hash the n-token compressed window
into per-head tables and average. Ported from v2 EngramModule."""
def __init__(self, num_heads, table_size, compress_dim, out_dim, max_ngram=3):
super().__init__()
self.num_heads, self.table_size = num_heads, table_size
self.max_ngram, self.out_dim = max_ngram, out_dim
self.tables = nn.ModuleList([nn.Embedding(table_size, out_dim) for _ in range(num_heads)])
for t in self.tables:
nn.init.normal_(t.weight, std=0.01)
for n in range(1, max_ngram + 1):
for k in range(n):
proj = torch.randn(num_heads, compress_dim)
proj = proj / (proj.norm(dim=1, keepdim=True) + 1e-8)
self.register_buffer(f"hash_proj_n{n}_p{k}", proj, persistent=True)
def forward(self, compressed):
B, S, _ = compressed.shape
dev = compressed.device
out = torch.zeros(B, S, self.out_dim, device=dev, dtype=compressed.dtype)
norm = torch.zeros(S, device=dev)
for n in range(1, self.max_ngram + 1):
if S < n:
continue
valid, start = S - n + 1, n - 1
h = torch.zeros(B, valid, self.num_heads, device=dev)
for k in range(n):
proj = getattr(self, f"hash_proj_n{n}_p{k}")
h = h + torch.matmul(compressed[:, k:k + valid, :].float(), proj.t())
idx = h.abs().long() % self.table_size
for hi, table in enumerate(self.tables):
out[:, start:, :] = out[:, start:, :] + table(idx[:, :, hi])
norm[start:] += self.num_heads
return (out / norm.view(1, -1, 1).clamp(min=1)).to(compressed.dtype)
class DERFContextGate(nn.Module):
def __init__(self, obs_size, init_bias=-4.0):
super().__init__()
self.proj = nn.Linear(obs_size * 2, obs_size)
self.alpha = nn.Parameter(torch.ones(obs_size))
self.bias = nn.Parameter(torch.full((obs_size,), init_bias))
self.gamma = nn.Parameter(torch.ones(obs_size))
def forward(self, retrieved, obs):
logits = self.proj(torch.cat([retrieved, obs], dim=-1))
gate = self.gamma * ((torch.erf(self.alpha * logits + self.bias) + 1.0) / 2.0)
return retrieved * gate
class PhaseAttentionRing(nn.Module):
"""Interstitial ATTENTION ring: attends causally over the sequence in
oscillator-PHASE space ([cos,sin] of the two neighbor oscillator rings) and
returns an injection current of width m = n_r + n_{r+1} for those neighbors.
Zero-init gate => no-op at start; soft_clamp bounds the injected drive."""
def __init__(self, m, n_heads=4, head_dim=16, bound=10.0):
super().__init__()
self.h, self.d, self.bound = n_heads, head_dim, bound
self.qkv = nn.Linear(2 * m, 3 * n_heads * head_dim, bias=False)
self.out = nn.Linear(n_heads * head_dim, m, bias=False)
self.gate = nn.Parameter(torch.zeros(1))
nn.init.normal_(self.qkv.weight, std=0.02)
nn.init.normal_(self.out.weight, std=0.02)
def forward(self, theta_slice): # [B,T,m] phases of the neighbor rings
B, T, m = theta_slice.shape
feat = torch.cat([torch.cos(theta_slice), torch.sin(theta_slice)], dim=-1)
q, k, v = self.qkv(feat).split(self.h * self.d, dim=-1)
shp = lambda z: z.view(B, T, self.h, self.d).transpose(1, 2)
y = F.scaled_dot_product_attention(shp(q), shp(k), shp(v), is_causal=True)
y = y.transpose(1, 2).reshape(B, T, self.h * self.d)
return soft_clamp(self.out(y) * torch.tanh(self.gate), self.bound)
class EngramRing(nn.Module):
"""Interstitial ENGRAM ring: absorbs n-gram context from the hidden state via
hash memory + DERF gate, projected to an injection current of width m for the
two neighbor oscillator rings. Two no-op gates at init (DERF bias -4 + scale)."""
def __init__(self, hidden_size, m, compress_dim=32, num_heads=2,
table_size=2048, max_ngram=3, bound=10.0):
super().__init__()
self.bound = bound
self.compressor = TokenCompressor(hidden_size, compress_dim)
self.lookup = MultiHeadHashLookup(num_heads, table_size, compress_dim, m, max_ngram)
self.to_obs = nn.Linear(hidden_size, m, bias=False)
self.gate = DERFContextGate(m, init_bias=-4.0)
self.scale = nn.Parameter(torch.zeros(1)) # extra no-op gate at init
nn.init.normal_(self.to_obs.weight, std=0.02)
def family_reinit(self):
"""Re-apply the inits the model's global self.apply would clobber (table
std 0.01, frozen-random compressor, DERF bias -4)."""
for t in self.lookup.tables:
nn.init.normal_(t.weight, std=0.01)
nn.init.normal_(self.compressor.proj.weight, std=0.02)
self.compressor.proj.weight.requires_grad_(False)
nn.init.constant_(self.gate.bias, -4.0)
def forward(self, h): # h: [B,T,hidden]
retrieved = self.lookup(self.compressor(h.detach()))
gated = self.gate(retrieved, self.to_obs(h))
return soft_clamp(gated * torch.tanh(self.scale), self.bound)
class RingController(nn.Module):
"""Tiny per-ring manager that OPTIMIZES ITSELF by a predictive / free-energy rule.
Core = a fast-weight linear predictor `W` (a BUFFER, excluded from the global
optimizer) updated online by the delta rule W += lr * (f - W@prev) outer prev,
which is exactly one gradient step on the squared prediction error -- the
controller learns to predict its ring's next state, minimizing surprise, with no
backprop. A small backprop-trained decoder maps the self-organized feature +
surprise into ring-control modulations; zero-init => exact no-op at start."""
def __init__(self, d_obs=4, feat=384, n_ctrl=4, local_lr=0.01):
super().__init__()
self.feat, self.local_lr = feat, local_lr
self.enc = nn.Linear(d_obs, feat)
self.dec = nn.Linear(feat + 1, n_ctrl)
nn.init.zeros_(self.dec.weight)
nn.init.zeros_(self.dec.bias) # control == 0 at init (no-op)
self.register_buffer("W", torch.zeros(feat, feat)) # self-organizing fast weights
self.register_buffer("prev_f", torch.zeros(feat))
def family_reinit(self):
nn.init.zeros_(self.dec.weight)
nn.init.zeros_(self.dec.bias)
def forward(self, obs): # obs: [d_obs] (detached ring stats)
f = torch.tanh(self.enc(obs)) # [feat]
pred = self.W @ self.prev_f # predicted current feature
surprise = F.mse_loss(f.detach(), pred)
if self.training:
with torch.no_grad(): # predictive self-organization (no global grad)
err = f.detach() - pred
self.W.add_(self.local_lr * torch.outer(err, self.prev_f)).clamp_(-3.0, 3.0)
self.prev_f.copy_(f.detach())
ctrl = self.dec(torch.cat([f, surprise.detach().reshape(1)])) # [n_ctrl]
return ctrl, surprise.detach()
class RingControllerBank(nn.Module):
"""One RingController per oscillator ring (shared across all layers)."""
def __init__(self, n_rings, d_obs=4, feat=384, local_lr=0.01):
super().__init__()
self.controllers = nn.ModuleList(
[RingController(d_obs, feat, 4, local_lr) for _ in range(n_rings)])
self.last_surprise = None
def forward(self, obs): # obs: [R, d_obs] -> ctrl [R, 4]
ctrls, surps = [], []
for r, c in enumerate(self.controllers):
ct, sp = c(obs[r])
ctrls.append(ct)
surps.append(sp)
self.last_surprise = torch.stack(surps).mean()
return torch.stack(ctrls, dim=0)
class RingSpecialists(nn.Module):
"""A MoE-style bank of `n_spec` MINI MEMORY SPECIALISTS for ONE oscillator ring.
Each specialist owns two fast-weight stores (test-time-mutable BUFFERS, like
RingController.W -- excluded from the optimizer):
* store_in -- a memory of the INPUT context that routes to it, and
* store_out -- the OUTPUT information it injects back into the ring.
Tokens are routed to the top-k specialists (a small MoE) by similarity to each
specialist's address = its learnable identity key + a read of its accumulated
input memory. The routed store_out is decoded into an injection current for the
ring, and BOTH stores are written online (train AND inference) by a gated EMA
rule -- so a generation accumulates an addressable context memory as it runs.
Slow (backprop) weights -- q_proj/in_enc/val_enc/out_dec/in_read/key/active --
learn to route, encode, retrieve and decode; the stores are the fast memory.
Family contract: zero-init `scale` => exact no-op at start, and an empty
store_out is zero anyway, so the block is doubly safe until it learns to write
and open the gate. `active` is a per-specialist usage gate biasing the router."""
def __init__(self, ring_size, hidden_size, n_spec=7, key_dim=32, slot_dim=64,
top_k=2, write_lr=0.1, bound=10.0):
super().__init__()
self.n_spec = n_spec
self.top_k = min(top_k, n_spec)
self.write_lr, self.bound = write_lr, bound
self.write_enabled = True
self.q_proj = nn.Linear(hidden_size, key_dim, bias=False) # router query (learned)
self.out_dec = nn.Linear(slot_dim, ring_size, bias=False) # retrieved -> injection (learned)
self.in_read = nn.Linear(slot_dim, key_dim, bias=False) # input-store -> addr (learned)
# write-side encoders are FROZEN RANDOM projections (cf. EngramRing's frozen
# compressor): they only ever run inside the no-grad write, so backprop can't
# train them -- as fixed random features the stores hold a stable encoding the
# learned read/route/decode path can address.
self.in_enc = nn.Linear(hidden_size, slot_dim, bias=False) # input -> input-store (frozen)
self.val_enc = nn.Linear(hidden_size, slot_dim, bias=False) # input -> output-store (frozen)
self.key = nn.Parameter(torch.randn(n_spec, key_dim) * 0.02) # specialist identity
self.active = nn.Parameter(torch.zeros(n_spec)) # per-specialist usage gate
self.scale = nn.Parameter(torch.zeros(1)) # no-op output gate at init
self.register_buffer("store_in", torch.zeros(n_spec, slot_dim))
self.register_buffer("store_out", torch.zeros(n_spec, slot_dim))
for m in (self.q_proj, self.out_dec, self.in_read, self.in_enc, self.val_enc):
nn.init.normal_(m.weight, std=0.02)
self.in_enc.weight.requires_grad_(False)
self.val_enc.weight.requires_grad_(False)
def family_reinit(self):
"""Re-apply inits the model's global self.apply clobbers (key, gates, and
the frozen write encoders)."""
nn.init.normal_(self.key, std=0.02)
nn.init.zeros_(self.active)
nn.init.zeros_(self.scale)
nn.init.normal_(self.in_enc.weight, std=0.02)
nn.init.normal_(self.val_enc.weight, std=0.02)
self.in_enc.weight.requires_grad_(False)
self.val_enc.weight.requires_grad_(False)
def reset_memory(self):
"""Clear both stores -- call between independent prompts/sequences so
context memory does not bleed across them."""
self.store_in.zero_()
self.store_out.zero_()
def forward(self, h): # h: [B,T,hidden]
B, T, _ = h.shape
# snapshot the fast-weight stores: the graph must hold an immutable copy
# because we mutate the buffers in-place for the online write below.
store_in, store_out = self.store_in.clone(), self.store_out.clone()
q = self.q_proj(h) # [B,T,key_dim]
addr = self.key + self.in_read(store_in) # [n_spec,key_dim]
logits = q @ addr.t() # [B,T,n_spec]
logits = logits + F.logsigmoid(self.active) # usage gate biases routing
if self.top_k < self.n_spec: # top-k MoE sparsity
tv = torch.topk(logits, self.top_k, dim=-1).values
logits = logits.masked_fill(logits < tv[..., [-1]], float("-inf"))
route = torch.softmax(logits, dim=-1) # [B,T,n_spec]
retrieved = route @ store_out # [B,T,slot_dim]
inject = self.out_dec(retrieved) * torch.tanh(self.scale) # [B,T,ring_size]
rec = _viz.get_rec()
if rec is not None and rec.enabled: # last-token routing
rec.push_spec(route[0, -1].tolist())
# online write: blend this step's input/value into the routed specialists
if self.write_enabled and self.write_lr > 0:
with torch.no_grad():
w = route.reshape(-1, self.n_spec) # [BT,n_spec]
denom = w.sum(0).clamp(min=1e-3).unsqueeze(1) # [n_spec,1]
in_info = (w.t() @ self.in_enc(h).reshape(-1, self.in_enc.out_features)) / denom
val_info = (w.t() @ self.val_enc(h).reshape(-1, self.val_enc.out_features)) / denom
a = self.write_lr
self.store_in.mul_(1 - a).add_(a * in_info).clamp_(-self.bound, self.bound)
self.store_out.mul_(1 - a).add_(a * val_info).clamp_(-self.bound, self.bound)
return soft_clamp(inject, self.bound)
class JEPAPredictorBlock(nn.Module):
"""Representation-space k-ahead prediction with stop-grad target (JEPA asymmetry).
Zero-init gate => identity at init; `up` normal so the gate gets gradient."""
def __init__(self, dim, pred_dim, horizon, eps=1e-3):
super().__init__()
self.horizon = horizon
self.norm = RMSNorm(dim, eps)
self.down = nn.Linear(dim, pred_dim, bias=False)
self.up = nn.Linear(pred_dim, dim, bias=False)
self.gate = nn.Parameter(torch.zeros(horizon))
nn.init.normal_(self.down.weight, std=0.02)
nn.init.normal_(self.up.weight, std=0.02)
def forward(self, h, k): # h: [B,T,dim]
update = self.up(F.silu(self.down(self.norm(h))))
return h + torch.tanh(self.gate[k - 1]) * update