Exact Multi-Output Linear-Probe Coreset / stage2_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, 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()