Skip to content

Commit 6d40f5a

Browse files
committed
fix: ensure we use any or all in calls for errors
1 parent df2a3d4 commit 6d40f5a

3 files changed

Lines changed: 15 additions & 11 deletions

File tree

jax_galsim/core/utils.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -45,11 +45,12 @@ def check_is_int_then_cast(val, msg):
4545
if val != int(val):
4646
raise TypeError(msg)
4747
val = int(val)
48-
elif not has_tracers(val):
48+
else:
4949
# otherwise we use more opaque checking upon jit via equinox
50+
val = jnp.array(val)
5051
val = equinox.error_if(
5152
val,
52-
val != jnp.trunc(val),
53+
np.any(val != jnp.trunc(val)),
5354
msg,
5455
)
5556
val = val.astype(int)

jax_galsim/moffat.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -74,9 +74,10 @@ def __init__(
7474
f"JAX-GalSim does not support Moffat beta values <= {self._beta_thr}."
7575
)
7676
elif not has_tracers(beta):
77+
beta = jnp.array(beta)
7778
beta = equinox.error_if(
78-
jnp.array(beta),
79-
beta <= self._beta_thr,
79+
beta,
80+
jnp.any(beta <= self._beta_thr),
8081
f"JAX-GalSim does not support Moffat beta values <= {self._beta_thr}.",
8182
)
8283

jax_galsim/random.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import numpy as np
1010
from jax.tree_util import register_pytree_node_class
1111

12-
from jax_galsim.core.utils import has_tracers, implements
12+
from jax_galsim.core.utils import implements
1313

1414
try:
1515
from jax.extend.random import wrap_key_data
@@ -313,10 +313,11 @@ def __init__(self, seed=None, mean=0.0, sigma=1.0):
313313
raise ValueError(
314314
f"Gaussian deviates must have a non-negative sigma. Got {sigma!r}."
315315
)
316-
elif not has_tracers(sigma):
316+
else:
317+
sigma = jnp.array(sigma)
317318
sigma = equinox.error_if(
318-
jnp.array(sigma),
319-
sigma < 0,
319+
sigma,
320+
jnp.any(sigma < 0),
320321
f"Gaussian deviates must have a non-negative sigma. Got {sigma!r}.",
321322
)
322323

@@ -466,10 +467,11 @@ def __init__(self, seed=None, mean=1.0):
466467
raise ValueError(
467468
f"Poisson deviates must have a non-negative mean. Got {mean!r}."
468469
)
469-
elif not has_tracers(mean):
470+
else:
471+
mean = jnp.array(mean)
470472
mean = equinox.error_if(
471-
jnp.array(mean),
472-
mean < 0,
473+
mean,
474+
jnp.any(mean < 0),
473475
f"Poisson deviates must have a non-negative mean. Got {mean!r}.",
474476
)
475477

0 commit comments

Comments
 (0)