[ENH]: Allow ParcelAggregation to apply multiple parcellations at once #131

Merged
fraimondo merged 4 commits from fraimondo/issue131 into main 2022-11-24 09:54:31 +00:00
7 changed files with 329 additions and 32 deletions

View file

@ -89,6 +89,8 @@ Enhancements
- Add support for "masks" (:gh:`79` by `Fede Raimondo`_).
- Allow :class:`junifer.markers.ParcelAggregation` to apply multiple parcellations at once (:gh:`131` by `Fede Raimondo`_).
Bugs
~~~~

View file

@ -213,6 +213,18 @@ def load_parcellation(
parcellation_img = None
if path_only is False:
parcellation_img = nib.load(parcellation_fname)
parcel_values = np.unique(parcellation_img.get_fdata())
if len(parcel_values) - 1 != len(parcellation_labels):
raise_error(
f"Parcellation {name} has {len(parcel_values) - 1} parcels "
f"but {len(parcellation_labels)} labels."
)
parcel_values.sort()
if np.any(np.diff(parcel_values) != 1):
raise_error(
f"Parcellation {name} must have all the values in the range "
f"[0, {len(parcel_values)}]."
)
return parcellation_img, parcellation_labels, parcellation_fname

View file

@ -11,6 +11,9 @@ from typing import List
import pytest
from numpy.testing import assert_array_almost_equal, assert_array_equal
import nibabel as nib
from nilearn.image import new_img_like
from junifer.data.parcellations import (
_retrieve_parcellation,
_retrieve_schaefer,
@ -72,6 +75,50 @@ def test_register_parcellation_already_registered() -> None:
)
def test_parcellation_wrong_labels_values(tmp_path: Path) -> None:
"""Test parcellation with wrong labels and values.
Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
schaefer, labels, schaefer_path = load_parcellation("Schaefer100x7")
assert schaefer is not None
# Test wrong number of labels
register_parcellation("WrongLabels", schaefer_path, labels[:10])
with pytest.raises(ValueError, match=r"has 100 parcels but 10"):
load_parcellation("WrongLabels")
# Test wrong number of labels
register_parcellation("WrongLabels2", schaefer_path, labels + ["wrong"])
with pytest.raises(ValueError, match=r"has 100 parcels but 101"):
load_parcellation("WrongLabels2")
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_parcellation("WrongValues", new_schaefer_path, labels[:-1])
with pytest.raises(ValueError, match=r"the range [0, 99]"):
load_parcellation("WrongValues")
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_parcellation("WrongValues2", new_schaefer_path, labels)
with pytest.raises(ValueError, match=r"the range [0, 100]"):
load_parcellation("WrongValues2")
@pytest.mark.parametrize(
"name, parcellation_path, parcels_labels, overwrite",
[

View file

@ -6,7 +6,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import numpy as np
@ -26,8 +26,8 @@ class RSSETSMarker(BaseMarker):
Parameters
----------
parcellation : str
The name of the parcellation. Check valid options by calling
parcellation : str or list of str
The name(s) of the parcellation(s). Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`.
agg_method : str, optional
The method to perform aggregation using. Check valid options in
@ -47,7 +47,7 @@ class RSSETSMarker(BaseMarker):
def __init__(
self,
parcellation: str,
parcellation: Union[str, List[str]],
agg_method: str = "mean",
agg_method_params: Optional[Dict] = None,
mask: Optional[str] = None,

View file

@ -4,7 +4,7 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance
@ -24,8 +24,8 @@ class FunctionalConnectivityParcels(BaseMarker):
Parameters
----------
parcellation : str
The name of the parcellation. Check valid options by calling
parcellation : str or list of str
The name(s) of the parcellation(s). Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`.
agg_method : str, optional
The method to perform aggregation using. Check valid options in
@ -51,7 +51,7 @@ class FunctionalConnectivityParcels(BaseMarker):
def __init__(
self,
parcellation: str,
parcellation: Union[str, List[str]],
agg_method: str = "mean",
agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance",

View file

@ -7,7 +7,7 @@
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import numpy as np
from nilearn.image import math_img, resample_to_img
from nilearn.image import math_img, resample_to_img, new_img_like
from nilearn.maskers import NiftiMasker
from ..api.decorators import register_marker
@ -26,8 +26,8 @@ class ParcelAggregation(BaseMarker):
Parameters
----------
parcellation : str
The name of the parcellation. Check valid options by calling
parcellation : str or list of str
The name(s) of the parcellation(s). Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`.
method : str
The method to perform aggregation using. Check valid options in
@ -50,13 +50,15 @@ class ParcelAggregation(BaseMarker):
def __init__(
self,
parcellation: str,
parcellation: Union[str, List[str]],
method: str,
method_params: Optional[Dict[str, Any]] = None,
mask: Optional[str] = None,
on: Union[List[str], str, None] = None,
name: Optional[str] = None,
) -> None:
if not isinstance(parcellation, list):
parcellation = [parcellation]
self.parcellation = parcellation
self.method = method
self.method_params = method_params or {}
@ -160,17 +162,50 @@ class ParcelAggregation(BaseMarker):
)
# Get the min of the voxels sizes and use it as the resolution
resolution = np.min(t_input.header.get_zooms()[:3])
t_parcellation, t_labels, _ = load_parcellation(
name=self.parcellation,
resolution=resolution,
)
parcellation_img_res = resample_to_img(
t_parcellation,
t_input,
interpolation="nearest",
copy=True,
)
# Load the parcellations
all_parcelations = []
all_labels = []
for t_parc_name in self.parcellation:
t_parcellation, t_labels, _ = load_parcellation(
name=t_parc_name,
resolution=resolution,
)
# Resample all of them to the image
t_parcellation_img_res = resample_to_img(
t_parcellation,
t_input,
interpolation="nearest",
copy=True,
)
all_parcelations.append(t_parcellation_img_res)
all_labels.append(t_labels)
# Avoid merging if there is only one parcellation
if len(all_parcelations) == 1:
parcellation_img_res = all_parcelations[0]
labels = all_labels[0]
else:
# Merge the parcellations
parc_data = all_parcelations[0].get_fdata()
labels = all_labels[0]
for t_parc, t_labels in zip(all_parcelations[1:], all_labels[1:]):
# Get the data from this parcellation
t_parc_data = t_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)
# 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.
parc_data[parc_data == 0] += t_parc_data[parc_data == 0]
labels.extend(t_labels)
parcellation_img_res = new_img_like(
all_parcelations[0],
parc_data,
)
parcellation_bin = math_img(
"img != 0",
@ -213,7 +248,7 @@ class ParcelAggregation(BaseMarker):
out_values.append(t_values)
# Update the labels just in case a parcel has no voxels
# in it
out_labels.append(t_labels[t_v - 1])
out_labels.append(labels[t_v - 1])
out_values = np.array(out_values).T
out = {"data": out_values, "columns": out_labels}

View file

@ -6,13 +6,14 @@
import nibabel as nib
import numpy as np
import pytest
from pathlib import Path
from nilearn import datasets
from nilearn.image import concat_imgs, math_img, resample_to_img
from nilearn.image import concat_imgs, math_img, resample_to_img, new_img_like
from nilearn.maskers import NiftiLabelsMasker, NiftiMasker
from numpy.testing import assert_array_almost_equal, assert_array_equal
from scipy.stats import trim_mean
from junifer.data import load_mask
from junifer.data import load_mask, load_parcellation, register_parcellation
from junifer.markers.parcel_aggregation import ParcelAggregation
@ -86,7 +87,7 @@ def test_ParcelAggregation_3D() -> None:
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == "Schaefer100x7"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
@ -111,7 +112,7 @@ def test_ParcelAggregation_3D() -> None:
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "std"
assert meta["parcellation"] == "Schaefer100x7"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation"
@ -144,7 +145,7 @@ def test_ParcelAggregation_3D() -> None:
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "trim_mean"
assert meta["parcellation"] == "Schaefer100x7"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation"
@ -178,7 +179,7 @@ def test_ParcelAggregation_4D():
meta = marker.get_meta("BOLD")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == "Schaefer100x7"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "BOLD_ParcelAggregation"
assert meta["class"] == "ParcelAggregation"
@ -202,8 +203,8 @@ def test_ParcelAggregation_3D_mask() -> None:
# Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker(
labels_img=parcellation.maps,
mask_img=mask_img)
labels_img=parcellation.maps, mask_img=mask_img
)
auto = nifti_masker.fit_transform(img)
# Use the ParcelAggregation object
@ -223,9 +224,209 @@ def test_ParcelAggregation_3D_mask() -> None:
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == "Schaefer100x7"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] == "GM_prob0.2"
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
"""Test ParcelAggregation with multiple non-overlapping parcellations.
Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
# Get the testing parcellation
parcellation, labels, _ = load_parcellation("Schaefer100x7")
assert parcellation is not None
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create two parcellations from it
parcellation_data = parcellation.get_fdata()
parcellation1_data = parcellation_data.copy()
parcellation1_data[parcellation1_data > 50] = 0
parcellation2_data = parcellation_data.copy()
parcellation2_data[parcellation2_data <= 50] = 0
parcellation2_data[parcellation2_data > 0] -= 50
labels1 = labels[:50]
labels2 = labels[50:]
parcellation1_img = new_img_like(parcellation, parcellation1_data)
parcellation2_img = new_img_like(parcellation, parcellation2_data)
parcellation1_path = tmp_path / "parcellation1.nii.gz"
parcellation2_path = tmp_path / "parcellation2.nii.gz"
nib.save(parcellation1_img, parcellation1_path)
nib.save(parcellation2_img, parcellation2_path)
register_parcellation("Schaefer100x7_low", parcellation1_path, labels1)
register_parcellation("Schaefer100x7_high", parcellation2_path, labels2)
# Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation(
parcellation="Schaefer100x7",
method="mean",
name="gmd_schaefer100x7_mean",
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img))
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2
assert orig_mean_data.shape[0] == 1
assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean)
meta = marker_original.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"],
method="mean",
name="gmd_schaefer100x7_mean",
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img))
split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2
assert split_mean_data.shape[0] == 1
assert split_mean_data.shape[1] == 100
meta = marker_split.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7_low", "Schaefer100x7_high"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Data and labels should be the same
assert_array_equal(orig_mean_data, split_mean_data)
assert orig_mean["columns"] == split_mean["columns"]
def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
"""Test ParcelAggregation with multiple overlapping parcellations.
Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
# Get the testing parcellation
parcellation, labels, _ = load_parcellation("Schaefer100x7")
assert parcellation is not None
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create two parcellations from it
parcellation_data = parcellation.get_fdata()
parcellation1_data = parcellation_data.copy()
parcellation1_data[parcellation1_data > 50] = 0
parcellation2_data = parcellation_data.copy()
# Make the second parcellation overlap with the first
parcellation2_data[parcellation2_data <= 45] = 0
parcellation2_data[parcellation2_data > 0] -= 45
labels1 = [f"low_{x}" for x in labels[:50]] # Change the labels
labels2 = [f"high_{x}" for x in labels[45:]] # Change the labels
parcellation1_img = new_img_like(parcellation, parcellation1_data)
parcellation2_img = new_img_like(parcellation, parcellation2_data)
parcellation1_path = tmp_path / "parcellation1.nii.gz"
parcellation2_path = tmp_path / "parcellation2.nii.gz"
nib.save(parcellation1_img, parcellation1_path)
nib.save(parcellation2_img, parcellation2_path)
register_parcellation("Schaefer100x7_low2", parcellation1_path, labels1)
register_parcellation("Schaefer100x7_high2", parcellation2_path, labels2)
# Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation(
parcellation="Schaefer100x7",
method="mean",
name="gmd_schaefer100x7_mean",
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img))
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2
assert orig_mean_data.shape[0] == 1
assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean)
meta = marker_original.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"],
method="mean",
name="gmd_schaefer100x7_mean",
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img))
split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2
assert split_mean_data.shape[0] == 1
assert split_mean_data.shape[1] == 100
meta = marker_split.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == [
"Schaefer100x7_low2",
"Schaefer100x7_high2",
]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Data should be the same
assert_array_equal(orig_mean_data, split_mean_data)
# Labels should be "low" for the first 50 and "high" for the second 50
assert all(x.startswith("low") for x in split_mean["columns"][:50])
assert all(x.startswith("high") for x in split_mean["columns"][50:])