File size: 8,001 Bytes
d8f9639
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""CIS 6270 Lecture 7. Small, executable OT and bridge calculations.

Run: python ot_sbm_examples.py --output results
The examples use synthetic distributions. Paper-scale benchmarks remain in
original repositories. Every displayed numerical result is recomputed here.
"""
from pathlib import Path
import argparse,itertools,json
import numpy as np
from scipy.special import logsumexp
from scipy.optimize import linprog
from scipy.linalg import expm


def sinkhorn(a,b,log_kernel,iterations=2000,tol=1e-12):
    """Scale a positive kernel to prescribed marginals in log space."""
    a,b=np.asarray(a,float),np.asarray(b,float)
    if np.any(a<=0) or np.any(b<=0):raise ValueError('Use positive marginals on the retained support.')
    if not np.isclose(a.sum(),b.sum()):raise ValueError('Marginals must have equal mass.')
    log_v=np.zeros_like(b);history=[]
    for k in range(iterations):
        log_u=np.log(a)-logsumexp(log_kernel+log_v[None,:],axis=1)
        log_v=np.log(b)-logsumexp(log_kernel+log_u[:,None],axis=0)
        pi=np.exp(log_u[:,None]+log_kernel+log_v[None,:])
        error=max(abs(pi.sum(1)-a).max(),abs(pi.sum(0)-b).max())
        history.append(float(error))
        if error<tol:break
    return pi,np.exp(log_u),np.exp(log_v),history


def transport_example():
    a=np.array([.6,.4]);b=np.array([.3,.7]);C=np.array([[1.,9.],[1.,1.]])
    A=np.array([[1,1,0,0],[0,0,1,1],[1,0,1,0],[0,1,0,1]])
    fit=linprog(C.ravel(),A_eq=A,b_eq=np.r_[a,b],bounds=(0,None),method='highs')
    assert fit.success
    pi,u,v,history=sinkhorn(a,b,-C/2)
    K=np.exp(-C/2);u1=a/K.sum(1);v1=b/(K.T@u1)
    f=np.array([0.,-8.]);g=np.array([1.,9.])
    dual=float(a@f+b@g)
    return dict(a=a,b=b,cost=C,plan=fit.x.reshape(2,2),ot_cost=fit.fun,dual=dual,
                entropic_plan=pi,entropic_cost=float((pi*C).sum()),kernel=K,
                first_u=u1,first_v=v1,first_plan=u1[:,None]*K*v1[None,:],
                marginal_errors=history)


def finite_bridge_example(steps=2):
    a=np.array([.5,.5]);b=np.array([.2,.8]);R=np.array([[.8,.2],[.3,.7]])
    K=np.linalg.matrix_power(R,steps)
    pi,u,v,_=sinkhorn(a,b,np.log(a[:,None]*K))
    h=[np.linalg.matrix_power(R,steps-k)@v for k in range(steps+1)]
    Q=[R*h[k+1][None,:]/h[k][:,None] for k in range(steps)]
    p=[a]
    for q in Q:p.append(p[-1]@q)
    paths=np.array(list(itertools.product(range(2),repeat=steps+1)))
    ref=np.array([a[x[0]]*np.prod([R[x[k],x[k+1]] for k in range(steps)]) for x in paths])
    star=np.array([ref[i]*u[x[0]]*v[x[-1]] for i,x in enumerate(paths)])
    kl=float(np.sum(star*np.log(star/ref)))
    return dict(R=R,K=K,pi=pi,u=u,v=v,h=h,Q=Q,marginals=p,paths=paths,reference=ref,bridge=star,kl=kl)


def discrete_imf_example(steps=3,iterations=80):
    """Exact Markovian and reciprocal projections on a 16-path space."""
    a=np.array([.5,.5]);b=np.array([.2,.8]);R=np.array([[.8,.2],[.3,.7]])
    paths=np.array(list(itertools.product(range(2),repeat=steps+1)))
    ref=np.array([a[x[0]]*np.prod([R[x[k],x[k+1]] for k in range(steps)]) for x in paths])
    joint=a[:,None]*np.linalg.matrix_power(R,steps)
    pi,_,_,_=sinkhorn(a,b,np.log(joint))
    star=np.array([ref[i]*pi[x[0],x[-1]]/joint[x[0],x[-1]] for i,x in enumerate(paths)])
    p=np.array([ref[i]*a[x[0]]*b[x[-1]]/joint[x[0],x[-1]] for i,x in enumerate(paths)])
    hist=[]
    for it in range(iterations):
        qs=[]
        for k in range(steps):
            J=np.zeros((2,2))
            for mass,x in zip(p,paths):J[x[k],x[k+1]]+=mass
            qs.append(J/J.sum(1,keepdims=True))
        m=np.array([a[x[0]]*np.prod([qs[k][x[k],x[k+1]] for k in range(steps)]) for x in paths])
        end=np.zeros((2,2))
        for mass,x in zip(m,paths):end[x[0],x[-1]]+=mass
        p=np.array([ref[i]*end[x[0],x[-1]]/joint[x[0],x[-1]] for i,x in enumerate(paths)])
        hist.append(float(np.sum(p*np.log(p/star))))
    return dict(kl_to_bridge=hist,final_path_l1=float(abs(p-star).sum()),paths=paths,bridge=star,final=p)


def ctmc_bridge_example():
    """Doob rates and exact marginal verification with matrix exponentials."""
    a=np.array([.5,.5]);b=np.array([.2,.8]);G=np.array([[-2.,2.],[1.,-1.]])
    K=expm(G);pi,u,v,_=sinkhorn(a,b,np.log(a[:,None]*K))
    values=[]
    for t in [0,.25,.5,.75,1]:
        h=expm((1-t)*G)@v
        forward=(a*u)@expm(t*G)
        p=forward*h
        rate=G*h[None,:]/h[:,None]
        np.fill_diagonal(rate,0);np.fill_diagonal(rate,-rate.sum(1))
        values.append(dict(t=t,h=h,p=p,rates=rate))
    return dict(generator=G,transition=K,coupling=pi,values=values)


def gaussian_example(epsilon=1.):
    """Brownian reference dX=sqrt(epsilon)dW, N(0,1) to N(2,1)."""
    covariance=(np.sqrt(epsilon**2+4)-epsilon)/2
    t=np.linspace(0,1,101)
    variance=(1-t)**2+t*t+2*t*(1-t)*covariance+epsilon*t*(1-t)
    # Forward Markov drift beta=(y-x)/(1-t) averaged conditionally on X_t.
    cov_yx=(1-t)*covariance+t
    slope=(cov_yx/variance-1)/np.maximum(1-t,1e-10)
    slope[-1]=1-covariance-epsilon
    intercept=2-slope*(2*t)
    return dict(t=t,mean=2*t,variance=variance,covariance=covariance,slope=slope,intercept=intercept)


def reward_bridge_example():
    """A fully masked root has one initial state, so terminal tilting preserves it."""
    base=np.array([.5,.3,.2]);reward=np.log(np.array([1.,2.,4.]));alpha=1.
    target=base*np.exp(reward/alpha);target/=target.sum()
    ref_path=np.array([.3,.2,.18,.12,.12,.08])
    terminal=np.array([0,0,1,1,2,2])
    tilted=ref_path*np.exp(reward[terminal]);tilted/=tilted.sum()
    proposal=np.array([.1,.1,.15,.15,.2,.3])
    log_weight=reward[terminal]+np.log(ref_path)-np.log(proposal)
    exact=np.sum(proposal*np.exp(log_weight))
    weighted=proposal*np.exp(log_weight)/exact
    return dict(base=base,reward=reward,target=target,reference_paths=ref_path,
                target_paths=tilted,proposal=proposal,log_weight=log_weight,weighted_paths=weighted,
                kl=float(np.sum(target*np.log(target/base))))


def branching_example():
    """Exact illustrative mass transfer; neural four-stage version is separate."""
    t=np.linspace(0,1,101);w=np.stack([1-t,.6*t,.4*t],1)
    growth=np.tile([-1.,.6,.4],(len(t),1))
    x=np.stack([t,.7*t+1.2*t*t,.7*t-1.2*t*t],1)
    return dict(t=t,weights=w,growth=growth,positions=x,total_mass=w.sum(1),midpoint_weights=w[50])


def entangled_geometry_example():
    """Check the cone guarantee for the bias increment, separately from noise."""
    s=np.array([3.,4.]);shat=s/np.linalg.norm(s);h=np.array([2.,-1.]);alpha=.8
    orthogonal=h-shat*(shat@h);bias=alpha*shat+orthogonal
    max_dt=2*(s@bias)/(bias@bias)
    dt=.1;before=s@s;after=(s-dt*bias)@(s-dt*bias)
    return dict(direction=s,unit_direction=shat,orthogonal=orthogonal,bias=bias,
                alignment=float(s@bias),max_dt=float(max_dt),dt=dt,distance2_before=float(before),distance2_after=float(after))


def serial(x):
    if isinstance(x,np.ndarray):return x.tolist()
    if isinstance(x,np.generic):return x.item()
    if isinstance(x,dict):return {k:serial(v) for k,v in x.items()}
    if isinstance(x,(tuple,list)):return [serial(v) for v in x]
    return x


def main():
    parser=argparse.ArgumentParser();parser.add_argument('--output',default='results');args=parser.parse_args()
    out=Path(args.output);out.mkdir(parents=True,exist_ok=True)
    results={k:f() for k,f in [('transport',transport_example),('finite_bridge',finite_bridge_example),
      ('discrete_imf',discrete_imf_example),('ctmc',ctmc_bridge_example),('gaussian',gaussian_example),
      ('reward',reward_bridge_example),('branch',branching_example),('entangled',entangled_geometry_example)]}
    (out/'results.json').write_text(json.dumps(serial(results),indent=2))
    for k,v in results.items():print(k,'complete')
    print('OT cost',results['transport']['ot_cost'],'bridge terminal',results['finite_bridge']['marginals'][-1],
          'IMF path L1',results['discrete_imf']['final_path_l1'])
if __name__=='__main__':main()