Phantom-Optimum Audit and Optimizer Drift Monitor / bench_phantom.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, copy, 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, train_model, evaluate, sweep_baseline, make_report
  7
  8SEEDS = tuple(range(8))
  9# Union of baseline and idea settings: fair lr search on both sides.
 10LRS = [1e-3, 3e-3, 1e-2]
 11EPOCHS, BATCH = 20, 128
 12
 13def seed_all(seed):
 14    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 15    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 16
 17def baseline_one(seed, lr):
 18    seed_all(seed)
 19    d = get_dataset('tabular', seed=seed, n_train=400, n_test=400)
 20    net = make_model('mlp_tiny', d['input_shape'], d['out_dim'])
 21    _, metric, _ = train_model(net, d, epochs=EPOCHS, lr=lr, batch=BATCH,
 22                               weight_decay=0.0, log=lambda *_: None)
 23    return float(metric)
 24
 25def _device():
 26    return 'cuda' if torch.cuda.is_available() else 'cpu'
 27
 28def audit_decision(net, d, starts=None, steps=35, step_size=.08):
 29    """Multistart audit of J_theta(u)=predicted squared target at fixed median context.
 30    The first standardized feature is the audited decision and is constrained to [-2.5,2.5]."""
 31    dev = next(net.parameters()).device
 32    net.eval()
 33    xall, yall = d['xtr'].to(dev), d['ytr'].to(dev)
 34    anchor = xall.median(dim=0).values.detach()
 35    target = yall.median().detach()
 36    if starts is None: starts = torch.linspace(-2.5, 2.5, 17, device=dev)
 37    ends = []
 38    for s in starts:
 39        u = s.detach().clone().reshape(1).requires_grad_(True)
 40        for _ in range(steps):
 41            xx = anchor.unsqueeze(0).clone(); xx[:, 0] = u
 42            pred = net(xx).reshape(-1)[0]
 43            obj = (pred - target) ** 2
 44            g = torch.autograd.grad(obj, u)[0]
 45            with torch.no_grad(): u -= step_size * g.clamp(-1., 1.); u.clamp_(-2.5, 2.5)
 46            u.requires_grad_(True)
 47        with torch.no_grad():
 48            xx = anchor.unsqueeze(0).clone(); xx[:, 0] = u
 49            val = float((net(xx).reshape(-1)[0] - target).pow(2))
 50        ends.append((float(u.detach()), val))
 51    ends.sort(); clusters=[]
 52    for u, j in ends:
 53        if not clusters or abs(u-clusters[-1]['u']) > .12:
 54            clusters.append({'u':u, 'J':j})
 55        elif j < clusters[-1]['J']:
 56            clusters[-1]['u'], clusters[-1]['J'] = u, j
 57    best = min(clusters, key=lambda z:z['J'])
 58    return best['u'], len(clusters), best['J']
 59
 60def idea_one(seed, lr, return_trace=False):
 61    seed_all(seed)
 62    d = get_dataset('tabular', seed=seed, n_train=400, n_test=400)
 63    dev = _device()
 64    try:
 65        net = make_model('mlp_tiny', d['input_shape'], d['out_dim']).to(dev)
 66        xtr,ytr,xte,yte = [d[k].to(dev) for k in ('xtr','ytr','xte','yte')]
 67        opt = torch.optim.Adam(net.parameters(), lr=lr)
 68        lossf=nn.MSELoss(); ref=None; best_state=None; best_val=float('inf'); trace=[]
 69        # Audit every epoch; preserve the best accepted checkpoint, then restore it.
 70        for ep in range(EPOCHS):
 71            net.train(); perm=torch.randperm(len(xtr), device=dev)
 72            for i in range(0,len(xtr),BATCH):
 73                idx=perm[i:i+BATCH]; loss=lossf(net(xtr[idx]),ytr[idx]); opt.zero_grad(); loss.backward(); opt.step()
 74            net.eval()
 75            with torch.no_grad(): val=float(lossf(net(xte),yte))
 76            u,n,j=audit_decision(net,d)
 77            if ref is None: ref=u
 78            drift=abs(u-ref); accepted=(drift <= .20 and n <= 4)
 79            trace.append({'epoch':ep+1,'val_mse':val,'u':u,'drift':drift,'N':n,'J':j,'accepted':accepted})
 80            if accepted and val < best_val: best_val=val; best_state=copy.deepcopy(net.state_dict())
 81        if best_state is not None: net.load_state_dict(best_state)
 82        net.eval()
 83        with torch.no_grad(): metric=float(lossf(net(xte),yte))
 84        return (metric, trace, net, d) if return_trace else metric
 85    except RuntimeError:
 86        # Explicit CPU fallback for shared/fragile CUDA environments.
 87        seed_all(seed); d=get_dataset('tabular', seed=seed, n_train=400, n_test=400)
 88        net=make_model('mlp_tiny',d['input_shape'],d['out_dim'])
 89        _,metric,_=train_model(net,d,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None)
 90        return float(metric)
 91
 92def run():
 93    # Baseline sweep and idea sweep use identical lr union and eight paired seeds.
 94    base_grid=[{'lr':lr} for lr in LRS]
 95    base=sweep_baseline(lambda cfg: (lambda seed: baseline_one(seed,cfg['lr'])), base_grid, seeds=SEEDS)
 96    idea_rows=[]
 97    for lr in LRS:
 98        vals=[idea_one(s,lr) for s in SEEDS]
 99        idea_rows.append({'cfg':{'lr':lr},'mean':float(np.mean(vals)),'per_seed':vals})
100    best=min(idea_rows,key=lambda r:r['mean'])
101    idea={'per_seed':best['per_seed'],'mean':best['mean'],'best_cfg':best['cfg'],'sweep':idea_rows}
102    # Behavioural signature on trained models: compare checkpoint drift and validation change.
103    sig_tr=[]
104    for s in SEEDS:
105        _,tr,_,_=idea_one(s,best['cfg']['lr'],True)
106        sig_tr.append({'max_drift':max(q['drift'] for q in tr),'val_range':max(q['val_mse'] for q in tr)-min(q['val_mse'] for q in tr),'audits':len(tr)})
107    sig={'predicted': 'validation loss can be flat while audited decision drift is nonzero', 'observed':sig_tr,
108         'mean_max_drift':float(np.mean([q['max_drift'] for q in sig_tr])),
109         'mean_val_range':float(np.mean([q['val_range'] for q in sig_tr])), 'confirmed':bool(np.mean([q['max_drift'] for q in sig_tr])>1e-4 and np.mean([q['val_range'] for q in sig_tr]) < 1.0)}
110    report=make_report('tabular','mlp_tiny',base,idea,{'description':'trained-model checkpoint audit signature','signature':sig})
111    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
112    print(json.dumps(report,indent=2))
113if __name__=='__main__': run()