From 64efd0ce4c3dd83d3b84f20480adfdb35a17460c Mon Sep 17 00:00:00 2001 From: Fede Date: Wed, 23 Nov 2022 23:08:19 +0100 Subject: [PATCH 1/4] [ENH]: Allow ParcelAggregation to apply multiple parcellations at once (WIP) --- junifer/data/parcellations.py | 13 ++++ junifer/data/tests/test_parcellations.py | 3 + junifer/markers/parcel_aggregation.py | 61 +++++++++++++++---- .../markers/tests/test_parcel_aggregation.py | 56 +++++++++++++++-- 4 files changed, 115 insertions(+), 18 deletions(-) diff --git a/junifer/data/parcellations.py b/junifer/data/parcellations.py index 26629d636..ad4a54170 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -213,6 +213,19 @@ 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 but " + f"{len(parcellation_labels)} labels." + ) + if parcel_values.min() != 0 and parcel_values.max() != len( + parcel_values + ) - 1: + raise_error( + f"Parcellation {name} has parcel values outside 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..5c012f084 100644 --- a/junifer/data/tests/test_parcellations.py +++ b/junifer/data/tests/test_parcellations.py @@ -72,6 +72,9 @@ def test_register_parcellation_already_registered() -> None: ) +# TODO: Add tests to verify wrong number of labels and wrong values + + @pytest.mark.parametrize( "name, parcellation_path, parcels_labels, overwrite", [ diff --git a/junifer/markers/parcel_aggregation.py b/junifer/markers/parcel_aggregation.py index f8b4ca8bc..52d34b842 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 @@ -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() + # 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..5536e1224 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -86,7 +86,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 +111,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 +144,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 +178,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" @@ -223,9 +223,55 @@ 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() -> None: + """Test ParcelAggregation object on 3D images, multiple parcellations.""" + + # Get the testing parcellation (for nilearn) + parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) + + # 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 NiftiLabelsMasker for schaefer 100 + nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) + auto_schaefer = nifti_masker.fit_transform(img) + + # Create NiftiLabelsMasker for Tian 2012 + nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) + auto_schaefer = nifti_masker.fit_transform(img) + + + # Use the ParcelAggregation object + marker = ParcelAggregation( + parcellation=["Schaefer100x7", "SUITxMNI"], + method="mean", + mask="GM_prob0.2", + name="gmd_schaefer100x7_mean", + on="VBM_GM", + ) # Test passing "on" as a keyword argument + input = dict(VBM_GM=dict(data=img)) + jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] + + assert jun_values3d_mean.ndim == 2 + assert jun_values3d_mean.shape[0] == 1 + # assert_array_almost_equal(auto, jun_values3d_mean) + + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "mean" + assert meta["parcellation"] == ["Schaefer100x7", "SUITxMNI"] + 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"] == {} + -- 2.52.0 From 3d0b534bd9309d24e4cfe95d71e17fc8156fd34a Mon Sep 17 00:00:00 2001 From: Fede Date: Wed, 23 Nov 2022 23:11:48 +0100 Subject: [PATCH 2/4] Test WIP --- junifer/data/parcellations.py | 11 ++++++----- junifer/markers/tests/test_parcel_aggregation.py | 16 +++++++--------- 2 files changed, 13 insertions(+), 14 deletions(-) diff --git a/junifer/data/parcellations.py b/junifer/data/parcellations.py index ad4a54170..221e8c3b8 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -216,12 +216,13 @@ def load_parcellation( 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 but " - f"{len(parcellation_labels)} labels." + f"Parcellation {name} has {len(parcel_values) - 1} parcels" + f"but {len(parcellation_labels)} labels." ) - if parcel_values.min() != 0 and parcel_values.max() != len( - parcel_values - ) - 1: + if ( + parcel_values.min() != 0 + and parcel_values.max() != len(parcel_values) - 1 + ): raise_error( f"Parcellation {name} has parcel values outside the range " f"[0, {len(parcel_values)}]." diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index 5536e1224..ce3accd5d 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -235,21 +235,20 @@ def test_ParcelAggregation_3D_multiple() -> None: """Test ParcelAggregation object on 3D images, multiple parcellations.""" # Get the testing parcellation (for nilearn) - parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) + # parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) # 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 NiftiLabelsMasker for schaefer 100 - nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) - auto_schaefer = nifti_masker.fit_transform(img) - - # Create NiftiLabelsMasker for Tian 2012 - nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) - auto_schaefer = nifti_masker.fit_transform(img) + # # Create NiftiLabelsMasker for schaefer 100 + # nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) + # auto_schaefer = nifti_masker.fit_transform(img) + # # Create NiftiLabelsMasker for Tian 2012 + # nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) + # auto_tian = nifti_masker.fit_transform(img) # Use the ParcelAggregation object marker = ParcelAggregation( @@ -274,4 +273,3 @@ def test_ParcelAggregation_3D_multiple() -> None: assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" assert meta["method_params"] == {} - -- 2.52.0 From bfd0e61dc5e3d4ee64d660030face28b46ce2baf Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Nov 2022 10:21:50 +0100 Subject: [PATCH 3/4] Add more tests + multiple parcellation support in FC markers --- docs/changes/latest.inc | 2 + junifer/data/parcellations.py | 10 +- junifer/data/tests/test_parcellations.py | 47 +++- junifer/markers/ets_rss.py | 8 +- .../functional_connectivity_parcels.py | 8 +- junifer/markers/parcel_aggregation.py | 6 +- .../markers/tests/test_parcel_aggregation.py | 205 ++++++++++++++++-- 7 files changed, 244 insertions(+), 42 deletions(-) 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 221e8c3b8..e9c7c5baf 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -216,15 +216,13 @@ def load_parcellation( 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"Parcellation {name} has {len(parcel_values) - 1} parcels " f"but {len(parcellation_labels)} labels." ) - if ( - parcel_values.min() != 0 - and parcel_values.max() != len(parcel_values) - 1 - ): + parcel_values.sort() + if np.any(np.diff(parcel_values) != 1): raise_error( - f"Parcellation {name} has parcel values outside the range " + f"Parcellation {name} must have all the values in the range " f"[0, {len(parcel_values)}]." ) diff --git a/junifer/data/tests/test_parcellations.py b/junifer/data/tests/test_parcellations.py index 5c012f084..778ebeaac 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,7 +75,49 @@ def test_register_parcellation_already_registered() -> None: ) -# TODO: Add tests to verify wrong number of labels and wrong values +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( 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 52d34b842..f76c6c8c4 100644 --- a/junifer/markers/parcel_aggregation.py +++ b/junifer/markers/parcel_aggregation.py @@ -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 @@ -191,7 +191,7 @@ class ParcelAggregation(BaseMarker): 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() + 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) diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index ce3accd5d..710caa41a 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 @@ -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 @@ -231,45 +232,201 @@ def test_ParcelAggregation_3D_mask() -> None: assert meta["method_params"] == {} -def test_ParcelAggregation_3D_multiple() -> None: - """Test ParcelAggregation object on 3D images, multiple parcellations.""" +def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None: + """Test ParcelAggregation with multiple non-overlapping parcellations. - # Get the testing parcellation (for nilearn) - # parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) + 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 NiftiLabelsMasker for schaefer 100 - # nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) - # auto_schaefer = nifti_masker.fit_transform(img) + # 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:] - # # Create NiftiLabelsMasker for Tian 2012 - # nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) - # auto_tian = nifti_masker.fit_transform(img) + parcellation1_img = new_img_like(parcellation, parcellation1_data) + parcellation2_img = new_img_like(parcellation, parcellation2_data) - # Use the ParcelAggregation object - marker = ParcelAggregation( - parcellation=["Schaefer100x7", "SUITxMNI"], + 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", - mask="GM_prob0.2", name="gmd_schaefer100x7_mean", on="VBM_GM", ) # Test passing "on" as a keyword argument input = dict(VBM_GM=dict(data=img)) - jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] + orig_mean = marker_original.fit_transform(input)["VBM_GM"] - assert jun_values3d_mean.ndim == 2 - assert jun_values3d_mean.shape[0] == 1 + 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.get_meta("VBM_GM")["marker"] + meta = marker_original.get_meta("VBM_GM")["marker"] assert meta["method"] == "mean" - assert meta["parcellation"] == ["Schaefer100x7", "SUITxMNI"] - assert meta["mask"] == "GM_prob0.2" + assert meta["parcellation"] == ["Schaefer100x7"] + assert meta["mask"] == 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"] == 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"] == 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"] == 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:]) -- 2.52.0 From 2d2a3e7cf3159c348719ad509b3934139210e6e7 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Nov 2022 10:37:42 +0100 Subject: [PATCH 4/4] damn flake8 --- junifer/data/tests/test_parcellations.py | 1 - junifer/markers/tests/test_parcel_aggregation.py | 8 ++++---- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/junifer/data/tests/test_parcellations.py b/junifer/data/tests/test_parcellations.py index 778ebeaac..1e04f10f4 100644 --- a/junifer/data/tests/test_parcellations.py +++ b/junifer/data/tests/test_parcellations.py @@ -119,7 +119,6 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None: load_parcellation("WrongValues2") - @pytest.mark.parametrize( "name, parcellation_path, parcels_labels, overwrite", [ diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index 710caa41a..cace74bba 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -293,7 +293,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None: meta = marker_original.get_meta("VBM_GM")["marker"] assert meta["method"] == "mean" assert meta["parcellation"] == ["Schaefer100x7"] - assert meta["mask"] == None + assert meta["mask"] is None assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" @@ -317,7 +317,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None: meta = marker_split.get_meta("VBM_GM")["marker"] assert meta["method"] == "mean" assert meta["parcellation"] == ["Schaefer100x7_low", "Schaefer100x7_high"] - assert meta["mask"] == None + assert meta["mask"] is None assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" @@ -391,7 +391,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None: meta = marker_original.get_meta("VBM_GM")["marker"] assert meta["method"] == "mean" assert meta["parcellation"] == ["Schaefer100x7"] - assert meta["mask"] == None + assert meta["mask"] is None assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" @@ -418,7 +418,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None: "Schaefer100x7_low2", "Schaefer100x7_high2", ] - assert meta["mask"] == None + assert meta["mask"] is None assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" -- 2.52.0