# JAX wrapping and element-wise LOO

**URL:** <https://discourse.pymc.io/t/jax-wrapping-and-element-wise-loo/14533>\
**Category:** General\
**Created:** [June 3, 2024, 9:37am UTC](https://discourse.pymc.io/t/jax-wrapping-and-element-wise-loo/14533 "2024-06-03T09:37:29Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![frrfpp](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/frrfpp/32/7362_2.png) [@frrfpp](https://discourse.pymc.io/u/frrfpp)\
**Post date:** [June 3, 2024, 9:37am UTC](https://discourse.pymc.io/t/jax-wrapping-and-element-wise-loo/14533/1 "2024-06-03T09:37:29Z")

</div>

Hi,  
I’m following the [tutorial](https://www.pymc.io/projects/examples/en/latest/howto/wrapping_jax_function.html) on how to use JAX functions into pymc and I have a question on how to obtain pointwise log likelihoods for each observation.

I modified the code in this [section](https://www.pymc.io/projects/examples/en/latest/howto/wrapping_jax_function.html#sampling-with-pymc) as follows:

```python

def logp(emission_observed, emission_signal,
         emission_noise, logp_initial_state, logp_transition):
    return hmm_logp_op(
            emission_observed,
            emission_signal,
            emission_noise,
            logp_initial_state,
            logp_transition,
        )

with pm.Model() as model:
    emission_signal = pm.Normal("emission_signal", 0, 1)
    emission_noise = pm.HalfNormal("emission_noise", 1)

    p_initial_state = pm.Dirichlet("p_initial_state", np.ones(3))
    logp_initial_state = pt.log(p_initial_state)

    p_transition = pm.Dirichlet("p_transition", np.ones(3), size=3)
    logp_transition = pt.log(p_transition)

    # use DensityDist instead of Potential
    hmm_logp_dist = pm.DensityDist('hmm_logp_dist', emission_signal, emission_noise, logp_initial_state, logp_transition, 
                           logp=logp, observed=emission_observed)

with model:
    idata = pm.sample(chains=2, cores=1, idata_kwargs={"log_likelihood": True})

```

Sampling works correctly and the results match the ones in the tutorial.

But, when I run `az.loo(idata)` I get `UserWarning: The point-wise LOO is the same with the sum LOO, please double check the Observed RV in your model to make sure it returns element-wise logp.` even though the `emission_observed` array has 70 observations.

I tried to change the following function so that it sums across observations rather than summing everything together, but it still results in the point-wise LOO being the same as the sum LOO.

```python
def vec_hmm_logp(*args):
    vmap = jax.vmap(
        hmm_logp,
        # Only the first argument, needs to be vectorized
        in_axes=(0, None, None, None, None),
    )
    # For simplicity we sum across observations
    return jnp.sum(vmap(*args), axis=0)

```

Can you provide any help on this? What needs to change in the tutorial to get the element-wise loo values?  
Thanks!

---

<div class="post-metadata">

**Author:** ![frrfpp](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/frrfpp/32/7362_2.png) [@frrfpp](https://discourse.pymc.io/u/frrfpp)\
**Post date:** [June 19, 2024, 2:52pm UTC](https://discourse.pymc.io/t/jax-wrapping-and-element-wise-loo/14533/2 "2024-06-19T14:52:26Z")

</div>

Hi,  
I was wondering if anybody can provide some help with this?  
Thanks,  
Filippo

---

<div class="post-metadata">

**Author:** ![ricardoV94](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/ricardov94/32/5775_2.png) [@ricardoV94](https://discourse.pymc.io/u/ricardoV94)\
**Post date:** [June 19, 2024, 3:12pm UTC](https://discourse.pymc.io/t/jax-wrapping-and-element-wise-loo/14533/3 "2024-06-19T15:12:06Z")

</div>

Do you have a single chain with 70 observation or 70 chains with x observations? Technically speaking 70 observations of a single chain would belong to a single multivariate distribution and hence have a single non-decomposable logp.

You can think of them as conditional probabilities. I am not sure whether LOO cares/handles that distinction. @OriolAbril may be able to chime in

---

<div class="post-metadata">

**Author:** ![frrfpp](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/frrfpp/32/7362_2.png) [@frrfpp](https://discourse.pymc.io/u/frrfpp)\
**Post date:** [June 20, 2024, 9:48am UTC](https://discourse.pymc.io/t/jax-wrapping-and-element-wise-loo/14533/4 "2024-06-20T09:48:40Z")

</div>

My understanding was that, when you have n observations, elpd (y\_1,...y\_n) is different to elpd(y\_1) + ... + elpd(y\_n)

I guess that in the tutorial, since you are trying to estimate the parameters of a single HMM process it makes to have elpd (y\_1,...y\_n) when computing loo. I may be wrong though?
