diff --git a/src/shapepipe/modules/sextractor_package/sextractor_script.py b/src/shapepipe/modules/sextractor_package/sextractor_script.py index 8d05d7edb..e41ddad57 100644 --- a/src/shapepipe/modules/sextractor_package/sextractor_script.py +++ b/src/shapepipe/modules/sextractor_package/sextractor_script.py @@ -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 @@ -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) @@ -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: @@ -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: diff --git a/src/shapepipe/modules/sextractor_runner.py b/src/shapepipe/modules/sextractor_runner.py index 493bc5c96..2a0df4713 100644 --- a/src/shapepipe/modules/sextractor_runner.py +++ b/src/shapepipe/modules/sextractor_runner.py @@ -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 @@ -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 diff --git a/src/shapepipe/pipeline/file_io.py b/src/shapepipe/pipeline/file_io.py index 565437038..6ae0dbf00 100644 --- a/src/shapepipe/pipeline/file_io.py +++ b/src/shapepipe/pipeline/file_io.py @@ -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() ] ) @@ -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 ---------- @@ -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 @@ -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. diff --git a/tests/module/test_sextractor_post_process.py b/tests/module/test_sextractor_post_process.py index 446e1fbf6..868559e36 100644 --- a/tests/module/test_sextractor_post_process.py +++ b/tests/module/test_sextractor_post_process.py @@ -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", "