Spectral-Band Dual-Timescale Network / stage2_bench.py

Failed on benchmark

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