[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
:members:
Masks
=====
.. automodule:: junifer.data.masks
:members:

View file

@ -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

View file

@ -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
~~~~

View file

@ -14,3 +14,11 @@ from .parcellations import (
load_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
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 = (

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
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:`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"]

View file

@ -29,9 +29,16 @@ class RSSETSMarker(BaseMarker):
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
: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)

View file

@ -40,6 +40,10 @@ class FunctionalConnectivityParcels(BaseMarker):
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
: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

View file

@ -43,6 +43,10 @@ class FunctionalConnectivitySpheres(BaseMarker):
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
: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",
)

View file

@ -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
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
: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

View file

@ -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").
synchon commented 2022-11-23 10:58:18 +00:00 (Migrated from github.com)

(default None).

`(default None).`
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

View file

@ -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"

View file

@ -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"] == {}

View file

@ -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"] == {}