from astropy.io import fits
from pathlib import Path
import numpy as np
from .logger import logger
[docs]
def call_mask(mask):
if mask is None:
return None
if isinstance(mask, np.ndarray):
return mask
mask = str(mask)
logger.info(f"Calling mask: {mask}")
if mask.endswith(".fits"):
mask = fits.getdata(mask).astype(bool)
elif mask.endswith(".npy"):
mask = np.load(mask).astype(bool)
else:
logger.error(
"Invalid scale mask: %s. "
"Expected as NumPy array, FITS file, or NPY file.",
mask
)
raise ValueError(
"mask must be a NumPy array, FITS file or NPY file."
)
return mask
[docs]
def shrink_fits(filename, extensions, replace=False,
strict=True):
"""
Shrink a FITS file by retaining data only in selected extensions.
Extensions not listed in ``extensions`` are replaced with empty HDUs,
preserving their headers and extension names.
Parameters
----------
filename : str or pathlib.Path
Input FITS file.
extensions : list of str or int
Extension names and/or extension numbers to retain.
The primary HDU (extension 0) is always preserved.
replace : bool, optional
If True, overwrite the original file. Otherwise create a new file
with suffix ``.shrink.fits``. Default is False.
strict : bool, optional
If True (default), raise an exception if any requested extension
does not exist. Otherwise, issue a warning and continue.
Returns
-------
str
Name of the output FITS file.
"""
filename = Path(filename)
logger.info("shrinking file {}".format(filename))
if replace:
outfile = filename
else:
outfile = filename.with_suffix("").with_suffix(".shrink.fits")
removed = []
with fits.open(filename) as hdul:
# ----------------------------------------------------------
# Validate requested extensions
# ----------------------------------------------------------
available_names = {hdu.name for hdu in hdul[1:]}
available_numbers = set(range(1, len(hdul)))
missing = []
for ext in extensions:
if ext == 0 or (isinstance(ext, str) and ext.upper() == "PRIMARY"):
continue
if isinstance(ext, str):
if ext not in available_names:
missing.append(ext)
elif isinstance(ext, int):
if ext not in available_numbers:
missing.append(ext)
if missing:
msg = (
"The following requested extensions do not exist: "
+ ", ".join(map(str, missing))
)
if strict:
raise ValueError(msg)
logger.warning(msg)
primary = hdul[0].copy()
primary.header.add_history("File shrunk using shrink_fits().")
extensions_set = set(extensions)
new_hdus = [primary] # Always keep the primary HDU
for i, hdu in enumerate(hdul[1:], start=1):
keep = (i in extensions_set) or (hdu.name in extensions_set)
if keep:
new_hdus.append(hdu.copy())
else:
removed.append(hdu.name)
header = hdu.header.copy()
if isinstance(hdu, fits.ImageHDU):
new_hdu = fits.ImageHDU(
data=None,
header=header,
name=hdu.name)
elif isinstance(hdu, fits.BinTableHDU):
new_hdu = fits.BinTableHDU(
data=None,
header=header,
name=hdu.name
)
else:
# Fallback for any other extension type
new_hdu = fits.ImageHDU(
data=None,
header=header,
name=hdu.name
)
new_hdus.append(new_hdu)
if removed:
primary.header.add_history(
"Removed data from extensions: "
+ ", ".join(removed)
)
fits.HDUList(new_hdus).writeto(outfile, overwrite=True)
return str(outfile)
[docs]
def create_fits(datadict, header_dict, filename="Avg_neid_data.fits"):
"""
Create a multi-extension FITS file from a dictionary of data arrays
and headers.
Parameters
----------
datadict : dict
Dictionary mapping extension names to their corresponding data.
The first entry is written as the primary HDU. Remaining entries
are written as ``ImageHDU`` objects, except those listed in
``tablehdu``, which are written as ``BinTableHDU`` objects.
header_dict : dict
Dictionary mapping extension names to FITS header information.
Each value must be compatible with ``astropy.io.fits.Header``.
filename : str, optional
Name of the output FITS file. Default is
``"Avg_neid_data.fits"``.
Notes
-----
The function automatically selects ``BinTableHDU`` for extensions
listed in ``tablehdu`` (currently only ``ACTIVITY``). All other
extensions are written as ``ImageHDU`` objects. Existing files with
the same name are overwritten.
Examples
--------
>>> datadict = {
... "PRIMARY": np.zeros((100, 100)),
... "SCIENCE": np.random.random((50, 50)),
... "ACTIVITY": structured_array,
... }
>>> header_dict = {
... "PRIMARY": {"OBSERVER": "Varghese"},
... "SCIENCE": {"EXTNAME": "SCIENCE"},
... "ACTIVITY": {"COMMENT": "Activity indices"},
... }
>>> create_fits(datadict, header_dict, filename="output.fits")
"""
header_names = list(datadict.keys())
hdus = []
# --- Primary HDU ---
primary_data = datadict[header_names[0]]
primary_header = fits.Header(header_dict[header_names[0]])
primary_hdu = fits.PrimaryHDU(data=primary_data, header=primary_header)
hdus.append(primary_hdu)
# --- Extensions ---
tablehdu = ['ACTIVITY']
for exts in header_names[1:]:
data = datadict[exts]
ext_header = fits.Header(header_dict[exts])
if exts in tablehdu:
hdu = fits.BinTableHDU(data=data, header=ext_header, name=exts)
else:
hdu = fits.ImageHDU(data=data, header=ext_header, name=exts)
hdus.append(hdu)
# --- Write FITS ---
hdul = fits.HDUList(hdus)
hdul.writeto(filename, overwrite=True)
# End