Fourier Replay-Mode Stabilizer / stage2_bench.py
Failed on benchmark
1import sys, json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, sweep_baseline, evaluate, make_report
8
9SEED0=1729
10DEVICE='cuda' if torch.cuda.is_available() else 'cpu'
11N=16; TAU_R=1.0; TAU_D=0.10
12MODES=(1,2,3,-1,-2,-3)
13
14def seed_all(s):
15 random.seed(s); np.random.seed(s); torch.manual_seed(s)
16 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
17
18def root_np(what,a=1.,tr=1.,td=TAU_D,it=20):
19 z=complex((-1+a*what)/tr)
20 for _ in range(it):
21 e=np.exp(-z*td); f=tr*z+1-a*what*e; df=tr+a*what*td*e
22 if abs(df)<1e-12: break
23 z-=f/df
24 return z
25
26def sanity():
27 # FFT formula and characteristic-root prediction against a direct Euler rollout.
28 rng=np.random.default_rng(3); w=rng.normal(size=N)*.04
29 hats=np.fft.fft(w); k=2; explicit=sum(w[j]*np.exp(-2j*np.pi*k*j/N) for j in range(N))
30 what=hats[k]; lam=root_np(what)
31 dt=.001; delay=round(TAU_D/dt); steps=3500
32 x=np.zeros(steps+delay+1,dtype=complex); x[:delay+1]=1e-3*np.exp(lam*np.arange(-delay,1)*dt)
33 for i in range(delay,delay+steps): x[i+1]=x[i]+dt*(-x[i]+what*x[i-delay])
34 t=np.arange(steps)*dt; sel=t>1.0
35 measured=np.polyfit(t[sel],np.log(np.abs(x[delay:delay+steps][sel])+1e-30),1)[0]
36 return {'fft_error':float(abs(what-explicit)), 'predicted_growth':float(lam.real),
37 'observed_growth':float(measured), 'abs_error':float(abs(lam.real-measured)),
38 'passed':bool(abs(what-explicit)<1e-10 and abs(lam.real-measured)<0.02)}
39
40class RingRNN(nn.Module):
41 def __init__(self, hidden=N):
42 super().__init__(); self.hidden=hidden
43 self.inp=nn.Linear(3,hidden); self.w=nn.Parameter(torch.randn(hidden)*0.05)
44 self.out=nn.Linear(hidden,1)
45 def recurrent(self,h):
46 # circulant convolution; FFT convention matches the proposed w_hat formula.
47 return torch.fft.ifft(torch.fft.fft(self.w)*torch.fft.fft(h,dim=-1),dim=-1).real
48 def forward(self,x,return_aux=False):
49 q=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device)
50 slopes=[]; preacts=[]
51 for t in range(q.shape[1]):
52 p=self.inp(q[:,t])+self.recurrent(h)
53 slopes.append(1-torch.tanh(p).pow(2)); preacts.append(p)
54 h=torch.tanh(p)
55 y=self.out(h)
56 if return_aux: return y, torch.stack(slopes).mean(), torch.stack(preacts)
57 return y
58
59def spectral_penalty(net,a):
60 hats=torch.fft.fft(net.w)
61 vals=[]
62 roots=[]
63 for k in MODES:
64 z=(-1+a*hats[k])/TAU_R
65 for _ in range(8):
66 e=torch.exp(-z*TAU_D); f=TAU_R*z+1-a*hats[k]*e
67 z=z-f/(TAU_R+a*hats[k]*TAU_D*e)
68 roots.append(z); vals.append(torch.nn.functional.softplus(z.real+0.02)**2)
69 return torch.stack(vals).mean(), torch.stack(roots)
70
71def train_one(seed, lr, wd, strength):
72 seed_all(seed)
73 d=get_dataset('dynamics',seed,n_train=1000,n_test=300)
74 net=RingRNN().to(DEVICE)
75 opt=torch.optim.Adam(net.parameters(),lr=lr,weight_decay=wd)
76 xtr,ytr=d['xtr'].to(DEVICE),d['ytr'].to(DEVICE)
77 lossf=nn.MSELoss(); batch=128; last_roots=None
78 for _ in range(12):
79 net.train(); perm=torch.randperm(len(xtr),device=DEVICE)
80 for i in range(0,len(xtr),batch):
81 ix=perm[i:i+batch]; pred,a,_=net(xtr[ix],True); loss=lossf(pred,ytr[ix].view(-1,1))
82 if strength>0:
83 sp,last_roots=spectral_penalty(net,a.detach()); loss=loss+strength*sp
84 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.0); opt.step()
85 net.eval()
86 with torch.no_grad(): metric=float(lossf(net(d['xte'].to(DEVICE)),d['yte'].to(DEVICE).view(-1,1)).cpu())
87 # Signature uses the trained model: predicted characteristic growth versus
88 # observed infinitesimal autonomous perturbation growth of its recurrence.
89 with torch.no_grad():
90 a=float(net.inp.weight.new_tensor(0.0))
91 q=d['xte'][:64].to(DEVICE).view(64,-1,3); h=torch.zeros(64,N,device=DEVICE)
92 ss=[]
93 for t in range(q.shape[1]):
94 p=net.inp(q[:,t])+net.recurrent(h); ss.append((1-torch.tanh(p).pow(2)).mean()); h=torch.tanh(p)
95 a=float(torch.stack(ss).mean().cpu()); hats=torch.fft.fft(net.w).cpu().numpy()
96 pred=max(root_np(h,a).real for h in hats[[k%N for k in MODES]])
97 h0=h[:1].clone(); eps=1e-4; hp=h0+eps*torch.randn_like(h0)
98 vals=[]
99 for _ in range(30):
100 hp=torch.tanh(net.recurrent(hp)); vals.append(float(torch.linalg.vector_norm(hp-h0).cpu())+1e-30)
101 obs=float(np.polyfit(np.arange(30),np.log(vals),1)[0])
102 return metric, {'predicted_max_growth':float(pred),'observed_perturbation_growth':obs,'local_slope':a,
103 'growth_abs_error':abs(float(pred)-obs)}
104
105def main():
106 sanity_result=sanity(); print('SANITY',json.dumps(sanity_result))
107 grid=[{'lr':lr,'wd':wd} for lr in (1e-3,3e-3,6e-3) for wd in (0.,1e-4)]
108 def base_fn(cfg): return lambda s: train_one(s,cfg['lr'],cfg['wd'],0.)[0]
109 base=sweep_baseline(base_fn,grid,seeds=(0,1,2,3))
110 best=base['best_cfg']; strengths=(0.01,0.03,0.10)
111 idea_cfgs=[{'lr':best['lr'],'wd':best['wd'],'strength':v} for v in strengths]
112 idea_trials=[]
113 for cfg in idea_cfgs:
114 r=evaluate(lambda s: train_one(s,cfg['lr'],cfg['wd'],cfg['strength'])[0])
115 idea_trials.append((cfg,r))
116 cfg,idea=min(idea_trials,key=lambda z:z[1]['mean'])
117 # Re-run best idea at all eight seeds already done above; collect signatures separately.
118 sig=[train_one(s,cfg['lr'],cfg['wd'],cfg['strength'])[1] for s in range(8)]
119 sigmean={k:float(np.mean([z[k] for z in sig])) for k in sig[0]}
120 sigmean['confirmed']=bool(sigmean['growth_abs_error']<0.15)
121 report=make_report('dynamics','ring_rnn',base,idea,{'predicted_vs_observed_growth':sigmean,
122 'stage1_sanity':sanity_result,'idea_sweep':[{'cfg':c,'mean':r['mean']} for c,r in idea_trials]})
123 report['parameter_count']=sum(p.numel() for p in RingRNN().parameters())
124 report['device']=DEVICE
125 Path('bench_report.json').write_text(json.dumps(report,indent=2))
126 print(json.dumps(report,indent=2))
127if __name__=='__main__': main()