Skip to content

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

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

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

Conversation

@xihaian251

@xihaian251 xihaian251 commented Sep 27, 2026 •

Copy link
Copy Markdown

Problem

The entropic unbalanced OT solvers report a value that is not the objective they solve.

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

Two independent defects, both in ot/unbalanced/_sinkhorn.py:

  1. reg_type='kl': the penalizations were evaluated with nx.kl_div(..., mass=False), i.e.
    Σ p log(p/q), which is not a divergence. That quantity is exactly the derivative of the
    objective along the scaling direction G → tG, so by first-order optimality it vanishes at
    the optimum: the reported value was numerically zero for every problem.
  2. reg_type='entropy': with c = 1 the solver minimizes reg · KL_gen(γ, 1), while the
    documented regularizer is Ω_ent(γ) = Σ γ log γ − Σ γ = KL_gen(γ, 1) − n_a · n_b. The offset
    is constant, so the plan was right 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).

Root cause

Both follow from the definition of the regularizer. Writing the objective as

J(γ) = ⟨γ, M⟩ + reg · Ω_reg_type(γ, c) + reg_m1 · KL(γ1, a) + reg_m2 · KL(γᵀ1, b)

Ω_reg_type(γ, c) = KL(γ, c)                       if reg_type = 'kl'
                 = Σ_{i,j} γ_ij log γ_ij − γ_ij   if reg_type = 'entropy'

KL(P, Q) = Σ_{i,j} P_ij log(P_ij / Q_ij) − P_ij + Q_ij      (generalized KL)

Σ p log(p/q) is KL_gen without its mass-correction term, and KL_gen(γ, 1) − Ω_ent(γ) = n_a n_b.

Fix

  • reg_cost = nx.kl_div(plan, c, mass=True), and for reg_type='entropy' the constant
    dim_a * dim_b is removed before weighting by reg.
  • Marginal penalizations untouched; c / reg_type coupling unchanged.
  • reg_type is now normalised once per solver through a _check_reg_type() helper: it is
    handled case insensitively by all three methods, and an unsupported name raises ValueError
    instead of silently falling back to 'kl'.
  • Docstrings: KL is defined as the generalized divergence, the objective uses
    Ω_reg_type(γ, c) with its two cases, the reg_type blocks are identical across the five
    public functions, a Notes section lists the log keys, and the multi-histogram behaviour is
    stated. The Returns headers that read if n_hists == 1: (code tests n_hists == 0) were
    corrected.
  • Multi-histogram mode: c and reg_type are ignored there (only the negative entropy
    regularizer is implemented, and only the linear cost is returned). This is a pre-existing,
    already-documented limitation (test/unbalanced/test_sinkhorn.py:572), so the algorithm was
    not changed; a UserWarning now fires when the request is silently ignored, and
    returnCost is validated in that branch ('total' warns instead of being ignored).

Tests

test/unbalanced/test_sinkhorn.py — six new test functions, i.e. 42 parametrised cases
over the numpy and torch backends. On the previous head 26 of those 42 fail and 16 pass
(the 16 include test_unbalanced_total_cost_kl and the reg_type='kl' case-insensitivity
cases, which confirms that the first round of this PR is preserved). After the change all
42 pass:/n

  • test_unbalanced_total_cost_kl — non-regression for the generalized KL objective;
  • test_unbalanced_total_cost_entropy — the value equals ⟨γ,M⟩ + reg·Ω_ent(γ) + marginals,
    recomputed with plain numpy rather than nx.kl_div;
  • test_unbalanced_entropy_constant — with c = 1, kl and entropy give the same plan and
    the two values differ by exactly reg · n_a · n_b;
  • test_unbalanced_reg_type_case_insensitive, test_unbalanced_unknown_reg_type_raises;
  • test_unbalanced_multiple_inputs_reg_type_warns — the 2d-b branch warns once and only when
    the request is silently ignored, and the message reports the actual n_hists;
  • test_unbalanced_multiple_inputs_returnCost — returnCost is honoured or reported in 2d mode
    instead of being silently ignored, and an invalid value raises.

test/test_solvers.py — test_solve_unbalanced_value removed as requested by the reviewer;
a global ot.solve wrapper test will be proposed separately.

Validation

Relevant suites: test/unbalanced/, test/test_solvers.py, test/test_da.py (results below).
An A/B matrix of 199 configurations (3 methods × 5 reg_type spellings × 3 c × 4 reg_m ×
{numpy, torch}, plus dispatcher and multi-histogram entries) was run against the previous
head. 83 of the 199 use reg_type='kl', and those 83 are bit-identical in plan, cost,
total_cost and log keys. The remaining 116 are entropy / capitalised spellings, where
the reported value is meant to change (and, for the capitalised spellings, the plan too —
see Compatibility).

jax, tensorflow, cupy and GPU were not available on the machine used for this patch; those are
left to the CI.

Compatibility

  • reg_type='kl': no change to plans, costs, log keys or gradients. This is the default,
    and it is the path taken by ot.solve(..., reg=..., unbalanced=...).
  • reg_type='entropy': the total_cost reported by these solvers decreases by
    reg · n_a · n_b, which is the documented objective; plans are unchanged. Note that
    ot.solve(..., reg=..., unbalanced=..., reg_type='entropy') is not affected: it is
    routed to lbfgsb_unbalanced, which uses the optimizer's own objective and never goes
    through these functions. The reg_type='entropy' value change is therefore visible
    through ot.unbalanced.sinkhorn_unbalanced / sinkhorn_unbalanced2 (and
    returnCost='total'), not through ot.solve.
  • reg_type='Entropy' / 'ENTROPY' now select the entropy regularizer on
    sinkhorn_stabilized_unbalanced and sinkhorn_unbalanced_translation_invariant (they used to
    fall back to kl and solve a different problem).
  • An unsupported reg_type now raises ValueError instead of being silently treated as 'kl'.
  • Multi-histogram mode: numbers unchanged, but a warning is now emitted when reg_type != 'entropy'
    or an explicit c is passed, and an invalid returnCost raises.

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
@codecov

codecov Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 97.01%. Comparing base (604c47f) to head (9aaf387).

Additional details and impacted files
@@            Coverage Diff             @@
##           master     #874      +/-   ##
==========================================
+ Coverage   97.00%   97.01%   +0.01%     
==========================================
  Files         128      128              
  Lines       26349    26490     +141     
==========================================
+ Hits        25559    25700     +141     
  Misses        790      790              
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@xihaian251

Copy link
Copy Markdown
Author

Hi, thanks for syncing the branch with master. The new workflow runs are currently waiting for approval. Could you please approve them when convenient? The previous run also reported full coverage of all modified lines. Thanks!

@cedricvincentcuaz
cedricvincentcuaz requested a review from 6Ulm October 8, 2026 11:14

@cedricvincentcuaz cedricvincentcuaz left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Dear @xihaian251, thank you for your PR !
Looking more globally at the functions in the unbalanced/_sinkhorn.py that you are modifying, I spotted some inconsistencies that it would be nice to resolve in this PR because i'm not entirely sure that your modifications lead to the intended behaviour when reg_type="entropy". Could you check that in details @6Ulm too as you coded most of these functions?
Modifications i'd consider relevant:

  • in the docstrings, systematically mention in the equations that KL relates to the generalized KL.
  • It would be clearier to systematically define the regularizer as $\Omega_{reg_type}(\gamma, c)$ with the distinct cases. So that we can harmonize how we refer to this regularization in the equations and the parameters description in particular reg_type.
  • the transformations reg_type.lower() and method.lower() are not consistently applied to all functions.
  • Would be nice to explicit keys of produced log dictionaries too.
  • The dependencies to c and reg_type are not clearly explained when n_hists > 1
  • Importantly, to properly fix the objective issues even in the simple case when reg_type = "entropy" I think we should remove a constant $n_a * n_b$ as $\Omega_{ent} = KL(\gamma, 1) - n_a*n_b$.

I would be very grateful that you both @xihaian251 and @6Ulm look into that.
Best,
Cédric

Comment thread test/test_solvers.py Outdated
pytest.skip("Not implemented")


def test_solve_unbalanced_value(nx):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You can remove this new test function because actually the problem is more generic and it would make sense to have a more global test that covers all options available in the ot.solve wrappers which go beyond these specific unbalanced problems - that could be covered in another PR

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Hi Cédric,

Thank you for the detailed review and for pointing out these inconsistencies.

I'll carefully revisit the objective definitions, particularly the constant term for reg_type="entropy" and the behavior with multiple histograms. I'll also address the documentation, case normalization, and log dictionary descriptions.

I'll remove the test_solve_unbalanced_value test as suggested and keep the broader ot.solve wrapper testing out of this PR.

I'll update the PR once the revised implementation and regression tests have been validated.

Best regards

B added 3 commits October 9, 2026 11:11
…y, 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.
- 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.
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.
@xihaian251

Copy link
Copy Markdown
Author

Hi Cédric,

Thanks for the review — you were right about the reg_type="entropy" objective.

I went through every point below. For reference I ran an A/B matrix of 199 configurations
(3 methods × 5 reg_type spellings × 3 c × 4 reg_m × {numpy, torch}, plus dispatcher and
multi-histogram entries) against the current head. 83 of the 199 use reg_type='kl', and those
83 are bit-identical in plan, cost, total_cost and log keys. The other 116 are the
entropy / capitalised spellings, where the value is meant to change.

1. Generalized KL in the docstrings.
All five functions now state
KL(P, Q) = Σ_{i,j} P_ij log(P_ij / Q_ij) − P_ij + Q_ij explicitly, and the objective is
written with Ω_reg_type(γ, c) instead of a bare KL(γ, c).

2. Unified Ω_reg_type(γ, c).
The problem is now written as

