diff --git a/RELEASES.md b/RELEASES.md index 98d226431..446e1c487 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -15,6 +15,7 @@ #### Closed issues +- `NumpyBackend.seed` adopts a `numpy.random.RandomState` instead of passing it to `RandomState.seed`, which rejected it (PR #881, 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) - Fix swapped arguments to `div_to_product` in `ot.gromov.fused_unbalanced_across_spaces_cost`: with `reg_type="independent"` (UCOOT) the entropic terms used the plan marginals as the reference measures and vice versa (PR #855, Issue #854) diff --git a/ot/backend.py b/ot/backend.py index fc087495c..118afc87b 100644 --- a/ot/backend.py +++ b/ot/backend.py @@ -1456,7 +1456,16 @@ def reshape(self, a, shape): return np.reshape(a, shape) def seed(self, seed=None): - if seed is not None: + r"""Set the random generator used by :meth:`rand` and :meth:`randn`. + + An integer seeds the current generator. A + :class:`numpy.random.RandomState` replaces it, matching the way the + torch and tensorflow backends adopt an external generator. ``None`` + leaves the generator unchanged. + """ + if isinstance(seed, np.random.RandomState): + self.rng_ = seed + elif seed is not None: self.rng_.seed(seed) def rand(self, *size, type_as=None): diff --git a/test/test_backend.py b/test/test_backend.py index c88ee5052..d3dbff996 100644 --- a/test/test_backend.py +++ b/test/test_backend.py @@ -982,3 +982,41 @@ def test_no_cuda_context_for_cpu_only_work(): f"interpreter exited with returncode {result.returncode}: " f"{result.stderr.decode(errors='replace')[-2000:]}" ) + + +def _random_state_equal(left, right): + return ( + left[0] == right[0] + and np.array_equal(left[1], right[1]) + and left[2:] == right[2:] + ) + + +def test_numpy_backend_adopts_random_state(): + # Non-regression for issue #848. RandomState.seed does not accept another + # RandomState, so the backend has to adopt the object. + shared = ot.backend.NumpyBackend.rng_ + shared_before = shared.get_state() + + adopted = np.random.RandomState(42) + expected = np.random.RandomState(42).rand(4) + backend = ot.backend.NumpyBackend() + backend.seed(adopted) + np.testing.assert_allclose(backend.rand(4), expected) + assert _random_state_equal(shared.get_state(), shared_before) + + continued = np.random.RandomState(7) + reference = np.random.RandomState(7) + backend.seed(continued) + np.testing.assert_allclose(backend.randn(2, 3), reference.randn(2, 3)) + + fresh = ot.backend.NumpyBackend() + fresh.seed(1) + first = fresh.rand(3) + fresh.seed(1) + second = fresh.rand(3) + np.testing.assert_allclose(first, second) + + unchanged = fresh.rng_.get_state() + fresh.seed(None) + assert _random_state_equal(fresh.rng_.get_state(), unchanged)