import json, math, time, random import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED = 158 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) def math_check(): rng = np.random.default_rng(SEED) d,m,q,K,r = 11,17,7,5,3 A = np.diag(np.tanh(rng.normal(size=q))) B = rng.normal(size=(q,d))*0.1 u = rng.normal(size=(80,d)); s=np.zeros(q); states=[] for x in u: s=A@s+B@x; states.append(s.copy()) spectral=float(max(abs(np.linalg.eigvals(A)))) z=rng.normal(size=q); norms=[] for _ in range(25): z=A@z; norms.append(np.linalg.norm(z)) U=rng.normal(size=(K,m,r)); V=rng.normal(size=(K,d,r)) ranks=[np.linalg.matrix_rank(U[k]@V[k].T) for k in range(K)] alpha=np.exp(rng.normal(size=K)); alpha/=alpha.sum() return {'spectral_radius':spectral, 'max_rank':int(max(ranks)), 'alpha_sum_error':float(abs(alpha.sum()-1)), 'homogeneous_norm_ratio_last':float(norms[-1]/norms[0]), 'stable':bool(spectral < 1 and max(ranks)<=r and abs(alpha.sum()-1)<1e-12)} class RoutedMLP(nn.Module): def __init__(self,d=24,m=48,q=8,K=4,r=3,routed=True): super().__init__(); self.d=d; self.m=m; self.q=q; self.K=K; self.r=r; self.routed=routed 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) if routed: self.a_raw=nn.Parameter(torch.full((q,), -1.4)); self.B=nn.Parameter(torch.randn(q,d)*.08) self.P=nn.Parameter(torch.randn(K,q)*.15); self.bias=nn.Parameter(torch.zeros(K)) self.U1=nn.Parameter(torch.randn(K,m,r)*.025); self.V1=nn.Parameter(torch.randn(K,d,r)*.025) self.U3=nn.Parameter(torch.randn(K,m,r)*.025); self.V3=nn.Parameter(torch.randn(K,d,r)*.025) self.U2=nn.Parameter(torch.randn(K,d,r)*.025); self.V2=nn.Parameter(torch.randn(K,m,r)*.025) def forward(self,x, return_stats=False): # x: batch, time, d; causal state recurrence is vectorized over time if not self.routed: h=F.silu(x@self.w1.T)*(x@self.w3.T); y=h@self.w2.T return (y,{}) if return_stats else y A=torch.tanh(self.a_raw); s=torch.zeros(x.shape[0],self.q,device=x.device); ys=[]; alphas=[]; deltas=[] for t in range(x.shape[1]): s=s@torch.diag(A)+x[:,t]@self.B.T alpha=F.softmax(s@self.P.T+self.bias,dim=-1); alphas.append(alpha) def dyn(base,U,V): # einsum forms weighted sum of rank-r matrices per batch item upd=torch.einsum('bk,kmr,kdr->bdm',alpha,U,V).transpose(1,2) return base.unsqueeze(0)+upd w1,w3,w2=dyn(self.w1,self.U1,self.V1),dyn(self.w3,self.U3,self.V3),dyn(self.w2,self.U2,self.V2) z1=torch.bmm(w1,x[:,t].unsqueeze(-1)).squeeze(-1); z3=torch.bmm(w3,x[:,t].unsqueeze(-1)).squeeze(-1) z=torch.bmm(w2,F.silu(z1).mul(z3).unsqueeze(-1)).squeeze(-1); ys.append(z) deltas.append((w1-self.w1).norm(dim=(1,2))) y=torch.stack(ys,1); stats={'alpha':torch.stack(alphas,1).detach(),'delta':torch.stack(deltas,1).detach()} return (y,stats) if return_stats else y def make_data(n=768,T=12,d=24): rng=np.random.default_rng(SEED+4); x=rng.normal(size=(n,T,d)).astype('float32'); y=np.empty_like(x) # Each sequence has a persistent hidden regime; token-local x alone is deliberately ambiguous. regime=rng.integers(0,2,size=n); signs=np.where(regime[:,None,None]==0,1.,-1.) y[:]=signs*x + .18*rng.normal(size=x.shape) return torch.tensor(x),torch.tensor(y) def train(routed, device, steps=260): torch.manual_seed(SEED); x,y=make_data(); x=x.to(device); y=y.to(device) model=RoutedMLP(routed=routed).to(device); opt=torch.optim.AdamW(model.parameters(),lr=3e-3) t0=time.time(); losses=[] for step in range(steps): 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())) with torch.no_grad(): pred,st=model(x[:192],return_stats=True); val=float(F.mse_loss(pred,y[:192]).cpu()) ent=float((-(st['alpha']*st['alpha'].clamp_min(1e-9).log()).sum(-1)).mean().cpu()) if routed else None delta=float(st['delta'].mean().cpu()) if routed else None 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())} def main(): check=math_check(); print('MATH',json.dumps(check)) 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) except Exception as e: 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'))} print('RESULTS',json.dumps(results)); print('SUMMARY',json.dumps({'math':check,'results':results})) if __name__=='__main__': main()