Skip to content

[MRG] Fix the total cost reported by the entropic unbalanced OT solvers - #874

Open
xihaian251 wants to merge 2 commits into
PythonOT:masterfrom
xihaian251:fix-uot-entropic-total-cost
Open

xihaian251 wants to merge 2 commits into
PythonOT:masterfrom
xihaian251:fix-uot-entropic-total-cost

Conversation

@xihaian251

Copy link
Copy Markdown

Problem

ot.solve(M, a, b, reg=..., unbalanced=...).value does not return the value of the unbalanced OT problem that the solver minimizes. It is numerically zero for every problem.

import numpy as np, ot
rng = np.random.RandomState(0)
a = rng.rand(7); a /= a.sum()
b = rng.rand(9); b /= b.sum()
M = rng.rand(7, 9); M /= M.max()

res = ot.solve(M, a, b, reg=0.1, unbalanced=1.0)
print(res.value)         # 7.58e-10, expected 0.23702381716399262
print(res.value_linear)  # 0.13066892631664662, correct

The same quantity is reported as log["total_cost"] by ot.unbalanced.sinkhorn_unbalanced (and through returnCost="total"), for the "sinkhorn", "sinkhorn_stabilized" and "sinkhorn_translation_invariant" methods.

Root cause

ot/unbalanced/_sinkhorn.py builds total_cost with the un-normalized KL

nx.kl_div(..., mass=False)   # sum(p log(p/q))

but the unbalanced OT penalization is the generalized KL divergence

nx.kl_div(..., mass=True)    # sum(p log(p/q)) - sum(p) + sum(q)

as already used by the exact unbalanced solver ot.unbalanced.mm_unbalanced (ot/unbalanced/_mm.py:189-196).

sum(p log(p/q)) is exactly the derivative of the true objective along the scaling direction G -> t G. The optimal plan of the entropic problem is interior, so by first-order optimality this derivative vanishes at the optimum: the reported value is the optimality residual, not the objective. That is why it is numerically zero for every problem.

Fix

Pass mass=True to the three nx.kl_div calls that build total_cost, in sinkhorn_knopp_unbalanced, sinkhorn_stabilized_unbalanced and sinkhorn_unbalanced_translation_invariant.

The transport plan, the solver iterations and the gradients are unchanged; only the reported total objective changes.

Tests

Two non-regression tests, both failing on master and passing with this change:

  • test_unbalanced_total_cost in test/unbalanced/test_sinkhorn.py, parametrized over the three methods and over backends through the nx fixture. It checks that the reported total cost is at least the linear cost (all penalizations are divergences) and that it matches the objective recomputed from the returned plan. Fails on master 6/6.
  • test_solve_unbalanced_value in test/test_solvers.py, checking the public API: ot.solve(M, a, b, reg=1.0, unbalanced=0.5).value matches the objective recomputed from res.plan. Fails on master 2/2.

Validation

Relevant test suites (test/unbalanced/ and test/test_solvers.py): 1034 passed, 25 skipped.

Additionally verified:

  • the returned plan is the minimizer of the generalized-KL objective, checked with an independent scipy L-BFGS-B minimization from several starting points (max |G_opt - G| = 0, while the mass=False objective is minimized at a different plan);
  • float32 inputs keep float32 outputs;
  • the gradient of value with respect to M matches finite differences to 3.7e-10;
  • the balanced limit (reg_m=inf) agrees with ot.sinkhorn to 7.1e-08, and the semi-relaxed reg_m=(inf, scalar) case behaves consistently.

Backends not available on the machine used for this patch: jax, tensorflow, cupy, and no GPU. Those are left to the CI.

Compatibility

This changes the reported total objective value, which was previously an incorrect quantity, to the documented objective. It does not change the transport plan, the solver iterations or any gradient. sinkhorn_unbalanced(..., returnCost="total") is affected in the same way.

The marginal penalization was computed with the un-normalized KL
divergence nx.kl_div(..., mass=False), i.e. sum(p log(p/q)), which is
not a divergence. At the optimum of the unbalanced OT problem this
quantity is the derivative of the objective along the scaling direction
G -> t G, so it vanishes: ot.solve(..., reg=..., unbalanced=...).value
was numerically zero for every problem, and log["total_cost"] of
ot.unbalanced.sinkhorn_unbalanced was too.

Use the generalized KL divergence (mass=True), as already done in
ot.unbalanced.mm_unbalanced, in sinkhorn_knopp_unbalanced,
sinkhorn_stabilized_unbalanced and
sinkhorn_unbalanced_translation_invariant. The transport plan, its
gradient and all other outputs are unchanged.

Add non-regression tests at the solver level and at the ot.solve level.
Both fail on master and pass with this change.
@xihaian251 xihaian251 changed the title [WIP] Fix the total cost reported by the entropic unbalanced OT solvers [MRG] Fix the total cost reported by the entropic unbalanced OT solvers Sep 27, 2026
@xihaian251
xihaian251 marked this pull request as ready for review September 27, 2026 08:35

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.

1 participant