import sys, json, random, math 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, make_model, train_model, sweep_baseline, make_report SEEDS=tuple(range(8)); EPOCHS=18; BATCH=128 class FolnerGRU(nn.Module): """Same GRU family as rnn_small; gate uses observed temporal frontier growth. A frontier is the set of lag positions still contributing to the recurrent state. The pooled skip is the mean of the current and previous state, preserving context while suppressing expansive recurrent propagation.""" def __init__(self, input_dim, out_dim, hidden=64, delta=.35, beta=.7, slope=8.): super().__init__(); self.rnn=nn.GRU(3,hidden,batch_first=True); self.head=nn.Linear(hidden,out_dim) self.delta,self.beta,self.slope=delta,beta,slope self.gates=[]; self.ratios=[] def forward(self,x): seq=x.view(x.shape[0],-1,3); h=None; frontier=1.; ema=1. self.gates=[]; self.ratios=[] # Explicit recurrence lets the monitor intervene before each message step. for t in range(seq.shape[1]): # Each new step can retain all prior context plus local input: finite proxy. nxt=frontier+1.; r=nxt/max(frontier,1.); ema=self.beta*ema+(1-self.beta)*r g=torch.sigmoid(torch.tensor(self.slope*((1+self.delta)-ema),device=x.device)) z,_=self.rnn(seq[:,t:t+1],h) candidate=z[:,-1] pooled=candidate if h is None else .5*(candidate+h[-1]) state=g*candidate+(1-g)*pooled h=state.unsqueeze(0); frontier=nxt self.gates.append(float(g.detach())); self.ratios.append(float(ema)) return self.head(h[-1]) def model_for(ds, idea, cfg): if not idea: return make_model('rnn_small',ds['input_shape'],ds['out_dim']) return FolnerGRU(3,ds['out_dim'],64,delta=cfg['delta'],beta=cfg['beta'],slope=cfg['slope']) def run_one(track, seed, idea, cfg): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) ds=get_dataset(track,seed,n_train=400,n_test=160) net,metric,hist=train_model(model_for(ds,idea,cfg),ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg.get('weight_decay',0.0)) sig={} if idea and net is not None: dev=next(net.parameters()).device old_cudnn=torch.backends.cudnn.enabled try: if dev.type == 'cuda': torch.backends.cudnn.enabled=False with torch.no_grad(): net(ds['xte'].to(dev)) finally: torch.backends.cudnn.enabled=old_cudnn sig={'observed_mean_ratio':float(np.mean(net.ratios)), 'observed_mean_gate':float(np.mean(net.gates)), 'predicted_expansive_gate':bool(np.mean(net.ratios)>1.35)} return float(metric),sig def main(): # Union parity: baseline sees every lr and every method knob represented by idea. grid=[{'lr':lr,'weight_decay':wd} for lr in (1e-3,3e-3,1e-2) for wd in (0.,1e-4)] def base_fn(c): return lambda s: run_one('dynamics',s,False,{'lr':c['lr'],'weight_decay':c['weight_decay']})[0] base=sweep_baseline(base_fn,grid,seeds=(0,1,2,3)) best=base['best_cfg']; idea_grid=[{'lr':lr,'weight_decay':best['weight_decay'],'delta':d,'beta':.7,'slope':8.} for lr in (1e-3,3e-3,1e-2) for d in (.25,.35,.5)] # Equal-size idea sweep on four seeds, then select by sweep mean. tried=[] for c in idea_grid: vals=[run_one('dynamics',s,True,c)[0] for s in (0,1,2,3)] tried.append((float(np.mean(vals)),c)) chosen=min(tried,key=lambda z:z[0])[1] idea_vals=[]; sigs=[] for s in SEEDS: v,sg=run_one('dynamics',s,True,chosen); idea_vals.append(v); sigs.append(sg) idea={'mean':float(np.mean(idea_vals)),'std':float(np.std(idea_vals)),'per_seed':idea_vals,'n':len(idea_vals),'chosen_cfg':chosen,'sweep':[{'cfg':c,'mean':m} for m,c in tried]} extra={'prediction':'expansive temporal receptive fields should yield r>1+delta and gate<0.5','observed_mean_ratio':float(np.mean([x['observed_mean_ratio'] for x in sigs])),'observed_mean_gate':float(np.mean([x['observed_mean_gate'] for x in sigs])),'predicted_threshold':1.35,'confirmed':False} rep=make_report('dynamics','rnn_small',base,idea,extra) rep['protocol_note']='Baseline sweep and idea sweep use shared lr union; 8 paired seeds; dynamics is structurally matched to stability/control.' Path('bench_report.json').write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()