fiesta.inference

Contents

fiesta.inference#

Components for Bayesian parameter estimation. Model classes (surrogate, analytical, and combined) now live in fiesta.models.

Likelihood#

Functions for computing likelihoods of data given a model.

class fiesta.inference.likelihood.EMLikelihood(model, data, trigger_time, data_tmin=0.0, data_tmax=999.0, filters=None, error_budget=0.3, conversion_function=<function EMLikelihood.<lambda>>, fixed_params={}, detection_limit=None)[source]#

Bases: LikelihoodBase

Likelihood object to compute likelihoods for the model parameters and a set of magnitude data points.

Parameters:
  • model (FiestaModel) – Light curve model that generates the estimated light curve from the parameters passed to evaluate.

  • data (dict[str, Float[Array, "ntimes 3"]]) – Dictionary with photometric filters as keys and arrays as values. The first column of the array are the detection times in MJD. The second column the magnitude data points. The third column are the Gaussian measurement errors. If an error is np.inf, the data point will be treated as an upper limit on the light curve.

  • trigger_time (Float) – Trigger time or start point of the light curve in MJD.

  • data_tmin (Float) – Time point (in observer frame, relative to trigger_time) before any data point from data will be cropped. Defaults to 0.0.

  • data_tmax (Float, default: 999.0) – Time point (in observer frame, relative to trigger_time) after which any data point from data will be cropped. Defaults to 999.0

  • filters (list[str]) – Filters that should be used for the likelihood evaluation. If None, will take filters from data. Defaults to None.

  • error_budget (Float) – Fixed error budget for the systematic uncertainty. Defaults to 0.3.

  • conversion_function (Callable) – Function that will be called on the params before model predicts the light curve. Defaults to the idenity.

  • fixed_params (dict[str, Float]) – Fixed parameters. These are added to the params before model predicts the light curve. Defaults to {}.

  • detection_limit (Float) – Detection limit of the telescope. If set, a truncated gaussian likelihood will be used. Defaults to None.

times_det#

The time points of the detected magnitudes per filter relative to the trigger time.

Type:

dict[str, Array]

times_nondet#

The time points of the non-detected magnitudes (upper limits) per filter relative to the trigger time.

Type:

dict[str, Array]

datapoints_det#

The detected magnitudes per filter.

Type:

dict[str, Array]

datapoints_nondet#

The non-detection magnitudes (upper limits) per filter.

Type:

dict[str, Array]

datapoints_err#

The gaussian measurement error of the detected magnitudes per filter.

Type:

dict[str, Array]

evaluate(theta)[source]#

Evaluate the log-likelihood of the data given the model and the parameters theta, at a single point.

Parameters:

theta (dict[str, Array]) – A dictionary containing the parameters used to generate the model light curve that is then used to compute the loglikelihood.

Returns:

The log-likelihood value at this parameter point.

Return type:

Float

class fiesta.inference.likelihood.FluxLikelihood(model, data, trigger_time, data_tmin=0.0, data_tmax=999.0, filters=None, error_budget=1, conversion_function=<function FluxLikelihood.<lambda>>, fixed_params={}, detection_limit=None, zero_point_mag=16.4)[source]#

Bases: LikelihoodBase

Likelihood object to compute likelihoods for the model parameters and a set of flux data points. Note that the data in the input argument still needs to be magnitudes. They will be converted internally to fluxes.

Parameters:
  • model (LightcurveModel | AnalyticalModel) – Light curve model that generates the estimated light curve from the parameters passed to evaluate.

  • data (dict[str, Float[Array, "ntimes 3"]]) – Dictionary with photometric filters as keys and arrays as values. The first column of the array are the detection times in MJD. The second column the magnitude data points. The third column are the Gaussian measurement errors. If an error is np.inf, the data point will be treated as an upper limit on the light curve.

  • trigger_time (Float) – Trigger time or start point of the light curve in MJD.

  • data_tmin (Float) – Time point (in observer frame, relative to trigger_time) before any data point from data will be cropped. Defaults to 0.0.

  • data_tmax (Float, default: 999.0) – Time point (in observer frame, relative to trigger_time) after which any data point from data will be cropped. Defaults to 999.0

  • filters (list[str]) – Filters that should be used for the likelihood evaluation. If None, will take filters from data. Defaults to None.

  • error_budget (Float) – Fixed error budget for the systematic uncertainty. Defaults to 1 mJy.

  • conversion_function (Callable) – Function that will be called on the params before model predicts the light curve. Defaults to the idenity.

  • fixed_params (dict[str, Float]) – Fixed parameters. These are added to the params before model predicts the light curve. Defaults to {}.

  • detection_limit (Float) – Detection limit of the telescope. If set, a truncated gaussian likelihood will be used. Defaults to None.

  • zero_point_mag (Float, default: 16.4) – Zero-point for mag-to-flux conversion, specifically to mJy (defaults to 16.4 for AB mag).

