|
|
| import math |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch.utils.checkpoint import checkpoint as grad_ckpt |
|
|
| def sinusoidal_pe(L: int, d: int, device) -> torch.Tensor: |
| pe = torch.zeros(1, L, d, device=device) |
| pos = torch.arange(L, device=device).unsqueeze(1).float() |
| div = torch.exp(torch.arange(0, d, 2, device=device).float() * (-math.log(10000.0) / d)) |
| pe[0, :, 0::2] = torch.sin(pos * div) |
| pe[0, :, 1::2] = torch.cos(pos * div) |
| return pe |
|
|
| def holo_encode(H: torch.Tensor, pe: torch.Tensor) -> torch.Tensor: |
| d = H.shape[-1] |
| Hf = torch.fft.rfft(F.normalize(H.float(), dim=-1), dim=-1) |
| Pf = torch.fft.rfft(F.normalize(pe.float(), dim=-1), dim=-1) |
| return torch.fft.irfft(Hf * Pf, n=d, dim=-1).to(H.dtype).sum(dim=1) |
|
|
| def holo_recall(query: torch.Tensor, trace: torch.Tensor) -> torch.Tensor: |
| d = query.shape[-1] |
| Q = torch.fft.rfft(F.normalize(query.float(), dim=-1), dim=-1) |
| T_ = torch.fft.rfft(trace.float(), dim=-1) |
| return torch.fft.irfft(Q.conj() * T_, n=d, dim=-1).to(query.dtype) |
|
|
| class ReactNet(nn.Module): |
| def __init__(self, d_in: int, d_hidden: int, d_out: int): |
| super().__init__() |
| self.main = nn.Sequential( |
| nn.Linear(d_in, d_hidden), nn.SiLU(), |
| nn.Linear(d_hidden, d_hidden), nn.SiLU(), |
| nn.Linear(d_hidden, d_out), |
| ) |
| self.skip = nn.Linear(d_in, d_out, bias=False) |
| for m in self.main: |
| if isinstance(m, nn.Linear): |
| nn.init.xavier_uniform_(m.weight, gain=0.1) |
| if m.bias is not None: |
| nn.init.zeros_(m.bias) |
| nn.init.xavier_uniform_(self.skip.weight, gain=0.05) |
| def forward(self, x): |
| return self.main(x) + 0.1 * self.skip(x) |
|
|
| class EnergyLandscape(nn.Module): |
| def __init__(self, d_model: int, n_attractors: int): |
| super().__init__() |
| self.attractors = nn.Parameter( |
| torch.randn(n_attractors, d_model) / math.sqrt(d_model)) |
| self.log_beta = nn.Parameter(torch.tensor(0.5)) |
| def gradient(self, h: torch.Tensor) -> torch.Tensor: |
| beta = torch.exp(self.log_beta).clamp(0.1, 5.0) |
| scores = beta * (h.float() @ self.attractors.float().T) |
| weights = torch.softmax(scores, dim=-1) |
| target = (weights @ self.attractors.float()).to(h.dtype) |
| return beta * (target - h) |
|
|
| class TuringDynamics(nn.Module): |
| def __init__(self, d_u, d_v, d_hidden, d_model, n_attractors): |
| super().__init__() |
| self.d_u = d_u |
| self.d_v = d_v |
| self.log_Du = nn.Parameter(torch.full((d_u,), -1.0)) |
| self.log_Dv = nn.Parameter(torch.full((d_v,), 1.0)) |
| self.react_U = ReactNet(d_model, d_hidden, d_u) |
| self.react_V = ReactNet(d_model, d_hidden, d_v) |
| self.energy = EnergyLandscape(d_model, n_attractors) |
| self.alpha_diff = nn.Parameter(torch.tensor(1.00)) |
| self.alpha_react = nn.Parameter(torch.tensor(0.15)) |
| self.alpha_global= nn.Parameter(torch.tensor(0.05)) |
| |
| @staticmethod |
| def causal_flow_1d(H: torch.Tensor) -> torch.Tensor: |
| left = F.pad(H[:, :-1, :], (0, 0, 1, 0)) |
| return left - H |
| |
| def forward(self, t: float, H: torch.Tensor) -> torch.Tensor: |
| B, L, _ = H.shape |
| D_u = torch.exp(self.log_Du).clamp(0.01, 5.0) |
| D_v = torch.exp(self.log_Dv).clamp(0.50, 20.0) |
| U = H[..., :self.d_u] |
| V = H[..., self.d_u:] |
| diff_U = D_u * self.causal_flow_1d(U) |
| diff_V = D_v * self.causal_flow_1d(V) |
| H_flat = H.reshape(B * L, -1) |
| rU = self.react_U(H_flat).reshape(B, L, self.d_u) |
| rV = self.react_V(H_flat).reshape(B, L, self.d_v) |
| H_mean = H.mean(dim=1) |
| hop_grad = self.energy.gradient(H_mean) |
| hop_U = hop_grad[:, :self.d_u].unsqueeze(1).expand(-1, L, -1) |
| hop_V = hop_grad[:, self.d_u:].unsqueeze(1).expand(-1, L, -1) |
| a_d = torch.abs(self.alpha_diff).clamp(0.01, 3.0) |
| a_r = torch.abs(self.alpha_react).clamp(0.001, 1.0) |
| a_g = torch.abs(self.alpha_global).clamp(0.001, 0.5) |
| dU = a_d * diff_U + a_r * rU + a_g * hop_U |
| dV = a_d * diff_V + a_r * rV + a_g * hop_V |
| return torch.cat([dU, dV], dim=-1).clamp(-10.0, 10.0) |
|
|
| def rk4_checkpointed(F_dyn, H0, T, steps): |
| dt = T / max(1, steps) |
| H = H0 |
| train = H.requires_grad or any(p.requires_grad for p in F_dyn.parameters()) |
| for i in range(steps): |
| t_i, dt_f = float(i) * dt, float(dt) |
| def make_step(tv, dv): |
| def step(H_in): |
| k1 = F_dyn(tv, H_in) |
| k2 = F_dyn(tv+dv/2, H_in + dv/2 * k1) |
| k3 = F_dyn(tv+dv/2, H_in + dv/2 * k2) |
| k4 = F_dyn(tv+dv, H_in + dv * k3) |
| return H_in + dv/6 * (k1 + 2*k2 + 2*k3 + k4) |
| return step |
| fn = make_step(t_i, dt_f) |
| H = grad_ckpt(fn, H, use_reentrant=False) if train else fn(H) |
| return H |
|
|
| class TuringLM(nn.Module): |
| def __init__(self, vocab_size, d_model, d_u, d_v, d_hidden, n_attractors, pad_id): |
| super().__init__() |
| self.d = d_model |
| self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=pad_id) |
| self.dynamics = TuringDynamics(d_u, d_v, d_hidden, d_model, n_attractors) |
| self.out_norm = nn.LayerNorm(d_model) |
| self.out_proj = nn.Linear(d_model, vocab_size, bias=False) |
| self.out_proj.weight = self.embedding.weight |
| nn.init.normal_(self.embedding.weight, std=0.02) |
| with torch.no_grad(): |
| self.embedding.weight[pad_id].zero_() |
| |
| def forward(self, input_ids, T, steps): |
| B, L = input_ids.shape |
| x = self.embedding(input_ids) |
| pe = sinusoidal_pe(L, self.d, x.device) |
| H0 = F.layer_norm(x + pe, [self.d]) |
| HT = rk4_checkpointed(self.dynamics, H0, T=T, steps=steps) |
| pe_T = sinusoidal_pe(L, self.d, HT.device) |
| trace = holo_encode(HT, pe_T.expand(B, -1, -1)) |
| query = HT.mean(dim=1) |
| readout = holo_recall(query, trace) |
| HT_ctx = HT + readout.unsqueeze(1) |
| out = self.out_norm(HT_ctx) |
| return self.out_proj(out) |
|
|