fiesta.models

Contents

fiesta.models#

Model classes implemented in fiesta: surrogate (neural-network) models, analytical (physics-based) models, and combinations thereof. All of them share the common FiestaModel interface (name, parameter_names, times, filters, predict()).

Base Interface#

class fiesta.models.base.FiestaModel(name, filters, times=None)[source]#

Bases: ABC

Filters: list[Filter]#
add_filters(filters)[source]#
add_name(x)[source]#

Turns an unnamed array into a dictionary.

filters: list[str]#
name: str#
parameter_names: list[str]#
abstractmethod predict(x)[source]#
Return type:

tuple[Array, dict[str, Array]]

predict_abs_mag(x)[source]#
Return type:

tuple[Array, dict[str, Array]]

times: Array | None#
vpredict(X)[source]#

Vectorized prediction function to calculate the apparent magnitudes for several inputs x at the same time.

Return type:

tuple[Array, dict[str, Array]]

Surrogate Models#

Store classes to load in trained machine-learning surrogates and give routines to let them generate lightcurves.

class fiesta.models.surrogate_models.FluxSurrogate(name, filters, directory=None)[source]#

Bases: Surrogate

Class of surrogate models that predicts the 2D spectral flux density array.

Parameters:
  • name (str) – Name of the model

  • filters (list[str]) – List of all the filters for which the model should be loaded.

  • directory (str) – Directory with trained model states and projection metadata such as scalers. Defaults to None, in which case there will be an attempt to load from the repo based on name.

compute_output(x)[source]#

Apply the trained flax neural network on the given input x.

Parameters:

x (dict[str, Array]) – Input array of parameters per filter

Returns:

_description_

Return type:

dict[str, Array]

convert_to_mag(y, x)[source]#
Return type:

tuple[Array, dict[str, Array]]

load_filters(filters=None)[source]#
Return type:

None

load_networks()[source]#
Return type:

None

nus: Array#
predict_log_flux(x)[source]#

Predict the total log10 flux array for the parameters x.

Parameters:

x (dict[str, Array]) – Input parameters, unnormalized and untransformed.

Returns:

times [Array]: time array in observer frame nus [Array]: frequency array in observer frame log10_flux [Array]: Array of log10-fluxes in mJy.

Return type:

tuple

project_input(x)[source]#

Project the given input to whatever preprocessed input space we are in.

Parameters:

x (Array) – Original input array

Returns:

Transformed input array

Return type:

Array

project_output(y)[source]#

Project the computed output to whatever preprocessed output space we are in.

Parameters:

y (dict[str, Array]) – Output array

Returns:

Output array transformed to the preprocessed space.

Return type:

dict[str, Array]

class fiesta.models.surrogate_models.LightcurveSurrogate(name, filters, directory=None)[source]#

Bases: Surrogate

Class of surrogate models that predicts the magnitudes per filter.

X_scaler: object#
compute_output(x)[source]#

Apply the trained flax neural network on the given input x.

Parameters:

x (dict[str, Array]) – Input array of parameters per filter

Returns:

_description_

Return type:

dict[str, Array]

convert_to_mag(y, x)[source]#
Return type:

tuple[Array, dict[str, Array]]

directory: str#
load_filters(filters_args=None)[source]#
Return type:

None

load_networks()[source]#
Return type:

None

metadata: dict#
models: dict[str, TrainState]#
project_input(x)[source]#

Project the given input to whatever preprocessed input space we are in.

Parameters:

x (dict[str, Array]) – Original input array

Returns:

Transformed input array

Return type:

dict[str, Array]

project_output(y)[source]#

Project the computed output to whatever preprocessed output space we are in.

Parameters:

y (dict[str, Array]) – Output array

Returns:

Output array transformed to the preprocessed space.

Return type:

dict[str, Array]

y_scaler: dict[str, object]#
class fiesta.models.surrogate_models.Surrogate(name, filters, directory=None)[source]#

Bases: FiestaModel

Abstract class for general surrogate models

compute_output(x)[source]#
Return type:

dict[str, Array]

directory: str#
load_metadata()[source]#
Return type:

None

predict(x)[source]#

Generate the apparent magnitudes from the unnormalized and untransformed input x. Chains the projections with the actual computation of the output. E.g. if the model is a trained surrogate neural network, they represent the map from x tilde to y tilde. The mappings from x to x tilde and y to y tilde take care of projections (e.g. SVD projections) and normalizations.

Parameters:

x (dict[str, Array]) – Input array, unnormalized and untransformed.

Returns:

times (Array): time array in observer frame mag (dict[str, Array]): The predicted magnitudes per filter

Return type:

tuple

project_input(x)[source]#
Return type:

dict[str, Array]

project_output(y)[source]#
Return type:

dict[str, Array]

fiesta.models.surrogate_models.get_default_directory(name)[source]#

Combined Models#

Class to combine several separate models.

class fiesta.models.combined_model.CombinedModel(models, sample_times)[source]#

Bases: FiestaModel

add_filters(filters)[source]#
predict(x)[source]#

Predict the joint light curve by combining several separate submodels.

Parameters:

x (dict[str, Array]) – Input array, unnormalized and untransformed. All model parameters from all models need to be specified here.

Returns:

times (Array): time array in observer frame mag (dict[str, Array]): The predicted magnitudes per filter

Return type:

tuple

Analytical Models#

Base classes, constants, and shared helpers for analytical light-curve models.

Each model is fully JIT-compilable and differentiable so that flowMC’s MALA sampler can compute jax.grad through the likelihood. The models follow the same predict() contract as the surrogate models:

(source_frame_times, {filter_name: apparent_mag_array})

This makes them drop-in replacements inside CombinedSurrogate and EMLikelihood.

All internal physics computations use log10 space to avoid float32 overflow (e.g. explosion energies ~1e49 erg exceed float32 max ~3.4e38).

class fiesta.models.analytical_models.base.AnalyticalModel(name, filters, times=None, temperature_floor=None)[source]#

Bases: FiestaModel

Base class for analytical (non-surrogate) light-curve models.

Subclasses must implement compute_log10_lbol_rphot(self, x, t_days) which returns (log10_Lbol, log10_Rphot) — log10 of bolometric luminosity in erg/s and photospheric radius in cm.

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

Return type:

tuple[Array, Array]

predict(x)[source]#
Return type:

tuple[Array, dict[str, Array]]

Phenomenological (flux-shape) light-curve models.

Reference:

Redback: nikhil-sarin/redback Boom: boom-astro/boom

class fiesta.models.analytical_models.phenomenological_models.AfterglowModel(filters, times=None, name=None)[source]#

Bases: PhenomenologicalModel

Smooth broken power-law afterglow model.

Reference:

Boom: boom-astro/boom

Transitions from r^(-alpha_1) at early times to r^(-alpha_2) at late times.

Shape parameters: t0, log10_t_break, alpha_1, alpha_2

compute_shape(x, t_days)[source]#

Return the temporal shape function S(t) >= 0.

has_baseline: bool = False#
shape_parameter_names: list[str] = ['t0', 'log10_t_break', 'alpha_1', 'alpha_2']#
class fiesta.models.analytical_models.phenomenological_models.BazinModel(filters, times=None, name=None)[source]#

Bases: PhenomenologicalModel

Bazin et al. phenomenological light-curve model.

Reference:

Boom: boom-astro/boom

Shape: exp(-dt/tau_fall) * sigmoid(dt/tau_rise)

Parameters (per-band): amp_mag_{filter}, base_mag_{filter} Shape parameters: t0, log10_tau_rise, log10_tau_fall

compute_shape(x, t_days)[source]#

Return the temporal shape function S(t) >= 0.

has_baseline: bool = True#
shape_parameter_names: list[str] = ['t0', 'log10_tau_rise', 'log10_tau_fall']#
class fiesta.models.analytical_models.phenomenological_models.EvolvingBlackbodyModel(filters, times=None, reference_time=1.0, name=None)[source]#

Bases: AnalyticalModel

Phenomenological model with piecewise power-law T and R evolution.

Reference:

Redback: nikhil-sarin/redback

Model-agnostic — useful for fast empirical fitting of any thermal transient. Based on the evolving_blackbody model from Redback.

Parameters (in x dict):

log10_temperature_0 – log10 initial temperature (K) at reference_time log10_radius_0 – log10 initial radius (cm) at reference_time temp_rise_index – T rise power-law index for t <= temp_peak_time temp_decline_index – T decline power-law index for t > temp_peak_time temp_peak_time – time (days) when temperature peaks radius_rise_index – R rise power-law index for t <= radius_peak_time radius_decline_index – R decline power-law index for t > radius_peak_time radius_peak_time – time (days) when radius peaks

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_temperature_0', 'log10_radius_0', 'temp_rise_index', 'temp_decline_index', 'temp_peak_time', 'radius_rise_index', 'radius_decline_index', 'radius_peak_time']#
class fiesta.models.analytical_models.phenomenological_models.PhenomenologicalModel(filters, times=None, name=None)[source]#

Bases: AnalyticalModel

Base class for phenomenological light-curve models.

Unlike physics-based models that compute L_bol + R_phot and pass through a blackbody SED, phenomenological models compute a temporal shape function S(t) and convert directly to per-band apparent magnitudes.

Subclasses must set:

shape_parameter_names : list[str] — shared temporal shape parameters has_baseline : bool — whether the model has a baseline flux component

and implement compute_shape(self, x, t_days) -> Array.

compute_shape(x, t_days)[source]#

Return the temporal shape function S(t) >= 0.

Return type:

Array

has_baseline: bool = False#
predict(x)[source]#
Return type:

tuple[Array, dict[str, Array]]

shape_parameter_names: list[str]#
class fiesta.models.analytical_models.phenomenological_models.PhenomenologicalTDEModel(filters, times=None, name=None)[source]#

Bases: PhenomenologicalModel

Phenomenological TDE light-curve model.

Reference:

Boom: boom-astro/boom

Sigmoid rise with power-law decay.

Shape parameters: t0, log10_tau_rise, log10_tau_fall, alpha_decay

compute_shape(x, t_days)[source]#

Return the temporal shape function S(t) >= 0.

has_baseline: bool = True#
shape_parameter_names: list[str] = ['t0', 'log10_tau_rise', 'log10_tau_fall', 'alpha_decay']#
class fiesta.models.analytical_models.phenomenological_models.VillarModel(filters, times=None, name=None)[source]#

Bases: PhenomenologicalModel

Villar et al. phenomenological light-curve model.

Reference:

Boom: boom-astro/boom

Piecewise shape with smooth sigmoid transition at gamma.

Shape parameters: t0, log10_tau_rise, log10_tau_fall, beta_slope, log10_gamma

compute_shape(x, t_days)[source]#

Return the temporal shape function S(t) >= 0.

constraint_penalty(x)[source]#

Physical validity penalty (de Soto et al. 2024).

Returns 0 for valid parameters, positive for invalid. Multiply by a large negative factor and add to log-likelihood to enforce.

has_baseline: bool = False#
shape_parameter_names: list[str] = ['t0', 'log10_tau_rise', 'log10_tau_fall', 'beta_slope', 'log10_gamma']#

Supernova analytical light-curve models.

Reference:

Redback: nikhil-sarin/redback NMMA: nuclear-multimessenger-astronomy/nmma

class fiesta.models.analytical_models.supernova_models.ArnettModel(filters, times=None, modified=False, name=None)[source]#

Bases: AnalyticalModel

Arnett (1982) Ni56/Co56-powered supernova bolometric model.

Reference:

Redback: nikhil-sarin/redback NMMA: nuclear-multimessenger-astronomy/nmma

Parameters (in x dict):

tau_m – diffusion timescale in days log10_mni – log10 of Ni56 mass in solar masses v_phot – photospheric velocity in units of 1e9 cm/s t_0 – (modified variant only) gamma-ray trapping timescale in days

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['tau_m', 'log10_mni', 'v_phot']#
class fiesta.models.analytical_models.supernova_models.CSMInteractionModel(filters, times=None, nn=12, delta=1, efficiency=0.5, temperature_floor=None, name=None)[source]#

Bases: AnalyticalModel

Circumstellar medium interaction model (Chevalier 1982).

Reference:

Redback: nikhil-sarin/redback

Forward + reverse shock luminosity from Chevalier self-similar solution, with optional CSM diffusion.

Parameters (in x dict):

log10_mej – log10 of ejecta mass in solar masses log10_csm_mass – log10 of CSM mass in solar masses log10_vej – log10 of ejecta velocity in km/s eta – CSM density profile exponent log10_rho – log10 of CSM density amplitude (g/cm^{eta+3}) log10_kappa – log10 of opacity (cm^2/g) log10_r0 – log10 of CSM inner radius in AU

Constructor kwargs:

nn – ejecta power-law index (default 12) delta – inner density exponent (default 1) efficiency – kinetic-to-luminosity conversion (default 0.5)

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_mej', 'log10_csm_mass', 'log10_vej', 'eta', 'log10_rho', 'log10_kappa', 'log10_r0']#
class fiesta.models.analytical_models.supernova_models.MagnetarPoweredSNModel(filters, times=None, temperature_floor=None, name=None)[source]#

Bases: AnalyticalModel

Magnetar spin-down powered supernova with Arnett (1982) diffusion.

Reference:

Redback: nikhil-sarin/redback

Parameters (in x dict):

