Composite Density-Power Loss / dpd_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 = 2548
  8
  9def seed_all(seed=SEED):
 10    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 11    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 12
 13def dpd_loss(logits, y, alpha, weights=None):
 14    # Exact categorical A_alpha=sum_c p_c^(1+alpha), averaged over examples/components.
 15    p = torch.softmax(logits.float(), dim=-1)
 16    a = float(alpha)
 17    A = (p.pow(1.0+a)).sum(dim=-1)
 18    observed = p.gather(-1, y.long().unsqueeze(-1)).squeeze(-1).clamp_min(1e-12)
 19    loss = A - (1.0 + 1.0/a) * observed.pow(a)
 20    if weights is not None:
 21        loss = loss * torch.as_tensor(weights, device=loss.device, dtype=loss.dtype)
 22    return loss.mean()
 23
 24def observed_term_loss(logits, y, alpha):
 25    p = torch.softmax(logits.float(), dim=-1)
 26    q = p.gather(-1, y.long().unsqueeze(-1)).squeeze(-1).clamp_min(1e-12)
 27    return -(1.0 + 1.0/alpha) * q.pow(alpha)
 28
 29def math_sweep():
 30    # Prediction 1: observed-score DPD/CE gradient norm ratio=(1+alpha)q^alpha.
 31    qs = np.array([1e-4, 3e-4, 1e-3, 3e-3, 1e-2, 3e-2, 1e-1])
 32    alphas = [0.1, 0.3, 0.5, 1.0]
 33    rows=[]
 34    for a in alphas:
 35        ratios=[]
 36        for q in qs:
 37            # logits chosen so class 0 has probability q, remaining mass uniform
 38            logits = torch.tensor([[math.log(q), math.log((1-q)/2), math.log((1-q)/2)]], requires_grad=True)
 39            y=torch.tensor([0])
 40            g= torch.autograd.grad(observed_term_loss(logits,y,a).sum(), logits)[0].norm().item()
 41            ce= torch.autograd.grad((-torch.log(torch.softmax(logits, -1)[:,0])).sum(), logits, retain_graph=True)[0].norm().item()
 42            ratios.append(g/ce)
 43        slope=np.polyfit(np.log(qs), np.log(ratios), 1)[0]
 44        pred=np.array([(1+a)*q**a for q in qs])
 45        rel=float(np.max(np.abs(np.array(ratios)-pred)/(pred+1e-12)))
 46        rows.append({'alpha':a,'max_relative_error':rel,'loglog_slope':float(slope),'predicted_slope':a,
 47                     'ratio_at_q_1e-3':float(ratios[2]),'predicted_ratio_at_q_1e-3':float((1+a)*1e-3**a)})
 48    # Prediction 2: alpha=0 limit of full DPD gradient approaches CE gradient.
 49    logits=torch.tensor([[1.2,-.7,.1]],dtype=torch.float64,requires_grad=True); y=torch.tensor([2])
 50    ce=torch.autograd.grad((-torch.log_softmax(logits,-1)[:,2]).sum(),logits,retain_graph=True)[0]
 51    limit=[]
 52    for a in [0.5,0.2,0.1,0.05,0.02,0.01]:
 53        l=dpd_loss(logits.float(),y,a)
 54        g=torch.autograd.grad(l,logits,retain_graph=True)[0].double()
 55        limit.append({'alpha':a,'gradient_relative_error':float((g-ce).norm()/ce.norm())})
 56    return {'observed_score_scaling':rows,'alpha_to_zero_full_gradient':limit}
 57
 58def make_data(n=1800, noise=0.20, seed=SEED):
 59    rng=np.random.default_rng(seed)
 60    centers=np.array([[-1.4,-1.0],[1.4,-1.0],[0,1.5]],dtype=np.float32)
 61    y=rng.integers(0,3,n); x=centers[y]+rng.normal(0,.85,(n,2)).astype(np.float32)
 62    clean=y.copy(); noisy=y.copy(); mask=rng.random(n)<noise
 63    for i in np.where(mask)[0]: noisy[i]=rng.integers(0,2) if y[i]==2 else 2 # deterministic different class
 64    idx=rng.permutation(n); tr=idx[:1200]; va=idx[1200:]
 65    return torch.tensor(x[tr]),torch.tensor(noisy[tr]),torch.tensor(clean[tr]),torch.tensor(x[va]),torch.tensor(clean[va])
 66
 67class TinyNet(nn.Module):
 68    def __init__(self):
 69        super().__init__(); self.net=nn.Sequential(nn.Linear(2,32),nn.Tanh(),nn.Linear(32,3))
 70    def forward(self,x): return self.net(x)
 71
 72def train(kind, alpha=.5, epochs=80):
 73    x,y,yclean,xv,yv=make_data()
 74    model=TinyNet(); opt=torch.optim.Adam(model.parameters(),lr=.025)
 75    for _ in range(epochs):
 76        perm=torch.randperm(len(x))
 77        for j in range(0,len(x),64):
 78            ix=perm[j:j+64]; logits=model(x[ix])
 79            loss=nn.functional.cross_entropy(logits,y[ix]) if kind=='ce' else dpd_loss(logits,y[ix],alpha)
 80            opt.zero_grad(); loss.backward(); opt.step()
 81    with torch.no_grad():
 82        pred=model(xv).argmax(1); acc=(pred==yv).float().mean().item()
 83        train_pred=model(x).argmax(1); train_acc=(train_pred==yclean).float().mean().item()
 84    # Per-example gradient norm for clean-vs-corrupted labels at the trained model.
 85    corrupted=(y!=yclean); vals=[]
 86    for i in range(len(x)):
 87        model.zero_grad(set_to_none=True); z=model(x[i:i+1])
 88        l=nn.functional.cross_entropy(z,y[i:i+1]) if kind=='ce' else dpd_loss(z,y[i:i+1],alpha)
 89        l.backward(); vals.append(sum((p.grad.detach()**2).sum().item() for p in model.parameters() if p.grad is not None)**.5)
 90    vals=np.array(vals)
 91    return {'clean_validation_accuracy':acc,'clean_train_accuracy':train_acc,
 92            'corrupted_gradient_norm':float(vals[corrupted.numpy()].mean()),
 93            'clean_gradient_norm':float(vals[~corrupted.numpy()].mean()),
 94            'gradient_corrupted_to_clean':float(vals[corrupted.numpy()].mean()/vals[~corrupted.numpy()].mean())}
 95
 96def main():
 97    seed_all(); math_results=math_sweep()
 98    # same data/model initialization stream is reset for an apples-to-apples comparison.
 99    seed_all(); ce=train('ce'); seed_all(); dpd=train('dpd',.5)
100    out={'seed':SEED,'math_verification':math_results,'toy_training':{'baseline_CE':ce,'idea_DPD_alpha_0.5':dpd}}
101    Path('results.json').write_text(json.dumps(out,indent=2))
102    print(json.dumps(out,indent=2))
103
104if __name__=='__main__': main()