Source code for autowisp.tests.fits_test_case

"""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}.", )