Primal-Dual Active-Set Optimizer Filter / pdas_bench.py
Mechanism confirmed, baseline not beaten
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()