Gaussian Barycentric Constraint Layer / experiment.py

Mechanism failed

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