Fourier Replay-Mode Stabilizer / stage2_bench.py

Failed on benchmark

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, 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()