Fourier-Mode Stability Shaping / stage2_fourier_bench.py
Beats tuned baseline
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, train_model, evaluate, sweep_baseline, make_report
8
9N, M, KAPPA, Q = 32, 4, 0.15, 0
10SEEDS = tuple(range(8))
11SWEEP_SEEDS = tuple(range(4))
12# Shared union: every idea learning rate is also evaluated by baseline.
13GRID = [{'lr': 1e-3, 'weight_decay': 0.0},
14 {'lr': 3e-3, 'weight_decay': 0.0},
15 {'lr': 1e-2, 'weight_decay': 0.0}]
16
17
18def mu_factors(weights, q=Q):
19 w = np.asarray(weights, dtype=float)
20 ell = np.arange(1, len(w)+1)
21 cq = np.cos(2*np.pi*q*ell/N)
22 return np.array([np.sum(w*cq*(1-np.cos(2*np.pi*k*ell/N))) for k in range(N)])
23
24class CyclicRNN(nn.Module):
25 def __init__(self, hidden=N):
26 super().__init__()
27 self.inp = nn.Linear(3, hidden)
28 self.weights = nn.Parameter(torch.ones(M))
29 self.head = nn.Linear(hidden, 1)
30 def transition(self, h, z):
31 v = torch.zeros_like(h)
32 for j in range(M):
33 # forward neighbor i+j+1, matching the paper's convention
34 v = v + self.weights[j] * (torch.roll(h, -(j+1), dims=1) - h)
35 return torch.tanh(h + KAPPA*v + z)
36 def forward(self, x):
37 seq = x.view(x.shape[0], -1, 3)
38 h = torch.zeros(x.shape[0], N, device=x.device, dtype=x.dtype)
39 for t in range(seq.shape[1]):
40 h = self.transition(h, self.inp(seq[:, t]))
41 return self.head(h)
42
43def seed_all(seed):
44 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
45 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
46
47def train_idea(model, ds, epochs, lr, weight_decay=0.0, lam=0.03, eps=0.02):
48 errs=[]
49 ladder = [('cuda', False), ('cuda', True)] if torch.cuda.is_available() else []
50 ladder.append(('cpu', False))
51 for device, no_cudnn in ladder:
52 try:
53 if no_cudnn: torch.backends.cudnn.enabled=False
54 net=model.to(device); opt=torch.optim.Adam(net.parameters(),lr=lr,weight_decay=weight_decay)
55 x,y=ds['xtr'].to(device),ds['ytr'].to(device)
56 for _ in range(epochs):
57 net.train(); perm=torch.randperm(len(x),device=device)
58 for i in range(0,len(x),128):
59 ix=perm[i:i+128]; pred=net(x[ix]); loss=((pred-y[ix])**2).mean()
60 mu=torch.stack([torch.sum(net.weights*torch.tensor(
61 np.cos(2*np.pi*Q*np.arange(1,M+1)/N)*(1-np.cos(2*np.pi*k*np.arange(1,M+1)/N)),
62 device=device,dtype=net.weights.dtype)) for k in range(1,N)])
63 barrier=torch.nn.functional.softplus(eps-mu).pow(2).mean()
64 loss=loss+lam*barrier
65 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),10.0); opt.step()
66 net.eval()
67 with torch.no_grad(): metric=float(((net(ds['xte'].to(device))-ds['yte'].to(device))**2).mean())
68 if no_cudnn: torch.backends.cudnn.enabled=True
69 return net, metric
70 except RuntimeError as e:
71 errs.append(str(e)[:100])
72 if no_cudnn: torch.backends.cudnn.enabled=True
73 return None, float('nan')
74
75def base_make(cfg):
76 def fn(seed):
77 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
78 net=CyclicRNN(); _,metric,_=train_model(net,ds,epochs=18,lr=cfg['lr'],batch=128,weight_decay=cfg['weight_decay'],log=lambda *_:None)
79 return metric
80 return fn
81
82def idea_make(cfg):
83 def fn(seed):
84 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
85 _,metric=train_idea(CyclicRNN(),ds,epochs=18,lr=cfg['lr'],weight_decay=cfg['weight_decay'])
86 return metric
87 return fn
88
89def choose_idea():
90 tried=[]
91 for cfg in GRID:
92 r=evaluate(idea_make(cfg),SWEEP_SEEDS); tried.append({'cfg':cfg,'mean':r['mean']})
93 best=min(tried,key=lambda x:x['mean'])['cfg']
94 return {'best_cfg':best,'sweep':tried,'full':evaluate(idea_make(best),SEEDS)}
95
96def signature(cfg):
97 seed_all(0); ds=get_dataset('dynamics',0,n_train=400,n_test=200)
98 b=CyclicRNN(); b, bm,_=train_model(b,ds,epochs=18,lr=cfg['lr'],batch=128,log=lambda *_:None)
99 seed_all(0); i,im=train_idea(CyclicRNN(),ds,epochs=18,lr=cfg['lr'])
100 rows=[]
101 for name,net in [('baseline',b),('idea',i)]:
102 w=net.weights.detach().cpu().numpy(); mu=mu_factors(w)
103 obs=[]
104 with torch.no_grad():
105 for k in (1,2,3):
106 phase=2*np.pi*k*np.arange(N)/N
107 h=(1e-4*torch.tensor(np.cos(phase),dtype=torch.float32)[None,:]).to(next(net.parameters()).device)
108 amps=[]
109 for _ in range(30):
110 z=torch.zeros_like(h); h=net.transition(h,z)
111 amps.append(float(torch.abs(torch.fft.fft(h)[0,k]).cpu()))
112 slope=float(np.polyfit(np.arange(len(amps))*1.0,np.log(np.maximum(amps,1e-30)),1)[0])
113 pred=float(-KAPPA*mu[k]); obs.append({'mode':k,'predicted':pred,'observed':slope})
114 rows.append({'system':name,'weights':w.tolist(),'modes':obs})
115 rel=[abs(x['observed']-x['predicted'])/max(abs(x['predicted']),1e-6) for r in rows for x in r['modes']]
116 return {'q':Q,'kappa':KAPPA,'rows':rows,'max_relative_error':float(max(rel)),'confirmed':bool(max(rel)<0.20),'task_metric_seed0':{'baseline':bm,'idea':im}}
117
118def main():
119 base=sweep_baseline(base_make,GRID,seeds=SWEEP_SEEDS)
120 idea=choose_idea()
121 # Ensure idea uses baseline-selected lr if its sweep happened to select another: report its best, all are shared.
122 rep=make_report('dynamics','cyclic_rnn',base,idea['full'],{'note':'trained-model zero-input Fourier perturbation; q=0','signature':signature(idea['best_cfg'])})
123 rep['idea_sweep']=idea['sweep']; rep['protocol']={'paired_seeds':list(SEEDS),'train_samples':400,'test_samples':200,'epochs':18,'architecture':'same CyclicRNN; barrier only differs'}
124 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
125 print(json.dumps(rep,indent=2))
126if __name__=='__main__': main()