Log-Scale Self-Similar Activation / experiment.py

Failed on benchmark

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