Adaptive Ballistic-to-Diffusive Propagation Schedule / bench_adaptive.py
Mechanism confirmed, baseline not beaten
1import json, sys
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8)); SWEEP_SEEDS = tuple(range(4))
10LRS = [1e-3, 3e-3, 1e-2]; GAMMAS = [0.0, 0.15, 0.30, 0.50]
11EPOCHS = 18; NTRAIN, NTEST = 400, 200
12
13class DephasedRNN(nn.Module):
14 def __init__(self, out_dim, gamma0=.5, target=.12, alpha=.5, adaptive=True,
15 depth=4, gmin=0., gmax=1.):
16 super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,out_dim)
17 self.gamma0=gamma0; self.target=target; self.alpha=alpha; self.adaptive=adaptive
18 self.depth=depth; self.gmin=gmin; self.gmax=gmax; self.last_stats={}
19 def forward(self,x):
20 h,_=self.rnn(x.view(x.shape[0],-1,3)); gamma=torch.full((x.shape[0],),float(self.gamma0),device=x.device)
21 ratios=[]; gammas=[]; hf=[]
22 for _ in range(self.depth):
23 var=h.var(dim=(1,2),unbiased=False)+1e-8; corr=(h[:,1:]*h[:,:-1]).mean(dim=(1,2)).abs(); r=corr/var
24 if self.adaptive: gamma=torch.clamp(gamma*torch.exp(self.alpha*(r.detach()-self.target)),self.gmin,self.gmax)
25 left=torch.cat([h[:,:1],h[:,:-1]],1); right=torch.cat([h[:,1:],h[:,-1:]],1)
26 h=h+gamma[:,None,None]*(.5*(left+right)-h)
27 ratios.append(r.detach().mean().item()); gammas.append(gamma.detach().mean().item())
28 z=torch.fft.rfft(h-h.mean(dim=1,keepdim=True),dim=1); hf.append((z[:,max(1,z.shape[1]//2):].abs()**2).mean().item())
29 self.last_stats={'ratio':float(np.mean(ratios)),'gamma_final':float(gamma.mean()),'gamma_mean':float(np.mean(gammas)),'hf_energy':float(np.mean(hf))}
30 return self.head(h[:,-1])
31
32def seed_all(seed):
33 np.random.seed(seed); torch.manual_seed(seed)
34 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
35
36def train_system(cfg,seed,adaptive):
37 seed_all(seed); d=get_dataset('dynamics',seed,NTRAIN,NTEST)
38 m=DephasedRNN(d['out_dim'],cfg['gamma0'],cfg['target'],cfg['alpha'],adaptive,cfg['depth'],gmax=cfg['gmax'])
39 return train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None)
40
41def baseline_fn(cfg):
42 return lambda seed: train_system({'lr':cfg['lr'],'gamma0':cfg['gamma'],'target':.12,'alpha':.5,'depth':4,'gmax':1.},seed,False)[1]
43
44def idea_fn(cfg,capture=False):
45 def run(seed):
46 net,metric,_=train_system(cfg,seed,True)
47 if capture and net is not None:
48 d=get_dataset('dynamics',seed,NTRAIN,NTEST); net.eval()
49 try:
50 dev=next(net.parameters()).device
51 with torch.no_grad(): net(d['xte'].to(dev))
52 except RuntimeError:
53 # Post-training signature must survive the shared-GPU/cuDNN limit.
54 net=net.to('cpu')
55 with torch.no_grad(): net(d['xte'].cpu())
56 return metric,dict(net.last_stats)
57 return metric
58 return run
59
60def main():
61 base_grid=[{'lr':lr,'gamma':g} for lr in LRS for g in GAMMAS]
62 base=sweep_baseline(baseline_fn,base_grid,seeds=SWEEP_SEEDS)
63 idea_grid=[{'lr':.001,'gamma0':.15,'target':.12,'alpha':.5,'depth':4,'gmax':1.},
64 {'lr':.003,'gamma0':.30,'target':.12,'alpha':.5,'depth':4,'gmax':1.},
65 {'lr':.01,'gamma0':.50,'target':.12,'alpha':.5,'depth':4,'gmax':1.}]
66 trials=[]
67 for cfg in idea_grid:
68 r=evaluate(idea_fn(cfg),SEEDS); trials.append({'cfg':cfg,**r})
69 best=min(trials,key=lambda z:z['mean']); idea={k:best[k] for k in ['mean','std','per_seed','n']}; idea['best_cfg']=best['cfg']; idea['sweep']=trials
70 sigs=[]
71 for s in SEEDS:
72 _,st=idea_fn(best['cfg'],True)(s)
73 if st and all(np.isfinite(v) for v in st.values()): sigs.append(st)
74 observed={k:float(np.mean([z[k] for z in sigs])) for k in ['ratio','gamma_final','gamma_mean','hf_energy']}
75 signature={'prediction':'trained adaptive propagation raises gamma when neighboring correlation exceeds target; diffusion then attenuates high-frequency hidden energy',
76 'predicted':{'mean_ratio_exceeds_target':True,'gamma_final_exceeds_initial':True,'hf_energy_finite':True},
77 'observed_means':observed,'initial_gamma':best['cfg']['gamma0'],'target':best['cfg']['target'],'n_models':len(sigs),
78 'confirmed':bool(observed['ratio']>best['cfg']['target'] and observed['gamma_final']>best['cfg']['gamma0'] and np.isfinite(observed['hf_energy']))}
79 report=make_report('dynamics','rnn_small',base,idea,signature)
80 report['protocol']={'paired_seeds':list(SEEDS),'n_train':NTRAIN,'n_test':NTEST,'epochs':EPOCHS,'track_rationale':'dynamics matches stability/control and multi-step rollout structure','baseline_grid':base_grid,'idea_grid':idea_grid}
81 Path('bench_report.json').write_text(json.dumps(report,indent=2,allow_nan=False)); print(json.dumps(report,indent=2,allow_nan=False))
82if __name__=='__main__': main()