Adaptive Barrier-Margin Regularization / experiment.py
Mechanism failed
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED=7
8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
9DEVICE='cuda' if torch.cuda.is_available() else 'cpu'
10try:
11 if DEVICE=='cuda': torch.cuda.get_device_properties(0)
12except Exception:
13 DEVICE='cpu'
14
15# Scalar reconstructed barrier: psi = safe_action(x)-u. psi>=0 means nominally safe.
16# Positive r=kappa*epsilon-psi is the proposed margin violation.
17def softplus_sq(r,tau):
18 return np.logaddexp(0.0,r/tau)**2
19
20def verify_update():
21 delta=.1; rho=.25; e=0.32; eps=0.02; vals=[]
22 for _ in range(100):
23 eps=np.clip(eps+rho*(e-(1-delta)*eps), .001, 2.)
24 vals.append(eps)
25 fixed=e/(1-delta)
26 # exact linear convergence: error after n = (1-rho*(1-delta))^n error0
27 pred_factor=1-rho*(1-delta)
28 observed_factor=np.mean(np.abs(np.diff(vals[-30:])[1:])/np.abs(np.diff(vals[-30:])[:-1]))
29 # coverage transition for Gaussian error magnitudes: eps at 90th percentile
30 rng=np.random.default_rng(SEED)
31 cover=[]
32 for sigma in [.05,.10,.20,.40]:
33 err=np.abs(rng.normal(0,sigma,200000))
34 q=float(np.quantile(err,.9)); rm_eps=float(np.mean(err)/(1-delta)); rm_cov=float(np.mean(err<=rm_eps)); pred_rm_cov=math.erf((math.sqrt(2/math.pi)/(1-delta))/math.sqrt(2)); cover.append({'sigma':sigma,'q90':q,'coverage_at_q90':float(np.mean(err<=q)), 'rm_epsilon_pred':rm_eps, 'rm_coverage_pred':pred_rm_cov, 'rm_coverage_obs':rm_cov})
35 return {'fixed_point_pred':fixed,'fixed_point_obs':vals[-1],
36 'convergence_factor_pred':pred_factor,'convergence_factor_obs':float(observed_factor),
37 'coverage_sweep':cover}
38
39class Policy(nn.Module):
40 def __init__(self):
41 super().__init__(); self.net=nn.Sequential(nn.Linear(1,16),nn.Tanh(),nn.Linear(16,1),nn.Tanh())
42 def forward(self,x): return 1.5*self.net(x)
43
44def train(method, sigma, seed=SEED, steps=700):
45 torch.manual_seed(seed); rng=np.random.default_rng(seed)
46 p=Policy().to(DEVICE); opt=torch.optim.Adam(p.parameters(),lr=.012)
47 eps=.03; delta=.1; rho=.08; kappa=1.0; tau=.08; beta=2.0
48 x=torch.linspace(-1,1,96,device=DEVICE).reshape(-1,1)
49 target=torch.full_like(x,.75)
50 for step in range(steps):
51 # deterministic known observer bias plus random true disturbance; detach observer quantities
52 d=torch.tensor(rng.normal(0,sigma,96),dtype=torch.float32,device=DEVICE).reshape(-1,1)
53 dh=torch.zeros_like(d)
54 err=float(torch.mean(torch.abs(d-dh)).item())
55 if method=='adaptive': eps=float(np.clip(eps+rho*(err-(1-delta)*eps),.001,1.5))
56 elif method=='fixed': eps=.22
57 else: eps=0.
58 u=p(x); safe=-.15*x + .05 # nominal action allowed below this
59 psi=safe-u
60 r=kappa*eps-psi
61 task=torch.mean((u-target)**2)
62 barrier=beta*torch.mean(torch.nn.functional.softplus(r/tau)**2)
63 loss=task + (barrier if method!='baseline' else 0.)
64 opt.zero_grad(); loss.backward(); opt.step()
65 # Evaluation over states/disturbances. Filter projects u+d_hat <= safe - k eps.
66 xx=torch.linspace(-1,1,256,device=DEVICE).reshape(-1,1)
67 with torch.no_grad(): raw=p(xx).cpu().numpy().ravel()
68 safe=(-.15*xx+.05).cpu().numpy().ravel()
69 # fresh disturbances, observer d_hat=0; projection is the inference-time barrier filter
70 erng=np.random.default_rng(seed+1000)
71 d=erng.normal(0,sigma,(40,256)); dh=0.
72 threshold=safe-kappa*eps
73 filt=np.minimum(raw,threshold)
74 intervention=np.abs(raw-filt)
75 actual=filt[None,:]+d
76 violations=float(np.mean(actual>safe[None,:]))
77 nominal_viol=float(np.mean(raw>safe))
78 track=float(np.mean((filt-.75)**2))
79 return {'eps':eps,'violations_after_filter':violations,'nominal_violations':nominal_viol,
80 'intervention_norm':float(np.mean(intervention)),'tracking_mse':track,
81 'barrier_fraction':float(np.mean(raw>threshold))}
82
83def main():
84 checks=verify_update(); results={'device':DEVICE,'verification':checks,'experiments':{}}
85 for sigma in [.05,.10,.20,.40]:
86 results['experiments'][str(sigma)]={m:train(m,sigma) for m in ['baseline','fixed','adaptive']}
87 Path('results.json').write_text(json.dumps(results,indent=2))
88 print(json.dumps(results,indent=2))
89if __name__=='__main__': main()