[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,
|
||||
) -> Union[
|
||||
tuple[ArrayLike, list[str]], # coordinates
|
||||
tuple["Nifti1Image", list[str]], # parcellation / maps
|
||||
tuple["Nifti1Image", dict[int, str]], # parcellation
|
||||
tuple["Nifti1Image", list[str]], # maps
|
||||
"Nifti1Image", # mask
|
||||
]:
|
||||
"""Get tailored ``kind`` for ``target_data``.
|
||||
|
|
@ -142,7 +143,7 @@ def get_data(
|
|||
Returns
|
||||
-------
|
||||
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
|
||||
|
||||
Raises
|
||||
|
|
|
|||
|
|
@ -403,7 +403,7 @@ class ParcellationRegistry(BasePipelineDataRegistry):
|
|||
if not path_only:
|
||||
# Load image via nibabel
|
||||
parcellation_img = nib.load(parcellation_fname)
|
||||
# Get unique values
|
||||
# Get unique values (returns sorted result)
|
||||
parcel_values = np.unique(parcellation_img.get_fdata())
|
||||
# Check for dimension
|
||||
if len(parcel_values) - 1 != len(parcellation_labels):
|
||||
|
|
@ -411,14 +411,6 @@ class ParcellationRegistry(BasePipelineDataRegistry):
|
|||
f"Parcellation {name} has {len(parcel_values) - 1} "
|
||||
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
|
||||
|
||||
|
|
@ -427,7 +419,7 @@ class ParcellationRegistry(BasePipelineDataRegistry):
|
|||
parcellations: Union[str, list[str]],
|
||||
target_data: dict[str, Any],
|
||||
extra_input: Optional[dict[str, Any]] = None,
|
||||
) -> tuple["Nifti1Image", list[str]]:
|
||||
) -> tuple["Nifti1Image", dict[int, str]]:
|
||||
"""Get parcellation, tailored for the target image.
|
||||
|
||||
Parameters
|
||||
|
|
@ -446,8 +438,8 @@ class ParcellationRegistry(BasePipelineDataRegistry):
|
|||
-------
|
||||
Nifti1Image
|
||||
The parcellation image.
|
||||
list of str
|
||||
Parcellation labels.
|
||||
dict
|
||||
Parcellation value to label mappings.
|
||||
|
||||
Raises
|
||||
------
|
||||
|
|
@ -561,7 +553,16 @@ class ParcellationRegistry(BasePipelineDataRegistry):
|
|||
# Avoid merging if there is only one parcellation
|
||||
if len(all_parcellations) == 1:
|
||||
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
|
||||
else:
|
||||
logger.debug("Merging parcellations.")
|
||||
|
|
@ -1310,7 +1311,7 @@ def merge_parcellations(
|
|||
parcellations_list: list["Nifti1Image"],
|
||||
parcellations_names: list[str],
|
||||
labels_lists: list[list[str]],
|
||||
) -> tuple["Nifti1Image", list[str]]:
|
||||
) -> tuple["Nifti1Image", dict[int, str]]:
|
||||
"""Merge multiple parcellations.
|
||||
|
||||
Parameters
|
||||
|
|
@ -1328,8 +1329,8 @@ def merge_parcellations(
|
|||
Niimg-like object
|
||||
The parcellation that results from merging the list of input
|
||||
parcellations.
|
||||
list of str
|
||||
List of labels for the resultant parcellation.
|
||||
dict of int and str
|
||||
Merged parcellation value to label mappings.
|
||||
|
||||
"""
|
||||
# Check for duplicated labels
|
||||
|
|
@ -1346,34 +1347,54 @@ def merge_parcellations(
|
|||
]
|
||||
overlapping_voxels = False
|
||||
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:]):
|
||||
if t_parc.shape != ref_parc.shape:
|
||||
for idx, (parc, labs) in enumerate(
|
||||
zip(parcellations_list[1:], labels_lists[1:])
|
||||
):
|
||||
# Resample to reference (1st in the list) parcellation
|
||||
if parc.shape != ref_parc.shape:
|
||||
warn_with_log(
|
||||
"The parcellations have different resolutions!"
|
||||
"The parcellations have different resolutions. "
|
||||
"Resampling all parcellations to the first one in the list."
|
||||
)
|
||||
t_parc = nimg.resample_to_img(
|
||||
t_parc, ref_parc, interpolation="nearest", copy=True
|
||||
parc = nimg.resample_to_img(
|
||||
source_img=parc,
|
||||
target_img=ref_parc,
|
||||
interpolation="nearest",
|
||||
copy=True,
|
||||
)
|
||||
|
||||
# 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
|
||||
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
|
||||
# This makes sure that the voxels that are in multiple
|
||||
# parcellations are assigned to the parcellation that was
|
||||
# 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
|
||||
|
||||
parc_data[parc_data == 0] += t_parc_data[parc_data == 0]
|
||||
labels.extend(t_labels)
|
||||
ref_parc_data[ref_parc_data == 0] += parc_data[ref_parc_data == 0]
|
||||
|
||||
if overlapping_voxels:
|
||||
warn_with_log(
|
||||
|
|
@ -1382,6 +1403,9 @@ def merge_parcellations(
|
|||
"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
|
||||
|
||||
import nibabel as nib
|
||||
import numpy as np
|
||||
import pytest
|
||||
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",
|
||||
)
|
||||
|
||||
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(
|
||||
"name, parcellation_path, parcels_labels, space, overwrite",
|
||||
|
|
@ -1172,7 +1131,7 @@ def test_get_single() -> None:
|
|||
tailored_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:
|
||||
|
|
|
|||
|
|
@ -213,7 +213,7 @@ class ParcelAggregation(BaseMarker):
|
|||
logger.debug("Computing ROI means")
|
||||
out_values = []
|
||||
# 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)
|
||||
out_values.append(t_values)
|
||||
# Update the labels just in case a parcel has no voxels
|
||||
|
|
@ -238,6 +238,6 @@ class ParcelAggregation(BaseMarker):
|
|||
return {
|
||||
"aggregation": {
|
||||
"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[parcellation2_data <= 8] = 0
|
||||
parcellation2_data[parcellation2_data > 0] -= 8
|
||||
labels1 = labels[:8]
|
||||
labels2 = labels[8:]
|
||||
labels1 = list(labels.values())[:8]
|
||||
labels2 = list(labels.values())[8:]
|
||||
|
||||
parcellation1_img = new_img_like(
|
||||
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
|
||||
parcellation2_data[parcellation2_data <= 6] = 0
|
||||
parcellation2_data[parcellation2_data > 0] -= 6
|
||||
labels1 = [f"low_{x}" for x in labels[:8]] # Change the labels
|
||||
labels2 = [f"high_{x}" for x in labels[6:]] # Change the labels
|
||||
# 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(
|
||||
testing_parcellation, parcellation1_data
|
||||
|
|
@ -621,8 +622,8 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
|
|||
parcellation2_data = parcellation_data.copy()
|
||||
parcellation2_data[parcellation2_data <= 8] = 0
|
||||
parcellation2_data[parcellation2_data > 0] -= 8
|
||||
labels1 = labels[:8]
|
||||
labels2 = labels[7:-1] # One label is duplicated
|
||||
labels1 = list(labels.values())[:8]
|
||||
labels2 = list(labels.values())[7:-1] # One label is duplicated
|
||||
|
||||
parcellation1_img = new_img_like(
|
||||
testing_parcellation, parcellation1_data
|
||||
|
|
|
|||
Loading…
Reference in a new issue