import json, math, time import numpy as np import torch from torch import nn SEED=7 np.random.seed(SEED); torch.manual_seed(SEED) try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device.type=='cuda': torch.zeros(1,device=device) except Exception: device=torch.device('cpu') # Exact Euclidean projection to the probability simplex. def simplex_proj(v): # v: (..., d) u,_=torch.sort(v, dim=-1, descending=True) cssv=torch.cumsum(u,dim=-1)-1 ind=torch.arange(1,v.shape[-1]+1,device=v.device,dtype=v.dtype) cond=u-cssv/ind>0 rho=cond.sum(dim=-1,keepdim=True).clamp_min(1) theta=cssv.gather(-1,rho.long()-1)/rho return torch.clamp(v-theta,min=0) def clip_renorm(v, eps=1e-8): z=torch.clamp(v,min=0) return z/(z.sum(-1,keepdim=True)+eps) def barycentric(x, K=32, delta=.01, mode='hard', tau=.03, noise=None): # Simplex C: nonnegative coordinates and sum exactly one. B,D=x.shape if noise is None: noise=torch.randn(B,K,D,device=x.device) y=x[:,None,:]+math.sqrt(delta)*noise if mode=='hard': q=(y.min(-1).values>=0) & ((y.sum(-1)-1).abs()<.035) w=q.float() else: # smooth convex violation surrogate; the resulting map is differentiable v=torch.relu(-y).pow(2).sum(-1)+(y.sum(-1)-1).pow(2) w=torch.exp(-v/tau) den=w.sum(1,keepdim=True) return (w[...,None]*y).sum(1)/(den+1e-8), w.mean(1), den.squeeze(1) class Net(nn.Module): def __init__(self,d,h=48): super().__init__(); self.net=nn.Sequential(nn.Linear(d,h),nn.Tanh(),nn.Linear(h,h),nn.Tanh(),nn.Linear(h,3)) def forward(self,x): return self.net(x) def run_model(mode, K=32, steps=500, delta=.01, tau=.03): # Fixed synthetic mapping to simplex, with noisy observations at test time. n,d=768,8 X=torch.randn(n,d,device=device) W=torch.randn(d,3,device=device) target= torch.softmax(X@W + .25*torch.randn(n,3,device=device),dim=-1) net=Net(d).to(device); opt=torch.optim.Adam(net.parameters(),lr=3e-3) t0=time.perf_counter() for step in range(steps): ix=torch.randint(0,n,(64,),device=device); raw=net(X[ix]) if mode=='clip': out=clip_renorm(raw) elif mode=='exact': out=simplex_proj(raw) else: out,_,_=barycentric(raw,K,delta,mode,tau) loss=((out-target[ix])**2).mean() opt.zero_grad(); loss.backward(); opt.step() elapsed=time.perf_counter()-t0 with torch.no_grad(): 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]) testX=torch.randn(256,d,device=device); testT=torch.softmax(testX@W,dim=-1) 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]) mse=((out2-testT)**2).mean().item() violation=torch.relu(-out2).max().item()+torch.abs(out2.sum(-1)-1).max().item() # noise robustness, using same perturbation for all methods is handled by caller approximately rate=(out2.min(-1).values>=-1e-6).float().mean().item() gradnorm=float(torch.sqrt(sum((p.grad.detach()**2).sum() for p in net.parameters() if p.grad is not None)).item()) return {'mse':mse,'max_violation':violation,'feasible_rate':rate,'train_sec':elapsed,'gradnorm':gradnorm} def math_check(): # Interval has a closed-form conditional Gaussian barycenter estimated by quadrature. # Monte Carlo also checks Dm=Cov(Y)/delta and common-random-number firmness. torch.manual_seed(SEED); N=1200000; delta=.16 eps=torch.randn(N,device=device); xs=torch.tensor([-0.7,-0.15,0.35,0.8],device=device) vals=[]; covs=[] for x in xs: y=x+math.sqrt(delta)*eps; keep=(y>=0)&(y<=1); z=y[keep] vals.append(z.mean()); covs.append(z.var(unbiased=False)) vals=torch.stack(vals); covs=torch.stack(covs) # finite difference around x=.35 with shared draws x=torch.tensor(.35,device=device); h=.01 def mcmean(a): y=a+math.sqrt(delta)*eps; return y[(y>=0)&(y<=1)].mean() deriv=(mcmean(x+h)-mcmean(x-h))/(2*h) identity=covs[2]/delta # firmness pairwise; monotone 1-Lipschitz equivalent in 1D lhs=[]; rhs=[] for i in range(len(xs)-1): dm=vals[i+1]-vals[i]; dx=xs[i+1]-xs[i]; lhs.append(dm*dm); rhs.append(dm*dx) 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())} def main(): print(json.dumps({'device':str(device),'math_check':math_check()},indent=2)) results={} for mode in ['clip','exact','soft','hard']: # Hard estimator is evaluated at a nearby point; its sparse rejection rate is reported separately. results[mode]=run_model(mode,K=32,steps=500,delta=.01,tau=.03) torch.manual_seed(SEED); x=torch.tensor([[.34,.33,.33]],device=device) _,rate,den=barycentric(x,K=32,delta=.01,mode='hard') results['hard']['feasible_samples_mean']=float(den.item()); results['hard']['sample_feasible_rate']=float(rate.item()) print(json.dumps({'results':results},indent=2)) if __name__=='__main__': main()