Exact Multi-Output Linear-Probe Coreset / stage2_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, evaluate, sweep_baseline, make_report
7
8SEEDS=tuple(range(8))
9GRID=[{'lr':1e-3,'epochs':10},{'lr':3e-3,'epochs':10},{'lr':1e-2,'epochs':10}]
10
11def fit_ls(X,Y,w=None):
12 if w is None: w=np.ones(len(X))
13 A=X.T@(w[:,None]*X); B=Y.T@(w[:,None]*X)
14 return B@np.linalg.pinv(A)
15
16def coreset(X,Y,tol=1e-10):
17 # Preserve normal equations at the full-data least-squares optimum.
18 W=fit_ls(X,Y); n,d=X.shape; m=Y.shape[1]
19 R=((Y-X@W.T)[:,:,None]*X[:,None,:]).reshape(n,m*d)
20 w=np.ones(n); active=list(range(n)); rank=np.linalg.matrix_rank(X,tol=1e-6*np.linalg.norm(X,2))
21 target=max(1,(m+1)*rank)
22 while len(active)>target:
23 block=np.asarray(active[:min(len(active),m*d+1)])
24 _,s,vh=np.linalg.svd(R[block].T,full_matrices=True)
25 c=vh[-1]
26 if np.linalg.norm(c)<tol: break
27 if np.all(c<=tol): c=-c
28 pos=c>tol
29 if not np.any(pos): break
30 t=np.min(w[block[pos]]/c[pos]); w[block]-=t*c
31 w[np.abs(w)<1e-10]=0
32 active=[i for i in active if w[i]>1e-10]
33 idx=np.asarray(active); return idx,w[idx],W,rank
34
35def seed_run(seed,cfg,idea):
36 torch.manual_seed(seed); np.random.seed(seed)
37 ds=get_dataset('tabular',seed,n_train=400,n_test=200)
38 # Canonical training path; the coreset intervention is applied to the trained readout.
39 net,_,hist=train_model(make_model('mlp_tiny',ds['input_shape'],ds['out_dim']),ds,
40 epochs=cfg['epochs'],lr=cfg['lr'],batch=128)
41 net.eval()
42 with torch.no_grad():
43 # mlp_tiny is Sequential: Linear-ReLU-Linear-ReLU-Linear; use penultimate activations.
44 dev=next(net.parameters()).device
45 xtr=ds['xtr'].to(dev); xte=ds['xte'].to(dev)
46 emb=net[:-1](xtr).detach().cpu().numpy()
47 et=net[:-1](xte).detach().cpu().numpy()
48 ytr=ds['ytr'].detach().cpu().numpy().reshape(-1,1)
49 yte=ds['yte'].detach().cpu().numpy().reshape(-1,1)
50 if idea:
51 idx,w,W,r=coreset(emb,ytr)
52 W2=fit_ls(emb[idx],ytr[idx],w)
53 support=len(idx)
54 else:
55 W2=fit_ls(emb,ytr); idx=np.arange(len(emb)); w=np.ones(len(emb)); r=np.linalg.matrix_rank(emb); support=len(idx)
56 pred=et@W2.T
57 mse=float(np.mean((pred-yte)**2))
58 fullW=fit_ls(emb,ytr)
59 normal=np.linalg.norm(((ytr[idx]-emb[idx]@fullW.T).T@(w[:,None]*emb[idx])) if idea else (ytr-emb@fullW.T).T@emb)
60 denom=max(np.linalg.norm(ytr.T@emb),1e-12)
61 return mse, {'support':support,'rank':int(r),'normal_residual_rel':float(normal/denom),'embedding_dim':int(emb.shape[1]),'train_final_loss':float(hist[-1]) if hist else None}
62
63def main():
64 # Common cache makes baseline and idea paired while preserving each system's own readout.
65 cache={}
66 def fn(cfg,idea):
67 def run(seed):
68 key=(seed,cfg['lr'],cfg['epochs'])
69 if key not in cache: cache[key]={}
70 # train separately per system to satisfy system parity; deterministic initialization/data
71 v,meta=seed_run(seed,cfg,idea)
72 cache[key][('idea' if idea else 'base')]=meta
73 return v
74 return run
75 base=sweep_baseline(lambda c: fn(c,False),GRID)
76 # Same grid on idea side: parity, and choose best using its 4-seed mean.
77 tried=[]
78 best=None; bm=float('inf')
79 for c in GRID:
80 r=evaluate(fn(c,True),seeds=(0,1,2,3)); tried.append({'cfg':c,'mean':r['mean']})
81 if r['mean']<bm: bm=r['mean']; best=c
82 idea=evaluate(fn(best,True),seeds=SEEDS)
83 sig=[]
84 for s in SEEDS:
85 # metadata was overwritten only by same-side repeated call, but deterministic and same values.
86 sig.append(cache[(s,best['lr'],best['epochs'])].get('idea',{}))
87 mean_support=float(np.mean([x['support'] for x in sig])); mean_rank=float(np.mean([x['rank'] for x in sig]))
88 predicted=(1+1)*mean_rank; observed=mean_support
89 report=make_report('tabular','mlp_tiny',{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':base['full']},idea,
90 {'prediction':'support <= (m+1)r and near exact normal-equation preservation at trained embedding scale',
91 'predicted_support_bound':predicted,'observed_mean_support':observed,
92 'observed_mean_normal_residual_rel':float(np.mean([x['normal_residual_rel'] for x in sig])),
93 'confirmed':bool(observed <= predicted+1e-6 and np.mean([x['normal_residual_rel'] for x in sig])<1e-6),
94 'idea_sweep':tried,'per_seed':sig})
95 report['idea']['best_cfg']=best
96 report['runtime_sec']=None
97 json.dump(report,open('bench_report.json','w'),indent=2)
98 print(json.dumps(report,indent=2))
99if __name__=='__main__': main()