Task-Visible Axis Pruning / run_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, random, time
  2import numpy as np
  3import torch
  4from torch import nn
  5from sklearn.datasets import load_digits
  6from sklearn.model_selection import train_test_split
  7from sklearn.preprocessing import StandardScaler
  8
  9SEED=582
 10random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 11torch.set_num_threads(8)
 12try:
 13    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 14    if device.type=='cuda': torch.cuda.get_device_name(0)
 15except Exception:
 16    device=torch.device('cpu')
 17
 18# Exact toy verification: a ReLU axis whose preactivation is strictly one-sided
 19# is either zero or affine on the entire patch, and can be absorbed downstream.
 20def identity_check():
 21    rng=np.random.default_rng(SEED)
 22    # Bounded task patch and deliberately chosen axes: unit 0 is positive,
 23    # unit 1 is negative, and units 2-4 cross zero.
 24    X=rng.uniform(-1.0,1.0,size=(300,4))
 25    A=rng.normal(size=(5,4)); b=np.zeros(5)
 26    A[0]=np.array([0.2,-0.1,0.1,0.05]); b[0]=2.0
 27    A[1]=np.array([0.2,-0.1,0.1,0.05]); b[1]=-2.0
 28    b[2:]=0.0
 29    z=X@A.T+b; H=np.maximum(z,0); W=rng.normal(size=(3,5)); c=rng.normal(size=3)
 30    y=H@W.T+c
 31    # Positive ReLU is affine, negative ReLU is identically zero.
 32    y_abs=(np.maximum(X@A[2:].T+b[2:],0)@W[:,2:].T
 33           +(X@A[0])[:,None]*W[:,0][None,:]+c+b[0]*W[:,0])
 34    maxerr=float(np.max(np.abs(y-y_abs)))
 35    qlo=np.quantile(z,.01,axis=0); qhi=np.quantile(z,.99,axis=0)
 36    crossing=((qlo<=0)&(qhi>=0)).astype(int)
 37    return {'max_absorption_error':maxerr,'q01':qlo.tolist(),'q99':qhi.tolist(),
 38            'crossing_mask':crossing.tolist(),
 39            'positive_or_negative_units':int((crossing==0).sum()),
 40            'expected_crossing_mask':[0,0,1,1,1]}
 41
 42class MLP(nn.Module):
 43    def __init__(self):
 44        super().__init__(); self.layers=nn.ModuleList([nn.Linear(64,96),nn.Linear(96,96),nn.Linear(96,64),nn.Linear(64,10)])
 45    def forward(self,x):
 46        for i,l in enumerate(self.layers):
 47            x=l(x)
 48            if i<3: x=torch.relu(x)
 49        return x
 50
 51def accuracy(model,X,y):
 52    model.eval()
 53    with torch.no_grad(): return float((model(X).argmax(1)==y).float().mean().cpu())
 54
 55def scores(model, X):
 56    hs=[]; zs=[]; x=X
 57    for i,l in enumerate(model.layers[:-1]):
 58        z=l(x); zs.append(z.detach().cpu().numpy()); x=torch.relu(z); hs.append(x.detach().cpu().numpy())
 59    # outgoing norm for each hidden layer is column norm of next Linear
 60    out=[]
 61    for i in range(3): out.append(model.layers[i+1].weight.detach().cpu().numpy().T)
 62    return zs,hs,out
 63
 64def make_masks(model,X,keep_frac,method):
 65    zs,hs,out=scores(model,X); masks=[]
 66    for z,h,w in zip(zs,hs,out):
 67        n=z.shape[1]; k=max(1,int(round(n*keep_frac)))
 68        if method=='magnitude': score=np.linalg.norm(w,axis=1) # keep largest
 69        elif method=='activation_mean': score=np.mean(np.abs(h),axis=0)
 70        else:
 71            lo=np.quantile(z,.01,axis=0); hi=np.quantile(z,.99,axis=0)
 72            visible=((lo<=0)&(hi>=0)&(np.linalg.norm(w,axis=1)>1e-8))
 73            # prioritize visible; within visible prefer larger downstream use
 74            score=visible.astype(float)*1e6 + np.linalg.norm(w,axis=1)
 75        idx=np.argsort(score)[-k:]
 76        mask=np.zeros(n,dtype=np.float32); mask[idx]=1; masks.append(mask)
 77    return masks
 78
 79def apply_masks(model,masks):
 80    # masks are hidden-unit masks; zero incoming/output connections and biases.
 81    with torch.no_grad():
 82        for i,m in enumerate(masks):
 83            mt=torch.tensor(m,device=device)
 84            model.layers[i].weight.mul_(mt[:,None]); model.layers[i].bias.mul_(mt)
 85            model.layers[i+1].weight.mul_(mt[None,:])
 86
 87def finetune(model,X,y,epochs=3):
 88    model.train(); opt=torch.optim.Adam(model.parameters(),lr=2e-3); lossfn=nn.CrossEntropyLoss()
 89    for _ in range(epochs):
 90        p=torch.randperm(len(X),device=device)
 91        for start in range(0,len(X),128):
 92            ix=p[start:start+128]; loss=lossfn(model(X[ix]),y[ix]); opt.zero_grad(); loss.backward(); opt.step()
 93
 94def main():
 95    d=load_digits(); X=d.data.astype('float32'); y=d.target.astype('int64'); X=StandardScaler().fit_transform(X).astype('float32')
 96    Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
 97    Xtr=torch.tensor(Xtr,device=device); Xte=torch.tensor(Xte,device=device); ytr=torch.tensor(ytr,device=device); yte=torch.tensor(yte,device=device)
 98    base=MLP().to(device); opt=torch.optim.Adam(base.parameters(),lr=2e-3); ce=nn.CrossEntropyLoss(); t=time.time()
 99    for ep in range(25):
100        p=torch.randperm(len(Xtr),device=device)
101        for st in range(0,len(Xtr),128):
102            ix=p[st:st+128]; loss=ce(base(Xtr[ix]),ytr[ix]); opt.zero_grad(); loss.backward(); opt.step()
103    base_acc=accuracy(base,Xte,yte); base_loss=float(ce(base(Xte),yte).detach().cpu())
104    results={'device':str(device),'identity_check':identity_check(),'baseline':{'accuracy':base_acc,'loss':base_loss,'train_seconds':time.time()-t},'pruning':{}}
105    for frac in [0.75,0.50,0.25]:
106        for method in ['magnitude','activation_mean','task_visible']:
107            m=MLP().to(device); m.load_state_dict(base.state_dict()); masks=make_masks(m,Xtr,frac,method); apply_masks(m,masks)
108            before=accuracy(m,Xte,yte); finetune(m,Xtr,ytr,epochs=3); after=accuracy(m,Xte,yte)
109            results['pruning'][f'{method}_{int(frac*100)}']={'accuracy_before':before,'accuracy_after_3ep':after,'retained_fraction':frac,'retained_units':[int(x.sum()) for x in masks]}
110    with open('results.json','w') as f: json.dump(results,f,indent=2)
111    print(json.dumps(results,indent=2))
112if __name__=='__main__': main()