Skip to content

Commit 29f83f3

Browse files
authored
fix: add type checking to seed method
Add type checking for seed parameter in the seed method.
1 parent 2610daa commit 29f83f3

1 file changed

Lines changed: 15 additions & 5 deletions

File tree

jax_galsim/random.py

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -91,12 +91,22 @@ def has_reliable_discard(self):
9191
def generates_in_pairs(self):
9292
return False
9393

94-
@implements(
95-
_galsim.BaseDeviate.seed,
96-
lax_description="The JAX version of this method does no type checking.",
97-
)
94+
@implements(_galsim.BaseDeviate.seed)
9895
def seed(self, seed=None):
99-
self._seed(seed=seed)
96+
if seed is None:
97+
self._seed(seed=seed)
98+
elif isinstance(seed, (int, float, np.integer, np.floating)):
99+
if seed == int(seed):
100+
self._seed(seed=int(seed))
101+
else:
102+
raise TypeError(f"BaseDeviate seed must be an integer. Got {seed!r}.")
103+
else:
104+
seed = equinox.error_if(
105+
seed,
106+
jnp.any(seed != jnp.trunc(seed)),
107+
"BaseDeviate seed must be an integer.",
108+
)
109+
self._seed(seed=seed)
100110

101111
@implements(_galsim.BaseDeviate._seed)
102112
def _seed(self, seed=None):

0 commit comments

Comments
 (0)