State-Routed Low-Rank MLP / experiment.py
Mechanism failed
1import json, math, time, random
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7SEED = 158
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9
10def math_check():
11 rng = np.random.default_rng(SEED)
12 d,m,q,K,r = 11,17,7,5,3
13 A = np.diag(np.tanh(rng.normal(size=q)))
14 B = rng.normal(size=(q,d))*0.1
15 u = rng.normal(size=(80,d)); s=np.zeros(q); states=[]
16 for x in u:
17 s=A@s+B@x; states.append(s.copy())
18 spectral=float(max(abs(np.linalg.eigvals(A))))
19 z=rng.normal(size=q); norms=[]
20 for _ in range(25): z=A@z; norms.append(np.linalg.norm(z))
21 U=rng.normal(size=(K,m,r)); V=rng.normal(size=(K,d,r))
22 ranks=[np.linalg.matrix_rank(U[k]@V[k].T) for k in range(K)]
23 alpha=np.exp(rng.normal(size=K)); alpha/=alpha.sum()
24 return {'spectral_radius':spectral, 'max_rank':int(max(ranks)),
25 'alpha_sum_error':float(abs(alpha.sum()-1)),
26 'homogeneous_norm_ratio_last':float(norms[-1]/norms[0]),
27 'stable':bool(spectral < 1 and max(ranks)<=r and abs(alpha.sum()-1)<1e-12)}
28
29class RoutedMLP(nn.Module):
30 def __init__(self,d=24,m=48,q=8,K=4,r=3,routed=True):
31 super().__init__(); self.d=d; self.m=m; self.q=q; self.K=K; self.r=r; self.routed=routed
32 self.w1=nn.Parameter(torch.randn(m,d)*.12); self.w3=nn.Parameter(torch.randn(m,d)*.12); self.w2=nn.Parameter(torch.randn(d,m)*.12)
33 if routed:
34 self.a_raw=nn.Parameter(torch.full((q,), -1.4)); self.B=nn.Parameter(torch.randn(q,d)*.08)
35 self.P=nn.Parameter(torch.randn(K,q)*.15); self.bias=nn.Parameter(torch.zeros(K))
36 self.U1=nn.Parameter(torch.randn(K,m,r)*.025); self.V1=nn.Parameter(torch.randn(K,d,r)*.025)
37 self.U3=nn.Parameter(torch.randn(K,m,r)*.025); self.V3=nn.Parameter(torch.randn(K,d,r)*.025)
38 self.U2=nn.Parameter(torch.randn(K,d,r)*.025); self.V2=nn.Parameter(torch.randn(K,m,r)*.025)
39 def forward(self,x, return_stats=False):
40 # x: batch, time, d; causal state recurrence is vectorized over time
41 if not self.routed:
42 h=F.silu(x@self.w1.T)*(x@self.w3.T); y=h@self.w2.T
43 return (y,{}) if return_stats else y
44 A=torch.tanh(self.a_raw); s=torch.zeros(x.shape[0],self.q,device=x.device); ys=[]; alphas=[]; deltas=[]
45 for t in range(x.shape[1]):
46 s=s@torch.diag(A)+x[:,t]@self.B.T
47 alpha=F.softmax(s@self.P.T+self.bias,dim=-1); alphas.append(alpha)
48 def dyn(base,U,V):
49 # einsum forms weighted sum of rank-r matrices per batch item
50 upd=torch.einsum('bk,kmr,kdr->bdm',alpha,U,V).transpose(1,2)
51 return base.unsqueeze(0)+upd
52 w1,w3,w2=dyn(self.w1,self.U1,self.V1),dyn(self.w3,self.U3,self.V3),dyn(self.w2,self.U2,self.V2)
53 z1=torch.bmm(w1,x[:,t].unsqueeze(-1)).squeeze(-1); z3=torch.bmm(w3,x[:,t].unsqueeze(-1)).squeeze(-1)
54 z=torch.bmm(w2,F.silu(z1).mul(z3).unsqueeze(-1)).squeeze(-1); ys.append(z)
55 deltas.append((w1-self.w1).norm(dim=(1,2)))
56 y=torch.stack(ys,1); stats={'alpha':torch.stack(alphas,1).detach(),'delta':torch.stack(deltas,1).detach()}
57 return (y,stats) if return_stats else y
58
59def make_data(n=768,T=12,d=24):
60 rng=np.random.default_rng(SEED+4); x=rng.normal(size=(n,T,d)).astype('float32'); y=np.empty_like(x)
61 # Each sequence has a persistent hidden regime; token-local x alone is deliberately ambiguous.
62 regime=rng.integers(0,2,size=n); signs=np.where(regime[:,None,None]==0,1.,-1.)
63 y[:]=signs*x + .18*rng.normal(size=x.shape)
64 return torch.tensor(x),torch.tensor(y)
65
66def train(routed, device, steps=260):
67 torch.manual_seed(SEED); x,y=make_data(); x=x.to(device); y=y.to(device)
68 model=RoutedMLP(routed=routed).to(device); opt=torch.optim.AdamW(model.parameters(),lr=3e-3)
69 t0=time.time(); losses=[]
70 for step in range(steps):
71 ix=((torch.arange(64,device=device)+step*64)%len(x)); pred=model(x[ix]); loss=F.mse_loss(pred,y[ix]); opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),1.0); opt.step(); losses.append(float(loss.detach().cpu()))
72 with torch.no_grad():
73 pred,st=model(x[:192],return_stats=True); val=float(F.mse_loss(pred,y[:192]).cpu())
74 ent=float((-(st['alpha']*st['alpha'].clamp_min(1e-9).log()).sum(-1)).mean().cpu()) if routed else None
75 delta=float(st['delta'].mean().cpu()) if routed else None
76 return {'final_train_mse':losses[-1],'validation_mse':val,'seconds':time.time()-t0,'alpha_entropy':ent,'mean_update_frobenius':delta,'params':sum(p.numel() for p in model.parameters())}
77
78def main():
79 check=math_check(); print('MATH',json.dumps(check))
80 try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu'); results={'device':str(device)}; results['baseline']=train(False,device); results['idea']=train(True,device)
81 except Exception as e:
82 print('CUDA/accelerator failed, retrying CPU:',repr(e)); results={'device':'cpu-fallback','error':repr(e),'baseline':train(False,torch.device('cpu')),'idea':train(True,torch.device('cpu'))}
83 print('RESULTS',json.dumps(results)); print('SUMMARY',json.dumps({'math':check,'results':results}))
84if __name__=='__main__': main()