Confidence-Tested LoRA Pruning / local_bench.py

Failed on benchmark

Raw ⬇ ZIP
  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()