Bernstein resolvent activation / experiment.py

Mechanism failed

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