import json, math, random import numpy as np import torch SEED=1381 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) DT=torch.float64 def sinkhorn_ot(a,b,eps=0.15,iters=35): n,m=a.shape[0],b.shape[0] C=((a[:,None,:]-b[None,:,:])**2).sum(-1) logK=-C/eps la=torch.full((n,),-math.log(n),dtype=a.dtype) lb=torch.full((m,),-math.log(m),dtype=a.dtype) u=torch.zeros_like(la); v=torch.zeros_like(lb) for _ in range(iters): u=la-torch.logsumexp(logK+v[None,:],dim=1) v=lb-torch.logsumexp(logK+u[:,None],dim=0) P=torch.exp(logK+u[:,None]+v[None,:]) return (P*C).sum() def sinkhorn_div(a,b,eps=0.15): return sinkhorn_ot(a,b,eps)-0.5*sinkhorn_ot(a,a,eps)-0.5*sinkhorn_ot(b,b,eps) def math_checks(): z=torch.linspace(-1.1,1.1,12,dtype=DT)[:,None] nominal=z deltas=np.linspace(0,1.2,7) vals=[] for d in deltas: vals.append(float(sinkhorn_div(nominal+float(d),nominal))) x=deltas[1:]**2; y=np.array(vals[1:]) slope=float(np.dot(x,y)/np.dot(x,x)) rel_err=float(max(abs(v-slope*d*d)/(abs(v)+1e-9) for d,v in zip(deltas[1:],vals[1:]))) rho=0.36 predicted=math.sqrt(rho/slope) grid=np.linspace(0,1.2,31) gv=np.array([float(sinkhorn_div(nominal+float(q),nominal)) for q in grid]) crossing=float(grid[np.argmin(np.abs(gv-rho))]) # Entropic regularization prediction: debiasing makes S(0) approximately zero. zero=float(sinkhorn_div(nominal,nominal)) # Radius violation and projected multiplier: lambda rises iff S > rho. lam=0.; eta=.8; seq=[] for d in [0.2,0.5,0.8,1.0]: s=float(sinkhorn_div(nominal+d,nominal)); lam=max(0.,lam+eta*(s-rho)); seq.append((d,s,lam)) return {'translation_scaling':{'prediction':'S approximately k*delta^2', 'fitted_k':slope, 'max_relative_error':rel_err, 'deltas':deltas.tolist(), 'S':vals}, 'radius_boundary':{'prediction':'boundary delta=sqrt(rho/k)', 'rho':rho, 'predicted_delta':predicted, 'observed_delta':crossing, 'absolute_error':abs(predicted-crossing)}, 'debiased_identity':{'prediction':'S(A,A)=0', 'observed':zero}, 'multiplier_projection':{'prediction':'lambda increases only for violations', 'sequence':seq}} def data(n, shift, seed): g=torch.Generator().manual_seed(seed) x=torch.rand(n,1,generator=g)*2-1 y=torch.sin(3*x)+0.12*torch.randn(n,1,generator=g)+shift return x,y def train(method, steps=100, seed=1381): torch.manual_seed(seed) # Nominal conditional generator y=sin(3x)+noise; adversary is a scalar residual shift. w=torch.zeros(2,1,dtype=DT,requires_grad=True) psi=torch.tensor(0.,dtype=DT,requires_grad=(method=='sinkhorn')) lam=0.; rho=.16; eps=.15 opt=torch.optim.SGD([w],lr=.055) for t in range(steps): x,base=data(16,0.,seed+t) z=torch.randn(16,1,dtype=DT)*.12 yn=base if method=='ordinary': ya=yn elif method=='unconstrained': # Fixed-budget unconstrained augmentation: maximum useful shift is large. ya=yn+0.55 else: ya=yn+psi+z*0.0 pred=w[0]*x+w[1] task=((pred-ya)**2).mean() if method=='sinkhorn': # Optimize adversary by ascent on task minus radius penalty. s=sinkhorn_div(yn+psi,yn,eps) adv=task-lam*torch.relu(s-rho) grad=torch.autograd.grad(adv,psi,retain_graph=True)[0] with torch.no_grad(): psi.add_(.16*grad).clamp_(-1.0,1.0) s=sinkhorn_div((yn+psi).detach(),yn,eps) lam=max(0.,lam+.35*(float(s)-rho)) ya=yn+psi.detach() task=((w[0]*x+w[1]-ya)**2).mean() opt.zero_grad(); task.backward(); opt.step() # Evaluate clean and shifted contexts. with torch.no_grad(): xt,yt=data(128,0.,9001); xs,ys=data(128,.55,9002) clean=float(((w[0]*xt+w[1]-yt)**2).mean()) shifted=float(((w[0]*xs+w[1]-ys)**2).mean()) if method=='sinkhorn': final_s=float(sinkhorn_div(yt+psi.detach(),yt,eps)); p=float(psi) else: final_s=float('nan'); p=float('nan') return {'clean_mse':clean,'shifted_mse':shifted,'psi':p,'final_sinkhorn':final_s,'lambda':lam} def main(): out={'math_checks':math_checks(),'training':{m:train(m) for m in ['ordinary','unconstrained','sinkhorn']}} print(json.dumps(out,indent=2)) with open('results.json','w') as f: json.dump(out,f,indent=2) if __name__=='__main__': main()