Adaptive Barrier-Margin Regularization / experiment.py

Mechanism failed

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