times_det#

The time points of the detected fluxes per filter relative to the trigger time.

Type:

dict[str, Array]

times_nondet#

The time points of the non-detected fluxes (upper limits) per filter relative to the trigger time.

Type:

dict[str, Array]

datapoints_det#

The detected fluxs per filter.

Type:

dict[str, Array]

datapoints_nondet#

The non-detection fluxes (upper limits) per filter.

Type:

dict[str, Array]

datapoints_err#

The gaussian measurement error of the detected fluxes per filter.

Type:

dict[str, Array]

evaluate(theta)[source]#

Evaluate the log-likelihood of the data given the model and the parameters theta, at a single point.

Parameters:

theta (dict[str, Array]) – A dictionary containing the parameters used to generate the model light curve that is then used to compute the loglikelihood.

Returns:

The log-likelihood value at this parameter point.

Return type:

Float

mag_to_flux(mag_arr)[source]#

Converts mag_arr to fluxes in mJy.

class fiesta.inference.likelihood.LikelihoodBase(model, data, trigger_time, data_tmin=0.0, data_tmax=999.0, filters=None, error_budget=0.3, conversion_function=<function LikelihoodBase.<lambda>>, fixed_params={}, detection_limit=None)[source]#

Bases: object

Base class for likelihoods.

static compute_gaussian_likelihood(y_est, y_data, sigma, lim)[source]#

Return the log likelihood of the chisquare part of the likelihood function, without truncation (no detection limit is given), i.e. a Gaussian pdf.

Return type:

Float

static compute_gaussian_survival(y_est, y_data, error_budget)[source]#
Return type:

Float

static compute_trunc_gaussian_likelihood(y_est, y_data, sigma, lim)[source]#

Return the log likelihood of the chisquare part of the likelihood function, with truncation of the Gaussian (detection limit is given).

Return type:

Float

cut_data_to_time_range(data, data_tmin, data_tmax)[source]#
Return type:

dict[str, Array, 'ntimes 3']]

data_tmax: Float#
data_tmin: Float#
datapoints_det: dict[str, Array]#
datapoints_err: dict[str, Array]#
datapoints_nondet: dict[str, Array]#
detection_limit: dict[str, Array]#
error_budget: dict[str, Array]#
evaluate(theta)[source]#

Evaluate the log-likelihood of the data given the parameters in theta and the underlying model.

Return type:

Float

filters: list[str]#
get_gaussprob_det(y_est, y_data, sigma, lim)[source]#

Return the log likelihood of the gaussian likelihood function for a single filter. Branch-off of jax.lax.cond is based on provided detection limit (lim). If the limit is infinite, the likelihood is calculated without truncation and without resorting to scipy for faster evaluation. If the limit is finite, the likelihood is calculated with truncation and with scipy.

Parameters:
  • y_est (Array) – The estimated data from the model at detection times.

  • y_data (Array) – The detected data.

  • sigma (Array) – The uncertainties on the detected apparent magnitudes, including the error budget.

  • lim (Float) – The detection limit for this filter.

Returns:

The gaussian log-likelihood for this filter.

Return type:

Float

get_gaussprob_nondet(y_est, y_data, error_budget)[source]#

Return the log likelihood of the gaussian likelihood function for a single filter. Branch-off of jax.lax.cond is based on provided detection limit (lim). If the limit is infinite, the likelihood is calculated without truncation and without resorting to scipy for faster evaluation. If the limit is finite, the likelihood is calculated with truncation and with scipy.

Parameters:
  • y_est (Array) – The estimated data from the model at detection times.

  • y_data (Array) – The nondetection data points.

  • sigma (Array) – The uncertainties on the detected apparent magnitudes, including the error budget.

  • lim (Float) – The detection limit for this filter.

Returns:

The gaussian log-likelihood for this filter.

Return type:

Float

model: FiestaModel#
setup_filters_and_data(filters, data)[source]#
Return type:

dict[str, Array, 'ntimes 3']]

times_det: dict[str, Array]#
times_nondet: dict[str, Array]#
trigger_time: Float#
vectorized_evaluate(theta)[source]#

