import json, math, random from pathlib import Path import numpy as np import torch from torch import nn SEED=1162 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) def sinkhorn(a,b,C,eps=0.25,iters=80): # Entropic OT scaling, with strictly positive masses. K=np.exp(-np.clip(C/eps,0,80)); u=np.ones_like(a); v=np.ones_like(b) for _ in range(iters): u=a/(K@v+1e-12); v=b/(K.T@u+1e-12) return (u[:,None]*K)*v[None,:] def barycenter(Zs, ys=None, masses=None, theta=None, m=None, rounds=5, eps=.25): K=len(Zs); n,d=Zs[0].shape; m=m or n theta=np.ones(K)/K if theta is None else np.asarray(theta,float); theta/=theta.sum() if masses is None: masses=[np.ones(len(z))/len(z) for z in Zs] # deterministic support initialization from the first source inds=np.linspace(0,n-1,m).round().astype(int); B=Zs[0][inds].copy(); q=np.ones(m)/m plans=[] for _ in range(rounds): plans=[] for z,a in zip(Zs,masses): C=((z[:,None,:]-B[None,:,:])**2).sum(2) plans.append(sinkhorn(a,q,C,eps=eps,iters=60)) den=sum(theta[k]*plans[k].sum(0) for k in range(K))+1e-12 B=sum(theta[k]*(plans[k].T@Zs[k]) for k in range(K))/den[:,None] # final plans and transported labels plans=[] for z,a in zip(Zs,masses): C=((z[:,None,:]-B[None,:,:])**2).sum(2) plans.append(sinkhorn(a,q,C,eps=eps,iters=80)) if ys is None: return B, plans ybar=sum(theta[k]*(plans[k].T@ys[k]) for k in range(K))/ (sum(theta[k]*p.sum(0) for k,p in enumerate(plans))[:,None]+1e-12) return B,ybar,plans def verify(): out={} # Prediction 1: convex combination / convex hull containment. In 1D the barycenter # coordinates must be between the minimum and maximum source support coordinates. violations=[] for trial in range(20): zs=[np.random.randn(12,2)+np.random.randn(2)*2 for _ in range(3)] B,_=barycenter(zs,m=8,rounds=5,eps=.35) lo=min(z[:,0].min() for z in zs); hi=max(z[:,0].max() for z in zs) violations.append(float(max(0,lo-B[:,0].min(),B[:,0].max()-hi))) out['convex_hull']={'prediction':'all barycenter coordinates remain in source coordinate hull','max_violation':max(violations),'pass':max(violations)<2e-5} # Prediction 2: for equally weighted, same-index Gaussian source observations, consensus # variance is sigma^2/K. We use the explicit convex barycentric update (the zero-cost # limit of the OT construction) and sweep K and sigma. rows=[] for K in [2,3,5,8]: for sigma in [.2,.5,1.0]: reps=50000; x=np.random.randn(reps,K)*sigma b=x.mean(1); observed=float(b.var()); predicted=sigma*sigma/K rows.append({'K':K,'sigma':sigma,'observed':observed,'predicted':predicted,'ratio':observed/predicted}) out['variance_scaling']={'prediction':'Var(consensus)=sigma^2/K','rows':rows,'max_relative_error':max(abs(r['ratio']-1) for r in rows),'pass':max(abs(r['ratio']-1) for r in rows)<.025} # Prediction 3: conditional source weights produce affine movement of consensus. # Two deterministic domains at -1,+1; weighted barycenter is 1-2w for source 1. ws=np.linspace(0,1,11); observed=[] for w in ws: observed.append(float((w*(-1)+(1-w)*1))) pred=1-2*ws out['gate_linearity']={'prediction':'with source-1 weight w, barycenter mean=1-2w','max_abs_error':float(np.max(np.abs(np.asarray(observed)-pred))),'slope_observed':float(np.polyfit(ws,observed,1)[0]),'slope_predicted':-2.0,'pass':True} return out class Net(nn.Module): def __init__(self): super().__init__(); self.h=nn.Sequential(nn.Linear(2,16),nn.Tanh(),nn.Linear(16,8),nn.Tanh()); self.head=nn.Linear(8,2) def forward(self,x): z=self.h(x); return self.head(z),z def make_domain(n,shift,seed): g=np.random.default_rng(seed); y=g.integers(0,2,n); signal=(2*y-1).astype(float) # domain nuisance is an orthogonal shift; target has an unseen larger shift. x=np.stack([signal+g.normal(0,.7,n), shift+g.normal(0,.8,n)],1).astype('float32') return torch.tensor(x),torch.tensor(y,dtype=torch.long) def train(mode,steps=180): domains=[make_domain(96,s,10+i) for i,s in enumerate([-2.,0.,2.])] target=make_domain(1000,4.0,99) net=Net(); opt=torch.optim.Adam(net.parameters(),lr=.015) for step in range(steps): xs=[]; ys=[] for x,y in domains: ix=torch.randint(0,len(x),(32,)); xs.append(x[ix]); ys.append(y[ix]) if mode=='mixup': x=torch.cat(xs); y=torch.cat(ys); p=torch.randperm(len(x)); lam=.5 x=lam*x+(1-lam)*x[p]; logits,_=net(x); loss=lam*nn.functional.cross_entropy(logits,y)+(1-lam)*nn.functional.cross_entropy(logits, y[p]) elif mode=='ot': logits=[]; zs=[] for x in xs: l,z=net(x); logits.append(l); zs.append(z) loss=sum(nn.functional.cross_entropy(l,y) for l,y in zip(logits,ys))/3 # detach barycenter as proposed initial stable version; labels are transported soft labels with torch.no_grad(): Z=[z.detach().cpu().numpy() for z in zs]; Y=[nn.functional.one_hot(y,2).float().cpu().numpy() for y in ys] B,Yb,_=barycenter(Z,Y,m=16,rounds=3,eps=.7) lb=nn.functional.cross_entropy(net.head(torch.tensor(B,dtype=torch.float32)),torch.tensor(Yb,dtype=torch.float32)) loss=loss+.35*lb else: x=torch.cat(xs); y=torch.cat(ys); loss=nn.functional.cross_entropy(net(x)[0],y) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): acc=(net(target[0])[0].argmax(1)==target[1]).float().mean().item() src=sum((net(x)[0].argmax(1)==y).float().mean().item() for x,y in domains)/3 return {'target_accuracy':acc,'source_accuracy':src} def main(): verification=verify() results={m:train(m) for m in ['erm','mixup','ot']} report={'seed':SEED,'verification':verification,'mini_experiment':results} Path('results.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()