Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -480,3 +480,5 @@ Artificial Intelligence.
\[92] Xie, Y., Wang, X., Wang, R., & Zha, H. (2020, August).
[A fast proximal point method for computing exact wasserstein distance.](https://proceedings.mlr.press/v115/xie20b/xie20b.pdf) In Uncertainty in artificial intelligence (pp. 433-453). PMLR.

\[93] Nguyen, K., Bariletto, N., & Ho, N. (2024). [Quasi-Monte Carlo for 3D Sliced Wasserstein](https://arxiv.org/abs/2309.11713). International Conference on Learning Representations (ICLR).

9 changes: 9 additions & 0 deletions RELEASES.md
Original file line number Diff line number Diff line change
@@ -1,4 +1,13 @@
# Releases
## 0.9.8dev
*August 2026*

#### New features

- Add Quasi-Monte Carlo sliced Wasserstein sampling (QSW/RQSW) via generalized
spiral points, selectable with `sampling_slices` in `sliced_wasserstein_distance`,
as described in [93] (PR #838)


## 0.9.7.post1

Expand Down
192 changes: 192 additions & 0 deletions examples/sliced-wasserstein/plot_qsw_3d.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,192 @@
# -*- coding: utf-8 -*-
"""
=========================================================
Quasi-Monte Carlo Sliced Wasserstein in 3D
=========================================================

This example illustrates the Quasi-Sliced Wasserstein (QSW) and Randomized
Quasi-Sliced Wasserstein (RQSW) sampling schemes introduced in [93], and
compares them to the default uniform (Monte Carlo) sampling of slicing
directions.

Sliced Wasserstein (SWD) approximates the Wasserstein distance by averaging
1D Wasserstein distances over projections onto random directions
:math:`\\theta` drawn uniformly on the sphere. By default these directions
are sampled purely at random (Monte Carlo), which introduces some variance
in the estimate for a given number of projections.

QSW replaces the random directions with a deterministic, low-discrepancy
point set on the sphere (generalized spiral points), which covers the
sphere more evenly than random sampling and reduces the approximation
error, especially in 3D. Since QSW is deterministic it cannot directly be
used as an unbiased estimator in stochastic settings (e.g. gradient-based
optimization) -- RQSW addresses this by applying a random rotation to the
same point set, which preserves both its low discrepancy and its
unbiasedness.

We first visualize the three sampling schemes on the sphere, then measure
how fast each one converges to the true Sliced Wasserstein distance
between two point clouds -- known here in closed form, with no
approximation error left except from the number of projections itself.

.. [93] Nguyen, K., Bariletto, N., & Ho, N. (2024). Quasi-Monte Carlo for
3D Sliced Wasserstein. International Conference on Learning
Representations (ICLR).
"""

# Author: Samuel Vangu <samuelvangu0@gmail.com>
#
# License: MIT License

# sphinx_gallery_thumbnail_number = 1

import numpy as np
import matplotlib.pylab as pl
from mpl_toolkits.mplot3d import Axes3D # noqa: F401 (registers the 3D projection)

import ot
from ot.sliced import get_random_projections, get_projections_spiral

##############################################################################
# Visualize the three sampling schemes on the sphere
# ----------------------------------------------------
# We draw a few hundred directions on :math:`S^2` with each scheme:
#
# - ``uniform``: directions are Gaussian vectors normalized to unit norm
# (standard Monte Carlo sampling of the sphere).
# - ``qsw``: deterministic generalized spiral points -- a simple,
# closed-form low-discrepancy point set (Rakhmanov, Saff & Zhou, 1994).
# The same call always returns the same points.
# - ``rqsw``: the same spiral point set, rotated by a random (3, 3)
# rotation matrix (drawn via QR decomposition of a Gaussian matrix).
# The rotation makes the estimator unbiased while keeping the points
# as evenly spread out as the deterministic QSW set.

n_projections = 500
d = 3
seed = 42

theta_uniform = get_random_projections(d, n_projections, seed=seed)
theta_qsw = get_projections_spiral(d, n_projections, randomized=False)
theta_rqsw = get_projections_spiral(d, n_projections, randomized=True, seed=seed)

fig = pl.figure(1, figsize=(15, 5))

schemes = [
(theta_uniform, "Uniform (Monte Carlo)"),
(theta_qsw, "QSW (deterministic spiral)"),
(theta_rqsw, "RQSW (randomly rotated spiral)"),
]

for i, (theta, title) in enumerate(schemes):
ax = fig.add_subplot(1, 3, i + 1, projection="3d")
ax.scatter(theta[0], theta[1], theta[2], c=theta[2], cmap="viridis", s=4, alpha=0.8)
ax.set_title(title)
ax.set_box_aspect([1, 1, 1])
ax.view_init(elev=20, azim=45)
ax.set_xticks([])
ax.set_yticks([])
ax.set_zticks([])

pl.tight_layout()
pl.show()

# Notice how the uniform sample leaves visible gaps and clusters, while QSW
# and RQSW spread the points much more evenly over the sphere -- this is
# exactly the low-discrepancy property that reduces the error of the Sliced
# Wasserstein estimate.

##############################################################################
# Convergence to the true Sliced Wasserstein distance
# ------------------------------------------------------
# We now compare how fast each sampling scheme converges to the *true*
# SWD as the number of projections grows. To get a reference value with
# **zero** approximation error -- not even from a finite number of
# samples -- we build ``Xt`` as a pure translation of ``Xs`` by a fixed
# vector :math:`\delta`: ``Xt = Xs + delta``.
#
# For a rigid translation, the classical 1D Wasserstein identity
# :math:`W_2(\mu, \mu + c) = |c|` holds *exactly*, for any distribution
# shape and any (even very small) sample size -- no law-of-large-numbers
# argument, no Gaussian assumption, just an algebraic identity of optimal
# transport on the line. Projected onto any direction :math:`\theta`, this
# gives :math:`W_2(\theta_\# \mu, \theta_\# \nu) = |\theta^T \delta|`
# exactly, and averaging the square over :math:`\theta` uniform on
# :math:`S^{d-1}` gives the closed-form identity
#
# .. math::
# \mathcal{SWD}_2(\mu, \nu) = \frac{\|\delta\|}{\sqrt{d}}
#
# Because this holds regardless of ``Xs``'s shape or size, the *only*
# remaining source of error in the experiment below is the number of
# projections -- exactly the quantity we want to study.

rng = np.random.RandomState(0)

n_samples = 200
delta = np.array([1.5, 1.0, -0.5])
Xs = rng.uniform(-2, 2, (n_samples, d))
Xt = Xs + delta

# Exact reference: no approximation at all, at any cost.
sw_true = np.linalg.norm(delta) / np.sqrt(d)

n_proj_list = [10, 20, 50, 100, 200, 500]
n_trials = 8

errors_uniform = np.zeros((n_trials, len(n_proj_list)))
errors_rqsw = np.zeros((n_trials, len(n_proj_list)))
errors_qsw = np.zeros(len(n_proj_list))

for j, n_proj in enumerate(n_proj_list):
for t in range(n_trials):
sw_uniform = ot.sliced_wasserstein_distance(
Xs, Xt, n_projections=n_proj, sampling_slices="uniform", seed=t
)
sw_rqsw = ot.sliced_wasserstein_distance(
Xs, Xt, n_projections=n_proj, sampling_slices="rqsw", seed=t
)
errors_uniform[t, j] = np.abs(sw_uniform - sw_true)
errors_rqsw[t, j] = np.abs(sw_rqsw - sw_true)

sw_qsw = ot.sliced_wasserstein_distance(
Xs, Xt, n_projections=n_proj, sampling_slices="qsw"
)
errors_qsw[j] = np.abs(sw_qsw - sw_true)

mean_err_uniform = errors_uniform.mean(axis=0)
std_err_uniform = errors_uniform.std(axis=0)
mean_err_rqsw = errors_rqsw.mean(axis=0)
std_err_rqsw = errors_rqsw.std(axis=0)

pl.figure(2, figsize=(6, 5))
pl.plot(n_proj_list, mean_err_uniform, "o-", label="Uniform (MC)")
pl.fill_between(
n_proj_list,
mean_err_uniform - std_err_uniform,
mean_err_uniform + std_err_uniform,
alpha=0.3,
)
pl.plot(n_proj_list, mean_err_rqsw, "s-", label="RQSW")
pl.fill_between(
n_proj_list,
mean_err_rqsw - std_err_rqsw,
mean_err_rqsw + std_err_rqsw,
alpha=0.3,
)
pl.plot(n_proj_list, errors_qsw, "^-", label="QSW (deterministic)")
pl.xscale("log")
pl.yscale("log")
pl.xlabel("Number of projections")
pl.ylabel("Absolute error to the true SWD")
pl.title("Convergence of the Sliced Wasserstein estimate (3D)")
pl.legend()
pl.show()

# QSW and RQSW reach a given accuracy with fewer projections than uniform
# sampling, and RQSW keeps the estimator unbiased -- so it is a drop-in
# replacement for uniform sampling in stochastic optimization settings
# (e.g. Sliced Wasserstein gradient flows) where a deterministic QSW
# estimate would not be appropriate.

# %%
2 changes: 2 additions & 0 deletions ot/sliced/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
get_random_projections,
get_projections_sphere,
projection_sphere_to_circle,
get_projections_spiral,
)
from ._sliced_distances import (
sliced_wasserstein_distance,
Expand All @@ -38,4 +39,5 @@
"sliced_wasserstein_sphere",
"sliced_wasserstein_sphere_unif",
"linear_sliced_wasserstein_sphere",
"get_projections_spiral",
]
49 changes: 43 additions & 6 deletions ot/sliced/_sliced_distances.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@

from ..backend import get_backend
from ..utils import list_to_array, apply_scaler
from ._utils import get_random_projections
from ._utils import get_random_projections, get_projections_spiral
from ..lp import wasserstein_1d


Expand All @@ -26,6 +26,7 @@ def sliced_wasserstein_distance(
seed=None,
log=False,
scaler=None,
sampling_slices="uniform",
):
r"""
Computes a Monte-Carlo approximation of the p-Sliced Wasserstein distance
Expand All @@ -38,6 +39,12 @@ def sliced_wasserstein_distance(

- :math:`\theta_\# \mu` stands for the pushforwards of the projection :math:`X \in \mathbb{R}^d \mapsto \langle \theta, X \rangle`

By default, the projection directions :math:`\theta` are sampled uniformly
at random. Setting ``sampling_slices`` to ``"qsw"`` or ``"rqsw"`` instead
uses Quasi-Monte Carlo point sets on the sphere (generalized spiral
points), which can reduce the approximation error for a given
``n_projections`` [93]. These two options are currently
only implemented for ``dim == 3``.

Parameters
----------
Expand All @@ -54,9 +61,12 @@ def sliced_wasserstein_distance(
p: float, optional
Power p used for computing the sliced Wasserstein
projections: shape (dim, n_projections), optional
Projection matrix (n_projections and seed are not used in this case)
Projection matrix (n_projections, seed and sampling_slices are not
used in this case)
seed: int or RandomState or None, optional
Seed used for random number generator
Seed used for random number generator. Ignored if
``sampling_slices="qsw"`` (the deterministic point set does not
depend on a seed).
log: bool, optional
if True, sliced_wasserstein_distance returns the projections used and their associated EMD.
scaler: None, object with .transform(), or callable, optional
Expand All @@ -73,6 +83,18 @@ def sliced_wasserstein_distance(

See :class:`ot.utils.DataScaler` for a backend-aware scaler that supports
joint fitting on multiple distributions.
sampling_slices: str, optional
Method used to sample the projection directions when ``projections``
is not provided directly. One of:

- ``"uniform"`` (default): directions sampled uniformly at random on
the sphere (Monte Carlo).
- ``"qsw"``: deterministic Quasi-Sliced Wasserstein directions via
generalized spiral points. Only implemented for ``dim == 3``.
- ``"rqsw"``: Randomized Quasi-Sliced Wasserstein -- the same spiral
point set as ``"qsw"``, with a random rotation applied, giving an
unbiased estimator suitable for stochastic optimization. Only
implemented for ``dim == 3``.

Returns
-------
Expand All @@ -94,6 +116,7 @@ def sliced_wasserstein_distance(
----------

.. [31] Bonneel, Nicolas, et al. "Sliced and radon wasserstein barycenters of measures." Journal of Mathematical Imaging and Vision 51.1 (2015): 22-45
.. [93] Nguyen, K., Bariletto, N., & Ho, N. (2024). "Quasi-Monte Carlo for 3D Sliced Wasserstein." International Conference on Learning Representations (ICLR).
"""

X_s, X_t = list_to_array(X_s, X_t)
Expand All @@ -120,9 +143,23 @@ def sliced_wasserstein_distance(
d = X_s.shape[1]

if projections is None:
projections = get_random_projections(
d, n_projections, seed, backend=nx, type_as=X_s
)
if sampling_slices == "uniform":
projections = get_random_projections(
d, n_projections, seed, backend=nx, type_as=X_s
)
elif sampling_slices == "qsw":
projections = get_projections_spiral(
d, n_projections, randomized=False, backend=nx, type_as=X_s
)
elif sampling_slices == "rqsw":
projections = get_projections_spiral(
d, n_projections, randomized=True, seed=seed, backend=nx, type_as=X_s
)
else:
raise ValueError(
f"Unknown sampling_slices method '{sampling_slices}', "
"must be one of 'uniform', 'qsw', 'rqsw'"
)
else:
n_projections = projections.shape[1]

Expand Down
Loading
Loading