Focus-Coefficient Switched Optimizer / stage2_focus_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()