Gaussian Barycentric Constraint Layer / experiment.py
Mechanism failed
1import json, math, time
2import numpy as np
3import torch
4from torch import nn
5
6SEED=7
7np.random.seed(SEED); torch.manual_seed(SEED)
8try:
9 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
10 if device.type=='cuda':
11 torch.zeros(1,device=device)
12except Exception:
13 device=torch.device('cpu')
14
15# Exact Euclidean projection to the probability simplex.
16def simplex_proj(v):
17 # v: (..., d)
18 u,_=torch.sort(v, dim=-1, descending=True)
19 cssv=torch.cumsum(u,dim=-1)-1
20 ind=torch.arange(1,v.shape[-1]+1,device=v.device,dtype=v.dtype)
21 cond=u-cssv/ind>0
22 rho=cond.sum(dim=-1,keepdim=True).clamp_min(1)
23 theta=cssv.gather(-1,rho.long()-1)/rho
24 return torch.clamp(v-theta,min=0)
25
26def clip_renorm(v, eps=1e-8):
27 z=torch.clamp(v,min=0)
28 return z/(z.sum(-1,keepdim=True)+eps)
29
30def barycentric(x, K=32, delta=.01, mode='hard', tau=.03, noise=None):
31 # Simplex C: nonnegative coordinates and sum exactly one.
32 B,D=x.shape
33 if noise is None: noise=torch.randn(B,K,D,device=x.device)
34 y=x[:,None,:]+math.sqrt(delta)*noise
35 if mode=='hard':
36 q=(y.min(-1).values>=0) & ((y.sum(-1)-1).abs()<.035)
37 w=q.float()
38 else:
39 # smooth convex violation surrogate; the resulting map is differentiable
40 v=torch.relu(-y).pow(2).sum(-1)+(y.sum(-1)-1).pow(2)
41 w=torch.exp(-v/tau)
42 den=w.sum(1,keepdim=True)
43 return (w[...,None]*y).sum(1)/(den+1e-8), w.mean(1), den.squeeze(1)
44
45class Net(nn.Module):
46 def __init__(self,d,h=48):
47 super().__init__(); self.net=nn.Sequential(nn.Linear(d,h),nn.Tanh(),nn.Linear(h,h),nn.Tanh(),nn.Linear(h,3))
48 def forward(self,x): return self.net(x)
49
50def run_model(mode, K=32, steps=500, delta=.01, tau=.03):
51 # Fixed synthetic mapping to simplex, with noisy observations at test time.
52 n,d=768,8
53 X=torch.randn(n,d,device=device)
54 W=torch.randn(d,3,device=device)
55 target= torch.softmax(X@W + .25*torch.randn(n,3,device=device),dim=-1)
56 net=Net(d).to(device); opt=torch.optim.Adam(net.parameters(),lr=3e-3)
57 t0=time.perf_counter()
58 for step in range(steps):
59 ix=torch.randint(0,n,(64,),device=device); raw=net(X[ix])
60 if mode=='clip': out=clip_renorm(raw)
61 elif mode=='exact': out=simplex_proj(raw)
62 else: out,_,_=barycentric(raw,K,delta,mode,tau)
63 loss=((out-target[ix])**2).mean()
64 opt.zero_grad(); loss.backward(); opt.step()
65 elapsed=time.perf_counter()-t0
66 with torch.no_grad():
67 raw=net(X); out=(clip_renorm(raw) if mode=='clip' else simplex_proj(raw) if mode=='exact' else barycentric(raw,K,delta,mode,tau)[0])
68 testX=torch.randn(256,d,device=device); testT=torch.softmax(testX@W,dim=-1)
69 raw2=net(testX); out2=(clip_renorm(raw2) if mode=='clip' else simplex_proj(raw2) if mode=='exact' else barycentric(raw2,K,delta,mode,tau)[0])
70 mse=((out2-testT)**2).mean().item()
71 violation=torch.relu(-out2).max().item()+torch.abs(out2.sum(-1)-1).max().item()
72 # noise robustness, using same perturbation for all methods is handled by caller approximately
73 rate=(out2.min(-1).values>=-1e-6).float().mean().item()
74 gradnorm=float(torch.sqrt(sum((p.grad.detach()**2).sum() for p in net.parameters() if p.grad is not None)).item())
75 return {'mse':mse,'max_violation':violation,'feasible_rate':rate,'train_sec':elapsed,'gradnorm':gradnorm}
76
77def math_check():
78 # Interval has a closed-form conditional Gaussian barycenter estimated by quadrature.
79 # Monte Carlo also checks Dm=Cov(Y)/delta and common-random-number firmness.
80 torch.manual_seed(SEED); N=1200000; delta=.16
81 eps=torch.randn(N,device=device); xs=torch.tensor([-0.7,-0.15,0.35,0.8],device=device)
82 vals=[]; covs=[]
83 for x in xs:
84 y=x+math.sqrt(delta)*eps; keep=(y>=0)&(y<=1); z=y[keep]
85 vals.append(z.mean()); covs.append(z.var(unbiased=False))
86 vals=torch.stack(vals); covs=torch.stack(covs)
87 # finite difference around x=.35 with shared draws
88 x=torch.tensor(.35,device=device); h=.01
89 def mcmean(a):
90 y=a+math.sqrt(delta)*eps; return y[(y>=0)&(y<=1)].mean()
91 deriv=(mcmean(x+h)-mcmean(x-h))/(2*h)
92 identity=covs[2]/delta
93 # firmness pairwise; monotone 1-Lipschitz equivalent in 1D
94 lhs=[]; rhs=[]
95 for i in range(len(xs)-1):
96 dm=vals[i+1]-vals[i]; dx=xs[i+1]-xs[i]; lhs.append(dm*dm); rhs.append(dm*dx)
97 return {'interval_means':[round(float(v),6) for v in vals.cpu()], 'jacobian_fd':float(deriv), 'cov_over_delta':float(identity), 'firm_max_lhs_minus_rhs':float((torch.stack(lhs)-torch.stack(rhs)).max().item()), 'acceptance_at_.35':float(((.35+math.sqrt(delta)*eps>=0)&(.35+math.sqrt(delta)*eps<=1)).float().mean().item())}
98
99def main():
100 print(json.dumps({'device':str(device),'math_check':math_check()},indent=2))
101 results={}
102 for mode in ['clip','exact','soft','hard']:
103 # Hard estimator is evaluated at a nearby point; its sparse rejection rate is reported separately.
104 results[mode]=run_model(mode,K=32,steps=500,delta=.01,tau=.03)
105 torch.manual_seed(SEED); x=torch.tensor([[.34,.33,.33]],device=device)
106 _,rate,den=barycentric(x,K=32,delta=.01,mode='hard')
107 results['hard']['feasible_samples_mean']=float(den.item()); results['hard']['sample_feasible_rate']=float(rate.item())
108 print(json.dumps({'results':results},indent=2))
109
110if __name__=='__main__': main()