log10_p0 – log10 initial spin period in ms log10_bp – log10 polar B-field in 1e14 G mass_ns – neutron star mass in solar masses theta_pb – angle between spin and B-field in radians log10_mej – log10 of ejecta mass in solar masses log10_vej – log10 of ejecta velocity in km/s log10_kappa – log10 of opacity (cm^2/g) log10_kappa_gamma – log10 of gamma-ray opacity (cm^2/g)

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_p0', 'log10_bp', 'mass_ns', 'theta_pb', 'log10_mej', 'log10_vej', 'log10_kappa', 'log10_kappa_gamma']#
class fiesta.models.analytical_models.supernova_models.NickelCobaltModel(filters, times=None, temperature_floor=None, name=None)[source]#

Bases: AnalyticalModel

Ni56/Co56 radioactive decay with Arnett (1982) diffusion.

Reference:

Redback: nikhil-sarin/redback

Parameters (in x dict):

f_nickel – fraction of ejecta mass in Ni56 log10_mej – log10 of ejecta mass in solar masses log10_vej – log10 of ejecta velocity in km/s log10_kappa – log10 of opacity (cm^2/g) log10_kappa_gamma – log10 of gamma-ray opacity (cm^2/g)

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['f_nickel', 'log10_mej', 'log10_vej', 'log10_kappa', 'log10_kappa_gamma']#

Kilonova analytical light-curve models.

Reference:

Redback: nikhil-sarin/redback NMMA: nuclear-multimessenger-astronomy/nmma

class fiesta.models.analytical_models.kilonova_models.MagnetarBoostedKilonovaModel(filters, times=None, neutron_precursor=True, pair_cascade=True, vmax=0.7, magnetar_heating='first_layer', name=None)[source]#

Bases: AnalyticalModel

Multi-shell kilonova with magnetar spin-down heating, matching Redback.

Reference:

Redback: _general_metzger_magnetar_driven_kilonova_model

200-shell ODE with magnetar injection into bottom layer, velocity evolution, optional pair cascade and neutron precursor.

Parameters (in x dict):

log10_mej – log10 ejecta mass in solar masses log10_vej – log10 ejecta velocity (vmin) in units of c beta – velocity power-law index log10_kappa_r – log10 opacity in cm^2/g log10_p0 – log10 initial spin period in ms log10_bp – log10 polar B-field in 1e14 G mass_ns – neutron star mass in solar masses theta_pb – angle between spin and B-field in radians thermalisation_efficiency – magnetar thermalisation efficiency

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_mej', 'log10_vej', 'beta', 'log10_kappa_r', 'log10_p0', 'log10_bp', 'mass_ns', 'theta_pb', 'thermalisation_efficiency']#
class fiesta.models.analytical_models.kilonova_models.MetzgerFullModel(filters, times=None, neutron_precursor=True, vmax=0.7, name=None)[source]#

Bases: AnalyticalModel

Multi-shell kilonova model (Metzger 2017), matching Redback exactly.

Reference:

Redback: _metzger_kilonova_model in kilonova_models.py

200 shells with linear velocity spacing, Barnes+16 thermalisation, optional neutron precursor, per-gram energy ODE.

Parameters (in x dict):

log10_mej – log10 ejecta mass in solar masses log10_vej – log10 ejecta velocity (vmin) in units of c beta – velocity power-law index log10_kappa_r – log10 opacity in cm^2/g

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_mej', 'log10_vej', 'beta', 'log10_kappa_r']#
class fiesta.models.analytical_models.kilonova_models.MetzgerModel(filters, times=None, name=None)[source]#

Bases: AnalyticalModel

300-shell kilonova model matching NMMA eff_metzger_lc.

Reference:

Redback: nikhil-sarin/redback NMMA: nuclear-multimessenger-astronomy/nmma

Parameters (in x dict):

log10_mej – log10 ejecta mass in solar masses log10_vej – log10 ejecta velocity in units of c beta – velocity power-law index log10_kappa_r – log10 opacity in cm^2/g

The ODE is solved per-shell in normalized units to avoid float32 overflow. Uses 300 mass shells with velocity profile, neutron fractions, and shell-dependent opacities matching the NMMA implementation.

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_mej', 'log10_vej', 'beta', 'log10_kappa_r']#
class fiesta.models.analytical_models.kilonova_models.OneComponentKilonovaModel(filters, times=None, temperature_floor=4000.0, name=None)[source]#

Bases: AnalyticalModel

Single-component kilonova with diffusion-integral heating.

Reference:

Redback: _one_component_kilonova_model in kilonova_models.py

Matches Redback’s cumulative trapezoid algorithm exactly, using a float32-safe damped recurrence that avoids exp(t^2/td^2) overflow.

Parameters (in x dict):

