|
1 | 1 | import secrets |
2 | 2 | from functools import partial |
3 | 3 |
|
| 4 | +import equinox |
4 | 5 | import galsim as _galsim |
5 | 6 | import jax |
6 | 7 | import jax.numpy as jnp |
7 | 8 | import jax.random as jrandom |
| 9 | +import numpy as np |
8 | 10 | from jax.tree_util import register_pytree_node_class |
9 | 11 |
|
10 | 12 | from jax_galsim.core.utils import implements |
@@ -122,9 +124,16 @@ def reset(self, seed=None): |
122 | 124 | self._state = _DeviateState( |
123 | 125 | wrap_key_data(jnp.array(seed, dtype=jnp.uint32)) |
124 | 126 | ) |
125 | | - else: |
| 127 | + elif ( |
| 128 | + isinstance(seed, (int, jnp.ndarray, jax.Array, np.ndarray)) or seed is None |
| 129 | + ): |
126 | 130 | _initial_seed = seed or secrets.randbelow(2**31) |
127 | 131 | 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 | + ) |
128 | 137 |
|
129 | 138 | @property |
130 | 139 | def _key(self): |
@@ -295,6 +304,19 @@ def __str__(self): |
295 | 304 | class GaussianDeviate(BaseDeviate): |
296 | 305 | def __init__(self, seed=None, mean=0.0, sigma=1.0): |
297 | 306 | 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 | + |
298 | 320 | self._params["mean"] = mean |
299 | 321 | self._params["sigma"] = sigma |
300 | 322 |
|
@@ -435,6 +457,19 @@ def __str__(self): |
435 | 457 | class PoissonDeviate(BaseDeviate): |
436 | 458 | def __init__(self, seed=None, mean=1.0): |
437 | 459 | 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 | + |
438 | 473 | self._params["mean"] = mean |
439 | 474 |
|
440 | 475 | @property |
@@ -484,6 +519,11 @@ def _generate_one(key, mean): |
484 | 519 |
|
485 | 520 | @implements(_galsim.PoissonDeviate.generate_from_expectation) |
486 | 521 | 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 | + ) |
487 | 527 | self._key, _array = self.__class__._generate_from_exp(self._key, array) |
488 | 528 | return _array |
489 | 529 |
|
|
0 commit comments