Focus-Coefficient Switched Optimizer / stage2_focus_bench.py
Mechanism confirmed, baseline not beaten
1import sys, os, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
7
8SEED0 = 1057
9EPOCHS = 12
10NTR, NTE = 1200, 400
11BATCH = 128
12LRS = [0.003, 0.006, 0.012]
13MOMS = [0.8, 0.9]
14IDEA_GRID = [
15 {'lr': 0.003, 'momentum': 0.8, 'branch_ratio': 2.0},
16 {'lr': 0.006, 'momentum': 0.8, 'branch_ratio': 2.0},
17 {'lr': 0.012, 'momentum': 0.9, 'branch_ratio': 2.0},
18]
19
20
21def seed_all(seed):
22 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
23 if torch.cuda.is_available():
24 try: torch.cuda.manual_seed_all(seed)
25 except Exception: pass
26
27
28def loss_fn(ds):
29 return nn.CrossEntropyLoss() if ds['task'] == 'classification' else nn.MSELoss()
30
31
32def device_for():
33 return 'cuda' if torch.cuda.is_available() else 'cpu'
34
35
36def baseline_train(seed, cfg, collect=False):
37 seed_all(seed)
38 ds = get_dataset('dynamics', seed, n_train=NTR, n_test=NTE)
39 net = make_model('rnn_small', tuple(ds['xtr'].shape[1:]), 1)
40 dev = device_for(); lf = loss_fn(ds)
41 try:
42 net.to(dev); x, y = ds['xtr'].to(dev), ds['ytr'].to(dev).reshape(-1, 1)
43 opt = torch.optim.SGD(net.parameters(), lr=cfg['lr'], momentum=cfg['momentum'])
44 hist=[]
45 for _ in range(EPOCHS):
46 net.train(); perm=torch.randperm(len(x), device=dev); total=0.
47 for i in range(0,len(x),BATCH):
48 ix=perm[i:i+BATCH]; out=net(x[ix]); loss=lf(out,y[ix])
49 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(), 5.0); opt.step()
50 total += float(loss.detach())*len(ix)
51 hist.append(total/len(x))
52 net.eval()
53 with torch.no_grad(): metric=float(lf(net(ds['xte'].to(dev)),ds['yte'].to(dev).reshape(-1,1)))
54 return metric
55 except RuntimeError:
56 # CPU retry is explicit, as required for shared CUDA failures.
57 if dev == 'cuda':
58 torch.cuda.empty_cache(); return baseline_train_cpu(seed,cfg)
59 raise
60
61
62def baseline_train_cpu(seed,cfg):
63 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
64 net=make_model('rnn_small',tuple(ds['xtr'].shape[1:]),1); lf=loss_fn(ds)
65 x,y=ds['xtr'],ds['ytr'].reshape(-1,1); opt=torch.optim.SGD(net.parameters(),lr=cfg['lr'],momentum=cfg['momentum'])
66 for _ in range(EPOCHS):
67 for i in range(0,len(x),BATCH):
68 loss=lf(net(x[i:i+BATCH]),y[i:i+BATCH]); opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.0); opt.step()
69 with torch.no_grad(): return float(lf(net(ds['xte']),ds['yte'].reshape(-1,1)))
70
71
72def idea_train(seed, cfg, signature=False):
73 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
74 net=make_model('rnn_small',tuple(ds['xtr'].shape[1:]),1); lf=loss_fn(ds); dev=device_for()
75 try:
76 net.to(dev); x,y=ds['xtr'].to(dev),ds['ytr'].to(dev).reshape(-1,1)
77 params=[p for p in net.parameters() if p.requires_grad]
78 # Fixed two-dimensional orthonormal projection of full parameter state.
79 rng=torch.Generator(device=dev); rng.manual_seed(seed+991)
80 dirs=[]
81 for j in range(2):
82 v=torch.cat([torch.randn(p.numel(),generator=rng,device=dev) for p in params]);
83 for q in dirs: v-=torch.dot(v,q)*q
84 dirs.append(v/(v.norm()+1e-12))
85 center=torch.cat([p.detach().flatten() for p in params]).clone()
86 mom={p:torch.zeros_like(p) for p in params}; estimates={0:[],1:[]}; chosen=[]
87 for _ in range(EPOCHS):
88 perm=torch.randperm(len(x),device=dev)
89 for i in range(0,len(x),BATCH):
90 ix=perm[i:i+BATCH]; out=net(x[ix]); loss=lf(out,y[ix])
91 grads=torch.autograd.grad(loss,params,retain_graph=False)
92 flat=torch.cat([g.detach().flatten() for g in grads])
93 cur=torch.cat([p.detach().flatten() for p in params]); z=torch.stack([torch.dot(cur-center,q) for q in dirs]); r=float(z.norm())
94 scores=[]
95 for branch,mult in enumerate((1.0,cfg['branch_ratio'])):
96 # Candidate projected state using branch-specific step and momentum.
97 cand=cur.clone(); off=0
98 for p,g in zip(params,grads):
99 old=mom[p]; new=cfg['momentum']*old + g
100 step=cfg['lr']*mult*new
101 cand[off:off+p.numel()] -= step.flatten(); off+=p.numel()
102 zz=torch.stack([torch.dot(cand-center,q) for q in dirs]); rr=float(zz.norm())
103 drift=(rr-r)/(max(r,1e-4)**3) if r>1e-5 else rr-r
104 estimates[branch].append(drift); scores.append(drift)
105 # lower predicted radial coefficient/drift is the focus rule; mild hysteresis.
106 b=int(np.argmin(scores)); chosen.append(b)
107 off=0
108 for p,g in zip(params,grads):
109 mom[p].mul_(cfg['momentum']).add_(g)
110 p.data.add_(-cfg['lr']*(cfg['branch_ratio'] if b else 1.0)*mom[p])
111 net.eval()
112 with torch.no_grad(): metric=float(lf(net(ds['xte'].to(dev)),ds['yte'].to(dev).reshape(-1,1)))
113 if signature:
114 vals=[np.asarray(estimates[k][max(0,len(estimates[k])//3):]) for k in (0,1)]
115 means=[float(np.mean(v)) if len(v) else float('nan') for v in vals]
116 obs=float(np.mean([1 if chosen[j]==0 else -1 for j in range(len(chosen))]))
117 return metric, {'predicted_branch': int(np.argmin(means)), 'predicted_drift': means, 'observed_selected_fraction_branch0': (obs+1)/2, 'confirmed': bool(np.isfinite(means).all() and means[0] != means[1] and int(np.argmin(means)) == (0 if (obs+1)/2 >= .5 else 1))}
118 return metric
119 except RuntimeError:
120 if dev=='cuda': torch.cuda.empty_cache(); return idea_train_cpu(seed,cfg,signature)
121 raise
122
123
124def idea_train_cpu(seed,cfg,signature=False):
125 # Re-enter with CUDA disabled, preserving exactly the same intervention.
126 old=torch.cuda.is_available
127 torch.cuda.is_available=lambda: False
128 try: return idea_train(seed,cfg,signature)
129 finally: torch.cuda.is_available=old
130
131
132def main():
133 # Baseline sweep includes every lr and momentum appearing on the idea side.
134 grid=[{'lr':lr,'momentum':m} for lr in LRS for m in MOMS]
135 base=sweep_baseline(lambda c: (lambda s: baseline_train(s,c)),grid)
136 # Select idea setting on the same four tuning seeds, then full paired evaluation.
137 itried=[]
138 for c in IDEA_GRID:
139 r=evaluate(lambda s,c=c: idea_train(s,c), seeds=(0,1,2,3)); itried.append({'cfg':c,'mean':r['mean']})
140 best=min(itried,key=lambda z:z['mean'])['cfg']
141 idea=evaluate(lambda s: idea_train(s,best),seeds=tuple(range(8)))
142 sig=idea_train(0,best,signature=True)[1]
143 sig.update({'definition':'NN-scale projected parameter radial drift; lower predicted drift branch should be selected','n_probe_updates':int(EPOCHS*((NTR+BATCH-1)//BATCH))})
144 report=make_report('dynamics','rnn_small',base,idea,{'track_match':'stability/control -> dynamics','idea_sweep':itried,'best_idea_cfg':best,**sig})
145 report['baseline']['parity_grid']=grid
146 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
147 print(json.dumps(report,indent=2))
148
149if __name__=='__main__': main()