Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 43 additions & 10 deletions TPTBox/segmentation/rib/_rib_assign.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -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.

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
]

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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)