[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`_). - 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 Bugs
~~~~ ~~~~

View file

@ -213,6 +213,18 @@ def load_parcellation(
parcellation_img = None parcellation_img = None
if path_only is False: if path_only is False:
parcellation_img = nib.load(parcellation_fname) 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 return parcellation_img, parcellation_labels, parcellation_fname

View file

@ -11,6 +11,9 @@ from typing import List
import pytest import pytest
from numpy.testing import assert_array_almost_equal, assert_array_equal 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 ( from junifer.data.parcellations import (
_retrieve_parcellation, _retrieve_parcellation,
_retrieve_schaefer, _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( @pytest.mark.parametrize(
"name, parcellation_path, parcels_labels, overwrite", "name, parcellation_path, parcels_labels, overwrite",
[ [

View file

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

View file

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

View file

@ -7,7 +7,7 @@
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import numpy as np 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 nilearn.maskers import NiftiMasker
from ..api.decorators import register_marker from ..api.decorators import register_marker
@ -26,8 +26,8 @@ class ParcelAggregation(BaseMarker):
Parameters Parameters
---------- ----------
parcellation : str parcellation : str or list of str
The name of the parcellation. Check valid options by calling The name(s) of the parcellation(s). Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`. :func:`junifer.data.parcellations.list_parcellations`.
method : str method : str
The method to perform aggregation using. Check valid options in The method to perform aggregation using. Check valid options in
@ -50,13 +50,15 @@ class ParcelAggregation(BaseMarker):
def __init__( def __init__(
self, self,
parcellation: str, parcellation: Union[str, List[str]],
method: str, method: str,
method_params: Optional[Dict[str, Any]] = None, method_params: Optional[Dict[str, Any]] = None,
mask: Optional[str] = None, mask: Optional[str] = None,
on: Union[List[str], str, None] = None, on: Union[List[str], str, None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
if not isinstance(parcellation, list):
parcellation = [parcellation]
self.parcellation = parcellation self.parcellation = parcellation
self.method = method self.method = method
self.method_params = method_params or {} 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 # Get the min of the voxels sizes and use it as the resolution
resolution = np.min(t_input.header.get_zooms()[:3]) resolution = np.min(t_input.header.get_zooms()[:3])
# Load the parcellations
all_parcelations = []
all_labels = []
for t_parc_name in self.parcellation:
t_parcellation, t_labels, _ = load_parcellation( t_parcellation, t_labels, _ = load_parcellation(
name=self.parcellation, name=t_parc_name,
resolution=resolution, resolution=resolution,
) )
# Resample all of them to the image
parcellation_img_res = resample_to_img( t_parcellation_img_res = resample_to_img(
t_parcellation, t_parcellation,
t_input, t_input,
interpolation="nearest", interpolation="nearest",
copy=True, 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( parcellation_bin = math_img(
"img != 0", "img != 0",
@ -213,7 +248,7 @@ class ParcelAggregation(BaseMarker):
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
# in it # 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_values = np.array(out_values).T
out = {"data": out_values, "columns": out_labels} out = {"data": out_values, "columns": out_labels}

View file

@ -6,13 +6,14 @@
import nibabel as nib import nibabel as nib
import numpy as np import numpy as np
import pytest import pytest
from pathlib import Path
from nilearn import datasets 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 nilearn.maskers import NiftiLabelsMasker, NiftiMasker
from numpy.testing import assert_array_almost_equal, assert_array_equal from numpy.testing import assert_array_almost_equal, assert_array_equal
from scipy.stats import trim_mean 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 from junifer.markers.parcel_aggregation import ParcelAggregation
@ -86,7 +87,7 @@ def test_ParcelAggregation_3D() -> None:
meta = marker.get_meta("VBM_GM")["marker"] meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean" assert meta["method"] == "mean"
assert meta["parcellation"] == "Schaefer100x7" assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
@ -111,7 +112,7 @@ def test_ParcelAggregation_3D() -> None:
meta = marker.get_meta("VBM_GM")["marker"] meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "std" assert meta["method"] == "std"
assert meta["parcellation"] == "Schaefer100x7" assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None assert meta["mask"] is None
assert meta["name"] == "VBM_GM_ParcelAggregation" assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
@ -144,7 +145,7 @@ def test_ParcelAggregation_3D() -> None:
meta = marker.get_meta("VBM_GM")["marker"] meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "trim_mean" assert meta["method"] == "trim_mean"
assert meta["parcellation"] == "Schaefer100x7" assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None assert meta["mask"] is None
assert meta["name"] == "VBM_GM_ParcelAggregation" assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
@ -178,7 +179,7 @@ def test_ParcelAggregation_4D():
meta = marker.get_meta("BOLD")["marker"] meta = marker.get_meta("BOLD")["marker"]
assert meta["method"] == "mean" assert meta["method"] == "mean"
assert meta["parcellation"] == "Schaefer100x7" assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None assert meta["mask"] is None
assert meta["name"] == "BOLD_ParcelAggregation" assert meta["name"] == "BOLD_ParcelAggregation"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
@ -202,8 +203,8 @@ def test_ParcelAggregation_3D_mask() -> None:
# Create NiftiLabelsMasker # Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker( nifti_masker = NiftiLabelsMasker(
labels_img=parcellation.maps, labels_img=parcellation.maps, mask_img=mask_img
mask_img=mask_img) )
auto = nifti_masker.fit_transform(img) auto = nifti_masker.fit_transform(img)
# Use the ParcelAggregation object # Use the ParcelAggregation object
@ -223,9 +224,209 @@ def test_ParcelAggregation_3D_mask() -> None:
meta = marker.get_meta("VBM_GM")["marker"] meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean" assert meta["method"] == "mean"
assert meta["parcellation"] == "Schaefer100x7" assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] == "GM_prob0.2" assert meta["mask"] == "GM_prob0.2"
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM" assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {} 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:])