Skip to content

[MRG] Add a PyTorch solver to ot.dr.wda (#806) - #858

Open
deeb01 wants to merge 12 commits into
PythonOT:masterfrom
deeb01:feat-806-wda-torch-solver
Open

deeb01 wants to merge 12 commits into
PythonOT:masterfrom
deeb01:feat-806-wda-torch-solver

Conversation

@deeb01

@deeb01 deeb01 commented Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

Types of changes

New feature (non-breaking change which adds functionality).

Motivation and context / Related issue

Addresses #806, following @rflamary's request for solver='torch' alongside the default solver='autograd'.

solver='torch' runs PyTorch autodiff with Riemannian gradient descent on Stiefel — QR retraction, backtracking with an adaptive initial step — mirroring pymanopt's SteepestDescent so both solvers target the same optimum. It needs only torch, and accepts torch tensors directly, keeping device and dtype.

ot.dr dependencies are now imported optionally, each function raising an ImportError naming what it needs. This changes import ot.dr from raising to succeeding on a partial install — flagging it in case you prefer the old behaviour.

On speed: pymanopt is 0.6–1.8% of runtime and this solver is 0.98–1.68× faster from n=600 up (slower at small n, where the per-call cost of ot.sinkhorn dominates), so the dependency choice is the real benefit. A frozen-plan gradient for a larger win stalls at a worse objective; numbers on the issue.

Also raises a clear ValueError when the between-class cost underflows to zero, previously a divide-by-zero.

How has this been tested (if it applies)

New tests assert 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 both solvers reach a comparable objective from the same start; the torch solver lands at 0.952–0.990 of pymanopt's objective across n, k and sinkhorn_method. Also covered: random_state, the line-search controls, device and dtype preservation, non-numeric labels, p > d, sinkhorn_log, input not mutated, error path.

test_dr.py 16 passed; full suite 2668 passed, 62 skipped, 6 xfailed; pre-commit clean.

PR checklist

  • I have read the CONTRIBUTING document.
  • The documentation is up-to-date with the changes I made (check build artifacts).
  • All tests passed, and additional code has been covered with new tests.
  • I have added the PR and Issue fix to the RELEASES.md file.

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.
@codecov

codecov Bot commented Sep 15, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 91.47541% with 26 lines in your changes missing coverage. Please review.
✅ Project coverage is 96.93%. Comparing base (604c47f) to head (b23f549).

Additional details and impacted files
@@            Coverage Diff             @@
##           master     #858      +/-   ##
==========================================
- Coverage   97.00%   96.93%   -0.07%     
==========================================
  Files         128      128              
  Lines       26349    26649     +300     
==========================================
+ Hits        25559    25833     +274     
- Misses        790      816      +26     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

rflamary and others added 8 commits September 15, 2026 10:22
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.
@cedricvincentcuaz cedricvincentcuaz self-assigned this Oct 6, 2026

@cedricvincentcuaz cedricvincentcuaz left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Dear @deeb01, thank you for your PR !
I started a code review, I let you check and let me know if you have any questions.
The main points to address relate to structuring the code to (1) minimize the number of newly hidden functions if they already exist in POT and/or whether they should actually be placed in a better place; (2) structure the code in order to proper shape wda as a wrapper. After that many tests should be added to not reduce that drastically coverage.
Best,
Cédric

Comment thread ot/dr.py Outdated
Comment thread ot/dr.py Outdated
Comment thread ot/dr.py Outdated
return G - P @ (0.5 * (W + W.T))


def _wda_cost_torch(P, xc, wc, regmean, reg, k, sinkhorn_solver):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

complete docstring

Comment thread ot/dr.py Outdated
Comment thread ot/dr.py Outdated
Comment thread ot/dr.py Outdated
if P0 is None:
P = torch.linalg.qr(torch.randn(d, p, dtype=dtype, device=device))[0]
else:
P = P0.clone().to(dtype)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

no control on device needed ?

Comment thread ot/dr.py Outdated


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)``.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

docstring to complete

Comment thread ot/dr.py Outdated
if P0 is None:
P0t = None
else:
P0t = (P0 if torch.is_tensor(P0) else torch.as_tensor(P0)).to(Xt.dtype)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

no control on device needed ?

Comment thread ot/dr.py Outdated

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

as now wda works as a wrapper that would make sense to the whole np solver out and call it here if needed.

Comment thread ot/dr.py

Parameters
----------
X : ndarray, shape (n, d)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

need to clarify and clean the logic to inherit the type and the device of the inputs

Following @cedricvincentcuaz's review on PythonOT#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.
@deeb01

deeb01 commented Oct 7, 2026

Copy link
Copy Markdown
Contributor Author

Thanks — the reuse points make the objective backend-agnostic rather than torch-only.

The three torch helpers are gone, replaced by ot.utils.dist and ot.bregman.sinkhorn with numItermax=k, stopThr=0. The graph is preserved: gradients agree to 1.1e-16 with the old versions. Overhead is 48% per call at n=200, within noise from n=600 up, so the description now says 0.98–1.68× rather than 1.0–1.75×. Stiefel projection and retraction use nx.dot/nx.qr/nx.sign.

wda is now a dispatcher, autograd path in _wda_autograd. random_state honoured, dtype and device followed, line-search constants exposed as step_growth/step_shrink/max_backtracks/gtol, docstrings completed. Seven tests added, including cross-backend agreement and the gradient against central differences.

On the line search: SteepestDescent defaults to BackTrackingLineSearcher, not Armijo. From the same start the torch solver reaches 0.952–0.990 of pymanopt's objective.

Kept _stiefel_projection as a named function for the tangency docstring; happy to inline.

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.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants