# The functions to perform mathematical operations
import numpy as np
from astropy.stats import biweight_location
from astropy.table import Table
from astropy.io.fits.fitsrec import FITS_rec
from .logger import logger
'''
Mathematical operations
'''
[docs]
def ari_operations(arr1, arr2, var_arr1=None, var_arr2=None, operation='+'):
"""
Perform element-wise arithmetic operations on two input arrays with
optional variance propagation.
Parameters
----------
arr1 : numpy.ndarray
First input array.
arr2 : numpy.ndarray
Second input array, must be broadcastable to the shape of `arr1`.
var_arr1 : numpy.ndarray or None, optional
Variance (uncertainty) array corresponding to `arr1`. Default is None.
var_arr2 : numpy.ndarray or None, optional
Variance (uncertainty) array corresponding to `arr2`. Default is None.
operation : str, optional
Arithmetic operation to apply (default is 'sum').
Supported values:
- '+' : element-wise addition
- '-' : element-wise subtraction (`arr1 - arr2`)
- '*' : element-wise multiplication
- '/' : element-wise division (`arr1 / arr2`)
Returns
-------
numpy.ndarray or tuple of numpy.ndarray
If `var_arr1` and `var_arr2` are not provided, returns the result of
element-wise operation on `arr1` and `arr2`.
If both variances are provided, returns the propagated variance array
computed according to the operation:
- For '+' and '-': variances are added.
- For '*' and '/': variance propagated as
product.
Raises
------
ZeroDivisionError
When division by zero occurs in 'div' operation.
ValueError
If `operation` is not one of the supported strings.
Examples
--------
>>> import numpy as np
>>> a = np.array([1.0, 2.0, 3.0])
>>> b = np.array([4.0, 5.0, 6.0])
>>> ari_operations(a, b, operation='+')
array([5., 7., 9.])
>>> var_a = np.array([0.1, 0.1, 0.1])
>>> var_b = np.array([0.2, 0.2, 0.2])
>>> ari_operations(a, b, var_a, var_b, operation='+')
array([0.3, 0.3, 0.3])
"""
if operation == '+':
answer = arr1 + arr2
elif operation == '-':
answer = arr1 - arr2
elif operation == '*':
answer = arr1 * arr2
elif operation == '/':
answer = arr1 / arr2
else:
raise ValueError(
f"Unsupported operation '{operation}'. Supported: ",
"'+', '-', '*', '/'.")
if (var_arr1 is not None) & (var_arr2 is not None):
if (operation == '+') or (operation == '-'):
var_tot = var_arr1 + var_arr2
elif (operation == '*') or (operation == '/'):
var_tot = answer**2 * ((var_arr1/arr1**2) + (var_arr2/arr2**2))
return answer, var_tot
return answer, None
'''
Combine
'''
[docs]
def combine_data(dataarr, var=None, method='mean',
mask=None):
"""
Combine multiple arrays along the first axis using a specified method.
Parameters
----------
dataarr : array_like
Input data array of shape (N, ...), where `N` is the number of
individual datasets to combine. The combination is performed
along axis=0.
var : array_like, optional
Variance array of the same shape as `dataarr`. If provided,
error propagation is performed assuming independent errors,
yielding the variance of the combined data. Default is None.
method : {'mean', 'median', 'biweight', 'weightedavg'}, optional
Method used for combining the data:
- 'mean' : arithmetic mean ignoring NaNs.
- 'median' : median ignoring NaNs.
- 'biweight' : robust biweight location (from `astropy.stats`).
Default is 'mean'.
Returns
-------
comb_data : ndarray
Combined data array, same shape as a single input array
(i.e., shape of `dataarr[0]`).
comb_var : ndarray, optional
Combined variance array of the same shape as `comb_data`.
Returned only if `var` is provided.
Notes
-----
- NaN values in `dataarr` are ignored during combination.
- Variance is propagated as if the combination method were the mean,
even if `median` or `biweight` are chosen. This provides an
approximate uncertainty estimate.
- The biweight method is less sensitive to outliers than the mean
or median.
"""
# Handle binary tables
if len(dataarr) > 0 and isinstance(dataarr[0], (FITS_rec, Table)):
return combine_bintable(dataarr, method=method)
if method == 'weightedavg':
comb_data, comb_var = weighted_mean_and_variance(dataarr, var)
return comb_data, comb_var
dataarr = np.array(dataarr)
N = dataarr.shape[0]
if mask is not None:
mask_full = np.broadcast_to(mask, dataarr.shape)
dataarr_ma = np.ma.array(dataarr, mask=mask_full)
if method == 'mean':
comb_data = np.ma.nanmean(dataarr_ma, axis=0).filled(np.nan)
elif method == 'median':
comb_data = np.ma.nanmedian(dataarr_ma, axis=0).filled(np.nan)
elif method == 'biweight':
comb_data = biweight_location(dataarr_ma, axis=0).filled(np.nan)
else:
if method == 'mean':
comb_data = np.nanmean(dataarr, axis=0)
elif method == 'median':
comb_data = np.nanmedian(dataarr, axis=0)
elif method == 'biweight':
comb_data = biweight_location(dataarr, axis=0)
# Propagating error.
# Treating the error propagation
# as mean for median also.
if var is not None:
if mask is None:
comb_var = np.sum(var, axis=0) / N**2
return comb_data, comb_var
else:
var_ma = np.ma.array(var, mask_full)
comb_var = np.ma.sum(var_ma, axis=0) / N**2
comb_var = comb_var.filled(np.nan)
return comb_data, comb_var
return comb_data, None
[docs]
def weighted_mean_and_variance(values, variances):
r"""
Compute the weighted mean and variance of the mean,
given measurements and their variances.
Parameters
----------
values : array-like
Measured values (x_i)
variances : array-like
Variances of the measurements.
Returns
----------
mean : float
Weighted mean.
variance_of_mean : float
Variance of the weighted mean
Raises
------
ValueError
If `variances` is None.
TypeError
If `values` or `variances` are not array-like.
Notes
-----
The weighted mean is computed as:
.. math::
\bar{x} = \frac{\sum_i w_i x_i}{\sum_i w_i}, \quad
w_i = \frac{1}{\sigma_i^2}
The variance of the weighted mean is:
.. math::
\sigma_{\bar{x}}^2 = \frac{1}{\sum_i w_i}
"""
if variances is None:
raise TypeError("variances must be an array-like object")
weights = 1.0 / variances
mean = np.sum(weights * values, axis=0) / np.sum(weights, axis=0)
variance_of_mean = 1.0 / np.sum(weights, axis=0)
return mean, variance_of_mean
[docs]
def combine_bintable(dataarr,
value_cols=None,
uncertainity_cols=None,
combine_cols=None,
method='mean'):
if value_cols is None:
value_cols = []
if uncertainity_cols is None:
uncertainity_cols = []
if combine_cols is None:
combine_cols = []
if len(value_cols) != len(uncertainity_cols):
raise ValueError(
"'value_cols' and 'uncertainity_cols' must have the same length"
)
tables = [Table(t) if not isinstance(t, Table) else t
for t in dataarr]
ref = tables[0]
comb = ref.copy(copy_data=True)
handled = set()
# Combine quantities with propagated uncertainities
for value_col, err_col in zip(value_cols, uncertainity_cols):
values = np.stack([t[value_col] for t in tables])
var = np.stack([t[err_col]**2 for t in tables])
comb[value_col], comb_var = combine_data(
values,
var=var,
method=method
)
comb[err_col] = np.sqrt(comb_var)
handled.update([value_col, err_col])
# Combine quantities without uncertainities
for col in combine_cols:
values = np.stack([t[col] for t in tables])
comb[col], _ = combine_data(
values,
method=method
)
handled.add(col)
# Verify all remaining columns are identical
for col in ref.colnames:
if col in handled:
continue
for table in tables[1:]:
if not np.array_equal(ref[col], table[col]):
raise ValueError(
f"Column '{col}' differs between tables."
)
return comb, None
[docs]
def combine_data_full(datadict, dataext=[1, 2, 3],
varext=[4, 5, 6],
extras=[],
table_info=None,
method='mean'):
"""
Combine data from multiple FITS files into a single dictionary.
This function combines data arrays, variance arrays, additional
numeric arrays, and binary tables stored in a dictionary (typically
produced by reading multiple FITS files). Flux and variance arrays
are combined using ``combine_data``, while binary tables are combined
using ``combine_bintable``.
Parameters
----------
datadict : dict
Dictionary containing data from multiple FITS files. Each key
corresponds to a FITS extension or metadata item. Arrays to be
combined are expected to be stacked along the first axis
(i.e., shape ``(n_files, ...)``).
dataext : list of int, optional
Indices of ``datadict.keys()`` corresponding to data arrays
(e.g., flux) that should be combined. Default is
``[1, 2, 3]``.
varext : list of int, optional
Indices of ``datadict.keys()`` corresponding to variance arrays.
Each entry must correspond to the matching entry in
``dataext``. Default is ``[4, 5, 6]``.
extras : list of int, optional
Indices of additional numeric arrays that should be combined
using ``combine_data`` without associated variance arrays.
Default is ``[]``.
table_info : dict, optional
Dictionary describing binary table extensions to combine. Keys
are indices into ``datadict.keys()`` and values are dictionaries
passed to ``combine_bintable``.
Each value may contain the following entries:
- ``value_cols`` : list of columns whose values are combined
with propagated uncertainties.
- ``uncertainty_cols`` : list of uncertainty columns
corresponding to ``value_cols``.
- ``combine_cols`` : list of numeric columns that are combined
without uncertainty propagation.
Columns not listed above are assumed to be identical in all
input tables and are copied from the first table after verifying
they are unchanged.
Default is ``None``.
method : {'mean', 'median', 'biweight'}, optional
Method used to combine the data.
- ``'mean'`` : arithmetic mean.
- ``'median'`` : median.
- ``'biweight'`` : biweight location.
Returns
-------
comb_dicts : dict
Dictionary containing the combined data.
- Data arrays in ``dataext`` are combined.
- Variance arrays in ``varext`` are propagated.
- Arrays in ``extras`` are combined.
- Binary tables in ``table_info`` are combined using
``combine_bintable``.
- All remaining entries are copied from the first input file.
Notes
-----
- The order of ``dataext`` and ``varext`` must correspond.
- Dictionary insertion order is assumed to match the FITS extension
order.
- Entries not listed in ``dataext``, ``varext``, ``extras``, or
``table_info`` are copied from the first input file.
- ``table_info`` provides a generic mechanism for combining
arbitrary FITS binary tables without requiring instrument-specific
code.
Examples
--------
>>> table_info = {
... 7: {
... "value_cols": ["VALUE"],
... "uncertainty_cols": ["UNCERTAINTY"],
... "combine_cols": []
... }
... }
>>> combined = combine_data_full(
... datadict,
... dataext=[1],
... varext=[2],
... table_info=table_info,
... method="mean",
... )
"""
dictkeys = list(datadict.keys())
comb_dicts = datadict.copy()
flux_keys = [dictkeys[int(i)] for i in dataext]
var_keys = [dictkeys[int(i)] for i in varext]
if len(extras) > 0:
extra_keys = [dictkeys[int(i)] for i in extras]
else:
extra_keys = []
table_keys = [dictkeys[i] for i in table_info] if table_info else []
# Avoiding the extensions that are not flux or variance.
# Taking only the first element of that. i.e.,
# The data from first fits file will
# be copied to the final output.
# For spectrum, wavelengths are interpolated to
# same array. So, that also
# copied in the same way.
for cro, keys in enumerate(dictkeys):
if keys not in flux_keys + var_keys + extra_keys + table_keys:
comb_dicts[keys] = comb_dicts[keys][0]
# Doing for flux and variance.
for index, extk in enumerate(flux_keys):
logger.info("Operation on %s", flux_keys[index])
logger.info("With variance in %s", var_keys[index])
fluxes = comb_dicts[flux_keys[index]]
variances = comb_dicts[var_keys[index]]
comb_flux, comb_var = combine_data(fluxes, variances,
method=method)
comb_dicts[flux_keys[index]] = comb_flux
comb_dicts[var_keys[index]] = comb_var
for index, ext in enumerate(extra_keys):
data = comb_dicts[ext]
logger.info("Operation on %s", ext)
comb_data, _ = combine_data(data,
method=method)
comb_dicts[ext] = comb_data
if table_info is not None:
for extnum, info in table_info.items():
key = dictkeys[extnum]
logger.info("Table operation on %s", key)
comb_table, _ = combine_bintable(
comb_dicts[key],
value_cols=info.get("value_cols", []),
uncertainity_cols=info.get("uncertainty_cols", []),
combine_cols=info.get("combine_cols", []),
method=method,
)
comb_dicts[key] = comb_table
return comb_dicts
# End