Source code for ot.sliced._utils

# -*- coding: utf-8 -*-
"""
Useful functions for solvers for the (balanced) sliced transport problem.
"""

# Author: Adrien Corenflos <adrien.corenflos@aalto.fi>
#         Nicolas Courty   <ncourty@irisa.fr>
#         Rémi Flamary <remi.flamary@polytechnique.edu>
#         Eloi Tanguy <eloi.tanguy@math.cnrs.fr>
#         Laetitia Chapel <laetitia.chapel@irisa.fr>
#         Clément Bonet <clement.bonet.mapp@polytechnique.edu>
#
# License: MIT License

import numpy as np
from ..backend import get_backend, NumpyBackend
from ..utils import get_coordinate_circle


[docs] def get_random_projections(d, n_projections, seed=None, backend=None, type_as=None): r""" Generates n_projections samples from the uniform on the unit sphere of dimension :math:`d-1`: :math:`\mathcal{U}(\mathcal{S}^{d-1})` Parameters ---------- d : int dimension of the space n_projections : int number of samples requested seed: int or RandomState, optional Seed used for numpy random number generator backend: Backend to use for random generation type_as: type, optional Type of the returned array Returns ------- out: ndarray, shape (d, n_projections) The uniform unit vectors on the sphere Examples -------- >>> n_projections = 100 >>> d = 5 >>> projs = get_random_projections(d, n_projections) >>> np.allclose(np.sum(np.square(projs), 0), 1.) # doctest: +NORMALIZE_WHITESPACE True """ if backend is None: nx = NumpyBackend() else: nx = backend if isinstance(seed, np.random.RandomState) and str(nx) == "numpy": projections = seed.randn(d, n_projections) else: if seed is not None: nx.seed(seed) projections = nx.randn(d, n_projections, type_as=type_as) projections = projections / nx.sqrt(nx.sum(projections**2, 0, keepdims=True)) return projections
[docs] def get_projections_sphere(d, n_projections, seed=None, backend=None, type_as=None): r""" Generates n_projections samples from the uniform distribution on the Stiefel manifold of dimension :math:`d\times 2`: :math:`\mathbb{V}_{d,2}=\{X \in \mathbb{R}^{d\times 2}, X^TX=I_2\}` Parameters ---------- d : int dimension of the space n_projections : int number of samples requested seed: int or RandomState, optional Seed used for numpy random number generator backend: Backend to use for random generation type_as: optional Type to use for random generation Returns ------- out: ndarray, shape (n_projections, d, 2) Examples -------- >>> n_projections = 100 >>> d = 5 >>> projs = get_projections_sphere(d, n_projections) >>> np.allclose(np.einsum("nij, nik -> njk", projs, projs), np.eye(2)) # doctest: +NORMALIZE_WHITESPACE True """ if backend is None: nx = NumpyBackend() else: nx = backend if isinstance(seed, np.random.RandomState) and str(nx) == "numpy": Z = seed.randn(n_projections, d, 2) else: if seed is not None: nx.seed(seed) Z = nx.randn(n_projections, d, 2, type_as=type_as) projections, _ = nx.qr(Z) return projections
[docs] def get_random_rotations(d, n_rotations, seed=None, backend=None, type_as=None): r""" Generates n_rotations samples from the uniform (Haar) distribution on the special orthogonal group :math:`\mathrm{SO}(d)=\{R \in \mathbb{R}^{d\times d}, R^TR=I_d, \mathrm{det}(R)=1\}`. The rotations are obtained from the QR factorization of Gaussian matrices, with the sign correction of :ref:`[94] <references-get-random-rotations>` and a sign flip of the first column of the matrices with negative determinant. Parameters ---------- d : int dimension of the space n_rotations : int number of samples requested seed: int or RandomState, optional Seed used for numpy random number generator backend: Backend to use for random generation type_as: optional Type to use for random generation Returns ------- out: ndarray, shape (n_rotations, d, d) Examples -------- >>> n_rotations = 100 >>> d = 5 >>> rotations = get_random_rotations(d, n_rotations) >>> np.allclose(np.einsum("nij, nkj -> nik", rotations, rotations), np.eye(d)) # doctest: +NORMALIZE_WHITESPACE True >>> np.allclose(np.linalg.det(rotations), 1.) # doctest: +NORMALIZE_WHITESPACE True .. _references-get-random-rotations: References ---------- .. [94] Mezzadri, F. (2007). How to generate random matrices from the classical compact groups. Notices of the American Mathematical Society, 54(5), 592-604. """ if backend is None: nx = NumpyBackend() else: nx = backend if isinstance(seed, np.random.RandomState) and str(nx) == "numpy": Z = seed.randn(n_rotations, d, d) else: if seed is not None: nx.seed(seed) Z = nx.randn(n_rotations, d, d, type_as=type_as) Q, R = nx.qr(Z) diagonal = nx.sum(R * nx.eye(d, type_as=R)[None, :, :], axis=-1) Q = Q * nx.sign(diagonal)[:, None, :] flip = nx.sign(nx.det(Q)) rotations = nx.concatenate( (Q[:, :, :1] * flip[:, None, None], Q[:, :, 1:]), axis=-1 ) return rotations
[docs] def projection_sphere_to_circle( x, n_projections=50, projections=None, seed=None, backend=None ): r""" Projection of :math:`x\in S^{d-1}` on circles using coordinates on [0,1[. To get the projection on the circle, we use the following formula: .. math:: P^U(x) = \frac{U^Tx}{\|U^Tx\|_2} where :math:`U` is a random matrix sampled from the uniform distribution on the Stiefel manifold of dimension :math:`d\times 2`: :math:`\mathbb{V}_{d,2}=\{X \in \mathbb{R}^{d\times 2}, X^TX=I_2\}` and :math:`x` is a point on the sphere. Then, we apply the function get_coordinate_circle to get the coordinates on :math:`[0,1[`. Parameters ---------- x : ndarray, shape (n_samples, dim) samples on the sphere n_projections : int, optional Number of projections used for the Monte-Carlo approximation projections: shape (n_projections, dim, 2), optional Projection matrix (n_projections and seed are not used in this case) seed: int or RandomState or None, optional Seed used for random number generator backend: Backend to use for random generation Returns ------- Xp_coords: ndarray, shape (n_projections, n_samples) Coordinates of the projections on the circle """ if backend is None: nx = get_backend(x) else: nx = backend n, d = x.shape if projections is None: projections = get_projections_sphere( d, n_projections, seed=seed, backend=nx, type_as=x ) # Projection on S^1 # Projection on plane Xp = nx.einsum("ikj, lk -> ilj", projections, x) # Projection on sphere Xp = Xp / nx.sqrt(nx.sum(Xp**2, -1, keepdims=True)) # Get coordinates on [0,1[ Xp_coords = nx.reshape( get_coordinate_circle(nx.reshape(Xp, (-1, 2))), (n_projections, n) ) return Xp_coords, projections
[docs] def get_projections_spiral( d, n_projections, randomized=True, seed=None, backend=None, type_as=None ): r""" Generates n_projections points on the sphere via generalized spiral points (Rakhmanov, Saff & Zhou, 1994) [96]. Only implemented for d=3 (the 2-sphere :math:`S^2`). Parameters ---------- d : int dimension of the space. Only d=3 is currently supported. n_projections : int number of samples requested randomized : bool, optional If True (default), applies a random (d, d) rotation to the deterministic spiral point set (RQSW), giving an unbiased estimator suitable for stochastic optimization. If False, returns the plain deterministic point set (QSW). seed: int or RandomState, optional Seed used for the random rotation. Ignored if randomized=False. backend: Backend to use for random generation type_as: type, optional Type of the returned array Returns ------- out: ndarray, shape (d, n_projections) The (optionally rotated) spiral points on the sphere Examples -------- >>> n_projections = 100 >>> d = 3 >>> projs = get_projections_spiral(d, n_projections, randomized=False) >>> np.allclose(np.sum(np.square(projs), 0), 1.) # doctest: +NORMALIZE_WHITESPACE True >>> rprojs = get_projections_spiral(d, n_projections, randomized=True, seed=0) >>> np.allclose(np.sum(np.square(rprojs), 0), 1.) # doctest: +NORMALIZE_WHITESPACE True References ---------- .. [95] Nguyen, K., Bariletto, N., & Ho, N. (2024). "Quasi-Monte Carlo for 3D Sliced Wasserstein." International Conference on Learning Representations (ICLR). .. [96] Rakhmanov, E. A., Saff, E. B., & Zhou, Y. M. (1994). "Minimal Discrete Energy on the Sphere." Mathematical Research Letters, 1(6), 647-662. """ if d != 3: raise ValueError( f"Generalized spiral points are only defined for d=3, got d={d}." ) if backend is None: nx = NumpyBackend() else: nx = backend i = np.arange(1, n_projections + 1) z = 1 - (2 * i - 1) / n_projections phi1 = np.arccos(z) phi2 = (1.8 * np.sqrt(n_projections) * phi1) % (2 * np.pi) theta_np = np.stack( [np.sin(phi1) * np.cos(phi2), np.sin(phi1) * np.sin(phi2), z], axis=0 ) # shape (d, n_projections) theta = nx.from_numpy(theta_np, type_as=type_as) if not randomized: return theta if isinstance(seed, np.random.RandomState) and str(nx) == "numpy": Z = seed.randn(d, d) else: if seed is not None: nx.seed(seed) Z = nx.randn(d, d, type_as=type_as) Q, R = nx.qr(Z) Q = Q * nx.sign(nx.diag(R))[None, :] return nx.matmul(Q, theta)
[docs] def projection_sphere_to_ball(x, eps=1e-6, backend=None): r""" Projection of :math:`x\in S^{d-1}` on the unit ball of :math:`\mathbb{R}^{d-1}` with :math:`\frac{1}{\pi}h_1\circ\phi_\epsilon`. To get the projection on the ball, we use the following closed form: .. math:: \frac{1}{\pi}(h_1\circ\phi_\epsilon)(x) = \frac{1}{\pi}\mathrm{arccos}\left(-\frac{3+5x_d}{5+3x_d}\right)\frac{x_{1:d-1}}{\|x_{1:d-1}\|} where :math:`\phi_\epsilon` is the stereographic projection :math:`\phi(x) = \frac{2 x_{1:d-1}}{1-x_d}` restricted to the sphere without the :math:`\epsilon`-cap around the north pole (points with :math:`x_d > 1-\epsilon` are first mapped to the circle :math:`x_d = 1-\epsilon`), and :math:`h_1(x) = \mathrm{arccos}\left(\frac{1-\|x\|^2}{1+\|x\|^2}\right)\frac{x}{\|x\|}` is the injective defining function of :ref:`[93] <references-s3w>`. Parameters ---------- x : ndarray, shape (..., dim) samples on the sphere eps: float, optional (default=1e-6) Size of the cap around the north pole excluded from the stereographic projection to ensure numerical stability backend: Backend to use for the computations Returns ------- Xp: ndarray, shape (..., dim-1) Images of the samples in the unit ball """ if backend is None: nx = get_backend(x) else: nx = backend x_d = nx.clip(x[..., -1:], -1.0, 1.0 - eps) x_azimuth = x[..., :-1] norm2 = nx.sum(x_azimuth**2, axis=-1, keepdims=True) # the azimuth of the poles is arbitrary, fix it for reproducibility x_azimuth = nx.where( norm2 > 0, x_azimuth, nx.ones(x_azimuth.shape, type_as=x_azimuth) ) norm2 = nx.sum(x_azimuth**2, axis=-1, keepdims=True) radius = nx.arccos(-(3.0 + 5.0 * x_d) / (5.0 + 3.0 * x_d)) / np.pi return radius * x_azimuth / nx.sqrt(norm2)
[docs] def get_random_orthogonal_directions( d, n_projections, seed=None, backend=None, type_as=None ): r""" Generates n_projections directions on the sphere via UnifOrtho (Rowland et al., 2019) [97]: independent blocks of d mutually orthogonal, unit-norm directions, each block drawn from the Haar measure on the special orthogonal group :math:`\mathrm{SO}(d)`. Unlike the generalized spiral points (:any:`get_projections_spiral`), which only cover :math:`d=3`, UnifOrtho is defined for any dimension d, and is recommended specifically for large d: Petrovic, Bardenet & Desolneux (2026) [98] show both empirically and theoretically that it reduces the variance of the Sliced Wasserstein estimator in high dimension, while quasi-Monte Carlo methods such as the spiral points remain preferable in low dimension (:math:`d \in \{2, 3\}`). Parameters ---------- d : int dimension of the space. Any d >= 1 is supported. n_projections : int number of samples requested seed: int or RandomState, optional Seed used for the underlying random rotations backend: Backend to use for random generation type_as: type, optional Type of the returned array Returns ------- out: ndarray, shape (d, n_projections) The UnifOrtho directions on the sphere Notes ----- ``n_projections`` need not be a multiple of ``d``: internally, ``ceil(n_projections / d)`` independent orthogonal blocks are drawn and concatenated, then truncated to the requested length. Every individual direction is still marginally uniform on the sphere, but if the last block is truncated, the directions coming from that specific block are no longer guaranteed to be mutually orthogonal with each other. Examples -------- >>> n_projections = 100 >>> d = 5 >>> projs = get_random_orthogonal_directions(d, n_projections, seed=0) >>> np.allclose(np.sum(np.square(projs), 0), 1.) # doctest: +NORMALIZE_WHITESPACE True >>> first_block = projs[:, :d] >>> np.allclose(first_block.T @ first_block, np.eye(d)) # doctest: +NORMALIZE_WHITESPACE True References ---------- .. [97] Rowland, M., Hron, J., Tang, Y., Choromanski, K., Sarlos, T., & Weller, A. (2019). "Orthogonal Estimation of Wasserstein Distances." Proceedings of the 22nd International Conference on Artificial Intelligence and Statistics (AISTATS), PMLR 89. .. [98] Petrovic, V., Bardenet, R., & Desolneux, A. (2026). "Repulsive Monte Carlo on the sphere for the sliced Wasserstein distance." Transactions on Machine Learning Research. """ if backend is None: nx = NumpyBackend() else: nx = backend # Number of independent orthogonal blocks needed so that k * d covers # at least n_projections directions (integer ceiling of n_projections / d). k = -(-n_projections // d) # k independent Haar-uniform (d, d) orthogonal matrices. Within each # (d, d) slice, the ROWS form d mutually orthogonal, unit-norm # directions (see get_random_rotations for the QR-based construction # and its Mezzadri (2007) sign correction, which is what makes the # rotation genuinely Haar-uniform rather than merely orthogonal). rotations = get_random_rotations(d, k, seed=seed, backend=nx, type_as=type_as) # Flatten the k blocks into a single list of k * d directions, one per row. directions = nx.reshape(rotations, (k * d, d)) # (d, k * d): match the (dim, n_projections) convention. directions = nx.transpose(directions) # Trim down to exactly n_projections when n_projections is not a # multiple of d (see Notes above for what this means for the last, # possibly incomplete, orthogonal block). return directions[:, :n_projections]