From 30d8ffea42a7e024d9f05dbac375d8c8d38793e6 Mon Sep 17 00:00:00 2001 From: Raashish Aggarwal <94279692+raashish1601@users.noreply.github.com> Date: Wed, 7 Oct 2026 20:27:37 +0530 Subject: [PATCH 1/2] [MRG] Accept 1D sample weights in the empirical Gaussian OT functions --- ot/gaussian.py | 26 ++++++++++++++++++ test/test_gaussian.py | 61 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 87 insertions(+) diff --git a/ot/gaussian.py b/ot/gaussian.py index 103d102d8..e137d0421 100644 --- a/ot/gaussian.py +++ b/ot/gaussian.py @@ -267,9 +267,13 @@ def empirical_bures_wasserstein_mapping( if ws is None: ws = nx.ones((xs.shape[0], 1), type_as=xs) / xs.shape[0] + else: + ws = nx.reshape(ws, (-1, 1)) if wt is None: wt = nx.ones((xt.shape[0], 1), type_as=xt) / xt.shape[0] + else: + wt = nx.reshape(wt, (-1, 1)) if bias: mxs = nx.dot(ws.T, xs) / nx.sum(ws) @@ -389,9 +393,13 @@ def empirical_bures_wasserstein_mapping_hd( if ws is None: ws = nx.ones((xs.shape[0], 1), type_as=xs) / xs.shape[0] + else: + ws = nx.reshape(ws, (-1, 1)) if wt is None: wt = nx.ones((xt.shape[0], 1), type_as=xt) / xt.shape[0] + else: + wt = nx.reshape(wt, (-1, 1)) if bias: mxs = nx.dot(ws.T, xs) / nx.sum(ws) @@ -760,9 +768,13 @@ def empirical_bures_wasserstein_distance( if ws is None: ws = nx.ones((xs.shape[0], 1), type_as=xs) / xs.shape[0] + else: + ws = nx.reshape(ws, (-1, 1)) if wt is None: wt = nx.ones((xt.shape[0], 1), type_as=xt) / xt.shape[0] + else: + wt = nx.reshape(wt, (-1, 1)) if bias: mxs = nx.dot(ws.T, xs) / nx.sum(ws) @@ -861,9 +873,13 @@ def empirical_bures_wasserstein_distance_hd( if ws is None: ws = nx.ones((xs.shape[0], 1), type_as=xs) / xs.shape[0] + else: + ws = nx.reshape(ws, (-1, 1)) if wt is None: wt = nx.ones((xt.shape[0], 1), type_as=xt) / xt.shape[0] + else: + wt = nx.reshape(wt, (-1, 1)) if bias: mxs = nx.dot(ws.T, xs) / nx.sum(ws) @@ -1343,6 +1359,8 @@ def empirical_bures_wasserstein_barycenter( w = [ nx.ones((X[i].shape[0], 1), type_as=X[i]) / X[i].shape[0] for i in range(k) ] + else: + w = [nx.reshape(w[i], (-1, 1)) for i in range(k)] if bias: m = [nx.dot(w[i].T, X[i]) / nx.sum(w[i]) for i in range(k)] @@ -1468,9 +1486,13 @@ def empirical_gaussian_gromov_wasserstein_distance(xs, xt, ws=None, wt=None, log if ws is None: ws = nx.ones((xs.shape[0], 1), type_as=xs) / xs.shape[0] + else: + ws = nx.reshape(ws, (-1, 1)) if wt is None: wt = nx.ones((xt.shape[0], 1), type_as=xt) / xt.shape[0] + else: + wt = nx.reshape(wt, (-1, 1)) mxs = nx.dot(ws.T, xs) / nx.sum(ws) mxt = nx.dot(wt.T, xt) / nx.sum(wt) @@ -1636,9 +1658,13 @@ def empirical_gaussian_gromov_wasserstein_mapping( if ws is None: ws = nx.ones((xs.shape[0], 1), type_as=xs) / xs.shape[0] + else: + ws = nx.reshape(ws, (-1, 1)) if wt is None: wt = nx.ones((xt.shape[0], 1), type_as=xt) / xt.shape[0] + else: + wt = nx.reshape(wt, (-1, 1)) # estimate mean and covariance mu_s = nx.dot(ws.T, xs) / nx.sum(ws) diff --git a/test/test_gaussian.py b/test/test_gaussian.py index 8beadea98..d0908a537 100644 --- a/test/test_gaussian.py +++ b/test/test_gaussian.py @@ -333,6 +333,51 @@ def test_empirical_bures_wasserstein_distance(nx, bias): np.testing.assert_allclose(10 * bias, nx.to_numpy(Wb), rtol=1e-2, atol=1e-2) +@pytest.mark.parametrize( + "func", + [ + ot.gaussian.empirical_bures_wasserstein_distance, + ot.gaussian.empirical_bures_wasserstein_mapping, + ot.gaussian.empirical_gaussian_gromov_wasserstein_distance, + ot.gaussian.empirical_gaussian_gromov_wasserstein_mapping, + ], +) +def test_empirical_gaussian_1d_weights(nx, func): + # sample weights of shape (n,) give the same result as weights of shape (n, 1) + rng = np.random.RandomState(0) + xs = rng.randn(20, 3) + xt = rng.randn(15, 3) * 2 + 1 + ws = rng.rand(20) + wt = rng.rand(15) + xsb, xtb, wsb, wtb = nx.from_numpy(xs, xt, ws, wt) + + expected = func(xsb, xtb, ws=wsb[:, None], wt=wtb[:, None]) + result = func(xsb, xtb, ws=wsb, wt=wtb) + + if not isinstance(result, tuple): + result, expected = (result,), (expected,) + for r, e in zip(result, expected): + np.testing.assert_allclose(nx.to_numpy(r), nx.to_numpy(e), rtol=1e-5) + + +def test_empirical_bures_wasserstein_distance_1d_weights(): + # with as many samples as dimensions, (n,) weights used to be silently + # applied to the dimensions instead of the samples + rng = np.random.RandomState(0) + xs = rng.randn(3, 3) + xt = rng.randn(3, 3) + 1 + ws = np.array([0.1, 0.3, 0.6]) + wt = np.array([0.5, 0.2, 0.3]) + + W = ot.gaussian.empirical_bures_wasserstein_distance(xs, xt, ws=ws, wt=wt) + + ms, mt = ws @ xs, wt @ xt + Cs = (xs - ms).T @ np.diag(ws) @ (xs - ms) + 1e-6 * np.eye(3) + Ct = (xt - mt).T @ np.diag(wt) @ (xt - mt) + 1e-6 * np.eye(3) + expected = ot.gaussian.bures_wasserstein_distance(ms, mt, Cs, Ct) + np.testing.assert_allclose(W, expected, rtol=1e-6) + + @pytest.mark.parametrize("bias", [True, False]) def test_empirical_bures_wasserstein_distance_hd(nx, bias): ns = 400 @@ -585,6 +630,22 @@ def test_empirical_bures_wasserstein_barycenter(nx, bias): np.testing.assert_allclose(mb, mblog, rtol=1e-2, atol=1e-2) +def test_empirical_bures_wasserstein_barycenter_1d_weights(nx): + rng = np.random.RandomState(0) + X = [rng.randn(20, 2), rng.randn(15, 2) + 1] + w = [rng.rand(20), rng.rand(15)] + Xb = nx.from_numpy(*X) + wb = nx.from_numpy(*w) + + mb, Cb = ot.gaussian.empirical_bures_wasserstein_barycenter(Xb, w=wb) + mb2, Cb2 = ot.gaussian.empirical_bures_wasserstein_barycenter( + Xb, w=[wi[:, None] for wi in wb] + ) + + np.testing.assert_allclose(nx.to_numpy(mb), nx.to_numpy(mb2), rtol=1e-5) + np.testing.assert_allclose(nx.to_numpy(Cb), nx.to_numpy(Cb2), rtol=1e-5) + + @pytest.mark.parametrize("d_target", [1, 2, 3, 10]) def test_gaussian_gromov_wasserstein_distance(nx, d_target): ns = 400 From e2499f091c7243ffb05b4160f156e39575dce9f4 Mon Sep 17 00:00:00 2001 From: Raashish Aggarwal <94279692+raashish1601@users.noreply.github.com> Date: Wed, 7 Oct 2026 20:27:50 +0530 Subject: [PATCH 2/2] Add release note --- RELEASES.md | 1 + 1 file changed, 1 insertion(+) diff --git a/RELEASES.md b/RELEASES.md index 860f2babb..e233546f8 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -15,6 +15,7 @@ #### Closed issues +- Accept sample weights of shape `(n,)`, as documented, in the empirical functions of `ot.gaussian`; they failed to broadcast, or were applied to the dimensions when there were as many samples as dimensions (PR #885) - 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)