setu.to_pymc_hierarchical#

setu.to_pymc_hierarchical(nle, theta, x_observed, *, conditions=None, theta_shape=None, chunk_size=None, mask=None, trial_dim=None, subject_dim=None, group_dim=None, name='nle_likelihood_hierarchical', dims=None)[source]#

Convert trained likelihood estimator (NLE or MixedNLE) to PyMC hierarchical likelihood.

Automatically handles 2-level or 3-level hierarchies via nested vmap.

Parameters:
  • nle (LikelihoodEstimator) – Trained NLE or MixedNLE instance (trained on SINGLE subject data)

  • theta (TensorVariable) – PyMC parameter variable with dims matching hierarchy

  • x_observed (Array) – Observed data with leading dimensions for trials/subjects/groups

  • conditions (Array | None) – Optional experimental conditions matching x_observed shape

  • theta_shape (tuple[int, ...] | None) – Override for theta shape. Required when theta uses dims instead of explicit shape (theta.type.shape contains None values). When not provided, inferred from theta.type.shape or x_observed.

  • chunk_size (int | None) – If provided, process trials in chunks of this size to reduce peak memory usage.

  • mask (Array | None) – Optional boolean/float mask for ragged data. Shape must match x_observed leading dims: (subjects, trials) for 2-level, (groups, subjects, trials) for 3-level. Where mask is 0/False, that trial’s logp contribution is zeroed out.

  • trial_dim (str | None) – Name of trial dimension (innermost level)

  • subject_dim (str | None) – Name of subject dimension (middle level)

  • group_dim (str | None) – Name of group dimension (outermost level, optional)

  • name (str) – Name for the PyMC potential

  • dims (tuple[str, ...] | None) – PyMC dimension names for the potential (for coords)

Return type:

Potential

Shape Conventions:
2-level (subject-trial):

theta: (num_subjects, theta_dim) with dims=(subject_dim,) x_observed: (num_subjects, num_trials, x_dim) conditions: (num_subjects, num_trials, condition_dim) mask: (num_subjects, num_trials)

3-level (group-subject-trial):

theta: (num_groups, num_subjects, theta_dim) x_observed: (num_groups, num_subjects, num_trials, x_dim) conditions: (num_groups, num_subjects, num_trials, condition_dim) mask: (num_groups, num_subjects, num_trials)

Note

When sampling with BlackJAX using chain_method="vectorized", set progressbar=False in pm.sample(). The BlackJAX progress bar uses IO callbacks inside jax.lax.cond, which are incompatible with the vmap used by vectorized chains.

Return type:

Potential

Returns:

PyMC Potential adding sum of log p(x_observed | theta, conditions) to model

Parameters:
  • nle (LikelihoodEstimator)

  • theta (TensorVariable)

  • x_observed (Array)

  • conditions (Array | None)

  • theta_shape (tuple[int, ...] | None)

  • chunk_size (int | None)

  • mask (Array | None)

  • trial_dim (str | None)

  • subject_dim (str | None)

  • group_dim (str | None)

  • name (str)

  • dims (tuple[str, ...] | None)

Example (2-level with ragged data):

# Subject 0 has 10 trials, subject 1 has 7 trials
# Pad x_observed to (2, 10, x_dim), mask out padding
mask = jnp.ones((2, 10))
mask = mask.at[1, 7:].set(0)  # Zero out trials 7-9 for subject 1
likelihood = to_pymc_hierarchical(
    nle,
    theta,
    x_observed,
    subject_dim="subject",
    mask=mask,
)