[BUG]: Cannot register parcellation with non-continuous values #257
7 changed files with 71 additions and 84 deletions
1
docs/changes/newsfragments/257.change
Normal file
1
docs/changes/newsfragments/257.change
Normal 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`_
|
||||||
1
docs/changes/newsfragments/257.enh
Normal file
1
docs/changes/newsfragments/257.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Allow parcellation registration with non-continuous ranges as values by `Synchon Mandal`_
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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()),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue