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 internal arrays created on the torch default device instead of the input device in `ot.partial.partial_wasserstein`, `ot.wasserstein_circle` (and `ot.sliced_wasserstein_sphere`), the high-dimensional Bures-Wasserstein functions in `ot.gaussian`, `ot.gmm`, `ot.lowrank_sinkhorn` with `init="deterministic"`, the semi-relaxed (F)GW barycenters and `ot.utils.projection_sparse_simplex`, which made them fail on GPU inputs. Add a `meta_default_device` test fixture to catch such allocations on CPU (PR #883, Issue #852)
- 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
10 changes: 5 additions & 5 deletions ot/gaussian.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,12 +175,12 @@ def bures_wasserstein_mapping_hd(ms, mt, Us, Ut, ls, lt, sigma2_s, sigma2_t, log

# source
Cs = nx.diag(nx.sqrt(ls + sigma2_s) - nx.sqrt(sigma2_s))
Ss_sq = dots(Us, Cs, Us.T) + nx.sqrt(sigma2_s) * nx.eye(p)
Ss_sq = dots(Us, Cs, Us.T) + nx.sqrt(sigma2_s) * nx.eye(p, type_as=Us)
Ds = nx.diag((nx.sqrt(ls + sigma2_s) - nx.sqrt(sigma2_s)) / nx.sqrt(ls + sigma2_s))
Ss_sqinv = (1 / nx.sqrt(sigma2_s)) * (nx.eye(p) - dots(Us, Ds, Us.T))
Ss_sqinv = (1 / nx.sqrt(sigma2_s)) * (nx.eye(p, type_as=Us) - dots(Us, Ds, Us.T))

# destination
St = dots(Ut, nx.diag(lt), Ut.T) + sigma2_t * nx.eye(p)
St = dots(Ut, nx.diag(lt), Ut.T) + sigma2_t * nx.eye(p, type_as=Ut)

M0 = nx.sqrtm(dots(Ss_sq, St, Ss_sq))

Expand Down Expand Up @@ -676,10 +676,10 @@ def bures_wasserstein_distance_hd(

# source
Cs = nx.diag(nx.sqrt(ls + sigma2_s) - nx.sqrt(sigma2_s))
Ss_sq = dots(Us, Cs, Us.T) + nx.sqrt(sigma2_s) * nx.eye(p)
Ss_sq = dots(Us, Cs, Us.T) + nx.sqrt(sigma2_s) * nx.eye(p, type_as=Us)

# destination
St = dots(Ut, nx.diag(lt), Ut.T) + sigma2_t * nx.eye(p)
St = dots(Ut, nx.diag(lt), Ut.T) + sigma2_t * nx.eye(p, type_as=Ut)
A = dots(Ss_sq, St, Ss_sq)
W2 = (
nx.sum((ms - mt) ** 2)
Expand Down
18 changes: 10 additions & 8 deletions ot/gmm.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ def gmm_pdf(x, m, C, w):
m.shape[0] == C.shape[0] == w.shape[0]
), "All GMM parameters must have the same amount of components"
nx = get_backend(x, m, C, w)
out = nx.zeros((x.shape[:-1]))
out = nx.zeros((x.shape[:-1]), type_as=x)
for k in range(m.shape[0]):
out = out + w[k] * gaussian_pdf(x, m[k], C[k])
return out
Expand Down Expand Up @@ -306,7 +306,7 @@ def gmm_ot_apply_map(
n_samples = x.shape[0]

if method == "bary":
out = nx.zeros(x.shape)
out = nx.zeros(x.shape, type_as=x)
logpdf = nx.stack(
[gaussian_logpdf(x, m_s[k], C_s[k])[:, None] for k in range(k_s)]
)
Expand Down Expand Up @@ -334,8 +334,8 @@ def gmm_ot_apply_map(
# i and j, b[i, j] is the translation part
rng = np.random.RandomState(seed)

A = nx.zeros((k_s, k_t, d, d))
b = nx.zeros((k_s, k_t, d))
A = nx.zeros((k_s, k_t, d, d), type_as=m_s)
b = nx.zeros((k_s, k_t, d), type_as=m_s)

# only need to compute for non-zero plan entries
for i, j in zip(*nx.where(plan > 0)):
Expand All @@ -350,13 +350,15 @@ def gmm_ot_apply_map(
[gaussian_logpdf(x, m_s[k], C_s[k]) for k in range(k_s)], axis=-1
)
# (n_samples, k_s)
out = nx.zeros(x.shape)
out = nx.zeros(x.shape, type_as=x)

for i_sample in range(n_samples):
log_g = logpdf[i_sample]
log_diff = log_g[:, None] - log_g[None, :]
weighted_exp = w_s[:, None] * nx.exp(log_diff)
denom = nx.sum(weighted_exp, axis=0)[:, None] * nx.ones(plan.shape[1])
denom = nx.sum(weighted_exp, axis=0)[:, None] * nx.ones(
plan.shape[1], type_as=plan
)
p_mat = plan / denom

p = p_mat.reshape(k_s * k_t) # stack line-by-line
Expand Down Expand Up @@ -418,8 +420,8 @@ def gmm_ot_plan_density(x, y, m_s, m_t, C_s, C_t, w_s, w_t, plan=None, atol=1e-2
nx = get_backend(x, y, m_s, m_t, C_s, C_t, w_s, w_t)

# hand-made d-variate meshgrid in ij indexing
xx = x[:, None, :] * nx.ones((1, m, 1)) # shapes (n, m, d)
yy = y[None, :, :] * nx.ones((n, 1, 1)) # shapes (n, m, d)
xx = x[:, None, :] * nx.ones((1, m, 1), type_as=x) # shapes (n, m, d)
yy = y[None, :, :] * nx.ones((n, 1, 1), type_as=y) # shapes (n, m, d)

if plan is None:
plan = gmm_ot_plan(m_s, m_t, C_s, C_t, w_s, w_t)
Expand Down
4 changes: 2 additions & 2 deletions ot/gromov/_semirelaxed.py
Original file line number Diff line number Diff line change
Expand Up @@ -1572,7 +1572,7 @@ def semirelaxed_gromov_barycenters(

S = len(Cs)
if lambdas is None:
lambdas = nx.ones(S) / S
lambdas = nx.ones(S, type_as=Cs[0]) / S
else:
lambdas = list_to_array(lambdas)
lambdas = nx.from_numpy(lambdas)
Expand Down Expand Up @@ -1886,7 +1886,7 @@ def semirelaxed_fgw_barycenters(

S = len(Cs)
if lambdas is None:
lambdas = nx.ones(S) / S
lambdas = nx.ones(S, type_as=Cs[0]) / S
else:
lambdas = list_to_array(lambdas)
lambdas = nx.from_numpy(lambdas)
Expand Down
2 changes: 1 addition & 1 deletion ot/lowrank.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ def _init_lr_sinkhorn(X_s, X_t, a, b, rank, init, reg_init, random_state, nx=Non

if init == "deterministic":
# Init g
g = nx.ones(rank) / rank
g = nx.ones(rank, type_as=X_s) / rank

lambda_1 = min(nx.min(a), nx.min(g), nx.min(b)) / 2
a1 = nx.arange(start=1, stop=ns + 1, type_as=X_s)
Expand Down
2 changes: 1 addition & 1 deletion ot/lp/solver_circle.py
Original file line number Diff line number Diff line change
Expand Up @@ -377,7 +377,7 @@ def binary_search_circle(
tp = nx.tile(tp, (1, m))
tc = (tm + tp) / 2

done = nx.zeros((u_values.shape[0], m))
done = nx.zeros((u_values.shape[0], m), type_as=u_values)

cpt = 0
while nx.any(1 - done):
Expand Down
8 changes: 6 additions & 2 deletions ot/partial/partial_solvers.py
Original file line number Diff line number Diff line change
Expand Up @@ -302,9 +302,13 @@ def partial_wasserstein(a, b, M, m=None, nb_dummies=1, log=False, **kwargs):
M_extension = nx.ones((nb_dummies, nb_dummies), type_as=M) * nx.max(M) * 2
M_extended = nx.concatenate(
(
nx.concatenate((M, nx.zeros((M.shape[0], M_extension.shape[1]))), axis=1),
nx.concatenate(
(nx.zeros((M_extension.shape[0], M.shape[1])), M_extension), axis=1
(M, nx.zeros((M.shape[0], M_extension.shape[1]), type_as=M)),
axis=1,
),
nx.concatenate(
(nx.zeros((M_extension.shape[0], M.shape[1]), type_as=M), M_extension),
axis=1,
),
),
axis=0,
Expand Down
8 changes: 4 additions & 4 deletions ot/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,18 +199,18 @@ def projection_sparse_simplex(V, max_nz, z=1, axis=None, nx=None):
max_nz_indices = nx.argsort(V, axis=1)[:, -max_nz:]
max_nz_indices = nx.flip(max_nz_indices, axis=1)

row_indices = nx.arange(V.shape[0])
row_indices = nx.arange(V.shape[0], type_as=max_nz_indices)
row_indices = row_indices.reshape(-1, 1)
# Extract the top max_nz values for each row
# and then project to simplex.
U = V[row_indices, max_nz_indices]
z = nx.ones(len(U)) * z
z = nx.ones(len(U), type_as=U) * z
cssv = nx.cumsum(U, axis=1) - z[:, None]
ind = nx.arange(max_nz) + 1
ind = nx.arange(max_nz, type_as=max_nz_indices) + 1
cond = U - cssv / ind > 0
# rho = nx.count_nonzero(cond, axis=1)
rho = nx.sum(cond, axis=1)
theta = cssv[nx.arange(len(U)), rho - 1] / rho
theta = cssv[nx.arange(len(U), type_as=rho), rho - 1] / rho
nz_projection = nx.maximum(U - theta[:, None], 0)

# Put the projection of max_nz_values to their original column indices
Expand Down
18 changes: 18 additions & 0 deletions test/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,24 @@ def nx(request):
yield backend


@pytest.fixture
def meta_default_device():
"""Make torch allocations without an explicit device land on ``meta``.

Inputs built with ``torch.from_numpy`` (or before the fixture is set up)
stay on CPU. A function that creates an internal buffer without
``type_as`` then mixes ``meta`` and CPU tensors and raises a device error,
as it would with GPU inputs, so this catches device bugs on CPU-only CI.
"""
torch = pytest.importorskip("torch")
previous = torch.get_default_device()
torch.set_default_device("meta")
try:
yield
finally:
torch.set_default_device(previous)


def skip_arg(arg, value, reason=None, getter=lambda x: x):
if isinstance(arg, (tuple, list)):
n = len(arg)
Expand Down
101 changes: 101 additions & 0 deletions test/test_device.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
"""Tests that solvers allocate their internal arrays on the input device"""

# License: MIT License

import numpy as np
import pytest

import ot
from ot.backend import torch


def _inputs():
# torch.from_numpy always returns CPU tensors, whatever the default device
rng = np.random.RandomState(0)
n, d = 6, 2
xs = torch.from_numpy(rng.randn(n, d))
xt = torch.from_numpy(rng.randn(n + 1, d))
a = ot.unif(n, type_as=xs)
b = ot.unif(n + 1, type_as=xs)
M = ot.dist(xs, xt)
M = M / M.max()
xs3 = torch.from_numpy(rng.randn(n, 3))
xt3 = torch.from_numpy(rng.randn(n + 1, 3))
return dict(
xs=xs,
xt=xt,
a=a,
b=b,
M=M,
C1=ot.dist(xs, xs),
C2=ot.dist(xt, xt),
u=torch.from_numpy(rng.rand(n)),
v=torch.from_numpy(rng.rand(n + 1)),
xs3=xs3 / xs3.norm(dim=1, keepdim=True),
xt3=xt3 / xt3.norm(dim=1, keepdim=True),
m=xs[:2],
mt=xt[:2],
C=torch.from_numpy(np.stack([np.eye(d), 2 * np.eye(d)])),
w=torch.from_numpy(np.array([0.3, 0.7])),
V=torch.from_numpy(rng.randn(3, 4)),
)


CASES = {
"emd": lambda i: ot.emd(i["a"], i["b"], i["M"]),
"sinkhorn": lambda i: ot.sinkhorn(i["a"], i["b"], i["M"], 1.0),
"sinkhorn_log": lambda i: ot.sinkhorn(
i["a"], i["b"], i["M"], 1.0, method="sinkhorn_log"
),
"sinkhorn_unbalanced": lambda i: ot.sinkhorn_unbalanced(
i["a"], i["b"], i["M"], 1.0, 1.0
),
"solve": lambda i: ot.solve(i["M"], i["a"], i["b"], reg=1.0).plan,
"solve_sample": lambda i: ot.solve_sample(i["xs"], i["xt"], i["a"], i["b"]).plan,
"solve_gromov": lambda i: ot.solve_gromov(i["C1"], i["C2"]).plan,
"partial_wasserstein": lambda i: ot.partial.partial_wasserstein(
i["a"], i["b"], i["M"], m=0.5
),
"wasserstein_1d": lambda i: ot.wasserstein_1d(i["u"], i["v"]),
"wasserstein_circle": lambda i: ot.wasserstein_circle(i["u"], i["v"]),
"sliced_wasserstein_sphere": lambda i: ot.sliced_wasserstein_sphere(
i["xs3"], i["xt3"], n_projections=5, seed=0
),
"empirical_bures_wasserstein_distance_hd": lambda i: (
ot.gaussian.empirical_bures_wasserstein_distance_hd(i["xs"], i["xt"], 1)
),
"empirical_bures_wasserstein_mapping_hd": lambda i: (
ot.gaussian.empirical_bures_wasserstein_mapping_hd(i["xs"], i["xt"], 1)
),
"gmm_pdf": lambda i: ot.gmm.gmm_pdf(i["xs"], i["m"], i["C"], i["w"]),
"gmm_ot_apply_map_bary": lambda i: ot.gmm.gmm_ot_apply_map(
i["xs"], i["m"], i["mt"], i["C"], i["C"], i["w"], i["w"], method="bary"
),
"gmm_ot_apply_map_rand": lambda i: ot.gmm.gmm_ot_apply_map(
i["xs"], i["m"], i["mt"], i["C"], i["C"], i["w"], i["w"], method="rand", seed=0
),
"gmm_ot_plan_density": lambda i: ot.gmm.gmm_ot_plan_density(
i["xs"], i["xt"], i["m"], i["mt"], i["C"], i["C"], i["w"], i["w"]
),
"lowrank_sinkhorn": lambda i: ot.lowrank_sinkhorn(
i["xs"], i["xt"], rank=2, init="deterministic"
),
"semirelaxed_gromov_barycenters": lambda i: (
ot.gromov.semirelaxed_gromov_barycenters(
3, [i["C1"], i["C2"]], max_iter=2, random_state=0
)
),
"projection_sparse_simplex": lambda i: ot.utils.projection_sparse_simplex(
i["V"], 2
),
}


@pytest.mark.skipif(not torch, reason="torch not installed")
@pytest.mark.parametrize("name", list(CASES))
def test_no_default_device_allocation(name, meta_default_device):
res = CASES[name](_inputs())
res = res if isinstance(res, (tuple, list)) else [res]
for r in res:
if torch.is_tensor(r):
assert r.device.type == "cpu"
Loading