Primal-Dual Active-Set Optimizer Filter / pdas_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, time
  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, make_report
  7
  8SEEDS = tuple(range(8))
  9LRS = [0.001, 0.003, 0.006]
 10EPOCHS, BATCH = 12, 128
 11
 12
 13def pdas(H, c, A, b, active_init=(), tol=1e-8, max_iter=30):
 14    n, m = len(c), len(b); active = set(active_init)
 15    for it in range(1, max_iter + 1):
 16        inds = sorted(active)
 17        if inds:
 18            K = np.block([[H, -A[inds].T], [A[inds], np.zeros((len(inds), len(inds)))]])
 19            try: sol = np.linalg.solve(K, np.r_[-c, b[inds]])
 20            except np.linalg.LinAlgError:
 21                active.remove(inds[-1]); continue
 22            v = sol[:n]; mu = np.zeros(m); mu[inds] = sol[n:]
 23        else:
 24            v = np.linalg.solve(H, -c); mu = np.zeros(m)
 25        r = A @ v - b
 26        bad = [j for j in range(m) if j not in active and r[j] < -tol]
 27        neg = [j for j in active if mu[j] < -tol]
 28        if not bad and not neg: return v, mu, tuple(sorted(active)), it
 29        if neg: active.remove(min(neg, key=lambda j: mu[j]))
 30        else: active.add(min(bad, key=lambda j: r[j]))
 31    return v, mu, tuple(sorted(active)), max_iter
 32
 33
 34def seed_model(ds):
 35    return make_model('mlp_tiny', tuple(ds['xtr'].shape[1:]), int(ds['out_dim']))
 36
 37
 38def baseline(ds, lr):
 39    torch.manual_seed(2265)
 40    return train_model(seed_model(ds), ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *_: None)
 41
 42
 43def idea(ds, lr, alpha=0.20):
 44    torch.manual_seed(2265)
 45    net = seed_model(ds)
 46    # One scalar control per parameter tensor. The linearized barrier is a
 47    # parameter-norm trust barrier: ||theta + eta*s*g|| <= cap.
 48    caps = [1.02 * p.detach().norm().item() + 1e-6 for p in net.parameters()]
 49    opt = torch.optim.Adam(net.parameters(), lr=lr)
 50    lossf = nn.MSELoss()
 51    x, y = ds['xtr'], ds['ytr']
 52    if x.ndim == 2: x = x.float()
 53    if y.ndim == 1: y = y.float().reshape(-1, 1)
 54    active = (); its=[]; changes=0; violations=[]; predicted=[]; observed=[]
 55    prev_active = ()
 56    for ep in range(EPOCHS):
 57        gen = torch.Generator().manual_seed(2265 + ep)
 58        order = torch.randperm(len(x), generator=gen)
 59        for start in range(0, len(x), BATCH):
 60            idx = order[start:start+BATCH]
 61            opt.zero_grad(set_to_none=True)
 62            out = net(x[idx]); loss = lossf(out, y[idx]); loss.backward()
 63            H = np.eye(len(list(net.parameters()))) * 1.01
 64            c = -np.ones(len(H)); A=[]; b=[]; nominal=[]
 65            params = list(net.parameters())
 66            for j,p in enumerate(params):
 67                g = p.grad.detach(); th=p.detach(); gn=float(g.norm())
 68                tn=float(th.norm())
 69                # scalar s multiplies Adam's current nominal direction;
 70                # linearized cap gives a_j*s >= b_j.
 71                direction = -g
 72                coeff = -lr * float((th * direction).sum()) / max(tn, 1e-8)
 73                h = caps[j] - tn
 74                # h(theta + eta*s*(-g)) >= (1-alpha)h(theta)
 75                # => coeff*s >= -alpha*h
 76                A.append([coeff if k == j else 0.0 for k in range(len(params))])
 77                b.append(-alpha * h)
 78                nominal.append(1.0)
 79            A=np.asarray(A); b=np.asarray(b); nominal=np.asarray(nominal)
 80            # If a nearly-zero gradient makes a true barrier infeasible, use
 81            # a harmless broad lower bound; this keeps the small MVP robust.
 82            for j in range(len(b)):
 83                if abs(A[j,j]) < 1e-10: A[j,j]=1e-10; b[j]=min(b[j], 0.)
 84            s, mu, active, it = pdas(H, -nominal, A, b, active)
 85            its.append(it); changes += int(active != prev_active); prev_active=active
 86            predicted.append(int(np.sum(np.asarray(b) > np.asarray(A).diagonal()))); observed.append(len(active))
 87            for j,p in enumerate(params):
 88                p.grad.mul_(float(s[j]))
 89            opt.step()
 90            # measured post-update barrier violations
 91            violations.extend(max(0., float(p.detach().norm())-caps[j]) for j,p in enumerate(params))
 92    net.eval()
 93    with torch.no_grad(): metric=float(lossf(net(ds['xte']), ds['yte'].float().reshape(-1,1)).item())
 94    return net, metric, {'iterations':its,'active_changes':changes,'violations':violations,'predicted_active':predicted,'observed_active':observed}
 95
 96
 97def main():
 98    # Correct structural match: optimizer/filter ideas belong to tabular.
 99    allres=[]; sweep=[]
