"""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