[BUG]: Cannot register parcellation with non-continuous values #257

Merged
synchon merged 9 commits from update/non-cont-parcel-support into main 2025-10-07 10:05:31 +00:00
7 changed files with 71 additions and 84 deletions

View file

@ -0,0 +1 @@
``ParcellationRegistry.get`` now returns labels as dictionary mapping from values to labels instead of a list of labels by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Allow parcellation registration with non-continuous ranges as values by `Synchon Mandal`_

View file

@ -120,7 +120,8 @@ def get_data(
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> Union[ ) -> Union[
tuple[ArrayLike, list[str]], # coordinates tuple[ArrayLike, list[str]], # coordinates
tuple["Nifti1Image", list[str]], # parcellation / maps tuple["Nifti1Image", dict[int, str]], # parcellation
tuple["Nifti1Image", list[str]], # maps
"Nifti1Image", # mask "Nifti1Image", # mask
]: ]:
"""Get tailored ``kind`` for ``target_data``. """Get tailored ``kind`` for ``target_data``.
@ -142,7 +143,7 @@ def get_data(
Returns Returns
------- -------
tuple of numpy.ndarray, list of str; \ tuple of numpy.ndarray, list of str; \
tuple of nibabel.nifti1.Nifti1Image, list of str; \ tuple of nibabel.nifti1.Nifti1Image, dict or list of str; \
nibabel.nifti1.Nifti1Image nibabel.nifti1.Nifti1Image
Raises Raises

View file

@ -403,7 +403,7 @@ class ParcellationRegistry(BasePipelineDataRegistry):
if not path_only: if not path_only:
# Load image via nibabel # Load image via nibabel
parcellation_img = nib.load(parcellation_fname) parcellation_img = nib.load(parcellation_fname)
# Get unique values # Get unique values (returns sorted result)
parcel_values = np.unique(parcellation_img.get_fdata()) parcel_values = np.unique(parcellation_img.get_fdata())
# Check for dimension # Check for dimension
if len(parcel_values) - 1 != len(parcellation_labels): if len(parcel_values) - 1 != len(parcellation_labels):
@ -411,14 +411,6 @@ class ParcellationRegistry(BasePipelineDataRegistry):
f"Parcellation {name} has {len(parcel_values) - 1} " f"Parcellation {name} has {len(parcel_values) - 1} "
f"parcels but {len(parcellation_labels)} labels." f"parcels but {len(parcellation_labels)} labels."
) )
# Sort values
parcel_values.sort()
# Check if value range is invalid
if np.any(np.diff(parcel_values) != 1):
raise_error(
f"Parcellation {name} must have all the values in the "
f"range [0, {len(parcel_values)}]"
)
return parcellation_img, parcellation_labels, parcellation_fname, space return parcellation_img, parcellation_labels, parcellation_fname, space
@ -427,7 +419,7 @@ class ParcellationRegistry(BasePipelineDataRegistry):
parcellations: Union[str, list[str]], parcellations: Union[str, list[str]],
target_data: dict[str, Any], target_data: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple["Nifti1Image", list[str]]: ) -> tuple["Nifti1Image", dict[int, str]]:
"""Get parcellation, tailored for the target image. """Get parcellation, tailored for the target image.
Parameters Parameters
@ -446,8 +438,8 @@ class ParcellationRegistry(BasePipelineDataRegistry):
------- -------
Nifti1Image Nifti1Image
The parcellation image. The parcellation image.
list of str dict
Parcellation labels. Parcellation value to label mappings.
Raises Raises
------ ------
@ -561,7 +553,16 @@ class ParcellationRegistry(BasePipelineDataRegistry):
# Avoid merging if there is only one parcellation # Avoid merging if there is only one parcellation
if len(all_parcellations) == 1: if len(all_parcellations) == 1:
resampled_parcellation_img = all_parcellations[0] resampled_parcellation_img = all_parcellations[0]
labels = all_labels[0] labels = dict(
zip(
np.trim_zeros(
np.unique(
resampled_parcellation_img.get_fdata().astype(int)
)
),
all_labels[0],
)
)
# Parcellations are already transformed to target standard space # Parcellations are already transformed to target standard space
else: else:
logger.debug("Merging parcellations.") logger.debug("Merging parcellations.")
@ -1310,7 +1311,7 @@ def merge_parcellations(
parcellations_list: list["Nifti1Image"], parcellations_list: list["Nifti1Image"],
parcellations_names: list[str], parcellations_names: list[str],
labels_lists: list[list[str]], labels_lists: list[list[str]],
) -> tuple["Nifti1Image", list[str]]: ) -> tuple["Nifti1Image", dict[int, str]]:
"""Merge multiple parcellations. """Merge multiple parcellations.
Parameters Parameters
@ -1328,8 +1329,8 @@ def merge_parcellations(
Niimg-like object Niimg-like object
The parcellation that results from merging the list of input The parcellation that results from merging the list of input
parcellations. parcellations.
list of str dict of int and str
List of labels for the resultant parcellation. Merged parcellation value to label mappings.
""" """
# Check for duplicated labels # Check for duplicated labels
@ -1346,34 +1347,54 @@ def merge_parcellations(
] ]
overlapping_voxels = False overlapping_voxels = False
ref_parc = parcellations_list[0] ref_parc = parcellations_list[0]
parc_data = ref_parc.get_fdata() ref_parc_data = ref_parc.get_fdata()
labels = labels_lists[0] # Get max value for the parcellation ROIs to get a reference for ROI value
# increment
max_val = np.max(
np.concatenate(
[np.unique(p.get_fdata().astype(int)) for p in parcellations_list]
)
)
# Setup parcellation value to label mapping
val_label_map = dict(
zip(
np.trim_zeros(np.unique(ref_parc_data.astype(int))),
labels_lists[0],
),
)
for t_parc, t_labels in zip(parcellations_list[1:], labels_lists[1:]): for idx, (parc, labs) in enumerate(
if t_parc.shape != ref_parc.shape: zip(parcellations_list[1:], labels_lists[1:])
):
# Resample to reference (1st in the list) parcellation
if parc.shape != ref_parc.shape:
warn_with_log( warn_with_log(
"The parcellations have different resolutions!" "The parcellations have different resolutions. "
"Resampling all parcellations to the first one in the list." "Resampling all parcellations to the first one in the list."
) )
t_parc = nimg.resample_to_img( parc = nimg.resample_to_img(
t_parc, ref_parc, interpolation="nearest", copy=True source_img=parc,
target_img=ref_parc,
interpolation="nearest",
copy=True,
) )
# Get the data from this parcellation # Get the data from this parcellation
t_parc_data = t_parc.get_fdata().copy() # must be copied parc_data = parc.get_fdata().copy() # must be copied
# Increase the values of each ROI to match the labels # Increase the values of each ROI to match the labels
t_parc_data[t_parc_data != 0] += len(labels) # and update label mapping
parc_data[parc_data != 0] += (idx + 1) * max_val
val_label_map.update(
dict(zip(np.trim_zeros(np.unique(parc_data.astype(int))), labs))
)
# Only set new values for the voxels that are 0 # Only set new values for the voxels that are 0
# This makes sure that the voxels that are in multiple # This makes sure that the voxels that are in multiple
# parcellations are assigned to the parcellation that was # parcellations are assigned to the parcellation that was
# first in the list. # first in the list.
if np.any(parc_data[t_parc_data != 0] != 0): if np.any(ref_parc_data[parc_data != 0] != 0):
overlapping_voxels = True overlapping_voxels = True
parc_data[parc_data == 0] += t_parc_data[parc_data == 0] ref_parc_data[ref_parc_data == 0] += parc_data[ref_parc_data == 0]
labels.extend(t_labels)
if overlapping_voxels: if overlapping_voxels:
warn_with_log( warn_with_log(
@ -1382,6 +1403,9 @@ def merge_parcellations(
"parcellation that was first in the list." "parcellation that was first in the list."
) )
parcellation_img_res = nimg.new_img_like(parcellations_list[0], parc_data) parcellation_img_res = nimg.new_img_like(
ref_niimg=parcellations_list[0],
data=ref_parc_data,
)
return parcellation_img_res, labels return parcellation_img_res, val_label_map

View file

@ -7,7 +7,6 @@
from pathlib import Path from pathlib import Path
import nibabel as nib
import numpy as np import numpy as np
import pytest import pytest
from nilearn.image import new_img_like, resample_to_img from nilearn.image import new_img_like, resample_to_img
@ -153,46 +152,6 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None:
target_space="MNI152NLin6Asym", target_space="MNI152NLin6Asym",
) )
schaefer_data = schaefer.get_fdata().copy()
schaefer_data[schaefer_data == 50] = 0
new_schaefer_path = tmp_path / "new_schaefer.nii.gz"
new_schaefer_img = new_img_like(schaefer, schaefer_data)
nib.save(new_schaefer_img, new_schaefer_path)
register_data(
kind="parcellation",
name="WrongValues",
parcellation_path=new_schaefer_path,
parcels_labels=labels[:-1],
space="MNI152Lin",
)
with pytest.raises(ValueError, match=r"must have all the values in the"):
load_data(
kind="parcellation",
name="WrongValues",
target_space="MNI152NLin6Asym",
)
schaefer_data = schaefer.get_fdata().copy()
schaefer_data[schaefer_data == 50] = 200
new_schaefer_path = tmp_path / "new_schaefer2.nii.gz"
new_schaefer_img = new_img_like(schaefer, schaefer_data)
nib.save(new_schaefer_img, new_schaefer_path)
register_data(
kind="parcellation",
name="WrongValues2",
parcellation_path=new_schaefer_path,
parcels_labels=labels,
space="MNI152Lin",
)
with pytest.raises(ValueError, match=r"must have all the values in the"):
load_data(
kind="parcellation",
name="WrongValues2",
target_space="MNI152NLin6Asym",
)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"name, parcellation_path, parcels_labels, space, overwrite", "name, parcellation_path, parcels_labels, space, overwrite",
@ -1172,7 +1131,7 @@ def test_get_single() -> None:
tailored_parcellation.get_fdata(), tailored_parcellation.get_fdata(),
resampled_raw_parcellation.get_fdata(), resampled_raw_parcellation.get_fdata(),
) )
assert tailored_labels == raw_labels assert list(tailored_labels.values()) == raw_labels
def test_get_multi_same_space() -> None: def test_get_multi_same_space() -> None:

