[ENH]: Allow ParcelAggregation to apply multiple parcellations at once #131
7 changed files with 329 additions and 32 deletions
|
|
@ -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
|
||||||
~~~~
|
~~~~
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
[
|
[
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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}
|
||||||
|
|
|
||||||
|
|
@ -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:])
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue