fiesta.train

Contents

fiesta.train#

Components for training surrogate models.

Trainers#

FluxSurrogateTrainer is the actively maintained training path: it trains a single spectral-flux surrogate covering all filters at once. LightcurveSurrogateTrainer (and its SVDTrainer subclass) predates it and is kept only so that already-trained fiesta.models.surrogate_models.LightcurveSurrogate models can still be reproduced or retrained; it is deprecated in favor of FluxSurrogateTrainer.

API class to train machine-learning surrogates on

class fiesta.train.trainers.FluxSurrogateTrainer.FluxSurrogateTrainer(surrogate_name, data, outdir, network, conversion=None, plots_dir=None, save_preprocessed_data=False)[source]#

Bases: object

Training API class for training a surrogate model that predicts a spectral flux density array.

data: DataLoader#
fit(verbose=True)[source]#

Method used to train the NN on the training data.

Parameters:

verbose (bool, optional) – Whether the train and validation loss is printed to terminal in certain intervals. Defaults to True.

Return type:

None

model_type: str#
network: NN#
outdir: str#
plot_example_lc(filters)[source]#
plot_learning_curve(train_losses, val_losses)[source]#
save()[source]#

Save the trained model and all the metadata to the outdir. The meta data is saved as a pickled dict to be read by fiesta.models.surrogate_models.Surrogate. The NN is saved as a pickled serialized dict using the NN.save_model method.

Return type:

None

surrogte_name: str#

DEPRECATED training API for per-filter lightcurve surrogates (SVD-decomposed magnitudes).

This predates fiesta’s spectral-flux training path and is kept only so that already trained fiesta.models.surrogate_models.LightcurveSurrogate models can still be reproduced/retrained. New surrogates should be trained with fiesta.train.trainers.FluxSurrogateTrainer instead, which trains a single spectral-flux surrogate covering all filters at once and is the actively maintained training path.

class fiesta.train.trainers.LightcurveSurrogateTrainer.LightcurveSurrogateTrainer(name, outdir, plots_dir=None, save_preprocessed_data=False)[source]#

Bases: object

DEPRECATED: Abstract class for training a collection of surrogate models per filter.

Use fiesta.train.trainers.FluxSurrogateTrainer instead. This class is no longer actively developed and is kept only for backwards compatibility.

filters: list[Filter]#
fit(config, key=Array([0, 0], dtype=uint32), verbose=True)[source]#

The config controls which architecture is built and therefore should not be specified here.

Parameters:

config (nn.NeuralnetConfig, optional) – _description_. Defaults to None.

Return type:

None

name: str#
outdir: str#
parameter_names: list[str]#
plot_example_lc(lc_model)[source]#
preprocess()[source]#
save()[source]#

Save the trained model and all the used metadata to the outdir.

class fiesta.train.trainers.LightcurveSurrogateTrainer.SVDTrainer(name, outdir, filters, data_manager_args, svd_ncoeff=50, conversion=None, plots_dir=None, save_preprocessed_data=False)[source]#

Bases: LightcurveSurrogateTrainer

DEPRECATED: see LightcurveSurrogateTrainer.

load_filters(filters)[source]#
preprocess()[source]#

Preprocessing method to get the SVD coefficients of the training and validation data. This includes scaling the inputs and outputs, as well as performing SVD decomposition.

Data#

DataLoader class to interact with the training data files

class fiesta.train.DataLoader.DataLoader(file, tmin, tmax, numin=1000000000.0, numax=2.5e+18, n_training=None, n_val=None, special_training=[])[source]#

Bases: object

load_from_file(group, index, special_label=None)[source]#

Loads raw data from the file and returns them as arrays.

Parameters:
  • group (str) – The data group from which the file to load from. Can be train, val, test, or special_train. If special_train, the argument special_label must also be provided.

  • index (int | slice) – Index or slice of indices to load (e.g. 5 or slice(5, 8)).

  • special_label (str) – Special data set to load from the special_train data group. Only relevant when group is "special_train". Defaults to None.

Raises:

IndexError – If index (or, for a slice, either of its bounds) falls outside the range of entries stored in group.

Return type:

tuple[Array, Array]

preprocess_cVAE(image_size, conversion=None)[source]#

