Differentially Passive Neural Blocks / safe_experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import json, random, math
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5from torch.nn.utils.parametrizations import spectral_norm
 6
 7SEED=2975
 8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 9device='cuda' if torch.cuda.is_available() else 'cpu'
10torch.set_default_dtype(torch.float64)
11
12class Block(nn.Module):
13    def __init__(self, d=2, hidden=16, eta=.15, alpha=.20, idea=False):
14        super().__init__(); self.d=d; self.eta=eta; self.alpha=alpha; self.idea=idea
15        self.l1=nn.Linear(2*d,hidden); self.l2=nn.Linear(hidden,d)
16        if idea:
17            # tanh and the concatenation projection are 1-Lipschitz. Spectral
18            # normalization makes each learned linear map 1-Lipschitz.
19            self.l1=spectral_norm(self.l1); self.l2=spectral_norm(self.l2)
20            # Reserve a Lipschitz budget for the residual correction.
21            self.skip=math.sqrt(1-alpha)-eta
22        else: self.skip=1.0
23    def forward(self,z,u):
24        h=torch.tanh(self.l1(torch.cat([z,u],-1)))
25        return self.skip*z+self.eta*self.l2(h)
26
27def jac(model,z,u):
28    z=z.detach().requires_grad_(True); y=model(z[None],u[None])[0]
29    return torch.stack([torch.autograd.grad(y[i],z,retain_graph=True)[0] for i in range(y.numel())])
30def eig(a): return torch.linalg.eigvalsh((a+a.T)/2)[-1]
31def report(model, alpha=.2):
32    g=torch.Generator(device=device).manual_seed(SEED+11)
33    zs=torch.rand(64,2,generator=g,device=device)*2-1; us=torch.rand(64,2,generator=g,device=device)*2-1
34    ev=[float(eig(jac(model,z,u).T@jac(model,z,u)-(1-alpha)*torch.eye(2,device=device)).detach().cpu()) for z,u in zip(zs,us)]
35    z1=torch.tensor([.65,-.45],device=device); z2=z1+torch.tensor([1e-3,-1.2e-3],device=device); u=torch.tensor([.2,-.1],device=device)
36    ratios=[]
37    for _ in range(20):
38        old=torch.linalg.vector_norm(z2-z1); z1=model(z1[None],u[None])[0]; z2=model(z2[None],u[None])[0]
39        ratios.append(float((torch.linalg.vector_norm(z2-z1)/old).detach().cpu()))
40    return {'max_violation':max(ev),'mean_violation':float(np.mean(ev)), 'max_ratio':max(ratios),'mean_ratio':float(np.mean(ratios)), 'distance_ratio_20':float(np.prod(ratios)), 'ratios':ratios}
41
42def train(model, constrained, epochs=250):
43    g=torch.Generator(device=device).manual_seed(SEED)
44    z=torch.rand(96,2,generator=g,device=device)*2-1; u=torch.rand(96,2,generator=g,device=device)*2-1
45    target=1.08*z+.25*u+.04*torch.sin(2*z)
46    opt=torch.optim.Adam(model.parameters(),lr=3e-3)
47    for _ in range(epochs):
48        opt.zero_grad(); pred=model(z,u); task=((pred-target)**2).mean()
49        # Exact Jacobian penalty, retained for the idea even though its
50        # spectral-normalized budget already gives a structural certificate.
51        qs=[]
52        for k in range(0,96,8):
53            M=jac(model,z[k],u[k]); qs.append(torch.nn.functional.softplus(eig(M.T@M-(1-model.alpha)*torch.eye(2,device=device))+.02)**2)
54        penalty=torch.stack(qs).mean(); (task+(2*penalty if constrained else 0)).backward(); opt.step()
55    return float(task.detach().cpu()),float(penalty.detach().cpu())
56
57def main():
58    global device
59    # Algebra sanity check: the discrete LMI is exactly equivalent here to
60    # the largest eigenvalue of M^T M-(1-alpha)I being nonpositive.
61    good=.8*torch.eye(2); bad=1.1*torch.eye(2); I=torch.eye(2)
62    out={'seed':SEED,'device':device,'math_check':{
63      'alpha':.3,'good_largest_eigenvalue':float(eig(good.T@good-(1-.3)*I)),
64      'bad_largest_eigenvalue':float(eig(bad.T@bad-(1-.3)*I)),
65      'good_distance_after_12':.8**12,'bad_distance_after_12':1.1**12}}
66    for name,flag in [('baseline',False),('idea',True)]:
67        try:
68            m=Block(idea=flag).to(device); loss,pen=train(m,flag); out[name]={'task_mse':loss,'penalty':pen,**report(m)}
69        except Exception:
70            if device!='cuda': raise
71            device='cpu'; m=Block(idea=flag); loss,pen=train(m,flag); out[name]={'task_mse':loss,'penalty':pen,**report(m)}
72    with open('safe_results.json','w') as f: json.dump(out,f,indent=2)
73    print(json.dumps(out,indent=2))
74if __name__=='__main__': main()