Repository navigation
[MRG] Fix the total cost reported by the entropic unbalanced OT solvers - #874
xihaian251 wants to merge 7 commits into
Conversation
The marginal penalization was computed with the un-normalized KL divergence nx.kl_div(..., mass=False), i.e. sum(p log(p/q)), which is not a divergence. At the optimum of the unbalanced OT problem this quantity is the derivative of the objective along the scaling direction G -> t G, so it vanishes: ot.solve(..., reg=..., unbalanced=...).value was numerically zero for every problem, and log["total_cost"] of ot.unbalanced.sinkhorn_unbalanced was too. Use the generalized KL divergence (mass=True), as already done in ot.unbalanced.mm_unbalanced, in sinkhorn_knopp_unbalanced, sinkhorn_stabilized_unbalanced and sinkhorn_unbalanced_translation_invariant. The transport plan, its gradient and all other outputs are unchanged. Add non-regression tests at the solver level and at the ot.solve level. Both fail on master and pass with this change.
Codecov Report✅ All modified and coverable lines are covered by tests. 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:
|
|
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! |
There was a problem hiding this comment.
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 particularreg_type. - the transformations
reg_type.lower()andmethod.lower()are not consistently applied to all functions. - Would be nice to explicit keys of produced
logdictionaries too. - The dependencies to
candreg_typeare not clearly explained whenn_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
| pytest.skip("Not implemented") | ||
|
|
||
|
|
||
| def test_solve_unbalanced_value(nx): |
There was a problem hiding this comment.
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
There was a problem hiding this comment.
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
…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.
|
Hi Cédric, Thanks for the review — you were right about the I went through every point below. For reference I ran an A/B matrix of 199 configurations 1. Generalized KL in the docstrings. 2. Unified and the five 3. 4. 5. Since changing the algorithm would change the numbers returned to existing users, I did not One point I would like your view on: because 6. The entropy constant 7. Tests. The new tests fail on the current head (26 of 42 cases) and all 42 pass after the
Relevant suites run locally: @6Ulm — the multi-histogram semantics is the one place where I deliberately did not change the 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.
Problem
The entropic unbalanced OT solvers report a value that is not the objective they solve.
Two independent defects, both in
ot/unbalanced/_sinkhorn.py:reg_type='kl': the penalizations were evaluated withnx.kl_div(..., mass=False), i.e.Σ p log(p/q), which is not a divergence. That quantity is exactly the derivative of theobjective along the scaling direction
G → tG, so by first-order optimality it vanishes atthe optimum: the reported value was numerically zero for every problem.
reg_type='entropy': withc = 1the solver minimizesreg · KL_gen(γ, 1), while thedocumented regularizer is
Ω_ent(γ) = Σ γ log γ − Σ γ = KL_gen(γ, 1) − n_a · n_b. The offsetis constant, so the plan was right but the reported value was
reg · n_a · n_btoo large(measured exactly:
+9.0forreg = 0.3,n_a = 5,n_b = 6).Root cause
Both follow from the definition of the regularizer. Writing the objective as
Σ p log(p/q)isKL_genwithout its mass-correction term, andKL_gen(γ, 1) − Ω_ent(γ) = n_a n_b.Fix
reg_cost = nx.kl_div(plan, c, mass=True), and forreg_type='entropy'the constantdim_a * dim_bis removed before weighting byreg.c/reg_typecoupling unchanged.reg_typeis now normalised once per solver through a_check_reg_type()helper: it ishandled case insensitively by all three methods, and an unsupported name raises
ValueErrorinstead of silently falling back to
'kl'.KLis defined as the generalized divergence, the objective usesΩ_reg_type(γ, c)with its two cases, thereg_typeblocks are identical across the fivepublic functions, a
Notessection lists thelogkeys, and the multi-histogram behaviour isstated. The
Returnsheaders that readif n_hists == 1:(code testsn_hists == 0) werecorrected.
candreg_typeare ignored there (only the negative entropyregularizer 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 wasnot changed; a
UserWarningnow fires when the request is silently ignored, andreturnCostis validated in that branch ('total'warns instead of being ignored).Tests
test/unbalanced/test_sinkhorn.py— six new test functions, i.e. 42 parametrised casesover the numpy and torch backends. On the previous head 26 of those 42 fail and 16 pass
(the 16 include
test_unbalanced_total_cost_kland thereg_type='kl'case-insensitivitycases, 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— withc = 1,klandentropygive the same plan andthe 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-bbranch warns once and only whenthe request is silently ignored, and the message reports the actual
n_hists;test_unbalanced_multiple_inputs_returnCost—returnCostis honoured or reported in 2d modeinstead of being silently ignored, and an invalid value raises.
test/test_solvers.py—test_solve_unbalanced_valueremoved as requested by the reviewer;a global
ot.solvewrapper 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_typespellings × 3c× 4reg_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_costandlogkeys. The remaining 116 areentropy/ capitalised spellings, wherethe 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,logkeys or gradients. This is the default,and it is the path taken by
ot.solve(..., reg=..., unbalanced=...).reg_type='entropy': thetotal_costreported by these solvers decreases byreg · n_a · n_b, which is the documented objective; plans are unchanged. Note thatot.solve(..., reg=..., unbalanced=..., reg_type='entropy')is not affected: it isrouted to
lbfgsb_unbalanced, which uses the optimizer's own objective and never goesthrough these functions. The
reg_type='entropy'value change is therefore visiblethrough
ot.unbalanced.sinkhorn_unbalanced/sinkhorn_unbalanced2(andreturnCost='total'), not throughot.solve.reg_type='Entropy'/'ENTROPY'now select the entropy regularizer onsinkhorn_stabilized_unbalancedandsinkhorn_unbalanced_translation_invariant(they used tofall back to
kland solve a different problem).reg_typenow raisesValueErrorinstead of being silently treated as'kl'.reg_type != 'entropy'or an explicit
cis passed, and an invalidreturnCostraises.