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