Skip to content
Closed
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

- `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)
Expand Down
11 changes: 10 additions & 1 deletion ot/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
38 changes: 38 additions & 0 deletions test/test_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading