[MRG] Fix the total cost reported by the entropic unbalanced OT solvers - #874
Open
xihaian251 wants to merge 2 commits into
Open
xihaian251 wants to merge 2 commits into
xihaian251 wants to merge 2 commits into
Conversation
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
marked this pull request as ready for review
September 27, 2026 08:35
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
ot.solve(M, a, b, reg=..., unbalanced=...).valuedoes not return the value of the unbalanced OT problem that the solver minimizes. It is numerically zero for every problem.The same quantity is reported as
log["total_cost"]byot.unbalanced.sinkhorn_unbalanced(and throughreturnCost="total"), for the"sinkhorn","sinkhorn_stabilized"and"sinkhorn_translation_invariant"methods.Root cause
ot/unbalanced/_sinkhorn.pybuildstotal_costwith the un-normalized KLbut the unbalanced OT penalization is the generalized KL divergence
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 directionG -> 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=Trueto the threenx.kl_divcalls that buildtotal_cost, insinkhorn_knopp_unbalanced,sinkhorn_stabilized_unbalancedandsinkhorn_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_costintest/unbalanced/test_sinkhorn.py, parametrized over the three methods and over backends through thenxfixture. 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_valueintest/test_solvers.py, checking the public API:ot.solve(M, a, b, reg=1.0, unbalanced=0.5).valuematches the objective recomputed fromres.plan. Fails on master 2/2.Validation
Relevant test suites (
test/unbalanced/andtest/test_solvers.py): 1034 passed, 25 skipped.Additionally verified:
mass=Falseobjective is minimized at a different plan);valuewith respect toMmatches finite differences to 3.7e-10;reg_m=inf) agrees withot.sinkhornto 7.1e-08, and the semi-relaxedreg_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.