Skip to content

Commit 3d645a2

Browse files
committed
fix: apparently this does not work on tracers
1 parent 15e1398 commit 3d645a2

2 files changed

Lines changed: 25 additions & 3 deletions

File tree

jax_galsim/random.py

Lines changed: 3 additions & 3 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 implements
12+
from jax_galsim.core.utils import has_tracers, implements
1313

1414
try:
1515
from jax.extend.random import wrap_key_data
@@ -313,7 +313,7 @@ def __init__(self, seed=None, mean=0.0, sigma=1.0):
313313
raise ValueError(
314314
f"Gaussian deviates must have a positive sigma. Got {sigma!r}."
315315
)
316-
else:
316+
elif not has_tracers(sigma):
317317
sigma = equinox.error_if(
318318
jnp.array(sigma),
319319
sigma <= 0,
@@ -466,7 +466,7 @@ def __init__(self, seed=None, mean=1.0):
466466
raise ValueError(
467467
f"Poisson deviates must have a non-negative mean. Got {mean!r}."
468468
)
469-
else:
469+
elif not has_tracers(mean):
470470
mean = equinox.error_if(
471471
jnp.array(mean),
472472
mean < 0,

tests/jax/test_random_jax.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
import jax
2+
import pytest
3+
4+
import jax_galsim
5+
6+
7+
def test_random_jax_gaussian_pos_sigma_jit():
8+
@jax.jit
9+
def _make_gauss(sigma):
10+
return jax_galsim.GaussianDeviate(seed=10, sigma=sigma)
11+
12+
with pytest.raises(Exception):
13+
_make_gauss(-1.0)
14+
15+
@jax.jit
16+
def _make_gauss(sigma):
17+
return jax_galsim.GaussianDeviate(seed=10, sigma=sigma)
18+
19+
_make_gauss(1.0)
20+
21+
with pytest.raises(Exception):
22+
_make_gauss(-1)

0 commit comments

Comments
 (0)