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 hierarchyx_observed (
Array) – Observed data with leading dimensions for trials/subjects/groupsconditions (
Array|None) – Optional experimental conditions matching x_observed shapetheta_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 potentialdims (
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", setprogressbar=Falseinpm.sample(). The BlackJAX progress bar uses IO callbacks insidejax.lax.cond, which are incompatible with thevmapused by vectorized chains.- Return type:
Potential- Returns:
PyMC Potential adding sum of log p(x_observed | theta, conditions) to model
- Parameters:
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, )