Priors#

class fiesta.inference.prior.CompositePrior(priors, transforms={}, **kwargs)[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

priors: list[Prior] = Field(name=None,type=None,default=<dataclasses._MISSING_TYPE object>,default_factory=<class 'list'>,init=True,repr=True,hash=None,compare=True,metadata=mappingproxy({}),kw_only=<dataclasses._MISSING_TYPE object>,_field_type=None)#
sample(rng_key, n_samples)[source]#
Return type:

dict[str, Array, 'n_samples']]

class fiesta.inference.prior.ConstrainedPrior(priors, conversion_function=<function ConstrainedPrior.<lambda>>, transforms={})[source]#

Bases: CompositePrior

constraints: list[Constraint]#
conversion: Callable#
evaluate_constraints(samples)[source]#
factor: Float#
log_prob(x)[source]#
Return type:

Float

sample(rng_key, n_samples)[source]#
Return type:

dict[str, Array, 'n_samples']]

class fiesta.inference.prior.Constraint(naming, xmin, xmax, transforms={})[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

xmax: float#
xmin: float#
class fiesta.inference.prior.InterpedPrior(xx, yy, naming)[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

sample(rng_key, n_samples)[source]#
Return type:

dict[str, Array, 'n_samples']]

xx: Array#
yy: Array#
class fiesta.inference.prior.LogUniform(xmin, xmax, naming, transforms={}, **kwargs)[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

sample(rng_key, n_samples)[source]#

Sample from a uniform distribution.

Parameters:
  • rng_key (PRNGKeyArray) – A random key to use for sampling.

  • n_samples (int) – The number of samples to draw.

Returns:

samples – Samples from the distribution. The keys are the names of the parameters.

Return type:

dict

xmax: float = 1.0#
xmin: float = 0.0#
class fiesta.inference.prior.Normal(mu, sigma, naming, transforms={}, **kwargs)[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

mu: float = 0.0#
sample(rng_key, n_samples)[source]#

Sample from a normal distribution.

Parameters:
  • rng_key (PRNGKeyArray) – A random key to use for sampling.

  • n_samples (int) – The number of samples to draw.

Returns:

samples – Samples from the distribution. The keys are the names of the parameters.

Return type:

dict

sigma: float = 1.0#
class fiesta.inference.prior.Prior(naming, transforms={})[source]#

Bases: object

A thin base clase to do book keeping.

Should not be used directly since it does not implement any of the real method.

The rationale behind this is to have a class that can be used to keep track of the names of the parameters and the transforms that are applied to them.

add_name(x)[source]#

Turn an array into a dictionary

Parameters:

x (Array) – An array of parameters. Shape (n_dim,).

Return type:

dict[str, Float]

log_prob(x)[source]#
Return type:

Float

property n_dim#
naming: list[str]#
sample(rng_key, n_samples)[source]#
Return type:

dict[str, Array, 'n_samples']]

transform(x)[source]#

Apply the transforms to the parameters.

Parameters:

x (dict) – A dictionary of parameters. Names should match the ones in the prior.

Returns:

x – A dictionary of parameters with the transforms applied.

Return type:

dict

transforms: dict[str, tuple[str, Callable]] = Field(name=None,type=None,default=<dataclasses._MISSING_TYPE object>,default_factory=<class 'dict'>,init=True,repr=True,hash=None,compare=True,metadata=mappingproxy({}),kw_only=<dataclasses._MISSING_TYPE object>,_field_type=None)#
class fiesta.inference.prior.Sine(naming, xmin=0, xmax=3.141592653589793, transforms={})[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

sample(rng_key, n_samples)[source]#

Sample from a uniform distribution.

Parameters:
  • rng_key (PRNGKeyArray) – A random key to use for sampling.

  • n_samples (int) – The number of samples to draw.

Returns:

samples – Samples from the distribution. The keys are the names of the parameters.

Return type:

dict

xmax: float = 1.0#
xmin: float = 0.0#
class fiesta.inference.prior.TruncatedNormal(mu, sigma, xmin, xmax, naming, transforms={}, **kwargs)[source]#

Bases: Prior

Truncated normal distribution with explicit bounds.

Useful for informed priors from population studies (e.g., superphot+). The SVISampler uses xmin/xmax for its guide constraints and mu/sigma for the model’s TruncatedNormal distribution.

log_prob(x)[source]#
Return type:

Float

mu: float = 0.0#
sample(rng_key, n_samples)[source]#
Return type:

dict[str, Array, 'n_samples']]

sigma: float = 1.0#
xmax: float = 10.0#
xmin: float = -10.0#
class fiesta.inference.prior.Uniform(xmin, xmax, naming, transforms={}, **kwargs)[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

sample(rng_key, n_samples)[source]#

Sample from a uniform distribution.

Parameters:
  • rng_key (PRNGKeyArray) – A random key to use for sampling.

  • n_samples (int) – The number of samples to draw.

Returns:

samples – Samples from the distribution. The keys are the names of the parameters.

Return type:

dict

xmax: float = 1.0#
xmin: float = 0.0#
class fiesta.inference.prior.UniformSourceFrame(dmin, dmax, naming, cosmology=FlatLambdaCDM(name='Planck18', H0=<Quantity 67.66 km / (Mpc s)>, Om0=0.30966, Tcmb0=<Quantity 2.7255 K>, Neff=3.046, m_nu=<Quantity [0., 0., 0.06] eV>, Ob0=0.04897), **kwargs)[source]#

Bases: InterpedPrior

xmax: float = 100000.0#

Prior that is uniform in comoving volume and source frame time, analogue to the corresponding bilby prior. Uses the default cosmology in fiesta which is Planck18.

xmin: float = 10.0#
class fiesta.inference.prior.UniformVolume(xmin, xmax, naming, transforms={}, **kwargs)[source]#

Bases: Prior

log_prob(x)[source]#
Return type:

Float

sample(rng_key, n_samples)[source]#

Sample luminosity distance from a distribution uniform in volume.

Parameters:
  • rng_key (PRNGKeyArray) – A random key to use for sampling.

  • n_samples (int) – The number of samples to draw.

Returns:

samples – Samples from the distribution. The keys are the names of the parameters.

Return type:

dict

xmax: float = 100000.0#
xmin: float = 10.0#

Systematics#

fiesta.inference.systematic.check_filter_compatability(yaml_dict, filters)[source]#
fiesta.inference.systematic.fetch_prior_params(yaml_entry)[source]#
fiesta.inference.systematic.process_file(systematic_file, filters)[source]#
fiesta.inference.systematic.setup_systematic_from_file(likelihood, prior, systematics_file)[source]#
fiesta.inference.systematic.setup_systematics_basic(likelihood, prior, error_budget=0.3)[source]#

Sampler#

Utilities#

Functions for creating and handling injections

class fiesta.inference.injection.InjectionAfterglowpy(jet_type=-1, *args, **kwargs)[source]#

Bases: InjectionBase

class fiesta.inference.injection.InjectionBase(filters, trigger_time, tmin=0.1, tmax=10.0, N_datapoints=10, t_detect=None, error_budget=1.0, nondetections=False, nondetections_fraction=0.1, detection_limit=inf)[source]#

Bases: object

Base class to create synthetic injection lightcurves. The injection model is first initialized with the following parameters:

filters (list): List of filters in which the synthetic data should be given out. trigger_time (float): Reference trigger time (e.g. MJD or GPS seconds) added as an offset to all detection time stamps. Required. tmin (float): Time of earliest synthetic detection possible in days. Defaults to 0.1. tmax (float): Time of latest synthetic detection possible in days. Defaults to 10.0 N_datapoints (int): Total number of datapoints (across all filters) for the synthetic lightcurve. Defaults to 10. t_detect (dict[str, Array]): Detection time points in each filter. If none is specified, then the detection times will be sampled randomly. error_budget (float): Typical measurement error scale of the synthetic data. Defaults to 1. detection_limit (float): Synthetic datapoints with mangnitude higher than this value (i.e. less brighter) will be turned into nondetections. Defaults to np.inf. nondetections (bool): Additional to detection_limit, this turns some of the synthetic datapoints to nondetections. Defaults to False. nondetections_fraction: If nondetections is True, then this will determine the fractions of N_datapoints turned into nondetections. Defaults to 0.1.

Then one can call the .create_injection() method to get synthetic lightcurve data. The method .write_to_file() writes the synthetic lightcurve data to file.

create_injection(injection_dict, file=None)[source]#

Creates an injection that is stored as a .data attribute.

Parameters:
  • injection_dict (dict) – Parameters for the synthetic light curve.

  • file (str, optional) – Training data file that stores light curves from the physical base model of the surrogate. If provided, the method will take a random test element and base the injection on it. In this case, the .injection_parameter attribute is updated to contain the real parameters used to generate the light curve.

create_injection_from_mags(times, mag_app)[source]#
create_t_detect(tmin, tmax, N)[source]#

Create a time grid for the injection data.

randomize_nondetections()[source]#
write_to_file(file)[source]#
class fiesta.inference.injection.InjectionKN(*args, **kwargs)[source]#

Bases: InjectionBase

class fiesta.inference.injection.InjectionPyblastafterglow(jet_type='tophat', *args, **kwargs)[source]#

Bases: InjectionBase

class fiesta.inference.injection.InjectionSurrogate(model, *args, **kwargs)[source]#

Bases: InjectionBase

Class to create synthetic injection lightcurves from a surrogate. After instantiation one can call the .create_injection() method to get synthetic lightcurve data. The method .write_to_file() writes the synthetic lightcurve data to file.

fiesta.inference.injection.get_parser(**kwargs)[source]#
class fiesta.inference.plot.LightcurvePlotter(posterior, likelihood, systematics_file=None, free_syserr=False)[source]#

Bases: object

Interface to plot lightcurves from a given posterior.

Parameters:
  • posterior (dict | pd.DataFrame) – Posterior samples for which the light curves should be plotted.

  • likelihood (EMLikelihood) – Likelihood object that was used to sample the posterior.

  • systematics_file (str) – Systematics file that was used to sample the posterior. Defaults to None.

  • free_syserr (bool) – Whether a global systematic uncertainty was sampled freely. Defaults to False. Will overwrite systematics_file.

get_chisquared(per_dof=False)[source]#

Get the total chisquared value and the chisquared values per filter. This is different from the log_likelihood value in the posterior, because the likelihood function contains (2 pi sigma)^(-1/2).

Parameters:

per_dof (bool) – Whether to return reduced chi-squared values, i.e., per number of data points.

Returns:

The total chi-squared value across all data points and a dict with the chi-squared value in each filter.

Return type:

tuple(float, dict)

plot_best_fit_lc(ax, filt, zorder=2, **kwargs)[source]#

Plots one filter from the best fit light curve from the posterior over ax.

Parameters:
  • ax (matplotlib.axes.Axes) – ax to plot the light curve onto.

  • filt (str) – Which filter from the best fit lightcurve should be plotted on ax.

  • zorder (int) – zorder with which the lightcurve should be plotted.

  • **kwargs – kwargs to be passed to plot.

plot_data(ax, filt, zorder=3, **kwargs)[source]#

Plots data points from a filter over ax.

Parameters:
  • ax (matplotlib.axes.Axes) – ax to plot the data points to.

  • filt (str) – Which filter from the data should be plotted on ax.

  • zorder (int) – zorder with which the data points should be plotted.

  • **kwargs – kwargs to be passed to errorbar and scatter.

plot_sample_lc(ax, filt, zorder=1)[source]#

Plots background light curves from the posterior over ax.

Parameters:
  • ax (matplotlib.axes.Axes) – ax to plot the light curve onto.

  • filt (str) – Which filter from the background light curves should be plotted on ax.

  • zorder (int) – zorder with which the lightcurve should be plotted.

plot_sys_uncertainty_band(ax, filt, zorder=2, **kwargs)[source]#

Plots systematic uncertainty band from the best fit light curve for one filter over ax.

Parameters:
  • ax (matplotlib.axes.Axes) – ax to plot the band onto.

  • filt (str) – Which filter from the band should be plotted on ax.

  • zorder (int) – zorder with which the band should be plotted.

  • **kwargs – kwargs to be passed to fill_between.

fiesta.inference.plot.corner_plot(posterior, parameter_names, truths=None, color='blue', legend_label=None, fig=None, ax=None, **kwargs)[source]#

Make a nice corner plot from the posterior with automated parameter labels.

Parameters:
  • posterior (dict | pd.DataFrame) – posterior samples for which to do the corner plot.

  • parameter_names (list[str]) – parameters from posterior that should be included in the corner plot.

  • truths (dict[str, float] | None) – True (injected values) for some of the parameters. Defaults to None.

  • color (str) – color for the corner plot contours. Defaults to blue.

  • legend_label (str) – Label for the legend. If not set, no legend will be shown. Defaults to None.

  • fig (matplotlib.figure.Figure) – Figure over which to do the corner plot. If set, ax must also be provided. Defaults to None.

  • ax (matplotlib.axes.Axes) – Axes over which to do the corner plot. If set, fig must also be provided. Defaults to None.

Returns:

Figure with the corner plot. ax (matplotlib.axes.Axes): array of axes

Return type:

fig (matplotlib.figure.Figure)