From f13e3ecc9a0e79232fa29f5310fccfb818148ead Mon Sep 17 00:00:00 2001 From: grape7 <3796320131@qq.com> Date: Sun, 27 Sep 2026 16:21:27 +0800 Subject: [PATCH 1/2] Fix the total cost reported by entropic unbalanced OT solvers 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. --- RELEASES.md | 1 + ot/unbalanced/_sinkhorn.py | 33 +++++++++++++++------ test/test_solvers.py | 34 ++++++++++++++++++++++ test/unbalanced/test_sinkhorn.py | 50 ++++++++++++++++++++++++++++++++ 4 files changed, 109 insertions(+), 9 deletions(-) diff --git a/RELEASES.md b/RELEASES.md index 063701229..01aede590 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -16,6 +16,7 @@ #### Closed issues - Remove a leftover debug `print` from `ot.utils.projection_sparse_simplex` with `axis=1`, and make the `ot.datasets.make_gauss_hd` docstring a raw string so importing `ot` no longer emits a `SyntaxWarning` (PR #860) +- Fix the total cost reported by the entropic unbalanced OT solvers (`ot.unbalanced.sinkhorn_unbalanced` with methods `"sinkhorn"`, `"sinkhorn_stabilized"` and `"sinkhorn_translation_invariant"`, also exposed as `ot.solve(..., reg=..., unbalanced=...).value`): the marginal penalization now uses the generalized KL divergence (`mass=True`), so the value is the objective actually minimized by the solver instead of its derivative along `G -> t G`, which vanishes at the optimum (PR #XXX) - Fix `ot.dist` ignoring the weights `w` for `metric="cityblock"`, which returned the unweighted distance although the weights are documented for this metric (PR #859) - Fix swapped arguments to `div_to_product` in `ot.gromov.fused_unbalanced_across_spaces_cost`: with `reg_type="independent"` (UCOOT) the entropic terms used the plan marginals as the reference measures and vice versa (PR #855, Issue #854) - Fix device placement in `ot.batch.bregman_projection_batch` so `ot.solve_batch(..., method="sinkhorn")` no longer crashes on GPU when the torch default device is CPU (PR #851) diff --git a/ot/unbalanced/_sinkhorn.py b/ot/unbalanced/_sinkhorn.py index d338e1652..b25aa19b4 100644 --- a/ot/unbalanced/_sinkhorn.py +++ b/ot/unbalanced/_sinkhorn.py @@ -806,11 +806,16 @@ def sinkhorn_knopp_unbalanced( linear_cost = nx.sum(plan * M) dict_log["cost"] = linear_cost - total_cost = linear_cost + reg * nx.kl_div(plan, c) + # mass=True: the penalization is the generalized KL divergence + total_cost = linear_cost + reg * nx.kl_div(plan, c, mass=True) if reg_m1 != float("inf"): - total_cost = total_cost + reg_m1 * nx.kl_div(nx.sum(plan, 1), a) + total_cost = total_cost + reg_m1 * nx.kl_div( + nx.sum(plan, 1), a, mass=True + ) if reg_m2 != float("inf"): - total_cost = total_cost + reg_m2 * nx.kl_div(nx.sum(plan, 0), b) + total_cost = total_cost + reg_m2 * nx.kl_div( + nx.sum(plan, 0), b, mass=True + ) dict_log["total_cost"] = total_cost return plan, dict_log @@ -1106,11 +1111,16 @@ def sinkhorn_stabilized_unbalanced( linear_cost = nx.sum(plan * M) dict_log["cost"] = linear_cost - total_cost = linear_cost + reg * nx.kl_div(plan, c) + # mass=True: the penalization is the generalized KL divergence + total_cost = linear_cost + reg * nx.kl_div(plan, c, mass=True) if reg_m1 != float("inf"): - total_cost = total_cost + reg_m1 * nx.kl_div(nx.sum(plan, 1), a) + total_cost = total_cost + reg_m1 * nx.kl_div( + nx.sum(plan, 1), a, mass=True + ) if reg_m2 != float("inf"): - total_cost = total_cost + reg_m2 * nx.kl_div(nx.sum(plan, 0), b) + total_cost = total_cost + reg_m2 * nx.kl_div( + nx.sum(plan, 0), b, mass=True + ) dict_log["total_cost"] = total_cost return plan, dict_log @@ -1389,11 +1399,16 @@ def sinkhorn_unbalanced_translation_invariant( linear_cost = nx.sum(plan * M) dict_log["cost"] = linear_cost - total_cost = linear_cost + reg * nx.kl_div(plan, c) + # mass=True: the penalization is the generalized KL divergence + total_cost = linear_cost + reg * nx.kl_div(plan, c, mass=True) if reg_m1 != float("inf"): - total_cost = total_cost + reg_m1 * nx.kl_div(nx.sum(plan, 1), a) + total_cost = total_cost + reg_m1 * nx.kl_div( + nx.sum(plan, 1), a, mass=True + ) if reg_m2 != float("inf"): - total_cost = total_cost + reg_m2 * nx.kl_div(nx.sum(plan, 0), b) + total_cost = total_cost + reg_m2 * nx.kl_div( + nx.sum(plan, 0), b, mass=True + ) dict_log["total_cost"] = total_cost return plan, dict_log diff --git a/test/test_solvers.py b/test/test_solvers.py index 6cc8137cc..0d73cd6a8 100644 --- a/test/test_solvers.py +++ b/test/test_solvers.py @@ -384,6 +384,40 @@ def df(G): pytest.skip("Not implemented") +def test_solve_unbalanced_value(nx): + # ot.solve must return the value of the unbalanced OT problem it solves. + # The marginal penalization is the generalized KL divergence, i.e. it + # includes the mass correction term (mass=True). With the un-normalized KL + # the returned value is the derivative of the objective along G -> t G, + # which vanishes at the optimum. + rng = np.random.RandomState(0) + + x = rng.randn(10, 2) + y = rng.randn(7, 2) + a = ot.utils.unif(10) + b = ot.utils.unif(7) + M = ot.dist(x, y) + a, b, M = nx.from_numpy(a, b, M) + + reg = 1.0 + unbalanced = 0.5 + + res = ot.solve(M, a, b, reg=reg, unbalanced=unbalanced) + + G = res.plan + c = a[:, None] * b[None, :] + expected = nx.sum(G * M) + expected = expected + reg * nx.kl_div(G, c, mass=True) + expected = expected + unbalanced * nx.kl_div(nx.sum(G, 1), a, mass=True) + expected = expected + unbalanced * nx.kl_div(nx.sum(G, 0), b, mass=True) + + # the penalizations are divergences: the value is at least the linear loss + np.testing.assert_array_less( + nx.to_numpy(res.value_linear) - 1e-5, nx.to_numpy(res.value) + ) + np.testing.assert_allclose(nx.to_numpy(res.value), nx.to_numpy(expected), atol=1e-6) + + def test_solve_not_implemented(nx): n_samples_s = 10 n_samples_t = 7 diff --git a/test/unbalanced/test_sinkhorn.py b/test/unbalanced/test_sinkhorn.py index be7694309..844cf200b 100644 --- a/test/unbalanced/test_sinkhorn.py +++ b/test/unbalanced/test_sinkhorn.py @@ -809,3 +809,53 @@ def test_implemented_methods(nx): ot.unbalanced.sinkhorn_unbalanced(a, b, M, epsilon, reg_m, method=method) ot.unbalanced.sinkhorn_unbalanced2(a, b, M, epsilon, reg_m, method=method) barycenter_unbalanced(A, M, reg=epsilon, reg_m=reg_m, method=method) + + +@pytest.mark.parametrize( + "method", + ["sinkhorn", "sinkhorn_stabilized", "sinkhorn_translation_invariant"], +) +def test_unbalanced_total_cost(nx, method): + # The total cost reported in the log must be the value of the unbalanced OT + # objective that the solver actually minimizes. The marginal penalization is + # the generalized KL divergence, i.e. it includes the mass correction term + # (mass=True). Without it the reported value is the derivative of the + # objective along G -> t G, which vanishes at the optimum. + n = 20 + rng = np.random.RandomState(42) + + x = rng.randn(n, 2) + a = ot.utils.unif(n) + b = ot.utils.unif(n) * 1.5 # make the problem unbalanced + M = ot.dist(x, x) + a, b, M = nx.from_numpy(a, b, M) + + reg = 1.0 + reg_m = 1.0 + + G, log = ot.unbalanced.sinkhorn_unbalanced( + a, + b, + M, + reg=reg, + reg_m=reg_m, + method=method, + numItermax=5000, + stopThr=1e-12, + log=True, + ) + + c = a[:, None] * b[None, :] + expected = nx.sum(G * M) + expected = expected + reg * nx.kl_div(G, c, mass=True) + expected = expected + reg_m * nx.kl_div(nx.sum(G, 1), a, mass=True) + expected = expected + reg_m * nx.kl_div(nx.sum(G, 0), b, mass=True) + + # all penalizations are divergences: the total cost is at least the + # linear cost of the optimal plan + np.testing.assert_array_less( + nx.to_numpy(log["cost"]) - 1e-5, nx.to_numpy(log["total_cost"]) + ) + np.testing.assert_allclose( + nx.to_numpy(log["total_cost"]), nx.to_numpy(expected), atol=1e-6 + ) From 7e6d472918c6151d989eb37244c638b66ae36179 Mon Sep 17 00:00:00 2001 From: grape7 <3796320131@qq.com> Date: Sun, 27 Sep 2026 16:30:16 +0800 Subject: [PATCH 2/2] Update release note with PR number --- RELEASES.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/RELEASES.md b/RELEASES.md index 01aede590..fe89b7a28 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -16,7 +16,7 @@ #### Closed issues - Remove a leftover debug `print` from `ot.utils.projection_sparse_simplex` with `axis=1`, and make the `ot.datasets.make_gauss_hd` docstring a raw string so importing `ot` no longer emits a `SyntaxWarning` (PR #860) -- Fix the total cost reported by the entropic unbalanced OT solvers (`ot.unbalanced.sinkhorn_unbalanced` with methods `"sinkhorn"`, `"sinkhorn_stabilized"` and `"sinkhorn_translation_invariant"`, also exposed as `ot.solve(..., reg=..., unbalanced=...).value`): the marginal penalization now uses the generalized KL divergence (`mass=True`), so the value is the objective actually minimized by the solver instead of its derivative along `G -> t G`, which vanishes at the optimum (PR #XXX) +- Fix the total cost reported by the entropic unbalanced OT solvers (`ot.unbalanced.sinkhorn_unbalanced` with methods `"sinkhorn"`, `"sinkhorn_stabilized"` and `"sinkhorn_translation_invariant"`, also exposed as `ot.solve(..., reg=..., unbalanced=...).value`): the marginal penalization now uses the generalized KL divergence (`mass=True`), so the value is the objective actually minimized by the solver instead of its derivative along `G -> t G`, which vanishes at the optimum (PR #874) - Fix `ot.dist` ignoring the weights `w` for `metric="cityblock"`, which returned the unweighted distance although the weights are documented for this metric (PR #859) - Fix swapped arguments to `div_to_product` in `ot.gromov.fused_unbalanced_across_spaces_cost`: with `reg_type="independent"` (UCOOT) the entropic terms used the plan marginals as the reference measures and vice versa (PR #855, Issue #854) - Fix device placement in `ot.batch.bregman_projection_batch` so `ot.solve_batch(..., method="sinkhorn")` no longer crashes on GPU when the torch default device is CPU (PR #851)