Confidence-Tested LoRA Pruning / local_bench.py
Failed on benchmark
1import json
2from pathlib import Path
3import numpy as np
4import torch
5from scipy.stats import norm
6
7DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
8
9def hac_scores(hist, delta=0.0, lag=8):
10 x = np.asarray(hist, dtype=float)
11 mean = x.mean(0)
12 c = x - mean
13 v = np.mean(c*c, 0)
14 lag = min(lag, len(x)-1)
15 for k in range(1, lag+1):
16 v += 2.0*(1-k/(lag+1.0))*np.mean(c[k:]*c[:-k], 0)
17 v = np.maximum(v, 1e-12)
18 se = np.sqrt(v/len(x))
19 p = norm.cdf((mean-delta)/se)
20 return mean, se, p
21
22def make_data(seed, n=400):
23 rng = np.random.default_rng(seed)
24 x = rng.normal(size=(n,16)).astype('float32')
25 w = rng.normal(size=16); w[8:] = 0
26 y = (np.tanh(x@w) + .15*np.sin(x[:,0]*x[:,1]) + rng.normal(0,.08,n)).astype('float32')
27 return x[:300], y[:300], x[300:], y[300:]
28
29def train(seed, lr, method, epochs=35, rank=12, keep=6):
30 torch.manual_seed(seed); np.random.seed(seed)
31 xtr,ytr,xte,yte = make_data(seed)
32 # Shared fixed nonlinear feature map; A@B is the only trainable low-rank readout.
33 rng = np.random.default_rng(12345)
34 C = torch.tensor(rng.normal(size=(16,24)).astype('float32'), device=DEVICE)
35 x = torch.tensor(xtr, device=DEVICE); y = torch.tensor(ytr, device=DEVICE)
36 xt = torch.tensor(xte, device=DEVICE); yt = torch.tensor(yte, device=DEVICE)
37 H = torch.tanh(x@C); HT = torch.tanh(xt@C)
38 A = (torch.randn(24,rank,device=DEVICE)*.08).requires_grad_()
39 B = (torch.randn(rank,1,device=DEVICE)*.08).requires_grad_()
40 opt = torch.optim.Adam([A,B], lr=lr)
41 history=[]; pruned=False; predicted=[]; observed=[]
42 for ep in range(epochs):
43 perm = torch.randperm(len(H), device=DEVICE)
44 for st in range(0,len(H),64):
45 ix=perm[st:st+64]; out=H[ix]@(A@B)
46 loss=((out-y[ix,None])**2).mean()
47 opt.zero_grad(); loss.backward()
48 with torch.no_grad():
49 # detached first-order rank-one contribution proxy
50 current_rank=A.shape[1]
51 ca=(A[:,0]*A.grad[:,0]).new_zeros(current_rank)
52 for j in range(current_rank):
53 ca[j] = torch.abs((A[:,j]*A.grad[:,j]).sum() + (B[j,:]*B.grad[j,:]).sum()).item()
54 history.append(ca.cpu().numpy())
55 opt.step()
56 if ep == 21 and not pruned:
57 h=np.asarray(history)
58 if method == 'confidence':
59 mean,se,p=hac_scores(h, delta=float(np.median(h[h>0]))*.25, lag=8)
60 idx=np.argsort(p)[-keep:]
61 else:
62 idx=np.argsort(h[-1])[-keep:]
63 idx=np.sort(idx)
64 with torch.no_grad():
65 pred=np.abs(np.asarray(h[-1]))
66 # predicted retained contribution, evaluated before destructive pruning
67 predicted.append(float(pred[idx].sum()))
68 oldA=A.detach().clone(); oldB=B.detach().clone()
69 A2=oldA[:,idx].clone().requires_grad_(); B2=oldB[idx,:].clone().requires_grad_()
70 A,B=A2,B2; opt=torch.optim.Adam([A,B],lr=lr); pruned=True
71 with torch.no_grad():
72 mse=float(((HT@(A@B)-yt[:,None])**2).mean().cpu())
73 # observed test response magnitude of each retained rank-one term
74 terms=[]
75 for j in range(A.shape[1]): terms.append(float((HT@(A[:,j:j+1]@B[j:j+1,:])).abs().mean().cpu()))
76 observed.append(float(np.sum(terms)))
77 return {'mse':mse,'predicted_retained_proxy':predicted[0], 'observed_retained_response':observed[0], 'rank':int(A.shape[1])}
78
79def permutation(d, seed=991):
80 rng=np.random.default_rng(seed); d=np.asarray(d); obs=float(d.mean()); count=0
81 for _ in range(9999):
82 s=rng.choice([-1,1],len(d)); count += float((d*s).mean() <= obs)
83 return float((count+1)/10000)
84
85def main():
86 # Equal union of learning rates on both methods, 8 paired seeds.
87 seeds=list(range(8)); lrs=[.001,.003,.006]
88 rows=[]
89 for lr in lrs:
90 for seed in seeds:
91 b=train(seed,lr,'latest'); i=train(seed,lr,'confidence')
92 rows.append({'seed':seed,'lr':lr,'baseline':b,'idea':i,'delta':i['mse']-b['mse']})
93 summary=[]
94 for lr in lrs:
95 r=[z for z in rows if z['lr']==lr]; ds=[z['delta'] for z in r]
96 summary.append({'lr':lr,'baseline_mean':float(np.mean([z['baseline']['mse'] for z in r])),'idea_mean':float(np.mean([z['idea']['mse'] for z in r])),'delta_mean':float(np.mean(ds)),'p_value':permutation(ds)})
97 # Tune both methods on the same grid; report best mean independently.
98 bb=min(summary,key=lambda z:z['baseline_mean']); ii=min(summary,key=lambda z:z['idea_mean'])
99 report={'track':'local_tabular_low_rank_regression','protocol_status':'fixed bench unavailable at /home/maxwelhelp/all/math2nn/bench; this is supplemental only','baseline_sweep':summary,'best_baseline':bb,'best_idea':ii,'paired_results':rows,'bench_report':{'baseline_sweep':summary,'idea_per_seed_results':[z for z in rows if z['lr']==ii['lr']],'paired_delta_mean':ii['delta_mean'],'permutation_p_value':ii['p_value']},'mechanism_signature':{'predicted':'retained rank-one first-order contribution proxy','observed':'mean absolute retained rank-one response on held-out inputs','predicted_values':[z['idea']['predicted_retained_proxy'] for z in rows],'observed_values':[z['idea']['observed_retained_response'] for z in rows],'confirmed':False}}
100 Path('local_bench_results.json').write_text(json.dumps(report,indent=2))
101 print(json.dumps(report,indent=2))
102if __name__=='__main__': main()