import sys, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report SEEDS = tuple(range(8)) SWEEP_SEEDS = (0, 1, 2, 3) NTR, NTE, EPOCHS, BATCH = 1200, 400, 15, 128 LR_GRID = [1e-3, 3e-3, 6e-3] def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): try: torch.cuda.manual_seed_all(seed) except Exception: pass class SingleBandRNN(nn.Module): """Single-timescale counterpart with the same dissipative state update.""" def __init__(self, hidden=64, gamma=0.4, out_dim=1): super().__init__() self.hidden = hidden self.inp = nn.Linear(3, hidden) self.L = nn.Parameter(0.05 * torch.randn(hidden, hidden)) self.alpha_raw = nn.Parameter(torch.tensor(float(np.log(np.expm1(gamma))))) self.head = nn.Linear(hidden, out_dim) self.dt = 0.1 def J(self): a = torch.nn.functional.softplus(self.alpha_raw) + 1e-4 return -(self.L.T @ self.L) - a * torch.eye(self.hidden, device=self.L.device) def forward(self, x, return_states=False): z = x.view(x.shape[0], -1, 3) h = torch.zeros(x.shape[0], self.hidden, device=x.device) states=[]; J=self.J() for k in range(z.shape[1]): h = h + self.dt * (h @ J.T + self.inp(z[:, k])) h = torch.tanh(h) states.append(h) if return_states: return self.head(h), torch.stack(states, 1) return self.head(h) class DualBandRNN(nn.Module): """Two dissipative bands plus bounded explicit cross-band exchange.""" def __init__(self, width=32, gamma_f=4.0, gamma_s=0.4, out_dim=1): super().__init__(); self.width=width; self.dt=.1 self.bf=nn.Linear(3,width); self.bs=nn.Linear(3,width) self.Lf=nn.Parameter(.05*torch.randn(width,width)); self.Ls=nn.Parameter(.05*torch.randn(width,width)) self.af=nn.Parameter(torch.tensor(float(np.log(np.expm1(gamma_f))))) self.a_s=nn.Parameter(torch.tensor(float(np.log(np.expm1(gamma_s))))) self.Efs=nn.Parameter(.03*torch.randn(width,width)); self.Esf=nn.Parameter(.03*torch.randn(width,width)) self.head=nn.Linear(2*width,out_dim) def Js(self): eye=torch.eye(self.width,device=self.Lf.device) jf=-(self.Lf.T@self.Lf)-(torch.nn.functional.softplus(self.af)+1e-4)*eye js=-(self.Ls.T@self.Ls)-(torch.nn.functional.softplus(self.a_s)+1e-4)*eye # bounded exchange prevents uncontrolled growth while retaining learned coupling ef=.15*torch.tanh(self.Efs); es=.15*torch.tanh(self.Esf) return jf,js,ef,es def forward(self,x,return_states=False): z=x.view(x.shape[0],-1,3); hf=torch.zeros(x.shape[0],self.width,device=x.device); hs=hf.clone() jf,js,ef,es=self.Js(); traces=[] for k in range(z.shape[1]): 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])) hf=torch.tanh(hf); hs=torch.tanh(hs); traces.append(torch.cat([hf,hs],1)) h=torch.cat([hf,hs],1) if return_states:return self.head(h),torch.stack(traces,1) return self.head(h) def train_one(kind, seed, lr, gamma=0.4): seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE) model=SingleBandRNN(gamma=gamma) if kind=='baseline' else DualBandRNN() net,metric,_=train_model(model,ds,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None) return float(metric), net, ds def fn(kind, cfg): def run(seed): return train_one(kind,seed,cfg['lr'],cfg.get('gamma',.4))[0] return run def signature(): pred=[]; obs=[] for seed in SEEDS: _,net,ds=train_one('idea',seed,3e-3) net.eval(); dev=next(net.parameters()).device x=torch.zeros(1,96,device=dev); x[0,:3]=ds['xtr'][0,:3].to(dev) with torch.no_grad(): _,st=net(x,return_states=True) 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 # observed discrete decay fit after the impulse, excluding first point 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] 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()) pred.append([abs(pf),abs(ps)]); obs.append([max(0.,-slf),max(0.,-sls)]) p=np.mean(pred,0); o=np.mean(obs,0) rel=np.abs(o-p)/np.maximum(p,1e-6) 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))} def main(): base_grid=[{'lr':lr,'gamma':g} for lr in LR_GRID for g in (.4,1.0)] base=sweep_baseline(lambda c:fn('baseline',c),base_grid,seeds=SWEEP_SEEDS) idea_cfgs=[{'lr':lr} for lr in LR_GRID] # Evaluate the prescribed three-setting idea sweep on all paired seeds. idea_runs=[] for c in idea_cfgs: r=evaluate(fn('idea',c),SEEDS); idea_runs.append((r,c)) idea,bestcfg=min(idea_runs,key=lambda q:q[0]['mean']) rep=make_report('dynamics','rnn_small',base,idea,extra=signature()) rep['idea_sweep']=[{'cfg':c,'mean':r['mean'],'per_seed':r['per_seed']} for r,c in idea_runs] 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.' Path('bench_report.json').write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()