setu.create_discrete_nle#

setu.create_discrete_nle(x_dim, theta_dim, *, distribution='categorical', num_categories=None, condition_dim=0, hidden_features=64, num_layers=2, dropout_rate=0.0, key)[source]#

Create a DiscreteNLE for categorical or count data.

Parameters:
  • x_dim (int) – Number of discrete variables.

  • theta_dim (int) – Dimension of parameters.

  • distribution (Literal['categorical', 'zinb', 'negbinom', 'poisson']) – Distribution family. “categorical” for categorical data (requires num_categories). “zinb”, “negbinom”, or “poisson” for count data.

  • num_categories (Array | None) – Categories per variable (x_dim,). Required for distribution=”categorical”.

  • condition_dim (int) – Dimension of experimental conditions. 0 if none.

  • hidden_features (int) – Hidden units per layer in MADE network.

  • num_layers (int) – Number of hidden layers.

  • dropout_rate (float) – Dropout rate for regularization.

  • key (Array) – JAX random key.

Return type:

DiscreteNLE

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.