import logging
import time
import warnings
from collections.abc import Callable
from dataclasses import dataclass
import numpy as np
from astropy.nddata.bitmask import bitfield_to_boolean_mask
from astropy.utils.decorators import lazyproperty
from scipy import ndimage
from scipy.interpolate import CubicSpline, UnivariateSpline
from stdatamodels.jwst import datamodels
from stdatamodels.jwst.datamodels import SossWaveGridModel, dqflags
from jwst.datamodels.utils.tso_multispec import make_tso_specmodel
from jwst.extract_1d.extract import populate_time_keywords
from jwst.extract_1d.soss_extract.atoca import ExtractionEngine, KernelShapeError, MaskOverlapError
from jwst.extract_1d.soss_extract.atoca_utils import (
WebbKernel,
get_wave_p_or_m,
grid_from_map_with_extrapolation,
make_combined_adaptive_grid,
oversample_grid,
throughput_soss,
)
from jwst.extract_1d.soss_extract.pastasoss import (
_find_spectral_order_index,
_verify_requested_orders,
get_soss_wavemaps,
)
from jwst.extract_1d.soss_extract.soss_boxextract import (
box_extract,
estim_error_nearest_data,
get_box_weights,
)
from jwst.extract_1d.soss_extract.soss_syscor import make_background_mask, soss_background
from jwst.lib import pipe_utils
log = logging.getLogger(__name__)
ORDER_STR_TO_INT = {f"Order {order}": order for order in [1, 2, 3]}
__all__ = ["DetectorModelOrder", "get_ref_file_args", "Integration", "run_extract1d"]
[docs]
@dataclass
class DetectorModelOrder:
"""
Model of detector properties for the ATOCA algorithm for a single spectral order.
Attributes
----------
spectral_order : int
The spectral order number.
wavemap : ndarray
The 2-D map of the expected wavelengths at each pixel for this order
as determined by PASTASOSS.
spectrace : ndarray
The 1-D spectral trace as determined by PASTASOSS.
specprofile : ndarray
The 2-D spatial profile of the spectral trace from the SPECPROFILE reference file.
throughput : callable
An interpolation function for the throughput of the order, computed from the throughput
in the PASTASOSS reference file.
kernel : `~jwst.extract_1d.soss_extract.atoca_utils.WebbKernel`, ndarray, or None
The spectral resolution kernel for this order on the input wavelength grid,
either as a `~jwst.extract_1d.soss_extract.atoca_utils.WebbKernel`
object (callable) or as a 2-D array.
kernel_native : `~jwst.extract_1d.soss_extract.atoca_utils.WebbKernel`, ndarray, or None
The spectral resolution kernel for this order on the native pixel grid,
either as a `~jwst.extract_1d.soss_extract.atoca_utils.WebbKernel`
object (callable) or as a 2-D array.
kernel_func : `~jwst.extract_1d.soss_extract.atoca_utils.WebbKernel` or None
The spectral resolution kernel for this order. This is intended not to be modified, so the
`~jwst.extract_1d.soss_extract.atoca.ExtractionEngine`
can use it to generate a kernel at any wavelength grid.
subarray : str or None
The name of the subarray used for the extraction.
trace_cutoff : int
The maximum x pixel to consider for this order, as determined by
the PASTASOSS reference file.
short_cutoff : float or None
The shortest wavelength to consider for this order, in microns,
as determined by the PASTASOSS reference file.
orderengine : `~jwst.extract_1d.soss_extract.atoca.ExtractionEngine` or None
Precomputed extraction engine for this order. This is intended to be
computed for the first integration and stored thereafter.
order12engine : `~jwst.extract_1d.soss_extract.atoca.ExtractionEngine` or None
Precomputed extraction engine for orders 1+2. This is intended to be
computed for the first integration and stored thereafter.
mederr : ndarray or None
Median error across integrations
m_inv : ndarray or None
Inverse of the design matrix of the regularized least squares problem.
Intended to be computed for the first integration and stored thereafter.
bmat : ndarray or None
Matrix of pixel values divided by the median error.
Intended to be computed for the first integration and stored thereafter.
mask : ndarray or None
Mask for pixels valid in at least some integrations.
Intended to be computed for the first integration and stored thereafter.
"""
spectral_order: int
wavemap: np.ndarray
spectrace: np.ndarray
specprofile: np.ndarray
throughput: Callable
kernel: WebbKernel | np.ndarray | None = None
kernel_native: WebbKernel | None = None
kernel_func: WebbKernel | None = None
subarray: str | None = None
trace_cutoff: int = 2048
short_cutoff: float | None = None
orderengine: ExtractionEngine | None = None
order12engine: ExtractionEngine | None = None
mederr: np.ndarray | None = None
m_inv: np.ndarray | None = None
bmat: np.ndarray | None = None
mask: np.ndarray | None = None
@lazyproperty
def trace(self):
"""
Get the x, y, wavelength of the trace after applying the transform.
Returns
-------
xtrace, ytrace, wavetrace : ndarray
The x, y, and wavelength of the trace, respectively.
"""
spectrace = self.spectrace
xtrace = np.arange(self.trace_cutoff)
# CubicSpline requires monotonically increasing x arr
if spectrace[0][0] - spectrace[0][1] > 0:
spectrace = np.flip(spectrace, axis=1)
trace_interp_y = CubicSpline(spectrace[0], spectrace[1])
trace_interp_wave = CubicSpline(spectrace[0], spectrace[2])
ytrace = trace_interp_y(xtrace)
wavetrace = trace_interp_wave(xtrace)
return xtrace, ytrace, wavetrace
@lazyproperty
def native_grid(self):
"""
Make a 1D grid of the pixels boundary based on the wavelength solution.
Returns
-------
wave : ndarray
Grid of the pixels boundaries at the native sampling (1D array)
col : ndarray
The column number of the pixel
"""
# From wavelength solution
col, _, wave = self.trace
# Keep only valid solution ...
idx_valid = np.isfinite(wave)
# ... and should correspond to subsequent columns
is_subsequent = np.diff(col[idx_valid]) == 1
if not is_subsequent.all():
msg = f"Wavelength solution for order {self.spectral_order} contains gaps."
log.warning(msg)
wave = wave[idx_valid]
col = col[idx_valid]
log.debug(f"Wavelength range for order {self.spectral_order}: ({wave[[0, -1]]})")
# Sort
idx_sort = np.argsort(wave)
wave = wave[idx_sort]
col = col[idx_sort]
return wave, col
[docs]
def get_grid_from_trace(self, n_os):
"""
Make a 1D grid of the pixels boundary based on the wavelength solution.
Parameters
----------
n_os : int or ndarray
The oversampling factor of the wavelength grid used when solving for
the uncontaminated flux.
Returns
-------
ndarray
Grid of the pixels boundaries at the native sampling (1D array).
"""
wave, _ = self.native_grid
# Use pixel boundaries instead of the center values
wv_upper_bnd, wv_lower_bnd = get_wave_p_or_m(wave[None, :])
# `get_wave_p_or_m` returns 2d array, so keep only 1d
wv_upper_bnd, wv_lower_bnd = wv_upper_bnd[0], wv_lower_bnd[0]
# Each upper boundary should correspond the the lower boundary
# of the following pixel, so only need to add the last upper boundary to complete the grid
wv_upper_bnd, wv_lower_bnd = np.sort(wv_upper_bnd), np.sort(wv_lower_bnd)
wave_grid = np.append(wv_lower_bnd, wv_upper_bnd[-1])
# Oversample as needed
return oversample_grid(wave_grid, n_os=n_os)
[docs]
def get_ref_file_args(ref_files, orders_requested=None):
"""
Prepare the reference files for the extraction engine.
Parameters
----------
ref_files : dict
A dictionary of the reference file datamodels, along with values
for the subarray and pwcpos, i.e., the pupil wheel position.
orders_requested : list or None, optional
A list of the spectral orders requested for extraction.
If None, all orders in the PASTASOSS reference file are used.
Returns
-------
list
A list of `DetectorModelOrder` objects containing PASTASOSS outputs and other reference
information, one per spectral order to be extracted.
"""
pastasoss_ref = ref_files["pastasoss"]
specprofile_ref = ref_files["spec_profiles"]
speckernel_ref = ref_files["speckernel"]
n_pix = 2 * speckernel_ref.meta.halfwidth + 1
speckernel_wv_range = [np.min(speckernel_ref.wavelengths), np.max(speckernel_ref.wavelengths)]
refmodel_orders = [int(trace.spectral_order) for trace in pastasoss_ref.traces]
if orders_requested is None:
orders_requested = refmodel_orders
else:
orders_requested = _verify_requested_orders(orders_requested, refmodel_orders)
wavemaps, spectraces = get_soss_wavemaps(
ref_files["pwcpos"],
refmodel=pastasoss_ref,
subarray=ref_files["subarray"],
spectraces=True,
orders_requested=orders_requested,
)
# Collect spectral profiles, wavemaps, throughputs, kernels for all the orders
detector_models = []
for order in orders_requested:
order_idx = _find_spectral_order_index(pastasoss_ref, order)
wavemap = wavemaps[order_idx]
specprofile = specprofile_ref.profile[order_idx].data
# apply padding to make specprofile and wavemap have same shape
prof_shape0, prof_shape1 = specprofile.shape
wavemap_shape0, wavemap_shape1 = wavemap.shape
if prof_shape0 != wavemap_shape0:
pad0 = (prof_shape0 - wavemap_shape0) // 2
if pad0 > 0:
specprofile = specprofile[pad0:-pad0, :]
elif pad0 < 0:
wavemap = wavemap[pad0:-pad0, :]
if prof_shape1 != wavemap_shape1:
pad1 = (prof_shape1 - wavemap_shape1) // 2
if pad1 > 0:
specprofile = specprofile[:, pad1:-pad1]
elif pad1 < 0:
wavemap = wavemap[:, pad1:-pad1]
# make throughput interpolator
thru = throughput_soss(
pastasoss_ref.throughputs[order_idx].wavelength[:],
pastasoss_ref.throughputs[order_idx].throughput[:],
)
trace_cutoff = pastasoss_ref.traces[order_idx].cutoff
short_cutoff = getattr(pastasoss_ref.traces[order_idx], "short_cutoff", None)
detector_model = DetectorModelOrder(
spectral_order=order,
subarray=ref_files["subarray"],
wavemap=wavemap,
specprofile=specprofile,
throughput=thru,
kernel=None,
spectrace=spectraces[order_idx],
trace_cutoff=trace_cutoff,
short_cutoff=short_cutoff,
)
# Build a kernel for this order
wv_cent = np.zeros(wavemap.shape[1])
# Get central wavelength as a function of columns
col, _row, wv = detector_model.trace
wv_cent[col] = wv
# Set invalid values to zero
idx_invalid = ~np.isfinite(wv_cent)
wv_cent[idx_invalid] = 0.0
kernel = WebbKernel(speckernel_ref.wavelengths, speckernel_ref.kernels, wv_cent, n_pix)
valid_wavemap = (speckernel_wv_range[0] <= wavemap) & (wavemap <= speckernel_wv_range[1])
wavemap = np.where(valid_wavemap, wavemap, 0.0)
detector_model.kernel = kernel
detector_model.kernel_native = kernel
detector_model.kernel_func = kernel
detector_models.append(detector_model)
return detector_models
def _infill_data_get_mederr(cube_model, refmask, ninterp=9):
"""
Impute missing data, compute median uncertainty.
Parameters
----------
cube_model : `~stdatamodels.jwst.datamodels.CubeModel`
The input datamodel.
refmask : ndarray
Boolean mask for the reference pixels.
ninterp : int, optional
Number of neighboring images (in time) to use to impute missing data.
Used with `scipy.ndimage.median_filter`. Should be odd.
Returns
-------
data_nanreplaced : ndarray
Array of shape ``(nintegrations, ny, nx)`` with NaNs replaced with the
mean of ``ninterp`` neighboring images.
medarr : ndarray
Median uncertainty from ``cube_model.err``, excluding NaNs.
"""
# Pixels that are bad in all integrations
allbad = np.sum(np.isfinite(cube_model.err) & np.isfinite(cube_model.data), axis=0) == 0
# Make a copy of the cube model where bad pixels in any individual
# integration are replaced with a running median in time.
data_nanreplaced = cube_model.data.copy()
data_infilled = cube_model.data.copy()
# We will use the median uncertainty throughout the calculation.
# Use of the median uncertainty means that the matrices used in
# the ATOCA algorithm are shared between all integrations.
with warnings.catch_warnings():
warnings.filterwarnings("ignore", "All-NaN slice encountered", RuntimeWarning)
mederr = np.nanmedian(cube_model.err, axis=0)
meddata = np.nanmedian(data_nanreplaced, axis=0)
mederr[allbad] = np.inf
mederr[refmask] = np.inf
meddata[allbad] = 0
# Fill in values for pixels that are bad in individual integrations
for i in range(data_nanreplaced.shape[0]):
indx = ~np.isfinite(data_nanreplaced[i])
data_infilled[i][indx] = meddata[indx]
# Apply a median filter in time to make the final replacements of bad
# pixels in individual integrations. Use a running median of ninterp
# time steps as a compromise between precision and time resolution.
nints = data_nanreplaced.shape[0]
medfilt = ndimage.median_filter(data_infilled, min(nints, ninterp), axes=(0))
indx = ~np.isfinite(data_nanreplaced)
data_nanreplaced[indx] = medfilt[indx]
return data_nanreplaced, mederr
def _refine_tikfac(tiktests, tikho_struct, engine, niter_refine=3):
"""
Refine the optimal Tikhonov regularization factor using iterations.
This algorithm takes the results of Tikhonov optimization over the
initial, coarse grid as input. It then finds the best factor using
the dchi2/dlog(factor) criterion. It chooses the closest point on
the grid, adds two new points on either side of this to increase the
resolution (or extends the grid if the best-fit was an endpoint),
and recomputes the best Tikhonov factor. This process is iterated
niter_refine times.
Parameters
----------
tiktests : `~jwst.extract_1d.soss_extract.atoca_utils.TikhoTests`
Dictionary-like object with Tikhonov test factors and associated
goodness-of-fit quantities.
tikho_struct : `~jwst.extract_1d.soss_extract.atoca_utils.Tikhonov`
A Tikhonov structure appropriately initialized to enable the
calculation of goodness-of-fit metrics for trial Tikhonov factors.
engine : `~jwst.extract_1d.soss_extract.atoca.ExtractionEngine`
Used to determine the best Tikhonov factor after the goodness-
of-fit quantities have been calculated.
niter_refine : int
Number of times to refine the calculation by adding two points.
Returns
-------
tifac : float
The best Tikhonov factor.
"""
tikfac = engine.best_tikho_factor(tiktests, fit_mode="d_chi2")
for _i in range(niter_refine):
f = np.sort(tiktests["factors"])
closest_point = np.abs(np.log(f) - np.log(tikfac)) == np.amin(
np.abs(np.log(f) - np.log(tikfac))
)
# Add two points, extrapolating if necessary.
j = np.where(closest_point)[0][0]
if j == 0:
newfac = np.array([f[0] ** 2 / f[1], np.sqrt(f[0] * f[1])])
elif j == len(f) - 1:
newfac = np.array([np.sqrt(f[-1] * f[-2]), f[-1] ** 2 / f[-2]])
else:
newfac = np.array([np.sqrt(f[j] * f[j - 1]), np.sqrt(f[j] * f[j + 1])])
# Merge results and recompute the best Tikhonov factor.
newtest = engine.get_tikho_tests(tikho_struct, newfac)
tiktests.merge(newtest)
tikfac = engine.best_tikho_factor(tiktests, fit_mode="d_chi2")
return tikfac, tiktests
def _populate_tikho_attr(spec, tiktests, idx, sp_ord):
spec.spectral_order = sp_ord
spec.meta.soss_extract1d.type = "TEST"
spec.meta.soss_extract1d.chi2 = tiktests["chi2"][idx]
spec.meta.soss_extract1d.chi2_soft_l1 = tiktests["chi2_soft_l1"][idx]
spec.meta.soss_extract1d.chi2_cauchy = tiktests["chi2_cauchy"][idx]
spec.meta.soss_extract1d.reg = np.nansum(tiktests["reg"][idx] ** 2)
spec.meta.soss_extract1d.factor = tiktests["factors"][idx]
spec.int_num = 0 # marks this as a test spectrum
def _build_tracemodel_order(engine, order_model, f_k, mask, force_recompute_engine=False):
"""
Build the trace model for a specific spectral order.
Parameters
----------
engine : `~jwst.extract_1d.soss_extract.atoca.ExtractionEngine`
The extraction engine used to extract the spectrum.
order_model : `DetectorModelOrder`
The model for the spectral order to be extracted.
f_k : ndarray
The extracted flux for the spectral order.
mask : ndarray
The global mask of pixels to be modeled. Bad pixels in the science data
should remain unmasked.
force_recompute_engine : bool, optional
Force the recomputation of the
`~jwst.extract_1d.soss_extract.atoca.ExtractionEngine` inside this function,
and prevent it from being saved to the order_model? Alternative is to
use or generate the engine already associated with ``order_model``. This
boolean should be `True` if reconstructing spectra within the Tikhonov
tests, and `False` if using the function within the main integration
routine. If `True`, ``order_model.orderengine`` will not be modified.
Default is `False`.
Returns
-------
tracemodel_ord : ndarray
The modeled detector image for the spectral order.
spec_ord : `~stdatamodels.jwst.datamodels.SpecModel`
The data model containing the extracted spectrum for the spectral order.
"""
i_order = engine.orders.index(order_model.spectral_order)
# Pre-convolve the extracted flux (f_k) at the order's resolution
# so that the convolution matrix must not be re-computed.
flux_order = engine.kernels[i_order].dot(f_k)
# Then must take the grid after convolution (smaller)
grid_order = engine.wave_grid_c(i_order)
# Keep only valid values to make sure there will be no Nans in the order model
idx_valid = np.isfinite(flux_order)
grid_order, flux_order = grid_order[idx_valid], flux_order[idx_valid]
# Build model of the order
# Give the identity kernel to the Engine (so no convolution)
# Load a precomputed engine if available and a recomputation is not
# forced. Otherwise, compute it now, and store it unless
# force_recompute_engine is True.
if (order_model.orderengine is None) or (force_recompute_engine):
engine = ExtractionEngine(
[order_model.wavemap],
[order_model.specprofile],
[order_model.throughput],
[np.array([1.0])],
wave_grid=grid_order,
mask_trace_profile=[mask],
orders=[order_model.spectral_order],
)
if not force_recompute_engine:
order_model.orderengine = engine
else:
engine = order_model.orderengine
# Project on detector and save in dictionary
tracemodel_ord = engine.rebuild(flux_order, fill_value=np.nan)
# Build 1d spectrum integrated over pixels
pixel_grid, valid_cols = order_model.native_grid
pixel_grid = pixel_grid[np.newaxis, :]
mask = np.all(mask, axis=0)[valid_cols]
engine = ExtractionEngine(
[pixel_grid],
[np.ones_like(pixel_grid)],
[order_model.throughput],
[order_model.kernel_native],
wave_grid=grid_order,
mask_trace_profile=[mask[np.newaxis, :]],
orders=[order_model.spectral_order],
)
order_model.kernel_native = engine.kernels[0]
# Rebuild on pixel grid
f_binned = engine.rebuild(flux_order, fill_value=np.nan)
pixel_grid = np.squeeze(pixel_grid)
f_binned = np.squeeze(f_binned)
# Remove Nans to save space
is_valid = np.isfinite(f_binned)
table_size = np.sum(is_valid)
out_table = np.zeros(table_size, dtype=datamodels.SpecModel().get_dtype("spec_table"))
out_table["WAVELENGTH"] = pixel_grid[is_valid]
out_table["FLUX"] = f_binned[is_valid]
spec = datamodels.SpecModel(spec_table=out_table)
spec.spectral_order = order_model.spectral_order
return tracemodel_ord, spec
def _build_null_spec_table(wave_grid, order, cut=None):
"""
Build a SpecModel of entirely bad values.
Parameters
----------
wave_grid : ndarray
Input wavelengths.
order : int
Spectral order.
cut : float or None
The shortest wavelength to consider for this order, in microns.
Returns
-------
spec : `~stdatamodels.jwst.datamodels.SpecModel`
Null data model. Flux values are NaN, DQ flags are 1,
but note that DQ gets overwritten at end of ``run_extract1d``.
"""
if cut is None:
wave_grid_cut = wave_grid
else:
wave_grid_cut = wave_grid[wave_grid > cut]
spec = datamodels.SpecModel()
spec.spectral_order = order
spec.meta.soss_extract1d.type = "OBSERVATION"
spec.meta.soss_extract1d.factor = np.nan
spec.spec_table = np.zeros(
(wave_grid_cut.size,), dtype=datamodels.SpecModel().get_dtype("spec_table")
)
spec.spec_table["WAVELENGTH"] = wave_grid_cut
spec.spec_table["FLUX"] = np.empty(wave_grid_cut.size) * np.nan
spec.spec_table["DQ"] = np.ones(wave_grid_cut.size)
spec.validate()
return spec
def _do_tiktests(
engine,
scidata_bkg,
scierr,
guess_factor,
order_models,
global_mask,
save_tiktests=False,
tikfac_log_range=None,
niter_refine=3,
):
"""
Test a grid of Tikhonov regularization factors to find the most appropriate one.
Parameters
----------
engine : `~jwst.extract_1d.soss_extract.atoca.ExtractionEngine`
The extraction engine to use for the tests.
scidata_bkg : ndarray
The science data with background subtracted.
scierr : ndarray
The error in the science data.
guess_factor : float
Initial guess for the Tikhonov factor.
order_models : list of `DetectorModelOrder`
Models of the detector and trace properties, one per spectral order.
global_mask : ndarray
Mask determining the aperture used for rebuilding the trace. This typically includes
only pixels that do not belong to either spectral trace, i.e., regions of the detector
where no real data could exist.
save_tiktests : bool, optional
If `True`, re-construct spectra for all tested Tikhonov factors and for all spectral orders
to provide a diagnostic product.
tikfac_log_range : list, optional
Logarithmic range around ``guess_factor`` to search. The default is [-2, 8], which was
chosen because we are looking for the smoothest (largest-factor) solution that
still provides a good fit to the data.
niter_refine : int
Number of times to add two points to the initial grid of trial Tikhonov
factors in order to better estimate the optimal factor.
Default 3.
Returns
-------
tikfac : float
The best-fitting Tikhonov factor.
spec_list : list of `~stdatamodels.jwst.datamodels.SpecModel`
If ``save_tiktests`` is `True`, a list of `~stdatamodels.jwst.datamodels.SpecModel`
for all tested Tikhonov factors
and all spectral orders; otherwise, an empty list.
"""
# reset the kernel so it can get a new shape here
for model in order_models:
model.kernel_native = model.kernel_func
if tikfac_log_range is None:
tikfac_log_range = [-2, 8]
# Find the tikhonov factor.
# Initial pass 8 orders of magnitude with 10 grid points.
log_guess = np.log10(guess_factor)
factors = np.logspace(log_guess + tikfac_log_range[0], log_guess + tikfac_log_range[1], 10)
tikho_struct = engine.get_tikho_test_structure(scidata_bkg, scierr)
all_tests = engine.get_tikho_tests(tikho_struct, factors)
tikfac = engine.best_tikho_factor(all_tests, fit_mode="d_chi2")
log.info("Coarse grid best tikfac: %.4e", tikfac)
# Refine to a final answer.
tikfac, all_tests = _refine_tikfac(all_tests, tikho_struct, engine, niter_refine=niter_refine)
log.info("Final best tikfac: %.4e", tikfac)
spec_list = []
if save_tiktests:
# Save spectra in a list of SingleSpecModels for optional output
for i, order in enumerate(engine.orders):
order_model = order_models[i]
for idx, fac in enumerate(all_tests["factors"]):
log.debug("Building diagnostic spectrum for order %d, factor %.4e", order, fac)
f_k = all_tests["solution"][idx, :]
if np.all(~np.isfinite(f_k)):
spec_ord = _build_null_spec_table(
engine.wave_grid, order, cut=order_model.short_cutoff
)
else:
_, spec_ord = _build_tracemodel_order(
engine, order_model, f_k, global_mask, force_recompute_engine=True
)
_populate_tikho_attr(spec_ord, all_tests, idx, order)
spec_list.append(spec_ord)
# reset the kernel, as it is set to an array inside _build_tracemodel_order
for model in order_models:
model.kernel_native = model.kernel_func
return tikfac, spec_list
def _compute_box_weights(order_models, shape, width, orders_requested):
"""
Determine the weights for the box extraction.
Parameters
----------
order_models : list[DetectorModelOrder]
Models of the detector and trace properties, one per spectral order.
shape : tuple
The shape of the detector image.
width : int
The width of the box aperture.
orders_requested : list[int]
List of orders to be extracted.
Returns
-------
box_weights : dict
A dictionary of the weights for each order.
wavelengths : dict
A dictionary of the wavelengths for each order.
"""
# Extract each order from order list
box_weights, wavelengths = {}, {}
order_str = {order: f"Order {order}" for order in orders_requested}
for order_integer in orders_requested:
# Order string-name is used more often than integer-name
order = order_str[order_integer]
order_idx = order_integer - 1
log.debug(f"Compute box weights for {order}.")
# Define the box aperture
xtrace, ytrace, wavelengths[order] = order_models[order_idx].trace
box_weights[order] = get_box_weights(ytrace, width, shape, cols=xtrace)
return box_weights, wavelengths
def _model_single_order(
data_order,
err_order,
order_model,
mask_fit,
mask_rebuild,
wave_grid,
valid_cols,
tikfac,
do_tiktests=False,
save_tiktests=False,
tikfac_log_range=None,
):
"""
Extract an output spectrum for a single spectral order using the ATOCA algorithm.
The Tikhonov factor is derived in two stages: first, ten factors are tested
spanning ``tikfac_log_range``, and then a further 10 factors are tested across
1 order of magnitude in each direction around the best factor from the first stage.
The best-fitting model and spectrum are reconstructed using the best-fit Tikhonov factor
and respecting ``mask_rebuild``.
Parameters
----------
data_order : ndarray
The 2D data array for the spectral order to be extracted.
err_order : ndarray
The 2D error array for the spectral order to be extracted.
order_model : `DetectorModelOrder`
The model for the spectral order to be extracted.
mask_fit : ndarray
Mask determining the aperture used for extraction. This typically includes
detector bad pixels and any pixels that are not part of the trace.
mask_rebuild : ndarray
Mask determining the aperture used for rebuilding the trace. This typically includes
only pixels that do not belong to either spectral trace, i.e., regions of the detector
where no real data could exist.
wave_grid : ndarray
The wavelength grid used to model the data.
tikfac : float
The Tikhonov factor to use. If ```do_tiktests`` is `True`, the best factor will be
determined by testing a range of factors around this value; otherwise, it will be taken
to be the best value.
do_tiktests : bool, optional
If `True`, test a range of Tikhonov factors to determine the best value.
Default is `False`.
save_tiktests : bool, optional
If `True`, save the intermediate models and spectra for each Tikhonov factor tested.
Has no effect if ``do_tiktests`` is `False`. Default is `False`.
tikfac_log_range : list, optional
The range in log10 space around the initial guess to test Tikhonov factors.
The default is [-2, 8], which was
chosen because we are looking for the smoothest (largest-factor) solution that
still provides a good fit to the data.
Returns
-------
model : ndarray
Model derived from the best Tikhonov factor, same shape as ``data_order``.
spec_list : list of `~stdatamodels.jwst.datamodels.SpecModel`
If ``save_tiktests`` is `True`, returns a list of the model spectra
for each Tikhonov factor tested,
with the best-fitting spectrum last in the list.
If ``save_tiktests`` is `False`, returns a one-element list with the best-fitting spectrum.
Notes
-----
The last spectrum in the list of `~stdatamodels.jwst.datamodels.SpecModel` lacks
the "chi2", "chi2_soft_l1", "chi2_cauchy", and "reg" attributes,
as these are only calculated for the intermediate models. The last spectrum is not
necessarily identical to any of the spectra in the list, as it is reconstructed according to
``mask_rebuild`` instead of fit respecting ``mask_fit``; that is, bad pixels are included.
"""
order = order_model.spectral_order
# The throughput and kernel is not needed here
# set them so they have no effect on the extraction.
def throughput(wavelength):
return np.ones_like(wavelength)
# Define wavelength grid with oversampling of 3 (should be enough)
wave_grid_os = oversample_grid(wave_grid, n_os=3)
# Initialize the Engine.
engine = ExtractionEngine(
[order_model.wavemap],
[order_model.specprofile],
[throughput],
[np.array([1.0])],
wave_grid=wave_grid_os,
mask_trace_profile=[mask_fit],
orders=[order],
)
# Find the tikhonov factor.
if do_tiktests:
tikfac, spec_list = _do_tiktests(
engine,
data_order,
err_order,
tikfac,
[order_model],
mask_rebuild,
save_tiktests=save_tiktests,
tikfac_log_range=tikfac_log_range,
)
else:
spec_list = []
# Use the best Tikhonov factor to build the final model and spectrum
f_k_final = engine(data_order, err_order, tikhonov=True, factor=tikfac)
# Rebuild trace, including bad pixels
engine = ExtractionEngine(
[order_model.wavemap],
[order_model.specprofile],
[throughput],
[np.array([1.0])],
wave_grid=wave_grid_os,
mask_trace_profile=[mask_rebuild],
orders=[order],
)
model = engine.rebuild(f_k_final, fill_value=np.nan)
# Build 1d spectrum integrated over pixels
mask = np.all(mask_rebuild, axis=0)[valid_cols]
wave_grid = wave_grid[np.newaxis, :]
engine = ExtractionEngine(
[wave_grid],
[np.ones_like(wave_grid)],
[throughput],
[np.array([1.0])],
wave_grid=wave_grid_os,
mask_trace_profile=[mask[np.newaxis, :]],
orders=[order_model.spectral_order],
)
f_binned = engine.rebuild(f_k_final, fill_value=np.nan)
wave_grid = np.squeeze(wave_grid)
f_binned = np.squeeze(f_binned)
# Remove Nans to save space
is_valid = np.isfinite(f_binned)
table_size = np.sum(is_valid)
out_table = np.zeros(table_size, dtype=datamodels.SpecModel().get_dtype("spec_table"))
out_table["WAVELENGTH"] = wave_grid[is_valid]
out_table["FLUX"] = f_binned[is_valid]
spec = datamodels.SpecModel(spec_table=out_table)
spec.spectral_order = order_model.spectral_order
spec.meta.soss_extract1d.factor = tikfac
spec.meta.soss_extract1d.type = "OBSERVATION"
# Add the result to spec_list
spec_list.append(spec)
return model, spec_list
[docs]
class Integration:
"""
Class to handle extraction of a single integration.
Parameters
----------
scidata : ndarray
2D science data array.
scierr : ndarray
2D science error array.
scimask : ndarray
2D boolean mask array for the science data. `True` values are masked.
refmask : ndarray
2D boolean mask array for the reference pixels. `True` values are masked.
order_models : list of `DetectorModelOrder`
Models of the detector and trace properties, one per spectral order.
box_weights : dict
A dictionary of the weights for each order.
do_bkgsub : bool, optional
Whether to perform background subtraction. Default is `True`.
extract_order3 : bool, optional
Whether to extract order 3. Default is `True`.
save_intermediate : bool, optional
Whether to save intermediate products for debugging. Default is `False`.
order2_separation_cutoff : tuple(float, float), optional
Wavelengths short of which to consider Order 2 to be well-separated from order 1.
Two values are defined. The first is the cutoff longwave of which should be included
in the combined extraction; the second is the cutoff shortwave of which should be
included in the extraction of just that order. Some overlap is helpful to avoid
numerical issues.
"""
def __init__(
self,
scidata,
scierr,
scimask,
refmask,
order_models,
box_weights,
do_bkgsub=True,
extract_order3=True,
save_intermediate=False,
order2_separation_cutoff=(0.77, 0.95),
):
self.scidata = scidata
self.scierr = scierr
self.scimask = scimask
self.refmask = refmask
self.box_weights = box_weights
self.save_intermediate = save_intermediate
if extract_order3:
self.order_list = [1, 2, 3]
else:
self.order_list = [1, 2]
self.order_strs = [f"Order {order}" for order in self.order_list]
self.order_indices = [o - 1 for o in self.order_list]
self._validate_masks()
self._subtract_bkg(do_bkgsub)
# unpack ref file args
self.order_models = order_models
self.subarray = order_models[0].subarray
self.order2_separation_cutoff = order2_separation_cutoff
# Define mask based on box aperture
# (we want to model each contaminated pixels that will be extracted)
self.mask_trace_profile = [
(~(self.box_weights[self.order_strs[i]] > 0)) | (self.refmask)
for i in self.order_indices
]
def _subtract_bkg(self, do_bkgsub):
# Perform background correction if requested
if do_bkgsub:
log.debug("Applying background subtraction.")
bkg_mask = make_background_mask(self.scidata, width=40)
self.scidata_bkg, self.col_bkg = soss_background(self.scidata, self.scimask, bkg_mask)
else:
log.debug("Skip background subtraction.")
self.scidata_bkg = self.scidata.copy()
self.col_bkg = np.zeros(self.scidata.shape[1])
def _validate_masks(self):
# Make sure there aren't any nans not flagged in scimask
not_finite = ~(np.isfinite(self.scidata) & np.isfinite(self.scierr))
if (not_finite & ~self.scimask).any():
log.warning(
"Input contains invalid values that "
"are not flagged correctly in the dq map. "
"Masking those values for extraction."
)
self.scimask |= not_finite
self.refmask &= ~not_finite
# Some error values are 0, we need to mask those pixels for the extraction engine.
self.scimask = self.scimask | ~(self.scierr > 0)
[docs]
def model_image(
self,
wave_grid,
estimate=None,
tikfacs_in=None,
threshold=1e-2,
):
"""
Extract the uncontaminated spectra of using the ATOCA algorithm.
The model consists of three separate extractions:
1. A simultaneous extraction of orders 1 and 2.
2. An extraction of the blue (uncontaminated) part of order 2.
3. An extraction of order 3, if requested.
For each of these extractions, the Tikhonov factor can either be provided
using the ``tikfacs_in`` dictionary, or else determined via a grid search.
Parameters
----------
wave_grid : ndarray or None
The wavelength grid to use for the simultaneous extraction of orders 1 and 2.
estimate : func, optional
An estimate of the target flux as a function of wavelength in microns.
This is used to generate the wavelength grid if ``wave_grid`` is None,
and to estimate the Tikhonov factor if not provided in ``tikfacs_in``.
If ``wave_grid`` is given and the Tikhonov factor for order 1 is provided
in ``tikfacs_in``, this is not needed. If either is not provided, this must be given.
tikfacs_in : dict, optional
A dictionary with keys "Order 1", "Order 2", and "Order 3" and values
giving the Tikhonov factor to use for each order. If None (default), the Tikhonov factor
that minimizes the chi-squared for each order will be determined via a grid search.
threshold : float, optional
The threshold (between 0 and 1) for determining which pixels to include in the
extraction. Pixels with a normalized flux contribution from the spectral trace
above this threshold will be included. Default is 1e-2.
Returns
-------
tracemodels : dict
A dictionary with keys "Order 1", "Order 2", and "Order 3" (if extracted)
and values giving the modeled detector image for each order.
spec_list : list of `~stdatamodels.jwst.datamodels.SpecModel`
A list of the extracted spectra for each order. If the Tikhonov factor
was determined via a grid search, this will include the intermediate spectra
for each factor tested, with the best-fitting spectrum last in the list.
tikfacs_out : dict
A dictionary with keys "Order 1", "Order 2", and "Order 3" (if extracted)
and values giving the Tikhonov factor used for each order.
"""
if tikfacs_in is None:
tikfacs_in = {"Order 1": None, "Order 2": None, "Order 3": None}
tikfacs_out = {"Order 1": None, "Order 2": None, "Order 3": None}
# Check that the inputs are valid
if (estimate is None) and (tikfacs_in["Order 1"] is None):
msg = (
"If the Tikhonov factor for Order 1 is not given, "
"an estimate of the target flux as a function of wavelength must be provided."
)
log.error(msg)
raise ValueError(msg)
# Define mask of pixel to model (all pixels inside box aperture)
global_mask = np.all(self.mask_trace_profile, axis=0).astype(bool)
# Initialize the Engine for combined extraction of orders 1 and 2
if self.order_models[0].order12engine is not None:
engine = self.order_models[0].order12engine
else:
engine = ExtractionEngine(
[om.wavemap for om in self.order_models][:2],
[om.specprofile for om in self.order_models][:2],
[om.throughput for om in self.order_models][:2],
[om.kernel for om in self.order_models][:2],
wave_grid=wave_grid,
mask_trace_profile=self.mask_trace_profile[:2],
global_mask=self.scimask,
threshold=threshold,
orders=[1, 2],
)
self.order_models[0].order12engine = engine
# set the kernels in the order models to those used in the engine
# these are now sparse matrices, so this avoids re-computing them in subsequent integrations
for i in range(2):
self.order_models[i].kernel = engine.kernels[i]
# Find the tikhonov factor for order 1 (and contaminated part of order 2)
if tikfacs_in["Order 1"] is None:
log.info("Solving for the optimal Tikhonov factor for overlapping orders 1 & 2.")
save_tiktests = self.save_intermediate
guess_factor = engine.estimate_tikho_factors(estimate)
tikfac, spec_list = _do_tiktests(
engine,
self.scidata_bkg,
self.scierr,
guess_factor,
self.order_models[:2],
global_mask,
save_tiktests=save_tiktests,
tikfac_log_range=[-4, 4],
)
tikfacs_out["Order 1"] = tikfac
for spec in spec_list:
spec.meta.soss_extract1d.color_range = "RED"
else:
save_tiktests = False
spec_list = []
tikfacs_out["Order 1"] = tikfacs_in["Order 1"]
# Run the extract method of the Engine.
log.debug("Running extraction engine for overlapping orders 1 & 2...")
# Precompute the inverse of the design matrix and the values divided
# by the uncertainties. This makes the solution of the matrix equation
# only a matrix multiplication.
if self.order_models[0].m_inv is None:
log.info("Precomputing the inverse of the design matrix for Order 1+2 modeling")
_m_inv, _bmat = engine.precompute_detector_model(
self.scidata_bkg, self.scierr, tikfac=tikfacs_out["Order 1"]
)
self.order_models[0].m_inv = _m_inv
self.order_models[0].bmat = _bmat
with warnings.catch_warnings():
warnings.filterwarnings("ignore", "invalid value encountered in divide", RuntimeWarning)
y_over_err = (self.scidata_bkg / self.order_models[0].mederr)[~engine.mask]
f_k = self.order_models[0].m_inv.dot(self.order_models[0].bmat.T @ y_over_err)
# Create a new instance of the engine for evaluating the trace model.
# This allows bad pixels and pixels below the threshold to be reconstructed as well.
# Model the traces for each order separately.
log.debug("Building decontaminated trace models and spectra for Order 1 and Order 2 red.")
tracemodels = {}
for i_order, _order in enumerate([1, 2]):
tracemodel_ord, spec_ord = _build_tracemodel_order(
engine, self.order_models[i_order], f_k, global_mask, force_recompute_engine=False
)
spec_ord.meta.soss_extract1d.factor = tikfacs_out["Order 1"]
spec_ord.meta.soss_extract1d.color_range = "RED"
spec_ord.meta.soss_extract1d.type = "OBSERVATION"
# Project on detector and save in dictionary
tracemodels[self.order_strs[i_order]] = tracemodel_ord
# Add the result to spec_list
spec_list.append(spec_ord)
# Make a null tracemodel to be overwritten later for order 3
if 3 in self.order_list:
tracemodels["Order 3"] = np.zeros_like(tracemodels["Order 2"]) * np.nan
# Model the blue part of order 2 and all of order 3 assuming they are well-separated
# from order 1
for order in self.order_list[1:]:
if self.subarray == "SUBSTRIP96":
continue
if order not in self.order_list:
continue
if order == 2:
cutoff = self.order2_separation_cutoff[1]
else:
cutoff = None
idx_order = np.array(self.order_indices)[np.array(self.order_list) == order][0]
order_str = self.order_strs[idx_order]
log.debug(f"Generate model for well-separated part of {order_str}")
order_model = self.order_models[idx_order]
# Use provided tikfac if given=
if tikfacs_in[order_str] is None:
# If not given just try what was done for Order 1
tikfac = tikfacs_out["Order 1"]
do_tiktests = True
else:
tikfac = tikfacs_in[order_str]
do_tiktests = False
# Mask for the fit. All valid pixels inside box aperture
mask_fit = self.mask_trace_profile[idx_order] | self.scimask
# Build 1d spectrum integrated over pixels
pixel_wave_grid, valid_cols = order_model.native_grid
# Hardcode wavelength highest boundary as well.
# Must overlap with lower limit in make_decontamination_grid
if cutoff is not None:
is_in_wv_range = pixel_wave_grid < cutoff
pixel_wave_grid = pixel_wave_grid[is_in_wv_range]
valid_cols = valid_cols[is_in_wv_range]
# Model with atoca
try:
model, spec_ord = _model_single_order(
self.scidata_bkg,
self.scierr,
order_model,
mask_fit,
global_mask,
pixel_wave_grid,
valid_cols,
tikfac,
do_tiktests=do_tiktests,
save_tiktests=save_tiktests,
tikfac_log_range=[-2, 8],
)
tikfacs_out[order_str] = spec_ord[-1].meta.soss_extract1d.factor
except MaskOverlapError:
log.error(
"Not enough unmasked pixels to model the remaining part of order 2."
" Model and spectrum will be NaN in that spectral region."
)
spec_ord = [
_build_null_spec_table(pixel_wave_grid, order, cut=order_model.short_cutoff)
]
model = np.nan * np.ones_like(self.scidata_bkg)
# Keep only pixels from which order 2 contribution
# is not already modeled.
already_modeled = np.isfinite(tracemodels[order_str])
model = np.where(already_modeled, 0.0, model)
# Add to tracemodels
both_nan = np.isnan(tracemodels[order_str]) & np.isnan(model)
tracemodels[order_str] = np.nansum([tracemodels[order_str], model], axis=0)
tracemodels[order_str][both_nan] = np.nan
# Add the result to spec_list
for sp in spec_ord:
if order == 2:
sp.meta.soss_extract1d.color_range = "BLUE"
else:
sp.meta.soss_extract1d.color_range = "ALL"
spec_list += spec_ord
return tracemodels, spec_list, tikfacs_out
[docs]
def estim_flux_first_order(self, threshold=1e-4):
"""
Roughly estimate the underlying flux of the target spectrum.
This is done by simply masking out order 2 and retrieving the flux from order 1.
Parameters
----------
threshold : float, optional
The pixels with an ``aperture[order 2] > threshold`` are considered contaminated
and will be masked.
Returns
-------
func
A spline estimator that provides the underlying flux as a function of wavelength
"""
# Define wavelength grid based on order 1 only (so first index)
model_order1 = self.order_models[0]
wave_grid = grid_from_map_with_extrapolation(
model_order1.wavemap, model_order1.specprofile, n_os=1
)
# Mask parts contaminated by order 2 based on its spatial profile
mask = (self.order_models[1].specprofile >= threshold) | self.mask_trace_profile[0]
# Init extraction without convolution kernel (so extract the spectrum at order 1 resolution)
engine = ExtractionEngine(
[model_order1.wavemap],
[model_order1.specprofile],
[model_order1.throughput],
[None],
wave_grid,
[mask],
global_mask=self.scimask,
orders=[1],
)
# Extract estimate
spec_estimate = engine(self.scidata_bkg, self.scierr)
# Interpolate
idx = np.isfinite(spec_estimate)
return UnivariateSpline(wave_grid[idx], spec_estimate[idx], k=3, s=0, ext=0)
[docs]
def make_decontamination_grid(self, estimate, rtol=1e-4, max_grid_size=20000, n_os=2):
"""
Create the grid to use for the simultaneous extraction of order 1 and 2.
The grid is made by:
1. requiring that it satisfies the oversampling ``n_os``,
2. trying to reach the specified tolerance for the spectral range
shared between order 1 and 2, and
3. trying to reach the specified tolerance in the rest of spectral range.
The ``max_grid_size`` overrules steps 2 and 3, so the precision may not be reached if
the grid size needed is too large.
Parameters
----------
estimate : `~scipy.interpolate.UnivariateSpline`
Estimate of the target flux as a function of wavelength in microns.
rtol : float, optional
The relative tolerance needed on a pixel model.
max_grid_size : int, optional
Maximum grid size allowed.
n_os : int, optional
The oversampling factor of the wavelength grid used when solving for
the uncontaminated flux.
Returns
-------
wave_grid : ndarray
The 1D grid of the pixels boundaries at the native sampling.
"""
# Build native grid for each order
spectral_orders = [2, 1]
grids_ord = {}
for sp_ord in spectral_orders:
model = self.order_models[sp_ord - 1]
grids_ord[sp_ord] = model.get_grid_from_trace(n_os=n_os)
# Build the list of grids given to make_combined_grid.
# It must be ordered in increasing priority.
# 1rst priority: shared wavelengths with order 1 and 2.
# 2nd priority: remaining red part of order 1
# 3rd priority: remaining blue part of order 2
# So, split order 2 in 2 parts, the shared wavelength and the bluemost part
is_shared = grids_ord[2] >= np.min(grids_ord[1])
# Make sure order 1 is not more in the blue than order 2
cond = grids_ord[1] > np.min(grids_ord[2][is_shared])
grids_ord[1] = grids_ord[1][cond]
# And make grid list
all_grids = [grids_ord[2][is_shared], grids_ord[1], grids_ord[2][~is_shared]]
# Cut order 2 at 0.77 (not smaller than that)
# because there is no contamination there. Can be extracted afterward.
# In the red, no cut.
wv_range = [self.order2_separation_cutoff[0], np.max(grids_ord[1])]
# Finally, build the list of corresponding estimates.
# The estimate for the overlapping part is the order 1 estimate.
# There is no estimate yet for the blue part of order 2, so give a flat spectrum.
def flat_fct(wv):
return np.ones_like(wv)
all_estimates = [estimate, estimate, flat_fct]
# Generate the combined grid
kwargs = {"rtol": rtol, "max_total_size": max_grid_size, "max_iter": 30}
return make_combined_adaptive_grid(all_grids, all_estimates, wv_range, **kwargs)
[docs]
def decontaminate_image(self, tracemodels):
"""
Perform decontamination of the image based on the trace models.
Parameters
----------
tracemodels : dict
Dictionary of the modeled detector images for each order.
Returns
-------
decontaminated_data : dict
Dictionary of the decontaminated data for each order.
"""
log.debug("Performing the decontamination.")
# List of modeled orders
mod_order_list = tracemodels.keys()
# Extract each order from order list
decontaminated_data = {}
for order in self.order_strs:
# Decontaminate using all other modeled orders
decont = self.scidata_bkg.copy()
for mod_order in mod_order_list:
if mod_order != order:
log.debug(f"Decontaminating {order} from {mod_order} using model.")
is_valid = np.isfinite(tracemodels[mod_order])
decont = decont - np.where(is_valid, tracemodels[mod_order], 0.0)
# Save results
decontaminated_data[order] = decont
return decontaminated_data
def _process_one_integration(
scidata,
scierr,
scimask,
refmask,
order_models,
box_weights,
wavelengths,
soss_kwargs,
wave_grid=None,
tikfacs_in=None,
generate_model=True,
int_num=None,
order2_separation_cutoff=(0.77, 0.95),
):
if type(int_num) is int:
log.info(f"Processing integration {int_num}")
# Print verbose log info only on the first integration.
verbose = int_num == 1
if tikfacs_in is None:
tikfacs_in = {"Order 1": None, "Order 2": None, "Order 3": None}
integration = Integration(
scidata,
scierr,
scimask,
refmask,
order_models,
box_weights,
extract_order3=soss_kwargs["order_3"],
do_bkgsub=soss_kwargs["subtract_background"],
save_intermediate=soss_kwargs["model"],
order2_separation_cutoff=order2_separation_cutoff,
)
# Model the traces based on optics filter configuration (CLEAR or F277W)
if generate_model:
if (tikfacs_in["Order 1"] is None or wave_grid is None) and soss_kwargs["estimate"] is None:
# If tikfac and wave_grid already given, no need to make an estimate
# So this if statement saves a bit of runtime
log.info("Estimating the target flux based on order 1 low-contamination pixels.")
estimate = integration.estim_flux_first_order()
else:
estimate = soss_kwargs["estimate"]
# Generate grid based on estimate if not given
if wave_grid is None:
log.info(f"wave_grid not given: generating grid based on rtol={soss_kwargs['rtol']}")
wave_grid = integration.make_decontamination_grid(
estimate,
rtol=soss_kwargs["rtol"],
max_grid_size=soss_kwargs["max_grid_size"],
n_os=soss_kwargs["n_os"],
)
log.info(
f"wave_grid covering from {wave_grid.min()} to {wave_grid.max()}"
f" with {wave_grid.size} points"
)
else:
if verbose:
log.info("Using previously computed or user specified wavelength grid.")
# Model the image.
try:
tracemodels, atoca_list, tikfacs_out = integration.model_image(
wave_grid,
estimate=estimate,
tikfacs_in=tikfacs_in,
threshold=soss_kwargs["threshold"],
)
except KernelShapeError:
# Revert to using kernel functions. This will happen if one integration has
# bad pixels covering an entire section of the detector.
for order_model in integration.order_models:
order_model.kernel = order_model.kernel_func
order_model.kernel_native = order_model.kernel_func
tracemodels, atoca_list, tikfacs_out = integration.model_image(
wave_grid,
estimate=estimate,
tikfacs_in=tikfacs_in,
threshold=soss_kwargs["threshold"],
)
# Convert atoca_list SpecModels to raw data for pickling
atoca_list_data = []
for spec in atoca_list:
spec_data = {
"wavelength": spec.spec_table["WAVELENGTH"].copy(),
"flux": spec.spec_table["FLUX"].copy(),
"spectral_order": spec.spectral_order,
"int_num": int_num if spec.int_num is None else spec.int_num,
}
atoca_list_data.append(spec_data)
else:
# Return empty tracemodels
tracemodels = {}
atoca_list_data = []
tikfacs_out = tikfacs_in
# Decontaminate the data using trace models (if tracemodels not empty)
data_to_extract = integration.decontaminate_image(tracemodels)
# Use the bad pixel models to perform a de-contaminated extraction.
fluxes, fluxerrs, npixels = integration.extract_image(
data_to_extract, bad_pix=soss_kwargs["bad_pix"], tracemodels=tracemodels, verbose=verbose
)
# Save trace models for output reference
for order in tracemodels:
# Put NaNs to zero
model_ord = tracemodels[order]
model_ord = np.where(np.isfinite(model_ord), model_ord, 0.0)
tracemodels[order] = model_ord
# Convert spec_list to raw data for pickling
spec_list_data = {}
for order in fluxes.keys():
table_size = len(wavelengths[order])
spec_list_data[order] = {
"wavelength": wavelengths[order][:table_size],
"flux": fluxes[order][:table_size],
"flux_error": fluxerrs[order][:table_size],
"background": integration.col_bkg[:table_size],
"npixels": npixels[order][:table_size],
"spectral_order": ORDER_STR_TO_INT[order],
}
return tracemodels, spec_list_data, atoca_list_data, tikfacs_out, wave_grid
def _reconstruct_spec_from_data(spec_data):
"""
Construct a SpecModel from raw data dictionary.
This allows passing SpecModel-like objects through multiprocessing.
Parameters
----------
spec_data : dict
Dictionary containing wavelength, flux, and metadata
Returns
-------
spec : `~stdatamodels.jwst.datamodels.SpecModel`
Reconstructed data model
"""
table_size = len(spec_data["wavelength"])
out_table = np.zeros(table_size, dtype=datamodels.SpecModel().get_dtype("spec_table"))
out_table["WAVELENGTH"] = spec_data["wavelength"]
out_table["FLUX"] = spec_data["flux"]
if "flux_error" in spec_data:
# Intermediate product candidate spectra coming from do_tiktests don't have these
out_table["FLUX_ERROR"] = spec_data["flux_error"]
out_table["BACKGROUND"] = spec_data["background"]
out_table["NPIXELS"] = spec_data["npixels"]
spec = datamodels.SpecModel(spec_table=out_table)
spec.spectral_order = spec_data["spectral_order"]
# Int_num is used only for ATOCA spectra. It will be None for regular output spectra.
spec.int_num = spec_data.get("int_num")
# Set units for the spectral table. Assume uncalibrated flux, wavelength in um.
spec.spec_table.columns["wavelength"].unit = "um"
spec.spec_table.columns["flux"].unit = "DN/s"
spec.spec_table.columns["flux_error"].unit = "DN/s"
spec.spec_table.columns["background"].unit = "DN/s"
return spec