Log-Scale Self-Similar Activation / experiment.py
Failed on benchmark
1import json, math, random, time
2from pathlib import Path
3import numpy as np
4
5SEED = 1234
6np.random.seed(SEED)
7random.seed(SEED)
8
9
10def h_np(z, eta, nu, sigma, kmin=-40, kmax=40):
11 z = np.asarray(z, dtype=np.float64)
12 kk = np.ceil(-np.log(z) / np.log(sigma)).astype(np.int64)
13 kk = np.clip(kk, kmin, kmax)
14 sk = sigma ** kk
15 y = ((eta-nu)/(sigma-1.0))*sk*z + (-eta+sigma*nu)/(sigma-1.0)
16 return y, kk
17
18
19def mechanism_checks():
20 rows = []
21 for sigma in [1.5, 2.0, 3.0, 4.0]:
22 eta, nu = 1.7, 0.35
23 # Prediction 1: k changes at b_k=sigma^{-k}.
24 boundary_err = []
25 boundary_fail = 0
26 for k in range(-5, 6):
27 b = sigma ** (-k)
28 kl = h_np([b*(1-1e-10)], eta, nu, sigma)[1][0]
29 kr = h_np([b*(1+1e-10)], eta, nu, sigma)[1][0]
30 if not (kl == k+1 and kr == k):
31 boundary_fail += 1
32 boundary_err.append(abs((sigma**(-k))-b)/b)
33 # Prediction 2: slope ratio between neighboring branches is sigma.
34 slopes=[]
35 for k in range(-4, 5):
36 z = sigma**(-k) * math.sqrt(1.0/sigma)
37 d = z*1e-6
38 yp = h_np([z+d], eta, nu, sigma)[0][0]
39 ym = h_np([z-d], eta, nu, sigma)[0][0]
40 slopes.append((yp-ym)/(2*d))
41 ratios=np.array(slopes[1:])/np.array(slopes[:-1])
42 slope_err=float(np.max(np.abs(ratios-sigma)))
43 # Prediction 3: equivariance error is roundoff, when clipping is inactive.
44 z=np.logspace(-3,3,2000)
45 y,_=h_np(z,eta,nu,sigma,kmin=-100,kmax=100)
46 ys,_=h_np(z/sigma,eta/sigma,nu/sigma,sigma,kmin=-100,kmax=100)
47 eq_rel=float(np.max(np.abs(ys-y/sigma)/(1e-12+np.abs(y/sigma))))
48 rows.append({"sigma":sigma,
49 "boundary_relative_error":float(max(boundary_err)),
50 "boundary_failures":boundary_fail,
51 "slope_ratio_observed_mean":float(np.mean(ratios)),
52 "slope_ratio_observed_min":float(np.min(ratios)),
53 "slope_ratio_observed_max":float(np.max(ratios)),
54 "slope_ratio_predicted":sigma,
55 "slope_ratio_max_abs_error":slope_err,
56 "equivariance_max_relative_error":eq_rel})
57 return rows
58
59
60def signed_log_activation_torch(x, eta=1.0, nu=0.0, sigma=2.0, eps=1e-6,
61 kmin=-12, kmax=12):
62 import torch
63 z=torch.abs(x)+eps
64 # Straight-through k: forward integer bins, derivative zero through k.
65 k=torch.ceil(-torch.log(z)/math.log(sigma)).clamp(kmin,kmax)
66 sk=torch.exp(k*math.log(sigma))
67 y=((eta-nu)/(sigma-1))*sk*z + (-eta+sigma*nu)/(sigma-1)
68 return torch.sign(x)*y
69
70
71def mlp_experiment():
72 import torch
73 from sklearn.datasets import load_digits
74 from sklearn.model_selection import train_test_split
75 from sklearn.preprocessing import StandardScaler
76 torch.manual_seed(SEED); np.random.seed(SEED)
77 device = "cuda" if torch.cuda.is_available() else "cpu"
78 try:
79 if device == "cuda": torch.cuda.empty_cache()
80 except Exception:
81 device="cpu"
82 d=load_digits()
83 X=d.data.astype(np.float32); y=d.target.astype(np.int64)
84 X=StandardScaler().fit_transform(X).astype(np.float32)
85 Xtr,Xv,ytr,yv=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
86 Xt=torch.tensor(Xtr,device=device); yt=torch.tensor(ytr,device=device)
87 Xval=torch.tensor(Xv,device=device); yval=torch.tensor(yv,device=device)
88 class Net(torch.nn.Module):
89 def __init__(self, kind):
90 super().__init__(); self.kind=kind
91 self.l1=torch.nn.Linear(64,96); self.l2=torch.nn.Linear(96,96); self.l3=torch.nn.Linear(96,10)
92 def forward(self,x):
93 x=self.l1(x)
94 x=torch.nn.functional.relu(x) if self.kind=='relu' else signed_log_activation_torch(x,eta=1.,nu=0.,sigma=2.)
95 x=self.l2(x)
96 x=torch.nn.functional.relu(x) if self.kind=='relu' else signed_log_activation_torch(x,eta=1.,nu=0.,sigma=2.)
97 return self.l3(x)
98 out={}
99 for kind in ['relu','logscale']:
100 torch.manual_seed(SEED); net=Net(kind).to(device)
101 opt=torch.optim.Adam(net.parameters(),lr=2e-3)
102 t0=time.time()
103 for step in range(400):
104 idx=torch.randint(0,len(Xt),(64,),device=device)
105 loss=torch.nn.functional.cross_entropy(net(Xt[idx]),yt[idx])
106 opt.zero_grad(); loss.backward(); opt.step()
107 with torch.no_grad():
108 logits=net(Xval); acc=(logits.argmax(1)==yval).float().mean().item()
109 scaled=[]
110 for q in [-2,-1,0,1,2]:
111 scaled.append((net(Xval*(2.0**q)).argmax(1)==yval).float().mean().item())
112 out[kind]={"validation_accuracy":acc,"rescaling_accuracy_q_-2..2":scaled,
113 "mean_scaled_accuracy":float(np.mean(scaled)),"seconds":time.time()-t0}
114 return {"device":device,"results":out}
115
116if __name__ == '__main__':
117 result={"seed":SEED,"mechanism_checks":mechanism_checks()}
118 try:
119 result["mlp"] = mlp_experiment()
120 except Exception as e:
121 result["mlp_error"] = repr(e)
122 Path('results.json').write_text(json.dumps(result,indent=2))
123 print(json.dumps(result,indent=2))