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:
- 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.