r-Deformed Power Divergence Loss / run_experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED=346
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10
11# Stable commuting r-deformed alpha divergence. Inputs are logits and target probabilities.
12def rlog(x, r):
13 if abs(r-1.0) < 1e-6:
14 return torch.log(x)
15 # expm1 is accurate near r=1; x is positive
16 return torch.expm1((1-r)*torch.log(x))/(1-r)
17
18def power_divergence(logits, target, alpha=0.5, r=1.0, eps=1e-8):
19 logp=torch.log_softmax(logits, dim=-1)
20 y=target.clamp_min(eps)
21 logT=torch.logsumexp(alpha*torch.log(y)+(1-alpha)*logp, dim=-1)
22 # Compute r-log from log T without forming unstable power sums where possible.
23 if abs(r-1.0) < 1e-6:
24 lr=logT
25 else:
26 lr=torch.expm1((1-r)*logT)/(1-r)
27 return (lr/(alpha-1)).mean()
28
29def ce(logits, target):
30 return -(target*torch.log_softmax(logits, -1)).sum(-1).mean()
31
32def synthetic(n=2400, k=3):
33 # Overlapping 2-D Gaussian classes, deliberately imbalanced.
34 counts=[n*2//3, n//6, n-(n*2//3+n//6)]
35 means=np.array([[-1.2,-.7],[1.2,-.5],[0,.95]],dtype=np.float32)
36 xs=[]; ys=[]
37 for c,num in enumerate(counts):
38 xs.append(np.random.randn(num,2).astype(np.float32)*0.85+means[c])
39 ys += [c]*num
40 x=np.concatenate(xs); y=np.array(ys)
41 ix=np.random.permutation(len(y)); return torch.tensor(x[ix]), torch.tensor(y[ix])
42
43def make_target(labels,k,smooth):
44 if smooth==0: return torch.nn.functional.one_hot(labels,k).float()
45 out=torch.full((len(labels),k),smooth/(k-1))
46 out.scatter_(1,labels[:,None],1-smooth); return out
47
48def train(kind, xtr,ytr,xva,yva, smooth=0.1, epochs=35):
49 kind_seed={'ce':1,'renyi':2,'r0':3,'r05':4}[kind]
50 torch.manual_seed(SEED+kind_seed)
51 model=nn.Sequential(nn.Linear(2,24),nn.Tanh(),nn.Linear(24,3))
52 opt=torch.optim.Adam(model.parameters(),lr=0.025)
53 bs=96; losses=[]; gradnorm=[]
54 for ep in range(epochs):
55 perm=torch.randperm(len(ytr))
56 for start in range(0,len(ytr),bs):
57 ix=perm[start:start+bs]; logits=model(xtr[ix]); target=make_target(ytr[ix],3,smooth)
58 if kind=='ce': loss=ce(logits,target)
59 elif kind=='renyi': loss=power_divergence(logits,target,.5,1.0)
60 elif kind=='r0': loss=power_divergence(logits,target,.5,0.0)
61 elif kind=='r05': loss=power_divergence(logits,target,.5,.5)
62 else: raise ValueError(kind)
63 opt.zero_grad(); loss.backward()
64 gn=float(torch.nn.utils.clip_grad_norm_(model.parameters(),100.0))
65 gradnorm.append(gn); opt.step(); losses.append(float(loss.detach()))
66 with torch.no_grad():
67 logits=model(xva); pred=logits.argmax(1); acc=float((pred==yva).float().mean())
68 probs=logits.softmax(-1); conf=probs.max(1).values
69 # 10-bin ECE, an auxiliary calibration observation.
70 ece=0.
71 for lo in torch.linspace(0,1,11)[:-1]:
72 hi=lo+0.1; mask=(conf>=lo)&(conf<hi)
73 if mask.any(): ece += float(mask.float().mean())*abs(float(conf[mask].mean())-float((pred[mask]==yva[mask]).float().mean()))
74 val_loss=float(ce(logits,make_target(yva,3,smooth)))
75 return {'acc':acc,'ece':ece,'val_ce':val_loss,'last_loss':float(np.mean(losses[-len(ytr)//bs:])),
76 'grad_mean':float(np.mean(gradnorm)),'grad_var':float(np.var(gradnorm)),'grad_max':float(np.max(gradnorm))}
77
78def math_check():
79 # Check claimed continuous r=1 limit against ordinary Renyi and finite autodiff.
80 torch.manual_seed(SEED)
81 logits=torch.randn(7,5,requires_grad=True); y=torch.softmax(torch.randn(7,5),-1)
82 vals=[]
83 for r in [0,.5,.99,1,1.01,1.5]:
84 z=power_divergence(logits,y,.5,r); vals.append(float(z.detach()))
85 assert torch.isfinite(z) and torch.isfinite(torch.autograd.grad(z,logits,retain_graph=True)[0]).all()
86 near=max(abs(vals[2]-vals[3]),abs(vals[4]-vals[3]))
87 renyi_exact=float(power_divergence(logits,y,.5,1.0).detach())
88 renyi_direct=float((torch.logsumexp(.5*torch.log(y)+.5*torch.log_softmax(logits,-1),-1)/(-.5)).mean().detach())
89 assert abs(renyi_exact-renyi_direct) < 1e-6
90 # Direct scalar loss profile demonstrates r changes the power-law penalty.
91 p=torch.tensor([.01,.1,.5,.9]); T=p**.5
92 profile={str(r):float((rlog(T,r)/(-.5))[0]) for r in [0,.5,1,1.5]}
93 return {'r_values':vals,'max_near_one_difference':near,'renyi_direct_difference':abs(renyi_exact-renyi_direct),'onehot_loss_at_p':profile}
94
95def main():
96 check=math_check()
97 x,y=synthetic(); split=1800; xtr,ytr=x[:split],y[:split]; xva,yva=x[split:],y[split:]
98 results={}
99 for kind in ['ce','renyi','r0','r05']:
100 results[kind]=train(kind,xtr,ytr,xva,yva,smooth=.1)
101 out={'seed':SEED,'alpha':.5,'r_candidates':[0,.5,1,1.5],
102 'parameterization':'commuting diagonal specialization; label smoothing=0.1; Adam; 35 epochs; batch=96',
103 'math_check':check,'results':results}
104 Path('results.json').write_text(json.dumps(out,indent=2))
105 print(json.dumps(out,indent=2))
106if __name__=='__main__': main()