Confidence-Tested LoRA Pruning / official_bench.py
Failed on benchmark
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5from scipy.stats import norm
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, evaluate, sweep_baseline, make_report
8
9TRACK='tabular'
10MODEL='mlp_tiny_factorized_first_layer'
11LR_GRID=[0.001,0.003,0.006]
12SEEDS=tuple(range(8))
13RANK=8
14KEEP=4
15EPOCHS=20
16BATCH=128
17PRUNE_EPOCH=11
18
19
20def hac_stats(samples, delta, lag=8):
21 x=np.asarray(samples, dtype=np.float64)
22 mean=x.mean(0); c=x-mean; v=np.mean(c*c,0)
23 lag=min(int(lag),len(x)-1)
24 for k in range(1,lag+1):
25 gamma=np.mean(c[k:]*c[:-k],axis=0)
26 v += 2.0*(1.0-k/(lag+1.0))*gamma
27 v=np.maximum(v,1e-12)
28 se=np.sqrt(v/len(x)); z=(mean-delta)/se; p=norm.cdf(z)
29 return mean,se,p
30
31
32class FactorizedMLP(torch.nn.Module):
33 # W = W0 + A B is a rank-one component decomposition.
34 def __init__(self, input_dim, hidden=64, rank=RANK, seed=0):
35 super().__init__()
36 g=torch.Generator().manual_seed(seed+100003)
37 self.register_buffer('w0', torch.randn(input_dim,hidden,generator=g)*0.08)
38 self.register_buffer('b0', torch.zeros(hidden))
39 self.A=torch.nn.Parameter(torch.randn(input_dim,rank,generator=g)*0.03)
40 self.B=torch.nn.Parameter(torch.randn(rank,hidden,generator=g)*0.03)
41 self.out=torch.nn.Linear(hidden,1)
42 def forward(self,x):
43 return torch.relu(x @ (self.w0+self.A@self.B)+self.b0) @ self.out.weight.t()+self.out.bias
44 def prune(self, idx):
45 with torch.no_grad():
46 oldA=self.A.detach().clone(); oldB=self.B.detach().clone()
47 self.A=torch.nn.Parameter(oldA[:,idx].clone())
48 self.B=torch.nn.Parameter(oldB[idx,:].clone())
49
50
51def train_one(seed, lr, mode, delta_fraction=0.25, collect_signature=False):
52 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
53 d=get_dataset(TRACK, seed, n_train=400, n_test=1000)
54 device='cuda' if torch.cuda.is_available() else 'cpu'
55 try:
56 net=FactorizedMLP(int(np.prod(d['input_shape'])),seed=seed).to(device)
57 x=d['xtr'].to(device); y=d['ytr'].to(device)
58 xt=d['xte'].to(device); yt=d['yte'].to(device)
59 opt=torch.optim.Adam(net.parameters(),lr=lr)
60 history=[]; selected=None; pre_proxy=None
61 for ep in range(EPOCHS):
62 perm=torch.randperm(x.shape[0],device=device)
63 for start in range(0,x.shape[0],BATCH):
64 ix=perm[start:start+BATCH]; xb=x[ix]; yb=y[ix]
65 w=net.w0+net.A@net.B
66 z=xb@w+net.b0; z.retain_grad()
67 pred=torch.relu(z) @ net.out.weight.t()+net.out.bias
68 loss=((pred-yb)**2).mean()
69 opt.zero_grad(); loss.backward()
70 with torch.no_grad():
71 # x_t,j = <grad_W, a_j b_j^T>, using grad_W = X^T grad_z.
72 gw=xb.t()@z.grad
73 vals=[]
74 for j in range(net.A.shape[1]):
75 vals.append(float(torch.abs((gw* (net.A[:,j:j+1]@net.B[j:j+1,:])).sum()).cpu()))
76 history.append(vals)
77 opt.step()
78 if ep==PRUNE_EPOCH:
79 h=np.asarray(history)
80 if mode=='confidence':
81 positive=h[h>0]
82 delta=float(np.median(positive)*delta_fraction) if positive.size else 0.0
83 mean,se,p=hac_stats(h,delta)
84 idx=np.argsort(p)[-KEEP:]
85 score_info={'mean':mean.tolist(),'se':se.tolist(),'p':p.tolist(),'delta':delta}
86 else:
87 idx=np.argsort(h[-1])[-KEEP:]
88 score_info={'latest':h[-1].tolist()}
89 idx=np.sort(idx)
90 pre_proxy=float(np.asarray(h[-1])[idx].sum())
91 net.prune(idx)
92 opt=torch.optim.Adam(net.parameters(),lr=lr)
93 selected=idx.tolist()
94 with torch.no_grad():
95 pred=net(xt); mse=float(((pred-yt)**2).mean().cpu())
96 terms=[]
97 hidden=torch.relu(xt@(net.w0+net.A@net.B)+net.b0)
98 for j in range(net.A.shape[1]):
99 terms.append(float(torch.abs((xt@net.A[:,j:j+1])@net.B[j:j+1,:]).mean().cpu()))
100 observed=float(np.sum(terms))
101 return {'metric':mse,'rank':int(net.A.shape[1]),'selected':selected,
102 'predicted_proxy':pre_proxy,'observed_response':observed,
103 'score_info':score_info if 'score_info' in locals() else {}}
104 except Exception:
105 if device=='cuda':
106 torch.cuda.empty_cache()
107 old=torch.cuda.is_available
108 raise
109
110
111def make_train(mode, cfg):
112 return lambda seed: train_one(seed,float(cfg['lr']),mode,float(cfg.get('delta_fraction',0.25)))['metric']
113
114
115def main():
116 # Baseline is tuned by the official sweep on exactly the union of idea lrs.
117 grid=[{'lr':lr,'delta_fraction':df} for lr in LR_GRID for df in [0.0]]
118 base=sweep_baseline(lambda cfg: make_train('latest',cfg),grid)
119 idea_runs={}
120 idea_summaries=[]
121 # Same lr grid, plus three a-priori confidence thresholds.
122 for cfg in [{'lr':lr,'delta_fraction':df} for lr in LR_GRID for df in [0.15,0.25,0.50]]:
123 res=evaluate(make_train('confidence',cfg),SEEDS)
124 idea_runs[str(cfg)]=res
125 idea_summaries.append({'cfg':cfg,'mean':res['mean']})
126 best_cfg=min(idea_summaries,key=lambda q:q['mean'])['cfg']
127 idea=idea_runs[str(best_cfg)]
128 # Re-run best baseline is already full-seed result returned by sweep_baseline.
129 example=train_one(0,float(best_cfg['lr']),'confidence',float(best_cfg['delta_fraction']))
130 signature={'predicted_quantity':'sum of retained rank-one gradient contributions at pruning','observed_quantity':'sum of retained rank-one absolute responses on held-out inputs','predicted_example':example['predicted_proxy'],'observed_example':example['observed_response'],'confirmed':False}
131 report=make_report(TRACK,MODEL,base,idea,signature)
132 report['idea_sweep']=idea_summaries
133 report['selected_idea_cfg']=best_cfg
134 report['protocol_note']='Official registered tabular track; baseline sweep and idea use shared lr union, 8 final paired seeds.'
135 Path('bench_report.json').write_text(json.dumps(report,indent=2))
136 print(json.dumps(report,indent=2))
137
138if __name__=='__main__': main()