"""Define a class to compare FITS files generated by AutoWISP."""
import numpy
from astropy.io import fits
from autowisp.fits_utilities import read_image_components
from autowisp.tests import AutoWISPTestCase
[docs]
class FITSTestCase(AutoWISPTestCase):
"""Add assert for comparing AutoWISP generated FITS files."""
[docs]
def _compare_data(self, fname1, fname2, fits_components):
"""Assert that the image/error/mask data (or table data) match."""
header1 = fits_components[0]["header"]
if (
header1.get("XTENSION", "").strip() == "BINTABLE"
and header1.get("EXTNAME", "").strip() != "COMPRESSED_IMAGE"
):
self.assert_fits_tables_match(fname1, fname2)
return
for component in ["image", "error"]:
self.assertTrue(
numpy.isclose(
fits_components[0][component],
fits_components[1][component],
rtol=5e-4,
atol=1e-4
* (
numpy.mean(
numpy.absolute(fits_components[0][component])
)
+ numpy.mean(
numpy.absolute(fits_components[1][component])
)
),
).all(),
f"{component.title()} pixels in {fname1} do not match "
f"{component} pixels in {fname2}!",
)
self.assertTrue(
(fits_components[0]["mask"] == fits_components[1]["mask"]).all(),
f"Pixel mask in {fname1} do not match pixel mask in {fname2}!",
)
[docs]
def assert_fits_match(self, fname1, fname2):
"""Check that two FITS files have matching headers and data."""
fits_components = [
dict(
zip(
["image", "error", "mask", "header"],
read_image_components(fits_fname),
)
)
for fits_fname in [fname1, fname2]
]
self._compare_headers(
fname1,
fname2,
fits_components[0]["header"],
fits_components[1]["header"],
)
self._compare_data(fname1, fname2, fits_components)
[docs]
def assert_fits_tables_match(self, fname1, fname2):
"""Check that all table extensions in two FITS files match."""
with (
fits.open(fname1, mode="readonly") as hdul1,
fits.open(fname2, mode="readonly") as hdul2,
):
tables1 = [
hdu for hdu in hdul1 if isinstance(hdu, fits.BinTableHDU)
]
tables2 = [
hdu for hdu in hdul2 if isinstance(hdu, fits.BinTableHDU)
]
self.assertEqual(
len(tables1),
len(tables2),
f"Number of table extensions differs: "
f"{fname1} has {len(tables1)}, {fname2} has {len(tables2)}.",
)
for table_index, tbl in enumerate(zip(tables1, tables2)):
self.assertEqual(
tbl[0].name,
tbl[1].name,
f"Table extension #{table_index} name differs: "
f"{tbl[0].name!r} in {fname1} vs {tbl[1].name!r} in "
f"{fname2}.",
)
cols1 = set(tbl[0].columns.names)
cols2 = set(tbl[1].columns.names)
self.assertEqual(
cols1,
cols2,
f"Column names differ in table {tbl[0].name!r}:\n"
f" Only in {fname1}: {cols1 - cols2}\n"
f" Only in {fname2}: {cols2 - cols1}",
)
self.assertEqual(
len(tbl[0].data),
len(tbl[1].data),
f"Row count differs in table {tbl[0].name!r}: "
f"{len(tbl[0].data)} in {fname1} vs "
f"{len(tbl[1].data)} in {fname2}.",
)
for col_name in tbl[0].columns.names:
col1 = tbl[0].data[col_name]
col2 = tbl[1].data[col_name]
if numpy.issubdtype(col1.dtype, numpy.floating):
max_diff_i = numpy.argmax(numpy.abs(col1 - col2))
self.assertTrue(
numpy.isclose(
col1, col2, rtol=1e-8, atol=1e-8, equal_nan=True
).all(),
f"Column {col_name!r} values differ in table "
f"{tbl[0].name!r} between {fname1} and {fname2}."
f"Max diff values: {col1[max_diff_i]!r} vs "
f"{col2[max_diff_i]!r}.",
)
else:
self.assertTrue(
(col1 == col2).all(),
f"Column {col_name!r} values differ in table "
f"{tbl[0].name!r} between {fname1} and {fname2}.",
)