Pole-safe rational neural layer / pole_safe_experiment.py
Beats tuned baseline
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 297
8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
9DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
10
11
12def make_operator(d=32):
13 r = np.random.default_rng(SEED)
14 qm2 = r.normal(size=(d,d))/math.sqrt(d) + .75*np.eye(d)
15 qm1 = r.normal(size=(d,d))/math.sqrt(d)
16 q0 = r.normal(size=(d,d))/math.sqrt(d)
17 return qm2, qm1, q0
18
19
20def math_check():
21 qm2,qm1,q0 = make_operator(); r=np.random.default_rng(SEED+1)
22 v=r.normal(size=32); v/=np.linalg.norm(v)
23 ds=np.array([1e-1,3e-2,1e-2,3e-3,1e-3,3e-4,1e-4,3e-5,1e-5])
24 unsafe=[]; safe=[]; errors=[]
25 for d in ds:
26 q=qm2/d**2+qm1/d+q0
27 unsafe.append(np.linalg.norm(q@v))
28 safe.append(np.linalg.norm(q@(d*d*v)))
29 errors.append(np.linalg.norm(q@(d*d*v)-qm2@v))
30 unsafe_slope=np.polyfit(np.log(ds),np.log(unsafe),1)[0]
31 safe_slope=np.polyfit(np.log(ds),np.log(safe),1)[0]
32 return {'unsafe_growth_slope':float(unsafe_slope),'safe_variation_slope':float(safe_slope),
33 'safe_limit_error_at_1e-5':float(errors[-1]),'unsafe_norm_ratio':float(unsafe[-1]/unsafe[0]),
34 'safe_norm_ratio':float(safe[-1]/safe[0])}
35
36
37class Branch(nn.Module):
38 def __init__(self,d=32, mode='unsafe'):
39 super().__init__(); self.mode=mode
40 self.h=nn.Sequential(nn.Linear(8,32),nn.Tanh(),nn.Linear(32,d))
41 def forward(self,x,z):
42 v=self.h(x)
43 if self.mode=='safe':
44 # Fixed beta=1, order m=2; normalize only to prevent feature scale cheating.
45 v=v/(v.norm(dim=1,keepdim=True)+1e-6)
46 return v*(z-1.0).pow(2)
47 return v
48
49
50def run_branch(mode, q, train_x, train_z, train_y, val_x, val_z, val_y, steps=450):
51 model=Branch(mode=mode).to(DEVICE)
52 opt=torch.optim.Adam(model.parameters(),lr=2e-3)
53 qm2,qm1,q0=[torch.tensor(a,dtype=torch.float32,device=DEVICE) for a in q]
54 def forward(x,z):
55 psi=model(x,z); t=z-1.0
56 return (psi@qm2.T)/(t*t)+(psi@qm1.T)/t+psi@q0.T
57 maxact=0.; maxgrad=0.; nan_count=0
58 for _ in range(steps):
59 opt.zero_grad(); pred=forward(train_x,train_z)
60 loss=((pred-train_y)**2).mean()
61 if not torch.isfinite(loss): nan_count+=1; break
62 loss.backward()
63 g=torch.nn.utils.clip_grad_norm_(model.parameters(),1e3) if mode=='clip' else 0.
64 maxgrad=max(maxgrad,float(g)); maxact=max(maxact,float(pred.detach().norm(dim=1).max()))
65 opt.step()
66 with torch.no_grad():
67 pred=forward(val_x,val_z); val_loss=((pred-val_y)**2).mean()
68 near=(torch.abs(val_z-1)<0.002).squeeze(1)
69 near_act=float(pred[near].norm(dim=1).max())
70 return {'val_mse':float(val_loss),'max_output_norm':maxact,'near_pole_output_norm':near_act,
71 'max_grad_norm':maxgrad,'nan_steps':nan_count,'parameters':sum(p.numel() for p in model.parameters())}
72
73
74def experiment():
75 qm2,qm1,q0=make_operator(); rng=np.random.default_rng(SEED+2)
76 ntr,nva=256,256; d=32
77 xtr=torch.tensor(rng.normal(size=(ntr,8)),dtype=torch.float32,device=DEVICE)
78 xva=torch.tensor(rng.normal(size=(nva,8)),dtype=torch.float32,device=DEVICE)
79 # Deliberately include very near-pole frequencies to test stability.
80 ztr=torch.tensor(1+rng.uniform(-.05,.05,ntr),dtype=torch.float32,device=DEVICE).view(-1,1)
81 zva=torch.tensor(1+rng.uniform(-.05,.05,nva),dtype=torch.float32,device=DEVICE).view(-1,1)
82 target=nn.Sequential(nn.Linear(8,d),nn.Tanh(),nn.Linear(d,d)).to(DEVICE)
83 with torch.no_grad():
84 ytr=target(xtr); yva=target(xva)
85 q=(qm2,qm1,q0); out={}
86 # clip is a standard control; its architecture is unconstrained and differs only in update clipping.
87 for mode in ('unsafe','clip','safe'):
88 out[mode]=run_branch(mode,q,xtr,ztr,ytr,xva,zva,yva)
89 return out
90
91
92def main():
93 try:
94 math_result=math_check(); exp=experiment()
95 except RuntimeError as e:
96 if 'CUDA' in str(e) or 'cuda' in str(e).lower():
97 global DEVICE; DEVICE='cpu'; math_result=math_check(); exp=experiment()
98 else: raise
99 result={'seed':SEED,'device':DEVICE,'math_check':math_result,'experiment':exp}
100 Path('results.json').write_text(json.dumps(result,indent=2))
101 print(json.dumps(result,indent=2))
102
103if __name__=='__main__': main()