setu.SimulationDataset#
- class setu.SimulationDataset(theta, x, conditions=None)[source]#
Bases:
objectDataset 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)
- split(val_fraction=0.1, key=None)[source]#
Split into training and validation sets.
- Parameters:
- Return type:
- Returns:
(train_dataset, val_dataset)