J(γ) = ⟨γ, M⟩ + reg · Ω_reg_type(γ, c) + reg_m1 · KL(γ1, a) + reg_m2 · KL(γᵀ1, b)

Ω_reg_type(γ, c) = KL(γ, c)                        if reg_type = 'kl'
                 = Σ_{i,j} γ_ij log γ_ij − γ_ij    if reg_type = 'entropy'

and the five reg_type parameter descriptions have been replaced by one identical block.

3. reg_type.lower() / method.lower().
method was already lower-cased at every dispatch site (12 sites, all consistent), so nothing
was needed there. reg_type was not: sinkhorn_stabilized_unbalanced and
sinkhorn_unbalanced_translation_invariant compared reg_type == "entropy" without .lower(),
so reg_type="Entropy" silently took the kl branch and solved a different problem
(measured: total_cost = 0.43617 instead of 8.11229, and a different plan). All three solvers
now go through a _check_reg_type() helper, and an unsupported name raises ValueError instead
of silently falling back to 'kl'.
This is a behaviour change for the capitalised spellings on those two methods; it is noted in
RELEASES.md.

4. log keys.
Each of the five functions now has a Notes section listing them. For the three solvers,
n_hists == 0 gives err, logu, logv, cost, total_cost, while n_hists > 0 gives only
err, logu, logv — cost and total_cost are simply not computed in that branch. The two
dispatchers forward the log of the solver selected by method. I also corrected the Returns
headers that said if n_hists == 1: whereas the code tests n_hists == 0.

5. n_hists > 1.
You are right that this was not explained. In the multi-histogram branch K = exp(−M / reg) and
c never enters, so the effective reference measure is the all-ones matrix and only the negative
entropy regularization is implemented. Measured: c ∈ {1, 2·1, a bᵀ} give exactly the same
result (max diff 0.0), reg_type 'kl' and 'entropy' give exactly the same result, and the
multi-histogram output equals the per-column single-histogram result computed with entropy
(≤ 1e-7). This matches the comment already present at
test/unbalanced/test_sinkhorn.py:572 ("reg_type="entropy" as multiple inputs does not work
for KL yet").

Since changing the algorithm would change the numbers returned to existing users, I did not
touch it. Instead the behaviour is documented, and a single UserWarning is emitted when the
2d-b branch silently ignores what was asked for (reg_type != 'entropy', or an explicit
c); the message reports the actual n_hists, since a 2d b with a single column takes that
branch too. reg_type='entropy' with the default c stays silent, and passing an explicit c
with entropy does not produce two warnings for the same thing. I also made returnCost
consistent in that branch: an invalid value now raises, and returnCost='total' warns instead
of being silently ignored. I am happy to open a separate issue for KL support in the
multi-histogram path.

One point I would like your view on: because 'kl' is the documented default, that warning now
fires on a plain sinkhorn_unbalanced(a, b_2d, M, reg, reg_m) call, which used to be silent
(sinkhorn_unbalanced2 did warn for that case). If you would rather keep that path silent and
rely on the docstrings only, it is a one-line change.

6. The entropy constant n_a · n_b.
Confirmed, and this was the real correctness problem. With c = 1 the solver minimizes
reg · KL(γ, 1) while the documented regularizer is
Ω_ent(γ) = Σ γ log γ − Σ γ = KL(γ, 1) − n_a · n_b. The offset is constant, so the plan was
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 all three methods. The three total_cost blocks now remove
it, and the marginal penalizations are untouched.
As you implied, total_cost ≥ linear_cost does not hold for entropy (the value can be
negative), so the test is split into a kl test and an entropy test.

7. test_solve_unbalanced_value.
Removed as requested. I agree a global ot.solve wrapper test is the right place for that
coverage and will propose it in a separate PR.

Tests. The new tests fail on the current head (26 of 42 cases) and all 42 pass after the
change:

  • test_unbalanced_total_cost_kl — non-regression for the generalized KL objective
    (passes before and after);
  • test_unbalanced_total_cost_entropy — the reported value equals
    ⟨γ, M⟩ + reg · Ω_ent(γ) + marginals, recomputed with plain numpy rather than nx.kl_div;
  • test_unbalanced_entropy_constant — with c = 1, kl and entropy return the same plan and
    the two values differ by exactly reg · n_a · n_b;
  • test_unbalanced_reg_type_case_insensitive, test_unbalanced_unknown_reg_type_raises;
  • test_unbalanced_multiple_inputs_reg_type_warns;
  • test_unbalanced_multiple_inputs_returnCost.

Relevant suites run locally: test/unbalanced/, test/test_solvers.py, test/test_da.py.
jax, tensorflow, cupy and GPU were not available on this machine, so those are left to the CI.

@6Ulm — the multi-histogram semantics is the one place where I deliberately did not change the
algorithm; I would be glad to hear whether you agree with documenting it instead.

Best,

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.

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.

2 participants