diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index 8d7ea61af..84fbc2d04 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -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 ~~~~ diff --git a/junifer/data/parcellations.py b/junifer/data/parcellations.py index 26629d636..e9c7c5baf 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -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 diff --git a/junifer/data/tests/test_parcellations.py b/junifer/data/tests/test_parcellations.py index f1278f769..1e04f10f4 100644 --- a/junifer/data/tests/test_parcellations.py +++ b/junifer/data/tests/test_parcellations.py @@ -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", [ diff --git a/junifer/markers/ets_rss.py b/junifer/markers/ets_rss.py index c6f0d6442..a7e7f939a 100644 --- a/junifer/markers/ets_rss.py +++ b/junifer/markers/ets_rss.py @@ -6,7 +6,7 @@ # Synchon Mandal # 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, diff --git a/junifer/markers/functional_connectivity_parcels.py b/junifer/markers/functional_connectivity_parcels.py index 3bb94216a..b42c1171d 100644 --- a/junifer/markers/functional_connectivity_parcels.py +++ b/junifer/markers/functional_connectivity_parcels.py @@ -4,7 +4,7 @@ # Kaustubh R. Patil # 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", diff --git a/junifer/markers/parcel_aggregation.py b/junifer/markers/parcel_aggregation.py index f8b4ca8bc..f76c6c8c4 100644 --- a/junifer/markers/parcel_aggregation.py +++ b/junifer/markers/parcel_aggregation.py @@ -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} diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index d8abd0d58..cace74bba 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -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:])