Loads in the training and validation data and performs data preprocessing for the CVAE using fiesta.utils.ImageScaler. Because of memory issues, the training data set is loaded in chunks. The X arrays (parameter values) are standardized with fiesta.utils.StandardScalerJax.

Parameters:
  • image_size (Array[Int]) – Image size the 2D flux arrays are down sampled to with jax.image.resize

  • conversion (str) – references how to convert the parameters for the training. Defaults to None, in which case it’s the identity.

Returns:

Standardized training parameters. train_y (Array): PCA coefficients of the training data. val_X (Array): Standardized validation parameters val_y (Array): PCA coefficients of the validation data. Xscaler (StandardScalerJax): Standardizer object fitted to the mean and sigma of the raw training data. Can be used to transform and inverse transform parameter points. yscaler (ImageScaler): ImageScaler object fitted to part of the raw training data. Can be used to transform and inverse transform log spectral flux densities.

Return type:

train_X (Array)

preprocess_data(X_scaler, y_scaler)[source]#
Return type:

tuple[Array, Array, Array, Array, ParameterScaler, DataScaler]

preprocess_fluxes(y_scaler)[source]#

Fits and transforms the fluxes (y data sets) from self.file.

Only the requested frequency/time window (nu_start:nu_stop, time_start:time_stop from set_up_domain_mask()) is ever read from disk.

Return type:

tuple[Array, Array, DataScaler]

preprocess_parameters(X_scaler)[source]#

Fits and transforms the parameters (X data sets) from self.file.

Return type:

tuple[Array, Array, ParameterScaler]

preprocess_svd(svd_ncoeff, filters, conversion=None)[source]#

Loads in the training and validation data and performs data preprocessing for the SVD decomposition using fiesta.utils.SVDDecomposer. This is done per filter supplied in the filters argument which is equivalent to the old NMMA procedure. The X arrays (parameter values) are scaled to [0,1] with MinMaxScalerJax()

Parameters:
  • svd_ncoeff (Int) – Number of SVD coefficients to keep

  • filters (Filter[list]) – List of fiesta.utils.filter instances that are used to convert the fluxes to magnitudes

  • conversion (str) – references how to convert the parameters for the training. Defaults to None, in which case it’s the identity.

Returns:

Scaled training parameters. train_y (dict[Array]): Dictionary of the SVD coefficients of the training magnitude lightcurves with the filter names as keys val_X (Array): Scaled validation parameters val_y (dict[Array]): Dictionary of the SVD coefficients of the validation magnitude lightcurves with the filter names as keys Xscaler (ParameterScaler): MinMaxScaler object fitted to the minimum and maximum of the training data parameters. Can be used to transform and inverse transform parameter points. yscaler (dict[str, SVDDecomposer]): Dictionary of SVDDecomposer objects with the filter names as keys. The SVDDecomposer objects are fitted to the magnitude training data. Can be used to transform and inverse transform magnitudes in this filter.

Return type:

train_X (Array)

print_file_info()[source]#

Prints the meta data of the raw data, i.e., time, frequencies, and parameter names to terminal. Also prints how many training, validation, and test data points are available.

Return type:

None

print_loaded_data_info()[source]#

Prints the meta data of the loaded data, i.e., the actually requested time and frequency range to terminal. Also prints how many training, validation, and test data points will actually be used.

Return type:

None

read_metadata_from_file()[source]#

Reads in the metadata of the raw data, i.e., times, frequencies and parameter names. Also determines how many training and validation data points are available.

Return type:

None

set_up_domain_mask()[source]#

Trims the stored data down to the time and frequency range desired for training. It sets the mask attribute which is a boolean mask used when loading the data arrays.

Return type:

None

fiesta.train.DataLoader.array_mask_from_interval(sorted_array, amin, amax)[source]#

Return a mask selecting the grid points spanning [amin, amax].

If a boundary exists exactly in the array, that exact value is used. Otherwise, the interval is expanded outward to the nearest grid point.

Parameters:
  • sorted_array (array) – A sorted array

  • amin (float) – Lower interval bound

  • amax (float) – Upper interval bound

Returns:

A boolean array mask

fiesta.train.DataLoader.concatenate_redshift(X_raw, max_z=0.5)[source]#
fiesta.train.DataLoader.redshifted_magnitude(filt, mJys, nus, redshifts)[source]#

This is a slow and inefficient implementation to get the redshifted magnitudes as training data.

