Conditional Sinkhorn Adversarial Augmentation / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2import numpy as np
  3import torch
  4
  5SEED=1381
  6random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  7torch.set_num_threads(4)
  8DT=torch.float64
  9
 10def sinkhorn_ot(a,b,eps=0.15,iters=35):
 11    n,m=a.shape[0],b.shape[0]
 12    C=((a[:,None,:]-b[None,:,:])**2).sum(-1)
 13    logK=-C/eps
 14    la=torch.full((n,),-math.log(n),dtype=a.dtype)
 15    lb=torch.full((m,),-math.log(m),dtype=a.dtype)
 16    u=torch.zeros_like(la); v=torch.zeros_like(lb)
 17    for _ in range(iters):
 18        u=la-torch.logsumexp(logK+v[None,:],dim=1)
 19        v=lb-torch.logsumexp(logK+u[:,None],dim=0)
 20    P=torch.exp(logK+u[:,None]+v[None,:])
 21    return (P*C).sum()
 22
 23def sinkhorn_div(a,b,eps=0.15):
 24    return sinkhorn_ot(a,b,eps)-0.5*sinkhorn_ot(a,a,eps)-0.5*sinkhorn_ot(b,b,eps)
 25
 26def math_checks():
 27    z=torch.linspace(-1.1,1.1,12,dtype=DT)[:,None]
 28    nominal=z
 29    deltas=np.linspace(0,1.2,7)
 30    vals=[]
 31    for d in deltas:
 32        vals.append(float(sinkhorn_div(nominal+float(d),nominal)))
 33    x=deltas[1:]**2; y=np.array(vals[1:])
 34    slope=float(np.dot(x,y)/np.dot(x,x))
 35    rel_err=float(max(abs(v-slope*d*d)/(abs(v)+1e-9) for d,v in zip(deltas[1:],vals[1:])))
 36    rho=0.36
 37    predicted=math.sqrt(rho/slope)
 38    grid=np.linspace(0,1.2,31)
 39    gv=np.array([float(sinkhorn_div(nominal+float(q),nominal)) for q in grid])
 40    crossing=float(grid[np.argmin(np.abs(gv-rho))])
 41    # Entropic regularization prediction: debiasing makes S(0) approximately zero.
 42    zero=float(sinkhorn_div(nominal,nominal))
 43    # Radius violation and projected multiplier: lambda rises iff S > rho.
 44    lam=0.; eta=.8; seq=[]
 45    for d in [0.2,0.5,0.8,1.0]:
 46        s=float(sinkhorn_div(nominal+d,nominal)); lam=max(0.,lam+eta*(s-rho)); seq.append((d,s,lam))
 47    return {'translation_scaling':{'prediction':'S approximately k*delta^2', 'fitted_k':slope, 'max_relative_error':rel_err, 'deltas':deltas.tolist(), 'S':vals},
 48            'radius_boundary':{'prediction':'boundary delta=sqrt(rho/k)', 'rho':rho, 'predicted_delta':predicted, 'observed_delta':crossing, 'absolute_error':abs(predicted-crossing)},
 49            'debiased_identity':{'prediction':'S(A,A)=0', 'observed':zero},
 50            'multiplier_projection':{'prediction':'lambda increases only for violations', 'sequence':seq}}
 51
 52def data(n, shift, seed):
 53    g=torch.Generator().manual_seed(seed)
 54    x=torch.rand(n,1,generator=g)*2-1
 55    y=torch.sin(3*x)+0.12*torch.randn(n,1,generator=g)+shift
 56    return x,y
 57
 58def train(method, steps=100, seed=1381):
 59    torch.manual_seed(seed)
 60    # Nominal conditional generator y=sin(3x)+noise; adversary is a scalar residual shift.
 61    w=torch.zeros(2,1,dtype=DT,requires_grad=True)
 62    psi=torch.tensor(0.,dtype=DT,requires_grad=(method=='sinkhorn'))
 63    lam=0.; rho=.16; eps=.15
 64    opt=torch.optim.SGD([w],lr=.055)
 65    for t in range(steps):
 66        x,base=data(16,0.,seed+t)
 67        z=torch.randn(16,1,dtype=DT)*.12
 68        yn=base
 69        if method=='ordinary':
 70            ya=yn
 71        elif method=='unconstrained':
 72            # Fixed-budget unconstrained augmentation: maximum useful shift is large.
 73            ya=yn+0.55
 74        else:
 75            ya=yn+psi+z*0.0
 76        pred=w[0]*x+w[1]
 77        task=((pred-ya)**2).mean()
 78        if method=='sinkhorn':
 79            # Optimize adversary by ascent on task minus radius penalty.
 80            s=sinkhorn_div(yn+psi,yn,eps)
 81            adv=task-lam*torch.relu(s-rho)
 82            grad=torch.autograd.grad(adv,psi,retain_graph=True)[0]
 83            with torch.no_grad(): psi.add_(.16*grad).clamp_(-1.0,1.0)
 84            s=sinkhorn_div((yn+psi).detach(),yn,eps)
 85            lam=max(0.,lam+.35*(float(s)-rho))
 86            ya=yn+psi.detach()
 87            task=((w[0]*x+w[1]-ya)**2).mean()
 88        opt.zero_grad(); task.backward(); opt.step()
 89    # Evaluate clean and shifted contexts.
 90    with torch.no_grad():
 91        xt,yt=data(128,0.,9001); xs,ys=data(128,.55,9002)
 92        clean=float(((w[0]*xt+w[1]-yt)**2).mean())
 93        shifted=float(((w[0]*xs+w[1]-ys)**2).mean())
 94        if method=='sinkhorn':
 95            final_s=float(sinkhorn_div(yt+psi.detach(),yt,eps)); p=float(psi)
 96        else: final_s=float('nan'); p=float('nan')
 97    return {'clean_mse':clean,'shifted_mse':shifted,'psi':p,'final_sinkhorn':final_s,'lambda':lam}
 98
 99def main():
100    out={'math_checks':math_checks(),'training':{m:train(m) for m in ['ordinary','unconstrained','sinkhorn']}}
101    print(json.dumps(out,indent=2))
102    with open('results.json','w') as f: json.dump(out,f,indent=2)
103if __name__=='__main__': main()