Bernstein resolvent activation / experiment.py
Mechanism failed
1import json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6from scipy.special import iv, jv
7from scipy.optimize import brentq
8from sklearn.datasets import load_digits
9from sklearn.model_selection import train_test_split
10from sklearn.preprocessing import StandardScaler
11
12SEED = 440
13random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
14try:
15 device = 'cuda' if torch.cuda.is_available() else 'cpu'
16 if device == 'cuda': torch.zeros(1, device='cuda')
17except Exception:
18 device = 'cpu'
19
20class BernsteinResolvent(nn.Module):
21 def __init__(self, K=8, tau=0.5, eps=1e-4):
22 super().__init__(); self.tau=tau; self.eps=eps
23 poles=torch.exp(torch.linspace(math.log(.03), math.log(30.), K))
24 self.ua=nn.Parameter(torch.tensor(math.log(math.expm1(.2))))
25 self.uc=nn.Parameter(torch.full((K,), math.log(math.expm1(.15))))
26 self.register_buffer('poles', poles)
27 def forward(self,x):
28 y=self.eps+F.relu(x); s=y.pow(self.tau)
29 a=F.softplus(self.ua); c=F.softplus(self.uc)
30 return a*s+(c*s[...,None]/(s[...,None]+self.poles)).sum(dim=-1)
31
32class MLP(nn.Module):
33 def __init__(self,kind='gelu',K=8,tau=.5):
34 super().__init__(); self.kind=kind
35 self.layers=nn.ModuleList([nn.Linear(64,128),nn.Linear(128,128),nn.Linear(128,128)])
36 self.out=nn.Linear(128,10)
37 self.act=BernsteinResolvent(K,tau) if kind=='resolvent' else nn.GELU()
38 def forward(self,x):
39 for layer in self.layers: x=self.act(layer(x))
40 return self.out(x)
41
42def math_checks():
43 # Check W_nu partial fraction against the Bessel ratio at positive s.
44 nu=.3; ss=np.array([.01,.1,1.,10.])
45 # Locate positive zeros for the fractional-order Bessel function J_{nu+1}.
46 order=nu+1; zeros=[]
47 step=math.pi/8
48 left=1e-7; fl=jv(order,left)
49 while len(zeros)<500:
50 right=left+step; fr=jv(order,right)
51 if fl*fr<0: zeros.append(brentq(lambda z:jv(order,z),left,right))
52 left,fl=right,fr
53 lam=np.asarray(zeros)**2
54 approx=2*(nu+1)+2*np.array([np.sum(s/(s+lam)) for s in ss])
55 exact=np.array([math.sqrt(s)*iv(nu,math.sqrt(s))/iv(nu+1,math.sqrt(s)) for s in ss])
56 relerr=float(np.max(np.abs(approx-exact)/np.abs(exact)))
57 # Dense autograd derivative signs for finite analogue on x>=0.
58 x=torch.linspace(0,100,4000,requires_grad=True)
59 vals=BernsteinResolvent(K=8,tau=.5)(x[:,None]).squeeze()
60 d=torch.autograd.grad(vals.sum(),x,create_graph=True)[0]
61 d2=torch.autograd.grad(d.sum(),x)[0]
62 mono=float(d.min()); conc=float(d2.max())
63 # Deliberately outside range: numerical concavity check for tau=.75.
64 x2=torch.linspace(1e-3,100,4000,requires_grad=True)
65 v2=BernsteinResolvent(K=8,tau=.75)(x2[:,None]).squeeze()
66 d21=torch.autograd.grad(v2.sum(),x2,create_graph=True)[0]
67 d22=torch.autograd.grad(d21.sum(),x2)[0]
68 return {'bessel_partial_fraction_max_relative_error':relerr,
69 'tau_0.5_min_first_derivative':mono,'tau_0.5_max_second_derivative':conc,
70 'tau_0.75_max_second_derivative':float(d22.max())}
71
72def train(kind,tau=.5,steps=500):
73 data=load_digits(); X=data.data.astype('float32'); y=data.target
74 X=StandardScaler().fit_transform(X).astype('float32')
75 Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
76 xt=torch.tensor(Xtr,device=device); yt=torch.tensor(ytr,device=device)
77 xv=torch.tensor(Xte,device=device); yv=torch.tensor(yte,device=device)
78 torch.manual_seed(SEED)
79 model=MLP(kind,tau=tau).to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3)
80 losses=[]; gradnorms=[]; bs=96; g=torch.Generator(device=device).manual_seed(SEED)
81 for step in range(steps):
82 ind=torch.randint(0,len(xt),(bs,),generator=g,device=device)
83 opt.zero_grad(set_to_none=True); loss=F.cross_entropy(model(xt[ind]),yt[ind]); loss.backward()
84 gn=float(torch.sqrt(sum((p.grad.detach()**2).sum() for p in model.parameters() if p.grad is not None)).cpu())
85 gradnorms.append(gn); losses.append(float(loss.detach().cpu())); opt.step()
86 with torch.no_grad():
87 pred=model(xv).argmax(1); acc=float((pred==yv).float().mean().cpu())
88 return {'final_train_loss':losses[-1],'test_accuracy':acc,
89 'mean_grad_norm':float(np.mean(gradnorms)),'grad_norm_std':float(np.std(gradnorms)),
90 'max_grad_norm':float(np.max(gradnorms))}
91
92def main():
93 result={'seed':SEED,'device':device,'math':math_checks(),'experiment':{}}
94 for kind,tau in [('gelu',.5),('resolvent',.5),('resolvent',.75)]:
95 key=kind if kind=='gelu' else kind+'_tau_'+str(tau).replace('.','_')
96 result['experiment'][key]=train(kind,tau)
97 print(json.dumps(result,indent=2))
98
99if __name__=='__main__': main()