import json, math, random from pathlib import Path import numpy as np import torch from torch import nn SEED=346 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) # Stable commuting r-deformed alpha divergence. Inputs are logits and target probabilities. def rlog(x, r): if abs(r-1.0) < 1e-6: return torch.log(x) # expm1 is accurate near r=1; x is positive return torch.expm1((1-r)*torch.log(x))/(1-r) def power_divergence(logits, target, alpha=0.5, r=1.0, eps=1e-8): logp=torch.log_softmax(logits, dim=-1) y=target.clamp_min(eps) logT=torch.logsumexp(alpha*torch.log(y)+(1-alpha)*logp, dim=-1) # Compute r-log from log T without forming unstable power sums where possible. if abs(r-1.0) < 1e-6: lr=logT else: lr=torch.expm1((1-r)*logT)/(1-r) return (lr/(alpha-1)).mean() def ce(logits, target): return -(target*torch.log_softmax(logits, -1)).sum(-1).mean() def synthetic(n=2400, k=3): # Overlapping 2-D Gaussian classes, deliberately imbalanced. counts=[n*2//3, n//6, n-(n*2//3+n//6)] means=np.array([[-1.2,-.7],[1.2,-.5],[0,.95]],dtype=np.float32) xs=[]; ys=[] for c,num in enumerate(counts): xs.append(np.random.randn(num,2).astype(np.float32)*0.85+means[c]) ys += [c]*num x=np.concatenate(xs); y=np.array(ys) ix=np.random.permutation(len(y)); return torch.tensor(x[ix]), torch.tensor(y[ix]) def make_target(labels,k,smooth): if smooth==0: return torch.nn.functional.one_hot(labels,k).float() out=torch.full((len(labels),k),smooth/(k-1)) out.scatter_(1,labels[:,None],1-smooth); return out def train(kind, xtr,ytr,xva,yva, smooth=0.1, epochs=35): kind_seed={'ce':1,'renyi':2,'r0':3,'r05':4}[kind] torch.manual_seed(SEED+kind_seed) model=nn.Sequential(nn.Linear(2,24),nn.Tanh(),nn.Linear(24,3)) opt=torch.optim.Adam(model.parameters(),lr=0.025) bs=96; losses=[]; gradnorm=[] for ep in range(epochs): perm=torch.randperm(len(ytr)) for start in range(0,len(ytr),bs): ix=perm[start:start+bs]; logits=model(xtr[ix]); target=make_target(ytr[ix],3,smooth) if kind=='ce': loss=ce(logits,target) elif kind=='renyi': loss=power_divergence(logits,target,.5,1.0) elif kind=='r0': loss=power_divergence(logits,target,.5,0.0) elif kind=='r05': loss=power_divergence(logits,target,.5,.5) else: raise ValueError(kind) opt.zero_grad(); loss.backward() gn=float(torch.nn.utils.clip_grad_norm_(model.parameters(),100.0)) gradnorm.append(gn); opt.step(); losses.append(float(loss.detach())) with torch.no_grad(): logits=model(xva); pred=logits.argmax(1); acc=float((pred==yva).float().mean()) probs=logits.softmax(-1); conf=probs.max(1).values # 10-bin ECE, an auxiliary calibration observation. ece=0. for lo in torch.linspace(0,1,11)[:-1]: hi=lo+0.1; mask=(conf>=lo)&(conf