"""Discrete Neural Likelihood Estimator for categorical and count data."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
if TYPE_CHECKING:
from setu.validation.checks import ValidationConfig, ValidationResult
from setu.validation.data import ValidationDataset
import equinox as eqx
import jax
import jax.numpy as jnp
from jax import Array
from setu.categorical_made import CategoricalMADE, create_categorical_made
from setu.count_made import CountMADE, create_count_made
from setu.data import SimulationDataset, _validate_fit_inputs
from setu.transforms import Standardizer
DiscreteDistribution = Literal["categorical", "zinb", "negbinom", "poisson"]
[docs]
@dataclass
class DiscreteNLETrainingResult:
"""Result of DiscreteNLE training.
Attributes:
discrete_nle: Trained DiscreteNLE instance with fitted transforms.
losses: Training loss per epoch.
val_losses: Validation loss per epoch.
best_epoch: Epoch with lowest validation loss.
"""
discrete_nle: DiscreteNLE
losses: Array
val_losses: Array
best_epoch: int
[docs]
def validate(
self,
val_data: ValidationDataset,
key: Array,
*,
checks: list[str] | None = None,
strict: bool = False,
config: ValidationConfig | None = None,
) -> ValidationResult:
"""Run validation checks on the trained estimator."""
from setu.validation import validate_nle
result: DiscreteNLETrainingResult = self
return validate_nle(
self.discrete_nle,
val_data,
key,
training_result=result,
checks=checks,
strict=strict,
config=config,
)
[docs]
class DiscreteNLE(eqx.Module):
"""Discrete Neural Likelihood Estimator for categorical or count data.
Handles pure discrete observations via CategoricalMADE (categorical data)
or CountMADE (count data with ZINB/NegBin/Poisson conditionals).
Use create_discrete_nle() to create and fit_discrete_nle() to train.
"""
discrete_net: CategoricalMADE | CountMADE
theta_transform: Standardizer | None = None
condition_transform: Standardizer | None = None
x_dim: int = 0
condition_dim: int = 0
distribution: str = "categorical"
num_categories: Array | None = None # Only for categorical
trained: bool = False
[docs]
def log_prob(
self, x: Array, theta: Array, conditions: Array | None = None
) -> Array:
"""Compute log p(x | theta) or log p(x | theta, conditions).
Args:
x: Discrete observation (x_dim,).
theta: Parameters (theta_dim,).
conditions: Experimental conditions (condition_dim,).
Returns:
Log probability as scalar.
Raises:
ValueError: If not trained or conditions mismatch.
"""
if not self.trained:
raise ValueError(
"DiscreteNLE has not been trained yet. Call fit_discrete_nle() first."
)
# Transform theta
if self.theta_transform is not None:
theta = self.theta_transform.transform(theta)
# Build context and delegate to net (handles unbatched input natively)
context = self._build_context(theta, conditions)
return self.discrete_net.log_prob(x, context)
[docs]
def sample(
self,
theta: Array,
key: Array,
conditions: Array | None = None,
*,
n_samples: int,
) -> Array:
"""Sample from p(x | theta).
Args:
theta: Single parameter vector (theta_dim,).
key: JAX random key.
conditions: Single condition vector (condition_dim,).
n_samples: Number of samples to draw.
Returns:
Samples (n_samples, x_dim).
Raises:
ValueError: If not trained or conditions mismatch.
"""
if not self.trained:
raise ValueError(
"DiscreteNLE has not been trained yet. Call fit_discrete_nle() first."
)
# Transform theta
if self.theta_transform is not None:
theta = self.theta_transform.transform(theta)
# Build context and broadcast across the sample batch (a view, no copy)
context = self._build_context(theta, conditions)
context_batch = jnp.broadcast_to(context, (n_samples, context.shape[-1]))
return self.discrete_net.sample(context_batch, key)
def _build_context(self, theta: Array, conditions: Array | None = None) -> Array:
"""Build context vector from theta and optional conditions."""
if self.condition_dim > 0:
if conditions is None:
raise ValueError(
f"DiscreteNLE was trained with conditions "
f"(condition_dim={self.condition_dim}). "
f"Must provide conditions argument."
)
if self.condition_transform is not None:
conditions = self.condition_transform.transform(conditions)
return jnp.concatenate([theta, conditions], axis=-1)
else:
if conditions is not None:
raise ValueError(
"DiscreteNLE was trained without conditions. "
"Do not provide conditions argument."
)
return theta
[docs]
def to_pymc(
self,
theta,
x_observed,
*,
conditions: Array | None = None,
chunk_size: int | None = None,
name: str = "discrete_nle_likelihood",
dims: tuple[str, ...] | None = None,
):
"""Convert to PyMC likelihood.
Args:
theta: PyMC parameter variable.
x_observed: Observed data (n_obs, x_dim) or (x_dim,).
conditions: Optional experimental conditions.
chunk_size: If provided, process observations in chunks.
name: Name for the PyMC potential.
dims: PyMC dimension names.
Returns:
PyMC Potential.
"""
from setu.integration.pymc import to_pymc
return to_pymc(
self,
theta,
x_observed,
conditions=conditions,
chunk_size=chunk_size,
name=name,
dims=dims,
)
[docs]
def to_numpyro(
self,
theta,
x_observed,
*,
conditions: Array | None = None,
name: str = "discrete_nle_likelihood",
):
"""Convert to NumPyro likelihood factor.
Args:
theta: JAX array from numpyro.sample().
x_observed: Observed data (n_obs, x_dim) or (x_dim,).
conditions: Optional conditions.
name: Name for the NumPyro factor site.
"""
from setu.integration.numpyro import to_numpyro
return to_numpyro(self, theta, x_observed, conditions=conditions, name=name)
[docs]
def to_pymc_hierarchical(
self,
theta,
x_observed,
*,
conditions: Array | None = None,
trial_dim: str | None = None,
subject_dim: str | None = None,
group_dim: str | None = None,
name: str = "discrete_nle_likelihood_hierarchical",
theta_shape: tuple[int, ...] | None = None,
dims: tuple[str, ...] | None = None,
chunk_size: int | None = None,
mask: Array | None = None,
):
"""Convert to PyMC hierarchical likelihood.
Automatically handles 2-level or 3-level hierarchies via nested vmap.
Args:
theta: PyMC parameter variable with dims matching hierarchy.
x_observed: Observed data ([groups], subjects, trials, x_dim).
conditions: Optional conditions matching x_observed leading dims.
trial_dim: Name of trial dimension (innermost level).
subject_dim: Name of subject dimension (middle level).
group_dim: Name of group dimension (outermost level, optional).
name: Name for the PyMC potential.
theta_shape: Explicit theta shape when theta.type.shape contains
None (e.g. when using PyMC dims). Inferred from x_observed
if not given.
dims: PyMC dimension names (for coords).
chunk_size: If set, evaluate log-prob in chunks to reduce memory.
mask: Boolean array matching x_observed leading dims. Where False,
trial logp is zeroed (for ragged/unbalanced data).
Returns:
PyMC Potential.
"""
from setu.integration.pymc import to_pymc_hierarchical
return to_pymc_hierarchical(
self,
theta,
x_observed,
conditions=conditions,
trial_dim=trial_dim,
subject_dim=subject_dim,
group_dim=group_dim,
name=name,
theta_shape=theta_shape,
dims=dims,
chunk_size=chunk_size,
mask=mask,
)
[docs]
def to_numpyro_hierarchical(
self,
theta,
x_observed,
*,
conditions: Array | None = None,
mask: Array | None = None,
name: str = "discrete_nle_likelihood_hierarchical",
):
"""Convert to NumPyro hierarchical likelihood factor.
Args:
theta: Parameters from numpyro.sample().
2-level: shape (n_subjects, theta_dim).
3-level: shape (n_groups, n_subjects, theta_dim).
x_observed: Observed data.
2-level: shape (n_subjects, n_trials, x_dim).
3-level: shape (n_groups, n_subjects, n_trials, x_dim).
conditions: Optional conditions matching x_observed leading dims.
mask: Boolean mask matching x_observed leading dims.
name: Name for the NumPyro factor site.
"""
from setu.integration.numpyro import to_numpyro_hierarchical
return to_numpyro_hierarchical(
self,
theta,
x_observed,
conditions=conditions,
mask=mask,
name=name,
)
[docs]
def create_discrete_nle(
x_dim: int,
theta_dim: int,
*,
distribution: DiscreteDistribution = "categorical",
num_categories: Array | None = None,
condition_dim: int = 0,
hidden_features: int = 64,
num_layers: int = 2,
dropout_rate: float = 0.0,
key: Array,
) -> DiscreteNLE:
"""Create a DiscreteNLE for categorical or count data.
Args:
x_dim: Number of discrete variables.
theta_dim: Dimension of parameters.
distribution: Distribution family. "categorical" for categorical data
(requires num_categories). "zinb", "negbinom", or "poisson" for
count data.
num_categories: Categories per variable (x_dim,). Required for
distribution="categorical".
condition_dim: Dimension of experimental conditions. 0 if none.
hidden_features: Hidden units per layer in MADE network.
num_layers: Number of hidden layers.
dropout_rate: Dropout rate for regularization.
key: JAX random key.
Returns:
Untrained DiscreteNLE instance.
Raises:
ValueError: If distribution="categorical" but num_categories not provided.
ValueError: If distribution is a count type but num_categories provided.
ValueError: If distribution is not recognized.
"""
context_dim = theta_dim + condition_dim
if distribution == "categorical":
if num_categories is None:
raise ValueError(
"num_categories is required for distribution='categorical'. "
"Provide an array of category counts per variable, "
"e.g. num_categories=jnp.array([3, 5, 2])."
)
num_categories = jnp.asarray(num_categories)
if len(num_categories) != x_dim:
raise ValueError(
f"num_categories length ({len(num_categories)}) must equal "
f"x_dim ({x_dim})."
)
discrete_net = create_categorical_made(
num_variables=x_dim,
num_categories=num_categories,
theta_dim=context_dim,
hidden_features=hidden_features,
num_layers=num_layers,
dropout_rate=dropout_rate,
key=key,
)
elif distribution in ("zinb", "negbinom", "poisson"):
if num_categories is not None:
raise ValueError(
f"num_categories should not be provided for "
f"distribution={distribution!r}. "
f"num_categories is only used for distribution='categorical'."
)
discrete_net = create_count_made(
num_variables=x_dim,
theta_dim=context_dim,
distribution=distribution,
hidden_features=hidden_features,
num_layers=num_layers,
dropout_rate=dropout_rate,
key=key,
)
else:
raise ValueError(
f"Unknown distribution: {distribution!r}. "
f"Expected 'categorical', 'zinb', 'negbinom', or 'poisson'."
)
return DiscreteNLE(
discrete_net=discrete_net,
x_dim=x_dim,
condition_dim=condition_dim,
distribution=distribution,
num_categories=num_categories,
)
[docs]
def fit_discrete_nle(
discrete_nle: DiscreteNLE,
dataset: SimulationDataset,
*,
learning_rate: float = 1e-3,
weight_decay: float = 1e-4,
max_epochs: int = 1000,
max_patience: int = 50,
batch_size: int = 256,
val_fraction: float = 0.1,
show_progress: bool = True,
epoch_callback: Callable[[int, float, float], None] | None = None,
key: Array,
) -> DiscreteNLETrainingResult:
"""Train DiscreteNLE on simulation data.
Uses early stopping based on validation loss. No x-standardization is
applied (discrete data should not be z-scored).
Args:
discrete_nle: Untrained DiscreteNLE from create_discrete_nle().
dataset: Training data with theta-x pairs.
learning_rate: AdamW optimizer learning rate.
weight_decay: L2 regularization weight.
max_epochs: Maximum training epochs.
max_patience: Epochs without improvement before early stopping.
batch_size: Mini-batch size.
val_fraction: Fraction of data for validation.
show_progress: Show tqdm progress bar.
epoch_callback: Called after each epoch with (epoch, train_loss, val_loss).
key: JAX random key.
Returns:
DiscreteNLETrainingResult with trained estimator and losses.
"""
import optax
from setu._made import partition_frozen_masks
from setu._training import train_with_early_stopping
_validate_fit_inputs(
estimator_name="DiscreteNLE",
dataset=dataset,
val_fraction=val_fraction,
x_dim=discrete_nle.x_dim,
condition_dim=discrete_nle.condition_dim,
)
# Validate discrete data
x = dataset.x
if discrete_nle.distribution == "categorical":
# Check integer-like values
if not jnp.allclose(x, jnp.round(x), atol=1e-6):
raise ValueError(
"x must contain integer-like values for categorical distribution."
)
# Check values in range
num_cats = discrete_nle.num_categories
assert num_cats is not None # guaranteed by create_discrete_nle
for i in range(discrete_nle.x_dim):
max_val = float(num_cats[i]) - 1
col = x[:, i]
if float(jnp.min(col)) < 0 or float(jnp.max(col)) > max_val:
raise ValueError(
f"Variable {i}: values must be in [0, {int(max_val)}] "
f"(num_categories={int(num_cats[i])}), "
f"got range [{float(jnp.min(col))}, {float(jnp.max(col))}]."
)
else:
# Count data: non-negative integer-like
if not jnp.allclose(x, jnp.round(x), atol=1e-6):
raise ValueError(
"x must contain integer-like values for count distribution."
)
if float(jnp.min(x)) < 0:
raise ValueError(
f"x must contain non-negative values for count distribution. "
f"Got min={float(jnp.min(x))}."
)
# Fit transforms (theta and conditions only, not x)
theta_transform = Standardizer.fit(dataset.theta)
theta_z = theta_transform.transform(dataset.theta)
condition_transform = None
if dataset.conditions is not None:
condition_transform = Standardizer.fit(dataset.conditions)
conditions_z = condition_transform.transform(dataset.conditions)
context = jnp.concatenate([theta_z, conditions_z], axis=-1)
else:
context = theta_z
# Update DiscreteNLE with transforms
nle_with_transforms = eqx.tree_at(
lambda m: (m.theta_transform, m.condition_transform),
discrete_nle,
(theta_transform, condition_transform),
is_leaf=lambda node: node is None,
)
# Split into trainable and static parts; freeze masks to preserve the
# autoregressive property
nle_trainable, nle_static = eqx.partition(nle_with_transforms, eqx.is_inexact_array)
nle_trainable, nle_static = partition_frozen_masks(
nle_trainable, nle_static, nle_with_transforms
)
# Loss functions
@eqx.filter_jit
def _loss(nle_trainable, nle_static, x_batch, context_batch, key, training):
nle_combined = eqx.combine(nle_trainable, nle_static)
log_probs = nle_combined.discrete_net.log_prob(
x_batch, context_batch, key=key, inference=not training
)
return -jnp.mean(log_probs)
def train_loss_fn(trainable, static, batch, key):
x_batch, context_batch = batch
return _loss(trainable, static, x_batch, context_batch, key, training=True)
def val_loss_fn(trainable, static, batch, key):
x_batch, context_batch = batch
return _loss(
trainable, static, x_batch, context_batch, key=None, training=False
)
# Split data
n_samples = len(dataset)
n_val = int(n_samples * val_fraction)
n_train = n_samples - n_val
key, shuffle_key = jax.random.split(key)
perm = jax.random.permutation(shuffle_key, n_samples)
x_shuffled = x[perm]
context_shuffled = context[perm]
x_train, x_val = x_shuffled[:n_train], x_shuffled[n_train:]
context_train, context_val = (
context_shuffled[:n_train],
context_shuffled[n_train:],
)
optimizer = optax.adamw(learning_rate, weight_decay=weight_decay)
key, train_key = jax.random.split(key)
result = train_with_early_stopping(
trainable=nle_trainable,
static=nle_static,
train_data=(x_train, context_train),
val_data=(x_val, context_val),
loss_fn=train_loss_fn,
val_loss_fn=val_loss_fn,
optimizer=optimizer,
batch_size=batch_size,
max_epochs=max_epochs,
max_patience=max_patience,
show_progress=show_progress,
epoch_callback=epoch_callback,
key=train_key,
)
# Combine best parameters with static parts
trained_nle = eqx.combine(result.best_trainable, nle_static)
trained_nle = eqx.tree_at(lambda m: m.trained, trained_nle, True)
return DiscreteNLETrainingResult(
discrete_nle=trained_nle,
losses=result.train_losses,
val_losses=result.val_losses,
best_epoch=result.best_epoch,
)