Skip to content

[MRG] Check for device mismatch with a meta default device and fix bare allocations (#852) - #883

Open
raashish1601 wants to merge 2 commits into
PythonOT:masterfrom
raashish1601:test/852-device-mismatch
Open

raashish1601 wants to merge 2 commits into
PythonOT:masterfrom
raashish1601:test/852-device-mismatch

Conversation

@raashish1601

Copy link
Copy Markdown

Types of changes

Bug fix and tests.

Motivation and context / Related issue

Closes #852.

This adds the meta_default_device fixture suggested in the issue. It sets the torch default device to meta for the test and puts it back afterwards. Inputs are built with torch.from_numpy, so they stay on CPU, and any array a solver creates without type_as ends up on meta and raises a device error. That is the same failure you get with GPU inputs, but it shows up on CPU-only CI.

test/test_device.py runs this check on a list of functions. Besides a few solvers that already worked (emd, sinkhorn, unbalanced sinkhorn, ot.solve*, 1D), it found these bare allocations, which this PR fixes by passing type_as:

  • ot.partial.partial_wasserstein / partial_wasserstein2: zero blocks of the extended cost matrix
  • ot.binary_search_circle, so also ot.wasserstein_circle and ot.sliced_wasserstein_sphere: the done mask
  • ot.gaussian.bures_wasserstein_mapping_hd and bures_wasserstein_distance_hd: nx.eye(p)
  • ot.gmm.gmm_pdf, gmm_ot_apply_map (both methods) and gmm_ot_plan_density
  • ot.lowrank_sinkhorn with init="deterministic": the initial g
  • ot.gromov.semirelaxed_gromov_barycenters / semirelaxed_fgw_barycenters: default lambdas
  • ot.utils.projection_sparse_simplex: the index aranges and z. The index aranges use an integer array as type_as so they stay integer even if arange starts following the dtype of type_as (Backend.arange(..., type_as=...) silently ignores type_as for dtype (and device, in some backends) #864).

I did not cover every function in the library here. The fixture makes it easy to add more cases to the list later.

How has this been tested (if it applies)

The new test fails for 12 of the 20 cases on master and passes with the fix. I also ran test_utils.py, test_gmm.py, test_gaussian.py, test_partial.py, test_circle_solver.py, test_lowrank.py, test_ot.py, test_solvers.py, test/gromov and test/sliced with the numpy and torch backends (CPU), and pre-commit on the changed files.

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.

TST check for device mismatch in solvers that support GPU input

1 participant