Conditional OT barycenter feature augmentation / ot_barycenter_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()