Quazim0t0 commited on
Commit
8abb25c
·
verified ·
1 Parent(s): 0b58465

Upload chimera.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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