Averaged Contractive State-Space Network / bench_stage2.py
Failed on benchmark
1import json, sys, time
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
8
9TRACK='dynamics'; MODEL='rnn_small'; SEEDS=tuple(range(8)); SWEEP_SEEDS=(0,1,2,3)
10# The union is shared by baseline and idea; no hidden hyperparameter is used.
11GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':1e-2}]
12EPOCHS=5; NTRAIN=400; NTEST=160
13
14class AveragedContractiveRNN(nn.Module):
15 """Discrete Euler sample of h'=F(t/eps,h,x), using quadrature over phase.
16 The recurrent matrix is spectrally bounded and the linear drift is strictly
17 contractive; phase averaging is the sole architectural intervention.
18 """
19 def __init__(self, hidden=16, phases=3, alpha=2.0, q=.30, dt=.05):
20 super().__init__(); self.hidden=hidden; self.phases=phases
21 self.alpha=alpha; self.q=q; self.dt=dt
22 self.w_raw=nn.Parameter(torch.randn(hidden,hidden)*.05)
23 self.inp=nn.Linear(3,hidden); self.bias=nn.Parameter(torch.zeros(hidden))
24 self.head=nn.Linear(hidden,1)
25 pattern=torch.ones(hidden); pattern[1::2]=-1
26 self.register_buffer('pattern',pattern)
27 def w_bound(self):
28 # A differentiable, uniform spectral bound, so tanh Jacobian <= ||W||.
29 return self.w_raw / (torch.linalg.matrix_norm(self.w_raw,2)+1e-6) * .35
30 def vector_field(self,h,x,phase):
31 W=self.w_bound(); a=-self.alpha + self.q*torch.sin(2*torch.pi*torch.as_tensor(phase, device=h.device, dtype=h.dtype))*self.pattern
32 return a*h + torch.tanh(h@W.T + self.inp(x) + self.bias)
33 def rollout(self,x, averaged=True, eps=1/8, return_states=False):
34 z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device)
35 states=[]
36 for k in range(z.shape[1]):
37 if averaged:
38 ps=torch.arange(self.phases,device=x.device,dtype=x.dtype)/self.phases
39 f=sum(self.vector_field(h,z[:,k],p) for p in ps)/self.phases
40 else:
41 phase=(k*self.dt/eps)
42 f=self.vector_field(h,z[:,k],phase)
43 h=h+self.dt*f; states.append(h)
44 out=self.head(h)
45 return (out,torch.stack(states,1)) if return_states else out
46 def forward(self,x): return self.rollout(x, averaged=True)
47
48def seed_all(seed):
49 np.random.seed(seed); torch.manual_seed(seed)
50 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
51
52def ds(seed): return get_dataset(TRACK,seed,n_train=NTRAIN,n_test=NTEST)
53
54def baseline_fn(cfg):
55 def run(seed):
56 seed_all(seed); d=ds(seed); m=make_model(MODEL,d['input_shape'],d['out_dim'])
57 _,metric,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=128)
58 return metric
59 return run
60
61def idea_fn(cfg, retain=False):
62 def run(seed):
63 seed_all(seed); d=ds(seed); m=AveragedContractiveRNN()
64 _,metric,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=128)
65 return metric
66 return run
67
68def signature():
69 # Re-test the stage-1 prediction on trained benchmark models, not a toy graph.
70 seed=0; seed_all(seed); d=ds(seed); m=AveragedContractiveRNN()
71 _,_,_=train_model(m,d,epochs=EPOCHS,lr=3e-3,batch=128)
72 device=next(m.parameters()).device; m.eval(); x=d['xte'][:96].to(device)
73 rows=[]
74 with torch.no_grad():
75 _,ha=m.rollout(x,averaged=True,return_states=True)
76 for eps in [0.5,.25,.125,.0625]:
77 _,hf=m.rollout(x,averaged=False,eps=eps,return_states=True)
78 e=torch.sqrt(((hf-ha)**2).mean()).item()
79 rows.append({'eps':eps,'rms_state_error':e})
80 slope=float(np.polyfit(np.log([r['eps'] for r in rows]),np.log([r['rms_state_error']+1e-12 for r in rows]),1)[0])
81 # Empirical matrix measure of the trained field at sampled hidden/input points.
82 mus=[]
83 for i in range(12):
84 h=torch.randn(1,m.hidden,device=device,requires_grad=True); u=x[i:i+1].view(1,-1,3)[:,0]
85 phase=float(i)/12
86 J=torch.autograd.functional.jacobian(lambda hh:m.vector_field(hh,u,phase),h).squeeze(0).squeeze(1)
87 mus.append(float(torch.linalg.eigvalsh((J+J.T)/2).max().detach().cpu()))
88 max_mu=max(mus); predicted_bound=-m.alpha+m.q+.35
89 return {'prediction':'trained averaged/fast state error decreases with eps; mu2 is negative',
90 'rows':rows,'observed_loglog_slope':slope,'predicted_error_slope':1.0,
91 'predicted_mu2_upper_bound':predicted_bound,'observed_max_mu2':max_mu,
92 'confirmed':bool(slope>0.5 and max_mu<0)}
93
94def main():
95 t=time.time()
96 # Baseline sweep on four seeds, then canonical full eight-seed reevaluation.
97 base=sweep_baseline(baseline_fn,GRID,seeds=SWEEP_SEEDS)
98 # Explicitly evaluate idea at best baseline lr and two nearby/shared settings.
99 idea_trials=[]
100 for cfg in GRID:
101 r=__import__('bench').evaluate(idea_fn(cfg),seeds=SEEDS)
102 idea_trials.append({'cfg':cfg,'result':r})
103 best=min(idea_trials,key=lambda z:z['result']['mean'])
104 rep=make_report(TRACK,MODEL,base,best['result'],signature())
105 rep['idea_sweep']=idea_trials; rep['runtime_seconds']=time.time()-t
106 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
107 print(json.dumps(rep,indent=2))
108if __name__=='__main__': main()