diff --git a/docs/api/data.rst b/docs/api/data.rst index 68f77b3b2..842ce47af 100644 --- a/docs/api/data.rst +++ b/docs/api/data.rst @@ -9,3 +9,10 @@ Coordinates .. automodule:: junifer.data.coordinates :members: + + +Masks +===== + +.. automodule:: junifer.data.masks + :members: diff --git a/docs/builtin.rst b/docs/builtin.rst index b0a79244f..70b95e5f3 100644 --- a/docs/builtin.rst +++ b/docs/builtin.rst @@ -321,5 +321,42 @@ Available Planned ~~~~~~~ + +Masks +----- + +.. + Provide a list of the masks that are implemented or planned. + + Version added: The Junifer version in which the mask was added. + +Available +~~~~~~~~~ + +.. list-table:: + :widths: auto + :header-rows: 1 + + * - Name + - Keys + - Version added + - Publication + * - Vickery-Patil (Gray Matter) + - | ``GM_prob0.2`` + - 0.0.1 + - | Vickery, Sam, & Patil, Kaustubh. (2022). + | Chimpanzee and Human Gray Matter Masks [Data set]. Zenodo. + | https://doi.org/10.5281/zenodo.6463123 + * - Vickery-Patil (Cortex + Basal Ganglia) + - | ``GM_prob0.2_cortex`` + - 0.0.1 + - | Vickery, Sam, & Patil, Kaustubh. (2022). + | Chimpanzee and Human Gray Matter Masks [Data set]. Zenodo. + | https://doi.org/10.5281/zenodo.6463123 + +Planned +~~~~~~~ + + .. helpful site for creating tables: https://rest-sphinx-memo.readthedocs.io/en/latest/ReST.html#tables diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index baa2ce7fb..c613a276a 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -87,6 +87,8 @@ Enhancements - Allow custom aggregation method for :class:`junifer.markers.SphereAggregation` (:gh:`102` by `Synchon Mandal`_). +- Add support for "masks" (:gh:`79` by `Fede Raimondo`_). + Bugs ~~~~ diff --git a/junifer/data/__init__.py b/junifer/data/__init__.py index 1ed8977f3..f1508945b 100644 --- a/junifer/data/__init__.py +++ b/junifer/data/__init__.py @@ -14,3 +14,11 @@ from .parcellations import ( load_parcellation, register_parcellation, ) + +from .masks import ( + list_masks, + load_mask, + register_mask, +) + +from . import utils \ No newline at end of file diff --git a/junifer/data/masks.py b/junifer/data/masks.py new file mode 100644 index 000000000..710ad0bde --- /dev/null +++ b/junifer/data/masks.py @@ -0,0 +1,186 @@ +"""Provide functions for masks.""" + +# Authors: Federico Raimondo +# License: AGPL + +from pathlib import Path +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union + +import nibabel as nib + +from .utils import closest_resolution +from ..utils.logging import logger, raise_error + + +if TYPE_CHECKING: + from nibabel import Nifti1Image + +# Path to the VOIs +_masks_path = Path(__file__).parent / "masks" + +""" +A dictionary containing all supported masks and their respective file or +data. + +The built-in masks are files that are shipped with the package in the +data/masks directory. The user can also register their own masks. +""" +_available_masks: Dict[str, Dict[str, Any]] = { + "GM_prob0.2": {"family": "Vickery-Patil"}, + "GM_prob0.2_cortex": {"family": "Vickery-Patil"}, +} + + +def register_mask( + name: str, + mask_path: Union[str, Path], + overwrite: bool = False, +) -> None: + """Register a custom user mask. + + Parameters + ---------- + name : str + The name of the mask. + mask_path : str or pathlib.Path + The path to the mask file. + overwrite : bool, optional + If True, overwrite an existing mask with the same name. + Does not apply to built-in mask (default False). + + Raises + ------ + ValueError + If the mask name is already registered and overwrite is set to + False or if the mask name is a built-in mask. + """ + # Check for attempt of overwriting built-in parcellations + if name in _available_masks: + if overwrite is True: + logger.info(f"Overwriting {name} mask") + if (_available_masks[name]["family"] != "CustomUserMask"): + raise_error( + f"Cannot overwrite {name} mask. " + "It is a built-in mask." + ) + else: + raise_error( + f"Mask {name} already registered. Set `overwrite=True`" + "to update its value." + ) + # Convert str to Path + if not isinstance(mask_path, Path): + mask_path = Path(mask_path) + # Add user parcellation info + _available_masks[name] = { + "path": str(mask_path.absolute()), + "family": "CustomUserMask", + } + + +def list_masks() -> List[str]: + """List all the available masks. + + Returns + ------- + list of str + A list with all available masks names. + """ + return sorted(_available_masks.keys()) + + +def load_mask( + name: str, + resolution: Optional[float] = None, + path_only: bool = False, +) -> Tuple[Optional["Nifti1Image"], Path]: + """Load mask. + + Parameters + ---------- + name : str + The name of the mask. + resolution : float, optional + The desired resolution of the mask to load. If it is not + available, the closest resolution will be loaded. Preferably, use a + resolution higher than the desired one. By default, will load the + highest one (default None). + path_only : bool, optional + If True, the mask image will not be loaded (default False). + + Returns + ------- + Nifti1Image or None + Loaded mask image. + pathlib.Path + File path to the mask image. + """ + if name not in _available_masks: + raise_error( + f"Mask {name} not found. " + f"Valid options are: {list_masks()}" + ) + + mask_definition = _available_masks[name].copy() + t_family = mask_definition.pop("family") + + if t_family == "CustomUserMask": + mask_fname = Path(mask_definition["path"]) + elif t_family == 'Vickery-Patil': + mask_fname = _load_vickery_patil_mask(name, resolution) + else: + raise_error( + f"I don't know about the {t_family} mask family." + ) + + logger.info(f"Loading mask {mask_fname.absolute()}") + + mask_img = None + if path_only is False: + mask_img = nib.load(mask_fname) + + return mask_img, mask_fname + + +def _load_vickery_patil_mask( + name: str, + resolution: Optional[float] = None, +) -> Path: + """Load Vickery-Patil mask. + + Parameters + ---------- + name : str + The name of the mask. + resolution : float, optional + The desired resolution of the mask to load. If it is not + available, the closest resolution will be loaded. Preferably, use a + resolution higher than the desired one. By default, will load the + highest one (default None). + + Returns + ------- + pathlib.Path + File path to the mask image. + """ + if name == "GM_prob0.2": + available_resolutions = [1.5, 3.0] + to_load = closest_resolution(resolution, available_resolutions) + if to_load == 3.0: + mask_fname = \ + "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz" + elif to_load == 1.5: + mask_fname = "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz" + else: + raise_error( + f"Cannot find a GM_prob0.2 mask for resolution {resolution}" + ) + elif name == "GM_prob0.2_cortex": + mask_fname = "GMprob0.2_cortex_3mm_NA_rm.nii.gz" + else: + raise_error( + f"Cannot find a Vickery-Patil mask called {name}" + ) + mask_fname = _masks_path / "vickery-patil" / mask_fname + + return mask_fname diff --git a/junifer/data/masks/vickery-patil/CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz b/junifer/data/masks/vickery-patil/CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz new file mode 100644 index 000000000..329197283 Binary files /dev/null and b/junifer/data/masks/vickery-patil/CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz differ diff --git a/junifer/data/masks/vickery-patil/CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz b/junifer/data/masks/vickery-patil/CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz new file mode 100644 index 000000000..597c1d44f Binary files /dev/null and b/junifer/data/masks/vickery-patil/CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz differ diff --git a/junifer/data/masks/vickery-patil/GMprob0.2_cortex_3mm_NA_rm.nii.gz b/junifer/data/masks/vickery-patil/GMprob0.2_cortex_3mm_NA_rm.nii.gz new file mode 100644 index 000000000..cd749ffbc Binary files /dev/null and b/junifer/data/masks/vickery-patil/GMprob0.2_cortex_3mm_NA_rm.nii.gz differ diff --git a/junifer/data/parcellations.py b/junifer/data/parcellations.py index e1ac64498..26629d636 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -18,6 +18,7 @@ import pandas as pd import requests from nilearn import datasets +from .utils import closest_resolution from ..utils.logging import logger, raise_error if TYPE_CHECKING: @@ -315,41 +316,6 @@ def _retrieve_parcellation( return parcellation_fname, parcellation_labesl -def _closest_resolution( - resolution: Optional[float], - valid_resolution: Union[List[float], List[int], np.ndarray], -) -> Union[float, int]: - """Find the closest resolution. - - Parameters - ---------- - resolution : float - The given resolution. - valid_resolution : list of float or np.ndarray - The array of valid resolutions. - - Returns - ------- - float - The closest valid resolution. - """ - # Convert list of int to numpy.ndarray - if not isinstance(valid_resolution, np.ndarray): - valid_resolution = np.array(valid_resolution) - - if resolution is None: - logger.info("Resolution set to None, using highest resolution.") - closest = np.min(valid_resolution) - elif any(x <= resolution for x in valid_resolution): - # Case 1: get the highest closest resolution - closest = np.max(valid_resolution[valid_resolution <= resolution]) - else: - # Case 2: get the lower closest resolution - closest = np.min(valid_resolution) - - return closest - - def _retrieve_schaefer( parcellations_dir: Path, resolution: Optional[float] = None, @@ -406,7 +372,7 @@ def _retrieve_schaefer( f"of the following: {_valid_networks}" ) - resolution = _closest_resolution(resolution, _valid_resolutions) + resolution = closest_resolution(resolution, _valid_resolutions) # define file names parcellation_fname = ( @@ -532,7 +498,7 @@ def _retrieve_tian( f"one of the following: 3T or 7T" ) - resolution = _closest_resolution(resolution, _valid_resolutions) + resolution = closest_resolution(resolution, _valid_resolutions) # define file names if magneticfield == "3T": @@ -667,7 +633,7 @@ def _retrieve_suit( # TODO: Validate this with Vera _valid_resolutions = [1] - resolution = _closest_resolution(resolution, _valid_resolutions) + resolution = closest_resolution(resolution, _valid_resolutions) # define file names parcellation_fname = ( diff --git a/junifer/data/tests/test_data_utils.py b/junifer/data/tests/test_data_utils.py new file mode 100644 index 000000000..80b2adf66 --- /dev/null +++ b/junifer/data/tests/test_data_utils.py @@ -0,0 +1,43 @@ +"""Provide tests for data utils.""" + +# Authors: Federico Raimondo +# License: AGPL + +from typing import List + +import pytest +import numpy as np + +from junifer.data.utils import closest_resolution + + +@pytest.mark.parametrize( + "resolution, valid_resolutions, expected", + [ + (1.0, [1.0, 2.0, 3.0], 1.0), + (1.1, [1.0, 2.0, 3.0], 1.0), + (0.9, [1.0, 2.0, 3.0], 1.0), + (2.1, [1.0, 2.0, 3.0], 2.0), + (2.0, [1.0, 2.0, 3.0], 2.0), + (4.0, [1.0, 2.0, 3.0], 3.0), + (None, [1.0, 2.0, 3.0], 1.0), + ], +) +def test_closest_resolution( + resolution: float, valid_resolutions: List[float], expected: float +): + """Test closest_resolution. + + Parameters + ---------- + resolution: float + The resolution to test. + valid_resolutions: list of float + The valid resolutions. + expected: float + The expected result. + """ + assert closest_resolution(resolution, valid_resolutions) == expected + assert ( + closest_resolution(resolution, np.array(valid_resolutions)) == expected + ) diff --git a/junifer/data/tests/test_masks.py b/junifer/data/tests/test_masks.py new file mode 100644 index 000000000..13926774f --- /dev/null +++ b/junifer/data/tests/test_masks.py @@ -0,0 +1,153 @@ +"""Provide tests for masks.""" + +# Authors: Federico Raimondo +# Vera Komeyer +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +import pytest +from numpy.testing import assert_array_almost_equal + +from junifer.data.masks import ( + load_mask, + register_mask, + list_masks, + _load_vickery_patil_mask, +) + + +def test_register_mask_built_in_check() -> None: + """Test mask registration check for built-in masks.""" + with pytest.raises(ValueError, match=r"built-in mask"): + register_mask( + name="GM_prob0.2", + mask_path="testmask.nii.gz", + overwrite=True, + ) + + +def test_list_masks_incorrect() -> None: + """Test incorrect information check for list masks.""" + masks = list_masks() + assert "testmask" not in masks + + +def test_register_mask_already_registered() -> None: + """Test mask registration check for already registered.""" + # Register custom mask + register_mask( + name="testmask", + mask_path="testmask.nii.gz", + ) + assert load_mask("testmask", path_only=True)[1].name == "testmask.nii.gz" + + # Try registering again + with pytest.raises(ValueError, match=r"already registered."): + register_mask( + name="testmask", + mask_path="testmask.nii.gz", + ) + register_mask( + name="testmask", + mask_path="testmask2.nii.gz", + overwrite=True, + ) + + assert load_mask("testmask", path_only=True)[1].name == "testmask2.nii.gz" + + +@pytest.mark.parametrize( + "name, mask_path, overwrite", + [ + ("testmask_1", "testmask_1.nii.gz", True), + ("testmask_2", "testmask_2.nii.gz", True), + ("testmask_3", Path("testmask_3.nii.gz"), True), + ], +) +def test_register_mask( + name: str, + mask_path: str, + overwrite: bool, +) -> None: + """Test mask registration. + + Parameters + ---------- + name : str + The parametrized mask name. + mask_path : str or pathlib.Path + The parametrized mask path. + overwrite : bool + The parametrized mask overwrite value. + + """ + # Register custom mask + register_mask( + name=name, + mask_path=mask_path, + overwrite=overwrite, + ) + # List available mask and check registration + masks = list_masks() + assert name in masks + # Load registered mask + _, fname = load_mask(name=name, path_only=True) + # Check values for registered mask + assert fname.name == f"{name}.nii.gz" + + +@pytest.mark.parametrize( + "mask_name", + [ + "GM_prob0.2", + "GM_prob0.2_cortex", + ], +) +def test_list_masks_correct(mask_name: str) -> None: + """Test correct information check for list masks. + + Parameters + ---------- + mask_name : str + The parametrized mask name. + + """ + masks = list_masks() + assert mask_name in masks + + +def test_load_mask_incorrect() -> None: + """Test loading of invalid masks.""" + with pytest.raises(ValueError, match=r"not found"): + load_mask("wrongmask") + + +def test_vickery_patil() -> None: + """Test Vickery-Patil mask.""" + mask, fname = load_mask("GM_prob0.2") + assert_array_almost_equal( + mask.header["pixdim"][1:4], [1.5, 1.5, 1.5] # type: ignore + ) + + assert fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz" + + mask, fname = load_mask("GM_prob0.2", resolution=3) + assert_array_almost_equal( + mask.header["pixdim"][1:4], [3.0, 3.0, 3.0] # type: ignore + ) + + assert ( + fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz" + ) + + mask, fname = load_mask("GM_prob0.2_cortex") + assert_array_almost_equal( + mask.header["pixdim"][1:4], [3.0, 3.0, 3.0] # type: ignore + ) + + assert fname.name == "GMprob0.2_cortex_3mm_NA_rm.nii.gz" + + with pytest.raises(ValueError, match=r"find a Vickery-Patil mask "): + _load_vickery_patil_mask("wrong", resolution=2) diff --git a/junifer/data/utils.py b/junifer/data/utils.py new file mode 100644 index 000000000..56450e0dd --- /dev/null +++ b/junifer/data/utils.py @@ -0,0 +1,42 @@ +"""Provide utilities for data module.""" +from typing import Optional, Union, List + +import numpy as np + +from ..utils.logging import logger + + +def closest_resolution( + resolution: Optional[float], + valid_resolution: Union[List[float], List[int], np.ndarray], +) -> Union[float, int]: + """Find the closest resolution. + + Parameters + ---------- + resolution : float, optional + The given resolution. If None, will return the highest resolution + (default None). + valid_resolution : list of float or int, or np.ndarray + The array of valid resolutions. + + Returns + ------- + float or int + The closest valid resolution. + """ + # Convert list of int to numpy.ndarray + if not isinstance(valid_resolution, np.ndarray): + valid_resolution = np.array(valid_resolution) + + if resolution is None: + logger.info("Resolution set to None, using highest resolution.") + closest = np.min(valid_resolution) + elif any(x <= resolution for x in valid_resolution): + # Case 1: get the highest closest resolution + closest = np.max(valid_resolution[valid_resolution <= resolution]) + else: + # Case 2: get the lower closest resolution + closest = np.min(valid_resolution) + + return closest diff --git a/junifer/markers/crossparcellation_functional_connectivity.py b/junifer/markers/crossparcellation_functional_connectivity.py index 36b2f133e..4e8ccbad7 100644 --- a/junifer/markers/crossparcellation_functional_connectivity.py +++ b/junifer/markers/crossparcellation_functional_connectivity.py @@ -32,6 +32,10 @@ class CrossParcellationFC(BaseMarker): correlation_method : str, optional Any method that can be passed to :any:`pandas.DataFrame.corr` (default "pearson"). + mask : str, optional + The name of the mask to apply to regions before extracting signals. + Check valid options by calling :func:`junifer.data.masks.list_masks` + (default None). name : str, optional The name of the marker. If None, will use the class name (default None). @@ -43,6 +47,7 @@ class CrossParcellationFC(BaseMarker): parcellation_two: str, aggregation_method: str = "mean", correlation_method: str = "pearson", + mask: Optional[str] = None, name: Optional[str] = None, ) -> None: if parcellation_one == parcellation_two: @@ -53,6 +58,7 @@ class CrossParcellationFC(BaseMarker): self.parcellation_two = parcellation_two self.aggregation_method = aggregation_method self.correlation_method = correlation_method + self.mask = mask super().__init__(on=["BOLD"], name=name) def get_valid_inputs(self) -> List[str]: @@ -145,10 +151,12 @@ class CrossParcellationFC(BaseMarker): parcellation_one_dict = ParcelAggregation( parcellation=self.parcellation_one, method=self.aggregation_method, + mask=self.mask, ).compute(input) parcellation_two_dict = ParcelAggregation( parcellation=self.parcellation_two, method=self.aggregation_method, + mask=self.mask, ).compute(input) parcellated_ts_one = parcellation_one_dict["data"] diff --git a/junifer/markers/ets_rss.py b/junifer/markers/ets_rss.py index 532bbb0e1..50aff68bd 100644 --- a/junifer/markers/ets_rss.py +++ b/junifer/markers/ets_rss.py @@ -29,9 +29,16 @@ class RSSETSMarker(BaseMarker): parcellation : str The name of the parcellation. Check valid options by calling :func:`junifer.data.parcellations.list_parcellations`. - aggregation_method : str, optional + agg_method : str, optional The method to perform aggregation using. Check valid options in :func:`junifer.stats.get_aggfunc_by_name` (default "mean"). + agg_method_params : dict, optional + Parameters to pass to the aggregation function. Check valid options in + :func:`junifer.stats.get_aggfunc_by_name` (default None). + mask : str, optional + The name of the mask to apply to regions before extracting signals. + Check valid options by calling :func:`junifer.data.masks.list_masks` + (default None). name : str, optional The name of the marker. If None, will use the class name (default None). @@ -41,11 +48,15 @@ class RSSETSMarker(BaseMarker): def __init__( self, parcellation: str, - aggregation_method: str = "mean", + agg_method: str = "mean", + agg_method_params: Optional[Dict] = None, + mask: Optional[str] = None, name: Optional[str] = None, ) -> None: self.parcellation = parcellation - self.aggregation_method = aggregation_method + self.agg_method = agg_method + self.agg_method_params = agg_method_params + self.mask = mask super().__init__(name=name) def get_valid_inputs(self) -> List[str]: @@ -136,7 +147,9 @@ class RSSETSMarker(BaseMarker): # Initialize a ParcelAggregation parcel_aggregation = ParcelAggregation( parcellation=self.parcellation, - method=self.aggregation_method, + method=self.agg_method, + method_params=self.agg_method_params, + mask=self.mask ) # Compute the parcel aggregation out = parcel_aggregation.compute(input=input, extra_input=extra_input) diff --git a/junifer/markers/functional_connectivity_parcels.py b/junifer/markers/functional_connectivity_parcels.py index 5c354c76d..3bb94216a 100644 --- a/junifer/markers/functional_connectivity_parcels.py +++ b/junifer/markers/functional_connectivity_parcels.py @@ -40,6 +40,10 @@ class FunctionalConnectivityParcels(BaseMarker): cor_method_params : dict, optional Parameters to pass to the correlation function. Check valid options in :class:`nilearn.connectome.ConnectivityMeasure` (default None). + mask : str, optional + The name of the mask to apply to regions before extracting signals. + Check valid options by calling :func:`junifer.data.masks.list_masks` + (default None). name : str, optional The name of the marker. If None, will use the class name (default None). @@ -52,21 +56,20 @@ class FunctionalConnectivityParcels(BaseMarker): agg_method_params: Optional[Dict] = None, cor_method: str = "covariance", cor_method_params: Optional[Dict] = None, + mask: Optional[str] = None, name: Optional[str] = None, ) -> None: self.parcellation = parcellation self.agg_method = agg_method - self.agg_method_params = ( - {} if agg_method_params is None else agg_method_params - ) + self.agg_method_params = agg_method_params self.cor_method = cor_method - self.cor_method_params = ( - {} if cor_method_params is None else cor_method_params - ) + self.cor_method_params = cor_method_params or {} + # default to nilearn behavior self.cor_method_params["empirical"] = self.cor_method_params.get( "empirical", False ) + self.mask = mask super().__init__(name=name) @@ -131,6 +134,7 @@ class FunctionalConnectivityParcels(BaseMarker): parcellation=self.parcellation, method=self.agg_method, method_params=self.agg_method_params, + mask=self.mask, on="BOLD", ) # get the 2D timeseries after parcel aggregation diff --git a/junifer/markers/functional_connectivity_spheres.py b/junifer/markers/functional_connectivity_spheres.py index ba3c6e3e4..193510220 100644 --- a/junifer/markers/functional_connectivity_spheres.py +++ b/junifer/markers/functional_connectivity_spheres.py @@ -43,6 +43,10 @@ class FunctionalConnectivitySpheres(BaseMarker): cor_method_params : dict, optional Parameters to pass to the correlation function. Check valid options in :class:`nilearn.connectome.ConnectivityMeasure` (default None). + mask : str, optional + The name of the mask to apply to regions before extracting signals. + Check valid options by calling :func:`junifer.data.masks.list_masks` + (default None). name : str, optional The name of the marker. By default, it will use KIND_FunctionalConnectivitySpheres where KIND is the kind of data it @@ -58,6 +62,7 @@ class FunctionalConnectivitySpheres(BaseMarker): agg_method_params: Optional[Dict] = None, cor_method: str = "covariance", cor_method_params: Optional[Dict] = None, + mask: Optional[str] = None, name: Optional[str] = None, ) -> None: self.coords = coords @@ -65,18 +70,17 @@ class FunctionalConnectivitySpheres(BaseMarker): if radius is None or radius <= 0: raise_error(f"radius should be > 0: provided {radius}") self.agg_method = agg_method - self.agg_method_params = ( - {} if agg_method_params is None else agg_method_params - ) + self.agg_method_params = agg_method_params self.cor_method = cor_method - self.cor_method_params = ( - {} if cor_method_params is None else cor_method_params - ) + self.cor_method_params = cor_method_params or {} + # default to nilearn behavior self.cor_method_params["empirical"] = self.cor_method_params.get( "empirical", False ) + self.mask = mask + super().__init__(name=name) def get_valid_inputs(self) -> List[str]: @@ -142,6 +146,7 @@ class FunctionalConnectivitySpheres(BaseMarker): radius=self.radius, method=self.agg_method, method_params=self.agg_method_params, + mask=self.mask, on="BOLD", ) diff --git a/junifer/markers/parcel_aggregation.py b/junifer/markers/parcel_aggregation.py index deb98e028..f8b4ca8bc 100644 --- a/junifer/markers/parcel_aggregation.py +++ b/junifer/markers/parcel_aggregation.py @@ -11,7 +11,7 @@ from nilearn.image import math_img, resample_to_img from nilearn.maskers import NiftiMasker from ..api.decorators import register_marker -from ..data import load_parcellation +from ..data import load_parcellation, load_mask from ..stats import get_aggfunc_by_name from ..utils import logger from .base import BaseMarker @@ -35,6 +35,10 @@ class ParcelAggregation(BaseMarker): method_params : dict, optional Parameters to pass to the aggregation function. Check valid options in :func:`junifer.stats.get_aggfunc_by_name`. + mask : str, optional + The name of the mask to apply to regions before extracting signals. + Check valid options by calling :func:`junifer.data.masks.list_masks` + (default None). on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} \ or list of the options, optional The data types to apply the marker to. If None, will work on all @@ -49,12 +53,14 @@ class ParcelAggregation(BaseMarker): parcellation: 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: self.parcellation = parcellation self.method = method - self.method_params = {} if method_params is None else method_params + self.method_params = method_params or {} + self.mask = mask super().__init__(on=on, name=name) def get_valid_inputs(self) -> List[str]: @@ -158,15 +164,34 @@ class ParcelAggregation(BaseMarker): name=self.parcellation, resolution=resolution, ) + parcellation_img_res = resample_to_img( t_parcellation, t_input, interpolation="nearest", + copy=True, ) + parcellation_bin = math_img( "img != 0", img=parcellation_img_res, ) + + if self.mask is not None: + logger.debug(f"Masking with {self.mask}") + mask_img, _ = load_mask(name=self.mask, resolution=resolution) + mask_img = resample_to_img( + mask_img, + t_input, + interpolation="nearest", + copy=True, + ) + parcellation_bin = math_img( + "np.logical_and(img, mask)", + img=parcellation_bin, + mask=mask_img, + ) + logger.debug("Masking") masker = NiftiMasker( parcellation_bin, target_affine=t_input.affine diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index ba3cff9ef..8fa648b4c 100644 --- a/junifer/markers/sphere_aggregation.py +++ b/junifer/markers/sphere_aggregation.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from ..api.decorators import register_marker -from ..data import load_coordinates +from ..data import load_coordinates, load_mask from ..external.nilearn import JuniferNiftiSpheresMasker from ..stats import get_aggfunc_by_name from ..utils import logger @@ -37,6 +37,10 @@ class SphereAggregation(BaseMarker): (default "mean"). method_params : dict, optional The parameters to pass to the aggregation method (default None). + mask : str, optional + The name of the mask to apply to regions before extracting signals. + Check valid options by calling :func:`junifer.data.masks.list_masks` + (default None). on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} or \ list of the options, optional The data types to apply the marker to. If None, will work on all @@ -53,13 +57,15 @@ class SphereAggregation(BaseMarker): radius: Optional[float] = None, method: str = "mean", method_params: Optional[Dict[str, Any]] = None, + mask: Optional[str] = None, on: Union[List[str], str, None] = None, name: Optional[str] = None, ) -> None: self.coords = coords self.radius = radius self.method = method - self.method_params = {} if method_params is None else method_params + self.method_params = method_params or {} + self.mask = mask super().__init__(on=on, name=name) def get_valid_inputs(self) -> List[str]: @@ -157,12 +163,17 @@ class SphereAggregation(BaseMarker): agg_func = get_aggfunc_by_name( self.method, func_params=self.method_params ) + # Load mask + mask_img = None + if self.mask is not None: + logger.debug(f"Masking with {self.mask}") + mask_img, _ = load_mask(self.mask) # Get seeds and labels coords, out_labels = load_coordinates(name=self.coords) masker = JuniferNiftiSpheresMasker( seeds=coords, radius=self.radius, - mask_img=None, # TODO: support this (needs #79) + mask_img=mask_img, agg_func=agg_func, ) # Fit and transform the marker on the data diff --git a/junifer/markers/tests/test_ets_rss.py b/junifer/markers/tests/test_ets_rss.py index ed6fd484e..d977b226b 100644 --- a/junifer/markers/tests/test_ets_rss.py +++ b/junifer/markers/tests/test_ets_rss.py @@ -45,7 +45,8 @@ def test_compute() -> None: # Assert the meta meta = ets_rss_marker.get_meta("BOLD")["marker"] assert meta["parcellation"] == "Schaefer100x17" - assert meta["aggregation_method"] == "mean" + assert meta["agg_method"] == "mean" + assert meta["agg_method_params"] is None assert meta["class"] == "RSSETSMarker" diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index 0bbc95109..d8abd0d58 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -12,6 +12,7 @@ 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.markers.parcel_aggregation import ParcelAggregation @@ -86,6 +87,7 @@ def test_ParcelAggregation_3D() -> None: meta = marker.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" @@ -110,6 +112,7 @@ def test_ParcelAggregation_3D() -> None: meta = marker.get_meta("VBM_GM")["marker"] assert meta["method"] == "std" assert meta["parcellation"] == "Schaefer100x7" + assert meta["mask"] is None assert meta["name"] == "VBM_GM_ParcelAggregation" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" @@ -142,6 +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["mask"] is None assert meta["name"] == "VBM_GM_ParcelAggregation" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "VBM_GM" @@ -175,7 +179,53 @@ def test_ParcelAggregation_4D(): meta = marker.get_meta("BOLD")["marker"] assert meta["method"] == "mean" assert meta["parcellation"] == "Schaefer100x7" + assert meta["mask"] is None assert meta["name"] == "BOLD_ParcelAggregation" assert meta["class"] == "ParcelAggregation" assert meta["kind"] == "BOLD" assert meta["method_params"] == {} + + +def test_ParcelAggregation_3D_mask() -> None: + """Test ParcelAggregation object on 3D images with mask.""" + + # Get the testing parcellation (for nilearn) + parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) + + # Get one mask + mask_img, _ = load_mask("GM_prob0.2") + + # 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 + nifti_masker = NiftiLabelsMasker( + labels_img=parcellation.maps, + mask_img=mask_img) + auto = nifti_masker.fit_transform(img) + + # Use the ParcelAggregation object + marker = 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"] + + 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" + 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"] == {} diff --git a/junifer/markers/tests/test_sphere_aggregation.py b/junifer/markers/tests/test_sphere_aggregation.py index 08ea00066..fd8f538a9 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -12,7 +12,7 @@ from nilearn.image import concat_imgs from nilearn.maskers import NiftiSpheresMasker from numpy.testing import assert_array_equal -from junifer.data import load_coordinates +from junifer.data import load_coordinates, load_mask from junifer.markers.sphere_aggregation import SphereAggregation from junifer.storage import SQLiteFeatureStorage @@ -44,7 +44,7 @@ def test_SphereAggregation_3D() -> None: vbm = oasis_dataset.gray_matter_maps[0] img = nib.load(vbm) - # Create NiftiLabelsMasker + # Create NiftSpheresMasker nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) auto4d = nifti_masker.fit_transform(img) @@ -63,6 +63,7 @@ def test_SphereAggregation_3D() -> None: assert meta["method"] == "mean" assert meta["coords"] == COORDS assert meta["radius"] == RADIUS + assert meta["mask"] is None assert meta["name"] == "VBM_GM_SphereAggregation" assert meta["class"] == "SphereAggregation" assert meta["kind"] == "VBM_GM" @@ -72,13 +73,13 @@ def test_SphereAggregation_3D() -> None: def test_SphereAggregation_4D() -> None: """Test SphereAggregation object on 4D images.""" # Get the testing coordinates (for nilearn) - coordinates, labels = load_coordinates(COORDS) + coordinates, _ = load_coordinates(COORDS) # Get the SPM auditory data subject_data = datasets.fetch_spm_auditory() fmri_img = concat_imgs(subject_data.func) # type: ignore - # Create NiftiLabelsMasker + # Create NiftSpheresMasker nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) auto4d = nifti_masker.fit_transform(fmri_img) @@ -97,6 +98,7 @@ def test_SphereAggregation_4D() -> None: assert meta["method"] == "mean" assert meta["coords"] == COORDS assert meta["radius"] == RADIUS + assert meta["mask"] is None assert meta["name"] == "BOLD_SphereAggregation" assert meta["class"] == "SphereAggregation" assert meta["kind"] == "BOLD" @@ -145,3 +147,43 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None: ) marker.fit_transform(input, storage=storage) + + +def test_SphereAggregation_3D_mask() -> None: + """Test SphereAggregation object on 3D images using mask.""" + # Get the testing coordinates (for nilearn) + coordinates, _ = load_coordinates(COORDS) + + # Get one mask + mask_img, _ = load_mask("GM_prob0.2") + + # 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 NiftSpheresMasker + nifti_masker = NiftiSpheresMasker( + seeds=coordinates, radius=RADIUS, mask_img=mask_img) + auto4d = nifti_masker.fit_transform(img) + + # Create SphereAggregation object + marker = SphereAggregation( + coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM", + mask="GM_prob0.2" + ) + input = {"VBM_GM": {"data": img}} + jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"] + + assert jun_values4d.ndim == 2 + assert_array_equal(auto4d.shape, jun_values4d.shape) + assert_array_equal(auto4d, jun_values4d) + + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "mean" + assert meta["coords"] == COORDS + assert meta["radius"] == RADIUS + assert meta["name"] == "VBM_GM_SphereAggregation" + assert meta["class"] == "SphereAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {}