setu.SimulationDataset#

class setu.SimulationDataset(theta, x, conditions=None)[source]#

Bases: object

Dataset of (theta, x) simulation pairs with optional conditions.

Container for simulator outputs used to train NLE/MixedNLE. Each row corresponds to one simulation: theta[i] produced x[i] (optionally with conditions[i]).

theta#

Parameter samples with shape (n_samples, theta_dim).

x#

Corresponding observations with shape (n_samples, x_dim).

conditions#

Experimental conditions with shape (n_samples, condition_dim). None if no conditions.

Raises:
  • ValueError – If theta and x have different sample counts.

  • ValueError – If conditions has wrong sample count or is not 2D.

Parameters:
  • theta (Array)

  • x (Array)

  • conditions (Array | None)

Example

>>> dataset = SimulationDataset(theta=theta, x=x)
>>> dataset_with_cond = SimulationDataset(theta, x, conditions=cond)
>>> train, val = dataset.split(val_fraction=0.1, key=key)
Parameters:
  • theta (Array)

  • x (Array)

  • conditions (Array | None)

property theta_dim: int#

Dimension of parameter space.

property x_dim: int#

Dimension of observation space.

property condition_dim: int#

Dimension of condition space, or 0 if no conditions.

split(val_fraction=0.1, key=None)[source]#

Split into training and validation sets.

Parameters:
  • val_fraction (float) – Fraction for validation (default 0.1)

  • key (Array | None) – JAX random key for shuffling. If None, no shuffle.

Return type:

tuple[SimulationDataset, SimulationDataset]

Returns:

(train_dataset, val_dataset)

save(path)[source]#

Save dataset to .npz file.

Parameters:

path (str | Path) – File path (e.g., “dataset.npz”).

Return type:

None

classmethod load(path)[source]#

Load dataset from .npz file.

Parameters:

path (str | Path) – File path to load from.

Return type:

SimulationDataset

Returns:

SimulationDataset with loaded arrays.