Conditional OT barycenter feature augmentation / ot_barycenter_experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED=1162
8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
9
10
11def sinkhorn(a,b,C,eps=0.25,iters=80):
12 # Entropic OT scaling, with strictly positive masses.
13 K=np.exp(-np.clip(C/eps,0,80)); u=np.ones_like(a); v=np.ones_like(b)
14 for _ in range(iters):
15 u=a/(K@v+1e-12); v=b/(K.T@u+1e-12)
16 return (u[:,None]*K)*v[None,:]
17
18
19def barycenter(Zs, ys=None, masses=None, theta=None, m=None, rounds=5, eps=.25):
20 K=len(Zs); n,d=Zs[0].shape; m=m or n
21 theta=np.ones(K)/K if theta is None else np.asarray(theta,float); theta/=theta.sum()
22 if masses is None: masses=[np.ones(len(z))/len(z) for z in Zs]
23 # deterministic support initialization from the first source
24 inds=np.linspace(0,n-1,m).round().astype(int); B=Zs[0][inds].copy(); q=np.ones(m)/m
25 plans=[]
26 for _ in range(rounds):
27 plans=[]
28 for z,a in zip(Zs,masses):
29 C=((z[:,None,:]-B[None,:,:])**2).sum(2)
30 plans.append(sinkhorn(a,q,C,eps=eps,iters=60))
31 den=sum(theta[k]*plans[k].sum(0) for k in range(K))+1e-12
32 B=sum(theta[k]*(plans[k].T@Zs[k]) for k in range(K))/den[:,None]
33 # final plans and transported labels
34 plans=[]
35 for z,a in zip(Zs,masses):
36 C=((z[:,None,:]-B[None,:,:])**2).sum(2)
37 plans.append(sinkhorn(a,q,C,eps=eps,iters=80))
38 if ys is None: return B, plans
39 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)
40 return B,ybar,plans
41
42
43def verify():
44 out={}
45 # Prediction 1: convex combination / convex hull containment. In 1D the barycenter
46 # coordinates must be between the minimum and maximum source support coordinates.
47 violations=[]
48 for trial in range(20):
49 zs=[np.random.randn(12,2)+np.random.randn(2)*2 for _ in range(3)]
50 B,_=barycenter(zs,m=8,rounds=5,eps=.35)
51 lo=min(z[:,0].min() for z in zs); hi=max(z[:,0].max() for z in zs)
52 violations.append(float(max(0,lo-B[:,0].min(),B[:,0].max()-hi)))
53 out['convex_hull']={'prediction':'all barycenter coordinates remain in source coordinate hull','max_violation':max(violations),'pass':max(violations)<2e-5}
54 # Prediction 2: for equally weighted, same-index Gaussian source observations, consensus
55 # variance is sigma^2/K. We use the explicit convex barycentric update (the zero-cost
56 # limit of the OT construction) and sweep K and sigma.
57 rows=[]
58 for K in [2,3,5,8]:
59 for sigma in [.2,.5,1.0]:
60 reps=50000; x=np.random.randn(reps,K)*sigma
61 b=x.mean(1); observed=float(b.var()); predicted=sigma*sigma/K
62 rows.append({'K':K,'sigma':sigma,'observed':observed,'predicted':predicted,'ratio':observed/predicted})
63 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}
64 # Prediction 3: conditional source weights produce affine movement of consensus.
65 # Two deterministic domains at -1,+1; weighted barycenter is 1-2w for source 1.
66 ws=np.linspace(0,1,11); observed=[]
67 for w in ws:
68 observed.append(float((w*(-1)+(1-w)*1)))
69 pred=1-2*ws
70 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}
71 return out
72
73class Net(nn.Module):
74 def __init__(self):
75 super().__init__(); self.h=nn.Sequential(nn.Linear(2,16),nn.Tanh(),nn.Linear(16,8),nn.Tanh()); self.head=nn.Linear(8,2)
76 def forward(self,x):
77 z=self.h(x); return self.head(z),z
78
79def make_domain(n,shift,seed):
80 g=np.random.default_rng(seed); y=g.integers(0,2,n); signal=(2*y-1).astype(float)
81 # domain nuisance is an orthogonal shift; target has an unseen larger shift.
82 x=np.stack([signal+g.normal(0,.7,n), shift+g.normal(0,.8,n)],1).astype('float32')
83 return torch.tensor(x),torch.tensor(y,dtype=torch.long)
84
85def train(mode,steps=180):
86 domains=[make_domain(96,s,10+i) for i,s in enumerate([-2.,0.,2.])]
87 target=make_domain(1000,4.0,99)
88 net=Net(); opt=torch.optim.Adam(net.parameters(),lr=.015)
89 for step in range(steps):
90 xs=[]; ys=[]
91 for x,y in domains:
92 ix=torch.randint(0,len(x),(32,)); xs.append(x[ix]); ys.append(y[ix])
93 if mode=='mixup':
94 x=torch.cat(xs); y=torch.cat(ys); p=torch.randperm(len(x)); lam=.5
95 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])
96 elif mode=='ot':
97 logits=[]; zs=[]
98 for x in xs:
99 l,z=net(x); logits.append(l); zs.append(z)
100 loss=sum(nn.functional.cross_entropy(l,y) for l,y in zip(logits,ys))/3
101 # detach barycenter as proposed initial stable version; labels are transported soft labels
102 with torch.no_grad():
103 Z=[z.detach().cpu().numpy() for z in zs]; Y=[nn.functional.one_hot(y,2).float().cpu().numpy() for y in ys]
104 B,Yb,_=barycenter(Z,Y,m=16,rounds=3,eps=.7)
105 lb=nn.functional.cross_entropy(net.head(torch.tensor(B,dtype=torch.float32)),torch.tensor(Yb,dtype=torch.float32))
106 loss=loss+.35*lb
107 else:
108 x=torch.cat(xs); y=torch.cat(ys); loss=nn.functional.cross_entropy(net(x)[0],y)
109 opt.zero_grad(); loss.backward(); opt.step()
110 with torch.no_grad():
111 acc=(net(target[0])[0].argmax(1)==target[1]).float().mean().item()
112 src=sum((net(x)[0].argmax(1)==y).float().mean().item() for x,y in domains)/3
113 return {'target_accuracy':acc,'source_accuracy':src}
114
115def main():
116 verification=verify()
117 results={m:train(m) for m in ['erm','mixup','ot']}
118 report={'seed':SEED,'verification':verification,'mini_experiment':results}
119 Path('results.json').write_text(json.dumps(report,indent=2))
120 print(json.dumps(report,indent=2))
121
122if __name__=='__main__': main()