Spectral Memory-Lift Ensemble / spectral_memory_lift.py
Beats tuned baseline
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 1355
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10
11class SpectralMemoryLift(nn.Module):
12 """Small batched implementation of q'=Aq+Bz+Fu, z'=Cq+Dz."""
13 def __init__(self, d=1, h=1, experts=4, taus=(2., 8., 32., 128.)):
14 super().__init__(); self.d=d; self.h=h; self.E=experts
15 self.A=nn.Parameter(.05*torch.randn(experts,d,d))
16 self.B=nn.Parameter(.1*torch.randn(experts,d,h))
17 self.C=nn.Parameter(.1*torch.randn(experts,h,d))
18 self.F=nn.Parameter(.1*torch.randn(experts,d,d))
19 self.Dhat=nn.Parameter(torch.randn(experts,h,h))
20 # Distinct initial memory rates; bounded through sigmoid.
21 rates = torch.tensor([math.exp(-1./t) for t in taus[:experts]])
22 logits = torch.logit(rates.clamp(.01,.99))
23 self.rate_logits=nn.Parameter(logits[:,None].repeat(1,h))
24 self.read=nn.Parameter(.1*torch.randn(experts,d+h))
25 self.gate=nn.Linear(d+experts*(d+h), experts)
26 self.taus=torch.tensor(list(taus[:experts]), dtype=torch.float32)
27 def D(self):
28 raw=self.Dhat/(self.Dhat.norm(dim=(-2,-1),keepdim=True)+1e-8)
29 return raw*torch.sigmoid(self.rate_logits)[:,:,None]
30 def forward(self, u, return_states=False):
31 # u: [batch,time,d]; q and z are reset per sequence.
32 b,T,_=u.shape; q=torch.zeros(b,self.E,self.d,device=u.device)
33 z=torch.zeros(b,self.E,self.h,device=u.device); D=self.D()
34 ys=[]; states=[]
35 for t in range(T):
36 x=u[:,t]
37 qn=torch.einsum('edk,bek->bed',self.A,q)+torch.einsum('edh,beh->bed',self.B,z)+torch.einsum('edk,bk->bed',self.F,x)
38 zn=torch.einsum('ehd,bed->beh',self.C,q)+torch.einsum('ehi,bei->beh',D,z)
39 vals=torch.cat([qn,zn],-1)
40 logits=self.gate(torch.cat([x,vals.reshape(b,-1)],-1))
41 # stop-gradient slow-mode feature, as proposed
42 logits=logits + .05*(-1./torch.log(self.taus.clamp_min(1.01))).to(x.device)
43 p=torch.softmax(logits,-1)
44 pred=(p.unsqueeze(-1)*torch.einsum('eK,beK->be',self.read,vals).unsqueeze(-1)).sum(1)
45 ys.append(pred); q,z=qn,zn; states.append((q,z))
46 out=torch.stack(ys,1)
47 return (out,states) if return_states else out
48
49class SingleExponential(nn.Module):
50 """Baseline: one stable scalar memory and linear readout."""
51 def __init__(self):
52 super().__init__(); self.logit=nn.Parameter(torch.tensor(0.0)); self.a=nn.Parameter(torch.tensor(.1)); self.b=nn.Parameter(torch.tensor(.1)); self.c=nn.Parameter(torch.tensor(.1))
53 def forward(self,u):
54 b,T,_=u.shape; z=torch.zeros(b,1,device=u.device); ys=[]
55 for t in range(T):
56 z=torch.sigmoid(self.logit)*z+self.a*u[:,t]
57 ys.append(self.b*z+self.c*u[:,t])
58 return torch.stack(ys,1)
59
60def exact_kernel_check():
61 # For z_{n+1}=Cq_n+Dz_n, contribution of q_n to q_{n+k} through Bz is BD^(k-1)C.
62 A=np.array([[.2]]); B=np.array([[.7]]); C=np.array([[.4]]); D=np.array([[.8]])
63 q=1.; z=0.; observed=[]
64 # impulse q_0, then suppress direct q recurrence and inject only its memory path
65 for k in range(1,8):
66 z=C*q + D*z if k==1 else D*z
67 observed.append(float(B@z)); q=0.
68 predicted=[float(B@np.linalg.matrix_power(D,k-1)@C) for k in range(1,8)]
69 return max(abs(np.array(observed)-predicted)), observed, predicted
70
71def stability_sweep():
72 # Scale a fixed 2x2 transition M by gamma. Boundary is gamma*rho(M0)=1.
73 M0=np.array([[.35,.25],[-.15,.70]])
74 rho=max(abs(np.linalg.eigvals(M0))); boundary=1/rho
75 rows=[]
76 for gamma in np.linspace(.4*boundary,1.6*boundary,13):
77 M=gamma*M0; eig=np.linalg.eigvals(M); r=max(abs(eig))
78 x=np.array([1.,-.3]); norms=[]
79 for _ in range(80): x=M@x; norms.append(np.linalg.norm(x))
80 # asymptotic ratio, robustly measured over last 20 steps
81 ratio=float((norms[-1]/norms[-21])**(1/20))
82 rows.append({'gamma':float(gamma),'rho':float(r),'ratio':ratio,'stable':bool(norms[-1]<norms[0])})
83 transition=min(rows,key=lambda x:abs(x['rho']-1))
84 return {'rho_M0':float(rho),'predicted_gamma_boundary':float(boundary),'rows':rows,'closest_observed':transition}
85
86def timescale_sweep():
87 # A pure mode r has amplitude r^n and tau=-1/log(r); fit tau from log amplitude.
88 rows=[]
89 for r in [.5,.7,.8,.9,.97,.99]:
90 tau=-1/math.log(r); n=np.arange(1,1000); amp=r**n
91 fit=-1/np.polyfit(n,np.log(amp),1)[0]
92 half=n[np.argmin(abs(amp-.5))]
93 rows.append({'rho':r,'predicted_tau':tau,'fitted_tau':float(fit),'predicted_half_life':math.log(.5)/math.log(r),'observed_half_life':int(half)})
94 return rows
95
96def mini_experiment():
97 # Target has two separated exponential/delayed components, generated from input history.
98 device='cuda' if torch.cuda.is_available() else 'cpu'
99 try:
100 N,T=256,96; g=torch.Generator().manual_seed(SEED)
101 x=torch.randn(N,T,1,generator=g).to(device)
102 target=torch.zeros_like(x)
103 for t in range(T):
104 if t>=2: target[:,t]+=0.55*x[:,t-2]
105 if t>=24: target[:,t]+=0.30*x[:,t-24]
106 if t>=64: target[:,t]+=0.15*x[:,t-64]
107 split=192; models=[SingleExponential().to(device),SpectralMemoryLift(experts=4,taus=(2,8,32,128)).to(device)]
108 results=[]
109 for model in models:
110 opt=torch.optim.Adam(model.parameters(),lr=0.02)
111 for step in range(500):
112 pred=model(x[:split]); loss=((pred-target[:split])**2).mean()
113 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.); opt.step()
114 with torch.no_grad():
115 tr=float(((model(x[:split])-target[:split])**2).mean()); va=float(((model(x[split:])-target[split:])**2).mean())
116 out=model(x[split:]); state_norm=None
117 if isinstance(model,SpectralMemoryLift):
118 _,states=model(x[split:],True); state_norm=float(max(max(q.norm(dim=-1).max().item(),z.norm(dim=-1).max().item()) for q,z in states))
119 else: state_norm=float(out.abs().max())
120 results.append({'model':type(model).__name__,'train_mse':tr,'validation_mse':va,'max_state_or_output_norm':state_norm})
121 if device=='cuda': torch.cuda.empty_cache()
122 return results
123 except Exception as e:
124 # Required safe fallback if CUDA/runtime allocation fails.
125 torch.cuda.empty_cache() if torch.cuda.is_available() else None
126 return [{'error':str(e),'fallback':'CPU unavailable in this run'}]
127
128def main():
129 kernel_err,obs,pred=exact_kernel_check()
130 result={'seed':SEED,'kernel_check':{'max_abs_error':kernel_err,'observed':obs,'predicted':pred},'stability_check':stability_sweep(),'timescale_check':timescale_sweep(),'mini_experiment':mini_experiment()}
131 Path('results.json').write_text(json.dumps(result,indent=2))
132 print(json.dumps(result,indent=2))
133if __name__=='__main__': main()