exoplanet-dev / exoplanet-dev/celerite2

celerite2 fails with latest numpyro

Open
#157 4 comments 0 reactions 0 assignees View on GitHub
Dominant language
C++
Stars
85
Forks
19
PR merge metrics
No merged PRs in 30d

Description

Installing a fresh copy of celerite2 in a fresh python env using numpyro and jax on my ARM64 macbook results in a failure of the simple GP model
```python
def test_gp(t, y, yerr):
mu = numpyro.sample('mu', dist.Normal(0, 10))
kernel = terms.UnderdampedSHOTerm(w0=2*np.pi, Q=10, sigma=1.0)
gp = GaussianProcess(kernel)
gp.compute(t, yerr=yerr, check_sorted=False)
numpyro.factor('log_likelihood', gp.log_likelihood(y-mu))

kernel = NUTS(test_gp)
mcmc = MCMC(kernel, num_warmup=1000, num_samples=1000)
rng_key = jrn.PRNGKey(np.random.randint(1<<32))
mcmc.run(rng_key, np.arange(100, dtype=np.float64), np.zeros(100), np.ones(100))
```

With the cryptic error message
```
Traceback (most recent call last):
File "", line 4, in
mcmc.run(rng_key, np.arange(100, dtype=np.float64), np.zeros(100), np.ones(100))
~~~~~~~~^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/mcmc.py", line 702, in run
states_flat, last_state = partial_map_fn(map_args)
~~~~~~~~~~~~~~^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/mcmc.py", line 465, in _single_chain_mcmc
new_init_state = self.sampler.init(
rng_key,
...<3 lines>...
model_kwargs=kwargs,
)
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/hmc.py", line 749, in init
init_params = self._init_state(
rng_key_init_model, model_args, model_kwargs, init_params
)
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/hmc.py", line 693, in _init_state
) = initialize_model(
~~~~~~~~~~~~~~~~^
rng_key,
^^^^^^^^
...<5 lines>...
forward_mode_differentiation=self._forward_mode_differentiation,
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
)
^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 750, in initialize_model
(init_params, pe, grad), is_valid = find_valid_initial_params(
~~~~~~~~~~~~~~~~~~~~~~~~~^
rng_key,
^^^^^^^^
...<14 lines>...
validate_grad=validate_grad,
^^^^^^^^^^^^^^^^^^^^^^^^^^^^
)
^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 472, in find_valid_initial_params
(init_params, pe, z_grad), is_valid = _find_valid_params(
~~~~~~~~~~~~~~~~~~^
rng_key, exit_early=True
^^^^^^^^^^^^^^^^^^^^^^^^
)
^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 458, in _find_valid_params
_, _, (init_params, pe, z_grad), is_valid = init_state = body_fn(init_state)
~~~~~~~^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 442, in body_fn
pe, z_grad = value_and_grad(potential_fn)(params)
~~~~~~~~~~~~~~~~~~~~~~~~~~~~^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 324, in potential_energy
log_joint, model_trace = log_density_(
~~~~~~~~~~~~^
substituted_model, model_args, model_kwargs, {}
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
)
^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 120, in log_density
log_joint, model_trace = compute_log_probs(model, model_args, model_kwargs, params)
~~~~~~~~~~~~~~~~~^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/infer/util.py", line 78, in compute_log_probs
model_trace = trace(model).get_trace(*model_args, **model_kwargs)
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/handlers.py", line 191, in get_trace
self(*args, **kwargs)
~~~~^^^^^^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/primitives.py", line 121, in __call__
return self.fn(*args, **kwargs)
~~~~~~~^^^^^^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/primitives.py", line 121, in __call__
return self.fn(*args, **kwargs)
~~~~~~~^^^^^^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/numpyro/primitives.py", line 121, in __call__
return self.fn(*args, **kwargs)
~~~~~~~^^^^^^^^^^^^^^^^^
[Previous line repeated 3 more times]
File "", line 6, in test_gp
numpyro.factor('log_likelihood', gp.log_likelihood(y-mu))
~~~~~~~~~~~~~~~~~^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/celerite2/core.py", line 428, in log_likelihood
return self._norm - 0.5 * self._do_norm(y - self._mean_value)
~~~~~~~~~~~~~^^^^^^^^^^^^^^^^^^^^^^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/celerite2/jax/celerite2.py", line 55, in _do_norm
alpha = ops.solve_lower(
~~~~~~~~~~~~~~~^
self._t, self._c, self._U, self._W, y[:, None]
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
)[:, 0]
^
File "/Users/wfarr/miniconda3/envs/celerite2-test/lib/python3.13/site-packages/celerite2/jax/ops.py", line 44, in solve_lower
Z, F = solve_lower_p.bind(t, c, U, W, Y)
~~~~~~~~~~~~~~~~~~^^^^^^^^^^^^^^^
ValueError: Invalid argument to dtype: None.
--------------------
For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these.
```

(The tutorial notebook in the docs also seems to have a similar error when it tries to show the numpyro sampling, so I don't think this is something specific to my setup.)

To my in-expert eye, this looks like maybe numpyro has changed the way it chooses an initialization point and somehow this has modified the concrete values that are passed to the various jax functions, causing this failure? Lmk if you need more information from me about my setup, or if you want me to test modifications of the code / setup.

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.