Source code for setu.discrete_nle

"""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, )