diff --git a/TPTBox/segmentation/rib/_rib_assign.py b/TPTBox/segmentation/rib/_rib_assign.py index 426880d..44302c9 100644 --- a/TPTBox/segmentation/rib/_rib_assign.py +++ b/TPTBox/segmentation/rib/_rib_assign.py @@ -59,7 +59,12 @@ def _split_ccs_by_side(rib_cc: NII, cms_cc: POI, cms_cc2: POI, cms_vert: POI) -> continue center = cms_cc[cc, 50] center2 = cms_cc2[cc, 50] if (cc, 50) in cms_cc2 else center # noqa: SIM401 - cand = RibCandidate(cc_label=cc, volume=vols.get(cc, 0), z=center[inf_axis], x=center2[right_axis]) + cand = RibCandidate( + cc_label=cc, + volume=vols.get(cc, 0), + z=center[inf_axis], + x=center2[right_axis], + ) distances = cms_vert.calculate_distances_cord(center2) min_key = min(distances, key=distances.get) # type: ignore # x < vertebra center => patient right in RAS-like orientation @@ -81,7 +86,11 @@ def _touching_surface(mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_wei All other axes have weight 1.0. """ if axis_weights is None: - axis_weights = {nii.get_axis("S"): 0.01, nii.get_axis("A"): 0.1, nii.get_axis("R"): 1} + axis_weights = { + nii.get_axis("S"): 0.01, + nii.get_axis("A"): 0.1, + nii.get_axis("R"): 1, + } surface = 0.0 for axis in range(mask_a.ndim): @@ -102,7 +111,12 @@ def _touching_surface(mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_wei def _try_erosion_split( - cc_label: int, binary_cc: NII, erosion_pixels: int, min_volume: int, _pass=0, verbose=True + cc_label: int, + binary_cc: NII, + erosion_pixels: int, + min_volume: int, + _pass=0, + verbose=True, ) -> tuple[int, list[np.ndarray]] | None: """Try to split one CC by erosion. @@ -157,7 +171,10 @@ def _try_erosion_split( target = remaining[0] # print("merge", target) else: - target = max(remaining, key=lambda candidate: _touching_surface(sub_mask, infected_arr == candidate, infected)) + target = max( + remaining, + key=lambda candidate: _touching_surface(sub_mask, infected_arr == candidate, infected), + ) # print("merge of many", target) infected_arr[infected_arr == sub] = target sub_labels.remove(sub) @@ -240,7 +257,14 @@ def split_touching_rib_ccs( with ThreadPoolExecutor(max_workers=num_workers) as executor: futures = [ - executor.submit(_try_erosion_split, cc_label, binary_cc, erosion_pixels, min_volume, _pass) + executor.submit( + _try_erosion_split, + cc_label, + binary_cc, + erosion_pixels, + min_volume, + _pass, + ) for cc_label, binary_cc in binary_ccs ] @@ -347,7 +371,12 @@ def assign_ribs_to_vert_segmentation( cms_vert = calc_centroids(vert_seg) if split_touching: rib_cc = split_touching_rib_ccs( - rib_cc, erosion_pixels=erosion_pixels, min_volume=min_volume, vert_ids=vert_ids, short_cut=short_cut, verbose=verbose + rib_cc, + erosion_pixels=erosion_pixels, + min_volume=min_volume, + vert_ids=vert_ids, + short_cut=short_cut, + verbose=verbose, ) cms_cc = calc_centroids(rib_cc * vert_pred.calc_convex_hull(None).dilate_msk_euclid(5)) cms_cc2 = calc_centroids(rib_cc) # .dilate_msk_euclid(5) @@ -405,7 +434,8 @@ def assign_ribs_to_vert_segmentation( logger.on_debug(f"{rib_inst.unique()=}") # Merge rib assignments back into the original vert / sem segmentations # without disturbing existing (non-rib) labels. Skip unmatched CCs (sentinel error_value=255). - matched = (rib_inst != 0) & (rib_inst != error_value) + print("rib_inst", rib_inst.dtype, rib_inst.seg) + matched = np.logical_and(rib_inst != 0, rib_inst != error_value) vert_seg[matched] = rib_inst[matched] sem_seg[rib_seg != 0] = rib_seg.map_labels({left_id: Location.Rib_Left.value, right_id: Location.Rib_Right.value})[rib_seg != 0] # type: ignore @@ -418,7 +448,10 @@ def assign_ribs_to_vert_segmentation( vert_seg2 = vert_seg.remove_labels(error_value).infect(vert_seg.extract_label(error_value), verbose=False) vert_seg[vert_seg != vert_seg2] = vert_seg2[vert_seg != vert_seg2] vert_seg[np.logical_and(rib_inst == error_value, vert_seg == 0)] = error_value - vert_seg[np.logical_and(vert_seg == 0, sem_seg.extract_label([Location.Rib_Left.value, Location.Rib_Right.value] == 1))] = ( - error_value - ) + vert_seg[ + np.logical_and( + vert_seg == 0, + sem_seg.extract_label([Location.Rib_Left.value, Location.Rib_Right.value]) == 1, + ) + ] = error_value return vert_seg.reorient_(ori), sem_seg.reorient_(ori)