State-Routed Low-Rank MLP / experiment.py

Mechanism failed

Raw ⬇ ZIP
 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()