File size: 13,369 Bytes
0ba9d09 | 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 167 168 169 | # Lecture 7: Optimal Transport and Schrodinger Bridges
Chapter 7 of the notes.
Every model of Chapters 2 through 6 trains on a pair of samples, and in almost all of them the two
were drawn independently. In this lecture we take the pairing as the quantity we choose. The
first four files price a pairing, relax Monge's map to Kantorovich's coupling, certify the optimum
with prices, smooth the problem with an entropy term that Sinkhorn's algorithm solves by rescaling
rows and columns, and read transport as motion in time, where the minimizer moves every particle in
a straight line. The next five replace the cost matrix by the endpoint law of a reference process,
which turns the same scaling into a Schrodinger bridge: the chain rule puts reference bridges
between every pair of ends, Doob's h-transform makes the bridge a process we can run forward, and
three learning methods differ only in what they simulate and what they regress. Three more carry the
construction to a finite state space, where a rate matrix replaces the drift, where a molecular
graph factorizes into forty-five categorical features, and where discrete time makes the whole path
law a table of sixteen numbers. The last file covers a population that splits into weighted branches
and a system of interacting particles whose own dynamics is the reference.
We compute almost every number below in closed form, following the chapter, and two files train,
`ot_flow_matching.py` and `branched.py`.
| file | command | what it computes |
| ---- | ------- | ---------------- |
| `optimal_transport.py` | `python lecture_7/optimal_transport.py` | the four points, the two-by-two plan and its prices, the sorted pairing on three points, and the Gaussian map |
| `sinkhorn.py` | `python lecture_7/sinkhorn.py` | the Gibbs kernel, the first sweep by hand, the converged entropic plan, and the plan at four values of epsilon |
| `displacement.py` | `python lecture_7/displacement.py` | the kinetic action of one particle and of a population, the displacement interpolation, and its velocity field |
| `ot_flow_matching.py` | `python lecture_7/ot_flow_matching.py` | **trains** two velocity networks on the four-Gaussian target, one per coupling, and measures the straightness of each |
| `static_bridge.py` | `python lecture_7/static_bridge.py` | the eight path probabilities, the chain rule, proportional fitting, the tilted matrices, and the bridge path law |
| `dynamic_bridge.py` | `python lecture_7/dynamic_bridge.py` | Girsanov on a constant control, the backward and Hamilton-Jacobi-Bellman residuals, and the Brownian bridge |
| `dsb.py` | `python lecture_7/dsb.py` | the two half-steps of proportional fitting on path laws, and the diffusion reversal on a Gaussian |
| `dsbm.py` | `python lecture_7/dsbm.py` | the averaged bridge drift at one point, and one round of Markovian fitting on the chain |
| `sf2m.py` | `python lecture_7/sf2m.py` | the two closed-form regression targets at a bridge sample, and the entropic plan of a batch |
| `discrete_bridge.py` | `python lecture_7/discrete_bridge.py` | the two-state continuous-time bridge, its Euler sampler and its reverse rates, and the cost computed twice |
| `ddsbm.py` | `python lecture_7/ddsbm.py` | the single bond with three types, the rates pinned at a terminal graph, and the counting over forty-five features |
| `csbm.py` | `python lecture_7/csbm.py` | the two projections in discrete time, and D-IMF on sixteen paths against the reference's mixing |
| `branched.py` | `python lecture_7/branched.py` | the two-branch energies, the tilt of a reference, and **trains** a bias force on two interacting particles |
Every file accepts `--seed` and `--quiet`, and `--steps` means the quantity that file iterates: the
number of quantile levels, Sinkhorn sweeps, times on a grid, optimizer steps, fitting rounds or
Euler steps. The slowest default run is about seven seconds; all thirteen run offline on a laptop
CPU.
## What a default run produces
`optimal_transport.py`. The ordered pairing of the four points costs 1 against 5 for the crossed
one and 3 for the independent coupling. The two-by-two optimum is `[[0.3, 0.3], [0, 0.4]]` at cost
3.4, the prices `f = (0, -8)` and `g = (1, 9)` have slack `[[0, 0], [8, 0]]` and dual value 3.4, so
the duality gap is zero and `W_2 = 1.8439`. The six permutations of the three-point example cost
11/3, 17/3, 35/3, 15, 77/3 and 27, and the quantile integral on 3072 levels returns 3.6667.
We transport N(0,1) to N(3,4) and get `W_2^2 = 10` under the map `T(x) = 3 + 2x`, and the same map
on a stratified sample gives 9.9996.
`sinkhorn.py`. The kernel is `[[0.606531, 0.011109], [0.606531, 0.606531]]`. The first sweep gives
`u = (0.971440, 0.329744)` and `v = (0.38013, 3.32081)`, and the marginal error falls 0.680, 0.382,
0.155, 0.047, 0.012, reaching 1e-9 after 17 sweeps. The entropic plan is
`[[0.293122, 0.306878], [0.006878, 0.393122]]` at cost 3.45502 against 3.4, and the regularized
objective is 1.208 at it, 1.222 at the exact optimum and 1.792 at independence. At epsilon = 10,
2, 0.5 and 0.1 the cost is 4.0540165, 3.4550204, 3.4000004 and 3.4000000, in 7, 17, 28 and 80
sweeps.
`displacement.py`. A steady particle has action 4 and a dawdling one 8, for the same endpoints. The
average squared displacement is 1, 3 and 5 under the ordered, independent and crossed pairings. At
t = 1/2 the displacement interpolation is N(1.5, 2.25) and the mixture of the two ends peaks at
0.13, with total variation 0.257 between them, and `W_2(rho_0, rho_1/2) = 1.5811`, half of 3.1623.
The straight paths have action 9.9987, which is their average squared displacement. The action
density is 9.9994 at every time on the grid, the stratified estimate of the exact 10, and the
continuity residual is 1.1e-16.
`ot_flow_matching.py`. On the two-point batch the independent coupling has average squared velocity
3 and the optimal one 1. Over training the batch coupling cost is 10.31 against 3.25, the learned
paths have straightness 1.80 against 0.065, the one-step error is 2.56 against 0.41, and the samples
from both models cover all four modes. At `--steps 1500`, about eighteen seconds, the optimal
coupling reaches straightness 0.024 and one-step error 0.28, and the independent one stays where it
was.
`static_bridge.py`. The eight reference paths sum to one with `R(1,1,2) = 0.105`. The chain rule
splits 0.2180 into 0.0201 and 0.1979. The reference lands 0.73 away from the requested terminal
law, and proportional fitting reaches `f = (1.072002, 0.937019)` and `g = (0.350517, 1.843637)`,
with the terminal error falling by 187.5 per round. The bridge coupling is
`[[0.114605, 0.385395], [0.085395, 0.414605]]` at `KL = 0.2820` against 0.2847 for independence.
The tilted matrices map (0.5, 0.5) to (0.449213, 0.550787) to (0.2, 0.8), the path ratios have
four values, and `KL(P* || R) = 0.281961`.
`dynamic_bridge.py`. The relative entropy of a constant control u = 2 is 2 from the control energy,
2 from the two endpoint laws and 1.9965 from the simulated paths. The backward residual is 2.2e-16
and the Hamilton-Jacobi-Bellman residual 8.9e-16, and `-grad V` equals the bridge drift to 4.4e-16.
At t = 0.25 the bridge has mean 0.5 and variance 0.1875, and at the state 0.6732 its drift is
1.7691 against the straight-line velocity 2. At t = 1/2 the simulated mean is 1.0095 and the
simulated variance 0.2563, against the exact 1 and 0.25.
`dsb.py`. The reversal of the last step is `[[0.681416, 0.318584], [0.379310, 0.620690]]`. The
backward chain started from (0.2, 0.8) reaches (0.466585, 0.533415), which is 0.0668 from the source
law, and the second half-step returns the terminal law (0.201967, 0.798033), off by 0.0039. That is
the round-1 error of proportional fitting on the coupling, and the error falls by 187.5 per round.
Run backward from its own terminal law, the diffusion reversal lands at variance 0.9836 against the
exact 1.
`dsbm.py`. Two bridges through x = 1 at t = 0.5 have drifts 4 and -2 and average 1, where the
derivative of the squared error vanishes. One round of Markovian fitting gives
`M_0 = [[0.591425, 0.408575], [0.307692, 0.692308]]` and
`M_1 = [[0.298457, 0.701543], [0.119588, 0.880412]]`, whose endpoint coupling has exact marginals
and brings the gap from 0.0584 to 0.0077. The gap then falls by 7.37 per round and nine rounds reach
1e-9.
`sf2m.py`. At t = 0.25 with z = 0.4 the sample is 0.6732, the drift 1.7691, the score -0.9238, the
flow target 2.2309 and the scaled score exactly -0.4. At t = 0.75 the coefficient changes sign and
the flow target is 1.7691, so the two lie symmetrically about the straight-line velocity 2.
`discrete_bridge.py`. We pay 0.1534 to halve every rate and 0.3863 to double it. The two-state
bridge has `g = (0.3349, 1.9643)`, `phi_1/2 = (0.7999, 1.2667)` and the generator
`[[-1.5836, 1.5836], [0.9472, -0.9472]]` at t = 1/2, where the marginal is (0.4668, 0.5332) and the
reverse rates 1.3862 and 1.0821 match the forward flux at 0.7392. The Euler sampler's terminal
error is 0.0223, 0.0106, 0.0051 and 0.0020 as the step is halved twice and then taken to a hundred
steps. The cost is 0.3235 on the endpoints and 0.3235 on the paths, with the integrand 0.1300 at
t = 1/2, and it rises 0.3148, 0.3235, 0.3337, 0.3348 toward the limit 0.3348 as the reference
speeds up.
`ddsbm.py`. The bond reference has `beta = 0.9163` and off-diagonal rate 0.3054, and the first
column scaling is (0.2083, 0.7143, 2.9167). Two hundred sweeps give `f = (1.1779, 0.9506, 0.5062)`
and `g = (0.1842, 0.6918, 3.0004)`, a plan that moves 0.4948 of the mass from no bond to double and
leaves 0.2612 unchanged, between 0.18 for independence and 0.4 for the edit-minimizing plan. At
r = 0.2, 0.4, 0.6, 0.8 and 0.99 the mass that stays is 0.2250, 0.2612, 0.2926, 0.3248 and 0.3812.
The survival ratio reaches 1/2 at t = 0.2435, where the pinned rates are 1.2217, 0.0764 and 0.3054,
a factor of 16 between a move onto the terminal type and a move off it, and one such edge against a
network rate of 1 contributes 0.0229. With two terminal types the network learns 0.8552 at a loss
of 0.1357, against 0.1467 at 1 and 0.2395 at 0.5. Nine heavy atoms give 45 features and 9.2e+27
graphs, at a total jump rate of 24.74 per unit time.
`csbm.py`. With alpha = 0.35 and three transitions, `A^3 = [[0.5135, 0.4865], [0.4865, 0.5135]]`
and the exact bridge coupling is `[[0.10432, 0.39568], [0.09568, 0.40432]]`. The path 1,1,1,1 has
reference bridge weight 0.5348, so its probability is 0.0535 at the start and 0.0558 under the
bridge. The two laws agree at the interior times, (0.4754, 0.5246) and (0.4107, 0.5893), and differ
by 0.0173 over the paths. One round of D-IMF brings that to 2.9e-04 and later rounds divide it by
58.9, reaching 2.4e-11 after five. The last conditional (0.3160, 0.6840) against a network at
(0.25, 0.75) costs 0.0110 nats. At alpha = 0.2, 0.35 and 0.4 the starting gap is 0.13589, 0.01728
and 0.00512 and the rounds divide it by 5, 59 and 284; on 8, 16 and 64 paths it is 0.05743, 0.01728
and 0.00156 and the rounds divide it by 7, 59 and 4674. The ordered reference out of the middle of
five categories is (0.0545, 0.2442, 0.4026, 0.2442, 0.0545).
`branched.py`. The two branches have weighted energies 1.7 and 0.3 for a total of 2.0, and with the
second target at -4 the linear schedule costs 2.9 against 2.6 for one that releases mass late. The
bent interpolant has energy 2.1667 against 2 for the straight line. The two-outcome tilt has
`Z = 0.109`, raises the chance of reaching B from 0.1 to 0.917, and costs 1.836 nats, with
`J(b*) = -log Z = 2.216`. We train the bias force and the mean terminal reward rises from -21.9
under the reference to -6.9, and that figure holds as the run lengthens, because the importance
weights concentrate as the control grows and the effective sample size falls with it.
## What is simplified
- Every example runs on the chapter's own two states, three bond types, four points or two
particles.
- The minibatch coupling of `ot_flow_matching.py` is an exact assignment through
`scipy.optimize.linear_sum_assignment`, which requires the two sides to be the same size. The
chapter's general `n` by `m` problem with unequal masses appears only in `sinkhorn.py`.
- We run the two training files with small networks and short budgets, so we read the direction of
an effect from their reports and stop short of its converged size.
- `dsb.py` and `dsbm.py` carry out their projections exactly on the full table of paths. We stop at
the exact projections, and the chapter goes on to replace them by a regression once the state
space is large.
- `discrete_bridge.py` and `ddsbm.py` work one categorical feature at a time. `ddsbm.py` counts the
forty-five features of a molecular graph and the rate of the product process, and the file stops
at the counting.
- `csbm.py` runs both projections exactly on the sixteen-entry table, and we stop there too. The
chapter carries on, replacing the conditionals by a network and the exact endpoint coupling by
simulation.
- `branched.py` uses constant growth rates taken from the worked example and trains the bias force
of the entangled construction. The file stops at that one force, and the four-stage training of
the branched construction and the fitting of the state cost `V_t` stay with the chapter.
|