import json, random import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED = 117 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' try: if DEVICE == 'cuda': torch.zeros(1, device='cuda') except Exception: DEVICE = 'cpu' def power_mean(a, weights, p, eps=1e-12): a = a.clamp_min(eps); w = weights.to(a.device) if abs(p) < 1e-10: return torch.exp((w * torch.log(a)).sum(-1)) return (w * a.pow(p)).sum(-1).clamp_min(eps).pow(1.0/p) def q_exponent(p, d): return p/(1.0-d*p) def softmin(u, tau): return -tau * torch.logsumexp(-u/tau, dim=0) def bbl_penalty(x1, x2, heads, p=0.0, d=1, tau_factor=.02): z = .5*x1 + .5*x2 f1 = F.softplus(heads[0](x1)).squeeze(-1) + 1e-6 f2 = F.softplus(heads[1](x2)).squeeze(-1) + 1e-6 g1 = F.softplus(heads[0](z)).squeeze(-1) + 1e-6 g2 = F.softplus(heads[1](z)).squeeze(-1) + 1e-6 r = torch.stack((g1/f1, g2/f2), -1) # Independent Monte Carlo mass estimates for f and g. fx1 = f1.mean(); fx2 = f2.mean() gx1 = (F.softplus(heads[0](x1)).squeeze(-1)+1e-6).mean() gx2 = (F.softplus(heads[1](x2)).squeeze(-1)+1e-6).mean() ratios = torch.stack((gx1/fx1, gx2/fx2)) w = torch.tensor([.5, .5], device=x1.device) u = power_mean(r, w, p) rhs = power_mean(ratios[None, :], w, q_exponent(p,d)).squeeze() tau = tau_factor * u.detach().median().clamp_min(1e-4) sm = softmin(u, tau) gap = sm-rhs return F.relu(gap).pow(2), gap, u.detach(), rhs.detach() class Branch(nn.Module): def __init__(self): super().__init__(); self.net=nn.Sequential(nn.Linear(1,16),nn.Tanh(),nn.Linear(16,1)) def forward(self,x): return self.net(x) def make_data(n, noise=.55, seed=0): gen=torch.Generator().manual_seed(seed) y=(torch.randint(0,2,(n,),generator=gen)*2-1).float() x1=(y+noise*torch.randn(n,generator=gen)).unsqueeze(1) x2=(y+noise*torch.randn(n,generator=gen)).unsqueeze(1) return x1,x2,(y>0).float() def ece(prob,y,bins=10): out=0.; edges=torch.linspace(0,1,bins+1) for j in range(bins): mask=(prob>=edges[j])&((prob0) if use_bbl: loss=loss+.15*pen opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): probs=torch.sigmoid((heads[0](tx1)+heads[1](tx2)).squeeze(1)); pred=(probs>.5).float() clean_acc=(pred==ty).float().mean().item() disagreement=(torch.sigmoid(heads[0](tx1))-torch.sigmoid(heads[1](tx2))).abs().mean().item() corrupt=tx1+2*torch.randn_like(tx1) cp=torch.sigmoid((heads[0](corrupt)+heads[1](tx2)).squeeze(1)) pen,gap,u,rhs=bbl_penalty(tx1[:96],tx2[:96],heads) return {'acc':clean_acc,'corrupt_acc':((cp>.5).float()==ty).float().mean().item(), 'ece':ece(probs,ty),'disagreement':disagreement,'eval_penalty':pen.item(), 'eval_gap':gap.item(),'ratio_q25':torch.quantile(u,.25).item(),'rhs':rhs.item(), 'train_active_steps':active,'train_max_penalty':max_pen,'device':DEVICE} def math_check(): grid=torch.linspace(-4,4,401); x1=grid[:,None]; x2=grid[None,:]; z=(x1+x2)/2 r1=torch.exp((-z*z+x1*x1)/2); r2=torch.exp((-z*z+x2*x2)/2); u=torch.sqrt(r1*r2) # For equal Gaussian fields and p=0, inf geometric ratio = mass-ratio RHS = 1. return {'gaussian_min':u.min().item(),'gaussian_rhs':1.0,'gaussian_max_violation':max(0.,u.min().item()-1.)} def gradient_check(): torch.manual_seed(SEED); h=[Branch().to(DEVICE),Branch().to(DEVICE)] x1,x2,_=make_data(96,seed=31); x1,x2=x1.to(DEVICE),x2.to(DEVICE) p,g,_,_=bbl_penalty(x1,x2,h); p.backward() return {'penalty':p.item(),'gap':g.item(),'gradient_norm':float(sum((v.grad.norm()**2 for q in h for v in q.parameters() if v.grad is not None),torch.tensor(0.,device=DEVICE)).sqrt().cpu())} def main(): print(json.dumps({'math_check':math_check(),'gradient_check':gradient_check(),'baseline':run(False),'idea':run(True)},indent=2)) if __name__=='__main__': main()