Pole-safe rational neural layer / pole_safe_experiment.py

✓✓ Beats tuned baseline

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