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/6] 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/6] 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) From ac97265fc227c319ebeacd84fff1db225bc49cad Mon Sep 17 00:00:00 2001 From: B Date: Fri, 9 Oct 2026 11:11:18 +0800 Subject: [PATCH 3/6] Address review of #874: entropy constant, reg_type consistency, docs Remove the constant dim_a * dim_b from the value reported by the entropic unbalanced OT solvers. With reg_type="entropy" the reference measure is the all-ones matrix, so the solver minimizes reg * KL_gen(gamma, 1) while the documented regularizer is Omega_ent(gamma) = sum(gamma log gamma - gamma) = KL_gen(gamma, 1) - n_a n_b. The offset is constant, so the plans were correct, but the reported value was reg * n_a * n_b too large (measured exactly +9.0 for reg=0.3, n_a=5, n_b=6, on the three methods). Also: - normalise reg_type once per solver through a _check_reg_type helper: it is now handled case insensitively by every method (sinkhorn_stabilized_unbalanced and sinkhorn_unbalanced_translation_invariant used to silently fall back to 'kl' for capitalised spellings, solving a different problem), and an unsupported name raises ValueError instead of being silently treated as 'kl'; - document the generalized KL convention, the regularizer Omega_reg_type(gamma, c) with its two cases, the keys of the log dictionary and the multi-histogram behaviour, with a single identical reg_type block for the five public functions; - warn when the multi-histogram path silently ignores reg_type or c, and validate returnCost in that branch; - remove test_solve_unbalanced_value as requested by the reviewer. Tests: five new tests in test/unbalanced/test_sinkhorn.py fail on the previous head (24 failed) and pass after the change (38 passed); every reg_type="kl" configuration is unchanged in plan, cost, total_cost and log keys over an A/B matrix of 199 configurations. --- RELEASES.md | 2 +- ot/unbalanced/_sinkhorn.py | 357 ++++++++++++++++++++++++++----- test/test_solvers.py | 34 --- test/unbalanced/test_sinkhorn.py | 246 ++++++++++++++++++--- 4 files changed, 520 insertions(+), 119 deletions(-) diff --git a/RELEASES.md b/RELEASES.md index f539ea1d6..abb8f04c8 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -17,7 +17,7 @@ - Allow `NumpyBackend.seed` to adopt an existing `np.random.RandomState` instance and remove NumPy-specific random sampling paths in sliced utilities (PR #849, Issue #848) - 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 #874) +- Fix the value reported by the entropic unbalanced OT solvers (`ot.unbalanced.sinkhorn_unbalanced`, `sinkhorn_unbalanced2`, `sinkhorn_knopp_unbalanced`, `sinkhorn_stabilized_unbalanced` and `sinkhorn_unbalanced_translation_invariant`, also exposed as `ot.solve(..., reg=..., unbalanced=...).value`): every penalization is now evaluated with the generalized KL divergence, and with `reg_type="entropy"` the constant `dim_a * dim_b` is removed so that the reported value matches the documented negative entropy regularizer. The transport plans are unchanged. `reg_type` is now handled case insensitively by all the methods and an unknown value raises a `ValueError` instead of silently falling back to `"kl"`. The docstrings now state the generalized KL convention, the regularizer `Omega_reg_type(gamma, c)` with its two cases, the keys of the `log` dictionary and the multi-histogram behaviour (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) diff --git a/ot/unbalanced/_sinkhorn.py b/ot/unbalanced/_sinkhorn.py index b25aa19b4..33d5bd1ff 100644 --- a/ot/unbalanced/_sinkhorn.py +++ b/ot/unbalanced/_sinkhorn.py @@ -15,6 +15,32 @@ from ..backend import get_backend from ..utils import list_to_array, get_parameter_pair +REG_TYPES = ("kl", "entropy") + + +def _check_reg_type(reg_type): + r"""Validate `reg_type` and return its lower case form. + + Only ``'kl'`` and ``'entropy'`` are implemented. Any other value used to be + silently handled as ``'kl'``, which silently changed the solved problem, so + it is now rejected explicitly. + + Parameters + ---------- + reg_type : str + Name of the regularizer, either 'kl' or 'entropy' (case insensitive). + + Returns + ------- + reg_type : str + The lower case name of the regularizer. + """ + if not isinstance(reg_type, str) or reg_type.lower() not in REG_TYPES: + raise ValueError( + "Unknown reg_type '{}'. Must be either 'kl' or 'entropy'.".format(reg_type) + ) + return reg_type.lower() + def sinkhorn_unbalanced( a, @@ -40,7 +66,7 @@ def sinkhorn_unbalanced( .. math:: W = \arg \min_\gamma \ \langle \gamma, \mathbf{M} \rangle_F + - \mathrm{reg} \cdot \mathrm{KL}(\gamma, \mathbf{c}) + + \mathrm{reg} \cdot \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) + \mathrm{reg_{m1}} \cdot \mathrm{KL}(\gamma \mathbf{1}, \mathbf{a}) + \mathrm{reg_{m2}} \cdot \mathrm{KL}(\gamma^T \mathbf{1}, \mathbf{b}) @@ -52,7 +78,17 @@ def sinkhorn_unbalanced( - :math:`\mathbf{M}` is the (`dim_a`, `dim_b`) metric cost matrix - :math:`\mathbf{a}` and :math:`\mathbf{b}` are source and target unbalanced distributions - :math:`\mathbf{c}` is a reference distribution for the regularization - - KL is the Kullback-Leibler divergence + - :math:`\mathrm{KL}` is the generalized Kullback-Leibler divergence + :math:`\mathrm{KL}(\mathbf{P}, \mathbf{Q}) = \sum_{i,j} \mathbf{P}_{i,j} \log(\mathbf{P}_{i,j} / \mathbf{Q}_{i,j}) - \mathbf{P}_{i,j} + \mathbf{Q}_{i,j}` + + and :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})` is the regularizer + + .. math:: + \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) = + \begin{cases} + \mathrm{KL}(\gamma, \mathbf{c}) & \text{if } \texttt{reg\_type} = \text{'kl'}, \\ + \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j} & \text{if } \texttt{reg\_type} = \text{'entropy'}. + \end{cases} The algorithm used for solving the problem is the generalized Sinkhorn-Knopp matrix scaling algorithm as proposed in :ref:`[10, 25] @@ -92,13 +128,19 @@ def sinkhorn_unbalanced( method used for the solver either 'sinkhorn', 'sinkhorn_stabilized', 'sinkhorn_translation_invariant' or 'sinkhorn_reg_scaling', see those function for specific parameters reg_type : string, optional - Regularizer term. Can take two values: + Name of the regularizer :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})`, + case insensitive. Can take two values: - Negative entropy: 'entropy': :math:`\Omega(\gamma) = \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j}`. - This is equivalent (up to a constant) to :math:`\Omega(\gamma) = \text{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)`. - - Kullback-Leibler divergence (default): 'kl': - :math:`\Omega(\gamma) = \text{KL}(\gamma, \mathbf{a} \mathbf{b}^T)`. + It is equal to :math:`\mathrm{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)` minus the + constant :math:`dim_a \times dim_b`, and the reference measure :math:`\mathbf{c}` + is overwritten by the all-ones matrix. + - Generalized Kullback-Leibler divergence (default): 'kl': + :math:`\Omega(\gamma, \mathbf{c}) = \mathrm{KL}(\gamma, \mathbf{c})` with + :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T` when `c` is None. + + Any other value raises a :class:`ValueError`. c : array-like, shape (dim_a, dim_b), optional (default=None) Reference measure for the regularization. If None, then use :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T`. @@ -118,7 +160,7 @@ def sinkhorn_unbalanced( Returns ------- - if n_hists == 1: + if n_hists == 0: - gamma : array-like, shape(dim_a, dim_b) Optimal transportation matrix for the given parameters - log : dict @@ -129,6 +171,18 @@ def sinkhorn_unbalanced( - log : dict log dictionary returned only if `log` is `True` + Notes + ----- + When `log=True`, the returned dictionary is the one of the solver selected by + `method`, see :any:`ot.unbalanced.sinkhorn_knopp_unbalanced`, + :any:`ot.unbalanced.sinkhorn_stabilized_unbalanced` and + :any:`ot.unbalanced.sinkhorn_unbalanced_translation_invariant`. + + .. note:: + When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, + only the negative entropy regularization is implemented: `reg_type` and `c` are + ignored, and only the linear cost is returned. + Examples -------- @@ -274,7 +328,7 @@ def sinkhorn_unbalanced2( .. math:: \min_\gamma \quad \langle \gamma, \mathbf{M} \rangle_F + - \mathrm{reg} \cdot \mathrm{KL}(\gamma, \mathbf{c}) + + \mathrm{reg} \cdot \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) + \mathrm{reg_{m1}} \cdot \mathrm{KL}(\gamma \mathbf{1}, \mathbf{a}) + \mathrm{reg_{m2}} \cdot \mathrm{KL}(\gamma^T \mathbf{1}, \mathbf{b}) @@ -285,7 +339,17 @@ def sinkhorn_unbalanced2( - :math:`\mathbf{M}` is the (`dim_a`, `dim_b`) metric cost matrix - :math:`\mathbf{a}` and :math:`\mathbf{b}` are source and target unbalanced distributions - :math:`\mathbf{c}` is a reference distribution for the regularization - - KL is the Kullback-Leibler divergence + - :math:`\mathrm{KL}` is the generalized Kullback-Leibler divergence + :math:`\mathrm{KL}(\mathbf{P}, \mathbf{Q}) = \sum_{i,j} \mathbf{P}_{i,j} \log(\mathbf{P}_{i,j} / \mathbf{Q}_{i,j}) - \mathbf{P}_{i,j} + \mathbf{Q}_{i,j}` + + and :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})` is the regularizer + + .. math:: + \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) = + \begin{cases} + \mathrm{KL}(\gamma, \mathbf{c}) & \text{if } \texttt{reg\_type} = \text{'kl'}, \\ + \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j} & \text{if } \texttt{reg\_type} = \text{'entropy'}. + \end{cases} The algorithm used for solving the problem is the generalized Sinkhorn-Knopp matrix scaling algorithm as proposed in :ref:`[10, 25] @@ -324,13 +388,19 @@ def sinkhorn_unbalanced2( method used for the solver either 'sinkhorn', 'sinkhorn_stabilized', 'sinkhorn_translation_invariant' or 'sinkhorn_reg_scaling', see those function for specific parameters reg_type : string, optional - Regularizer term. Can take two values: + Name of the regularizer :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})`, + case insensitive. Can take two values: - Negative entropy: 'entropy': :math:`\Omega(\gamma) = \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j}`. - This is equivalent (up to a constant) to :math:`\Omega(\gamma) = \text{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)`. - - Kullback-Leibler divergence: 'kl': - :math:`\Omega(\gamma) = \text{KL}(\gamma, \mathbf{a} \mathbf{b}^T)`. + It is equal to :math:`\mathrm{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)` minus the + constant :math:`dim_a \times dim_b`, and the reference measure :math:`\mathbf{c}` + is overwritten by the all-ones matrix. + - Generalized Kullback-Leibler divergence (default): 'kl': + :math:`\Omega(\gamma, \mathbf{c}) = \mathrm{KL}(\gamma, \mathbf{c})` with + :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T` when `c` is None. + + Any other value raises a :class:`ValueError`. c : array-like, shape (dim_a, dim_b), optional (default=None) Reference measure for the regularization. If None, then use :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T`. @@ -358,6 +428,18 @@ def sinkhorn_unbalanced2( log : dict log dictionary returned only if `log` is `True` + Notes + ----- + When `log=True`, the returned dictionary is the one of the solver selected by + `method`, see :any:`ot.unbalanced.sinkhorn_knopp_unbalanced`, + :any:`ot.unbalanced.sinkhorn_stabilized_unbalanced` and + :any:`ot.unbalanced.sinkhorn_unbalanced_translation_invariant`. + + .. note:: + When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, + only the negative entropy regularization is implemented: `reg_type` and `c` are + ignored, and only the linear cost is returned. + Examples -------- @@ -487,8 +569,13 @@ def sinkhorn_unbalanced2( return cost else: - if reg_type == "kl": - warnings.warn("Reg_type not implemented yet. Use entropy.") + if returnCost not in ("linear", "total"): + raise ValueError("Unknown returnCost = {}".format(returnCost)) + if returnCost != "linear": + warnings.warn( + "returnCost='total' is not available with multiple histograms " + "(n_hists > 1): the linear cost is returned." + ) if method.lower() == "sinkhorn": return sinkhorn_knopp_unbalanced( @@ -585,7 +672,7 @@ def sinkhorn_knopp_unbalanced( .. math:: W = \arg \min_\gamma \quad \langle \gamma, \mathbf{M} \rangle_F + - \mathrm{reg} \cdot \mathrm{KL}(\gamma, \mathbf{c}) + + \mathrm{reg} \cdot \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) + \mathrm{reg_{m1}} \cdot \mathrm{KL}(\gamma \mathbf{1}, \mathbf{a}) + \mathrm{reg_{m2}} \cdot \mathrm{KL}(\gamma^T \mathbf{1}, \mathbf{b}) @@ -597,7 +684,17 @@ def sinkhorn_knopp_unbalanced( - :math:`\mathbf{M}` is the (`dim_a`, `dim_b`) metric cost matrix - :math:`\mathbf{a}` and :math:`\mathbf{b}` are source and target unbalanced distributions - :math:`\mathbf{c}` is a reference distribution for the regularization - - KL is the Kullback-Leibler divergence + - :math:`\mathrm{KL}` is the generalized Kullback-Leibler divergence + :math:`\mathrm{KL}(\mathbf{P}, \mathbf{Q}) = \sum_{i,j} \mathbf{P}_{i,j} \log(\mathbf{P}_{i,j} / \mathbf{Q}_{i,j}) - \mathbf{P}_{i,j} + \mathbf{Q}_{i,j}` + + and :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})` is the regularizer + + .. math:: + \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) = + \begin{cases} + \mathrm{KL}(\gamma, \mathbf{c}) & \text{if } \texttt{reg\_type} = \text{'kl'}, \\ + \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j} & \text{if } \texttt{reg\_type} = \text{'entropy'}. + \end{cases} The algorithm used for solving the problem is the generalized Sinkhorn-Knopp matrix scaling algorithm as proposed in :ref:`[10, 25] ` @@ -632,13 +729,19 @@ def sinkhorn_knopp_unbalanced( If :math:`\mathrm{reg_{m}}` is an array, it must have the same backend as input arrays `(a, b, M)`. reg_type : string, optional - Regularizer term. Can take two values: + Name of the regularizer :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})`, + case insensitive. Can take two values: - Negative entropy: 'entropy': :math:`\Omega(\gamma) = \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j}`. - This is equivalent (up to a constant) to :math:`\Omega(\gamma) = \text{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)`. - - Kullback-Leibler divergence: 'kl': - :math:`\Omega(\gamma) = \text{KL}(\gamma, \mathbf{a} \mathbf{b}^T)`. + It is equal to :math:`\mathrm{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)` minus the + constant :math:`dim_a \times dim_b`, and the reference measure :math:`\mathbf{c}` + is overwritten by the all-ones matrix. + - Generalized Kullback-Leibler divergence (default): 'kl': + :math:`\Omega(\gamma, \mathbf{c}) = \mathrm{KL}(\gamma, \mathbf{c})` with + :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T` when `c` is None. + + Any other value raises a :class:`ValueError`. c : array-like, shape (dim_a, dim_b), optional (default=None) Reference measure for the regularization. If None, then use :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T`. @@ -658,7 +761,7 @@ def sinkhorn_knopp_unbalanced( Returns ------- - if n_hists == 1: + if n_hists == 0: - gamma : array-like, shape (dim_a, dim_b) Optimal transportation matrix for the given parameters - log : dict @@ -669,6 +772,24 @@ def sinkhorn_knopp_unbalanced( - log : dict log dictionary returned only if `log` is `True` + Notes + ----- + The `log` dictionary returned when `log=True` contains: + + - 'err' : list of float, the error at each iteration; + - 'logu', 'logv' : array-like, the log of the scaling vectors; + - 'cost' : float, the linear cost :math:`\langle \gamma, \mathbf{M} \rangle_F`; + - 'total_cost' : float, the value of the optimization problem above. + + 'cost' and 'total_cost' are only computed when `b` is a single histogram. + + .. note:: + When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, + the function returns the cost of each column and **only the negative entropy + regularization is implemented**: `reg_type` and `c` are ignored and the + reference measure is the all-ones matrix. In that case `log` only contains + 'err', 'logu' and 'logv'. + Examples -------- @@ -702,6 +823,8 @@ def sinkhorn_knopp_unbalanced( M, a, b = list_to_array(M, a, b) nx = get_backend(M, a, b) + reg_type = _check_reg_type(reg_type) + dim_a, dim_b = M.shape if len(a) == 0: @@ -714,6 +837,13 @@ def sinkhorn_knopp_unbalanced( else: n_hists = 0 + if n_hists and (reg_type != "entropy" or c is not None): + warnings.warn( + "With multiple histograms (n_hists > 1) only the negative entropy " + "regularization is implemented: reg_type and c are ignored and the " + "reference measure is the all-ones matrix." + ) + reg_m1, reg_m2 = get_parameter_pair(reg_m) if log: @@ -732,10 +862,12 @@ def sinkhorn_knopp_unbalanced( else: u, v = nx.exp(warmstart[0]), nx.exp(warmstart[1]) - if reg_type.lower() == "entropy": - warnings.warn( - "If reg_type = entropy, then the matrix c is overwritten by the one matrix." - ) + if reg_type == "entropy": + if c is not None: + warnings.warn( + "reg_type='entropy' ignores the provided c: the reference measure " + "of the regularization is the all-ones matrix." + ) c = nx.ones((dim_a, dim_b), type_as=M) if n_hists: @@ -806,8 +938,17 @@ def sinkhorn_knopp_unbalanced( linear_cost = nx.sum(plan * M) dict_log["cost"] = linear_cost - # mass=True: the penalization is the generalized KL divergence - total_cost = linear_cost + reg * nx.kl_div(plan, c, mass=True) + # The regularizer is the generalized KL divergence + # KL(plan, c) = sum(plan * log(plan / c) - plan + c). + # With reg_type="entropy" the reference measure is c = 1 and the + # regularizer is the negative entropy + # Omega(plan) = sum(plan * log(plan) - plan) + # = KL(plan, 1) - dim_a * dim_b, + # so the constant dim_a * dim_b must be removed. + reg_cost = nx.kl_div(plan, c, mass=True) + if reg_type == "entropy": + reg_cost = reg_cost - dim_a * dim_b + total_cost = linear_cost + reg * reg_cost if reg_m1 != float("inf"): total_cost = total_cost + reg_m1 * nx.kl_div( nx.sum(plan, 1), a, mass=True @@ -848,7 +989,7 @@ def sinkhorn_stabilized_unbalanced( .. math:: W = \arg \min_\gamma \quad \langle \gamma, \mathbf{M} \rangle_F + - \mathrm{reg} \cdot \mathrm{KL}(\gamma, \mathbf{c}) + + \mathrm{reg} \cdot \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) + \mathrm{reg_{m1}} \cdot \mathrm{KL}(\gamma \mathbf{1}, \mathbf{a}) + \mathrm{reg_{m2}} \cdot \mathrm{KL}(\gamma^T \mathbf{1}, \mathbf{b}) @@ -860,7 +1001,17 @@ def sinkhorn_stabilized_unbalanced( - :math:`\mathbf{M}` is the (`dim_a`, `dim_b`) metric cost matrix - :math:`\mathbf{a}` and :math:`\mathbf{b}` are source and target unbalanced distributions - :math:`\mathbf{c}` is a reference distribution for the regularization - - KL is the Kullback-Leibler divergence + - :math:`\mathrm{KL}` is the generalized Kullback-Leibler divergence + :math:`\mathrm{KL}(\mathbf{P}, \mathbf{Q}) = \sum_{i,j} \mathbf{P}_{i,j} \log(\mathbf{P}_{i,j} / \mathbf{Q}_{i,j}) - \mathbf{P}_{i,j} + \mathbf{Q}_{i,j}` + + and :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})` is the regularizer + + .. math:: + \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) = + \begin{cases} + \mathrm{KL}(\gamma, \mathbf{c}) & \text{if } \texttt{reg\_type} = \text{'kl'}, \\ + \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j} & \text{if } \texttt{reg\_type} = \text{'entropy'}. + \end{cases} The algorithm used for solving the problem is the generalized Sinkhorn-Knopp matrix scaling algorithm as proposed in :ref:`[10, 25] ` @@ -895,13 +1046,19 @@ def sinkhorn_stabilized_unbalanced( method used for the solver either 'sinkhorn', 'sinkhorn_stabilized' or 'sinkhorn_reg_scaling', see those function for specific parameters reg_type : string, optional - Regularizer term. Can take two values: + Name of the regularizer :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})`, + case insensitive. Can take two values: - Negative entropy: 'entropy': :math:`\Omega(\gamma) = \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j}`. - This is equivalent (up to a constant) to :math:`\Omega(\gamma) = \text{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)`. - - Kullback-Leibler divergence: 'kl': - :math:`\Omega(\gamma) = \text{KL}(\gamma, \mathbf{a} \mathbf{b}^T)`. + It is equal to :math:`\mathrm{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)` minus the + constant :math:`dim_a \times dim_b`, and the reference measure :math:`\mathbf{c}` + is overwritten by the all-ones matrix. + - Generalized Kullback-Leibler divergence (default): 'kl': + :math:`\Omega(\gamma, \mathbf{c}) = \mathrm{KL}(\gamma, \mathbf{c})` with + :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T` when `c` is None. + + Any other value raises a :class:`ValueError`. c : array-like, shape (dim_a, dim_b), optional (default=None) Reference measure for the regularization. If None, then use :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T`. @@ -927,7 +1084,7 @@ def sinkhorn_stabilized_unbalanced( Returns ------- - if n_hists == 1: + if n_hists == 0: - gamma : array-like, shape (dim_a, dim_b) Optimal transportation matrix for the given parameters - log : dict @@ -937,6 +1094,24 @@ def sinkhorn_stabilized_unbalanced( the OT cost between :math:`\mathbf{a}` and each of the histograms :math:`\mathbf{b}_i` - log : dict log dictionary returned only if `log` is `True` + Notes + ----- + The `log` dictionary returned when `log=True` contains: + + - 'err' : list of float, the error at each iteration; + - 'logu', 'logv' : array-like, the log of the scaling vectors; + - 'cost' : float, the linear cost :math:`\langle \gamma, \mathbf{M} \rangle_F`; + - 'total_cost' : float, the value of the optimization problem above. + + 'cost' and 'total_cost' are only computed when `b` is a single histogram. + + .. note:: + When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, + the function returns the cost of each column and **only the negative entropy + regularization is implemented**: `reg_type` and `c` are ignored and the + reference measure is the all-ones matrix. In that case `log` only contains + 'err', 'logu' and 'logv'. + Examples -------- @@ -969,6 +1144,8 @@ def sinkhorn_stabilized_unbalanced( a, b, M = list_to_array(a, b, M) nx = get_backend(M, a, b) + reg_type = _check_reg_type(reg_type) + dim_a, dim_b = M.shape if len(a) == 0: @@ -981,6 +1158,13 @@ def sinkhorn_stabilized_unbalanced( else: n_hists = 0 + if n_hists and (reg_type != "entropy" or c is not None): + warnings.warn( + "With multiple histograms (n_hists > 1) only the negative entropy " + "regularization is implemented: reg_type and c are ignored and the " + "reference measure is the all-ones matrix." + ) + reg_m1, reg_m2 = get_parameter_pair(reg_m) if log: @@ -1000,9 +1184,11 @@ def sinkhorn_stabilized_unbalanced( u, v = nx.exp(warmstart[0]), nx.exp(warmstart[1]) if reg_type == "entropy": - warnings.warn( - "If reg_type = entropy, then the matrix c is overwritten by the one matrix." - ) + if c is not None: + warnings.warn( + "reg_type='entropy' ignores the provided c: the reference measure " + "of the regularization is the all-ones matrix." + ) c = nx.ones((dim_a, dim_b), type_as=M) if n_hists: @@ -1111,8 +1297,17 @@ def sinkhorn_stabilized_unbalanced( linear_cost = nx.sum(plan * M) dict_log["cost"] = linear_cost - # mass=True: the penalization is the generalized KL divergence - total_cost = linear_cost + reg * nx.kl_div(plan, c, mass=True) + # The regularizer is the generalized KL divergence + # KL(plan, c) = sum(plan * log(plan / c) - plan + c). + # With reg_type="entropy" the reference measure is c = 1 and the + # regularizer is the negative entropy + # Omega(plan) = sum(plan * log(plan) - plan) + # = KL(plan, 1) - dim_a * dim_b, + # so the constant dim_a * dim_b must be removed. + reg_cost = nx.kl_div(plan, c, mass=True) + if reg_type == "entropy": + reg_cost = reg_cost - dim_a * dim_b + total_cost = linear_cost + reg * reg_cost if reg_m1 != float("inf"): total_cost = total_cost + reg_m1 * nx.kl_div( nx.sum(plan, 1), a, mass=True @@ -1151,7 +1346,7 @@ def sinkhorn_unbalanced_translation_invariant( .. math:: W = \arg \min_\gamma \ \langle \gamma, \mathbf{M} \rangle_F + - \mathrm{reg} \cdot \mathrm{KL}(\gamma, \mathbf{c}) + + \mathrm{reg} \cdot \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) + \mathrm{reg_{m1}} \cdot \mathrm{KL}(\gamma \mathbf{1}, \mathbf{a}) + \mathrm{reg_{m2}} \cdot \mathrm{KL}(\gamma^T \mathbf{1}, \mathbf{b}) @@ -1163,7 +1358,17 @@ def sinkhorn_unbalanced_translation_invariant( - :math:`\mathbf{M}` is the (`dim_a`, `dim_b`) metric cost matrix - :math:`\Omega` is the entropic regularization term,KL divergence - :math:`\mathbf{a}` and :math:`\mathbf{b}` are source and target unbalanced distributions - - KL is the Kullback-Leibler divergence + - :math:`\mathrm{KL}` is the generalized Kullback-Leibler divergence + :math:`\mathrm{KL}(\mathbf{P}, \mathbf{Q}) = \sum_{i,j} \mathbf{P}_{i,j} \log(\mathbf{P}_{i,j} / \mathbf{Q}_{i,j}) - \mathbf{P}_{i,j} + \mathbf{Q}_{i,j}` + + and :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})` is the regularizer + + .. math:: + \Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c}) = + \begin{cases} + \mathrm{KL}(\gamma, \mathbf{c}) & \text{if } \texttt{reg\_type} = \text{'kl'}, \\ + \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j} & \text{if } \texttt{reg\_type} = \text{'entropy'}. + \end{cases} The algorithm used for solving the problem is the translation invariant Sinkhorn algorithm as proposed in :ref:`[73] ` @@ -1187,11 +1392,19 @@ def sinkhorn_unbalanced_translation_invariant( `reg_m=(float("inf"), scalar)` or `reg_m=(scalar, float("inf"))`. If reg_m is an array, it must have the same backend as input arrays (a, b, M). reg_type : string, optional - Regularizer term. Can take two values: - 'entropy' (negative entropy) - :math:`\Omega(\gamma) = \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j}`, or - 'kl' (Kullback-Leibler) - :math:`\Omega(\gamma) = \text{KL}(\gamma, \mathbf{a} \mathbf{b}^T)`. + Name of the regularizer :math:`\Omega_{\texttt{reg\_type}}(\gamma, \mathbf{c})`, + case insensitive. Can take two values: + + - Negative entropy: 'entropy': + :math:`\Omega(\gamma) = \sum_{i,j} \gamma_{i,j} \log(\gamma_{i,j}) - \sum_{i,j} \gamma_{i,j}`. + It is equal to :math:`\mathrm{KL}(\gamma, 1_{dim_a} 1_{dim_b}^T)` minus the + constant :math:`dim_a \times dim_b`, and the reference measure :math:`\mathbf{c}` + is overwritten by the all-ones matrix. + - Generalized Kullback-Leibler divergence (default): 'kl': + :math:`\Omega(\gamma, \mathbf{c}) = \mathrm{KL}(\gamma, \mathbf{c})` with + :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T` when `c` is None. + + Any other value raises a :class:`ValueError`. c : array-like, shape (dim_a, dim_b), optional (default=None) Reference measure for the regularization. If None, then use :math:`\mathbf{c} = \mathbf{a} \mathbf{b}^T`. @@ -1211,7 +1424,7 @@ def sinkhorn_unbalanced_translation_invariant( Returns ------- - if n_hists == 1: + if n_hists == 0: - gamma : array-like, shape (dim_a, dim_b) Optimal transportation matrix for the given parameters - log : dict @@ -1222,6 +1435,24 @@ def sinkhorn_unbalanced_translation_invariant( - log : dict log dictionary returned only if `log` is `True` + Notes + ----- + The `log` dictionary returned when `log=True` contains: + + - 'err' : list of float, the error at each iteration; + - 'logu', 'logv' : array-like, the log of the scaling vectors; + - 'cost' : float, the linear cost :math:`\langle \gamma, \mathbf{M} \rangle_F`; + - 'total_cost' : float, the value of the optimization problem above. + + 'cost' and 'total_cost' are only computed when `b` is a single histogram. + + .. note:: + When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, + the function returns the cost of each column and **only the negative entropy + regularization is implemented**: `reg_type` and `c` are ignored and the + reference measure is the all-ones matrix. In that case `log` only contains + 'err', 'logu' and 'logv'. + Examples -------- @@ -1245,6 +1476,8 @@ def sinkhorn_unbalanced_translation_invariant( M, a, b = list_to_array(M, a, b) nx = get_backend(M, a, b) + reg_type = _check_reg_type(reg_type) + dim_a, dim_b = M.shape if len(a) == 0: @@ -1257,6 +1490,13 @@ def sinkhorn_unbalanced_translation_invariant( else: n_hists = 0 + if n_hists and (reg_type != "entropy" or c is not None): + warnings.warn( + "With multiple histograms (n_hists > 1) only the negative entropy " + "regularization is implemented: reg_type and c are ignored and the " + "reference measure is the all-ones matrix." + ) + reg_m1, reg_m2 = get_parameter_pair(reg_m) if log: @@ -1278,9 +1518,11 @@ def sinkhorn_unbalanced_translation_invariant( u_, v_ = u, v if reg_type == "entropy": - warnings.warn( - "If reg_type = entropy, then the matrix c is overwritten by the one matrix." - ) + if c is not None: + warnings.warn( + "reg_type='entropy' ignores the provided c: the reference measure " + "of the regularization is the all-ones matrix." + ) c = nx.ones((dim_a, dim_b), type_as=M) if n_hists: @@ -1399,8 +1641,17 @@ def sinkhorn_unbalanced_translation_invariant( linear_cost = nx.sum(plan * M) dict_log["cost"] = linear_cost - # mass=True: the penalization is the generalized KL divergence - total_cost = linear_cost + reg * nx.kl_div(plan, c, mass=True) + # The regularizer is the generalized KL divergence + # KL(plan, c) = sum(plan * log(plan / c) - plan + c). + # With reg_type="entropy" the reference measure is c = 1 and the + # regularizer is the negative entropy + # Omega(plan) = sum(plan * log(plan) - plan) + # = KL(plan, 1) - dim_a * dim_b, + # so the constant dim_a * dim_b must be removed. + reg_cost = nx.kl_div(plan, c, mass=True) + if reg_type == "entropy": + reg_cost = reg_cost - dim_a * dim_b + total_cost = linear_cost + reg * reg_cost if reg_m1 != float("inf"): total_cost = total_cost + reg_m1 * nx.kl_div( nx.sum(plan, 1), a, mass=True diff --git a/test/test_solvers.py b/test/test_solvers.py index 0d73cd6a8..6cc8137cc 100644 --- a/test/test_solvers.py +++ b/test/test_solvers.py @@ -384,40 +384,6 @@ 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 844cf200b..00a493240 100644 --- a/test/unbalanced/test_sinkhorn.py +++ b/test/unbalanced/test_sinkhorn.py @@ -7,6 +7,8 @@ # License: MIT License import itertools +import warnings + import numpy as np import ot import pytest @@ -811,27 +813,51 @@ def test_implemented_methods(nx): 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) +METHODS_UNBALANCED = [ + "sinkhorn", + "sinkhorn_stabilized", + "sinkhorn_translation_invariant", +] - 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 +def _unbalanced_problem(nx, n=6, m=7, seed=0): + rng = np.random.RandomState(seed) + a = rng.rand(n) + a /= a.sum() + b = rng.rand(m) + b /= b.sum() + M = rng.rand(n, m) + M /= M.max() + return nx.from_numpy(a, b, M) + + +def _gen_kl(P, Q): + r"""Generalized KL divergence, computed with plain numpy. + + Deliberately independent from ``nx.kl_div`` so that the reference value does + not reuse the implementation under test. + """ + P = np.asarray(P, dtype=np.float64) + Q = np.asarray(Q, dtype=np.float64) + return float(np.sum(P * np.log(P / Q) - P + Q)) + + +def _negative_entropy(P): + r"""Negative entropy :math:`\sum_{ij} P_{ij} \log(P_{ij}) - P_{ij}`.""" + P = np.asarray(P, dtype=np.float64) + return float(np.sum(P * np.log(P) - P)) + + +@pytest.mark.parametrize("method", METHODS_UNBALANCED) +def test_unbalanced_total_cost_kl(nx, method): + """Non-regression test for the generalized KL objective (PR #874). + + The penalizations of the unbalanced OT problem are generalized KL + divergences, so the reported total cost must be at least the linear cost and + must match the objective recomputed from the returned plan. + """ + a, b, M = _unbalanced_problem(nx) + reg, reg_m = 0.3, 1.0 G, log = ot.unbalanced.sinkhorn_unbalanced( a, @@ -840,22 +866,180 @@ def test_unbalanced_total_cost(nx, method): reg=reg, reg_m=reg_m, method=method, + reg_type="kl", + log=True, numItermax=5000, - stopThr=1e-12, + stopThr=1e-13, + ) + + G_np = nx.to_numpy(G) + a_np = nx.to_numpy(a) + b_np = nx.to_numpy(b) + M_np = nx.to_numpy(M) + c_np = a_np[:, None] * b_np[None, :] + + expected = ( + float(np.sum(G_np * M_np)) + + reg * _gen_kl(G_np, c_np) + + reg_m * _gen_kl(G_np.sum(1), a_np) + + reg_m * _gen_kl(G_np.sum(0), b_np) + ) + + total_cost = float(nx.to_numpy(log["total_cost"])) + np.testing.assert_allclose(total_cost, expected, atol=1e-6) + # with reg_type="kl" every penalization is a divergence, hence non negative + np.testing.assert_array_less(float(nx.to_numpy(log["cost"])) - 1e-5, total_cost) + + +@pytest.mark.parametrize("method", METHODS_UNBALANCED) +def test_unbalanced_total_cost_entropy(nx, method): + """The entropy regularizer must not carry the constant ``dim_a * dim_b``. + + The reported value must be the objective built on + :math:`\\Omega(\\gamma) = \\sum \\gamma \\log \\gamma - \\gamma`, which is + ``KL(gamma, 1) - dim_a * dim_b``. It may be negative, so no comparison with + the linear cost is made here. + """ + a, b, M = _unbalanced_problem(nx) + reg, reg_m = 0.3, 1.0 + + G, log = ot.unbalanced.sinkhorn_unbalanced( + a, + b, + M, + reg=reg, + reg_m=reg_m, + method=method, + reg_type="entropy", log=True, + numItermax=5000, + stopThr=1e-13, ) - 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) + G_np = nx.to_numpy(G) + a_np = nx.to_numpy(a) + b_np = nx.to_numpy(b) + M_np = nx.to_numpy(M) - # 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"]) + expected = ( + float(np.sum(G_np * M_np)) + + reg * _negative_entropy(G_np) + + reg_m * _gen_kl(G_np.sum(1), a_np) + + reg_m * _gen_kl(G_np.sum(0), b_np) ) - np.testing.assert_allclose( - nx.to_numpy(log["total_cost"]), nx.to_numpy(expected), atol=1e-6 + total_cost = float(nx.to_numpy(log["total_cost"])) + np.testing.assert_allclose(total_cost, expected, atol=1e-6) + + # the constant reg * dim_a * dim_b must have been removed + without_constant_removal = expected + reg * G_np.size + assert abs(total_cost - without_constant_removal) > 1e-3 + + +@pytest.mark.parametrize("method", METHODS_UNBALANCED) +def test_unbalanced_entropy_constant(nx, method): + """The entropy constant is exactly ``reg * dim_a * dim_b``. + + With ``c`` set to the all-ones matrix, ``reg_type="kl"`` and + ``reg_type="entropy"`` minimize the same objective up to that constant, so + the plans must coincide and the two reported total costs must differ by + exactly ``reg * dim_a * dim_b``. + """ + a, b, M = _unbalanced_problem(nx) + reg, reg_m = 0.3, 1.0 + n, m = nx.to_numpy(M).shape + c_ones = nx.from_numpy(np.ones((n, m))) + + G_kl, log_kl = ot.unbalanced.sinkhorn_unbalanced( + a, + b, + M, + reg=reg, + reg_m=reg_m, + method=method, + reg_type="kl", + c=c_ones, + log=True, + numItermax=5000, + stopThr=1e-13, + ) + G_en, log_en = ot.unbalanced.sinkhorn_unbalanced( + a, + b, + M, + reg=reg, + reg_m=reg_m, + method=method, + reg_type="entropy", + log=True, + numItermax=5000, + stopThr=1e-13, ) + + # same objective up to a constant -> same plan + np.testing.assert_allclose(nx.to_numpy(G_kl), nx.to_numpy(G_en), atol=1e-8) + diff = float(nx.to_numpy(log_kl["total_cost"])) - float( + nx.to_numpy(log_en["total_cost"]) + ) + np.testing.assert_allclose(diff, reg * n * m, rtol=1e-6) + + +@pytest.mark.parametrize("method", METHODS_UNBALANCED) +@pytest.mark.parametrize("reg_type", ["kl", "entropy"]) +def test_unbalanced_reg_type_case_insensitive(nx, method, reg_type): + """`reg_type` must be handled case insensitively by every method.""" + a, b, M = _unbalanced_problem(nx) + reg, reg_m = 0.3, 1.0 + + ref_plan, ref_cost = None, None + for variant in [reg_type, reg_type.capitalize(), reg_type.upper()]: + G, log = ot.unbalanced.sinkhorn_unbalanced( + a, + b, + M, + reg=reg, + reg_m=reg_m, + method=method, + reg_type=variant, + log=True, + numItermax=5000, + stopThr=1e-13, + ) + plan = nx.to_numpy(G) + cost = float(nx.to_numpy(log["total_cost"])) + if ref_plan is None: + ref_plan, ref_cost = plan, cost + else: + np.testing.assert_allclose(plan, ref_plan, atol=1e-10) + np.testing.assert_allclose(cost, ref_cost, atol=1e-10) + + +@pytest.mark.parametrize("method", METHODS_UNBALANCED) +def test_unbalanced_unknown_reg_type_raises(nx, method): + """An unknown `reg_type` must be rejected instead of silently using 'kl'.""" + a, b, M = _unbalanced_problem(nx) + with pytest.raises(ValueError, match="Unknown reg_type"): + ot.unbalanced.sinkhorn_unbalanced( + a, b, M, reg=0.3, reg_m=1.0, method=method, reg_type="cryptic divergence" + ) + + +def test_unbalanced_multiple_inputs_reg_type_warns(nx): + """Multi histogram mode only implements the negative entropy regularization. + + `reg_type` and `c` are ignored in that mode, which must be advertised with a + warning. Asking for 'entropy' (the implemented regularizer) must stay silent. + """ + a, b, M = _unbalanced_problem(nx) + n_hists = 3 + B = nx.from_numpy(np.random.RandomState(1).rand(nx.to_numpy(M).shape[1], n_hists)) + + with pytest.warns(UserWarning, match="multiple histograms"): + ot.unbalanced.sinkhorn_unbalanced( + a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="kl" + ) + + with warnings.catch_warnings(): + warnings.simplefilter("error") + ot.unbalanced.sinkhorn_unbalanced( + a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy" + ) From b0fd7d4dfc8c401c596cef01177372534f68c660 Mon Sep 17 00:00:00 2001 From: B Date: Fri, 9 Oct 2026 11:37:35 +0800 Subject: [PATCH 4/6] Refine the 2d-b warning and its message - the warning now reports the actual n_hists instead of "n_hists > 1", so it is also accurate for a 2d b with a single column, which takes the same branch and returns a cost array instead of a plan; - passing an explicit c together with reg_type="entropy" in 2d mode no longer emits two warnings for the same thing; - the docstring Notes state that any 2d b selects that branch, including n_hists == 1; - the multi-histogram test now asserts the number of warnings and the reported n_hists, and covers the (dim_b, 1) case. --- ot/unbalanced/_sinkhorn.py | 78 ++++++++++++++++++-------------- test/unbalanced/test_sinkhorn.py | 45 +++++++++++++----- 2 files changed, 77 insertions(+), 46 deletions(-) diff --git a/ot/unbalanced/_sinkhorn.py b/ot/unbalanced/_sinkhorn.py index 33d5bd1ff..c9890bd83 100644 --- a/ot/unbalanced/_sinkhorn.py +++ b/ot/unbalanced/_sinkhorn.py @@ -179,9 +179,10 @@ def sinkhorn_unbalanced( :any:`ot.unbalanced.sinkhorn_unbalanced_translation_invariant`. .. note:: - When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, - only the negative entropy regularization is implemented: `reg_type` and `c` are - ignored, and only the linear cost is returned. + When `b` is a 2d array of shape (`dim_b`, `n_hists`) -- which includes + :math:`n_{hists} = 1` -- only the negative entropy regularization is + implemented: `reg_type` and `c` are ignored, and only the linear cost is + returned. Examples -------- @@ -436,9 +437,10 @@ def sinkhorn_unbalanced2( :any:`ot.unbalanced.sinkhorn_unbalanced_translation_invariant`. .. note:: - When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, - only the negative entropy regularization is implemented: `reg_type` and `c` are - ignored, and only the linear cost is returned. + When `b` is a 2d array of shape (`dim_b`, `n_hists`) -- which includes + :math:`n_{hists} = 1` -- only the negative entropy regularization is + implemented: `reg_type` and `c` are ignored, and only the linear cost is + returned. Examples -------- @@ -573,8 +575,8 @@ def sinkhorn_unbalanced2( raise ValueError("Unknown returnCost = {}".format(returnCost)) if returnCost != "linear": warnings.warn( - "returnCost='total' is not available with multiple histograms " - "(n_hists > 1): the linear cost is returned." + "returnCost='total' is not available with a 2d b (n_hists={}): " + "the linear cost is returned.".format(n_hists) ) if method.lower() == "sinkhorn": @@ -784,11 +786,12 @@ def sinkhorn_knopp_unbalanced( 'cost' and 'total_cost' are only computed when `b` is a single histogram. .. note:: - When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, - the function returns the cost of each column and **only the negative entropy - regularization is implemented**: `reg_type` and `c` are ignored and the - reference measure is the all-ones matrix. In that case `log` only contains - 'err', 'logu' and 'logv'. + When `b` is a 2d array of shape (`dim_b`, `n_hists`) -- which includes + :math:`n_{hists} = 1`, since any 2d `b` selects this branch -- the function + returns the cost of each column instead of a plan, and **only the negative + entropy regularization is implemented**: `reg_type` and `c` are ignored and + the reference measure is the all-ones matrix. In that case `log` only + contains 'err', 'logu' and 'logv'. Examples -------- @@ -839,9 +842,9 @@ def sinkhorn_knopp_unbalanced( if n_hists and (reg_type != "entropy" or c is not None): warnings.warn( - "With multiple histograms (n_hists > 1) only the negative entropy " - "regularization is implemented: reg_type and c are ignored and the " - "reference measure is the all-ones matrix." + "With a 2d b (n_hists={}) only the negative entropy regularization is " + "implemented: reg_type and c are ignored and the reference measure is " + "the all-ones matrix.".format(n_hists) ) reg_m1, reg_m2 = get_parameter_pair(reg_m) @@ -863,7 +866,8 @@ def sinkhorn_knopp_unbalanced( u, v = nx.exp(warmstart[0]), nx.exp(warmstart[1]) if reg_type == "entropy": - if c is not None: + # in 2d mode the warning above already reports that c is ignored + if c is not None and not n_hists: warnings.warn( "reg_type='entropy' ignores the provided c: the reference measure " "of the regularization is the all-ones matrix." @@ -1106,11 +1110,12 @@ def sinkhorn_stabilized_unbalanced( 'cost' and 'total_cost' are only computed when `b` is a single histogram. .. note:: - When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, - the function returns the cost of each column and **only the negative entropy - regularization is implemented**: `reg_type` and `c` are ignored and the - reference measure is the all-ones matrix. In that case `log` only contains - 'err', 'logu' and 'logv'. + When `b` is a 2d array of shape (`dim_b`, `n_hists`) -- which includes + :math:`n_{hists} = 1`, since any 2d `b` selects this branch -- the function + returns the cost of each column instead of a plan, and **only the negative + entropy regularization is implemented**: `reg_type` and `c` are ignored and + the reference measure is the all-ones matrix. In that case `log` only + contains 'err', 'logu' and 'logv'. Examples -------- @@ -1160,9 +1165,9 @@ def sinkhorn_stabilized_unbalanced( if n_hists and (reg_type != "entropy" or c is not None): warnings.warn( - "With multiple histograms (n_hists > 1) only the negative entropy " - "regularization is implemented: reg_type and c are ignored and the " - "reference measure is the all-ones matrix." + "With a 2d b (n_hists={}) only the negative entropy regularization is " + "implemented: reg_type and c are ignored and the reference measure is " + "the all-ones matrix.".format(n_hists) ) reg_m1, reg_m2 = get_parameter_pair(reg_m) @@ -1184,7 +1189,8 @@ def sinkhorn_stabilized_unbalanced( u, v = nx.exp(warmstart[0]), nx.exp(warmstart[1]) if reg_type == "entropy": - if c is not None: + # in 2d mode the warning above already reports that c is ignored + if c is not None and not n_hists: warnings.warn( "reg_type='entropy' ignores the provided c: the reference measure " "of the regularization is the all-ones matrix." @@ -1447,11 +1453,12 @@ def sinkhorn_unbalanced_translation_invariant( 'cost' and 'total_cost' are only computed when `b` is a single histogram. .. note:: - When `b` is a 2d array of shape (`dim_b`, `n_hists`) with :math:`n_{hists} > 1`, - the function returns the cost of each column and **only the negative entropy - regularization is implemented**: `reg_type` and `c` are ignored and the - reference measure is the all-ones matrix. In that case `log` only contains - 'err', 'logu' and 'logv'. + When `b` is a 2d array of shape (`dim_b`, `n_hists`) -- which includes + :math:`n_{hists} = 1`, since any 2d `b` selects this branch -- the function + returns the cost of each column instead of a plan, and **only the negative + entropy regularization is implemented**: `reg_type` and `c` are ignored and + the reference measure is the all-ones matrix. In that case `log` only + contains 'err', 'logu' and 'logv'. Examples -------- @@ -1492,9 +1499,9 @@ def sinkhorn_unbalanced_translation_invariant( if n_hists and (reg_type != "entropy" or c is not None): warnings.warn( - "With multiple histograms (n_hists > 1) only the negative entropy " - "regularization is implemented: reg_type and c are ignored and the " - "reference measure is the all-ones matrix." + "With a 2d b (n_hists={}) only the negative entropy regularization is " + "implemented: reg_type and c are ignored and the reference measure is " + "the all-ones matrix.".format(n_hists) ) reg_m1, reg_m2 = get_parameter_pair(reg_m) @@ -1518,7 +1525,8 @@ def sinkhorn_unbalanced_translation_invariant( u_, v_ = u, v if reg_type == "entropy": - if c is not None: + # in 2d mode the warning above already reports that c is ignored + if c is not None and not n_hists: warnings.warn( "reg_type='entropy' ignores the provided c: the reference measure " "of the regularization is the all-ones matrix." diff --git a/test/unbalanced/test_sinkhorn.py b/test/unbalanced/test_sinkhorn.py index 00a493240..48741a847 100644 --- a/test/unbalanced/test_sinkhorn.py +++ b/test/unbalanced/test_sinkhorn.py @@ -1027,19 +1027,42 @@ def test_unbalanced_multiple_inputs_reg_type_warns(nx): """Multi histogram mode only implements the negative entropy regularization. `reg_type` and `c` are ignored in that mode, which must be advertised with a - warning. Asking for 'entropy' (the implemented regularizer) must stay silent. + single warning. Asking for 'entropy' (the implemented regularizer) with the + default reference measure must stay silent. """ a, b, M = _unbalanced_problem(nx) n_hists = 3 - B = nx.from_numpy(np.random.RandomState(1).rand(nx.to_numpy(M).shape[1], n_hists)) - - with pytest.warns(UserWarning, match="multiple histograms"): - ot.unbalanced.sinkhorn_unbalanced( - a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="kl" - ) - - with warnings.catch_warnings(): - warnings.simplefilter("error") + m = nx.to_numpy(M).shape[1] + B = nx.from_numpy(np.random.RandomState(1).rand(m, n_hists)) + c_custom = nx.from_numpy(np.full((nx.to_numpy(M).shape[0], m), 0.7)) + + def n_warnings(**kwargs): + with warnings.catch_warnings(record=True) as rec: + warnings.simplefilter("always") + ot.unbalanced.sinkhorn_unbalanced( + a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", **kwargs + ) + return [str(w.message) for w in rec if issubclass(w.category, UserWarning)] + + # the default reg_type is "kl", which is not implemented in that branch + msgs = n_warnings(reg_type="kl") + assert len(msgs) == 1, msgs + assert "2d b" in msgs[0] and "n_hists=3" in msgs[0] + + # "entropy" is the implemented regularizer and c is not given -> silent + assert n_warnings(reg_type="entropy") == [] + + # an explicit c is ignored too, but must not be reported twice + msgs = n_warnings(reg_type="entropy", c=c_custom) + assert len(msgs) == 1, msgs + + # a 2d b with a single column also takes that branch: the message must be exact + B1 = nx.from_numpy(np.random.RandomState(2).rand(m, 1)) + with warnings.catch_warnings(record=True) as rec: + warnings.simplefilter("always") ot.unbalanced.sinkhorn_unbalanced( - a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy" + a, B1, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="kl" ) + msgs = [str(w.message) for w in rec if issubclass(w.category, UserWarning)] + assert len(msgs) == 1, msgs + assert "n_hists=1" in msgs[0] From 1c0b9a0535a0e848e5b9e70ebacb85b999575c05 Mon Sep 17 00:00:00 2001 From: B Date: Fri, 9 Oct 2026 12:00:12 +0800 Subject: [PATCH 5/6] Fix undefined n_hists in the 2d-b returnCost warning The message rewritten for the 2d-b branch of sinkhorn_unbalanced2 called format(n_hists), but n_hists is not defined in that function, so sinkhorn_unbalanced2(a, b_2d, M, reg, reg_m, returnCost="total") raised NameError instead of warning. Use b.shape[1], which is the value the message is about. Caught by ruff (default rules, F821) run with the same version as the pre-commit hook in CI. Adds test_unbalanced_multiple_inputs_returnCost covering the three returnCost paths in 2d mode. --- ot/unbalanced/_sinkhorn.py | 2 +- test/unbalanced/test_sinkhorn.py | 35 ++++++++++++++++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/ot/unbalanced/_sinkhorn.py b/ot/unbalanced/_sinkhorn.py index c9890bd83..b32507cec 100644 --- a/ot/unbalanced/_sinkhorn.py +++ b/ot/unbalanced/_sinkhorn.py @@ -576,7 +576,7 @@ def sinkhorn_unbalanced2( if returnCost != "linear": warnings.warn( "returnCost='total' is not available with a 2d b (n_hists={}): " - "the linear cost is returned.".format(n_hists) + "the linear cost is returned.".format(b.shape[1]) ) if method.lower() == "sinkhorn": diff --git a/test/unbalanced/test_sinkhorn.py b/test/unbalanced/test_sinkhorn.py index 48741a847..9369d8ad5 100644 --- a/test/unbalanced/test_sinkhorn.py +++ b/test/unbalanced/test_sinkhorn.py @@ -1066,3 +1066,38 @@ def n_warnings(**kwargs): msgs = [str(w.message) for w in rec if issubclass(w.category, UserWarning)] assert len(msgs) == 1, msgs assert "n_hists=1" in msgs[0] + + +def test_unbalanced_multiple_inputs_returnCost(nx): + """`returnCost` must be honoured or reported, never silently ignored. + + With a 2d `b` only the linear cost is available. `sinkhorn_unbalanced2` used to + ignore `returnCost` entirely in that branch. + """ + a, b, M = _unbalanced_problem(nx) + m = nx.to_numpy(M).shape[1] + B = nx.from_numpy(np.random.RandomState(1).rand(m, 3)) + + # "linear" is the only supported value and must not warn + with warnings.catch_warnings(): + warnings.simplefilter("error") + loss = ot.unbalanced.sinkhorn_unbalanced2( + a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy", + returnCost="linear", + ) + assert nx.to_numpy(loss).shape == (3,) + + # "total" is not available there: warn and still return the linear cost + with pytest.warns(UserWarning, match="returnCost='total' is not available"): + loss = ot.unbalanced.sinkhorn_unbalanced2( + a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy", + returnCost="total", + ) + assert nx.to_numpy(loss).shape == (3,) + + # an invalid value raises, as in the single histogram branch + with pytest.raises(ValueError, match="Unknown returnCost"): + ot.unbalanced.sinkhorn_unbalanced2( + a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy", + returnCost="invalid", + ) From 9aaf3871f4e301189f85dc3ec1143f6426f14292 Mon Sep 17 00:00:00 2001 From: B Date: Fri, 9 Oct 2026 21:57:45 +0800 Subject: [PATCH 6/6] style: Apply ruff formatting to pass CI The ruff-format pre-commit hook reformatted test/unbalanced/test_sinkhorn.py on CI. The three calls to sinkhorn_unbalanced2 added in test_unbalanced_multiple_inputs_returnCost were written with a magic trailing comma on a single line, so the formatter explodes the arguments one per line. Formatting only: ast.dump() is identical (6597 nodes), the token stream is identical (6447 tokens, ignoring whitespace and comments), and the counts of assertions (50), numeric tolerances (42), parameter identifiers (403) and test functions (21) are unchanged. Ruff 0.5.2, the version pinned by .pre-commit-config.yaml. --- test/unbalanced/test_sinkhorn.py | 24 +++++++++++++++++++++--- 1 file changed, 21 insertions(+), 3 deletions(-) diff --git a/test/unbalanced/test_sinkhorn.py b/test/unbalanced/test_sinkhorn.py index 9369d8ad5..1088f4d03 100644 --- a/test/unbalanced/test_sinkhorn.py +++ b/test/unbalanced/test_sinkhorn.py @@ -1082,7 +1082,13 @@ def test_unbalanced_multiple_inputs_returnCost(nx): with warnings.catch_warnings(): warnings.simplefilter("error") loss = ot.unbalanced.sinkhorn_unbalanced2( - a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy", + a, + B, + M, + reg=0.3, + reg_m=1.0, + method="sinkhorn", + reg_type="entropy", returnCost="linear", ) assert nx.to_numpy(loss).shape == (3,) @@ -1090,7 +1096,13 @@ def test_unbalanced_multiple_inputs_returnCost(nx): # "total" is not available there: warn and still return the linear cost with pytest.warns(UserWarning, match="returnCost='total' is not available"): loss = ot.unbalanced.sinkhorn_unbalanced2( - a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy", + a, + B, + M, + reg=0.3, + reg_m=1.0, + method="sinkhorn", + reg_type="entropy", returnCost="total", ) assert nx.to_numpy(loss).shape == (3,) @@ -1098,6 +1110,12 @@ def test_unbalanced_multiple_inputs_returnCost(nx): # an invalid value raises, as in the single histogram branch with pytest.raises(ValueError, match="Unknown returnCost"): ot.unbalanced.sinkhorn_unbalanced2( - a, B, M, reg=0.3, reg_m=1.0, method="sinkhorn", reg_type="entropy", + a, + B, + M, + reg=0.3, + reg_m=1.0, + method="sinkhorn", + reg_type="entropy", returnCost="invalid", )