[ENH]: Add support for masks (GM/WM/etc) #79

Merged
fraimondo merged 8 commits from enh/masks into main 2022-11-23 16:07:50 +00:00
21 changed files with 667 additions and 64 deletions

View file

@ -9,3 +9,10 @@ Coordinates
.. automodule:: junifer.data.coordinates .. automodule:: junifer.data.coordinates
:members: :members:
Masks
=====
.. automodule:: junifer.data.masks
:members:

View file

@ -321,5 +321,42 @@ Available
Planned 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 helpful site for creating tables: https://rest-sphinx-memo.readthedocs.io/en/latest/ReST.html#tables

View file

@ -87,6 +87,8 @@ Enhancements
- Allow custom aggregation method for :class:`junifer.markers.SphereAggregation` (:gh:`102` by `Synchon Mandal`_). - Allow custom aggregation method for :class:`junifer.markers.SphereAggregation` (:gh:`102` by `Synchon Mandal`_).
- Add support for "masks" (:gh:`79` by `Fede Raimondo`_).
Bugs Bugs
~~~~ ~~~~

View file

@ -14,3 +14,11 @@ from .parcellations import (
load_parcellation, load_parcellation,
register_parcellation, register_parcellation,
) )
from .masks import (
list_masks,
load_mask,
register_mask,
)
from . import utils

186
junifer/data/masks.py Normal file
View file

@ -0,0 +1,186 @@
"""Provide functions for masks."""
synchon commented 2022-11-23 10:44:50 +00:00 (Migrated from github.com)

Dict[str, Dict[str, str]]?

`Dict[str, Dict[str, str]]`?
fraimondo commented 2022-11-23 14:35:16 +00:00 (Migrated from github.com)

not entirely. Inner values could be anything

not entirely. Inner values could be anything
synchon commented 2022-11-23 15:04:58 +00:00 (Migrated from github.com)

But shouldn't be the inner dict keys be string?

But shouldn't be the inner dict keys be string?
synchon commented 2022-11-23 15:05:57 +00:00 (Migrated from github.com)

Okay now, saw the change.

Okay now, saw the change.
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# 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

View file

