Fourier-Mode Stability Shaping / stage2_fourier_bench.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()