View file

@ -213,7 +213,7 @@ class ParcelAggregation(BaseMarker):
logger.debug("Computing ROI means") logger.debug("Computing ROI means")
out_values = [] out_values = []
# Iterate over the parcels (existing) # Iterate over the parcels (existing)
for t_v in range(1, len(labels) + 1): for t_v in labels.keys():
t_values = agg_func(data[:, parcellation_values == t_v], axis=-1) t_values = agg_func(data[:, parcellation_values == t_v], axis=-1)
out_values.append(t_values) out_values.append(t_values)
# Update the labels just in case a parcel has no voxels # Update the labels just in case a parcel has no voxels
@ -238,6 +238,6 @@ class ParcelAggregation(BaseMarker):
return { return {
"aggregation": { "aggregation": {
"data": out_values, "data": out_values,
"col_names": labels, "col_names": list(labels.values()),
}, },
} }

View file

@ -408,8 +408,8 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
parcellation2_data = parcellation_data.copy() parcellation2_data = parcellation_data.copy()
parcellation2_data[parcellation2_data <= 8] = 0 parcellation2_data[parcellation2_data <= 8] = 0
parcellation2_data[parcellation2_data > 0] -= 8 parcellation2_data[parcellation2_data > 0] -= 8
labels1 = labels[:8] labels1 = list(labels.values())[:8]
labels2 = labels[8:] labels2 = list(labels.values())[8:]
parcellation1_img = new_img_like( parcellation1_img = new_img_like(
testing_parcellation, parcellation1_data testing_parcellation, parcellation1_data
@ -512,8 +512,9 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
# Make the second parcellation overlap with the first # Make the second parcellation overlap with the first
parcellation2_data[parcellation2_data <= 6] = 0 parcellation2_data[parcellation2_data <= 6] = 0
parcellation2_data[parcellation2_data > 0] -= 6 parcellation2_data[parcellation2_data > 0] -= 6
labels1 = [f"low_{x}" for x in labels[:8]] # Change the labels # Change the labels
labels2 = [f"high_{x}" for x in labels[6:]] # Change the labels labels1 = [f"low_{x}" for x in list(labels.values())[:8]]
labels2 = [f"high_{x}" for x in list(labels.values())[6:]]
parcellation1_img = new_img_like( parcellation1_img = new_img_like(
testing_parcellation, parcellation1_data testing_parcellation, parcellation1_data
@ -621,8 +622,8 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
parcellation2_data = parcellation_data.copy() parcellation2_data = parcellation_data.copy()
parcellation2_data[parcellation2_data <= 8] = 0 parcellation2_data[parcellation2_data <= 8] = 0
parcellation2_data[parcellation2_data > 0] -= 8 parcellation2_data[parcellation2_data > 0] -= 8
labels1 = labels[:8] labels1 = list(labels.values())[:8]
labels2 = labels[7:-1] # One label is duplicated labels2 = list(labels.values())[7:-1] # One label is duplicated
parcellation1_img = new_img_like( parcellation1_img = new_img_like(
testing_parcellation, parcellation1_data testing_parcellation, parcellation1_data