From 4dd460dd07b0afd48df3e56fd5b5f9d00f21a389 Mon Sep 17 00:00:00 2001 From: robert Date: Thu, 24 Sep 2026 10:36:00 +0200 Subject: [PATCH 1/6] fix min_value, smaller overlap nii to save RAM --- TPTBox/stitching/stitching.py | 184 +++++++++++++++++++++------- TPTBox/stitching/stitching_tools.py | 17 ++- 2 files changed, 157 insertions(+), 44 deletions(-) diff --git a/TPTBox/stitching/stitching.py b/TPTBox/stitching/stitching.py index 1fcf676..0f1fb33 100755 --- a/TPTBox/stitching/stitching.py +++ b/TPTBox/stitching/stitching.py @@ -464,15 +464,25 @@ def _auto_output_dtype(niis: list[nib.nifti1.Nifti1Image]) -> type: """Pick the smallest lossless output dtype from the inputs' on-disk dtypes. Reads only the NIfTI headers — no pixel data is loaded. If every input is - integer-typed, returns the widest integer dtype that covers them all; - otherwise falls back to float32. This lets the stitcher store magnitude - MR outputs as uint16 (or int16) when the source already fit in 16 bits, - halving the on-disk and downstream RAM footprint versus float64. + integer-typed AND has a trivial slope/inter (so the raw storage range is + the true value range), returns the widest integer dtype that covers them + all; otherwise falls back to float32. This lets the stitcher store + magnitude MR outputs as uint16 (or int16) when the source already fit in + 16 bits, halving the on-disk and downstream RAM footprint versus float64. """ dtypes = [np.dtype(nii.get_data_dtype()) for nii in niis] - if all(np.issubdtype(d, np.integer) for d in dtypes): - return max(dtypes, key=lambda d: d.itemsize).type # e.g. np.uint16 - return np.float32 + if not all(np.issubdtype(d, np.integer) for d in dtypes): + return np.float32 + for nii in niis: + slope = getattr(nii.dataobj, "slope", 1.0) + inter = getattr(nii.dataobj, "inter", 0.0) + if slope is None or inter is None: + return np.float32 + if not (np.isfinite(slope) and np.isfinite(inter)): + return np.float32 + if float(slope) != 1.0 or float(inter) != 0.0: + return np.float32 + return max(dtypes, key=lambda d: d.itemsize).type # e.g. np.uint16 def main( # noqa: C901 @@ -481,7 +491,7 @@ def main( # noqa: C901 match_histogram: bool = False, store_ramp: bool = False, verbose: bool = False, - min_value: float = 0, + min_value: float | None = None, bias_field: bool = True, crop_to_bias_field: bool = False, crop_empty: bool = False, @@ -517,8 +527,16 @@ def main( # noqa: C901 store_ramp: If True, also saves the per-volume blend weights as a 4-D NIfTI alongside the stitched output. verbose: If True, prints progress messages to stdout. - min_value: Background value (0 for MRI, -1024 for CT). Voxels at or - below this value are replaced by it in the output. + min_value: Background fill used as ``cval`` when resampling each chunk + into the target space, and — when explicitly set — a hard floor + applied to the stitched output. Pass ``0`` for MR, ``-1024`` for + CT. Default ``None``: use ``0`` internally as the background / + NaN-fill and only apply a hard floor when the output dtype cannot + represent negatives (unsigned integer types), which prevents + negative-to-huge-positive wraparound on the cast. Explicit values + (segmentations force ``0``, CT callers pass ``-1024``) are always + enforced; ``None`` lets signed / float outputs keep legitimate + negative signal (Philips-scaled fat-fraction, phase, B0 offsets). bias_field: If True, applies N4 bias-field correction to each input before stitching. Forced to False for segmentations. crop_to_bias_field: If True, crops each bias-corrected volume to the @@ -554,6 +572,7 @@ def main( # noqa: C901 min_value = 0 match_histogram = False histogram = None + _bg_value: float = 0.0 if min_value is None else float(min_value) if len(images) == 0 or len(images) == 1: print("!!! Need at least two images (-i ...nii.gz ...nii.gz) to stitch!!!\n Got " + str(images)) return None, None @@ -564,8 +583,20 @@ def main( # noqa: C901 for f_name in images: if isinstance(f_name, (Path, str)): print("Load ", f_name, Path(f_name)) if verbose else None - # Load Nii - nii: nib.nifti1.Nifti1Image = nib.load(f_name) # type: ignore + # Load NII + try: + from TPTBox.core.nii_wrapper import to_nii as _to_nii_load + + _nii_wrap = _to_nii_load(Path(f_name)) + _arr_corrected = _nii_wrap.get_array() + nii = nib.Nifti1Image(_arr_corrected, _nii_wrap.affine) + nii.set_data_dtype(_arr_corrected.dtype) + try: + nii.header.set_slope_inter(1.0, 0.0) + except Exception: # noqa: BLE001 + pass + except Exception: # noqa: BLE001 + nii: nib.nifti1.Nifti1Image = nib.load(f_name) # type: ignore else: nii = f_name @@ -589,7 +620,7 @@ def main( # noqa: C901 image = get_array(nii) matched = match_histograms(image.astype(float), reference.astype(float)) - matched[matched <= min_value] = min_value + matched[matched <= _bg_value] = _bg_value nii = set_array(nii, matched) niis.append(nii) @@ -617,11 +648,12 @@ def main( # noqa: C901 dtype = dtype2 else: # Auto-detect the output dtype from the input headers when the caller - # asked for "auto". Blending math still runs in float; only the final - # cast at save time uses the picked dtype (e.g. uint16 for magnitude MR). + # asked for "auto". Blending math runs in float32 (was float64) — + # halves peak RAM of target_list/occupancy_list on large stitched + # volumes; the final cast at save time uses `dtype` (uint16 etc.). if isinstance(dtype, str) and dtype == "auto": dtype = _auto_output_dtype(niis) - dtype2 = float + dtype2 = np.float32 nii_out = get_max_affine_and_shape(corners_current, affines, min_spacing=min_spacing, dtype=dtype2, verbose=verbose) target_list = [] occupancy_list = [] @@ -629,10 +661,10 @@ def main( # noqa: C901 print("### resample to new space ###") if verbose else None for i, nii in enumerate(niis, 1): print(f"{i:2}/{len(niis):2} resampled", end="\r") if verbose else None - nii_new = nip.resample_from_to(nii, nii_out, 0 if is_segmentation else 3, mode="constant", cval=min_value) + nii_new = nip.resample_from_to(nii, nii_out, 0 if is_segmentation else 3, mode="constant", cval=_bg_value) arr_new = get_array(nii_new) if not is_segmentation and np.issubdtype(arr_new.dtype, np.floating): - np.nan_to_num(arr_new, copy=False, nan=min_value, posinf=min_value, neginf=min_value) + np.nan_to_num(arr_new, copy=False, nan=_bg_value, posinf=_bg_value, neginf=_bg_value) target_list.append(arr_new) b = nib.Nifti1Image(np.ones(nii.shape, dtype=np.float32), affine=nii.affine) # type: ignore b = nip.resample_from_to(b, nii_new, 0, cval=0, mode="constant") @@ -649,22 +681,50 @@ def main( # noqa: C901 # inside the ramp loop for pairs that don't touch. `_occupancy_bbox` # returns None for an empty occupancy — treated as "no overlap possible". bboxes = [_occupancy_bbox(occ) for occ in occupancy_list] + grid_shape_arr = np.asarray(occupancy_list[0].shape, dtype=np.int64) + # Padding around the joint AABB: `ramp_edge_min_value` so binary_opening's + # erode+dilate at the crop boundary yields the same result as on the full + # volume; +1 slack for the distance transform. + _ramp_pad = max(int(ramp_edge_min_value), 1) + 1 # ramp stitching combinations = list(itertools.combinations(range(len(target_list)), 2)) + _ramp_done = 0 + _ramp_skipped_aabb = 0 + _ramp_skipped_no_voxel_overlap = 0 for idx, item in enumerate(combinations, 1): print(f"{idx:2}/{len(combinations):2} ramp stitching", end="\r") if verbose else None # Skip disjoint pairs before touching the full-volume arrays. if not _aabb_overlaps(bboxes[item[0]], bboxes[item[1]]): + _ramp_skipped_aabb += 1 continue + # Work on the union-AABB sub-volume of the two chunks. Outside the + # union `arr_i / sum_` equals the original `arr_i_full` (there, + # overlap = 0, arr_i_ = binary mask of arr_i and the "other" mask + # is 0, so sum_ = arr_i_ and arr_i / sum_ = arr_i). And chunk i's + # occupancy support is fully contained in bboxes[i] ⊆ union, so + # writing the result back only at the sub-slice is functionally + # identical to the previous full-volume compute — but the ramp's + # peak RAM drops from ~full-volume to ~union-AABB size (typically + # 2 adjacent chunks tall). + lo_i, hi_i = bboxes[item[0]] + lo_j, hi_j = bboxes[item[1]] + lo = np.maximum(np.minimum(lo_i, lo_j) - _ramp_pad, 0) + hi = np.minimum(np.maximum(hi_i, hi_j) + _ramp_pad, grid_shape_arr - 1) + sub = ( + slice(int(lo[0]), int(hi[0]) + 1), + slice(int(lo[1]), int(hi[1]) + 1), + slice(int(lo[2]), int(hi[2]) + 1), + ) # TODO fix intersection with more than two occupancies arr_1_full = occupancy_list[item[0]] arr_2_full = occupancy_list[item[1]] ### structure = np.ones((ramp_edge_min_value, ramp_edge_min_value, ramp_edge_min_value), dtype=bool) - arr_1: np.ndarray = arr_1_full.copy() - arr_2: np.ndarray = arr_2_full.copy() + arr_1: np.ndarray = arr_1_full[sub].astype(np.float32, copy=True) + arr_2: np.ndarray = arr_2_full[sub].astype(np.float32, copy=True) overlap = (arr_1 * arr_2) > 0.0 if overlap.sum() > 0: + _ramp_done += 1 arr_1_ = (arr_1 > 0.0).astype(np.float32) - overlap arr_2_ = (arr_2 > 0.0).astype(np.float32) - overlap if ramp_edge_min_value == 0: @@ -680,28 +740,32 @@ def main( # noqa: C901 arr_2_[overlap] = arr_2[overlap] sum_ = arr_1_ + arr_2_ sum_[sum_ == 0] = 1.0 - arr_1_full = arr_1 / sum_ - arr_2_full = arr_2 / sum_ - if arr_1_full.max() != 1: + arr_1_sub = arr_1 / sum_ + arr_2_sub = arr_2 / sum_ + # Chunk i's occupancy is fully contained inside bboxes[i] ⊆ sub, + # so the sub-array max equals the volume-wide max. + max_1 = float(arr_1_sub.max()) + max_2 = float(arr_2_sub.max()) + if max_1 != 1: import warnings warnings.warn( - str((arr_1_full.min(), arr_1_full.max())) + " the image in fully incorporated insight of an other " + str(images), + str((float(arr_1_sub.min()), max_1)) + " the image in fully incorporated insight of an other " + str(images), stacklevel=4, ) if kick_out_fully_integrated_images: images.pop(item[0]) - elif arr_2_full.max() != 1: + elif max_2 != 1: import warnings warnings.warn( - str((arr_2_full.min(), arr_2_full.max())) + " the image in fully incorporated insight of an other " + str(images), + str((float(arr_2_sub.min()), max_2)) + " the image in fully incorporated insight of an other " + str(images), stacklevel=4, ) if kick_out_fully_integrated_images: images.pop(item[1]) - if (arr_1_full.max() != 1 or arr_2_full.max() != 1) and kick_out_fully_integrated_images: + if (max_1 != 1 or max_2 != 1) and kick_out_fully_integrated_images: print("kick_out_fully_integrated_images") print(images) @@ -721,20 +785,45 @@ def main( # noqa: C901 kick_out_fully_integrated_images, save, ) - # assert arr_1_full.max() == 1, (arr_1_full.min(), arr_1_full.max()) - # assert arr_2_full.max() == 1, (arr_2_full.min(), arr_2_full.max()) - occupancy_list[item[0]] = arr_1_full - occupancy_list[item[1]] = arr_2_full + arr_1_full[sub] = arr_1_sub.astype(arr_1_full.dtype, copy=False) + arr_2_full[sub] = arr_2_sub.astype(arr_2_full.dtype, copy=False) else: + _ramp_skipped_no_voxel_overlap += 1 continue - occupancy_arr = np.stack(occupancy_list) - if is_segmentation: - occupancy_arr = np.round(occupancy_arr) # TODO assuming only two intersecting regions - target_arr = np.stack(target_list) * occupancy_arr + if verbose: + print( + f"\nramp summary: {_ramp_done} computed, " + f"{_ramp_skipped_aabb} skipped (disjoint AABB), " + f"{_ramp_skipped_no_voxel_overlap} skipped (no voxel overlap) " + f"of {len(combinations)} pairs" + ) + # Aggregate: in-place accumulate `t * occupancy` per chunk instead of + # `np.stack(target_list) * np.stack(occupancy_list)` which would peak at + # ~2 × N × volume of temporary float arrays. + target_arr = np.zeros(target_list[0].shape, dtype=dtype2) if is_segmentation: - target_arr = target_arr.astype(dtype2) - target_arr = target_arr.sum(0) - target_arr[target_arr <= min_value] = min_value + for t, o in zip(target_list, occupancy_list): + target_arr += (t * np.round(o)).astype(dtype2) # TODO assuming only two intersecting regions + else: + for t, o in zip(target_list, occupancy_list): + target_arr += t * o + # Hard-floor policy: + # * If the caller passed `min_value` explicitly (segmentations force 0, + # CT typically passes -1024), always apply that floor. + # * If `min_value is None` (the default), only apply a floor when the + # output dtype cannot represent negatives (unsigned integer types) — + # otherwise negative-to-huge-positive wraparound would silently corrupt + # the save. Signed/float outputs keep legitimate negative signal + # (Philips-scaled fat-fraction, phase, B0 offsets). + _floor: float | None + if min_value is not None: + _floor = float(min_value) + elif np.issubdtype(np.dtype(dtype), np.unsignedinteger): + _floor = 0.0 + else: + _floor = None + if _floor is not None: + target_arr[target_arr <= _floor] = _floor print("\n### Save ###") if verbose else None if output is not None: output = str(output) @@ -748,9 +837,22 @@ def main( # noqa: C901 if bias_field: nii_out = n4_bias_field_correction(nii_out) if crop_empty: - nii_occ = set_array(nii_out, occupancy_arr) - ex_slice = compute_crop_slice(nii_occ) - nii_out = nii_out.slicer[ex_slice] + # Crop to the union of per-chunk occupancy AABBs. The previous path + # went through compute_crop_slice on a 4-D (N, X, Y, Z) stack, which + # sliced the wrong axes when applied to 3-D nii_out; deriving the + # crop directly from `bboxes` is both cheaper and correct. + valid_bboxes = [b for b in bboxes if b is not None] + if valid_bboxes: + lo = np.stack([b[0] for b in valid_bboxes]).min(axis=0) + hi = np.stack([b[1] for b in valid_bboxes]).max(axis=0) + ex_slice = ( + slice(int(lo[0]), int(hi[0]) + 1), + slice(int(lo[1]), int(hi[1]) + 1), + slice(int(lo[2]), int(hi[2]) + 1), + ) + nii_out = nii_out.slicer[ex_slice] + else: + ex_slice = () else: ex_slice = () diff --git a/TPTBox/stitching/stitching_tools.py b/TPTBox/stitching/stitching_tools.py index a59d8ec..d42c67d 100755 --- a/TPTBox/stitching/stitching_tools.py +++ b/TPTBox/stitching/stitching_tools.py @@ -23,6 +23,7 @@ def stitching( match_histogram: bool = False, store_ramp: bool = False, ramp_path=None, + min_value: float | None = None, ) -> tuple: """Stitch a list of BIDS/NII volumes into a single output NIfTI file. @@ -37,8 +38,18 @@ def stitching( :class:`BIDS_FILE` is provided, the ``"nii.gz"`` file path is used. is_seg: If True, treats the inputs as segmentation images (disables bias field and histogram matching, uses integer dtypes). - is_ct: If True, sets the background ``min_value`` to ``-1024`` (CT - air) instead of ``0`` (MRI). + is_ct: If True and ``min_value`` was not passed explicitly, sets the + background ``min_value`` to ``-1024`` (CT air). + min_value: Explicit background/floor for :func:`stitching_raw`. + ``None`` (the default, and what non-CT MR normally uses) lets the + stitcher only apply a hard floor when the output dtype cannot + represent negatives (unsigned int), so quantitative maps + (fat-fraction, phase, B0) keep their legitimate negatives. Pass + ``0`` to floor the output at 0 — the right choice for magnitude + MR (in/out-phase, water, fat, per-echo magnitude), where the + cubic-spline resample can leave small negative ringing artefacts + that a downstream network wasn't trained on. Overrides ``is_ct`` + when both are given. verbose_stitching: If True, forwards verbose output from the low-level stitching routine. bias_field: If True, applies N4 bias-field correction to each input. @@ -66,7 +77,7 @@ def stitching( match_histogram=match_histogram, store_ramp=store_ramp, verbose=verbose_stitching, - min_value=-1024 if is_ct else 0, + min_value=(min_value if min_value is not None else (-1024 if is_ct else None)), bias_field=bias_field, kick_out_fully_integrated_images=kick_out_fully_integrated_images, is_segmentation=is_seg, From 03108d924b344b6d393984b3cb03aa5e81bfc03b Mon Sep 17 00:00:00 2001 From: robert Date: Fri, 25 Sep 2026 09:17:28 +0200 Subject: [PATCH 2/6] add c_val option --- TPTBox/core/internal/nii_help.py | 3 ++- TPTBox/core/nii_wrapper.py | 8 ++++---- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/TPTBox/core/internal/nii_help.py b/TPTBox/core/internal/nii_help.py index 5f8b081..ba79b0d 100644 --- a/TPTBox/core/internal/nii_help.py +++ b/TPTBox/core/internal/nii_help.py @@ -205,6 +205,7 @@ def _resample_from_to( mode: MODES = "nearest", align_corners: bool | Sentinel = Sentinel(), # noqa: B008 out_dtype: np.dtype | type | str | None = None, + c_val: float | None = None, ) -> tuple[np.ndarray, np.ndarray, object]: """Resample *from_img* into the voxel space defined by *to_img*. @@ -323,7 +324,7 @@ def _resample_from_to( to_shape, order=order, mode=mode, - cval=from_img.get_c_val(), + cval=from_img.get_c_val(c_val), output=scipy_out, ) if post_cast is not None: diff --git a/TPTBox/core/nii_wrapper.py b/TPTBox/core/nii_wrapper.py index 3cbb826..93c6512 100755 --- a/TPTBox/core/nii_wrapper.py +++ b/TPTBox/core/nii_wrapper.py @@ -999,7 +999,7 @@ def pad_to(self, target_shape: list[int] | tuple[int, int, int] | Self, mode: MO s = s.apply_crop(tuple(crop),inplace=inplace) return s.apply_pad(padding,inplace=inplace,mode=mode) - def apply_pad(self, padd: Sequence[tuple[int | None, int | None]] | int | None, mode: MODES = "constant", inplace=False, verbose: logging = True,) -> Self: + def apply_pad(self, padd: Sequence[tuple[int | None, int | None]] | int | None, mode: MODES = "constant", inplace=False, verbose: logging = True, c_val: float | None = None) -> Self: """Pads the image with explicit per-axis ``(before, after)`` amounts. The affine is updated so that the world-space origin is preserved (i.e. the @@ -1064,7 +1064,7 @@ def apply_pad(self, padd: Sequence[tuple[int | None, int | None]] | int | None, args = {} if mode == "constant": - args["constant_values"] = self.get_c_val() + args["constant_values"] = self.get_c_val(c_val) if mode == "nearest": mode = "edge" @@ -1249,7 +1249,7 @@ def resample_from_to(self, to_vox_map:Image_Reference|Has_Grid|tuple[SHAPE,AFFIN pad_after = dst_shape - shift - src_shape pad = tuple((int(b), int(a)) for b, a in zip(pad_before, pad_after)) try: - ret = s.apply_pad(pad,mode=mode,inplace=inplace,verbose=verbose) + ret = s.apply_pad(pad,mode=mode,inplace=inplace,verbose=verbose,c_val=c_val) valid = ret.assert_affine(mapping,raise_error=False,origin_tolerance=0.0001,error_tolerance=0.0001,shape_tolerance=0) if valid: log.print(f"resample_from_to only needs padding/cropping {pad}",verbose=verbose) @@ -1263,7 +1263,7 @@ def resample_from_to(self, to_vox_map:Image_Reference|Has_Grid|tuple[SHAPE,AFFIN log.print(f"resample_from_to: {self} to {mapping}",verbose=verbose) if order is None: order = 0 if self.seg else 3 - nii = _resample_from_to(self, mapping,order=order, mode=mode,align_corners=align_corners, out_dtype=out_dtype) + nii = _resample_from_to(self, mapping,order=order, mode=mode,align_corners=align_corners, out_dtype=out_dtype, c_val=c_val) if inplace: From bbd89646b5b94bbe658b82bd147c39957f22079c Mon Sep 17 00:00:00 2001 From: robert Date: Fri, 25 Sep 2026 09:17:56 +0200 Subject: [PATCH 3/6] only warn if memory is not trival --- .../segmentation/VibeSeg/inference_nnunet.py | 43 +++++++++++++++---- 1 file changed, 34 insertions(+), 9 deletions(-) diff --git a/TPTBox/segmentation/VibeSeg/inference_nnunet.py b/TPTBox/segmentation/VibeSeg/inference_nnunet.py index 3a14dae..0f0711b 100644 --- a/TPTBox/segmentation/VibeSeg/inference_nnunet.py +++ b/TPTBox/segmentation/VibeSeg/inference_nnunet.py @@ -267,13 +267,6 @@ def run_inference_on_file( if "memory_factor" not in ds_info: missing_mem_keys.append("memory_factor") memory_factor = float(ds_info.get("memory_factor", 160)) - if missing_mem_keys: - _suggest_memory_estimation_script( - idx, - model_path, - f"Memory parameter(s) {missing_mem_keys} not set in the model's dataset.json; falling back to defaults. {memory_base=}, {memory_factor=}", - logger=logger, - ) use_folds_arg = tuple(folds) if len(folds) != 5 else None # Include every setting that changes the loaded predictor so a cache hit is always equivalent @@ -362,11 +355,43 @@ def run_inference_on_file( if padd != 0: p = (padd, padd) input_nii = [i.apply_pad([p, p, p], mode="reflect") for i in input_nii] + + def _defaults_fit_gpu(memory: float, gpu: int | None, safety_factor: float = 2.0) -> bool: + """Whether the fallback ``memory_base`` reservation is a small fraction of GPU memory. + + Returns ``True`` when ``memory_base * safety_factor <= total_gpu_memory_mb`` — + i.e. the fallback reservation leaves ample headroom for tile scheduling and + the missing ``memory_base``/``memory_factor`` defaults are safe to use + without warning. Returns ``False`` on CUDA-unavailable systems or when GPU + memory can't be probed, so the load-time warning still fires on + memory-constrained setups (CPU-only, small GPUs, or driver errors). + """ + try: + import torch + + if not torch.cuda.is_available(): + return False + device = torch.device(f"cuda:{gpu}") if gpu is not None else torch.device("cuda:0") + _, total_bytes = torch.cuda.mem_get_info(device) + total_mb = total_bytes / (1024**2) + except Exception: # noqa: BLE001 + return False + return memory * safety_factor <= total_mb + + num_classes = int(nnunet.label_manager.num_segmentation_heads) + est_full_mb = estimate_peak_ram_mb(input_nii[0].shape, num_classes, len(input_nii)) + # Only warn if memory requirement is non-trivial. + if missing_mem_keys and not _defaults_fit_gpu(est_full_mb, gpu): + _suggest_memory_estimation_script( + idx, + model_path, + f"Memory parameter(s) {missing_mem_keys} not set in the model's dataset.json; falling back to defaults. {memory_base=}, {memory_factor=}", + logger=logger, + ) + if _cpu_chunks is None or _cpu_chunks <= 1: - num_classes = int(nnunet.label_manager.num_segmentation_heads) total_ram_mb = _get_total_ram_mb() target_ram_mb = total_ram_mb * 0.5 - est_full_mb = estimate_peak_ram_mb(input_nii[0].shape, num_classes, len(input_nii)) if est_full_mb > target_ram_mb: shape = input_nii[0].shape split_axis = int(np.argmax(shape)) From 0ebd092be082619e55b5e8df148a37de639eb50c Mon Sep 17 00:00:00 2001 From: robert Date: Fri, 25 Sep 2026 09:18:18 +0200 Subject: [PATCH 4/6] Refactor stiching to us TPTBox function which are more efficent. --- TPTBox/stitching/stitching.py | 589 +++++++++++++--------------------- 1 file changed, 218 insertions(+), 371 deletions(-) diff --git a/TPTBox/stitching/stitching.py b/TPTBox/stitching/stitching.py index 0f1fb33..e69e011 100755 --- a/TPTBox/stitching/stitching.py +++ b/TPTBox/stitching/stitching.py @@ -1,17 +1,45 @@ from __future__ import annotations +import contextlib import itertools +import warnings +from functools import cache from pathlib import Path -import nibabel as nib -import nibabel.processing as nip import numpy as np from nibabel.affines import apply_affine -from nibabel.nifti1 import Nifti1Image from scipy.ndimage import binary_opening, distance_transform_edt from scipy.spatial import ConvexHull from skimage.exposure import match_histograms +from TPTBox.core.nii_wrapper import NII, to_nii +from TPTBox.core.np_utils import np_bbox_binary +from TPTBox.logger import Print_Logger + +logger = Print_Logger() + + +@contextlib.contextmanager +def _suppress_dtype_warning(): + """Silence the "Loaded NIfTY: incorrect dtype detected" ``UserWarning``. + + ``NII._check_if_nifty_is_lying_about_its_dtype`` (nii_wrapper.py, three + ``warnings.warn`` sites around line 151/163/172) fires this warning + during ``_unpack`` whenever the on-disk dtype doesn't match the array's + actual range — typical for Philips-scaled magnitude MR arriving as + ``int16`` with a large ``scl_slope``. The stitching pipeline handles + that case explicitly (rebuilds a scale-1 image with the widened dtype + in ``main``'s loading loop), so the warning is noise here — but we + only suppress it for the current call's stack, not globally. + """ + with warnings.catch_warnings(): + warnings.filterwarnings( + "ignore", + message=r"Loaded NIfTY: incorrect dtype detected.*", + category=UserWarning, + ) + yield + def get_rotation_and_spacing_from_affine(affine: np.ndarray) -> tuple[np.ndarray, np.ndarray]: """Decompose a NIfTI affine into its rotation matrix and voxel spacing. @@ -72,69 +100,23 @@ def get_all_corner_points(affine: np.ndarray, shape: tuple[int, ...]) -> np.ndar return apply_affine(affine, lst) -def get_array(nii: Nifti1Image) -> np.ndarray: - """Extract the voxel data from a NIfTI image as a writable NumPy array. - - Args: - nii: Source NIfTI image. - - Returns: - A copy of the image data array with the original dtype preserved. - """ - return np.asanyarray(nii.dataobj, dtype=nii.dataobj.dtype).copy() # type: ignore - - -def set_array(nii: Nifti1Image, arr: np.ndarray) -> Nifti1Image: - """Return a new NIfTI image with the given array, preserving header and affine. - - If the dtype of ``arr`` differs from the existing image, the header dtype is - updated accordingly. - - Args: - nii: Source NIfTI image whose header and affine are reused. - arr: Replacement voxel data array. - - Returns: - A new :class:`Nifti1Image` backed by ``arr``. - """ - if nii.dataobj.dtype == arr.dtype: # type: ignore - nii = Nifti1Image(arr, nii.affine, nii.header) - else: - nii = Nifti1Image(get_array(nii), nii.affine, nii.header) - nii.set_data_dtype(arr.dtype) - nii = Nifti1Image(arr, nii.affine, nii.header) - return nii - - -def argmin(lst: list) -> int: - """Return the index of the minimum element in a list. - - Args: - lst: Input list of comparable elements. - - Returns: - Zero-based index of the smallest element. - """ - return lst.index(min(lst)) - - def _occupancy_bbox(occ: np.ndarray) -> tuple[np.ndarray, np.ndarray] | None: """Axis-aligned bounding box (min, max inclusive) of the nonzero region of `occ`. - Uses three 1-D `np.any` reductions rather than `np.where`, so cost is O(N) - with cache-friendly access. Returns ``None`` for an all-zero occupancy. + Thin wrapper around :func:`TPTBox.core.np_utils.np_bbox_binary` (which + shares a 2-D projection between two of the three axis reductions for 3-D + inputs, so ~2× faster than three independent ``np.any`` reductions). The + slice tuple is repacked as ``(lo, hi)`` int64 arrays because the ramp + math further down works on plain vectors, and ``None`` is returned for + an all-zero occupancy so ``_aabb_overlaps`` can treat it as "no overlap + possible". """ - mask = occ > 0 - axes = mask.ndim - lo = np.empty(axes, dtype=np.int64) - hi = np.empty(axes, dtype=np.int64) - for ax in range(axes): - collapsed = np.any(mask, axis=tuple(i for i in range(axes) if i != ax)) - idx = np.flatnonzero(collapsed) - if idx.size == 0: - return None - lo[ax] = idx[0] - hi[ax] = idx[-1] + try: + slices = np_bbox_binary(occ > 0) + except ValueError: + return None + lo = np.array([s.start for s in slices], dtype=np.int64) + hi = np.array([s.stop - 1 for s in slices], dtype=np.int64) return lo, hi @@ -153,7 +135,7 @@ def get_max_affine_and_shape( min_spacing: float | None = None, dtype: type = float, verbose: bool = False, -) -> Nifti1Image: +) -> NII: """Determine the optimal output affine and shape that encloses all input volumes. Iterates over all input affines and selects the rotation that minimises the @@ -170,7 +152,7 @@ def get_max_affine_and_shape( rotation to stdout. Returns: - A zeroed :class:`Nifti1Image` with the computed affine and shape, ready + A zeroed :class:`NII` with the computed affine and shape, ready to be used as a resampling target. Raises: @@ -219,248 +201,50 @@ def get_max_affine_and_shape( new_spacing = np.maximum(min_spacing, new_spacing) shape: np.ndarray = np.ceil(min_shape / new_spacing) - print("Choose the following spacing:", new_spacing) if verbose else None - print(f"Output shape is {shape}, which utilizes {min_possible_volume / min_volume * 100:.1f} % of all voxels.") if verbose else None + logger.on_neutral("Choose the following spacing:", new_spacing, verbose=verbose) + logger.on_neutral( + f"Output shape is {shape}, which utilizes {min_possible_volume / min_volume * 100:.1f} % of all voxels.", + verbose=verbose, + ) affine = get_ras_affine(min_rotation, new_spacing, origen[0]) - print("The new origin is ", np.round(affine[:3, 3], 2)) if verbose else None - print("The optimal rotation came from file number ", opt_id, " ", np.round(min_rotation.reshape(-1), 2)) if verbose else None - return nib.Nifti1Image(np.zeros(shape.astype(int), dtype=dtype), affine) # type: ignore - - -def compute_crop_slice(nii: Nifti1Image, minimum: float = 0, dist: int = 0) -> tuple[slice, slice, slice]: - """Compute the tight 3-D crop slice that removes empty space from a NIfTI volume. - - A voxel is considered "filled" when its value is strictly greater than - ``minimum``. The returned slice tuple can be applied via ``nii.slicer`` - to crop the image. - - Args: - nii: Input NIfTI image whose bounding box is computed. - minimum: Background threshold. Voxels above this value delimit the crop - (0 for MRI, -1024 for CT). - dist: Padding in millimetres added on every side of the crop. Converted - to voxels using the image's zooms. - - Returns: - A 3-tuple of ``slice`` objects to apply along the (X, Y, Z) axes. - - Raises: - ValueError: If no voxels exceed ``minimum`` (crop would be empty). - """ - shp = nii.shape - zms = nii.header.get_zooms() # type: ignore - d = np.around(dist / np.asarray(zms)).astype(int) - array = get_array(nii) # + minimum - msk_bin = np.zeros(array.shape, dtype=bool) - # bool_arr[array minimum] = 1 - # msk_bin = np.asanyarray(bool_arr, dtype=bool) - msk_bin[np.isnan(msk_bin)] = 0 - cor_msk = np.where(msk_bin > 0) - if cor_msk[0].shape[0] == 0: - raise ValueError("Array would be reduced to zero size") - c_min = [cor_msk[0].min(), cor_msk[1].min(), cor_msk[2].min()] - c_max = [cor_msk[0].max(), cor_msk[1].max(), cor_msk[2].max()] - x0 = max(0, c_min[0] - d[0]) - y0 = max(0, c_min[1] - d[1]) - z0 = max(0, c_min[2] - d[2]) - x1 = min(shp[0], c_max[0] + d[0]) - y1 = min(shp[1], c_max[1] + d[1]) - z1 = min(shp[2], c_max[2] + d[2]) - ex_slice = (slice(x0, x1 + 1), slice(y0, y1 + 1), slice(z0, z1 + 1)) - return ex_slice - - -def dilate_msk(msk_i_data: np.ndarray, mm: int = 5, connectivity: int = 3) -> np.ndarray: - """Dilate each label in a segmentation mask by a fixed number of voxels. - - Args: - msk_i_data: Integer-valued 3-D segmentation array. Label 0 is background. - mm: Number of dilation iterations to apply per label. - connectivity: Structuring-element connectivity (1 = face-connected, - 3 = fully-connected including diagonals). - - Returns: - A dilated ``uint8`` array of the same shape as ``msk_i_data``. - """ - from scipy.ndimage import binary_dilation, generate_binary_structure - - struct = generate_binary_structure(3, connectivity) - out = msk_i_data.copy() * 0 - for i in np.unique(msk_i_data): - if i == 0: - continue - data = msk_i_data.copy() - data[i != data] = 0 - msk_ibe_data = binary_dilation(data, structure=struct, iterations=mm) - out[out == 0] = msk_ibe_data[out == 0] - return out.astype(np.uint8) - - -def n4_bias_field_correction( - nib: Nifti1Image, - mask: np.ndarray | None = None, - threshold: int = 60, - shrink_factor: int = 4, - convergence: dict | None = None, - spline_param: int = 150, - verbose: bool = False, - weight_mask: np.ndarray | None = None, - crop: bool = False, -) -> Nifti1Image: - """Apply N4 bias-field correction to a NIfTI image using ANTsPy. - - A binary mask is derived automatically from voxels above ``threshold`` - and dilated by 3 voxels before correction is applied. - - Args: - nib: Input NIfTI image to correct. - mask: Optional pre-computed binary mask passed to ANTsPy. Overridden - when ``threshold != 0``. - threshold: Voxel intensity threshold for automatic mask generation. - Set to 0 to disable automatic masking. - shrink_factor: Image downsampling factor used inside ANTsPy to - speed up computation. - convergence: ANTsPy convergence dict with keys ``"iters"`` and - ``"tol"``. Defaults to ``{"iters": [50, 50, 50, 50], "tol": 1e-7}``. - spline_param: B-spline control point spacing for the bias field model. - verbose: If True, ANTsPy prints progress information. - weight_mask: Optional spatial weight mask passed to ANTsPy. - crop: If True, crops the corrected image to the region where the bias - field differed from the input. - - Returns: - The bias-field-corrected NIfTI image. - - Raises: - ModuleNotFoundError: If ``antspyx`` is not installed. - """ - try: - import ants - import ants.utils.bias_correction as bc # pip install antspyx==0.4.2 - from ants.utils.convert_nibabel import from_nibabel, to_nibabel - - except ModuleNotFoundError as err: - raise ModuleNotFoundError("n4 bias field correction uses ants install it with pip install antspyx==0.4.2") from err - # 5.3 or higher - import ants - import ants.ops.bias_correction as bc # pip install antspyx - - # TODO add conversion and remove this - - def from_nibabel(nib_image): - """Converts a given Nifti image into an ANTsPy image. - - Parameters - ---------- - nib_image: nibabel Nifti1Image - - Returns: - ------- - ants_image: ANTsImage - """ - ndim = nib_image.ndim - - if ndim < 3: - print("Dimensionality is less than 3.") - return None - - q_form = nib_image.get_qform() - spacing = nib_image.header["pixdim"][1 : ndim + 1] - - origin = np.zeros(ndim) - origin[:3] = q_form[:3, 3] - - direction = np.diag(np.ones(ndim)) - direction[:3, :3] = q_form[:3, :3] / spacing[:3] - - ants_img = ants.from_numpy(data=nib_image.get_fdata(), origin=origin.tolist(), spacing=spacing.tolist(), direction=direction) - - return ants_img - - def to_nibabel(img: ants.core.ants_image.ANTsImage): - try: - from nibabel.nifti1 import Nifti1Image - except ModuleNotFoundError as e: - raise ModuleNotFoundError( - "Could not import nibabel, for conversion to nibabel. Install nibabel with pip install nibabel" - ) from e - affine = get_ras_affine(rotation=img.direction, spacing=img.spacing, origin=img.origin) - return Nifti1Image(img.numpy(), affine, nib.header) - - if convergence is None: - convergence = {"iters": [50, 50, 50, 50], "tol": 1e-07} - input_ants = from_nibabel(nib) - - if threshold != 0: - mask = get_array(nib) - mask[mask < threshold] = 0 - mask[mask != 0] = 1 - mask = mask.astype(np.uint8) - mask = dilate_msk(mask, mm=3) - mask = from_nibabel(set_array(nib, mask)) - - out = bc.n4_bias_field_correction( - input_ants, - mask=mask, - shrink_factor=shrink_factor, - convergence=convergence, - spline_param=spline_param, + logger.on_neutral("The new origin is ", np.round(affine[:3, 3], 2), verbose=verbose) + logger.on_neutral( + "The optimal rotation came from file number ", + opt_id, + " ", + np.round(min_rotation.reshape(-1), 2), verbose=verbose, - weight_mask=weight_mask, ) - out_nib = to_nibabel(out) - if crop: - # Crop to regions that had a normalization applied. Removes a lot of dead space - dif = to_nibabel(input_ants - out) - da = get_array(dif) - da[da != 0] = 1 - dif = set_array(dif, da) - ex_slice = compute_crop_slice(dif) - out_nib = out_nib.slicer[ex_slice] - - return out_nib - + return NII((np.zeros(shape.astype(int), dtype=dtype), affine, None)) # type: ignore -buffer_references = {} - -def buffer_reference(path: str | Path, bias_field: bool, crop: bool = False) -> np.ndarray | Nifti1Image: +@cache +def buffer_reference(path: str | Path, bias_field: bool, crop: bool = False) -> NII: """Load (and optionally bias-correct) a NIfTI file, caching the result. - Subsequent calls with the same ``path`` return the cached result without - re-reading or re-correcting the file. + Subsequent calls with the same ``(path, bias_field, crop)`` return the + cached NII without re-reading or re-correcting the file. Caching by all + three arguments (rather than just ``path`` as the previous module-level + dict did) means a call with ``crop=True`` no longer poisons a later call + with ``crop=False`` for the same path. ``lru_cache`` is also thread-safe, + so this is safe under a ``ThreadPoolExecutor``-based dispatcher. Args: path: File path of the NIfTI image. bias_field: If True, applies N4 bias-field correction before caching. - crop: Passed to :func:`n4_bias_field_correction` when ``bias_field`` is True. + crop: Passed to :meth:`NII.n4_bias_field_correction` when ``bias_field`` + is True. Returns: - The image data array (if ``bias_field`` is False) or the corrected - :class:`Nifti1Image` (if ``bias_field`` is True). + The loaded (and optionally bias-corrected) :class:`NII`. """ - if path in buffer_references: - return buffer_references[path] - reference = n4_bias_field_correction(nib.load(path), crop) if bias_field else get_array(nib.load(path)) # type: ignore - buffer_references[path] = reference + reference = NII.load(path, False) + if bias_field: + reference = reference.n4_bias_field_correction(crop=crop) return reference -type_mapping = { - "float": float, - "uint8": np.uint8, - "uint16": np.uint16, - "uint32": np.uint32, - "uint64": np.uint64, - "int8": np.int8, - "int16": np.int16, - "int32": np.int32, - "int64": np.int64, -} - - -def _auto_output_dtype(niis: list[nib.nifti1.Nifti1Image]) -> type: +def _auto_output_dtype(niis: list[NII]) -> type: """Pick the smallest lossless output dtype from the inputs' on-disk dtypes. Reads only the NIfTI headers — no pixel data is loaded. If every input is @@ -469,13 +253,26 @@ def _auto_output_dtype(niis: list[nib.nifti1.Nifti1Image]) -> type: all; otherwise falls back to float32. This lets the stitcher store magnitude MR outputs as uint16 (or int16) when the source already fit in 16 bits, halving the on-disk and downstream RAM footprint versus float64. + + The slope check is load-bearing: Philips exports the axial VIBE-DIXON as + ``int16`` with a large ``scl_slope`` (~641 or ~2260) so the raw bytes span + the signed range but ``fdata = raw * slope`` reaches ~1e6. Returning + ``int16`` (or ``uint32`` after `_check_if_nifty_is_lying_about_its_dtype` + widens the array) and casting the blended float back to it wraps values + outside the target range into garbage — the stitched output ends up as a + static-noise pattern or an unexpectedly wide dtype (e.g. ``uint32`` for + the 6-echo mDIX magnitudes). When ANY input has slope != 1 or inter != 0 + we bail to float32 so the accumulator's true range survives the save. + ``dataobj.slope`` / ``dataobj.inter`` on the underlying ``Nifti1Image`` + give the ORIGINAL header values (before NII's dtype-widening pass), which + is what we need for this decision. """ - dtypes = [np.dtype(nii.get_data_dtype()) for nii in niis] + dtypes = [np.dtype(nii.dtype) for nii in niis] if not all(np.issubdtype(d, np.integer) for d in dtypes): return np.float32 for nii in niis: - slope = getattr(nii.dataobj, "slope", 1.0) - inter = getattr(nii.dataobj, "inter", 0.0) + slope = getattr(nii.nii.dataobj, "slope", 1.0) + inter = getattr(nii.nii.dataobj, "inter", 0.0) if slope is None or inter is None: return np.float32 if not (np.isfinite(slope) and np.isfinite(inter)): @@ -486,7 +283,7 @@ def _auto_output_dtype(niis: list[nib.nifti1.Nifti1Image]) -> type: def main( # noqa: C901 - images: list[str] | list[Path] | list[nib.nifti1.Nifti1Image], + images: list[str] | list[Path] | list[NII], output: str | None, match_histogram: bool = False, store_ramp: bool = False, @@ -503,7 +300,7 @@ def main( # noqa: C901 dtype: type | str = float, save: bool = True, ramp_path=None, -) -> tuple[Nifti1Image | None, Nifti1Image | None]: +) -> tuple[NII | None, NII | None]: """Stitch multiple overlapping NIfTI volumes into a single output volume. The algorithm: @@ -517,8 +314,8 @@ def main( # noqa: C901 5. Combines all resampled volumes with those weights and saves the result. Args: - images: Input volumes as file paths or pre-loaded :class:`Nifti1Image` - objects. At least two are required. + images: Input volumes as file paths or pre-loaded :class:`NII` (or + :class:`Nifti1Image`) objects. At least two are required. output: Output file path (``".nii.gz"`` extension is appended if absent). If None the result is returned without writing to disk (``save`` must also be False). @@ -565,6 +362,54 @@ def main( # noqa: C901 unless ``store_ramp`` is True. Returns ``(None, None)`` when fewer than two images are supplied. """ + # Suppress `_check_if_nifty_is_lying_about_its_dtype` UserWarnings for + # the whole main() call — stitching handles the "dtype mismatches + # actual range" case explicitly, and the warning is just noise here. + # See `_suppress_dtype_warning` for scope / rationale. + with _suppress_dtype_warning(): + return _main( + images, + output, + match_histogram, + store_ramp, + verbose, + min_value, + bias_field, + crop_to_bias_field, + crop_empty, + histogram, + ramp_edge_min_value, + min_spacing, + kick_out_fully_integrated_images, + is_segmentation, + dtype, + save, + ramp_path, + ) + + +def _main( # noqa: C901 + images: list[str] | list[Path] | list[NII], + output: str | None, + match_histogram: bool = False, + store_ramp: bool = False, + verbose: bool = False, + min_value: float | None = None, + bias_field: bool = True, + crop_to_bias_field: bool = False, + crop_empty: bool = False, + histogram: str | None = None, + ramp_edge_min_value: int = 5, + min_spacing: int | None = None, + kick_out_fully_integrated_images: bool = False, + is_segmentation: bool = False, + dtype: type | str = float, + save: bool = True, + ramp_path=None, +) -> tuple[NII | None, NII | None]: + """Body of :func:`main`, split out so :func:`main` can wrap it in the + dtype-warning suppression context. + """ np.set_printoptions(precision=2, floatmode="fixed") if is_segmentation: bias_field = False @@ -574,54 +419,44 @@ def main( # noqa: C901 histogram = None _bg_value: float = 0.0 if min_value is None else float(min_value) if len(images) == 0 or len(images) == 1: - print("!!! Need at least two images (-i ...nii.gz ...nii.gz) to stitch!!!\n Got " + str(images)) + logger.on_fail("Need at least two images (-i ...nii.gz ...nii.gz) to stitch. Got:", images) return None, None corners = [] affines = [] - niis: list[nib.nifti1.Nifti1Image] = [] - print("### loading ###") if verbose else None + niis: list[NII] = [] + logger.on_log("### loading ###", verbose=verbose) for f_name in images: if isinstance(f_name, (Path, str)): - print("Load ", f_name, Path(f_name)) if verbose else None + logger.on_neutral("Load ", f_name, Path(f_name), verbose=verbose) # Load NII - try: - from TPTBox.core.nii_wrapper import to_nii as _to_nii_load - - _nii_wrap = _to_nii_load(Path(f_name)) - _arr_corrected = _nii_wrap.get_array() - nii = nib.Nifti1Image(_arr_corrected, _nii_wrap.affine) - nii.set_data_dtype(_arr_corrected.dtype) - try: - nii.header.set_slope_inter(1.0, 0.0) - except Exception: # noqa: BLE001 - pass - except Exception: # noqa: BLE001 - nii: nib.nifti1.Nifti1Image = nib.load(f_name) # type: ignore - - else: + nii = to_nii(Path(f_name), seg=is_segmentation) + elif isinstance(f_name, (NII)): nii = f_name + nii.seg = is_segmentation + else: + nii = NII(f_name, seg=is_segmentation) if bias_field: - nii = n4_bias_field_correction(nii, crop=crop_to_bias_field) + nii = nii.n4_bias_field_correction(crop=crop_to_bias_field) ## Histogram equalization. if match_histogram: if histogram is None: if len(niis) == 0: reference = None else: - print("Histogram equalization with previous file") if verbose else None - reference = get_array(niis[-1]) + logger.on_neutral("Histogram equalization with previous file", verbose=verbose) + reference = niis[-1].get_array() elif histogram.isdigit(): - print("Histogram equalization", images[int(histogram)]) if verbose else None + logger.on_neutral("Histogram equalization", images[int(histogram)], verbose=verbose) reference = buffer_reference(images[int(histogram)], bias_field=bias_field, crop=crop_to_bias_field) # type: ignore else: - print("Histogram equalization with file", histogram) if verbose else None + logger.on_neutral("Histogram equalization with file", histogram, verbose=verbose) reference = buffer_reference(histogram, bias_field=bias_field, crop=crop_to_bias_field) # type: ignore if reference is not None: - image = get_array(nii) + image = nii.get_array() matched = match_histograms(image.astype(float), reference.astype(float)) matched[matched <= _bg_value] = _bg_value - nii = set_array(nii, matched) + nii = nii.set_array(matched) niis.append(nii) # Get affine and points for minimum enclosing Rectangle calculation @@ -634,9 +469,9 @@ def main( # noqa: C901 corners_current = np.concatenate(corners, axis=0) # compute output shape and affine - print("### compute output shape and affine ###") if verbose else None + logger.on_log("### compute output shape and affine ###", verbose=verbose) if is_segmentation: - max_value = max([x.get_fdata().max() for x in niis]) + max_value = max([x.max() for x in niis]) if max_value < 256: dtype2 = np.uint8 elif max_value < 256 * 256: @@ -658,24 +493,24 @@ def main( # noqa: C901 target_list = [] occupancy_list = [] # get resampled arrays and occupancy - print("### resample to new space ###") if verbose else None + logger.on_log("### resample to new space ###", verbose=verbose) for i, nii in enumerate(niis, 1): - print(f"{i:2}/{len(niis):2} resampled", end="\r") if verbose else None - nii_new = nip.resample_from_to(nii, nii_out, 0 if is_segmentation else 3, mode="constant", cval=_bg_value) - arr_new = get_array(nii_new) + logger.on_neutral(f"{i:2}/{len(niis):2} resampled", end="\r", verbose=verbose) + nii_new = nii.resample_from_to(nii_out, order=0 if is_segmentation else 3, mode="constant", c_val=_bg_value, verbose=False) + arr_new = nii_new.get_array() if not is_segmentation and np.issubdtype(arr_new.dtype, np.floating): np.nan_to_num(arr_new, copy=False, nan=_bg_value, posinf=_bg_value, neginf=_bg_value) target_list.append(arr_new) - b = nib.Nifti1Image(np.ones(nii.shape, dtype=np.float32), affine=nii.affine) # type: ignore - b = nip.resample_from_to(b, nii_new, 0, cval=0, mode="constant") + b = NII((np.ones(nii.shape, dtype=np.float32), nii.affine, None)) # type: ignore + b = b.resample_from_to(nii_new, order=0, c_val=0, mode="constant", verbose=False) if is_segmentation: x = arr_new > 0 - occupancy_list.append((get_array(b) * x.astype(np.int8)).astype(np.float32)) # Keep segmentation if other is 0 + occupancy_list.append((b.get_array() * x.astype(np.int8)).astype(np.float32)) # Keep segmentation if other is 0 else: - occupancy_list.append(get_array(b).astype(np.float32)) + occupancy_list.append(b.get_array().astype(np.float32)) - print("\n### ramp stitching ###") if verbose else None + logger.on_log("\n### ramp stitching ###", verbose=verbose) # Per-chunk axis-aligned bounding box in target-space voxel coords. # Precomputing once avoids the O(N_chunks^2) full-volume copy + multiply # inside the ramp loop for pairs that don't touch. `_occupancy_bbox` @@ -692,7 +527,7 @@ def main( # noqa: C901 _ramp_skipped_aabb = 0 _ramp_skipped_no_voxel_overlap = 0 for idx, item in enumerate(combinations, 1): - print(f"{idx:2}/{len(combinations):2} ramp stitching", end="\r") if verbose else None + logger.on_neutral(f"{idx:2}/{len(combinations):2} ramp stitching", end="\r", verbose=verbose) # Skip disjoint pairs before touching the full-volume arrays. if not _aabb_overlaps(bboxes[item[0]], bboxes[item[1]]): _ramp_skipped_aabb += 1 @@ -766,37 +601,44 @@ def main( # noqa: C901 if kick_out_fully_integrated_images: images.pop(item[1]) if (max_1 != 1 or max_2 != 1) and kick_out_fully_integrated_images: - print("kick_out_fully_integrated_images") - - print(images) + logger.on_warning("kick_out_fully_integrated_images") + logger.on_warning(images) + # Pass EVERY argument by keyword — the positional form used + # to silently drop `is_segmentation`, `dtype`, `ramp_path` and + # shift `save` onto `is_segmentation`, which flipped the + # recursion into the segmentation code path and produced an + # unexpected uint dtype for what was actually magnitude MR. return main( - images, - output, - match_histogram, - store_ramp, - verbose, - min_value, - bias_field, - crop_to_bias_field, - crop_empty, - histogram, - ramp_edge_min_value, - min_spacing, - kick_out_fully_integrated_images, - save, + images=images, + output=output, + match_histogram=match_histogram, + store_ramp=store_ramp, + verbose=verbose, + min_value=min_value, + bias_field=bias_field, + crop_to_bias_field=crop_to_bias_field, + crop_empty=crop_empty, + histogram=histogram, + ramp_edge_min_value=ramp_edge_min_value, + min_spacing=min_spacing, + kick_out_fully_integrated_images=kick_out_fully_integrated_images, + is_segmentation=is_segmentation, + dtype=dtype, + save=save, + ramp_path=ramp_path, ) arr_1_full[sub] = arr_1_sub.astype(arr_1_full.dtype, copy=False) arr_2_full[sub] = arr_2_sub.astype(arr_2_full.dtype, copy=False) else: _ramp_skipped_no_voxel_overlap += 1 continue - if verbose: - print( - f"\nramp summary: {_ramp_done} computed, " - f"{_ramp_skipped_aabb} skipped (disjoint AABB), " - f"{_ramp_skipped_no_voxel_overlap} skipped (no voxel overlap) " - f"of {len(combinations)} pairs" - ) + logger.on_neutral( + f"\nramp summary: {_ramp_done} computed, " + f"{_ramp_skipped_aabb} skipped (disjoint AABB), " + f"{_ramp_skipped_no_voxel_overlap} skipped (no voxel overlap) " + f"of {len(combinations)} pairs", + verbose=verbose, + ) # Aggregate: in-place accumulate `t * occupancy` per chunk instead of # `np.stack(target_list) * np.stack(occupancy_list)` which would peak at # ~2 × N × volume of temporary float arrays. @@ -824,7 +666,7 @@ def main( # noqa: C901 _floor = None if _floor is not None: target_arr[target_arr <= _floor] = _floor - print("\n### Save ###") if verbose else None + logger.on_log("\n### Save ###", verbose=verbose) if output is not None: output = str(output) if not output.endswith(".nii.gz"): @@ -832,10 +674,15 @@ def main( # noqa: C901 if "/" not in output and "\\" not in output: assert isinstance(images[0], (str, Path)), "automatic path fetching only possible if images are strings or Path, not objects" output = str(Path(Path(images[0]).parent, output)) - dtype = type_mapping.get(dtype, dtype) # type: ignore - nii_out = set_array(nii_out, target_arr.astype(dtype)) + # Accept a string dtype name ("uint8", "float", …) from the CLI argparse + # path. `np.dtype(...)` covers every alias `type_mapping` used to define + # explicitly. Non-string dtypes (types picked by `_auto_output_dtype` or + # passed in directly) pass through unchanged. + if isinstance(dtype, str): + dtype = np.dtype(dtype).type + nii_out = nii_out.set_array(target_arr.astype(dtype)) if bias_field: - nii_out = n4_bias_field_correction(nii_out) + nii_out = nii_out.n4_bias_field_correction() if crop_empty: # Crop to the union of per-chunk occupancy AABBs. The previous path # went through compute_crop_slice on a 4-D (N, X, Y, Z) stack, which @@ -850,31 +697,31 @@ def main( # noqa: C901 slice(int(lo[1]), int(hi[1]) + 1), slice(int(lo[2]), int(hi[2]) + 1), ) - nii_out = nii_out.slicer[ex_slice] + nii_out = nii_out[ex_slice] else: ex_slice = () else: ex_slice = () - nii_out.set_data_dtype(dtype) + nii_out.set_dtype_(dtype) if save: - nib.save(nii_out, output) # type: ignore - print("Saved ", output) if verbose else None + nii_out.save(output) # type: ignore + logger.on_save("Saved ", output, verbose=verbose) if store_ramp: occupancy_arr = np.stack(occupancy_list, -1) if crop_empty: occupancy_arr = occupancy_arr[ex_slice] assert output is not None - nii_occ = set_array(nii_out, occupancy_arr) - nii_occ.set_data_dtype(np.int8) + nii_occ = nii_out.set_array(occupancy_arr) + nii_occ.set_dtype_(np.int8) output = output.replace(".nii.gz", "_ramps.nii.gz").replace("_msk_", "_") if ramp_path is None else ramp_path if save: - nib.save(nii_occ, output) # type: ignore - print("Saved ", output) if verbose else None + nii_occ.save(output) # type: ignore + logger.on_save("Saved ", output, verbose=verbose) return nii_out, nii_occ - print("\n### Finished ###") if verbose else None + logger.on_ok("\n### Finished ###", verbose=verbose) return nii_out, None From 07c61b628cc171c2716007f05c5812a5de158c3b Mon Sep 17 00:00:00 2001 From: robert Date: Fri, 25 Sep 2026 10:02:15 +0200 Subject: [PATCH 5/6] style: fix ruff D205 and RUF003 in stitching.py - Reflow docstring to satisfy D205 (blank line between summary and body) - Replace ambiguous MULTIPLICATION SIGN with ASCII '*' in comment Co-Authored-By: Claude Opus 4.7 --- TPTBox/stitching/stitching.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/TPTBox/stitching/stitching.py b/TPTBox/stitching/stitching.py index e69e011..0c6d65a 100755 --- a/TPTBox/stitching/stitching.py +++ b/TPTBox/stitching/stitching.py @@ -407,8 +407,9 @@ def _main( # noqa: C901 save: bool = True, ramp_path=None, ) -> tuple[NII | None, NII | None]: - """Body of :func:`main`, split out so :func:`main` can wrap it in the - dtype-warning suppression context. + """Body of :func:`main`, split out for dtype-warning context wrapping. + + Split out so :func:`main` can wrap it in the dtype-warning suppression context. """ np.set_printoptions(precision=2, floatmode="fixed") if is_segmentation: @@ -641,7 +642,7 @@ def _main( # noqa: C901 ) # Aggregate: in-place accumulate `t * occupancy` per chunk instead of # `np.stack(target_list) * np.stack(occupancy_list)` which would peak at - # ~2 × N × volume of temporary float arrays. + # ~2 * N * volume of temporary float arrays. target_arr = np.zeros(target_list[0].shape, dtype=dtype2) if is_segmentation: for t, o in zip(target_list, occupancy_list): From 2350a5ab136ef5d882ac785a50386597a3e323a0 Mon Sep 17 00:00:00 2001 From: robert Date: Fri, 25 Sep 2026 10:23:36 +0200 Subject: [PATCH 6/6] add test and update documentaiton --- TPTBox/stitching/README.md | 17 +++++++++-------- unit_tests/test_stiching.py | 2 +- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/TPTBox/stitching/README.md b/TPTBox/stitching/README.md index 57bab4e..b4ffbfc 100644 --- a/TPTBox/stitching/README.md +++ b/TPTBox/stitching/README.md @@ -8,8 +8,8 @@ You can verify alignment by opening the images in ITKSnap with "open additional | Function | Description | |---|---| -| `stitching(nii_list, out, ...)` | Stitch a list of `NII` objects; returns `(result_nii, ramp_nii)` | -| `stitching_raw(paths, out, ...)` | Stitch from file paths directly | +| `stitching(inputs, out, ...)` | High-level wrapper. `inputs` accepts any mix of `BIDS_FILE`, `NII`, `str`, or `Path`. Resolves the output path from a `BIDS_FILE` when one is passed. Returns `(result_nii, ramp_nii)` as `NII` objects. | +| `stitching_raw(images, out, ...)` | Low-level driver. `images` accepts file paths, pre-loaded `NII` objects, or (fallback) `Nifti1Image` objects. Returns `(result_nii, ramp_nii)` as `NII` objects. | | `NAKO_stitch_T2w(HWS, BWS, LWS, n4_after_stitch=False)` | Stitch the three NAKO sagittal T2w spine stations (HWS cervical, BWS thoracic, LWS lumbar) into one volume | ![Example of a stitching](https://raw.githubusercontent.com/Hendrik-code/TPTBox/main/TPTBox/stitching/stitching.jpg "Example of a stitching") @@ -24,7 +24,7 @@ stitching.py [-i IMAGES [IMAGES ...]] a list of input image paths [-o OUTPUT] The output image path [-v] verbose - if set, there will be more printouts. -[-min_value MIN_VALUE] New pixels not present will get this value. Recommended 0 for MRI and for CT -1024 or the known min-value. +[-min_value MIN_VALUE] Background fill used when resampling each chunk, and — when set explicitly — a hard floor applied to the stitched output. Pass 0 for MRI magnitude, -1024 for CT. Omitting it (the Python-API default `None`) uses 0 as the internal background and only applies a hard floor when the output dtype cannot represent negatives (unsigned integer types) — signed / float outputs then keep legitimate negatives (Philips-scaled fat-fraction, phase, B0 offsets). [-seg] This flag is required if you merge segmentation Niftis. Switches: [-no_bias] If set: Do not use n4_bias_field_correction. It speeds up the process, but n4_bias_field_correction helps in roughly aligning the histogram. @@ -67,14 +67,15 @@ list_of_files = [ # Call the stitching function # This will combine your images into a single NIfTI file stitching( - list_of_files, # List of input files - out="out_path_stitched_image.nii.gz", # Path to save stitched output - is_seg=False, # Set True if these are segmentation masks - is_ct=False, # True for CT, min_value will by -1024 instead of 0 + list_of_files, # BIDS_FILE / NII / str / Path (any mix) + out="out_path_stitched_image.nii.gz", # Path or BIDS_FILE for the stitched output + is_seg=False, # Set True for segmentation masks (forces integer dtype, nearest-neighbour resample) + is_ct=False, # Sets min_value to -1024 (CT air) when min_value is not passed explicitly kick_out_fully_integrated_images=True, - dtype=float, # Data type of the output image + dtype=float, # Output dtype; "auto" picks the smallest lossless type from the inputs (float32 when any input has a non-trivial scl_slope/inter) match_histogram=False, # Match intensity histograms across images store_ramp=False, # Store blending ramp (optional) + min_value=None, # Explicit background/floor. None (default) = clip only when the output dtype is unsigned integer (protects against negative-to-huge wraparound) and let signed / float outputs keep legitimate negatives. Pass 0 to floor magnitude MR at 0 so cubic-spline resample ringing doesn't leak small negatives into what should be a non-negative volume; pass -1024 for CT. ) ``` diff --git a/unit_tests/test_stiching.py b/unit_tests/test_stiching.py index 75f5aaf..f987754 100755 --- a/unit_tests/test_stiching.py +++ b/unit_tests/test_stiching.py @@ -120,7 +120,7 @@ def test_stitching( self.assertTrue(output.exists(), output) output.unlink(missing_ok=True) # Assertions - self.assertIsInstance(result, nib.Nifti1Image) # Check if result is a Nifti1Image instance + self.assertIsInstance(result, NII) # Check the result is a NII (post-migration return type) # Add more assertions based on your requirements def test_stitching2(self):