import json, random, time import numpy as np import torch from torch import nn from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler SEED=582 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(8) try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device.type=='cuda': torch.cuda.get_device_name(0) except Exception: device=torch.device('cpu') # Exact toy verification: a ReLU axis whose preactivation is strictly one-sided # is either zero or affine on the entire patch, and can be absorbed downstream. def identity_check(): rng=np.random.default_rng(SEED) # Bounded task patch and deliberately chosen axes: unit 0 is positive, # unit 1 is negative, and units 2-4 cross zero. X=rng.uniform(-1.0,1.0,size=(300,4)) A=rng.normal(size=(5,4)); b=np.zeros(5) A[0]=np.array([0.2,-0.1,0.1,0.05]); b[0]=2.0 A[1]=np.array([0.2,-0.1,0.1,0.05]); b[1]=-2.0 b[2:]=0.0 z=X@A.T+b; H=np.maximum(z,0); W=rng.normal(size=(3,5)); c=rng.normal(size=3) y=H@W.T+c # Positive ReLU is affine, negative ReLU is identically zero. y_abs=(np.maximum(X@A[2:].T+b[2:],0)@W[:,2:].T +(X@A[0])[:,None]*W[:,0][None,:]+c+b[0]*W[:,0]) maxerr=float(np.max(np.abs(y-y_abs))) qlo=np.quantile(z,.01,axis=0); qhi=np.quantile(z,.99,axis=0) crossing=((qlo<=0)&(qhi>=0)).astype(int) return {'max_absorption_error':maxerr,'q01':qlo.tolist(),'q99':qhi.tolist(), 'crossing_mask':crossing.tolist(), 'positive_or_negative_units':int((crossing==0).sum()), 'expected_crossing_mask':[0,0,1,1,1]} class MLP(nn.Module): def __init__(self): super().__init__(); self.layers=nn.ModuleList([nn.Linear(64,96),nn.Linear(96,96),nn.Linear(96,64),nn.Linear(64,10)]) def forward(self,x): for i,l in enumerate(self.layers): x=l(x) if i<3: x=torch.relu(x) return x def accuracy(model,X,y): model.eval() with torch.no_grad(): return float((model(X).argmax(1)==y).float().mean().cpu()) def scores(model, X): hs=[]; zs=[]; x=X for i,l in enumerate(model.layers[:-1]): z=l(x); zs.append(z.detach().cpu().numpy()); x=torch.relu(z); hs.append(x.detach().cpu().numpy()) # outgoing norm for each hidden layer is column norm of next Linear out=[] for i in range(3): out.append(model.layers[i+1].weight.detach().cpu().numpy().T) return zs,hs,out def make_masks(model,X,keep_frac,method): zs,hs,out=scores(model,X); masks=[] for z,h,w in zip(zs,hs,out): n=z.shape[1]; k=max(1,int(round(n*keep_frac))) if method=='magnitude': score=np.linalg.norm(w,axis=1) # keep largest elif method=='activation_mean': score=np.mean(np.abs(h),axis=0) else: lo=np.quantile(z,.01,axis=0); hi=np.quantile(z,.99,axis=0) visible=((lo<=0)&(hi>=0)&(np.linalg.norm(w,axis=1)>1e-8)) # prioritize visible; within visible prefer larger downstream use score=visible.astype(float)*1e6 + np.linalg.norm(w,axis=1) idx=np.argsort(score)[-k:] mask=np.zeros(n,dtype=np.float32); mask[idx]=1; masks.append(mask) return masks def apply_masks(model,masks): # masks are hidden-unit masks; zero incoming/output connections and biases. with torch.no_grad(): for i,m in enumerate(masks): mt=torch.tensor(m,device=device) model.layers[i].weight.mul_(mt[:,None]); model.layers[i].bias.mul_(mt) model.layers[i+1].weight.mul_(mt[None,:]) def finetune(model,X,y,epochs=3): model.train(); opt=torch.optim.Adam(model.parameters(),lr=2e-3); lossfn=nn.CrossEntropyLoss() for _ in range(epochs): p=torch.randperm(len(X),device=device) for start in range(0,len(X),128): ix=p[start:start+128]; loss=lossfn(model(X[ix]),y[ix]); opt.zero_grad(); loss.backward(); opt.step() def main(): d=load_digits(); X=d.data.astype('float32'); y=d.target.astype('int64'); X=StandardScaler().fit_transform(X).astype('float32') Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y) Xtr=torch.tensor(Xtr,device=device); Xte=torch.tensor(Xte,device=device); ytr=torch.tensor(ytr,device=device); yte=torch.tensor(yte,device=device) base=MLP().to(device); opt=torch.optim.Adam(base.parameters(),lr=2e-3); ce=nn.CrossEntropyLoss(); t=time.time() for ep in range(25): p=torch.randperm(len(Xtr),device=device) for st in range(0,len(Xtr),128): ix=p[st:st+128]; loss=ce(base(Xtr[ix]),ytr[ix]); opt.zero_grad(); loss.backward(); opt.step() base_acc=accuracy(base,Xte,yte); base_loss=float(ce(base(Xte),yte).detach().cpu()) results={'device':str(device),'identity_check':identity_check(),'baseline':{'accuracy':base_acc,'loss':base_loss,'train_seconds':time.time()-t},'pruning':{}} for frac in [0.75,0.50,0.25]: for method in ['magnitude','activation_mean','task_visible']: m=MLP().to(device); m.load_state_dict(base.state_dict()); masks=make_masks(m,Xtr,frac,method); apply_masks(m,masks) before=accuracy(m,Xte,yte); finetune(m,Xtr,ytr,epochs=3); after=accuracy(m,Xte,yte) 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]} with open('results.json','w') as f: json.dump(results,f,indent=2) print(json.dumps(results,indent=2)) if __name__=='__main__': main()