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

- 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)
Expand Down
26 changes: 26 additions & 0 deletions ot/gaussian.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)]
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
61 changes: 61 additions & 0 deletions test/test_gaussian.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading