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 @@ -15,6 +15,7 @@

#### Closed issues

- Fix `ot.dist` ignoring the weights `w` for the `sqeuclidean`, `euclidean`, `cosine` and `correlation` metrics and for the other metrics computed with `scipy.spatial.distance.cdist` (PR #884)
- 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 `ot.dist` ignoring the weights `w` for `metric="cityblock"`, which returned the unweighted distance although the weights are documented for this metric (PR #859)
Expand Down
33 changes: 19 additions & 14 deletions ot/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -505,10 +505,13 @@ def dist(
if w is not None:
return nx.from_numpy(cdist(x1, x2, metric=metric, w=w))
return nx.from_numpy(cdist(x1, x2, metric=metric))
elif metric == "sqeuclidean":
return euclidean_distances(x1, x2, squared=True, nx=nx)
elif metric == "euclidean":
return euclidean_distances(x1, x2, squared=False, nx=nx)
elif metric in ("sqeuclidean", "euclidean"):
if w is not None:
# sum_k w_k (x1_k - x2_k)^2 is the squared distance between the
# samples scaled by sqrt(w)
x1 = x1 * nx.sqrt(w)[None, :]
x2 = x2 * nx.sqrt(w)[None, :]
return euclidean_distances(x1, x2, squared=metric == "sqeuclidean", nx=nx)
elif metric == "cityblock":
if w is None:
if use_tensor:
Expand Down Expand Up @@ -557,13 +560,17 @@ def dist(
for i in range(x1.shape[1]):
M += w[i] * nx.abs(x1[:, i][:, None] - x2[:, i][None, :]) ** p
return M ** (1 / p)
elif metric == "cosine":
nx1 = nx.sqrt(nx.einsum("ij,ij->i", x1, x1))
nx2 = nx.sqrt(nx.einsum("ij,ij->i", x2, x2))
return 1.0 - (nx.dot(x1, nx.transpose(x2)) / nx1[:, None] / nx2[None, :])
elif metric == "correlation":
x1 = x1 - nx.mean(x1, axis=1)[:, None]
x2 = x2 - nx.mean(x2, axis=1)[:, None]
elif metric in ("cosine", "correlation"):
if metric == "correlation":
if w is None:
x1 = x1 - nx.mean(x1, axis=1)[:, None]
x2 = x2 - nx.mean(x2, axis=1)[:, None]
else:
x1 = x1 - (nx.dot(x1, w) / nx.sum(w))[:, None]
x2 = x2 - (nx.dot(x2, w) / nx.sum(w))[:, None]
if w is not None:
x1 = x1 * nx.sqrt(w)[None, :]
x2 = x2 * nx.sqrt(w)[None, :]
nx1 = nx.sqrt(nx.einsum("ij,ij->i", x1, x1))
nx2 = nx.sqrt(nx.einsum("ij,ij->i", x2, x2))
return 1.0 - (nx.dot(x1, nx.transpose(x2)) / nx1[:, None] / nx2[None, :])
Expand All @@ -573,9 +580,7 @@ def dist(
else:
if isinstance(metric, str) and metric.endswith("minkowski"):
return cdist(x1, x2, metric=metric, p=p, w=w)
# Only pass w parameter for metrics that support it
# According to SciPy docs, only 'minkowski' and 'wminkowski' support w
if w is not None and metric in ["minkowski", "wminkowski"]:
if w is not None:
return cdist(x1, x2, metric=metric, w=w)
return cdist(x1, x2, metric=metric)

Expand Down
39 changes: 39 additions & 0 deletions test/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -320,6 +320,45 @@ def test_dist_weighted_cityblock(use_tensor):
np.testing.assert_allclose(D, expected, atol=1e-12)


@pytest.mark.parametrize(
"metric",
[
"sqeuclidean",
"euclidean",
"cosine",
"correlation",
"braycurtis",
"canberra",
],
)
def test_dist_weighted_vs_cdist(metric):
rng = np.random.RandomState(0)
x1 = rng.randn(5, 3)
x2 = rng.randn(4, 3)
w = rng.rand(3)

expected = scipy.spatial.distance.cdist(x1, x2, metric=metric, w=w)
D = ot.dist(x1, x2, metric=metric, w=w)

np.testing.assert_allclose(D, expected, atol=1e-12)


@pytest.mark.parametrize(
"metric", ["sqeuclidean", "euclidean", "cosine", "correlation"]
)
def test_dist_weighted_backends(nx, metric):
rng = np.random.RandomState(0)
x1 = rng.randn(5, 3)
x2 = rng.randn(4, 3)
w = rng.rand(3)

expected = scipy.spatial.distance.cdist(x1, x2, metric=metric, w=w)
D = ot.dist(nx.from_numpy(x1), nx.from_numpy(x2), metric=metric, w=nx.from_numpy(w))

# low atol because jax forces float32
np.testing.assert_allclose(nx.to_numpy(D), expected, atol=1e-5)


def test_sparse_ot_dist_uses_pair_weights():
x1 = np.array([[0.0], [1.0]])
x2 = np.array([[0.0], [5.0]])
Expand Down
Loading