Task-Visible Axis Pruning / run_experiment.py
Failed on benchmark
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()