Unbalanced Optimal Transport¶
Allow transport to create or remove mass when exact marginal matching is inappropriate. Compare generalized scaling updates with the balanced case and inspect how marginal penalties change the resulting coupling.
Run this tour¶
Run the cells in order with a Python 3 kernel. The first cell locates the companion data and toolbox and installs missing dependencies when needed. All worked examples include their implementation directly in this notebook. Random seeds make comparisons reproducible; you can change them to explore other samples.
# Locate the companion toolbox locally, or fetch it for a standalone/Colab copy.
from pathlib import Path
import importlib.util
import os
import subprocess
import sys
working = Path.cwd()
candidates = [working, working / "python", working.parent / "python"]
python_dir = next((p for p in candidates if (p / "nt_toolbox").is_dir()), None)
if python_dir is None:
checkout = working / "numerical-tours-support"
if not checkout.exists():
subprocess.run(
[
"git",
"clone",
"--depth",
"1",
"--branch",
"master",
"https://github.com/gpeyre/numerical-tours.git",
str(checkout),
],
check=True,
)
python_dir = checkout / "python"
os.chdir(python_dir)
if str(python_dir) not in sys.path:
sys.path.insert(0, str(python_dir))
requirements = python_dir / "requirements.txt"
if any(
importlib.util.find_spec(name) is None
for name in [
"numpy",
"scipy",
"matplotlib",
"skimage",
"sklearn",
"pywt",
"ipywidgets",
"cvxpy",
"skfmm",
"autograd",
"progressbar",
"celer",
]
):
subprocess.run(
[sys.executable, "-m", "pip", "install", "-r", str(requirements)], check=True
)
import numpy as np
import matplotlib.pyplot as plt
np.random.seed(0)
plt.rcParams.update(
{
"figure.figsize": (8, 4),
"figure.dpi": 100,
"axes.spines.top": False,
"axes.spines.right": False,
"font.size": 11,
"image.cmap": "gray",
}
)
%matplotlib inline
$\newcommand{\dotp}[2]{\langle #1, #2 \rangle}$ $\newcommand{\enscond}[2]{\lbrace #1, #2 \rbrace}$ $\newcommand{\pd}[2]{ \frac{ \partial #1}{\partial #2} }$ $\newcommand{\umin}[1]{\underset{#1}{\min}\;}$ $\newcommand{\umax}[1]{\underset{#1}{\max}\;}$ $\newcommand{\uargmin}[1]{\underset{#1}{argmin}\;}$ $\newcommand{\norm}[1]{\|#1\|}$ $\newcommand{\abs}[1]{\left|#1\right|}$ $\newcommand{\choice}[1]{ \left\{ \begin{array}{l} #1 \end{array} \right. }$ $\newcommand{\pa}[1]{\left(#1\right)}$ $\newcommand{\diag}[1]{{diag}\left( #1 \right)}$ $\newcommand{\qandq}{\quad\text{and}\quad}$ $\newcommand{\qwhereq}{\quad\text{where}\quad}$ $\newcommand{\qifq}{ \quad \text{if} \quad }$ $\newcommand{\qarrq}{ \quad \Longrightarrow \quad }$ $\newcommand{\ZZ}{\mathbb{Z}}$ $\newcommand{\CC}{\mathbb{C}}$ $\newcommand{\RR}{\mathbb{R}}$ $\newcommand{\EE}{\mathbb{E}}$ $\newcommand{\Zz}{\mathcal{Z}}$ $\newcommand{\Ww}{\mathcal{W}}$ $\newcommand{\Vv}{\mathcal{V}}$ $\newcommand{\Nn}{\mathcal{N}}$ $\newcommand{\NN}{\mathcal{N}}$ $\newcommand{\Hh}{\mathcal{H}}$ $\newcommand{\Bb}{\mathcal{B}}$ $\newcommand{\Ee}{\mathcal{E}}$ $\newcommand{\Cc}{\mathcal{C}}$ $\newcommand{\Gg}{\mathcal{G}}$ $\newcommand{\Ss}{\mathcal{S}}$ $\newcommand{\Pp}{\mathcal{P}}$ $\newcommand{\Ff}{\mathcal{F}}$ $\newcommand{\Xx}{\mathcal{X}}$ $\newcommand{\Mm}{\mathcal{M}}$ $\newcommand{\Ii}{\mathcal{I}}$ $\newcommand{\Dd}{\mathcal{D}}$ $\newcommand{\Ll}{\mathcal{L}}$ $\newcommand{\Tt}{\mathcal{T}}$ $\newcommand{\si}{\sigma}$ $\newcommand{\al}{\alpha}$ $\newcommand{\la}{\lambda}$ $\newcommand{\ga}{\gamma}$ $\newcommand{\Ga}{\Gamma}$ $\newcommand{\La}{\Lambda}$ $\newcommand{\Si}{\Sigma}$ $\newcommand{\be}{\beta}$ $\newcommand{\de}{\delta}$ $\newcommand{\De}{\Delta}$ $\newcommand{\phi}{\varphi}$ $\newcommand{\th}{\theta}$ $\newcommand{\om}{\omega}$ $\newcommand{\Om}{\Omega}$
This numerical tour details how to perform "unbalanced" OT, which allows one to compare measures with different total mass, and also leads to more regular transportation plans by enabling creation/destruction of mass. This extension of classicl ("balanced") OT is crucial for applications of OT to imaging sciences and machine learning, since it makes OT more robust to noise and outliers. The modification with respect to the usual OT is very minor, since it corresponds a penalization of the mass conservation constraint.
The original idea can be found in the paper of Matthias Liero, Alexander Mielke and Giuseppe Savaré. The entropic regularized version with the corresponding Sinkhorn's algorithm can be found in the paper of Lenaic Chizat, Gabriel Peyré, Bernhard Schmitzer and François-Xavier Vialard.
We now install CVXPY. Warning: seems to not be working on Python 3.7, use rather 3.6.
import numpy as np
import matplotlib.pyplot as plt
import cvxpy as cp
Definition of the input measures¶
For the sake of concreteness and to ease the display, we consider the transport between 1D distributions. But this can be applied to any OT problem.
We consider two dicretes distributions $$ \sum_{i=1}^n a_i \de_{x_i} \qandq \sum_{j=1}^m b_j \de_{y_j}, $$ where $n,m$ are the number of points, $\de_x$ is the Dirac at location $x$, and $(x_i)_i, (y_j)_j$ are the positions of the diracs (in some metric space, in the following we consider the space to be $\RR$).
n = int(120)
m = int(110)
We consider two Gaussian measures sampled on a 1-D grid.
Gaussian = lambda t0, sigma, N: np.exp(
-((np.arange(0, N) / N - t0) ** 2) / (2 * sigma**2)
)
normalize = lambda p: p / np.sum(p)
sigma = 0.06
a = Gaussian(0.25, sigma, n)
b = Gaussian(0.8, sigma, m)
Add some minimal mass and normalize. Here we do not use the same total mass for $a$ and $b$.
vmin = 0.01
a = 0.95 * normalize(a + np.max(a) * vmin)
b = 1.05 * normalize(b + np.max(b) * vmin)
x = np.arange(0, n) / n
y = np.arange(0, m) / m
Display the histograms.
plt.bar(x, a, width=1 / n, color="b")
plt.bar(y, b, width=1 / m, color="r");
Compute the cost matrix, here we use Euclidean distance squared, $C_{i,j} = \norm{x_i-x_j}^2$.
C = np.abs(x[:, None] - y[None, :]) ** 2
plt.imshow(C);
Kantorovitch-Hellinger / Wasserstein-Fisher-Rao Transport¶
The unbalanced OT problem corresponds to $$ W_\rho(a,b) \triangleq \umin{P \in \RR_+^{n \times m}} \dotp{P}{C} + \rho D_\phi( P 1_m|a ) + \rho D_\phi( P^\top 1_n|b ), $$
where here $D_\phi$ is a so-called Cizarr f-divergence $$ D_\phi(h|b) \triangleq \sum_{i} \phi(h_i/b_i) b_i. $$
The most well known are the KL divergence obaind when using $\phi(s)=s \log(s)-s+1$ and the total variation for $\phi(s)=|s-1|$. These are the two examples we will consider in this tour.
The parameter $\rho$ controls the amount of mass conservation relaxation. When $\rho \rightarrow +\infty$ one recovers the usual (balanced) OT. When $\rho \rightarrow 0$, no transport is performed.
Define the optimiztion variable $P$, the OT coupling.
P = cp.Variable((n, m))
We first consider the case of $\phi(s)=s \log(s)-s+1$, in which case the unbalanced OT problem is called either "Kantorovitch-Hellinger" or "Wasserstein-Fisher-Rao".
In this case, assuming $n=m$ and that the sampling points are equal, $x_i=y_i$, one has that $W_\rho/\rho$ converges toward the squared Hellinger distance as $\rho \to +\infty$ $$ W_\rho(a,b)/\rho \longrightarrow \sum_{i} (\sqrt{a_i}-\sqrt{b_i})^2. $$
We define the CVXPY problem and solve it.
u = np.ones((n, 1))
v = np.ones((m, 1))
q = cp.sum(cp.kl_div(cp.matmul(P, v), a[:, None]))
r = cp.sum(cp.kl_div(cp.matmul(P.T, u), b[:, None]))
constr = [0 <= P]
# uncomment to perform balanced OT
# constr = [0 <= P, cp.matmul(P,u)==a[:,None], cp.matmul(P.T,v)==b[:,None]]
rho = 0.1
objective = cp.Minimize(cp.sum(cp.multiply(P, C)) + rho * q + rho * r)
prob = cp.Problem(objective, constr)
result = prob.solve(solver="CLARABEL", warm_start=False)
assert prob.status in {"optimal", "optimal_inaccurate"}
assert np.isfinite(result) and np.isfinite(P.value).all()
Display the solution coupling.
def remap_plan(P): # boost contrast
return np.log(0.001 + P)
plt.figure(figsize=(5, 5))
plt.imshow(remap_plan(P.value));
Display the marginals.
a1 = np.sum(P.value, axis=1)
b1 = np.sum(P.value.T, axis=1)
plt.bar(x, a1, width=1 / n, color="b")
plt.bar(y, b1, width=1 / m, color="r")
plt.bar(x, a, width=1 / n, color="b", alpha=0.2)
plt.bar(y, b, width=1 / m, color="r", alpha=0.2);
Shows the impact of $\rho$ on the solution.
rho_list = np.array([0.03, 0.1, 0.5, 1])
for k in range(len(rho_list)):
rho = rho_list[k]
objective = cp.Minimize(cp.sum(cp.multiply(P, C)) + rho * q + rho * r)
prob = cp.Problem(objective, constr)
result = prob.solve(solver="CLARABEL", warm_start=False)
assert prob.status in {"optimal", "optimal_inaccurate"}
assert np.isfinite(result) and np.isfinite(P.value).all()
ax = plt.subplot(1, len(rho_list), k + 1)
plt.imshow(remap_plan(P.value))
ax.set(xticks=[], yticks=[])
plt.tight_layout()
for k in range(len(rho_list)):
rho = rho_list[k]
objective = cp.Minimize(cp.sum(cp.multiply(P, C)) + rho * q + rho * r)
prob = cp.Problem(objective, constr)
result = prob.solve(solver="CLARABEL", warm_start=False)
assert prob.status in {"optimal", "optimal_inaccurate"}
assert np.isfinite(result) and np.isfinite(P.value).all()
a1 = np.sum(P.value, axis=1)
b1 = np.sum(P.value.T, axis=1)
ax = plt.subplot(len(rho_list), 1, k + 1)
plt.bar(x, a1, width=1 / n, color="b")
plt.bar(y, b1, width=1 / m, color="r")
plt.bar(x, a, width=1 / n, color="b", alpha=0.2)
plt.bar(y, b, width=1 / m, color="r", alpha=0.2)
ax.set(xticks=[], yticks=[])
plt.tight_layout()
Partial Optimal Transport¶
We can consider other divergences, such as the total variation, which corresponds to the $\ell^1$ norm of densities, obtained for $\phi(s)=|s-1|$ $$ D_\phi(h|a) = \norm{a-h}_1 = \sum_i |a_i-h_i|. $$ The resulting OT problem corresponds to a penalized version of the celebrated partial transport problem. In sharp contrast to the KL problem, this partial OT either transport the mass or detroys it.
q = cp.sum(cp.abs(cp.matmul(P, v) - a[:, None]))
r = cp.sum(cp.abs(cp.matmul(P.T, u) - b[:, None]))
Display the marginals of the optimal plan. One can see the presence of small spikes, which are caused by the discretization of the problem. We have displayed up-side down the densities to highlight that these error actually almost cancel.
rho = 0.1
objective = cp.Minimize(cp.sum(cp.multiply(P, C)) + rho * q + rho * r)
prob = cp.Problem(objective, constr)
result = prob.solve(solver="CLARABEL", warm_start=False)
assert prob.status in {"optimal", "optimal_inaccurate"}
assert np.isfinite(result) and np.isfinite(P.value).all()
a1 = np.sum(P.value, axis=1)
b1 = np.sum(P.value.T, axis=1)
plt.bar(x, a1, width=1 / n, color="b")
plt.bar(y, -b1, width=1 / m, color="r")
plt.bar(x, a, width=1 / n, color="b", alpha=0.2)
plt.bar(y, -b, width=1 / m, color="r", alpha=0.2);
Display the impact of $\rho$.
rho_list = np.array([0.05, 0.1, 0.2, 5])
for k in range(len(rho_list)):
rho = rho_list[k]
objective = cp.Minimize(cp.sum(cp.multiply(P, C)) + rho * q + rho * r)
prob = cp.Problem(objective, constr)
result = prob.solve(solver="CLARABEL", warm_start=False)
assert prob.status in {"optimal", "optimal_inaccurate"}
assert np.isfinite(result) and np.isfinite(P.value).all()
a1 = np.sum(P.value, axis=1)
b1 = np.sum(P.value.T, axis=1)
ax = plt.subplot(len(rho_list), 1, k + 1)
plt.bar(x, a1, width=1 / n, color="b")
plt.bar(y, b1, width=1 / m, color="r")
plt.bar(x, a, width=1 / n, color="b", alpha=0.2)
plt.bar(y, b, width=1 / m, color="r", alpha=0.2)
ax.set(xticks=[], yticks=[])
plt.tight_layout()
We can compare several $\phi$-divergence.
Entropic Regularization and Sinkhorn¶
It is possible to regularized the initial problem using entropic regularization and consider $$ \umin{P \in \RR_+^{n \times m}} \dotp{P}{C} + \rho KL( P 1_m|a ) + \rho KL( P^\top 1_n|b ) + \varepsilon KL(P|ab^\top). $$ Here $\varepsilon>0$ controls the strength of the regularization, increasing it results in faster algorithms but degrades the approximation.
The solution $P$ can be shown to be of the form $$ P_{i,j} = e^{ \frac{f_i+g_j-C_{i,j}}{\epsilon} } a_i b_j $$ where the dual variables $f \in \RR^n, g \in \RR^m$ satisfies the following scaled Sinkhorn iterations (written here in log domain): $$ f_i = -\varepsilon \kappa \log \sum_{j} \exp\pa{ \frac{g_j-C_{i,j}}{\epsilon} } b_j $$ $$ g_j = -\varepsilon\kappa \log \sum_{i} \exp\pa{ \frac{f_i-C_{i,j}}{\epsilon} } a_j $$ where we noted $$ \kappa \triangleq \frac{\rho}{\varepsilon + \rho} . $$ Sinkhorn's algoritm simply iterates these two fixed points.
We define the log-sum-exp operator (which corresponds to soft $C$-transforms).
def mina_u(H, epsilon):
return -epsilon * np.log(np.sum(a[:, None] * np.exp(-H / epsilon), 0))
def minb_u(H, epsilon):
return -epsilon * np.log(np.sum(b[None, :] * np.exp(-H / epsilon), 1))
They can be stabilized using the usual log-sum-exp trick.
def mina(H, epsilon):
return mina_u(H - np.min(H, 0), epsilon) + np.min(H, 0)
def minb(H, epsilon):
return minb_u(H - np.min(H, 1)[:, None], epsilon) + np.min(H, 1)
Values of $\varepsilon, \rho$ and $\kappa$.
epsilon = 0.001
rho = 0.2
kappa = rho / (rho + epsilon)
Implement Sinkhorn's iterates.
f = np.zeros(n)
niter = 1000
for it in range(niter):
g = kappa * mina(C - f[:, None], epsilon)
f = kappa * minb(C - g[None, :], epsilon)
# generate the coupling
P = a[:, None] * np.exp((f[:, None] + g[None, :] - C) / epsilon) * b[None, :]
Display the optimal plan.
plt.imshow(remap_plan(P))
Display the marginals.
a1 = np.sum(P, axis=1)
b1 = np.sum(P.T, axis=1)
plt.bar(x, a1, width=1 / n, color="b")
plt.bar(y, b1, width=1 / m, color="r")
plt.bar(x, a, width=1 / n, color="b", alpha=0.2)
plt.bar(y, b, width=1 / m, color="r", alpha=0.2);
Worked example (easy): Experiments with different values of $\varepsilon$ and $\rho$. Study the rate of convergence of the method.
Worked example (hard): Extend Sinkhorn for other type of divergence, starting with TV.
References and further reading¶
Lénaïc Chizat, Gabriel Peyré, Bernhard Schmitzer, and François-Xavier Vialard. Scaling Algorithms for Unbalanced Transport Problems. 2018, Mathematics of Computation 87, 2563–2609. Relaxed marginal constraints and generalized Sinkhorn iterations.
Gabriel Peyré and Marco Cuturi. Computational Optimal Transport. 2019, Foundations and Trends in Machine Learning 11(5–6), 355–607. Transport plans, duality, entropic regularization, and applications.
Marco Cuturi. Sinkhorn Distances: Lightspeed Computation of Optimal Transport. 2013, NeurIPS. Entropic regularization and scalable matrix scaling.
Jean-David Benamou, Guillaume Carlier, Marco Cuturi, Luca Nenna, and Gabriel Peyré. Iterative Bregman Projections for Regularized Transportation Problems. 2015, SIAM Journal on Scientific Computing 37(2), A1111–A1138. Projection-based algorithms for transport and barycenters.
Neal Parikh and Stephen Boyd. Proximal Algorithms. 2014, Foundations and Trends in Optimization 1(3), 127–239. Proximity operators and splitting methods for nonsmooth objectives.