Source code for fiesta.utils

import copy
from multiprocessing import Value

import warnings
warnings.filterwarnings("ignore", module="erfa")

import numpy as np
import h5py
from astropy.time import Time
import scipy.interpolate as interp

import jax
import jax.numpy as jnp
from jax.scipy.stats import truncnorm
from jaxtyping import Array, Float, Int




##########################
### I/O DATA UTILITIES ###
##########################

[docs] def load_event_data(filename): """ Takes a file and outputs a magnitude dict with filters as keys. Args: filename (str): path to file to be read in Returns: data (dict[str, Array]): Data dictionary with filters as keys. The array has the structure [[mjd, mag, err]]. """ mjd, filters, mags, mag_errors = [], [], [], [] with open(filename, "r") as input: for line in input: line = line.rstrip("\n") t, filter, mag, mag_err = line.split(" ") mjd.append(Time(t, format="isot").mjd) # convert to mjd filters.append(filter) mags.append(float(mag)) mag_errors.append(float(mag_err)) mjd = np.array(mjd) filters = np.array(filters) mags = np.array(mags) mag_errors = np.array(mag_errors) data = {} unique_filters = np.unique(filters) for filt in unique_filters: filt_inds = np.where(filters==filt)[0] data[filt] = np.array([ mjd[filt_inds], mags[filt_inds], mag_errors[filt_inds] ]).T return data
[docs] def write_event_data(filename: str, data: dict): """ Takes a magnitude dict and writes it to filename. The magnitude dict should have filters as keys, the arrays should have the structure [[mjd, mag, err]]. """ with open(filename, "w") as out: for filt in data.keys(): for data_point in data[filt]: time = Time(data_point[0], format = "mjd") filt_name = filt.replace("_", ":") line = f"{time.isot} {filt_name} {data_point[1]:f} {data_point[2]:f}" out.write(line +"\n")
[docs] def truncated_gaussian(mag_det: Array, mag_err: Array, mag_est: Array, lim: Float = jnp.inf): """ Evaluate log PDF of a truncated Gaussian with loc at mag_est and scale mag_err, truncated at lim above. Returns: _type_: _description_ """ loc, scale = mag_est, mag_err a_trunc = -999 # TODO: OK if we just fix this to a large number, to avoid infs? a, b = (a_trunc - loc) / scale, (lim - loc) / scale logpdf = truncnorm.logpdf(mag_det, a, b, loc=loc, scale=scale) return logpdf
############## ### LEGACY ### ##############
[docs] def interpolate_nans(data: dict[str, Float[Array, " n_files n_times"]], times: Array, output_times: Array = None) -> dict[str, Float[Array, " n_files n_times"]]: """ Interpolate NaNs and infs in the raw light curve data. Args: data (dict[str, Float[Array, 'n_files n_times']]): The raw light curve data diagnose (bool): If True, print out the number of NaNs and infs in the data etc to inform about quality of the grid. Returns: dict[str, Float[Array, 'n_files n_times']]: Raw light curve data but with NaNs and infs interpolated """ if output_times is None: output_times = times # TODO: improve this function overall! copy_data = copy.deepcopy(data) output = {} for filt, lc_array in copy_data.items(): n_files = np.shape(lc_array)[0] if filt == "t": continue for i in range(n_files): lc = lc_array[i] # Get NaN or inf indices nan_idx = np.isnan(lc) inf_idx = np.isinf(lc) bad_idx = nan_idx | inf_idx good_idx = ~bad_idx # Interpolate through good values on given time grid if len(good_idx) > 1: # Make interpolation routine at the good idx good_times = times[good_idx] good_mags = lc[good_idx] interpolator = interp.interp1d(good_times, good_mags, fill_value="extrapolate") # Apply it to all times to interpolate mag_interp = interpolator(output_times) else: raise ValueError("No good values to interpolate from") if filt in output: output[filt] = np.vstack((output[filt], mag_interp)) else: output[filt] = np.array(mag_interp) return output