Classes to create afterglow training data. Not well maintained, mostly used at the very beginning for training data creation.

class fiesta.train.AfterglowData.AfterglowData(outfile, n_training, n_val, n_test, parameter_distributions=None, jet_type=-1, tmin=1.0, tmax=1000.0, n_times=100, use_log_spacing=True, numin=1000000000.0, numax=2.5e+18, n_nu=256, fixed_parameters=None)[source]#

Bases: object

create_raw_data(n, training=True)[source]#

Create draws X in the parameter space and run the afterglow model on it.

create_special_data(X_raw, label, comment=None)[source]#

Create special training data with pre-specified parameters X. These will be stored in the ‘special_train’ hdf5 group.

fix_nans(X, y)[source]#
get_raw_data(n, group)[source]#
initialize_nus(numin, numax, n_nu)[source]#
initialize_times(tmin, tmax, n_times, use_log_spacing=True)[source]#
run_afterglow_model(X)[source]#
class fiesta.train.AfterglowData.AfterglowpyData(n_pool, *args, **kwargs)[source]#

Bases: AfterglowData

run_afterglow_model(X)[source]#

Uses multiprocessing to run afterglowpy on the supplied parameters in X.

class fiesta.train.AfterglowData.BlastwaveData(*args, n_pool=None, **kwargs)[source]#

Bases: AfterglowData

create_raw_data(n, training=True)[source]#

Create draws X in the parameter space and run the blastwave model on it.

run_afterglow_model(X)[source]#

Run blastwave model on the supplied parameters in X.

No multiprocessing.Pool needed — the Rust extension uses rayon for automatic parallelism within each FluxDensity call (releases the GIL via py.allow_threads).

class fiesta.train.AfterglowData.BlastwaveRSData(*args, n_pool=None, **kwargs)[source]#

Bases: BlastwaveData

BlastwaveData variant with reverse shock enabled.

Following Japelj+ 2014 (1402.3701), the RS microphysics are tied to the FS values via a magnetization ratio RB:

eps_e_rs = eps_e_f
eps_b_rs = RB * eps_b_f
p_rs     = p_f

Extra sampled parameter: log10_RB, log10_duration. sigma is kept fixed at 0.0 (unmagnetized ejecta).

create_raw_data(n, training=True)[source]#

Sample parameters, enforce FS and RS energy constraints, then run model.

run_afterglow_model(X)[source]#

Run blastwave model on the supplied parameters in X.

No multiprocessing.Pool needed — the Rust extension uses rayon for automatic parallelism within each FluxDensity call (releases the GIL via py.allow_threads).

fiesta.train.AfterglowData.JetsimpyData#

alias of BlastwaveData

class fiesta.train.AfterglowData.PyblastafterglowData(path_to_exec, pbag_kwargs=None, rank=0, *args, **kwargs)[source]#

Bases: AfterglowData

run_afterglow_model(X)[source]#

Should be run in parallel with different mpi processes to run pyblastafterglow on the parameters in the array X.

supplement_time(t_supp)[source]#

WARNING: NOT READY TO BE USED

class fiesta.train.AfterglowData.RunAfterglowpy(jet_type, times, nus, X, parameter_names, fixed_parameters=None)[source]#

Bases: object

class fiesta.train.AfterglowData.RunBlastwave(times, nus, X, parameter_names, fixed_parameters=None, ncells=33, spread_mode='ode')[source]#

Bases: object

class fiesta.train.AfterglowData.RunBlastwaveRS(times, nus, X, parameter_names, fixed_parameters=None, ncells=33, spread_mode='ode')[source]#

Bases: object

Like RunBlastwave but with reverse shock enabled.

Following Japelj+ 2014: RS microphysics derived from FS values via magnetization ratio RB = eps_b_rs / eps_b_f, with eps_e_rs = eps_e_f and p_rs = p_f.

fiesta.train.AfterglowData.RunJetsimpy#

alias of RunBlastwave

class fiesta.train.AfterglowData.RunPyblastafterglow(jet_type, times, nus, X, parameter_names, fixed_parameters=None, rank=0, path_to_exec='./pba.out', grb_resolution=12, ntb=1000, tb0=10.0, tb1=100000000000.0, rtol=0.1, loglevel='err')[source]#

Bases: object

fiesta.train.utils.append_training_data_file(outfile, train_X, train_y, val_X, val_y, test_X, test_y)[source]#
fiesta.train.utils.convert_POSSIS_outputs_to_h5(dirs, outfile, parameter_names, log_arguments, train_size=0.8, clip=6.5144)[source]#

Merges a directory (or several directories) full of POSSIS .h5 outputs to a single training data file in the fiesta format.

Parameters:
  • dirs (str | list[str]) – directory or list of directories with the outputs that should be merged into a single training data file.

  • outfile (str) – Name of the .h5 training data file to create.

  • parameter_names (list[str]) – Parameter names in the order they appear in the file names. These will be the parameters of the trained surrogate in the end. Note that function expects the possis file to contain different inclinations.

  • log_arguments (list[int]) – Parameters, which are not log10 when read from the filenames, but should be converted to log10 for the training.

  • train_size (float) – Relative proportion of the training data. Defaults to 0.8.

  • clip (float) – Lower floor value for the minimum log10(mJy) at 10 pc. Every flux density below that will be set to that value. Defaults to 6.5144 (appr. 0 abs. mag).

Return type:

None

fiesta.train.utils.convert_SEDONA_outputs_to_h5(dirs, outfile, parameter_names, log_arguments, train_size=0.8, clip=6.5144)[source]#

Merges a directory (or several directories) full of SEDONA .h5 outputs to a single training data file in the fiesta format.

Parameters:
  • dirs (str | list[str]) – directory or list of directories with the outputs that should be merged into a single training data file.

  • outfile (str) – Name of the .h5 training data file to create.

  • parameter_names (list[str]) – Parameter names in the order they appear in the file names. These will be the parameters of the trained surrogate in the end.

  • log_arguments (list[int]) – Parameters, which are not log10 when read from the filenames, but should be converted to log10 for the training.

  • train_size (float) – Relative proportion of the training data. Defaults to 0.8.

  • clip (float) – Lower floor value for the minimum log10(mJy) at 10 pc. Every flux density below that will be set to that value. Defaults to 6.5144 (appr. 0 abs. mag).

Return type:

None

fiesta.train.utils.read_POSSIS_file(filename)[source]#
fiesta.train.utils.read_SEDONA_file(filename)[source]#
fiesta.train.utils.read_SEDONA_parameters(filename)[source]#
fiesta.train.utils.read_parameters_POSSIS(filename)[source]#
fiesta.train.utils.train_test_split(X, y, train_size)[source]#

Split arrays into training and test sets.

Parameters:
  • X (Array) – Input features, with samples along the first axis.

  • y (Array) – Target values corresponding to the samples in X.

  • train_size (float | int) – Number or fraction of samples to include in the training set. If a float, must be between 0 and 1. If an int, specifies the exact number of training samples.

Returns:

A tuple containing X_train, X_test, y_train, y_test. The training arrays contain train_size samples, while the test arrays contain the remaining samples.

Return type:

tuple[Array, Array, Array, Array]

Raises:

ValueError – If X and y have different numbers of samples, or if train_size is invalid.

fiesta.train.utils.write_training_data(outfile, train_X, train_y, val_X, val_y, test_X, test_y, times, nus, parameter_names, parameter_distributions)[source]#

Neural Networks#

fiesta.train.neuralnets wraps the raw Flax network definitions in fiesta.train.nn_architectures with a common training-loop interface (the NN base class) and a shared configuration object (NeuralnetConfig).

Base class for the neural network API

class fiesta.train.neuralnets.base.NN[source]#

Bases: object

Abstract base class for the NN architecture wrappers of flax neural networks.

save_model(outfile)[source]#

Serialize and save the model to a file.

Raises:

ValueError – If the provided file extension is not .pkl or .pickle.

Parameters:

outfile (str) – The pickle file to which we save the serialized model.

Return type:

None

train_loop(train_X, train_y, val_X=None, val_y=None, verbose=True)[source]#
Return type:

tuple[TrainState, Array, Array]

Class for the multi-layer perceptron.

class fiesta.train.neuralnets.mlp.MLP(config, key=Array((), dtype=key<fry>) overlaying: [ 0 21])[source]#

Bases: NN

Classical multi-layer perceptron using the flax-interface.

Parameters:
  • config (NeuralnetConfig) – NN config dictionary. Its output_size will determine to the number of PCA components kept after data preprocessing.

  • key (PRNGKey, optional) – Random key for initialization. Defaults to 21.

static eval_step(state, X, y, component_weights)[source]#
static load_model(filename)[source]#

Load an MLP from file.

Parameters:

filename (str) – Filename of the model to be loaded.

Raises:

ValueError – If there is something wrong with loading, since lots of things can go wrong here.

Returns:

The TrainState object loaded from the file and the NeuralnetConfig object.

Return type:

tuple[TrainState, NeuralnetConfig]

preprocess_data(data, conversion)[source]#

Preprocesses the training and validation data. Returns rescaled training and validation data arrays as well as the scaler objects. For the MLP, we perform PCA decomposition and keep the number of PCA components specified through the output_size of the NN.

Parameters:
  • data (DataLoader) – Data file with training and validation data.

  • conversion (callable) – Special conversion function to generate parameter combinations that assist the training.

Raises:

ValueError – If nan values are introduced when rescaling the flux densities.

Return type:

None

train_loop(train_X, train_y, val_X=None, val_y=None, verbose=True)[source]#
static train_step(state, batch_X, batch_y, dropout_rng, component_weights)[source]#

Class for the conditional variational autoencoder.

class fiesta.train.neuralnets.cvae.CVAE(config, image_size, key=Array((), dtype=key<fry>) overlaying: [ 0 21])[source]#

Bases: NN

Conditional variational autoencoder using the flax-interface.

Parameters:
  • config (NeuralnetConfig) – NN config dictionary. Its latent_dim will determine the size of the latent layer.

  • image_size (tuple[int]) – Tuple of length two that will determine to which size the 2D arrays for the flux densities are down scaled to when preprocessing the data. This also then becomes the input and output dimension of the CVAE.

  • key (PRNGKey, optional) – Random key for initialization. Defaults to 21.

static load_full_model(filename)[source]#
Return type:

tuple[TrainState, NeuralnetConfig]

static load_model(filename)[source]#

Load a model from a file.

Parameters:

filename (str) – Filename of the model to be loaded.

Raises:

ValueError – If there is something wrong with loading, since lots of things can go wrong here.

Returns:

The TrainState object loaded from the file and the NeuralnetConfig object.

Return type:

tuple[TrainState, NeuralnetConfig]

preprocess_data(data, conversion)[source]#

Preprocesses the training and validation data. Returns rescaled training and validation data arrays as well as the scaler objects. For the CVAE, we scale the 2D flux arrays down to the shape determined through image_size and standardize them.

Parameters:
  • data (DataLoader) – Data file with training and validation data.

  • conversion (callable) – Special conversion function to generate parameter combinations that assist the training.

Raises:

ValueError – If nan values are introduced when rescaling the flux densities.

Return type:

None

train_loop(train_X, train_y, val_X=None, val_y=None, verbose=True)[source]#
static train_step(state, train_X, train_y, rng, val_X=None, val_y=None)[source]#
Return type:

tuple[TrainState, Array, 'n_batch_train'], Array, 'n_batch_val']]

Utils for dealing with the neural networks

class fiesta.train.neuralnets.utils.NeuralnetConfig(name='MLP', output_size=10, input_size=10, hidden_layer_sizes=[64, 128, 64], learning_rate=0.001, conditional_dim=None, latent_dim=20, weight_decay=0.0, batch_size=128, nb_epochs=1000, nb_report=None, dropout_rate=0.0, use_cosine_schedule=False, cosine_alpha=0.01, max_grad_norm=0.0, pca_smoothness_weight=0.0, pca_smoothness_start=0)[source]#

Bases: ConfigDict

Configuration for a neural network model. For type hinting

hidden_layer_sizes: list[int]#
input_size: Int#
learning_rate: Float#
name: str#
output_size: Int#
fiesta.train.neuralnets.utils.bce(y, pred)[source]#

binary cross entropy between y and the predicted array pred

fiesta.train.neuralnets.utils.kld(mean, logvar)[source]#

Kullback-Leibler divergence of a normal distribution with arbitrary mean and log variance to the standard normal distribution with mean 0 and unit variance.

fiesta.train.neuralnets.utils.mse(y, pred)[source]#

square error between y and the predicted array pred

fiesta.train.neuralnets.utils.serialize(state, config=None)[source]#

Serialize function to save the model and its configuration.

Parameters:
  • state (TrainState) – The TrainState object to be serialized.

  • config (NeuralnetConfig, optional) – The config to be serialized. Defaults to None.

Returns:

_description_

Return type:

_type_

class fiesta.train.nn_architectures.BaseNeuralnet(layer_sizes, act_func=<jax._src.custom_derivatives.custom_jvp object>, parent=<flax.linen.module._Sentinel object>, name=None)[source]#

Bases: Module

Abstract base class. Needs layer sizes and activation function used

layer_sizes: Sequence[int]#
name: str | None = None#
parent: Module | Scope | _Sentinel | None = None#
scope: Scope | None = None#
setup()[source]#

Initializes a Module lazily (similar to a lazy __init__).

setup is called once lazily on a module instance when a module is bound, immediately before any other methods like __call__ are invoked, or before a setup-defined attribute on self is accessed.

This can happen in three cases:

  1. Immediately when invoking apply(), init() or init_and_output().

  2. Once the module is given a name by being assigned to an attribute of another module inside the other module’s setup method (see __setattr__()):

    >>> class MyModule(nn.Module):
    ...   def setup(self):
    ...     submodule = nn.Conv(...)
    
    ...     # Accessing `submodule` attributes does not yet work here.
    
    ...     # The following line invokes `self.__setattr__`, which gives
    ...     # `submodule` the name "conv1".
    ...     self.conv1 = submodule
    
    ...     # Accessing `submodule` attributes or methods is now safe and
    ...     # either causes setup() to be called once.
    
  3. Once a module is constructed inside a method wrapped with compact(), immediately before another method is called or setup defined attribute is accessed.

class fiesta.train.nn_architectures.CNN(dense_layer_sizes, kernel_sizes, conv_layer_sizes, output_shape, spatial=32, act_func=<jax._src.custom_derivatives.custom_jvp object>, parent=<flax.linen.module._Sentinel object>, name=None)[source]#

Bases: Module

Convolutional Neural Network

conv_layer_sizes: Sequence[Int]#
dense_layer_sizes: Sequence[Int]#
kernel_sizes: Sequence[Int]#
name: str | None = None#
output_shape: tuple[Int, Int]#
parent: Module | Scope | _Sentinel | None = None#
scope: Scope | None = None#
setup()[source]#

Initializes a Module lazily (similar to a lazy __init__).

setup is called once lazily on a module instance when a module is bound, immediately before any other methods like __call__ are invoked, or before a setup-defined attribute on self is accessed.

This can happen in three cases:

  1. Immediately when invoking apply(), init() or init_and_output().

  2. Once the module is given a name by being assigned to an attribute of another module inside the other module’s setup method (see __setattr__()):

    >>> class MyModule(nn.Module):
    ...   def setup(self):
    ...     submodule = nn.Conv(...)
    
    ...     # Accessing `submodule` attributes does not yet work here.
    
    ...     # The following line invokes `self.__setattr__`, which gives
    ...     # `submodule` the name "conv1".
    ...     self.conv1 = submodule
    
    ...     # Accessing `submodule` attributes or methods is now safe and
    ...     # either causes setup() to be called once.
    
  3. Once a module is constructed inside a method wrapped with compact(), immediately before another method is called or setup defined attribute is accessed.

spatial: Int = 32#
class fiesta.train.nn_architectures.CVAE(hidden_layer_sizes, output_size, latent_dim=20, parent=<flax.linen.module._Sentinel object>, name=None)[source]#

Bases: Module

Conditional Variational Autoencoder consisting of an Encoder and a Decoder.

hidden_layer_sizes: Sequence[Int]#
latent_dim: Int = 20#
name: str | None = None#
output_size: Int#
parent: Module | Scope | _Sentinel | None = None#
scope: Scope | None = None#
setup()[source]#

Initializes a Module lazily (similar to a lazy __init__).

setup is called once lazily on a module instance when a module is bound, immediately before any other methods like __call__ are invoked, or before a setup-defined attribute on self is accessed.

This can happen in three cases:

  1. Immediately when invoking apply(), init() or init_and_output().

  2. Once the module is given a name by being assigned to an attribute of another module inside the other module’s setup method (see __setattr__()):

    >>> class MyModule(nn.Module):
    ...   def setup(self):
    ...     submodule = nn.Conv(...)
    
    ...     # Accessing `submodule` attributes does not yet work here.
    
    ...     # The following line invokes `self.__setattr__`, which gives
    ...     # `submodule` the name "conv1".
    ...     self.conv1 = submodule
    
    ...     # Accessing `submodule` attributes or methods is now safe and
    ...     # either causes setup() to be called once.
    
  3. Once a module is constructed inside a method wrapped with compact(), immediately before another method is called or setup defined attribute is accessed.

class fiesta.train.nn_architectures.Decoder(layer_sizes, act_func=<jax._src.custom_derivatives.custom_jvp object>, dropout_rate=0.0, parent=<flax.linen.module._Sentinel object>, name=None)[source]#

Bases: MLP

name: str | None = None#
parent: Module | Scope | _Sentinel | None = None#
scope: Scope | None = None#
class fiesta.train.nn_architectures.Encoder(layer_sizes, act_func=<jax._src.custom_derivatives.custom_jvp object>, parent=<flax.linen.module._Sentinel object>, name=None)[source]#

Bases: Module

layer_sizes: Sequence[int]#
name: str | None = None#
parent: Module | Scope | _Sentinel | None = None#
scope: Scope | None = None#
setup()[source]#

Initializes a Module lazily (similar to a lazy __init__).

setup is called once lazily on a module instance when a module is bound, immediately before any other methods like __call__ are invoked, or before a setup-defined attribute on self is accessed.

This can happen in three cases:

  1. Immediately when invoking apply(), init() or init_and_output().

  2. Once the module is given a name by being assigned to an attribute of another module inside the other module’s setup method (see __setattr__()):

    >>> class MyModule(nn.Module):
    ...   def setup(self):
    ...     submodule = nn.Conv(...)
    
    ...     # Accessing `submodule` attributes does not yet work here.
    
    ...     # The following line invokes `self.__setattr__`, which gives
    ...     # `submodule` the name "conv1".
    ...     self.conv1 = submodule
    
    ...     # Accessing `submodule` attributes or methods is now safe and
    ...     # either causes setup() to be called once.
    
  3. Once a module is constructed inside a method wrapped with compact(), immediately before another method is called or setup defined attribute is accessed.

class fiesta.train.nn_architectures.MLP(layer_sizes, act_func=<jax._src.custom_derivatives.custom_jvp object>, dropout_rate=0.0, parent=<flax.linen.module._Sentinel object>, name=None)[source]#

Bases: BaseNeuralnet

Basic multi-layer perceptron: a feedforward neural network with multiple Dense layers.

dropout_rate: float = 0.0#
name: str | None = None#
parent: Module | Scope | _Sentinel | None = None#
scope: Scope | None = None#
setup()[source]#

Initializes a Module lazily (similar to a lazy __init__).

setup is called once lazily on a module instance when a module is bound, immediately before any other methods like __call__ are invoked, or before a setup-defined attribute on self is accessed.

This can happen in three cases:

  1. Immediately when invoking apply(), init() or init_and_output().

  2. Once the module is given a name by being assigned to an attribute of another module inside the other module’s setup method (see __setattr__()):

    >>> class MyModule(nn.Module):
    ...   def setup(self):
    ...     submodule = nn.Conv(...)
    
    ...     # Accessing `submodule` attributes does not yet work here.
    
    ...     # The following line invokes `self.__setattr__`, which gives
    ...     # `submodule` the name "conv1".
    ...     self.conv1 = submodule
    
    ...     # Accessing `submodule` attributes or methods is now safe and
    ...     # either causes setup() to be called once.
    
  3. Once a module is constructed inside a method wrapped with compact(), immediately before another method is called or setup defined attribute is accessed.

Benchmarking#

class fiesta.train.Benchmarker.Benchmarker(model, data, filters=None, outdir='./benchmarks', output_format='pdf')[source]#

Bases: object

benchmark()[source]#
calculate_error()[source]#
get_data()[source]#
lightcurves_mismatch(metric_key='highest_lc_error')[source]#
plot_error_distribution(metric_key='highest_lc_error')[source]#
plot_error_over_time()[source]#
plot_lightcurves_mismatch()[source]#
plot_worst_lightcurves()[source]#
print_correlations(metric_key='highest_lc_error')[source]#
worst_lightcurves(metric_key='highest_lc_error')[source]#