From 449ebcd87f5e631505b94500d0dc98fa0d2fe30c Mon Sep 17 00:00:00 2001 From: gu53lif Date: Wed, 30 Sep 2026 10:19:46 +0200 Subject: [PATCH 1/2] bug fix _rib_assingn --- TPTBox/segmentation/rib/_rib_assign.py | 133 +++++++++++++++++++------ 1 file changed, 104 insertions(+), 29 deletions(-) diff --git a/TPTBox/segmentation/rib/_rib_assign.py b/TPTBox/segmentation/rib/_rib_assign.py index 426880d..39025bc 100644 --- a/TPTBox/segmentation/rib/_rib_assign.py +++ b/TPTBox/segmentation/rib/_rib_assign.py @@ -47,7 +47,9 @@ def _thoracic_like_labels(labels: Iterable[int]) -> list[int]: return [*sorted(v for v in labels if 7 <= int(v) <= 20), 28] -def _split_ccs_by_side(rib_cc: NII, cms_cc: POI, cms_cc2: POI, cms_vert: POI) -> tuple[list[RibCandidate], list[RibCandidate]]: +def _split_ccs_by_side( + rib_cc: NII, cms_cc: POI, cms_cc2: POI, cms_vert: POI +) -> tuple[list[RibCandidate], list[RibCandidate]]: right_axis = cms_vert.get_axis("R") inf_axis = cms_vert.get_axis("I") # ref_x = cms_vert[ref_vert, 50][right_axis] @@ -59,7 +61,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 @@ -74,14 +81,20 @@ def _split_ccs_by_side(rib_cc: NII, cms_cc: POI, cms_cc2: POI, cms_vert: POI) -> return left, right -def _touching_surface(mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_weights=None) -> float: +def _touching_surface( + mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_weights=None +) -> float: """Return weighted number of voxel faces shared by two masks. Touching along ``up_down_axis`` is weighted by ``up_down_weight``. 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): @@ -92,7 +105,8 @@ def _touching_surface(mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_wei sl2[axis] = slice(1, None) touching = ( - np.logical_and(mask_a[tuple(sl1)], mask_b[tuple(sl2)]).sum() + np.logical_and(mask_b[tuple(sl1)], mask_a[tuple(sl2)]).sum() + np.logical_and(mask_a[tuple(sl1)], mask_b[tuple(sl2)]).sum() + + np.logical_and(mask_b[tuple(sl1)], mask_a[tuple(sl2)]).sum() ) weight = axis_weights.get(axis, 1) @@ -102,7 +116,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. @@ -123,9 +142,13 @@ def _try_erosion_split( elif _pass == 1: eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False) elif _pass == 2: - eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False, ignore_direction="R") + eroded = cc_nii.erode_msk( + n_pixel=erosion_pixels, verbose=False, ignore_direction="R" + ) elif _pass == 3: - eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False, ignore_direction="A") + eroded = cc_nii.erode_msk( + n_pixel=erosion_pixels, verbose=False, ignore_direction="A" + ) else: eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False) except Exception as e: @@ -144,7 +167,9 @@ def _try_erosion_split( for sub in sub_labels.copy(): count = int((infected_arr == sub).sum()) - should_merge = count * cc_nii.voxel_volume() < min_volume or count / full_size < 0.2 + should_merge = ( + count * cc_nii.voxel_volume() < min_volume or count / full_size < 0.2 + ) if not should_merge: continue # print("merge", count, full_size, count / full_size) @@ -157,7 +182,12 @@ 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) @@ -219,7 +249,9 @@ def split_touching_rib_ccs( next_label = int(rib_cc.max()) + 1 if rib_cc.shape and rib_cc.max() > 0 else 1 u = {int(x) for x in rib_cc.unique() if x != 0} # ribs counte twice (left/right) - expected_number_of_ccs = len([a for a in Vertebra_Instance.thoracic() if a.value in vert_ids]) * 2 + expected_number_of_ccs = ( + len([a for a in Vertebra_Instance.thoracic() if a.value in vert_ids]) * 2 + ) if expected_number_of_ccs == 0: expected_number_of_ccs = 24 @@ -231,8 +263,12 @@ def split_touching_rib_ccs( binary_ccs = [(cc_label, rib_cc.extract_label(cc_label)) for cc_label in labels] if num_workers == 1: results = [] - for cc_label, binary_cc in tqdm(binary_ccs, total=len(binary_ccs), desc="Separate RIBs"): - result = _try_erosion_split(cc_label, binary_cc, erosion_pixels, min_volume, _pass) + for cc_label, binary_cc in tqdm( + binary_ccs, total=len(binary_ccs), desc="Separate RIBs" + ): + result = _try_erosion_split( + cc_label, binary_cc, erosion_pixels, min_volume, _pass + ) if result is not None: results.append(result) else: @@ -240,12 +276,21 @@ 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 ] # tqdm advances immediately whenever an individual future completes. - for future in tqdm(as_completed(futures), total=len(futures), desc="Separate RIBs"): + for future in tqdm( + as_completed(futures), total=len(futures), desc="Separate RIBs" + ): result = future.result() if result is not None: results.append(result) @@ -268,7 +313,9 @@ def split_touching_rib_ccs( u.add(new_lbl) if verbose: - logger.print(f"split CC {cc_label} via erosion -> {new_lbl} (voxels={int(m.sum())})") + logger.print( + f"split CC {cc_label} via erosion -> {new_lbl} (voxels={int(m.sum())})" + ) if short_cut and rib_cc.max() >= expected_number_of_ccs: break @@ -338,18 +385,29 @@ def assign_ribs_to_vert_segmentation( vert_seg = vert_seg.reorient() sem_seg = sem_seg.reorient() rib_seg = rib_seg.reorient(verbose=verbose) - vert_pred = vert_seg.extract_label([Vertebra_Instance.C7, *Vertebra_Instance.thoracic(), Vertebra_Instance.L1]) + vert_pred = vert_seg.extract_label( + [Vertebra_Instance.C7, *Vertebra_Instance.thoracic(), Vertebra_Instance.L1] + ) rib_seg = rib_seg.extract_label([left_id, right_id], keep_label=True) # type: ignore logger.on_debug(f"{rib_seg.unique()=}") # Remove rib voxels overlapping vertebrae rib_seg[vert_pred != 0] = 0 - rib_cc = rib_seg.filter_connected_components(None, min_volume=min_volume, keep_label=False) + rib_cc = rib_seg.filter_connected_components( + None, min_volume=min_volume, keep_label=False + ) 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_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) vert_labels = _thoracic_like_labels(vert_seg.unique()) @@ -373,12 +431,17 @@ def assign_ribs_to_vert_segmentation( assert isinstance(vert_labels, (tuple, list)), type(vert_labels) if vert_labels and ( - (vert_labels[0] != 28 and vert_labels[1] != 28) or ((vert_labels[0], 50) in vert_labels and (vert_labels[1], 50) in vert_labels) + (vert_labels[0] != 28 and vert_labels[1] != 28) + or ((vert_labels[0], 50) in vert_labels and (vert_labels[1], 50) in vert_labels) ): inf_axis = cms_vert.get_axis("I") t1_z = cms_vert[vert_labels[0], 50][inf_axis] # use median vertebral spacing as adaptive threshold - spacing = abs(cms_vert[vert_labels[1], 50][inf_axis] - t1_z) if len(vert_labels) > 1 else 30 + spacing = ( + abs(cms_vert[vert_labels[1], 50][inf_axis] - t1_z) + if len(vert_labels) > 1 + else 30 + ) threshold = max(5, spacing * 0.6) nearest_left = abs(left_ccs[0].z - t1_z) if left_ccs else float("inf") @@ -405,20 +468,32 @@ 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 + 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 undefined = sum(v == error_value for v in rib_vert_map.values()) logger.print(f"Unmatched rib CCs: {undefined}") if add_error: # or split_touching: vert_seg[rib_inst == error_value] = error_value if np.any(vert_seg.get_seg_array() == error_value): - vert_seg2 = vert_seg.remove_labels(error_value).infect(vert_seg.extract_label(error_value), verbose=False) + 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(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 return vert_seg.reorient_(ori), sem_seg.reorient_(ori) From 41e4a15aa7e9ede360e24813e1e8a6c6aed68253 Mon Sep 17 00:00:00 2001 From: TanjaLe <105376287+TanjaLe@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:21:34 +0000 Subject: [PATCH 2/2] style fixes by ruff --- TPTBox/segmentation/rib/_rib_assign.py | 84 +++++++------------------- 1 file changed, 21 insertions(+), 63 deletions(-) diff --git a/TPTBox/segmentation/rib/_rib_assign.py b/TPTBox/segmentation/rib/_rib_assign.py index 39025bc..44302c9 100644 --- a/TPTBox/segmentation/rib/_rib_assign.py +++ b/TPTBox/segmentation/rib/_rib_assign.py @@ -47,9 +47,7 @@ def _thoracic_like_labels(labels: Iterable[int]) -> list[int]: return [*sorted(v for v in labels if 7 <= int(v) <= 20), 28] -def _split_ccs_by_side( - rib_cc: NII, cms_cc: POI, cms_cc2: POI, cms_vert: POI -) -> tuple[list[RibCandidate], list[RibCandidate]]: +def _split_ccs_by_side(rib_cc: NII, cms_cc: POI, cms_cc2: POI, cms_vert: POI) -> tuple[list[RibCandidate], list[RibCandidate]]: right_axis = cms_vert.get_axis("R") inf_axis = cms_vert.get_axis("I") # ref_x = cms_vert[ref_vert, 50][right_axis] @@ -81,9 +79,7 @@ def _split_ccs_by_side( return left, right -def _touching_surface( - mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_weights=None -) -> float: +def _touching_surface(mask_a: np.ndarray, mask_b: np.ndarray, nii: NII, axis_weights=None) -> float: """Return weighted number of voxel faces shared by two masks. Touching along ``up_down_axis`` is weighted by ``up_down_weight``. @@ -105,8 +101,7 @@ def _touching_surface( sl2[axis] = slice(1, None) touching = ( - np.logical_and(mask_a[tuple(sl1)], mask_b[tuple(sl2)]).sum() - + np.logical_and(mask_b[tuple(sl1)], mask_a[tuple(sl2)]).sum() + np.logical_and(mask_a[tuple(sl1)], mask_b[tuple(sl2)]).sum() + np.logical_and(mask_b[tuple(sl1)], mask_a[tuple(sl2)]).sum() ) weight = axis_weights.get(axis, 1) @@ -142,13 +137,9 @@ def _try_erosion_split( elif _pass == 1: eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False) elif _pass == 2: - eroded = cc_nii.erode_msk( - n_pixel=erosion_pixels, verbose=False, ignore_direction="R" - ) + eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False, ignore_direction="R") elif _pass == 3: - eroded = cc_nii.erode_msk( - n_pixel=erosion_pixels, verbose=False, ignore_direction="A" - ) + eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False, ignore_direction="A") else: eroded = cc_nii.erode_msk(n_pixel=erosion_pixels, verbose=False) except Exception as e: @@ -167,9 +158,7 @@ def _try_erosion_split( for sub in sub_labels.copy(): count = int((infected_arr == sub).sum()) - should_merge = ( - count * cc_nii.voxel_volume() < min_volume or count / full_size < 0.2 - ) + should_merge = count * cc_nii.voxel_volume() < min_volume or count / full_size < 0.2 if not should_merge: continue # print("merge", count, full_size, count / full_size) @@ -184,9 +173,7 @@ def _try_erosion_split( else: target = max( remaining, - key=lambda candidate: _touching_surface( - sub_mask, infected_arr == candidate, infected - ), + key=lambda candidate: _touching_surface(sub_mask, infected_arr == candidate, infected), ) # print("merge of many", target) infected_arr[infected_arr == sub] = target @@ -249,9 +236,7 @@ def split_touching_rib_ccs( next_label = int(rib_cc.max()) + 1 if rib_cc.shape and rib_cc.max() > 0 else 1 u = {int(x) for x in rib_cc.unique() if x != 0} # ribs counte twice (left/right) - expected_number_of_ccs = ( - len([a for a in Vertebra_Instance.thoracic() if a.value in vert_ids]) * 2 - ) + expected_number_of_ccs = len([a for a in Vertebra_Instance.thoracic() if a.value in vert_ids]) * 2 if expected_number_of_ccs == 0: expected_number_of_ccs = 24 @@ -263,12 +248,8 @@ def split_touching_rib_ccs( binary_ccs = [(cc_label, rib_cc.extract_label(cc_label)) for cc_label in labels] if num_workers == 1: results = [] - for cc_label, binary_cc in tqdm( - binary_ccs, total=len(binary_ccs), desc="Separate RIBs" - ): - result = _try_erosion_split( - cc_label, binary_cc, erosion_pixels, min_volume, _pass - ) + for cc_label, binary_cc in tqdm(binary_ccs, total=len(binary_ccs), desc="Separate RIBs"): + result = _try_erosion_split(cc_label, binary_cc, erosion_pixels, min_volume, _pass) if result is not None: results.append(result) else: @@ -288,9 +269,7 @@ def split_touching_rib_ccs( ] # tqdm advances immediately whenever an individual future completes. - for future in tqdm( - as_completed(futures), total=len(futures), desc="Separate RIBs" - ): + for future in tqdm(as_completed(futures), total=len(futures), desc="Separate RIBs"): result = future.result() if result is not None: results.append(result) @@ -313,9 +292,7 @@ def split_touching_rib_ccs( u.add(new_lbl) if verbose: - logger.print( - f"split CC {cc_label} via erosion -> {new_lbl} (voxels={int(m.sum())})" - ) + logger.print(f"split CC {cc_label} via erosion -> {new_lbl} (voxels={int(m.sum())})") if short_cut and rib_cc.max() >= expected_number_of_ccs: break @@ -385,16 +362,12 @@ def assign_ribs_to_vert_segmentation( vert_seg = vert_seg.reorient() sem_seg = sem_seg.reorient() rib_seg = rib_seg.reorient(verbose=verbose) - vert_pred = vert_seg.extract_label( - [Vertebra_Instance.C7, *Vertebra_Instance.thoracic(), Vertebra_Instance.L1] - ) + vert_pred = vert_seg.extract_label([Vertebra_Instance.C7, *Vertebra_Instance.thoracic(), Vertebra_Instance.L1]) rib_seg = rib_seg.extract_label([left_id, right_id], keep_label=True) # type: ignore logger.on_debug(f"{rib_seg.unique()=}") # Remove rib voxels overlapping vertebrae rib_seg[vert_pred != 0] = 0 - rib_cc = rib_seg.filter_connected_components( - None, min_volume=min_volume, keep_label=False - ) + rib_cc = rib_seg.filter_connected_components(None, min_volume=min_volume, keep_label=False) cms_vert = calc_centroids(vert_seg) if split_touching: rib_cc = split_touching_rib_ccs( @@ -405,9 +378,7 @@ def assign_ribs_to_vert_segmentation( short_cut=short_cut, verbose=verbose, ) - cms_cc = calc_centroids( - rib_cc * vert_pred.calc_convex_hull(None).dilate_msk_euclid(5) - ) + 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) vert_labels = _thoracic_like_labels(vert_seg.unique()) @@ -431,17 +402,12 @@ def assign_ribs_to_vert_segmentation( assert isinstance(vert_labels, (tuple, list)), type(vert_labels) if vert_labels and ( - (vert_labels[0] != 28 and vert_labels[1] != 28) - or ((vert_labels[0], 50) in vert_labels and (vert_labels[1], 50) in vert_labels) + (vert_labels[0] != 28 and vert_labels[1] != 28) or ((vert_labels[0], 50) in vert_labels and (vert_labels[1], 50) in vert_labels) ): inf_axis = cms_vert.get_axis("I") t1_z = cms_vert[vert_labels[0], 50][inf_axis] # use median vertebral spacing as adaptive threshold - spacing = ( - abs(cms_vert[vert_labels[1], 50][inf_axis] - t1_z) - if len(vert_labels) > 1 - else 30 - ) + spacing = abs(cms_vert[vert_labels[1], 50][inf_axis] - t1_z) if len(vert_labels) > 1 else 30 threshold = max(5, spacing * 0.6) nearest_left = abs(left_ccs[0].z - t1_z) if left_ccs else float("inf") @@ -472,28 +438,20 @@ def assign_ribs_to_vert_segmentation( 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 + 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 undefined = sum(v == error_value for v in rib_vert_map.values()) logger.print(f"Unmatched rib CCs: {undefined}") if add_error: # or split_touching: vert_seg[rib_inst == error_value] = error_value if np.any(vert_seg.get_seg_array() == error_value): - vert_seg2 = vert_seg.remove_labels(error_value).infect( - vert_seg.extract_label(error_value), verbose=False - ) + 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(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, + sem_seg.extract_label([Location.Rib_Left.value, Location.Rib_Right.value]) == 1, ) ] = error_value return vert_seg.reorient_(ori), sem_seg.reorient_(ori)