100    for lr in LRS:
101        vals=[]
102        for seed in SEEDS:
103            ds=get_dataset('tabular', seed, n_train=400, n_test=200)
104            vals.append(baseline(ds,lr)[1])
105        sweep.append({'config':{'lr':lr,'epochs':EPOCHS},'mean':float(np.mean(vals)),'per_seed':vals})
106    best=min(sweep,key=lambda z:z['mean']); idea_grid=[]; idea_diag=[]
107    for lr in LRS:
108        vals=[]
109        for seed in SEEDS:
110            ds=get_dataset('tabular', seed, n_train=400, n_test=200)
111            net,v,d=idea(ds,lr); vals.append(v)
112            if lr==best['config']['lr']: idea_diag.append(d)
113        idea_grid.append({'config':{'lr':lr,'alpha':0.20,'epochs':EPOCHS},'mean':float(np.mean(vals)),'per_seed':vals})
114    best_i=min(idea_grid,key=lambda z:z['mean'])
115    base_block={'best_cfg':best['config'],'sweep':[{'cfg':z['config'],'mean':z['mean']} for z in sweep], 'full':{'mean':float(np.mean(best['per_seed'])),'std':float(np.std(best['per_seed'])),'per_seed':best['per_seed'],'n':len(best['per_seed'])}}
116    idea_res={'config':best_i['config'],'mean':best_i['mean'],'std':float(np.std(best_i['per_seed'])),'per_seed':best_i['per_seed'],'n':len(best_i['per_seed']), 'grid':idea_grid}
117    pred=np.concatenate([d['predicted_active'] for d in idea_diag]); obs=np.concatenate([d['observed_active'] for d in idea_diag])
118    sig={'predicted_mean_active':float(pred.mean()),'observed_mean_active':float(obs.mean()),'mean_iterations':float(np.mean([x for d in idea_diag for x in d['iterations']])), 'one_iteration_fraction':float(np.mean([x==1 for d in idea_diag for x in d['iterations']])), 'active_set_changes':int(sum(d['active_changes'] for d in idea_diag)), 'max_violation':float(max([max(d['violations']) for d in idea_diag],default=0.0))}
119    sig['confirmed']=bool(abs(sig['predicted_mean_active']-sig['observed_mean_active']) <= max(0.5,0.25*max(1,sig['observed_mean_active'])))
120    report=make_report('tabular','mlp_tiny',base_block,idea_res,extra=sig)
121    report['baseline_sweep']=sweep; report['idea_grid']=idea_grid; report['search_space_parity']=True
122    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
123    print(json.dumps(report,indent=2))
124
125if __name__=='__main__': main()