From 6ffce16072b98b5102ab9bbfa429f15d9ed8e24e Mon Sep 17 00:00:00 2001 From: Ahmed Eldeeb <62363199+deeb01@users.noreply.github.com> Date: Mon, 14 Sep 2026 13:49:51 -0700 Subject: [PATCH 1/5] Add a PyTorch solver to ot.dr.wda (#806) Adds solver='torch' to wda: PyTorch autodiff with Riemannian gradient descent on the Stiefel manifold, using a QR retraction and backtracking with an adaptive initial step. It mirrors what pymanopt's SteepestDescent does, so both solvers target the same optimum rather than two different algorithms. The torch path needs only torch, so it works on installations without autograd or pymanopt, and accepts torch tensors directly, keeping their device and dtype. To make that possible, ot.dr's dependencies are now imported optionally and each function raises an ImportError naming what it needs, rather than the module failing to import unless all of them are present. Verified that the torch objective and its gradient match the autograd ones at the same point, and that both solvers reach a comparable objective from the same starting point. Also raises a clear ValueError when the between-class transport cost underflows to zero, which previously produced a divide-by-zero warning and an undefined objective. --- RELEASES.md | 7 ++ ot/dr.py | 249 ++++++++++++++++++++++++++++++++++++++++++++++-- test/test_dr.py | 147 +++++++++++++++++++++++++++- 3 files changed, 394 insertions(+), 9 deletions(-) diff --git a/RELEASES.md b/RELEASES.md index d636e8cf8..657bf11de 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -4,6 +4,13 @@ #### 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 #853, 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 #853) + +## 0.9.8dev + +#### New features + - 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 spiral points, selectable with `sampling_slices` in `sliced_wasserstein_distance`, diff --git a/ot/dr.py b/ot/dr.py index 9cfabff5f..a3922ee22 100644 --- a/ot/dr.py +++ b/ot/dr.py @@ -18,22 +18,54 @@ 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 .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 +111,184 @@ def split_classes(X, y): return [X[y == i, :].astype(np.float32) for i in lstsclass] +def _dist_torch(x1, x2): + r"""Squared euclidean distance between samples (torch).""" + return ( + torch.sum(x1**2, 1).reshape((-1, 1)) + + torch.sum(x2**2, 1).reshape((1, -1)) + - 2 * (x1 @ x2.T) + ) + + +def _sinkhorn_torch(w1, w2, M, reg, k): + r"""Sinkhorn algorithm with fixed number of iterations (torch).""" + K = torch.exp(-M / reg) + ui = torch.ones(M.shape[0], dtype=M.dtype, device=M.device) + vi = torch.ones(M.shape[1], dtype=M.dtype, device=M.device) + for _ in range(k): + vi = w2 / (K.T @ ui + 1e-50) + ui = w1 / (K @ vi + 1e-50) + return ui.reshape((-1, 1)) * K * vi.reshape((1, -1)) + + +def _sinkhorn_log_torch(w1, w2, M, reg, k): + r"""Sinkhorn algorithm in log-domain with fixed iterations (torch).""" + Mr = -M / reg + ui = torch.zeros(M.shape[0], dtype=M.dtype, device=M.device) + vi = torch.zeros(M.shape[1], dtype=M.dtype, device=M.device) + log_w1, log_w2 = torch.log(w1), torch.log(w2) + for _ in range(k): + vi = log_w2 - torch.logsumexp(Mr + ui[:, None], 0) + ui = log_w1 - torch.logsumexp(Mr + vi[None, :], 1) + return torch.exp(ui[:, None] + Mr + vi[None, :]) + + +def _stiefel_retract(P, X): + r"""QR retraction onto the Stiefel manifold, with a sign convention.""" + Q, R = torch.linalg.qr(P + X) + return Q * torch.sign(torch.sign(torch.diagonal(R)) + 0.5) + + +def _stiefel_project(P, G): + r"""Project a euclidean gradient onto the tangent space of Stiefel.""" + W = P.T @ G + return G - P @ (0.5 * (W + W.T)) + + +def _wda_cost_torch(P, xc, wc, regmean, reg, k, sinkhorn_solver): + r"""WDA objective: within-class transport cost over between-class.""" + loss_b, loss_w = 0.0, 0.0 + for i, xi in enumerate(xc): + xi = xi @ P + for j, xj in enumerate(xc[i:]): + xj = xj @ P + M = _dist_torch(xi, xj) + G = sinkhorn_solver(wc[i], wc[j + i], M, reg * regmean[i, j], k) + term = torch.sum(G * M) + if j == 0: + loss_w = loss_w + term + else: + loss_b = loss_b + term + if float(loss_b.detach() if torch.is_tensor(loss_b) else 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): + r"""WDA solved with PyTorch autodiff and Riemannian gradient descent. + + Mirrors the pymanopt ``SteepestDescent`` path: projected gradient, QR + retraction and backtracking, so both solvers target the same optimum. + """ + dtype = X.dtype + device = 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 P0 is None: + P = torch.linalg.qr(torch.randn(d, p, dtype=dtype, device=device))[0] + else: + P = P0.clone().to(dtype) + + 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_torch(xi, xj)) / ( + xi.shape[0] * xj.shape[0] + ) + + if sinkhorn_method.lower() == "sinkhorn": + solver_fn = _sinkhorn_torch + elif sinkhorn_method.lower() == "sinkhorn_log": + solver_fn = _sinkhorn_log_torch + else: + raise ValueError("Unknown Sinkhorn method '%s'." % sinkhorn_method) + + def value(Q): + with torch.no_grad(): + return _wda_cost_torch(Q, xc, wc, regmean, reg, k, solver_fn) + + f = value(P) + step = 1.0 + for it in range(maxiter): + Q = P.detach().requires_grad_(True) + v = _wda_cost_torch(Q, xc, wc, regmean, reg, k, solver_fn) + (g,) = torch.autograd.grad(v, Q) + direction = -_stiefel_project(P, g) + gnorm = float(torch.linalg.norm(direction)) + if verbose: + print(f"{it + 1:<6d} {float(v.detach()):+.16e} {gnorm:.8e}") + if gnorm <= 1e-12: + break + # start from twice the last accepted step, as pymanopt's backtracking + # line search does, so progress is not throttled by a fixed unit step + step = min(2.0 * step, 1e4 / (gnorm + 1e-12)) + improved = False + for _ in range(40): + Pn = _stiefel_retract(P, step * direction) + fn = value(Pn) + if fn < f: + improved = True + break + step *= 0.5 + 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): + r"""Convert inputs, centre, run the torch solver, return ``(P, proj)``. + + numpy in gives numpy out; a torch tensor in keeps its device and dtype. + """ + was_numpy = not torch.is_tensor(X) + Xt = torch.as_tensor(X) if was_numpy else X + if not torch.is_floating_point(Xt): + Xt = Xt.to(torch.float64) + yt = y if torch.is_tensor(y) else torch.as_tensor(np.asarray(y)) + if P0 is None: + P0t = None + else: + P0t = (P0 if torch.is_tensor(P0) else torch.as_tensor(P0)).to(Xt.dtype) + + 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 + ) + + 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 @@ -184,9 +394,17 @@ def wda( 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 + 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) @@ -211,6 +429,20 @@ def wda( Wasserstein Discriminant Analysis. arXiv preprint arXiv:1608.08063. """ # noqa + 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 + ) + + _require( + HAS_AUTOGRAD and HAS_PYMANOPT, + "wda(solver='autograd')", + "autograd and pymanopt", + ) + if solver == "autograd": + solver = None + if sinkhorn_method.lower() == "sinkhorn": sinkhorn_solver = sinkhorn elif sinkhorn_method.lower() == "sinkhorn_log": @@ -495,6 +727,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..7a7fdb92e 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,141 @@ 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(nogo or notorch, reason="Missing modules") +def test_wda_backends_agree_on_cost_and_gradient(): + """The torch objective and its gradient must match the autograd ones.""" + import autograd + import autograd.numpy as anp + + 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)]) + y = np.repeat(np.arange(C), n // C) + P0 = np.linalg.qr(rng.randn(d, 2))[0] + Xc = X - X.mean(0) + + xc = [np.ascontiguousarray(Xc[y == c]) for c in range(C)] + wc = [np.ones(x.shape[0]) / x.shape[0] for x in xc] + + def cost_autograd(P): + loss_b, loss_w = 0.0, 0.0 + for i, xi in enumerate(xc): + xi = anp.dot(xi, P) + for j, xj in enumerate(xc[i:]): + xj = anp.dot(xj, P) + M = ot.dr.dist(xi, xj) + G = ot.dr.sinkhorn(wc[i], wc[j + i], M, reg, k) + term = anp.sum(G * M) + if j == 0: + loss_w = loss_w + term + else: + loss_b = loss_b + term + return loss_w / loss_b + + 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, requires_grad=True) + v = ot.dr._wda_cost_torch(Pt, xct, wct, rmt, reg, k, ot.dr._sinkhorn_torch) + (g,) = torch.autograd.grad(v, Pt) + + np.testing.assert_allclose(float(v.detach()), float(cost_autograd(P0)), rtol=1e-10) + np.testing.assert_allclose( + g.numpy(), autograd.grad(cost_autograd)(P0), rtol=1e-8, 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) + + def objective(P): + return float( + ot.dr._wda_cost_torch( + torch.tensor(P, dtype=torch.float64), + xc, + wc, + rm, + 1, + 10, + ot.dr._sinkhorn_torch, + ) + ) + + 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") From 227d5f68614851936f1c88ffed71550c93b05ab3 Mon Sep 17 00:00:00 2001 From: Ahmed Eldeeb <62363199+deeb01@users.noreply.github.com> Date: Mon, 14 Sep 2026 21:31:40 -0700 Subject: [PATCH 2/5] Correct PR number in RELEASES.md --- RELEASES.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/RELEASES.md b/RELEASES.md index 657bf11de..a3da914a0 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -4,8 +4,8 @@ #### 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 #853, 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 #853) +- `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) ## 0.9.8dev From 522802ad6cf2ba2e27935177fbcc0941a5a4de2a Mon Sep 17 00:00:00 2001 From: Ahmed Eldeeb <62363199+deeb01@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:32:06 -0700 Subject: [PATCH 3/5] Fix three differences between the torch and autograd wda solvers Found while reviewing the diff, all three cases where solver='torch' behaved differently from the default solver: - Non-numeric labels raised TypeError. numpy's split_classes indexes classes by value, so string labels work there, but torch.unique cannot hold them. Labels are now mapped to positional codes first. - p > d silently returned a wrongly shaped P. torch.linalg.qr on a (d, p) matrix with p > d returns a (d, d) factor, and nothing downstream objected. pymanopt's Stiefel(d, p) raises for this, so the torch path now checks 1 <= p <= d explicitly. - float32 numpy input returned float32 while the autograd path returns float64. numpy input is now promoted to float64; a torch tensor still keeps its own dtype, as documented. Adds a regression test for each. --- ot/dr.py | 18 ++++++++++++++---- test/test_dr.py | 32 ++++++++++++++++++++++++++++++++ 2 files changed, 46 insertions(+), 4 deletions(-) diff --git a/ot/dr.py b/ot/dr.py index a3922ee22..e55f92115 100644 --- a/ot/dr.py +++ b/ot/dr.py @@ -196,6 +196,9 @@ def _wda_torch(X, y, p, reg, k, sinkhorn_method, maxiter, verbose, P0, normalize 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: P = torch.linalg.qr(torch.randn(d, p, dtype=dtype, device=device))[0] else: @@ -258,10 +261,17 @@ def _wda_torch_entry(X, y, p, reg, k, sinkhorn_method, maxiter, verbose, P0, nor numpy in gives numpy out; a torch tensor in keeps its device and dtype. """ was_numpy = not torch.is_tensor(X) - Xt = torch.as_tensor(X) if was_numpy else X - if not torch.is_floating_point(Xt): - Xt = Xt.to(torch.float64) - yt = y if torch.is_tensor(y) else torch.as_tensor(np.asarray(y)) + if was_numpy: + # match the autograd path, which promotes to float64 via P + Xt = torch.as_tensor(np.asarray(X, dtype=np.float64)) + else: + Xt = X if torch.is_floating_point(X) else X.to(torch.float64) + # labels may be strings or any hashable, which torch cannot hold, so index + # them by position the way numpy's split_classes does + if torch.is_tensor(y): + yt = y + else: + yt = torch.as_tensor(np.unique(np.asarray(y), return_inverse=True)[1]) if P0 is None: P0t = None else: diff --git a/test/test_dr.py b/test/test_dr.py index 7a7fdb92e..86ae9ce32 100644 --- a/test/test_dr.py +++ b/test/test_dr.py @@ -396,3 +396,35 @@ def test_wda_torch_unknown_sinkhorn_method(): 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 From b74468dbe5b5e39ca9cc2bf0bdf96dbc6d2a1f81 Mon Sep 17 00:00:00 2001 From: Ahmed Eldeeb <62363199+deeb01@users.noreply.github.com> Date: Wed, 7 Oct 2026 15:15:54 -0700 Subject: [PATCH 4/5] Address review: reuse POT primitives, make wda a dispatcher Following @cedricvincentcuaz's review on #858. Reuse rather than reimplement. _dist_torch, _sinkhorn_torch and _sinkhorn_log_torch are gone; the objective now calls ot.utils.dist and ot.bregman.sinkhorn with numItermax=k and stopThr=0 to keep the fixed-depth iteration that WDA differentiates through. Verified against the previous hand-rolled versions: identical value and gradients agreeing to 1.1e-16, with autograd flowing through ot.sinkhorn. The Stiefel projection and retraction now go through nx.dot, nx.qr and nx.sign, so the objective is backend-agnostic rather than torch-only; only the gradient step is torch-specific. Per-call cost of routing through ot.sinkhorn is 48% at n=200 but within noise from n=600 up (-0.6% at n=600, 4.4% at n=1800, -2.9% at n=1800 k=50), since the extra work is fixed per call rather than per iteration. wda is now a dispatcher: the autograd path moved into _wda_autograd, and every private function has a full parameter docstring. random_state is now honoured for the random starting point, P0 and the labels follow the dtype and device of X, and the line-search constants are exposed as step_growth, step_shrink, max_backtracks and gtol with their previous values as defaults. New tests: the objective agrees between the numpy and torch backends, its autodiff gradient matches central differences, the Stiefel projection is tangent and the retraction stays on the manifold and fixes a zero step, random_state is reproducible, the line-search controls take effect, and the device and dtype of a torch input are preserved. --- ot/dr.py | 517 ++++++++++++++++++++++++++++++++++-------------- test/test_dr.py | 174 ++++++++++++---- 2 files changed, 500 insertions(+), 191 deletions(-) diff --git a/ot/dr.py b/ot/dr.py index e55f92115..491b7726b 100644 --- a/ot/dr.py +++ b/ot/dr.py @@ -53,6 +53,7 @@ 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 @@ -111,65 +112,108 @@ def split_classes(X, y): return [X[y == i, :].astype(np.float32) for i in lstsclass] -def _dist_torch(x1, x2): - r"""Squared euclidean distance between samples (torch).""" - return ( - torch.sum(x1**2, 1).reshape((-1, 1)) - + torch.sum(x2**2, 1).reshape((1, -1)) - - 2 * (x1 @ x2.T) - ) +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. -def _sinkhorn_torch(w1, w2, M, reg, k): - r"""Sinkhorn algorithm with fixed number of iterations (torch).""" - K = torch.exp(-M / reg) - ui = torch.ones(M.shape[0], dtype=M.dtype, device=M.device) - vi = torch.ones(M.shape[1], dtype=M.dtype, device=M.device) - for _ in range(k): - vi = w2 / (K.T @ ui + 1e-50) - ui = w1 / (K @ vi + 1e-50) - return ui.reshape((-1, 1)) * K * vi.reshape((1, -1)) + 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 _sinkhorn_log_torch(w1, w2, M, reg, k): - r"""Sinkhorn algorithm in log-domain with fixed iterations (torch).""" - Mr = -M / reg - ui = torch.zeros(M.shape[0], dtype=M.dtype, device=M.device) - vi = torch.zeros(M.shape[1], dtype=M.dtype, device=M.device) - log_w1, log_w2 = torch.log(w1), torch.log(w2) - for _ in range(k): - vi = log_w2 - torch.logsumexp(Mr + ui[:, None], 0) - ui = log_w1 - torch.logsumexp(Mr + vi[None, :], 1) - return torch.exp(ui[:, None] + Mr + vi[None, :]) +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 _stiefel_retract(P, X): - r"""QR retraction onto the Stiefel manifold, with a sign convention.""" - Q, R = torch.linalg.qr(P + X) - return Q * torch.sign(torch.sign(torch.diagonal(R)) + 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}`. -def _stiefel_project(P, G): - r"""Project a euclidean gradient onto the tangent space of Stiefel.""" - W = P.T @ G - return G - P @ (0.5 * (W + W.T)) + 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. -def _wda_cost_torch(P, xc, wc, regmean, reg, k, sinkhorn_solver): - r"""WDA objective: within-class transport cost over between-class.""" + 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 = xi @ P + xi = nx.dot(xi, P) for j, xj in enumerate(xc[i:]): - xj = xj @ P - M = _dist_torch(xi, xj) - G = sinkhorn_solver(wc[i], wc[j + i], M, reg * regmean[i, j], k) - term = torch.sum(G * M) + 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(loss_b.detach() if torch.is_tensor(loss_b) else loss_b) == 0.0: + 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 " @@ -179,14 +223,70 @@ def _wda_cost_torch(P, xc, wc, regmean, reg, k, sinkhorn_solver): return loss_w / loss_b -def _wda_torch(X, y, p, reg, k, sinkhorn_method, maxiter, verbose, P0, normalize): - r"""WDA solved with PyTorch autodiff and Riemannian gradient descent. +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. - Mirrors the pymanopt ``SteepestDescent`` path: projected gradient, QR - retraction and backtracking, so both solvers 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. """ - dtype = X.dtype - device = X.device + nx = get_backend(X) + dtype, device = X.dtype, X.device labels = torch.unique(y) xc = [X[y == c] for c in labels] wc = [ @@ -200,9 +300,12 @@ def _wda_torch(X, y, p, reg, k, sinkhorn_method, maxiter, verbose, P0, normalize raise ValueError(f"Need d >= p >= 1. Values supplied were d = {d} and p = {p}") if P0 is None: - P = torch.linalg.qr(torch.randn(d, p, dtype=dtype, device=device))[0] + 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) + P = P0.clone().to(dtype=dtype, device=device) regmean = torch.ones((nc, nc), dtype=dtype, device=device) if P0 is not None and normalize: @@ -211,77 +314,121 @@ def _wda_torch(X, y, p, reg, k, sinkhorn_method, maxiter, verbose, P0, normalize xi = xi @ P for j, xj in enumerate(xc[i:]): xj = xj @ P - regmean[i, j] = torch.sum(_dist_torch(xi, xj)) / ( + regmean[i, j] = torch.sum(dist_utils(xi, xj)) / ( xi.shape[0] * xj.shape[0] ) - if sinkhorn_method.lower() == "sinkhorn": - solver_fn = _sinkhorn_torch - elif sinkhorn_method.lower() == "sinkhorn_log": - solver_fn = _sinkhorn_log_torch - else: - raise ValueError("Unknown Sinkhorn method '%s'." % sinkhorn_method) - def value(Q): with torch.no_grad(): - return _wda_cost_torch(Q, xc, wc, regmean, reg, k, solver_fn) + 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_torch(Q, xc, wc, regmean, reg, k, solver_fn) + v = _wda_cost(Q, xc, wc, regmean, reg, k, sinkhorn_method, nx) (g,) = torch.autograd.grad(v, Q) - direction = -_stiefel_project(P, g) - gnorm = float(torch.linalg.norm(direction)) + 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 <= 1e-12: + if gnorm <= gtol: break - # start from twice the last accepted step, as pymanopt's backtracking - # line search does, so progress is not throttled by a fixed unit step - step = min(2.0 * step, 1e4 / (gnorm + 1e-12)) + step = step_growth * step improved = False - for _ in range(40): - Pn = _stiefel_retract(P, step * direction) + for _ in range(max_backtracks): + Pn = _stiefel_retraction(P, step * direction, nx) fn = value(Pn) if fn < f: improved = True break - step *= 0.5 + 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): - r"""Convert inputs, centre, run the torch solver, return ``(P, proj)``. +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. - numpy in gives numpy out; a torch tensor in keeps its device and dtype. + 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: - # match the autograd path, which promotes to float64 via P Xt = torch.as_tensor(np.asarray(X, dtype=np.float64)) else: Xt = X if torch.is_floating_point(X) else X.to(torch.float64) - # labels may be strings or any hashable, which torch cannot hold, so index - # them by position the way numpy's split_classes does + + # 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 + yt = y.to(Xt.device) else: - yt = torch.as_tensor(np.unique(np.asarray(y), return_inverse=True)[1]) + 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(P0)).to(Xt.dtype) + 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 + Xc, + yt, + p, + reg, + k, + sinkhorn_method, + maxiter, + verbose, + P0t, + normalize, + random_state, + step_growth, + step_shrink, + max_backtracks, + gtol, ) if was_numpy: @@ -358,6 +505,106 @@ def proj(X): return Popt, proj +def _wda_autograd( + X, y, p, reg, k, solver, sinkhorn_method, maxiter, verbose, P0, normalize +): + r"""Solve WDA with autograd and a pymanopt optimizer. + + Parameters + ---------- + X : ndarray, shape (n, d) + Training samples, not modified. + y : ndarray, shape (n,) + Labels for training samples. + 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 projection. + proj : callable + Projection function including mean centering. + """ + if solver == "autograd": + solver = None + + if sinkhorn_method.lower() == "sinkhorn": + sinkhorn_solver = sinkhorn + elif sinkhorn_method.lower() == "sinkhorn_log": + sinkhorn_solver = sinkhorn_log + else: + raise ValueError("Unknown Sinkhorn method '%s'." % sinkhorn_method) + + mx = np.mean(X, axis=0) + X = X - mx.reshape((1, -1)) + + # data split between classes + d = X.shape[1] + xc = split_classes(X, y) + # compute uniform weighs + wc = [np.ones((x.shape[0]), dtype=np.float32) / x.shape[0] for x in xc] + + # pre-compute reg_c,c' + if P0 is not None and normalize: + regmean = np.zeros((len(xc), len(xc))) + for i, xi in enumerate(xc): + xi = np.dot(xi, P0) + for j, xj in enumerate(xc[i:]): + xj = np.dot(xj, P0) + M = dist(xi, xj) + regmean[i, j] = np.sum(M) / (len(xi) * len(xj)) + else: + regmean = np.ones((len(xc), len(xc))) + + manifold = pymanopt.manifolds.Stiefel(d, p) + + @pymanopt.function.autograd(manifold) + def cost(P): + # wda loss + loss_b = 0 + loss_w = 0 + + for i, xi in enumerate(xc): + xi = np.dot(xi, P) + for j, xj in enumerate(xc[i:]): + xj = np.dot(xj, P) + M = dist(xi, xj) + G = sinkhorn_solver(wc[i], wc[j + i], M, reg * regmean[i, j], k) + if j == 0: + loss_w += np.sum(G * M) + else: + loss_b += np.sum(G * M) + + # loss inversed because minimization + return loss_w / loss_b + + # declare manifold and problem + + problem = pymanopt.Problem(manifold=manifold, cost=cost) + + # declare solver and solve + if solver is None: + solver = pymanopt.optimizers.SteepestDescent( + max_iterations=maxiter, log_verbosity=verbose + ) + elif solver in ["tr", "TrustRegions"]: + solver = pymanopt.optimizers.TrustRegions( + max_iterations=maxiter, log_verbosity=verbose + ) + + Popt = solver.run(problem, initial_point=P0) + + def proj(X): + return (X - mx.reshape((1, -1))).dot(Popt.point) + + return Popt.point, proj + + def wda( X, y, @@ -370,6 +617,11 @@ def wda( 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] ` @@ -423,6 +675,21 @@ def wda( 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 ------- @@ -442,7 +709,21 @@ def wda( 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 + X, + y, + p, + reg, + k, + sinkhorn_method, + maxiter, + verbose, + P0, + normalize, + random_state, + step_growth, + step_shrink, + max_backtracks, + gtol, ) _require( @@ -450,79 +731,9 @@ def wda( "wda(solver='autograd')", "autograd and pymanopt", ) - if solver == "autograd": - solver = None - - if sinkhorn_method.lower() == "sinkhorn": - sinkhorn_solver = sinkhorn - elif sinkhorn_method.lower() == "sinkhorn_log": - sinkhorn_solver = sinkhorn_log - else: - raise ValueError("Unknown Sinkhorn method '%s'." % sinkhorn_method) - - mx = np.mean(X, axis=0) - X = X - mx.reshape((1, -1)) - - # data split between classes - d = X.shape[1] - xc = split_classes(X, y) - # compute uniform weighs - wc = [np.ones((x.shape[0]), dtype=np.float32) / x.shape[0] for x in xc] - - # pre-compute reg_c,c' - if P0 is not None and normalize: - regmean = np.zeros((len(xc), len(xc))) - for i, xi in enumerate(xc): - xi = np.dot(xi, P0) - for j, xj in enumerate(xc[i:]): - xj = np.dot(xj, P0) - M = dist(xi, xj) - regmean[i, j] = np.sum(M) / (len(xi) * len(xj)) - else: - regmean = np.ones((len(xc), len(xc))) - - manifold = pymanopt.manifolds.Stiefel(d, p) - - @pymanopt.function.autograd(manifold) - def cost(P): - # wda loss - loss_b = 0 - loss_w = 0 - - for i, xi in enumerate(xc): - xi = np.dot(xi, P) - for j, xj in enumerate(xc[i:]): - xj = np.dot(xj, P) - M = dist(xi, xj) - G = sinkhorn_solver(wc[i], wc[j + i], M, reg * regmean[i, j], k) - if j == 0: - loss_w += np.sum(G * M) - else: - loss_b += np.sum(G * M) - - # loss inversed because minimization - return loss_w / loss_b - - # declare manifold and problem - - problem = pymanopt.Problem(manifold=manifold, cost=cost) - - # declare solver and solve - if solver is None: - solver = pymanopt.optimizers.SteepestDescent( - max_iterations=maxiter, log_verbosity=verbose - ) - elif solver in ["tr", "TrustRegions"]: - solver = pymanopt.optimizers.TrustRegions( - max_iterations=maxiter, log_verbosity=verbose - ) - - Popt = solver.run(problem, initial_point=P0) - - def proj(X): - return (X - mx.reshape((1, -1))).dot(Popt.point) - - return Popt.point, proj + return _wda_autograd( + X, y, p, reg, k, solver, sinkhorn_method, maxiter, verbose, P0, normalize + ) def projection_robust_wasserstein( diff --git a/test/test_dr.py b/test/test_dr.py index 86ae9ce32..7ac38e47f 100644 --- a/test/test_dr.py +++ b/test/test_dr.py @@ -312,47 +312,88 @@ def test_wda_torch_sinkhorn_log(): np.testing.assert_allclose(np.sum(P**2, 0), np.ones(p), rtol=1e-6) -@pytest.mark.skipif(nogo or notorch, reason="Missing modules") -def test_wda_backends_agree_on_cost_and_gradient(): - """The torch objective and its gradient must match the autograd ones.""" - import autograd - import autograd.numpy as anp - +@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 = X - X.mean(0) - xc = [np.ascontiguousarray(Xc[y == c]) for c in range(C)] + xc = [np.ascontiguousarray(X[y == c]) for c in range(C)] wc = [np.ones(x.shape[0]) / x.shape[0] for x in xc] - - def cost_autograd(P): - loss_b, loss_w = 0.0, 0.0 - for i, xi in enumerate(xc): - xi = anp.dot(xi, P) - for j, xj in enumerate(xc[i:]): - xj = anp.dot(xj, P) - M = ot.dr.dist(xi, xj) - G = ot.dr.sinkhorn(wc[i], wc[j + i], M, reg, k) - term = anp.sum(G * M) - if j == 0: - loss_w = loss_w + term - else: - loss_b = loss_b + term - return loss_w / loss_b + 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, requires_grad=True) - v = ot.dr._wda_cost_torch(Pt, xct, wct, rmt, reg, k, ot.dr._sinkhorn_torch) - (g,) = torch.autograd.grad(v, Pt) + 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(float(v.detach()), float(cost_autograd(P0)), rtol=1e-10) + 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( - g.numpy(), autograd.grad(cost_autograd)(P0), rtol=1e-8, atol=1e-10 + ot.dr._stiefel_retraction(P, np.zeros_like(P), nx), P, atol=1e-10 ) @@ -372,19 +413,22 @@ def test_wda_solvers_reach_comparable_objective(): 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): - return float( - ot.dr._wda_cost_torch( - torch.tensor(P, dtype=torch.float64), - xc, - wc, - rm, - 1, - 10, - ot.dr._sinkhorn_torch, + 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) @@ -428,3 +472,57 @@ def test_wda_torch_numpy_input_returns_float64(): 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 From b23f549264f8b57613b3b5c979e54e33d51b1d75 Mon Sep 17 00:00:00 2001 From: Ahmed Eldeeb <62363199+deeb01@users.noreply.github.com> Date: Wed, 7 Oct 2026 15:25:29 -0700 Subject: [PATCH 5/5] Reject single-class input in wda With one class there is no between-class transport cost, so the WDA objective is 0/0. The autograd solver ran through a cascade of divide-by-zero warnings and returned a NaN projection without raising, and the torch solver raised, but blamed a too-small reg. wda now checks the number of classes before dispatching and raises a ValueError saying so, for both solvers. Correction to the previous commit message, which said the line-search constants were exposed "with their previous values as defaults". The old cap on the initial trial step, min(2 * step, 1e4 / gnorm), was dropped rather than exposed. Measured on 300-iteration runs across five configurations, every run still stops at a tangent gradient norm between 1e-10 and 5e-10 with an objective equal to or below pymanopt's, so the behaviour is kept as is. --- ot/dr.py | 11 +++++++++++ test/test_dr.py | 14 ++++++++++++++ 2 files changed, 25 insertions(+) diff --git a/ot/dr.py b/ot/dr.py index 491b7726b..8336ccfeb 100644 --- a/ot/dr.py +++ b/ot/dr.py @@ -706,6 +706,17 @@ def wda( 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( diff --git a/test/test_dr.py b/test/test_dr.py index 7ac38e47f..501e14781 100644 --- a/test/test_dr.py +++ b/test/test_dr.py @@ -526,3 +526,17 @@ def test_wda_torch_preserves_device(): 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)