From 69278d5b6658415833e9938145a6b5c8989e1a3b Mon Sep 17 00:00:00 2001 From: Raashish Aggarwal <94279692+raashish1601@users.noreply.github.com> Date: Tue, 6 Oct 2026 22:25:57 +0530 Subject: [PATCH 1/2] TST check for device mismatch in solvers and fix default-device allocations (#852) --- RELEASES.md | 1 + ot/gaussian.py | 10 ++-- ot/gmm.py | 18 +++--- ot/gromov/_semirelaxed.py | 4 +- ot/lowrank.py | 2 +- ot/lp/solver_circle.py | 2 +- ot/partial/partial_solvers.py | 8 ++- ot/utils.py | 8 +-- test/conftest.py | 18 ++++++ test/test_device.py | 101 ++++++++++++++++++++++++++++++++++ 10 files changed, 149 insertions(+), 23 deletions(-) create mode 100644 test/test_device.py diff --git a/RELEASES.md b/RELEASES.md index 860f2babb..82fec0ba4 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -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 (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) diff --git a/ot/gaussian.py b/ot/gaussian.py index 103d102d8..5f4e19b95 100644 --- a/ot/gaussian.py +++ b/ot/gaussian.py @@ -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)) @@ -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) diff --git a/ot/gmm.py b/ot/gmm.py index 3cb6cea6a..fde2528aa 100644 --- a/ot/gmm.py +++ b/ot/gmm.py @@ -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 @@ -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)] ) @@ -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)): @@ -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 @@ -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) diff --git a/ot/gromov/_semirelaxed.py b/ot/gromov/_semirelaxed.py index d28939b51..440d3228c 100644 --- a/ot/gromov/_semirelaxed.py +++ b/ot/gromov/_semirelaxed.py @@ -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) @@ -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) diff --git a/ot/lowrank.py b/ot/lowrank.py index fef964949..228e980bd 100644 --- a/ot/lowrank.py +++ b/ot/lowrank.py @@ -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) diff --git a/ot/lp/solver_circle.py b/ot/lp/solver_circle.py index af118d1a7..a5c75849f 100644 --- a/ot/lp/solver_circle.py +++ b/ot/lp/solver_circle.py @@ -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): diff --git a/ot/partial/partial_solvers.py b/ot/partial/partial_solvers.py index f6e1c65a6..637accedd 100755 --- a/ot/partial/partial_solvers.py +++ b/ot/partial/partial_solvers.py @@ -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, diff --git a/ot/utils.py b/ot/utils.py index 8a40515b1..e05d39318 100644 --- a/ot/utils.py +++ b/ot/utils.py @@ -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 diff --git a/test/conftest.py b/test/conftest.py index a03a15873..bccc5d7f5 100644 --- a/test/conftest.py +++ b/test/conftest.py @@ -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) diff --git a/test/test_device.py b/test/test_device.py new file mode 100644 index 000000000..5a6cd3922 --- /dev/null +++ b/test/test_device.py @@ -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" From 5ff0a9614663e0c346d8fdcef00813a379fa6c68 Mon Sep 17 00:00:00 2001 From: Raashish Aggarwal <94279692+raashish1601@users.noreply.github.com> Date: Tue, 6 Oct 2026 22:39:21 +0530 Subject: [PATCH 2/2] DOC add PR number to RELEASES entry --- RELEASES.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/RELEASES.md b/RELEASES.md index 82fec0ba4..b4d447ac7 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -15,7 +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 (Issue #852) +- 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)