Skip to content

[MRG] Accept 1D sample weights in the empirical Gaussian OT functions - #885

Open
raashish1601 wants to merge 2 commits into
PythonOT:masterfrom
raashish1601:fix/gaussian-1d-weights
Open

raashish1601 wants to merge 2 commits into
PythonOT:masterfrom
raashish1601:fix/gaussian-1d-weights

Conversation

@raashish1601

Copy link
Copy Markdown

Types of changes

  • Bug fix

Motivation and context / Related issue

ot.gaussian.empirical_bures_wasserstein_distance documents ws/wt with shape (ns), and empirical_bures_wasserstein_barycenter documents w as a list of (n,) arrays, but the code only works with column vectors of shape (n, 1). With 1D weights:

  • if the number of samples differs from the dimension, xs * ws fails with a broadcasting error;
  • if they are equal, xs * ws scales the dimensions instead of the samples and ws.T @ xs gives a wrong mean, so a wrong value is returned without any error.
import numpy as np, ot
rng = np.random.RandomState(0)
xs, xt = rng.randn(20, 3), rng.randn(15, 3)
ot.gaussian.empirical_bures_wasserstein_distance(xs, xt, ws=rng.rand(20), wt=rng.rand(15))
# ValueError: operands could not be broadcast together with shapes (20,3) (20,)

The other empirical functions in ot.gaussian (mappings, _hd variants, Gaussian Gromov-Wasserstein) have the same code. This PR reshapes user-given sample weights to (n, 1) in all of them, so both (n,) and (n, 1) work. Weights that are already (n, 1) are unchanged.

How has this been tested (if it applies)

New tests in test/test_gaussian.py:

  • test_empirical_gaussian_1d_weights: for the distance, mapping, GGW distance and GGW mapping, 1D weights give the same result as (n, 1) weights (all backends);
  • test_empirical_bures_wasserstein_distance_1d_weights: with 3 samples in dimension 3, the result equals the Bures-Wasserstein distance computed from the weighted means and covariances;
  • test_empirical_bures_wasserstein_barycenter_1d_weights: same check for the empirical barycenter.

All 11 new cases fail on master. pytest test/test_gaussian.py passes (numpy and torch backends 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