Upload chimera.py with huggingface_hub
Browse files- chimera.py +215 -0
chimera.py
ADDED
|
@@ -0,0 +1,215 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
chimera.py -- the ChimeraBlock: the grand-finale channel mixer that fuses EVERY project
|
| 3 |
+
in the family into one block. A creature made of parts of many beasts.
|
| 4 |
+
|
| 5 |
+
Per token, three PHYSICS CORES each propose a candidate update of the hidden state:
|
| 6 |
+
|
| 7 |
+
* KuramotoCore -- coupled phase oscillators (Quazimoto): mean-field Kuramoto with
|
| 8 |
+
learnable frustration; a few Euler steps; readout [cos, sin].
|
| 9 |
+
* GrowthCore -- Neighbour-Sensing fungal growth (Mycel): tips in a bounded latent
|
| 10 |
+
region sense a low-rank density field and steer (negative autotropism).
|
| 11 |
+
* WaveCore -- Wheeler-DeWitt wave (Wheeler): K minisuperspace modes under a
|
| 12 |
+
LORENTZIAN supermetric, leapfrog wave steps; exposes the Hamiltonian
|
| 13 |
+
constraint <H^2> so the block can be pressured onto H Psi = 0.
|
| 14 |
+
|
| 15 |
+
Then the GRPO GROUP-RELATIVE SELECTION (grpo_lm) is the META-MIXER: an internal critic
|
| 16 |
+
scores the three candidates, group-relative advantages A_g = (r_g - mean)/std decide the
|
| 17 |
+
winner RELATIVE to its peers, the mixing weights are CLIPPED around uniform (the PPO clip)
|
| 18 |
+
and ANCHORED back toward uniform by beta (the KL-to-reference anchor). The output is the
|
| 19 |
+
selected convex combination -- so the network LEARNS WHICH LAW OF PHYSICS to apply to each
|
| 20 |
+
token. Fractal phase (fractal.py) seeds all three cores; Elo attention lives in the
|
| 21 |
+
attention block. Every idea we designed, in one place, behind one family gate.
|
| 22 |
+
"""
|
| 23 |
+
import math
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
import torch.nn.functional as F
|
| 27 |
+
|
| 28 |
+
from family import soft_clamp, RMSNorm
|
| 29 |
+
import instrument as _viz
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
# --------------------------------------------------------------------------------------
|
| 33 |
+
# physics cores -- each takes normed hidden h [B,T,d] (+ optional fractal seed) and returns
|
| 34 |
+
# a candidate update [B,T,d]. Compact readouts (feat -> d -> d) keep the three-core cost sane.
|
| 35 |
+
# --------------------------------------------------------------------------------------
|
| 36 |
+
class KuramotoCore(nn.Module):
|
| 37 |
+
"""Quazimoto: a single mean-field Kuramoto ring with learnable frustration alpha."""
|
| 38 |
+
|
| 39 |
+
def __init__(self, cfg):
|
| 40 |
+
super().__init__()
|
| 41 |
+
self.N = cfg.chim_osc
|
| 42 |
+
self.steps, self.dt = cfg.chim_osc_steps, cfg.osc_dt
|
| 43 |
+
self.to_theta = nn.Linear(cfg.d_model, self.N)
|
| 44 |
+
self.to_omega = nn.Linear(cfg.d_model, self.N)
|
| 45 |
+
self.k_coupling = nn.Parameter(torch.tensor(1.0)) # global coupling strength
|
| 46 |
+
self.alpha = nn.Parameter(torch.zeros(1)) # Sakaguchi frustration
|
| 47 |
+
self.read = nn.Sequential(nn.Linear(2 * self.N, cfg.d_model), nn.GELU(),
|
| 48 |
+
nn.Linear(cfg.d_model, cfg.d_model))
|
| 49 |
+
|
| 50 |
+
def forward(self, h, seed=None):
|
| 51 |
+
theta = self.to_theta(h)
|
| 52 |
+
if seed is not None:
|
| 53 |
+
theta = theta + seed[..., :self.N]
|
| 54 |
+
omega = torch.tanh(self.to_omega(h))
|
| 55 |
+
a = self.alpha
|
| 56 |
+
for _ in range(self.steps):
|
| 57 |
+
c, s = torch.cos(theta), torch.sin(theta)
|
| 58 |
+
zc = c.mean(-1, keepdim=True) # mean-field order parameter
|
| 59 |
+
zs = s.mean(-1, keepdim=True)
|
| 60 |
+
# K * Im( e^{i(alpha - theta)} * z ) = K*(sin(a-θ)zc + cos(a-θ)zs)
|
| 61 |
+
coupling = self.k_coupling * (torch.sin(a - theta) * zc + torch.cos(a - theta) * zs)
|
| 62 |
+
theta = theta + self.dt * (omega + coupling)
|
| 63 |
+
return self.read(torch.cat([torch.cos(theta), torch.sin(theta)], dim=-1))
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
class GrowthCore(nn.Module):
|
| 67 |
+
"""Mycel: Neighbour-Sensing tips growing in a bounded latent region, sensing a
|
| 68 |
+
low-rank density field (O(N*F)) and steering away from their own density."""
|
| 69 |
+
|
| 70 |
+
def __init__(self, cfg):
|
| 71 |
+
super().__init__()
|
| 72 |
+
self.N, self.pd, self.F = cfg.chim_tips, cfg.chim_pos_dim, cfg.chim_centers
|
| 73 |
+
self.steps, self.dt, self.bound = cfg.chim_growth_steps, cfg.chim_growth_dt, cfg.osc_bound
|
| 74 |
+
self.to_pos = nn.Linear(cfg.d_model, self.N * self.pd)
|
| 75 |
+
self.to_dir = nn.Linear(cfg.d_model, self.N * self.pd)
|
| 76 |
+
self.centers = nn.Parameter(torch.randn(self.F, self.pd) * 0.5)
|
| 77 |
+
self.log_persist = nn.Parameter(torch.tensor(1.4))
|
| 78 |
+
self.tropism = nn.Parameter(torch.zeros(1))
|
| 79 |
+
self.log_bw = nn.Parameter(torch.zeros(1))
|
| 80 |
+
self.read = nn.Sequential(nn.Linear(self.N * (2 * self.pd + 1), cfg.d_model), nn.GELU(),
|
| 81 |
+
nn.Linear(cfg.d_model, cfg.d_model))
|
| 82 |
+
|
| 83 |
+
def _clamp(self, p):
|
| 84 |
+
return self.bound * torch.tanh(p / self.bound)
|
| 85 |
+
|
| 86 |
+
def _cdist2(self, p, c):
|
| 87 |
+
p2 = (p * p).sum(-1, keepdim=True)
|
| 88 |
+
c2 = (c * c).sum(-1)
|
| 89 |
+
return (p2 + c2 - 2.0 * (p @ c.t())).clamp(min=0.0)
|
| 90 |
+
|
| 91 |
+
def forward(self, h, seed=None):
|
| 92 |
+
B, T, _ = h.shape
|
| 93 |
+
p = self._clamp(self.to_pos(h).view(B, T, self.N, self.pd))
|
| 94 |
+
if seed is not None:
|
| 95 |
+
p = self._clamp(p + seed[..., :self.N * self.pd].view(B, T, self.N, self.pd))
|
| 96 |
+
v = self.to_dir(h).view(B, T, self.N, self.pd)
|
| 97 |
+
pers, trop = torch.sigmoid(self.log_persist), torch.tanh(self.tropism)
|
| 98 |
+
bw = F.softplus(self.log_bw).clamp(min=1e-3)
|
| 99 |
+
for _ in range(self.steps):
|
| 100 |
+
K = torch.exp(-self._cdist2(p, self.centers) / bw)
|
| 101 |
+
w = K.mean(-2, keepdim=True) * K
|
| 102 |
+
away = p * w.sum(-1, keepdim=True) - torch.matmul(w, self.centers)
|
| 103 |
+
v = pers * v + trop * away
|
| 104 |
+
p = self._clamp(p + self.dt * v)
|
| 105 |
+
dens = torch.exp(-self._cdist2(p, self.centers) / bw).mean(-1, keepdim=True)
|
| 106 |
+
feat = torch.cat([p, v, dens], dim=-1).flatten(2)
|
| 107 |
+
return self.read(feat)
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
class WaveCore(nn.Module):
|
| 111 |
+
"""Wheeler: K minisuperspace modes under a learnable LORENTZIAN DeWitt supermetric,
|
| 112 |
+
leapfrog wave dynamics, curvature potential. Exposes the Hamiltonian constraint."""
|
| 113 |
+
|
| 114 |
+
def __init__(self, cfg):
|
| 115 |
+
super().__init__()
|
| 116 |
+
self.K = cfg.wdw_modes
|
| 117 |
+
self.steps, self.dt = cfg.wdw_steps, cfg.wdw_dt
|
| 118 |
+
self.to_psi = nn.Linear(cfg.d_model, self.K)
|
| 119 |
+
self.to_pi = nn.Linear(cfg.d_model, self.K)
|
| 120 |
+
self.to_curv = nn.Linear(cfg.d_model, self.K)
|
| 121 |
+
self.ginv_raw = nn.Parameter(torch.zeros(self.K, self.K))
|
| 122 |
+
sig = torch.ones(self.K); sig[0] = -1.0
|
| 123 |
+
self.register_buffer("signature", sig)
|
| 124 |
+
self.log_lapse = nn.Parameter(torch.zeros(1))
|
| 125 |
+
self.read = nn.Sequential(nn.Linear(2 * self.K + 2, cfg.d_model), nn.GELU(),
|
| 126 |
+
nn.Linear(cfg.d_model, cfg.d_model))
|
| 127 |
+
self.last_H = None
|
| 128 |
+
|
| 129 |
+
def _supermetric(self):
|
| 130 |
+
A = 0.5 * (self.ginv_raw + self.ginv_raw.t())
|
| 131 |
+
return A + torch.diag(self.signature)
|
| 132 |
+
|
| 133 |
+
def forward(self, h, seed=None):
|
| 134 |
+
psi = self.to_psi(h)
|
| 135 |
+
if seed is not None:
|
| 136 |
+
psi = psi + seed[..., :self.K]
|
| 137 |
+
pi = self.to_pi(h)
|
| 138 |
+
r = self.to_curv(h)
|
| 139 |
+
Ginv = self._supermetric()
|
| 140 |
+
dt = self.dt * F.softplus(self.log_lapse).clamp(max=4.0)
|
| 141 |
+
pi = pi - 0.5 * dt * (r * psi)
|
| 142 |
+
for j in range(self.steps):
|
| 143 |
+
psi = psi + dt * torch.matmul(pi, Ginv.t())
|
| 144 |
+
pi = pi - (0.5 if j == self.steps - 1 else 1.0) * dt * (r * psi)
|
| 145 |
+
H = 0.5 * (torch.matmul(pi, Ginv.t()) * pi).sum(-1) + 0.5 * (r * psi * psi).sum(-1)
|
| 146 |
+
self.last_H = (H ** 2).mean() if self.training else None
|
| 147 |
+
t_comp, space_norm = psi[..., :1], psi[..., 1:].norm(dim=-1, keepdim=True)
|
| 148 |
+
return self.read(torch.cat([psi, pi, t_comp, space_norm], dim=-1))
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
# --------------------------------------------------------------------------------------
|
| 152 |
+
# the Chimera block: council of the three cores, combined by GRPO group-relative selection
|
| 153 |
+
# --------------------------------------------------------------------------------------
|
| 154 |
+
class ChimeraBlock(nn.Module):
|
| 155 |
+
"""Drop-in for QuazimotoBlock.forward(x, ring_ctl, phase_seed). Runs three physics
|
| 156 |
+
cores and GRPO-selects among them (critic -> group-relative advantage -> clip -> anchor),
|
| 157 |
+
behind one family gate. Exposes last_constraint (the wave core's <H^2>)."""
|
| 158 |
+
|
| 159 |
+
def __init__(self, cfg):
|
| 160 |
+
super().__init__()
|
| 161 |
+
self.cfg = cfg
|
| 162 |
+
self.norm = RMSNorm(cfg.d_model)
|
| 163 |
+
self.cores = nn.ModuleList([KuramotoCore(cfg), GrowthCore(cfg), WaveCore(cfg)])
|
| 164 |
+
self.G = len(self.cores)
|
| 165 |
+
self.critic = nn.Linear(cfg.d_model, 1) # scores each candidate (the verifier)
|
| 166 |
+
self.adv_clip, self.temp = cfg.chim_adv_clip, cfg.chim_select_temp
|
| 167 |
+
self.clip, self.anchor = cfg.chim_clip, cfg.chim_anchor
|
| 168 |
+
self.drop = nn.Dropout(cfg.dropout)
|
| 169 |
+
go = math.atanh(min(cfg.gate_init_open, 0.9)) if cfg.gate_init_open > 0 else 0.0
|
| 170 |
+
self.gate = nn.Parameter(torch.full((1,), go))
|
| 171 |
+
for m in self.modules():
|
| 172 |
+
if isinstance(m, nn.Linear):
|
| 173 |
+
nn.init.normal_(m.weight, std=0.02)
|
| 174 |
+
if m.bias is not None:
|
| 175 |
+
nn.init.zeros_(m.bias)
|
| 176 |
+
if cfg.use_fractal_phase_seed:
|
| 177 |
+
self.seed_gate = nn.Parameter(torch.zeros(1)) # zero-init -> no fractal seed at start
|
| 178 |
+
self.last_constraint = None
|
| 179 |
+
self.last_select = None # mean selection weights (for viz)
|
| 180 |
+
|
| 181 |
+
def forward(self, x, ring_ctl=None, phase_seed=None):
|
| 182 |
+
cfg = self.cfg
|
| 183 |
+
h = self.norm(x)
|
| 184 |
+
seed = None
|
| 185 |
+
if phase_seed is not None and cfg.use_fractal_phase_seed:
|
| 186 |
+
seed = torch.tanh(self.seed_gate) * phase_seed
|
| 187 |
+
|
| 188 |
+
cands = torch.stack([core(h, seed) for core in self.cores], dim=2) # [B,T,G,d]
|
| 189 |
+
|
| 190 |
+
# ---- GRPO group-relative selection over the G candidates ----
|
| 191 |
+
r = self.critic(cands).squeeze(-1) # [B,T,G] critic reward per core
|
| 192 |
+
mean = r.mean(-1, keepdim=True)
|
| 193 |
+
std = r.std(-1, keepdim=True)
|
| 194 |
+
adv = ((r - mean) / (std + 1e-4)).clamp(-self.adv_clip, self.adv_clip) # group-relative
|
| 195 |
+
sel = torch.softmax(adv / max(self.temp, 1e-6), dim=-1) # winners get mass
|
| 196 |
+
u = 1.0 / self.G # uniform = "old policy"
|
| 197 |
+
sel = sel.clamp(u * (1.0 - self.clip), u * (1.0 + self.clip)) # PPO clip around uniform
|
| 198 |
+
sel = sel / sel.sum(-1, keepdim=True)
|
| 199 |
+
w = (1.0 - self.anchor) * sel + self.anchor * u # KL-to-uniform anchor
|
| 200 |
+
out = (w.unsqueeze(-1) * cands).sum(2) # [B,T,d]
|
| 201 |
+
out = self.drop(soft_clamp(out * torch.tanh(self.gate), cfg.osc_bound))
|
| 202 |
+
|
| 203 |
+
# wave core's Hamiltonian constraint -> trunk aux loss (H Psi = 0 pressure)
|
| 204 |
+
wave = self.cores[2]
|
| 205 |
+
self.last_constraint = wave.last_H
|
| 206 |
+
self.last_select = w.detach().mean((0, 1)) if not self.training else None
|
| 207 |
+
|
| 208 |
+
rec = _viz.get_rec()
|
| 209 |
+
if rec is not None and rec.enabled: # live-viz: which law won this token
|
| 210 |
+
wsel = w[0, -1].tolist() # [G] selection weights
|
| 211 |
+
hval = float(wave.last_H) if wave.last_H is not None else 0.0
|
| 212 |
+
rec.log_ring([hval], wsel, wsel)
|
| 213 |
+
rec.log_quaz_norm(out[0, -1].norm().item())
|
| 214 |
+
rec.flush_spec()
|
| 215 |
+
return out
|