[ENH]: Add support for masks (GM/WM/etc) #79
21 changed files with 667 additions and 64 deletions
|
|
@ -9,3 +9,10 @@ Coordinates
|
|||
|
||||
.. automodule:: junifer.data.coordinates
|
||||
:members:
|
||||
|
||||
|
||||
Masks
|
||||
=====
|
||||
|
||||
.. automodule:: junifer.data.masks
|
||||
:members:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
~~~~
|
||||
|
||||
|
|
|
|||
|
|
@ -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
186
junifer/data/masks.py
Normal file
|
|
@ -0,0 +1,186 @@
|
|||
"""Provide functions for masks."""
|
||||
|
not entirely. Inner values could be anything not entirely. Inner values could be anything
But shouldn't be the inner dict keys be string? But shouldn't be the inner dict keys be string?
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
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
|
@ -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 = (
|
||||
|
|
|
|||
43
junifer/data/tests/test_data_utils.py
Normal file
43
junifer/data/tests/test_data_utils.py
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
"""Provide tests for data utils."""
|
||||
|
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(
|
||||
|
`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
|
||||
)
|
||||
153
junifer/data/tests/test_masks.py
Normal file
153
junifer/data/tests/test_masks.py
Normal 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
42
junifer/data/utils.py
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
"""Provide utilities for data module."""
|
||||
|
float or None? float or None?
list of float or int, or np.ndarray? list of float or int, or np.ndarray?
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
|
||||
|
|
@ -32,6 +32,10 @@ class CrossParcellationFC(BaseMarker):
|
|||
correlation_method : str, optional
|
||||
|
`(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"]
|
||||
|
|
|
|||
|
|
@ -29,9 +29,16 @@ class RSSETSMarker(BaseMarker):
|
|||
parcellation : str
|
||||
|
`(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)
|
||||
|
|
|
|||
|
|
@ -40,6 +40,10 @@ class FunctionalConnectivityParcels(BaseMarker):
|
|||
cor_method_params : dict, optional
|
||||
|
`(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
|
||||
|
|
|
|||
|
|
@ -43,6 +43,10 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
|||
cor_method_params : dict, optional
|
||||
|
`(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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
`(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
|
||||
|
|
|
|||
|
|
@ -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").
|
||||
|
`(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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
|
|||
Loading…
Reference in a new issue
Dict[str, Dict[str, str]]?