Skip to content

Commit 0b366de

Browse files
committed
fix: raise errors for RNG inits
1 parent c892e04 commit 0b366de

2 files changed

Lines changed: 42 additions & 2 deletions

File tree

jax_galsim/random.py

Lines changed: 41 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
11
import secrets
22
from functools import partial
33

4+
import equinox
45
import galsim as _galsim
56
import jax
67
import jax.numpy as jnp
78
import jax.random as jrandom
9+
import numpy as np
810
from jax.tree_util import register_pytree_node_class
911

1012
from jax_galsim.core.utils import implements
@@ -122,9 +124,16 @@ def reset(self, seed=None):
122124
self._state = _DeviateState(
123125
wrap_key_data(jnp.array(seed, dtype=jnp.uint32))
124126
)
125-
else:
127+
elif (
128+
isinstance(seed, (int, jnp.ndarray, jax.Array, np.ndarray)) or seed is None
129+
):
126130
_initial_seed = seed or secrets.randbelow(2**31)
127131
self._state = _DeviateState(jrandom.key(_initial_seed))
132+
else:
133+
raise TypeError(
134+
"Seeds for BaseDeviate must be an int-like, str, tuple, or another BaseDeviate."
135+
f"Got seed {seed!r}."
136+
)
128137

129138
@property
130139
def _key(self):
@@ -295,6 +304,19 @@ def __str__(self):
295304
class GaussianDeviate(BaseDeviate):
296305
def __init__(self, seed=None, mean=0.0, sigma=1.0):
297306
super().__init__(seed=seed)
307+
308+
if isinstance(sigma, (int, float)):
309+
if sigma <= 0:
310+
raise ValueError(
311+
f"Gaussian deviates must have a positive sigma. Got {sigma!r}."
312+
)
313+
else:
314+
sigma = equinox.error_if(
315+
jnp.array(sigma),
316+
sigma <= 0,
317+
f"Gaussian deviates must have a positive sigma. Got {sigma!r}.",
318+
)
319+
298320
self._params["mean"] = mean
299321
self._params["sigma"] = sigma
300322

@@ -435,6 +457,19 @@ def __str__(self):
435457
class PoissonDeviate(BaseDeviate):
436458
def __init__(self, seed=None, mean=1.0):
437459
super().__init__(seed=seed)
460+
461+
if isinstance(mean, (int, float)):
462+
if mean < 0:
463+
raise ValueError(
464+
f"Poisson deviates must have a non-negative mean. Got {mean!r}."
465+
)
466+
else:
467+
mean = equinox.error_if(
468+
jnp.array(mean),
469+
mean < 0,
470+
f"Poisson deviates must have a non-negative mean. Got {mean!r}.",
471+
)
472+
438473
self._params["mean"] = mean
439474

440475
@property
@@ -484,6 +519,11 @@ def _generate_one(key, mean):
484519

485520
@implements(_galsim.PoissonDeviate.generate_from_expectation)
486521
def generate_from_expectation(self, array):
522+
array = equinox.error_if(
523+
jnp.array(array),
524+
jnp.any(jnp.array(array) < 0),
525+
"Poission deviates must have a non-negative mean.",
526+
)
487527
self._key, _array = self.__class__._generate_from_exp(self._key, array)
488528
return _array
489529

tests/GalSim

0 commit comments

Comments
 (0)