Skip to content

[MRG] Use the weights w in ot.dist for all metrics - #884

Open
raashish1601 wants to merge 2 commits into
PythonOT:masterfrom
raashish1601:fix/dist-weights
Open

raashish1601 wants to merge 2 commits into
PythonOT:masterfrom
raashish1601:fix/dist-weights

Conversation

@raashish1601

Copy link
Copy Markdown

Types of changes

  • Bug fix

Motivation and context / Related issue

ot.dist silently ignored w for every metric except cityblock and minkowski, so it returned the unweighted distance:

import numpy as np, ot
from scipy.spatial.distance import cdist
x, y, w = np.array([[0., 0.]]), np.array([[1., 1.]]), np.array([1., 2.])
ot.dist(x, y, metric="sqeuclidean", w=w)   # [[2.]]
cdist(x, y, metric="sqeuclidean", w=w)     # [[3.]]

This affected the backend implementations of sqeuclidean, euclidean, cosine and correlation, and the numpy fallback to cdist, which only passed w for minkowski/wminkowski. With backend="scipy" the weights were used, so the result depended on the backend. test_dist already lists these metrics as "those that support weights" but did not check the values. #859 fixed the same problem for cityblock.

Changes:

  • sqeuclidean/euclidean: scale both point sets by sqrt(w) before euclidean_distances.
  • cosine: same scaling. correlation: center with the weighted mean, then the same scaling. This matches the scipy definitions.
  • numpy fallback: pass w to cdist whenever it is given (scipy still raises for metrics that do not accept weights).

Unweighted results are unchanged.

How has this been tested (if it applies)

Added test_dist_weighted_vs_cdist (numpy, compared with scipy.spatial.distance.cdist for sqeuclidean, euclidean, cosine, correlation, braycurtis and canberra) and test_dist_weighted_backends (all backends, for the four backend metrics). All new cases fail on master. pytest test/test_utils.py passes (numpy and torch backends installed locally), and ruff 0.5.2 lint and format are clean.

PR checklist

  • I have read the CONTRIBUTING document.
  • The documentation is up-to-date with the changes I made (check build artifacts).
  • All tests passed, and additional code has been covered with new tests.
  • I have added the PR and Issue fix to the RELEASES.md file.

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.

1 participant