Spectral-Band Dual-Timescale Network / stage2_bench.py
Failed on benchmark
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report
9
10SEEDS = tuple(range(8))
11SWEEP_SEEDS = (0, 1, 2, 3)
12NTR, NTE, EPOCHS, BATCH = 1200, 400, 15, 128
13LR_GRID = [1e-3, 3e-3, 6e-3]
14
15
16def seed_all(seed):
17 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
18 if torch.cuda.is_available():
19 try: torch.cuda.manual_seed_all(seed)
20 except Exception: pass
21
22
23class SingleBandRNN(nn.Module):
24 """Single-timescale counterpart with the same dissipative state update."""
25 def __init__(self, hidden=64, gamma=0.4, out_dim=1):
26 super().__init__()
27 self.hidden = hidden
28 self.inp = nn.Linear(3, hidden)
29 self.L = nn.Parameter(0.05 * torch.randn(hidden, hidden))
30 self.alpha_raw = nn.Parameter(torch.tensor(float(np.log(np.expm1(gamma)))))
31 self.head = nn.Linear(hidden, out_dim)
32 self.dt = 0.1
33
34 def J(self):
35 a = torch.nn.functional.softplus(self.alpha_raw) + 1e-4
36 return -(self.L.T @ self.L) - a * torch.eye(self.hidden, device=self.L.device)
37
38 def forward(self, x, return_states=False):
39 z = x.view(x.shape[0], -1, 3)
40 h = torch.zeros(x.shape[0], self.hidden, device=x.device)
41 states=[]; J=self.J()
42 for k in range(z.shape[1]):
43 h = h + self.dt * (h @ J.T + self.inp(z[:, k]))
44 h = torch.tanh(h)
45 states.append(h)
46 if return_states: return self.head(h), torch.stack(states, 1)
47 return self.head(h)
48
49
50class DualBandRNN(nn.Module):
51 """Two dissipative bands plus bounded explicit cross-band exchange."""
52 def __init__(self, width=32, gamma_f=4.0, gamma_s=0.4, out_dim=1):
53 super().__init__(); self.width=width; self.dt=.1
54 self.bf=nn.Linear(3,width); self.bs=nn.Linear(3,width)
55 self.Lf=nn.Parameter(.05*torch.randn(width,width)); self.Ls=nn.Parameter(.05*torch.randn(width,width))
56 self.af=nn.Parameter(torch.tensor(float(np.log(np.expm1(gamma_f)))))
57 self.a_s=nn.Parameter(torch.tensor(float(np.log(np.expm1(gamma_s)))))
58 self.Efs=nn.Parameter(.03*torch.randn(width,width)); self.Esf=nn.Parameter(.03*torch.randn(width,width))
59 self.head=nn.Linear(2*width,out_dim)
60
61 def Js(self):
62 eye=torch.eye(self.width,device=self.Lf.device)
63 jf=-(self.Lf.T@self.Lf)-(torch.nn.functional.softplus(self.af)+1e-4)*eye
64 js=-(self.Ls.T@self.Ls)-(torch.nn.functional.softplus(self.a_s)+1e-4)*eye
65 # bounded exchange prevents uncontrolled growth while retaining learned coupling
66 ef=.15*torch.tanh(self.Efs); es=.15*torch.tanh(self.Esf)
67 return jf,js,ef,es
68
69 def forward(self,x,return_states=False):
70 z=x.view(x.shape[0],-1,3); hf=torch.zeros(x.shape[0],self.width,device=x.device); hs=hf.clone()
71 jf,js,ef,es=self.Js(); traces=[]
72 for k in range(z.shape[1]):
73 hf=hf+self.dt*(hf@jf.T+hs@ef.T+self.bf(z[:,k])); hs=hs+self.dt*(hs@js.T+hf@es.T+self.bs(z[:,k]))
74 hf=torch.tanh(hf); hs=torch.tanh(hs); traces.append(torch.cat([hf,hs],1))
75 h=torch.cat([hf,hs],1)
76 if return_states:return self.head(h),torch.stack(traces,1)
77 return self.head(h)
78
79
80def train_one(kind, seed, lr, gamma=0.4):
81 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
82 model=SingleBandRNN(gamma=gamma) if kind=='baseline' else DualBandRNN()
83 net,metric,_=train_model(model,ds,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None)
84 return float(metric), net, ds
85
86
87def fn(kind, cfg):
88 def run(seed): return train_one(kind,seed,cfg['lr'],cfg.get('gamma',.4))[0]
89 return run
90
91
92def signature():
93 pred=[]; obs=[]
94 for seed in SEEDS:
95 _,net,ds=train_one('idea',seed,3e-3)
96 net.eval(); dev=next(net.parameters()).device
97 x=torch.zeros(1,96,device=dev); x[0,:3]=ds['xtr'][0,:3].to(dev)
98 with torch.no_grad(): _,st=net(x,return_states=True)
99 st=st[0].cpu().numpy(); wf=np.linalg.norm(st[:,:net.width],axis=1)+1e-8; ws=np.linalg.norm(st[:,net.width:],axis=1)+1e-8
100 # observed discrete decay fit after the impulse, excluding first point
101 tt=np.arange(len(wf))*net.dt; slf=np.polyfit(tt[1:],np.log(wf[1:]),1)[0]; sls=np.polyfit(tt[1:],np.log(ws[1:]),1)[0]
102 jf,js,_,_=net.Js(); pf=float(torch.linalg.eigvalsh((jf+jf.T)/2).max().detach().cpu()); ps=float(torch.linalg.eigvalsh((js+js.T)/2).max().detach().cpu())
103 pred.append([abs(pf),abs(ps)]); obs.append([max(0.,-slf),max(0.,-sls)])
104 p=np.mean(pred,0); o=np.mean(obs,0)
105 rel=np.abs(o-p)/np.maximum(p,1e-6)
106 return {'prediction':'trained fast/slow modal decay rates agree within 20%', 'predicted_rates':p.tolist(),'observed_rates':o.tolist(),'relative_errors':rel.tolist(),'confirmed':bool(np.all(rel<.2))}
107
108
109def main():
110 base_grid=[{'lr':lr,'gamma':g} for lr in LR_GRID for g in (.4,1.0)]
111 base=sweep_baseline(lambda c:fn('baseline',c),base_grid,seeds=SWEEP_SEEDS)
112 idea_cfgs=[{'lr':lr} for lr in LR_GRID]
113 # Evaluate the prescribed three-setting idea sweep on all paired seeds.
114 idea_runs=[]
115 for c in idea_cfgs:
116 r=evaluate(fn('idea',c),SEEDS); idea_runs.append((r,c))
117 idea,bestcfg=min(idea_runs,key=lambda q:q[0]['mean'])
118 rep=make_report('dynamics','rnn_small',base,idea,extra=signature())
119 rep['idea_sweep']=[{'cfg':c,'mean':r['mean'],'per_seed':r['per_seed']} for r,c in idea_runs]
120 rep['protocol_note']='Matched dynamics task; baseline and idea share 64-dimensional dissipative recurrent state, input/output, optimizer, epochs, batch, and all learning-rate values. Only one versus two relaxation bands and explicit exchange differ.'
121 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
122 print(json.dumps(rep,indent=2))
123
124if __name__=='__main__': main()