g2-model / turing_model.py
gijl's picture
Upload turing_model.py with huggingface_hub
f5f96a6 verified
Raw
History Blame Contribute Delete
6.17 kB
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)