@ -18,6 +18,7 @@ import pandas as pd
import requests import requests
from nilearn import datasets from nilearn import datasets
from .utils import closest_resolution
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
if TYPE_CHECKING: if TYPE_CHECKING:
@ -315,41 +316,6 @@ def _retrieve_parcellation(
return parcellation_fname, parcellation_labesl 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( def _retrieve_schaefer(
parcellations_dir: Path, parcellations_dir: Path,
resolution: Optional[float] = None, resolution: Optional[float] = None,
@ -406,7 +372,7 @@ def _retrieve_schaefer(
f"of the following: {_valid_networks}" f"of the following: {_valid_networks}"
) )
resolution = _closest_resolution(resolution, _valid_resolutions) resolution = closest_resolution(resolution, _valid_resolutions)
# define file names # define file names
parcellation_fname = ( parcellation_fname = (
@ -532,7 +498,7 @@ def _retrieve_tian(
f"one of the following: 3T or 7T" f"one of the following: 3T or 7T"
) )
resolution = _closest_resolution(resolution, _valid_resolutions) resolution = closest_resolution(resolution, _valid_resolutions)
# define file names # define file names
if magneticfield == "3T": if magneticfield == "3T":
@ -667,7 +633,7 @@ def _retrieve_suit(
# TODO: Validate this with Vera # TODO: Validate this with Vera
_valid_resolutions = [1] _valid_resolutions = [1]
resolution = _closest_resolution(resolution, _valid_resolutions) resolution = closest_resolution(resolution, _valid_resolutions)
# define file names # define file names
parcellation_fname = ( parcellation_fname = (

View file

@ -0,0 +1,43 @@
"""Provide tests for data utils."""
synchon commented 2022-11-23 10:48:35 +00:00 (Migrated from github.com)

For consistency: the docstring is missing the Parameters section.

For consistency: the docstring is missing the Parameters section.
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# 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(
synchon commented 2022-11-23 15:07:07 +00:00 (Migrated from github.com)

list of float

`list of float`
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
)

View file

@ -0,0 +1,153 @@
"""Provide tests for masks."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Vera Komeyer <v.komeyer@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

42
junifer/data/utils.py Normal file
View file

@ -0,0 +1,42 @@
"""Provide utilities for data module."""
synchon commented 2022-11-23 10:52:49 +00:00 (Migrated from github.com)

float or None?

float or None?
synchon commented 2022-11-23 10:53:26 +00:00 (Migrated from github.com)

list of float or int, or np.ndarray?

list of float or int, or np.ndarray?
synchon commented 2022-11-23 10:53:36 +00:00 (Migrated from github.com)

float or int

float or int
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

View file

@ -32,6 +32,10 @@ class CrossParcellationFC(BaseMarker):
correlation_method : str, optional correlation_method : str, optional
synchon commented 2022-11-23 10:57:06 +00:00 (Migrated from github.com)

(default None).

`(default None).`
Any method that can be passed to Any method that can be passed to
:any:`pandas.DataFrame.corr` (default "pearson"). :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 name : str, optional
The name of the marker. If None, will use the class name The name of the marker. If None, will use the class name
(default None). (default None).
@ -43,6 +47,7 @@ class CrossParcellationFC(BaseMarker):
parcellation_two: str, parcellation_two: str,
aggregation_method: str = "mean", aggregation_method: str = "mean",
correlation_method: str = "pearson", correlation_method: str = "pearson",
mask: Optional[str] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
if parcellation_one == parcellation_two: if parcellation_one == parcellation_two:
@ -53,6 +58,7 @@ class CrossParcellationFC(BaseMarker):
self.parcellation_two = parcellation_two self.parcellation_two = parcellation_two
self.aggregation_method = aggregation_method self.aggregation_method = aggregation_method
self.correlation_method = correlation_method self.correlation_method = correlation_method
self.mask = mask
super().__init__(on=["BOLD"], name=name) super().__init__(on=["BOLD"], name=name)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
@ -145,10 +151,12 @@ class CrossParcellationFC(BaseMarker):
parcellation_one_dict = ParcelAggregation( parcellation_one_dict = ParcelAggregation(
parcellation=self.parcellation_one, parcellation=self.parcellation_one,
method=self.aggregation_method, method=self.aggregation_method,
mask=self.mask,
).compute(input) ).compute(input)
parcellation_two_dict = ParcelAggregation( parcellation_two_dict = ParcelAggregation(
parcellation=self.parcellation_two, parcellation=self.parcellation_two,
method=self.aggregation_method, method=self.aggregation_method,
mask=self.mask,
).compute(input) ).compute(input)
parcellated_ts_one = parcellation_one_dict["data"] parcellated_ts_one = parcellation_one_dict["data"]

View file

@ -29,9 +29,16 @@ class RSSETSMarker(BaseMarker):
parcellation : str parcellation : str
synchon commented 2022-11-23 10:57:29 +00:00 (Migrated from github.com)

(default None).

`(default None).`
The name of the parcellation. Check valid options by calling The name of the parcellation. Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`. :func:`junifer.data.parcellations.list_parcellations`.
aggregation_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
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). :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 name : str, optional
The name of the marker. If None, will use the class name (default The name of the marker. If None, will use the class name (default
None). None).
@ -41,11 +48,15 @@ class RSSETSMarker(BaseMarker):
def __init__( def __init__(
self, self,
parcellation: str, parcellation: str,
aggregation_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None,
mask: Optional[str] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation 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) super().__init__(name=name)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
@ -136,7 +147,9 @@ class RSSETSMarker(BaseMarker):
# Initialize a ParcelAggregation # Initialize a ParcelAggregation
parcel_aggregation = ParcelAggregation( parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.aggregation_method, method=self.agg_method,
method_params=self.agg_method_params,
mask=self.mask
) )
# Compute the parcel aggregation # Compute the parcel aggregation
out = parcel_aggregation.compute(input=input, extra_input=extra_input) out = parcel_aggregation.compute(input=input, extra_input=extra_input)

View file

@ -40,6 +40,10 @@ class FunctionalConnectivityParcels(BaseMarker):
cor_method_params : dict, optional cor_method_params : dict, optional
synchon commented 2022-11-23 10:57:40 +00:00 (Migrated from github.com)

(default None).

`(default None).`
Parameters to pass to the correlation function. Check valid options in Parameters to pass to the correlation function. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default None). :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 name : str, optional
The name of the marker. If None, will use the class name (default The name of the marker. If None, will use the class name (default
None). None).
@ -52,21 +56,20 @@ class FunctionalConnectivityParcels(BaseMarker):
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance", cor_method: str = "covariance",
cor_method_params: Optional[Dict] = None, cor_method_params: Optional[Dict] = None,
mask: Optional[str] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation self.parcellation = parcellation
self.agg_method = agg_method self.agg_method = agg_method
self.agg_method_params = ( self.agg_method_params = agg_method_params
{} if agg_method_params is None else agg_method_params
)
self.cor_method = cor_method self.cor_method = cor_method
self.cor_method_params = ( self.cor_method_params = cor_method_params or {}
{} if cor_method_params is None else cor_method_params
)
# default to nilearn behavior # default to nilearn behavior
self.cor_method_params["empirical"] = self.cor_method_params.get( self.cor_method_params["empirical"] = self.cor_method_params.get(
"empirical", False "empirical", False
) )
self.mask = mask
super().__init__(name=name) super().__init__(name=name)
@ -131,6 +134,7 @@ class FunctionalConnectivityParcels(BaseMarker):
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
mask=self.mask,
on="BOLD", on="BOLD",
) )
# get the 2D timeseries after parcel aggregation # get the 2D timeseries after parcel aggregation

View file

@ -43,6 +43,10 @@ class FunctionalConnectivitySpheres(BaseMarker):
cor_method_params : dict, optional cor_method_params : dict, optional
synchon commented 2022-11-23 10:58:03 +00:00 (Migrated from github.com)

(default None).

`(default None).`
Parameters to pass to the correlation function. Check valid options in Parameters to pass to the correlation function. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default None). :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 name : str, optional
The name of the marker. By default, it will use The name of the marker. By default, it will use
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
@ -58,6 +62,7 @@ class FunctionalConnectivitySpheres(BaseMarker):
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance", cor_method: str = "covariance",
cor_method_params: Optional[Dict] = None, cor_method_params: Optional[Dict] = None,
mask: Optional[str] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.coords = coords self.coords = coords
@ -65,18 +70,17 @@ class FunctionalConnectivitySpheres(BaseMarker):
if radius is None or radius <= 0: if radius is None or radius <= 0:
raise_error(f"radius should be > 0: provided {radius}") raise_error(f"radius should be > 0: provided {radius}")
self.agg_method = agg_method self.agg_method = agg_method
self.agg_method_params = ( self.agg_method_params = agg_method_params
{} if agg_method_params is None else agg_method_params
)
self.cor_method = cor_method self.cor_method = cor_method
self.cor_method_params = ( self.cor_method_params = cor_method_params or {}
{} if cor_method_params is None else cor_method_params
)
# default to nilearn behavior # default to nilearn behavior
self.cor_method_params["empirical"] = self.cor_method_params.get( self.cor_method_params["empirical"] = self.cor_method_params.get(
"empirical", False "empirical", False
) )
self.mask = mask
super().__init__(name=name) super().__init__(name=name)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
@ -142,6 +146,7 @@ class FunctionalConnectivitySpheres(BaseMarker):
radius=self.radius, radius=self.radius,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
mask=self.mask,
on="BOLD", on="BOLD",
) )

View file

@ -11,7 +11,7 @@ from nilearn.image import math_img, resample_to_img
from nilearn.maskers import NiftiMasker from nilearn.maskers import NiftiMasker
from ..api.decorators import register_marker 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 ..stats import get_aggfunc_by_name
from ..utils import logger from ..utils import logger
from .base import BaseMarker from .base import BaseMarker
@ -35,6 +35,10 @@ class ParcelAggregation(BaseMarker):
method_params : dict, optional method_params : dict, optional
synchon commented 2022-11-23 10:58:12 +00:00 (Migrated from github.com)

(default None).

`(default None).`
Parameters to pass to the aggregation function. Check valid options in Parameters to pass to the aggregation function. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name`. :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"} \ on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} \
or list of the options, optional or list of the options, optional
The data types to apply the marker to. If None, will work on all The data types to apply the marker to. If None, will work on all
@ -49,12 +53,14 @@ class ParcelAggregation(BaseMarker):
parcellation: str, parcellation: str,
method: str, method: str,
method_params: Optional[Dict[str, Any]] = None, method_params: Optional[Dict[str, Any]] = 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:
self.parcellation = parcellation self.parcellation = parcellation
self.method = method 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) super().__init__(on=on, name=name)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
@ -158,15 +164,34 @@ class ParcelAggregation(BaseMarker):
name=self.parcellation, name=self.parcellation,
resolution=resolution, resolution=resolution,
) )
parcellation_img_res = resample_to_img( parcellation_img_res = resample_to_img(
t_parcellation, t_parcellation,
t_input, t_input,
interpolation="nearest", interpolation="nearest",
copy=True,
) )
parcellation_bin = math_img( parcellation_bin = math_img(
"img != 0", "img != 0",
img=parcellation_img_res, 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") logger.debug("Masking")
masker = NiftiMasker( masker = NiftiMasker(
parcellation_bin, target_affine=t_input.affine parcellation_bin, target_affine=t_input.affine

View file

@ -7,7 +7,7 @@
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..data import load_coordinates from ..data import load_coordinates, load_mask
from ..external.nilearn import JuniferNiftiSpheresMasker from ..external.nilearn import JuniferNiftiSpheresMasker
from ..stats import get_aggfunc_by_name from ..stats import get_aggfunc_by_name
from ..utils import logger from ..utils import logger
@ -37,6 +37,10 @@ class SphereAggregation(BaseMarker):
(default "mean"). (default "mean").
synchon commented 2022-11-23 10:58:18 +00:00 (Migrated from github.com)

(default None).

`(default None).`
method_params : dict, optional method_params : dict, optional
The parameters to pass to the aggregation method (default None). 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 \ on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} or \
list of the options, optional list of the options, optional
The data types to apply the marker to. If None, will work on all 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, radius: Optional[float] = None,
method: str = "mean", method: str = "mean",
method_params: Optional[Dict[str, Any]] = None, method_params: Optional[Dict[str, Any]] = 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:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.method = method 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) super().__init__(on=on, name=name)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
@ -157,12 +163,17 @@ class SphereAggregation(BaseMarker):
agg_func = get_aggfunc_by_name( agg_func = get_aggfunc_by_name(
self.method, func_params=self.method_params 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 # Get seeds and labels
coords, out_labels = load_coordinates(name=self.coords) coords, out_labels = load_coordinates(name=self.coords)
masker = JuniferNiftiSpheresMasker( masker = JuniferNiftiSpheresMasker(
seeds=coords, seeds=coords,
radius=self.radius, radius=self.radius,
mask_img=None, # TODO: support this (needs #79) mask_img=mask_img,
agg_func=agg_func, agg_func=agg_func,
) )
# Fit and transform the marker on the data # Fit and transform the marker on the data

View file

@ -45,7 +45,8 @@ def test_compute() -> None:
# Assert the meta # Assert the meta
meta = ets_rss_marker.get_meta("BOLD")["marker"] meta = ets_rss_marker.get_meta("BOLD")["marker"]
assert meta["parcellation"] == "Schaefer100x17" 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" assert meta["class"] == "RSSETSMarker"

View file

@ -12,6 +12,7 @@ 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.markers.parcel_aggregation import ParcelAggregation from junifer.markers.parcel_aggregation import ParcelAggregation
@ -86,6 +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["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"
@ -110,6 +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["name"] == "VBM_GM_ParcelAggregation" assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM" assert meta["kind"] == "VBM_GM"
@ -142,6 +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["name"] == "VBM_GM_ParcelAggregation" assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM" assert meta["kind"] == "VBM_GM"
@ -175,7 +179,53 @@ 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["name"] == "BOLD_ParcelAggregation" assert meta["name"] == "BOLD_ParcelAggregation"
assert meta["class"] == "ParcelAggregation" assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "BOLD" assert meta["kind"] == "BOLD"
assert meta["method_params"] == {} 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"] == {}

View file

@ -12,7 +12,7 @@ from nilearn.image import concat_imgs
from nilearn.maskers import NiftiSpheresMasker from nilearn.maskers import NiftiSpheresMasker
from numpy.testing import assert_array_equal 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.markers.sphere_aggregation import SphereAggregation
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
@ -44,7 +44,7 @@ def test_SphereAggregation_3D() -> None:
vbm = oasis_dataset.gray_matter_maps[0] vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm) img = nib.load(vbm)
# Create NiftiLabelsMasker # Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS)
auto4d = nifti_masker.fit_transform(img) auto4d = nifti_masker.fit_transform(img)
@ -63,6 +63,7 @@ def test_SphereAggregation_3D() -> None:
assert meta["method"] == "mean" assert meta["method"] == "mean"
assert meta["coords"] == COORDS assert meta["coords"] == COORDS
assert meta["radius"] == RADIUS assert meta["radius"] == RADIUS
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_SphereAggregation" assert meta["name"] == "VBM_GM_SphereAggregation"
assert meta["class"] == "SphereAggregation" assert meta["class"] == "SphereAggregation"
assert meta["kind"] == "VBM_GM" assert meta["kind"] == "VBM_GM"
@ -72,13 +73,13 @@ def test_SphereAggregation_3D() -> None:
def test_SphereAggregation_4D() -> None: def test_SphereAggregation_4D() -> None:
"""Test SphereAggregation object on 4D images.""" """Test SphereAggregation object on 4D images."""
# Get the testing coordinates (for nilearn) # Get the testing coordinates (for nilearn)
coordinates, labels = load_coordinates(COORDS) coordinates, _ = load_coordinates(COORDS)
# Get the SPM auditory data # Get the SPM auditory data
subject_data = datasets.fetch_spm_auditory() subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore fmri_img = concat_imgs(subject_data.func) # type: ignore
# Create NiftiLabelsMasker # Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS)
auto4d = nifti_masker.fit_transform(fmri_img) auto4d = nifti_masker.fit_transform(fmri_img)
@ -97,6 +98,7 @@ def test_SphereAggregation_4D() -> None:
assert meta["method"] == "mean" assert meta["method"] == "mean"
assert meta["coords"] == COORDS assert meta["coords"] == COORDS
assert meta["radius"] == RADIUS assert meta["radius"] == RADIUS
assert meta["mask"] is None
assert meta["name"] == "BOLD_SphereAggregation" assert meta["name"] == "BOLD_SphereAggregation"
assert meta["class"] == "SphereAggregation" assert meta["class"] == "SphereAggregation"
assert meta["kind"] == "BOLD" assert meta["kind"] == "BOLD"
@ -145,3 +147,43 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
) )
marker.fit_transform(input, storage=storage) 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"] == {}