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 TRACK='dynamics'; MODEL='rnn_small'; N=4; H=16; EPOCHS=12; BATCH=128 ROOT=Path(__file__).parent def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def lap(n=N): a=np.zeros((n,n)) for i in range(1,n): a[i,i]=1.; a[i,0]=-1. return a def math_check(): L=lap(); x0=np.linspace(-1,1,N); dt=.001; rows=[] for eta in (.5,1.): x=x0.copy(); v=[] for _ in range(4000): m=x.mean(); v.append(.5*np.sum((x-m)**2)); x += dt*(.8*x-eta*L@x) obs=float(np.polyfit(np.arange(len(v))*dt,np.log(np.maximum(v,1e-300)),1)[0]) pred=-2*(eta-.8) rows.append({'eta':eta,'predicted_slope':pred,'observed_slope':obs,'relative_error':abs(obs-pred)/max(abs(pred),1e-9)}) return {'gamma_G':1.,'L':.8,'threshold':.8,'rows':rows,'passed':rows[1]['observed_slope']<0 and rows[0]['observed_slope']>0} class ParallelGRU(nn.Module): def __init__(self, coupled=False, eta=0.): super().__init__(); self.coupled=coupled; self.eta=float(eta) self.grus=nn.ModuleList([nn.GRUCell(3,H) for _ in range(N)]) self.head=nn.Linear(H,1) def forward(self,x, return_signature=False): z=x.view(x.shape[0],-1,3); hs=[torch.zeros(x.shape[0],H,device=x.device) for _ in range(N)] for t in range(z.shape[1]): hs=[g(z[:,t],h) for g,h in zip(self.grus,hs)] if self.coupled: root=hs[0] hs=[hs[0]]+[h+self.eta*(root-h) for h in hs[1:]] stack=torch.stack(hs,1); out=self.head(stack.mean(1)) if return_signature: return out, float((stack-stack.mean(1,keepdim=True)).pow(2).mean().sqrt().detach().cpu()) return out def run(cfg, seed, coupled): seed_all(seed); ds=get_dataset(TRACK,seed,n_train=400,n_test=400) model=ParallelGRU(coupled=coupled,eta=cfg['eta']) _, metric, _=train_model(model,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None) return float(metric) def train_and_measure(cfg,seed,coupled): seed_all(seed); ds=get_dataset(TRACK,seed,n_train=400,n_test=400) model=ParallelGRU(coupled=coupled,eta=cfg['eta']) net,metric,_=train_model(model,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None) with torch.no_grad(): dev=next(net.parameters()).device _, disagreement=net(ds['xte'].to(dev),return_signature=True) return float(metric), disagreement def main(): check=math_check(); lrs=[0.0015,0.003,0.006]; base_grid=[{'lr':lr,'eta':0.0} for lr in lrs] def base_fn(cfg): return lambda s: run(cfg,s,False) baseline=sweep_baseline(lambda c: base_fn(c),base_grid,seeds=(0,1,2,3)) best_lr=baseline['best_cfg']['lr']; idea_grid=[{'lr':best_lr,'eta':e} for e in (.02,.08,.20)] # Baseline was evaluated at every union learning rate; idea settings share best lr. idea_records=[] for cfg in idea_grid: r=evaluate(lambda s: run(cfg,s,True),seeds=tuple(range(8))) idea_records.append({'cfg':cfg,'result':r}) best=min(idea_records,key=lambda q:q['result']['mean']); idea=best['result']; cfg=best['cfg'] base_full=baseline['full']; sig=[] for s in range(8): bm,bd=train_and_measure({'lr':best_lr,'eta':0.},s,False) im,idg=train_and_measure(cfg,s,True); sig.append({'seed':s,'baseline_disagreement':bd,'idea_disagreement':idg}) obs=float(np.mean([q['idea_disagreement'] for q in sig])); bobs=float(np.mean([q['baseline_disagreement'] for q in sig])) predicted='negative disagreement trend when eta*gamma_G exceeds L'; confirmed=bool(cfg['eta']>0 and obs < bobs) rep=make_report(TRACK,MODEL,baseline,idea,{'math_check':check,'idea_grid':idea_records,'trained_behavior':sig,'prediction':predicted,'predicted_threshold':.8,'observed_mean_baseline_disagreement':bobs,'observed_mean_idea_disagreement':obs,'confirmed':confirmed}) rep['parameter_counts']={'baseline':sum(p.numel() for p in ParallelGRU(False).parameters()),'idea':sum(p.numel() for p in ParallelGRU(True,cfg['eta']).parameters())} rep['protocol_notes']='8 paired seeds; baseline sweep uses 4 seeds and full re-evaluation; n_train=n_test=400; same ParallelGRU base and optimizer.' (ROOT/'bench_report.json').write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()