Spectral Memory-Lift Ensemble / bench_experiment.py
Beats tuned baseline
1import sys, json, math, random
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
7
8EPOCHS = 12
9NTRAIN, NTEST = 1000, 300
10LRS = [1e-3, 3e-3, 1e-2]
11SEEDS = tuple(range(8))
12
13
14def seed_all(seed):
15 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
16 if torch.cuda.is_available():
17 try: torch.cuda.manual_seed_all(seed)
18 except Exception: pass
19
20
21class GRUBaseline(nn.Module):
22 def __init__(self, hidden=32):
23 super().__init__(); self.rnn=nn.GRU(3, hidden, batch_first=True); self.head=nn.Linear(hidden,1); self.no_cudnn=False
24 def forward(self,x):
25 seq=x.view(x.shape[0],-1,3)
26 try: _,h=self.rnn(seq)
27 except RuntimeError:
28 self.no_cudnn=True
29 if self.no_cudnn:
30 old=torch.backends.cudnn.enabled; torch.backends.cudnn.enabled=False
31 try: _,h=self.rnn(seq)
32 finally: torch.backends.cudnn.enabled=old
33 return self.head(h[-1])
34
35
36class SpectralMemoryLift(nn.Module):
37 def __init__(self, qdim=32, hdim=8, experts=4):
38 super().__init__(); self.qdim=qdim; self.hdim=hdim; self.E=experts
39 self.inp=nn.Linear(3,qdim)
40 self.A=nn.Parameter(.04*torch.randn(experts,qdim,qdim))
41 self.B=nn.Parameter(.06*torch.randn(experts,qdim,hdim))
42 self.C=nn.Parameter(.06*torch.randn(experts,hdim,qdim))
43 self.F=nn.Parameter(.04*torch.randn(experts,qdim,qdim))
44 self.Dhat=nn.Parameter(torch.randn(experts,hdim,hdim))
45 rates=torch.tensor([math.exp(-1/t) for t in (2.,8.,32.,128.)])
46 self.rate_logits=nn.Parameter(torch.logit(rates)[:,None].repeat(1,hdim))
47 self.read=nn.Linear(qdim+hdim,1)
48 self.gate=nn.Sequential(nn.Linear(3+experts*(qdim+hdim),32),nn.Tanh(),nn.Linear(32,experts))
49 def D(self):
50 raw=self.Dhat/(self.Dhat.flatten(1).norm(dim=1,keepdim=True).view(self.E,1,1)+1e-8)
51 return raw*torch.sigmoid(self.rate_logits).unsqueeze(-1)
52 def forward(self,x,return_states=False):
53 u=self.inp(x.view(x.shape[0],-1,3)); b,t,_=u.shape
54 q=torch.zeros(b,self.E,self.qdim,device=x.device); z=torch.zeros(b,self.E,self.hdim,device=x.device); D=self.D(); outs=[]; norms=[]
55 rawx=x.view(b,t,3)
56 for k in range(t):
57 uk=u[:,k]
58 qn=torch.einsum('eij,bej->bei',self.A,q)+torch.einsum('eih,beh->bei',self.B,z)+torch.einsum('eij,bj->bei',self.F,uk)
59 zn=torch.einsum('ehi,bei->beh',self.C,q)+torch.einsum('ehj,bej->beh',D,z)
60 vals=torch.cat([qn,zn],-1); logits=self.gate(torch.cat([rawx[:,k],vals.reshape(b,-1)],-1))
61 p=torch.softmax(logits,-1); ev=self.read(vals).squeeze(-1); outs.append((p*ev).sum(-1)); norms.append(vals.norm(dim=-1).mean())
62 q,z=qn,zn
63 y=torch.stack(outs,1)[:,-1:]
64 if return_states: return y, float(torch.stack(norms).max().detach().cpu())
65 return y
66
67
68def train_one(kind, seed, lr):
69 seed_all(seed); ds=get_dataset('dynamics',seed,NTRAIN,NTEST)
70 model=GRUBaseline() if kind=='baseline' else SpectralMemoryLift()
71 _,metric,hist=train_model(model,ds,epochs=EPOCHS,lr=lr,batch=128)
72 return metric
73
74
75def signature(seed, lr):
76 seed_all(seed); ds=get_dataset('dynamics',seed,NTRAIN,NTEST); model=SpectralMemoryLift(); net,_,_=train_model(model,ds,epochs=EPOCHS,lr=lr,batch=128)
77 net.eval(); device=next(net.parameters()).device
78 x=ds['xte'][:128].to(device)
79 with torch.no_grad():
80 _, observed=net(x,True)
81 D=net.D().detach().cpu().numpy(); predicted=float(np.max(np.abs(np.linalg.eigvals(D)),axis=1).max())
82 return {'predicted_slowest_memory_mode':predicted,'observed_max_state_norm':observed,'bounded_observed_state':bool(observed < 100.),'confirmed':bool(predicted < 1.0 and observed < 100.)}
83
84
85def main():
86 # Baseline sweep and idea sweep use identical union of learning rates.
87 base=sweep_baseline(lambda cfg: lambda seed: train_one('baseline',seed,cfg['lr']), [{'lr':lr} for lr in LRS])
88 # Idea receives the same three-value learning-rate sweep; select on SWEEP_SEEDS.
89 idea_trials=[]; idea_best=None; idea_best_mean=float('inf')
90 for lr in LRS:
91 r=evaluate(lambda seed, lr=lr: train_one('idea',seed,lr), seeds=(0,1,2,3))
92 idea_trials.append({'cfg': {'lr': lr}, 'mean': r['mean']})
93 if r['mean'] < idea_best_mean:
94 idea_best_mean=r['mean']; idea_best={'lr': lr}
95 idea=evaluate(lambda seed: train_one('idea',seed,idea_best['lr']), seeds=SEEDS)
96 rep=make_report('dynamics','rnn_small',base,idea,extra=signature(0,idea_best['lr']))
97 rep['idea_sweep']=idea_trials
98 rep['idea_best_cfg']=idea_best
99 rep['track_justification']='Dynamics is structurally matched: the task is actuated pendulum forecasting and the intervention is a stable recurrent memory state.'
100 rep['search_space_parity']={'baseline_grid':LRS,'idea_grid':LRS,'selected_idea_lr':idea_best['lr']}
101 rep['parameterization']={'baseline_params':sum(p.numel() for p in GRUBaseline().parameters()),'idea_params':sum(p.numel() for p in SpectralMemoryLift().parameters()),'epochs':EPOCHS,'n_train':NTRAIN,'n_test':NTEST}
102 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
103 print(json.dumps(rep,indent=2))
104
105if __name__=='__main__': main()