Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions RELEASES.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 #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)
Expand Down
33 changes: 24 additions & 9 deletions ot/unbalanced/_sinkhorn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
34 changes: 34 additions & 0 deletions test/test_solvers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
50 changes: 50 additions & 0 deletions test/unbalanced/test_sinkhorn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Loading