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
7 changes: 6 additions & 1 deletion meridian/backend/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -1072,6 +1072,7 @@ def _jax_adstock_process(
roll = _jax_roll
split = _jax_split
stack = _ops.stack
squeeze = _ops.squeeze
tile = _jax_tile
transpose = _jax_transpose
unique_with_counts = _jax_unique_with_counts
Expand Down Expand Up @@ -1254,6 +1255,7 @@ def _tf_adstock_process(
set_random_seed = tf_backend.keras.utils.set_random_seed
split = _ops.split
stack = _ops.stack
squeeze = _ops.squeeze
tile = _ops.tile
transpose = _ops.transpose
unique_with_counts = _tf_unique_with_counts
Expand Down Expand Up @@ -1379,7 +1381,10 @@ def __init__(self, seed: SeedType):
self._key: Optional["_jax.Array"] = None

if seed is None:
return
# Automatically generate a seed if none is provided, allowing JAX
# to function similarly to other backends where None is acceptable.
seed = np.random.randint(_MAX_INT32)
self._int_seed = seed

if (
isinstance(seed, jax.Array) # pylint: disable=undefined-variable
Expand Down
54 changes: 40 additions & 14 deletions meridian/backend/backend_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -1615,15 +1615,29 @@ def test_correct_class_is_exposed(self, backend_name):
self.assertIs(backend.RNGHandler, backend._TFRNGHandler)
# pylint: enable=protected-access

@parameterized.named_parameters(("tensorflow", _TF), ("jax", _JAX))
def test_initialization_with_none_seed_is_noop(self, backend_name):
"""Verifies that a None seed creates a handler that returns None."""
@parameterized.named_parameters(
dict(
testcase_name="tensorflow",
backend_name=_TF,
assert_fn_name="assertIsNone",
),
dict(
testcase_name="jax",
backend_name=_JAX,
assert_fn_name="assertIsNotNone",
),
)
def test_initialization_with_none_seed_is_noop(
self, backend_name, assert_fn_name
):
"""Verifies behavior when initialized with None."""
self._set_backend_for_test(backend_name)
handler = backend.RNGHandler(None)
assertion = getattr(self, assert_fn_name)

self.assertIsNone(handler._seed_input)
self.assertIsNone(handler.get_next_seed())
self.assertIsNone(handler.get_kernel_seed())
assertion(handler.get_next_seed())
assertion(handler.get_kernel_seed())

@parameterized.named_parameters(("tensorflow", _TF), ("jax", _JAX))
def test_initialization_with_integer_seed(self, backend_name):
Expand Down Expand Up @@ -1770,17 +1784,29 @@ def test_get_next_seed_is_reproducible(self, backend_name):
else:
test_utils.assert_allequal(s1, s2)

@parameterized.named_parameters(("tensorflow", _TF), ("jax", _JAX))
def test_advance_handler_with_none_seed(self, backend_name):
"""Tests that advancing a no-op handler produces another no-op handler."""
@parameterized.named_parameters(
dict(
testcase_name="tensorflow",
backend_name=_TF,
assert_fn_name="assertIsNone",
),
dict(
testcase_name="jax",
backend_name=_JAX,
assert_fn_name="assertIsNotNone",
),
)
def test_advance_handler_with_none_seed(self, backend_name, assert_fn_name):
"""Tests advancing a handler initialized with None."""
self._set_backend_for_test(backend_name)
handler = backend.RNGHandler(None)
new_handler = handler.advance_handler()
assertion = getattr(self, assert_fn_name)

self.assertIsNot(handler, new_handler)
self.assertIsNone(new_handler._seed_input)
self.assertIsNone(handler.get_next_seed())
self.assertIsNone(new_handler.get_kernel_seed())
assertion(new_handler._seed_input)
assertion(handler.get_next_seed())
assertion(new_handler.get_kernel_seed())

@parameterized.named_parameters(("tensorflow", _TF), ("jax", _JAX))
def test_advance_handler_provides_independent_handlers(self, backend_name):
Expand Down Expand Up @@ -1958,9 +1984,9 @@ def _get_test_model(self, dims=2):
tfd = backend.tfd
loc = backend.zeros(dims, dtype=backend.float32)
scale_diag = backend.ones(dims, dtype=backend.float32)
return tfd.JointDistributionNamed(
{"x": tfd.MultivariateNormalDiag(loc=loc, scale_diag=scale_diag)}
)
return tfd.JointDistributionNamed({
"x": lambda: tfd.MultivariateNormalDiag(loc=loc, scale_diag=scale_diag)
})

def _run_sampling(
self,
Expand Down
Loading
Loading