Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 25 additions & 23 deletions src/shapepipe/modules/sextractor_package/sextractor_script.py
Original file line number Diff line number Diff line change
Expand Up @@ -338,12 +338,9 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None):
If SQL file not found

"""
cat = file_io.FITSCatalogue(
cat_path,
SEx_catalogue=True,
open_mode=file_io.BaseCatalogue.OpenMode.ReadWrite,
)
cat.open()
with fits.open(cat_path) as hdul:
hdus = [hdu.copy() for hdu in hdul]
objects = next(h for h in hdus if h.name == "LDAC_OBJECTS")

# One lock-free read of the whole header log: merge_headers wrote and
# closed it in an earlier step, and keyed SqliteDict reads would take an
Expand All @@ -358,7 +355,7 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None):
n_hdu = len(f_wcs[exp_keys[0]])

history = []
for idx in cat.get_data(1)[0][0]:
for idx in hdus[1].data[0][0]:
if re.split("HISTORY", idx)[0] == "":
history.append(idx)

Expand All @@ -374,11 +371,13 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None):

exp_list = list(dict.fromkeys(exp_list))

obj_id = np.copy(cat.get_data()["NUMBER"])
obj_id = np.copy(objects.data["NUMBER"])

ra = np.copy(cat.get_data()[pos_params[0]])
dec = np.copy(cat.get_data()[pos_params[1]])
ra = np.copy(objects.data[pos_params[0]])
dec = np.copy(objects.data[pos_params[1]])

column = file_io.FITSCatalogue.fits_column
epochs = []
n_epoch = np.zeros(len(obj_id), dtype="int32")
for idx, exp in enumerate(exp_list):
if exp not in f_wcs:
Expand Down Expand Up @@ -418,21 +417,24 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None):
ind[np.where(near)[0][ind_near]] = True
pos_tmp[ind] = idx_j
n_epoch[ind] += 1
exp_name = np.array([exp_list[idx] for n in range(len(obj_id))])
a = np.array(
[(obj_id[ii], exp_name[ii], pos_tmp[ii]) for ii in range(len(exp_name))],
dtype=[
("NUMBER", obj_id.dtype),
("EXP_NAME", exp_name.dtype),
("CCD_N", pos_tmp.dtype),
epochs.append(fits.BinTableHDU.from_columns(
[
column("NUMBER", obj_id),
column("EXP_NAME", np.full(len(obj_id), exp)),
column("CCD_N", pos_tmp),
],
)
cat.save_as_fits(data=a, ext_name=f"EPOCH_{idx}")
cat.open()

cat.add_col("N_EPOCH", n_epoch)
name=f"EPOCH_{idx}",
))

cat.close()
# One write of the whole catalogue: LDAC_OBJECTS with N_EPOCH appended
# (laid out as FITSCatalogue.add_col lays it out), then the EPOCH HDUs.
new = fits.BinTableHDU.from_columns(
objects.data.columns + fits.ColDefs([column("N_EPOCH", n_epoch)]),
name="LDAC_OBJECTS",
)
fits.HDUList(
[new if h is objects else h for h in hdus] + epochs
).writeto(cat_path, overwrite=True)


class SExtractorCaller:
Expand Down
167 changes: 98 additions & 69 deletions src/shapepipe/modules/sextractor_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,11 @@

"""

import contextlib
import os
import re
import shutil
import tempfile

from shapepipe.modules.module_decorator import module_runner
from shapepipe.modules.sextractor_package import match_catalogue as mc
Expand Down Expand Up @@ -97,79 +101,104 @@ def sextractor_runner(
f"{run_dirs['tmp']}/seg_vignet{file_number_string}.param",
)

# Create sextractor caller class instance
ss_inst = ss.SExtractorCaller(
input_file_list,
run_dirs["output"],
file_number_string,
dot_sex,
dot_param,
dot_conv,
weight_file,
flag_file,
psf_file,
detection_image,
detection_weight,
zp_from_header,
bkg_from_header,
zero_point_key=zp_key,
background_key=bkg_key,
check_image=check_image,
output_prefix=prefix,
)

# Generate sextractor command line
command_line = ss_inst.make_command_line(exec_path)
w_log.info(f"Calling command: {command_line}")

# Execute command line
stderr, stdout = execute(command_line)

# Parse SExtractor errors
stdout, stderr = ss_inst.parse_errors(stderr, stdout)

# SEG_VIGNET is cut before the join, which relabels its stamps with the
# rows' new NUMBERs.
if seg_vignet:
ss.add_seg_vignet(
ss_inst.path_output_file,
ss_inst.check_paths["SEGMENTATION"],
w_log=w_log,
# WORK_DIR (optional, environment-expanded): SExtractor writes its
# catalogue and check images in a fresh directory there, the SEG_VIGNET
# cut, the join and the post-processing each rewrite the catalogue there,
# and the files are moved to the run's output directory once complete.
# On node-local disk this keeps the whole-catalogue rewrites off NFS,
# where each costs minutes under load. The directory is removed on the way
# out, also on failure. Without WORK_DIR everything is written in the
# output directory.
if config.has_option(module_config_sec, "WORK_DIR"):
work = tempfile.TemporaryDirectory(
prefix=f"sp-detect{file_number_string}.",
dir=config.getexpanded(module_config_sec, "WORK_DIR"),
)

# MATCH_CATALOGUE (optional, environment-expanded path; empty for none):
# take membership and NUMBER from that external catalogue of the same
# image, before the post-processing keys the epoch HDUs on NUMBER.
match_path = (
config.getexpanded(module_config_sec, "MATCH_CATALOGUE")
if config.has_option(module_config_sec, "MATCH_CATALOGUE")
else ""
)
if match_path:
mc.match_catalogue(
ss_inst.path_output_file,
match_path,
radius=config.getfloat(module_config_sec, "MATCH_RADIUS"),
min_fraction=config.getfloat(
module_config_sec, "MATCH_MIN_FRACTION"
),
tolerated_unpaired=config.getint(
module_config_sec, "MATCH_TOLERATED_UNPAIRED"
),
w_log=w_log,
else:
work = contextlib.nullcontext(run_dirs["output"])

with work as work_dir:
# Create sextractor caller class instance
ss_inst = ss.SExtractorCaller(
input_file_list,
work_dir,
file_number_string,
dot_sex,
dot_param,
dot_conv,
weight_file,
flag_file,
psf_file,
detection_image,
detection_weight,
zp_from_header,
bkg_from_header,
zero_point_key=zp_key,
background_key=bkg_key,
check_image=check_image,
output_prefix=prefix,
)

# Run sextractor post processing
if config.getboolean(module_config_sec, "MAKE_POST_PROCESS"):
pos_params = config.getlist(module_config_sec, "WORLD_POSITION")
ccd_size = config.getlist(module_config_sec, "CCD_SIZE")
ss.make_post_process(
ss_inst.path_output_file,
f_wcs_path,
pos_params,
ccd_size,
w_log=w_log,
# Generate sextractor command line
command_line = ss_inst.make_command_line(exec_path)
w_log.info(f"Calling command: {command_line}")

# Execute command line
stderr, stdout = execute(command_line)

# Parse SExtractor errors
stdout, stderr = ss_inst.parse_errors(stderr, stdout)

# SEG_VIGNET is cut before the join, which relabels its stamps with the
# rows' new NUMBERs.
if seg_vignet:
ss.add_seg_vignet(
ss_inst.path_output_file,
ss_inst.check_paths["SEGMENTATION"],
w_log=w_log,
)

# MATCH_CATALOGUE (optional, environment-expanded path; empty for none):
# take membership and NUMBER from that external catalogue of the same
# image, before the post-processing keys the epoch HDUs on NUMBER.
match_path = (
config.getexpanded(module_config_sec, "MATCH_CATALOGUE")
if config.has_option(module_config_sec, "MATCH_CATALOGUE")
else ""
)
if match_path:
mc.match_catalogue(
ss_inst.path_output_file,
match_path,
radius=config.getfloat(module_config_sec, "MATCH_RADIUS"),
min_fraction=config.getfloat(
module_config_sec, "MATCH_MIN_FRACTION"
),
tolerated_unpaired=config.getint(
module_config_sec, "MATCH_TOLERATED_UNPAIRED"
),
w_log=w_log,
)

# Run sextractor post processing
if config.getboolean(module_config_sec, "MAKE_POST_PROCESS"):
pos_params = config.getlist(module_config_sec, "WORLD_POSITION")
ccd_size = config.getlist(module_config_sec, "CCD_SIZE")
ss.make_post_process(
ss_inst.path_output_file,
f_wcs_path,
pos_params,
ccd_size,
w_log=w_log,
)

if work_dir != run_dirs["output"]:
for path in [*ss_inst.check_paths.values(),
ss_inst.path_output_file]:
shutil.move(
path,
os.path.join(run_dirs["output"], os.path.basename(path)),
)

# Return stdout and stderr
return stdout, stderr
15 changes: 9 additions & 6 deletions src/shapepipe/pipeline/file_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -1441,7 +1441,7 @@ def add_cols(

col_list = self._cat_data[hdu_no].data.columns + fits.ColDefs(
[
self._make_fits_col(col_name, col_data)
self.fits_column(col_name, col_data)
for col_name, col_data in columns.items()
]
)
Expand All @@ -1462,13 +1462,15 @@ def add_cols(
memmap=self.use_memmap,
)

def _make_fits_col(self, col_name, col_data):
"""Make FITS Column.
@staticmethod
def fits_column(col_name, col_data):
"""FITS Column.

Build the ``astropy.io.fits.Column`` that :meth:`add_cols` appends
for one array: the FITS type from :meth:`_get_fits_col_type`, a
repeat count and ``TDIM`` for multi-dimensional arrays, and a width
set by the longest entry for strings.
set by the longest entry for strings. :meth:`save_as_fits` lays out
the columns of a new table HDU the same way.

Parameters
----------
Expand All @@ -1483,7 +1485,7 @@ def _make_fits_col(self, col_name, col_data):
The column

"""
data_type = self._get_fits_col_type(col_data)
data_type = FITSCatalogue._get_fits_col_type(col_data)
data_shape = col_data.shape[1:]
dim = None
mem_size = 1
Expand Down Expand Up @@ -1575,7 +1577,8 @@ def _append_col(self, column, hdu_no=None):
else:
raise BaseCatalogue.catalogueNotOpen(self.fullpath)

def _get_fits_col_type(self, col_data):
@staticmethod
def _get_fits_col_type(col_data):
"""Get FITS Column Type.

Get the FITS data type of a given column.
Expand Down
57 changes: 57 additions & 0 deletions tests/module/test_sextractor_post_process.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,3 +195,60 @@ def test_exposure_missing_from_header_log_raises(tmp_path):
["XWIN_WORLD", "YWIN_WORLD"],
["0", str(CCD_NPIX), "0", str(CCD_NPIX)],
)


def test_post_process_writes_the_file_catalogue_appends_would(tmp_path):
"""make_post_process writes the catalogue once; the bytes are those of
one FITSCatalogue.save_as_fits per EPOCH HDU followed by add_col of
N_EPOCH, the per-step route that rewrote the whole file each time."""
from shapepipe.pipeline import file_io

exp_names = ["123456", "654321"]
header_files = []
for name in exp_names:
npy_path = tmp_path / f"headers-{name}.npy"
_write_exposure_headers(npy_path)
header_files.append([str(npy_path)])
merge_headers(header_files, str(tmp_path), tile_number="54")
sqlite_path = tmp_path / "log_exp_headers54.sqlite"

# Objects on CCDs 0 and 2 and one on no CCD.
positions = np.array(
[
_make_ccd_wcs(ccd)[0].all_pix2world([[50.0, 50.0]], 0)[0]
for ccd in (0, 2)
]
+ [[170.0, 10.0]]
)
cat_path = tmp_path / "sexcat.fits"
_write_sex_ldac(cat_path, exp_names, positions)
ref_path = tmp_path / "reference.fits"
ref_path.write_bytes(cat_path.read_bytes())

sextractor_script.make_post_process(
str(cat_path),
str(sqlite_path),
["XWIN_WORLD", "YWIN_WORLD"],
["0", str(CCD_NPIX), "0", str(CCD_NPIX)],
)

number = np.arange(1, 4, dtype=">i4")
ccd_n = np.array([0, 2, -1], dtype="int32")
ref = file_io.FITSCatalogue(
str(ref_path),
SEx_catalogue=True,
open_mode=file_io.BaseCatalogue.OpenMode.ReadWrite,
)
ref.open()
for idx, exp in enumerate(exp_names):
epoch = np.array(
list(zip(number, [exp] * 3, ccd_n)),
dtype=[("NUMBER", number.dtype), ("EXP_NAME", "<U6"),
("CCD_N", ccd_n.dtype)],
)
ref.save_as_fits(data=epoch, ext_name=f"EPOCH_{idx}")
ref.open()
ref.add_col("N_EPOCH", np.array([2, 2, 0], dtype="int32"))
ref.close()

assert cat_path.read_bytes() == ref_path.read_bytes()
Loading
Loading