diff --git a/RELEASES.md b/RELEASES.md index 860f2babb..2e2112375 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -4,6 +4,9 @@ #### New features + +- `ot.dr.wda` gains `solver='torch'`, a PyTorch autodiff solver with Riemannian gradient descent, usable on installations without autograd or pymanopt, and accepting torch tensors directly (PR #858, Issue #806) +- `ot.dr` dependencies (autograd, pymanopt, scikit-learn, torch) are now imported optionally, so importing `ot.dr` no longer requires all of them; each function raises an explicit `ImportError` naming what it needs (PR #858) - Use `ot.utils.check_marginal` (and shape-tuple support in `ot.utils.unif`) to fill and validate default marginals consistently across solvers (Gromov, low-rank, stochastic, barycenter, factored) (PR #856) - Add stereographic spherical sliced Wasserstein distance in `ot.sliced.stereographic_sliced_wasserstein_sphere`, with its rotationally invariant extension (PR #836) - Add Quasi-Monte Carlo sliced Wasserstein sampling (QSW/RQSW) via generalized diff --git a/ot/dr.py b/ot/dr.py index 9cfabff5f..8336ccfeb 100644 --- a/ot/dr.py +++ b/ot/dr.py @@ -18,22 +18,55 @@ from scipy import linalg +# ot.dr offers solvers with different dependencies. Each is imported optionally +# so that, for instance, the PyTorch WDA solver works on an installation with no +# autograd or pymanopt. Functions raise an ImportError naming what they need. try: import autograd.numpy as np - from sklearn.decomposition import PCA + HAS_AUTOGRAD = True +except ImportError: # pragma: no cover - depends on the installation + import numpy as np + + HAS_AUTOGRAD = False + +try: import pymanopt import pymanopt.manifolds import pymanopt.optimizers -except ImportError: - raise ImportError( - "Missing dependency for ot.dr. Requires autograd, pymanopt, scikit-learn. You can install with install with 'pip install POT[dr]', or 'conda install autograd pymanopt scikit-learn'" - ) + HAS_PYMANOPT = True +except ImportError: # pragma: no cover - depends on the installation + HAS_PYMANOPT = False + +try: + import torch + + HAS_TORCH = True +except ImportError: # pragma: no cover - depends on the installation + HAS_TORCH = False + +try: + from sklearn.decomposition import PCA + + HAS_SKLEARN = True +except ImportError: # pragma: no cover - depends on the installation + HAS_SKLEARN = False + +from .backend import get_backend from .bregman import sinkhorn as sinkhorn_bregman from .utils import dist as dist_utils, check_random_state +def _require(condition, function, dependencies): + if not condition: + raise ImportError( + f"Missing dependency for ot.dr.{function}. Requires {dependencies}. " + "You can install with 'pip install POT[dr]', or " + "'conda install autograd pymanopt scikit-learn'" + ) + + def dist(x1, x2): r"""Compute squared euclidean distance between samples (autograd)""" x1p2 = np.sum(np.square(x1), 1) @@ -79,6 +112,340 @@ def split_classes(X, y): return [X[y == i, :].astype(np.float32) for i in lstsclass] +def _stiefel_projection(P, G, nx): + r"""Project a euclidean gradient onto the tangent space of the Stiefel manifold. + + Parameters + ---------- + P : array-like, shape (d, p) + Current point, with orthonormal columns. + G : array-like, shape (d, p) + Euclidean gradient at :math:`\mathbf{P}`. + nx : backend + Backend to use for the computation. + + Returns + ------- + G_tangent : array-like, shape (d, p) + Component of :math:`\mathbf{G}` tangent to Stiefel at :math:`\mathbf{P}`, + that is :math:`\mathbf{G} - \mathbf{P}\,\mathrm{sym}(\mathbf{P}^\top \mathbf{G})`. + """ + W = nx.dot(P.T, G) + return G - nx.dot(P, 0.5 * (W + W.T)) + + +def _stiefel_retraction(P, X, nx): + r"""Retract :math:`\mathbf{P} + \mathbf{X}` onto the Stiefel manifold by QR. + + The sign of each column is fixed so that the diagonal of :math:`\mathbf{R}` + is non-negative, which makes the retraction a continuous map. + + Parameters + ---------- + P : array-like, shape (d, p) + Current point, with orthonormal columns. + X : array-like, shape (d, p) + Tangent step to apply. + nx : backend + Backend to use for the computation. + + Returns + ------- + P_next : array-like, shape (d, p) + Point on Stiefel with orthonormal columns. + """ + Q, R = nx.qr(P + X) + d = nx.diag(R) + return Q * nx.sign(nx.sign(d) + 0.5) + + +def _wda_cost(P, xc, wc, regmean, reg, k, sinkhorn_method, nx): + r"""Wasserstein Discriminant Analysis objective at a projection :math:`\mathbf{P}`. + + Computes the ratio of the within-class to the between-class entropic + transport cost of the projected samples, which is the quantity `wda` + minimizes. The Sinkhorn solver is run for exactly `k` iterations rather + than to a tolerance, so the objective is a fixed-depth, differentiable + function of :math:`\mathbf{P}`. + + Parameters + ---------- + P : array-like, shape (d, p) + Projection, with orthonormal columns. + xc : list of array-like + Samples split by class, each of shape (n_i, d). + wc : list of array-like + Uniform weights for each class, each of shape (n_i,). + regmean : array-like, shape (n_classes, n_classes) + Per-pair scaling of `reg`, all ones unless `normalize` was requested. + reg : float + Entropic regularization term > 0. + k : int + Number of Sinkhorn iterations. + sinkhorn_method : str + Either 'sinkhorn' or 'sinkhorn_log', passed to :py:func:`ot.sinkhorn`. + nx : backend + Backend to use for the computation. + + Returns + ------- + loss : float or array-like scalar + Within-class cost divided by between-class cost. + """ + loss_b, loss_w = 0.0, 0.0 + for i, xi in enumerate(xc): + xi = nx.dot(xi, P) + for j, xj in enumerate(xc[i:]): + xj = nx.dot(xj, P) + M = dist_utils(xi, xj) + G = sinkhorn_bregman( + wc[i], + wc[j + i], + M, + reg * regmean[i, j], + method=sinkhorn_method, + numItermax=k, + stopThr=0.0, + warn=False, + ) + term = nx.sum(G * M) + if j == 0: + loss_w = loss_w + term + else: + loss_b = loss_b + term + if float(nx.to_numpy(loss_b)) == 0.0: + raise ValueError( + "The between-class transport cost underflowed to zero, so the WDA " + "objective is undefined. reg is too small for the scale of the " + "data: exp(-M/reg) underflows. Increase reg, or use " + "sinkhorn_method='sinkhorn_log'." + ) + return loss_w / loss_b + + +def _wda_torch( + X, + y, + p, + reg, + k, + sinkhorn_method, + maxiter, + verbose, + P0, + normalize, + random_state, + step_growth, + step_shrink, + max_backtracks, + gtol, +): + r"""Solve WDA with PyTorch autodiff and Riemannian gradient descent. + + Follows the same scheme as pymanopt's ``SteepestDescent`` with its default + ``BackTrackingLineSearcher``: project the euclidean gradient onto the + tangent space, retract by QR, and backtrack from a step proportional to the + last accepted one. Both solvers therefore target the same optimum. + + Parameters + ---------- + X : torch.Tensor, shape (n, d) + Mean-centred training samples. + y : torch.Tensor, shape (n,) + Integer class codes. + p : int + Size of the dimensionality reduction. + reg : float + Entropic regularization term > 0. + k : int + Number of Sinkhorn iterations per objective evaluation. + sinkhorn_method : str + Either 'sinkhorn' or 'sinkhorn_log'. + maxiter : int + Maximum number of descent iterations. + verbose : int + Print the objective and gradient norm at each iteration if non-zero. + P0 : torch.Tensor, shape (d, p) or None + Starting point. Drawn at random when None. + normalize : bool + Scale `reg` per class pair by the mean projected distance at `P0`. + random_state : int, RandomState or None + Seeds the starting point when `P0` is None. + step_growth : float + Factor applied to the last accepted step to start the line search. + step_shrink : float + Factor applied on each backtracking step. + max_backtracks : int + Maximum number of backtracking steps before the solver stops. + gtol : float + Stop once the tangent gradient norm falls below this value. + + Returns + ------- + P : torch.Tensor, shape (d, p) + Optimal projection. + """ + nx = get_backend(X) + dtype, device = X.dtype, X.device + labels = torch.unique(y) + xc = [X[y == c] for c in labels] + wc = [ + torch.full((x.shape[0],), 1.0 / x.shape[0], dtype=dtype, device=device) + for x in xc + ] + d = X.shape[1] + nc = len(xc) + + if not 1 <= p <= d: + raise ValueError(f"Need d >= p >= 1. Values supplied were d = {d} and p = {p}") + + if P0 is None: + rng = check_random_state(random_state) + P = torch.linalg.qr(torch.tensor(rng.randn(d, p), dtype=dtype, device=device))[ + 0 + ] + else: + P = P0.clone().to(dtype=dtype, device=device) + + regmean = torch.ones((nc, nc), dtype=dtype, device=device) + if P0 is not None and normalize: + with torch.no_grad(): + for i, xi in enumerate(xc): + xi = xi @ P + for j, xj in enumerate(xc[i:]): + xj = xj @ P + regmean[i, j] = torch.sum(dist_utils(xi, xj)) / ( + xi.shape[0] * xj.shape[0] + ) + + def value(Q): + with torch.no_grad(): + return _wda_cost(Q, xc, wc, regmean, reg, k, sinkhorn_method, nx) + + f = value(P) + step = 1.0 + for it in range(maxiter): + Q = P.detach().requires_grad_(True) + v = _wda_cost(Q, xc, wc, regmean, reg, k, sinkhorn_method, nx) + (g,) = torch.autograd.grad(v, Q) + direction = -_stiefel_projection(P, g, nx) + gnorm = float(nx.norm(direction)) + if verbose: + print(f"{it + 1:<6d} {float(v.detach()):+.16e} {gnorm:.8e}") + if gnorm <= gtol: + break + step = step_growth * step + improved = False + for _ in range(max_backtracks): + Pn = _stiefel_retraction(P, step * direction, nx) + fn = value(Pn) + if fn < f: + improved = True + break + step = step_shrink * step + if not improved: + break + P, f = Pn, fn + return P.detach() + + +def _wda_torch_entry( + X, + y, + p, + reg, + k, + sinkhorn_method, + maxiter, + verbose, + P0, + normalize, + random_state, + step_growth, + step_shrink, + max_backtracks, + gtol, +): + r"""Prepare inputs for the torch solver and wrap its output. + + A torch tensor keeps its dtype and device throughout. Any other array type + is promoted to float64 on the CPU, matching what the autograd solver + returns, and the projection comes back as a numpy array. + + Parameters + ---------- + X : array-like, shape (n, d) + Training samples, not modified. + y : array-like, shape (n,) + Labels, of any type numpy can take the unique values of. + + Other parameters are those of :py:func:`ot.dr.wda`. + + Returns + ------- + P : array-like, shape (d, p) + Optimal projection, of the same array type as `X`. + proj : callable + Projection function including mean centering. + """ + was_numpy = not torch.is_tensor(X) + if was_numpy: + Xt = torch.as_tensor(np.asarray(X, dtype=np.float64)) + else: + Xt = X if torch.is_floating_point(X) else X.to(torch.float64) + + # torch cannot hold arbitrary label types, so index classes by position the + # way split_classes does for numpy + if torch.is_tensor(y): + yt = y.to(Xt.device) + else: + yt = torch.as_tensor( + np.unique(np.asarray(y), return_inverse=True)[1], device=Xt.device + ) + + if P0 is None: + P0t = None + else: + P0t = (P0 if torch.is_tensor(P0) else torch.as_tensor(np.asarray(P0))).to( + dtype=Xt.dtype, device=Xt.device + ) + + mx = Xt.mean(dim=0) + Xc = Xt - mx.reshape((1, -1)) + + Popt = _wda_torch( + Xc, + yt, + p, + reg, + k, + sinkhorn_method, + maxiter, + verbose, + P0t, + normalize, + random_state, + step_growth, + step_shrink, + max_backtracks, + gtol, + ) + + if was_numpy: + Pn = Popt.detach().cpu().numpy() + mxn = mx.detach().cpu().numpy() + + def proj(Z): + return (Z - mxn.reshape((1, -1))).dot(Pn) + + return Pn, proj + + def proj(Z): + return (Z - mx.reshape((1, -1))) @ Popt + + return Popt, proj + + def fda(X, y, p=2, reg=1e-16): r"""Fisher Discriminant Analysis @@ -138,78 +505,33 @@ def proj(X): return Popt, proj -def wda( - X, - y, - p=2, - reg=1, - k=10, - solver=None, - sinkhorn_method="sinkhorn", - maxiter=100, - verbose=0, - P0=None, - normalize=False, +def _wda_autograd( + X, y, p, reg, k, solver, sinkhorn_method, maxiter, verbose, P0, normalize ): - r""" - Wasserstein Discriminant Analysis :ref:`[11] ` - - The function solves the following optimization problem: - - .. math:: - \mathbf{P} = \mathop{\arg \min}_\mathbf{P} \quad - \frac{\sum\limits_i W(P \mathbf{X}^i, P \mathbf{X}^i)}{\sum\limits_{i, j \neq i} W(P \mathbf{X}^i, P \mathbf{X}^j)} - - where : - - - :math:`P` is a linear projection operator in the Stiefel(`p`, `d`) manifold - - :math:`W` is entropic regularized Wasserstein distances - - :math:`\mathbf{X}^i` are samples in the dataset corresponding to class i - - **Choosing a Sinkhorn solver** - - By default and when using a regularization parameter that is not too small - the default sinkhorn solver should be enough. If you need to use a small - regularization to get sparse cost matrices, you should use the - :py:func:`ot.dr.sinkhorn_log` solver that will avoid numerical - errors, but can be slow in practice. + r"""Solve WDA with autograd and a pymanopt optimizer. Parameters ---------- X : ndarray, shape (n, d) - Training samples. + Training samples, not modified. y : ndarray, shape (n,) Labels for training samples. - p : int, optional - Size of dimensionality reduction. - reg : float, optional - Regularization term >0 (entropic regularization) - solver : None | str, optional - None for steepest descent or 'TrustRegions' for trust regions algorithm - else should be a pymanopt.solvers - sinkhorn_method : str - method used for the Sinkhorn solver, either 'sinkhorn' or 'sinkhorn_log' - P0 : ndarray, shape (d, p) - Initial starting point for projection. - normalize : bool, optional - Normalize the Wasserstaiun distance by the average distance on P0 (default : False) - verbose : int, optional - Print information along iterations. + solver : None | str | pymanopt.optimizers + ``None`` or ``'autograd'`` selects ``SteepestDescent``, ``'tr'`` or + ``'TrustRegions'`` selects ``TrustRegions``, and any other value is + used as a pymanopt optimizer instance. + + Other parameters are those of :py:func:`ot.dr.wda`. Returns ------- P : ndarray, shape (d, p) - Optimal transportation matrix for the given parameters + Optimal projection. proj : callable Projection function including mean centering. - - - .. _references-wda: - References - ---------- - .. [11] Flamary, R., Cuturi, M., Courty, N., & Rakotomamonjy, A. (2016). - Wasserstein Discriminant Analysis. arXiv preprint arXiv:1608.08063. - """ # noqa + """ + if solver == "autograd": + solver = None if sinkhorn_method.lower() == "sinkhorn": sinkhorn_solver = sinkhorn @@ -283,6 +605,148 @@ def proj(X): return Popt.point, proj +def wda( + X, + y, + p=2, + reg=1, + k=10, + solver=None, + sinkhorn_method="sinkhorn", + maxiter=100, + verbose=0, + P0=None, + normalize=False, + random_state=None, + step_growth=2.0, + step_shrink=0.5, + max_backtracks=40, + gtol=1e-12, +): + r""" + Wasserstein Discriminant Analysis :ref:`[11] ` + + The function solves the following optimization problem: + + .. math:: + \mathbf{P} = \mathop{\arg \min}_\mathbf{P} \quad + \frac{\sum\limits_i W(P \mathbf{X}^i, P \mathbf{X}^i)}{\sum\limits_{i, j \neq i} W(P \mathbf{X}^i, P \mathbf{X}^j)} + + where : + + - :math:`P` is a linear projection operator in the Stiefel(`p`, `d`) manifold + - :math:`W` is entropic regularized Wasserstein distances + - :math:`\mathbf{X}^i` are samples in the dataset corresponding to class i + + **Choosing a Sinkhorn solver** + + By default and when using a regularization parameter that is not too small + the default sinkhorn solver should be enough. If you need to use a small + regularization to get sparse cost matrices, you should use the + :py:func:`ot.dr.sinkhorn_log` solver that will avoid numerical + errors, but can be slow in practice. + + Parameters + ---------- + X : ndarray, shape (n, d) + Training samples. + y : ndarray, shape (n,) + Labels for training samples. + p : int, optional + Size of dimensionality reduction. + reg : float, optional + Regularization term >0 (entropic regularization) + solver : None | str | pymanopt.optimizers, optional + Chooses both the autodiff framework and the optimizer. + + - `None` or `'autograd'` (default): autograd and pymanopt + `SteepestDescent`. + - `'TrustRegions'` (or `'tr'`): autograd and pymanopt `TrustRegions`. + - a `pymanopt.optimizers` instance: autograd with that optimizer. + - `'torch'`: PyTorch autodiff with Riemannian gradient descent and a QR + retraction. Requires only `torch`, so it works on installations + without autograd or pymanopt, and accepts torch tensors directly, + keeping their device and dtype. + sinkhorn_method : str + method used for the Sinkhorn solver, either 'sinkhorn' or 'sinkhorn_log' + P0 : ndarray, shape (d, p) + Initial starting point for projection. + normalize : bool, optional + Normalize the Wasserstaiun distance by the average distance on P0 (default : False) + verbose : int, optional + Print information along iterations. + random_state : int, RandomState instance or None, optional + Seeds the random starting point when `P0` is None. Only used by + `solver='torch'`; the pymanopt solvers draw their own starting point. + step_growth : float, optional + Factor applied to the last accepted step size to start the next line + search (default 2.0). Only used by `solver='torch'`. + step_shrink : float, optional + Factor applied to the step size on each backtracking step + (default 0.5). Only used by `solver='torch'`. + max_backtracks : int, optional + Maximum number of backtracking steps per iteration before the solver + stops (default 40). Only used by `solver='torch'`. + gtol : float, optional + Stop once the norm of the tangent gradient falls below this value + (default 1e-12). Only used by `solver='torch'`. + + Returns + ------- + P : ndarray, shape (d, p) + Optimal transportation matrix for the given parameters + proj : callable + Projection function including mean centering. + + + .. _references-wda: + References + ---------- + .. [11] Flamary, R., Cuturi, M., Courty, N., & Rakotomamonjy, A. (2016). + Wasserstein Discriminant Analysis. arXiv preprint arXiv:1608.08063. + """ # noqa + + if HAS_TORCH and torch.is_tensor(y): + n_classes = int(torch.unique(y).numel()) + else: + n_classes = np.unique(np.asarray(y)).size + if n_classes < 2: + raise ValueError( + f"WDA needs at least two classes, got {n_classes}: with a single " + "class the between-class transport cost is zero and the objective " + "is undefined." + ) + + if solver == "torch": + _require(HAS_TORCH, "wda(solver='torch')", "torch") + return _wda_torch_entry( + X, + y, + p, + reg, + k, + sinkhorn_method, + maxiter, + verbose, + P0, + normalize, + random_state, + step_growth, + step_shrink, + max_backtracks, + gtol, + ) + + _require( + HAS_AUTOGRAD and HAS_PYMANOPT, + "wda(solver='autograd')", + "autograd and pymanopt", + ) + return _wda_autograd( + X, y, p, reg, k, solver, sinkhorn_method, maxiter, verbose, P0, normalize + ) + + def projection_robust_wasserstein( X, Y, @@ -495,6 +959,7 @@ def ewca( X = X - X.mean(0) if U0 is None: + _require(HAS_SKLEARN, "ewca", "scikit-learn") pca_fitted = PCA(n_components=k).fit(X) U = pca_fitted.components_.T if method == "MM": diff --git a/test/test_dr.py b/test/test_dr.py index dcb477717..501e14781 100644 --- a/test/test_dr.py +++ b/test/test_dr.py @@ -13,10 +13,17 @@ try: # test if autograd and pymanopt are installed import ot.dr - nogo = False + nogo = not (ot.dr.HAS_AUTOGRAD and ot.dr.HAS_PYMANOPT) except ImportError: nogo = True +try: + import torch + + notorch = False +except ImportError: + notorch = True + @pytest.mark.skipif(nogo, reason="Missing modules (autograd or pymanopt)") def test_fda(): @@ -251,3 +258,285 @@ def test_ewca(): U_last_eigvec = np.linalg.svd(X.T, full_matrices=False)[0][:, -k:] _, cos, _ = np.linalg.svd(U.T @ U_last_eigvec, full_matrices=False) assert np.allclose(cos, np.ones(k), atol=1e-3) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_solver(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 90, random_state=rng) + xs = np.hstack((xs, rng.randn(90, 4))) + p = 2 + + P, proj = ot.dr.wda(xs, ys, p, maxiter=10, solver="torch") + + np.testing.assert_allclose(np.sum(P**2, 0), np.ones(p), rtol=1e-6) + assert proj(xs).shape == (90, p) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_accepts_torch_tensors(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 90, random_state=rng) + xt = torch.tensor(xs, dtype=torch.float64) + yt = torch.tensor(ys) + + P, proj = ot.dr.wda(xt, yt, 2, maxiter=5, solver="torch") + + assert torch.is_tensor(P) + assert P.dtype == torch.float64 + assert proj(xt).shape == (90, 2) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_does_not_modify_input(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 90, random_state=rng) + xs = xs + 10.0 + xs_copy = xs.copy() + + ot.dr.wda(xs, ys, 2, maxiter=5, solver="torch") + + np.testing.assert_allclose(xs, xs_copy) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_sinkhorn_log(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 90, random_state=rng) + p = 2 + + P, _ = ot.dr.wda( + xs, ys, p, maxiter=10, solver="torch", sinkhorn_method="sinkhorn_log" + ) + + np.testing.assert_allclose(np.sum(P**2, 0), np.ones(p), rtol=1e-6) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_cost_agrees_across_backends(): + """_wda_cost goes through the POT backend, so numpy and torch must agree.""" + rng = np.random.RandomState(0) + n, d, C, reg, k = 180, 6, 3, 1.0, 10 + X = np.vstack([rng.randn(n // C, d) + 3 * rng.randn(1, d) for _ in range(C)]) + X = X - X.mean(0) + y = np.repeat(np.arange(C), n // C) + P0 = np.linalg.qr(rng.randn(d, 2))[0] + + xc = [np.ascontiguousarray(X[y == c]) for c in range(C)] + wc = [np.ones(x.shape[0]) / x.shape[0] for x in xc] + rm = np.ones((C, C)) + v_np = ot.dr._wda_cost( + P0, xc, wc, rm, reg, k, "sinkhorn", ot.backend.NumpyBackend() + ) + + xct = [torch.tensor(x, dtype=torch.float64) for x in xc] + wct = [torch.tensor(w, dtype=torch.float64) for w in wc] + rmt = torch.ones((C, C), dtype=torch.float64) + Pt = torch.tensor(P0, dtype=torch.float64) + nxt = ot.backend.TorchBackend() + v_t = ot.dr._wda_cost(Pt, xct, wct, rmt, reg, k, "sinkhorn", nxt) + + np.testing.assert_allclose(float(v_t), float(v_np), rtol=1e-10) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_cost_gradient_matches_finite_differences(): + """The autodiff gradient of _wda_cost must match a central difference.""" + rng = np.random.RandomState(0) + n, d, C, reg, k = 120, 4, 2, 1.0, 10 + X = np.vstack([rng.randn(n // C, d) + 3 * rng.randn(1, d) for _ in range(C)]) + X = X - X.mean(0) + y = np.repeat(np.arange(C), n // C) + P0 = np.linalg.qr(rng.randn(d, 2))[0] + + xc = [torch.tensor(X[y == c], dtype=torch.float64) for c in range(C)] + wc = [torch.full((x.shape[0],), 1.0 / x.shape[0], dtype=torch.float64) for x in xc] + rm = torch.ones((C, C), dtype=torch.float64) + nx = ot.backend.TorchBackend() + + def f(P): + return ot.dr._wda_cost(P, xc, wc, rm, reg, k, "sinkhorn", nx) + + P = torch.tensor(P0, dtype=torch.float64, requires_grad=True) + (grad,) = torch.autograd.grad(f(P), P) + + eps = 1e-6 + fd = np.zeros_like(P0) + for i in range(P0.shape[0]): + for j in range(P0.shape[1]): + Pp, Pm = P0.copy(), P0.copy() + Pp[i, j] += eps + Pm[i, j] -= eps + with torch.no_grad(): + vp = float(f(torch.tensor(Pp, dtype=torch.float64))) + vm = float(f(torch.tensor(Pm, dtype=torch.float64))) + fd[i, j] = (vp - vm) / (2 * eps) + + np.testing.assert_allclose(grad.numpy(), fd, rtol=1e-4, atol=1e-8) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_stiefel_projection_and_retraction(): + """The projection is tangent to Stiefel and the retraction stays on it.""" + rng = np.random.RandomState(0) + d, p = 6, 2 + nx = ot.backend.NumpyBackend() + P = np.linalg.qr(rng.randn(d, p))[0] + G = rng.randn(d, p) + + T = ot.dr._stiefel_projection(P, G, nx) + # tangency: P^T T must be skew-symmetric + W = P.T @ T + np.testing.assert_allclose(W + W.T, np.zeros((p, p)), atol=1e-10) + + Pn = ot.dr._stiefel_retraction(P, 0.1 * T, nx) + np.testing.assert_allclose(Pn.T @ Pn, np.eye(p), atol=1e-10) + # a zero step must return the same point + np.testing.assert_allclose( + ot.dr._stiefel_retraction(P, np.zeros_like(P), nx), P, atol=1e-10 + ) + + +@pytest.mark.skipif(nogo or notorch, reason="Missing modules") +def test_wda_solvers_reach_comparable_objective(): + """Both solvers minimise the same objective, so neither should be much worse.""" + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 120, random_state=rng) + xs = np.hstack((xs, rng.randn(120, 3))) + P0 = np.linalg.qr(rng.randn(xs.shape[1], 2))[0] + + Pa, _ = ot.dr.wda(xs, ys, 2, maxiter=40, P0=P0) + Pt, _ = ot.dr.wda(xs, ys, 2, maxiter=40, P0=P0, solver="torch") + + Xc = torch.tensor(xs - xs.mean(0), dtype=torch.float64) + yt = torch.tensor(ys) + xc = [Xc[yt == c] for c in torch.unique(yt)] + wc = [torch.full((x.shape[0],), 1.0 / x.shape[0], dtype=torch.float64) for x in xc] + rm = torch.ones((len(xc), len(xc)), dtype=torch.float64) + nx = ot.backend.TorchBackend() + + def objective(P): + with torch.no_grad(): + return float( + ot.dr._wda_cost( + torch.tensor(P, dtype=torch.float64), + xc, + wc, + rm, + 1, + 10, + "sinkhorn", + nx, + ) + ) + + assert objective(Pt) < 1.15 * objective(Pa) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_unknown_sinkhorn_method(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 60, random_state=rng) + + with pytest.raises(ValueError): + ot.dr.wda(xs, ys, 2, maxiter=2, solver="torch", sinkhorn_method="nope") + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_non_numeric_labels(): + """Labels need not be numeric: numpy indexes by value, torch cannot.""" + rng = np.random.RandomState(0) + xs = np.vstack([rng.randn(40, 5) + 3 * rng.randn(1, 5) for _ in range(2)]) + ys = np.array(["cat"] * 40 + ["dog"] * 40) + + P, _ = ot.dr.wda(xs, ys, 2, maxiter=3, solver="torch") + + assert P.shape == (5, 2) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_rejects_p_larger_than_d(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 60, random_state=rng) + + with pytest.raises(ValueError): + ot.dr.wda(xs, ys, xs.shape[1] + 1, maxiter=3, solver="torch") + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_numpy_input_returns_float64(): + """numpy in gives float64 out, matching the autograd solver.""" + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 60, random_state=rng) + + P, _ = ot.dr.wda(xs.astype(np.float32), ys, 2, maxiter=3, solver="torch") + + assert P.dtype == np.float64 + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_random_state_is_reproducible(): + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 90, random_state=rng) + + kw = dict(p=2, maxiter=5, solver="torch") + P_a, _ = ot.dr.wda(xs, ys, random_state=0, **kw) + P_b, _ = ot.dr.wda(xs, ys, random_state=0, **kw) + P_c, _ = ot.dr.wda(xs, ys, random_state=1, **kw) + + np.testing.assert_allclose(P_a, P_b) + assert not np.allclose(P_a, P_c) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_line_search_parameters(): + """The line-search controls are plumbed through and bound the work done.""" + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 90, random_state=rng) + P0 = np.linalg.qr(rng.randn(xs.shape[1], 2))[0] + + # a single backtracking step with no growth still returns a valid point + P, _ = ot.dr.wda( + xs, + ys, + 2, + maxiter=5, + P0=P0, + solver="torch", + step_growth=1.0, + step_shrink=0.1, + max_backtracks=1, + ) + np.testing.assert_allclose(np.sum(P**2, 0), np.ones(2), rtol=1e-6) + + # a gtol above the initial gradient norm stops immediately at P0 + P_stop, _ = ot.dr.wda(xs, ys, 2, maxiter=50, P0=P0, solver="torch", gtol=1e9) + np.testing.assert_allclose(P_stop, P0, atol=1e-10) + + +@pytest.mark.skipif(notorch, reason="Missing module (torch)") +def test_wda_torch_preserves_device(): + P0 = None + rng = np.random.RandomState(0) + xs, ys = ot.datasets.make_data_classif("gaussrot", 60, random_state=rng) + xt = torch.tensor(xs, dtype=torch.float32) + yt = torch.tensor(ys) + + P, _ = ot.dr.wda(xt, yt, 2, maxiter=3, solver="torch", P0=P0, random_state=0) + + assert P.device == xt.device + assert P.dtype == torch.float32 + + +@pytest.mark.skipif(nogo, reason="Missing modules (autograd or pymanopt)") +@pytest.mark.parametrize("solver", [None, "torch"]) +def test_wda_rejects_a_single_class(solver): + """One class has no between-class cost; both solvers must say so.""" + if solver == "torch" and notorch: + pytest.skip("Missing module (torch)") + rng = np.random.RandomState(0) + xs = rng.randn(40, 4) + ys = np.zeros(40, dtype=int) + + with pytest.raises(ValueError, match="at least two classes"): + ot.dr.wda(xs, ys, 2, maxiter=2, solver=solver)