r-Deformed Power Divergence Loss / run_experiment.py

Mechanism confirmed, baseline not beaten

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