Conditional Sinkhorn Adversarial Augmentation / experiment.py
Failed on benchmark
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()