Repository navigation
Conversation
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 Report❌ Patch coverage is 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:
|
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.
…feat-806-wda-torch-solver
cedricvincentcuaz
left a comment
There was a problem hiding this comment.
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
| return G - P @ (0.5 * (W + W.T)) | ||
|
|
||
|
|
||
| def _wda_cost_torch(P, xc, wc, regmean, reg, k, sinkhorn_solver): |
There was a problem hiding this comment.
complete docstring
| if P0 is None: | ||
| P = torch.linalg.qr(torch.randn(d, p, dtype=dtype, device=device))[0] | ||
| else: | ||
| P = P0.clone().to(dtype) |
There was a problem hiding this comment.
no control on device needed ?
|
|
||
|
|
||
| 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)``. |
There was a problem hiding this comment.
docstring to complete
| if P0 is None: | ||
| P0t = None | ||
| else: | ||
| P0t = (P0 if torch.is_tensor(P0) else torch.as_tensor(P0)).to(Xt.dtype) |
There was a problem hiding this comment.
no control on device needed ?
There was a problem hiding this comment.
as now wda works as a wrapper that would make sense to the whole np solver out and call it here if needed.
|
|
||
| Parameters | ||
| ---------- | ||
| X : ndarray, shape (n, d) |
There was a problem hiding this comment.
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.
|
Thanks — the reuse points make the objective backend-agnostic rather than torch-only. The three torch helpers are gone, replaced by
On the line search: Kept |
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.
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 defaultsolver='autograd'.solver='torch'runs PyTorch autodiff with Riemannian gradient descent on Stiefel — QR retraction, backtracking with an adaptive initial step — mirroring pymanopt'sSteepestDescentso both solvers target the same optimum. It needs only torch, and accepts torch tensors directly, keeping device and dtype.ot.drdependencies are now imported optionally, each function raising anImportErrornaming what it needs. This changesimport ot.drfrom 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.sinkhorndominates), 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
ValueErrorwhen 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.py16 passed; full suite 2668 passed, 62 skipped, 6 xfailed; pre-commit clean.PR checklist