log10_mej – log10 ejecta mass in solar masses log10_vej – log10 ejecta velocity in units of c log10_kappa – log10 gray opacity in cm^2/g

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_mej', 'log10_vej', 'log10_kappa']#

Shock-powered analytical light-curve models.

Reference:

Redback: nikhil-sarin/redback NMMA: nuclear-multimessenger-astronomy/nmma

class fiesta.models.analytical_models.shock_powered_models.ShockCoolingModel(filters, times=None, name=None)[source]#

Bases: AnalyticalModel

Shock-cooling emission following Piro (2021).

Reference:

Redback: nikhil-sarin/redback NMMA: nuclear-multimessenger-astronomy/nmma

Parameters (all in x dict):

log10_Menv – log10 envelope mass in solar masses log10_Renv – log10 envelope radius in solar radii log10_Ee – log10 explosion energy in erg

compute_log10_lbol_rphot(x, t_days)[source]#

Full Piro (2021) shock cooling with n=10, delta=1.1.

Matches NMMA sc_bol_lc exactly. All overflow-prone quantities are computed in log10 space to stay within float32 range.

parameter_names: list[str] = ['log10_Menv', 'log10_Renv', 'log10_Ee']#
class fiesta.models.analytical_models.shock_powered_models.ShockedCocoonModel(filters, times=None, name=None)[source]#

Bases: AnalyticalModel

Analytical jet cocoon cooling model.

Reference:

Redback: nikhil-sarin/redback

Fully algebraic (no ODE) — power-law luminosity decay with diffusion timescale. Based on the shocked cocoon model from Redback.

Parameters (in x dict):

log10_mej – log10 ejecta mass in solar masses log10_vej – log10 ejecta velocity in units of c eta – slope for ejecta density profile log10_tshock – log10 shock time in seconds shocked_fraction – fraction of ejecta mass shocked cos_theta_cocoon – cosine of cocoon opening angle log10_kappa – log10 gray opacity in cm^2/g

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_mej', 'log10_vej', 'eta', 'log10_tshock', 'shocked_fraction', 'cos_theta_cocoon', 'log10_kappa']#

Tidal disruption event (TDE) analytical light-curve models.

Reference:

Redback: nikhil-sarin/redback

class fiesta.models.analytical_models.tde_models.TDEAnalyticalModel(filters, times=None, temperature_floor=None, name=None)[source]#

Bases: AnalyticalModel

TDE analytical model with t^{-5/3} fallback + Arnett diffusion.

Reference:

Redback: nikhil-sarin/redback

Parameters (in x dict):

log10_l0 – log10 of luminosity at 1 second (erg/s) t_0_turn – turn-on time in days log10_mej – log10 of ejecta mass in solar masses log10_vej – log10 of ejecta velocity in km/s log10_kappa – log10 of opacity (cm^2/g) log10_kappa_gamma – log10 of gamma-ray opacity (cm^2/g)

compute_log10_lbol_rphot(x, t_days)[source]#

Return (log10_L_bol, log10_R_phot) arrays at each time in t_days.

L_bol in erg/s, R_phot in cm.

parameter_names: list[str] = ['log10_l0', 't_0_turn', 'log10_mej', 'log10_vej', 'log10_kappa', 'log10_kappa_gamma']#

SALT3 spectral-template supernova model via jax-bandflux.

Uses jax_supernovae (PyPI: jax-bandflux) for JAX-native, JIT-compiled, differentiable SALT3 light-curve evaluation. Unlike the physics-based models that compute L_bol + R_phot -> blackbody SED, SALT3 uses spectral templates (M0, M1, colour law) to compute per-band fluxes directly.

The jax_supernovae import is kept lazy to avoid loading heavy dependencies for users who don’t use SALT3.

class fiesta.models.analytical_models.salt3_models.SALT3Model(filters, times=None, redshift=0.0)[source]#

Bases: FiestaModel

SALT3 spectral-template model for Type Ia supernova light curves.

Parameters:
  • filters (list[str]) – Band names recognised by jax_supernovae.bandpasses (e.g. "ztfg", "ztfr", "bessellb").

  • times (Array, optional) – Observer-frame times (days) at which to evaluate the model.

  • redshift (float) – Source redshift (fixed, not sampled).

  • predict(x)) (Sampled parameters (passed via) – log10_x0 – log10 of the SALT3 amplitude parameter x0 x1 – SALT3 stretch c – SALT3 colour t0 – time of B-band maximum (days, same frame as times)

parameter_names: list[str]#
predict(x)[source]#
Return type:

tuple[Array, dict[str, Array]]