Confidence-Tested LoRA Pruning / official_bench.py

Failed on benchmark

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