Download family.py from Quazim0t0/Chimera-64M: direct link, hf CLI and curl.
- Browser
- Download file 22 kB
-
https://huggingface.co/Quazim0t0/Chimera-64M/resolve/main/family.py
- Command line
-
hf download hf://Quazim0t0/Chimera-64M/family.py
-
curl -L -o family.py https://huggingface.co/Quazim0t0/Chimera-64M/resolve/main/family.py
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 | |