[ENH]: Support for handling multiple masks #174

Merged
fraimondo merged 13 commits from enh/multiple_masks into main 2023-02-27 10:37:36 +00:00
26 changed files with 585 additions and 191 deletions

View file

@ -16,3 +16,17 @@ the results. Finally, we will show how to use the ``queue`` command to interact
codeless codeless
running running
queueing queueing
.. _using_components:
Using junifer common components
synchon commented 2023-01-30 10:02:46 +00:00 (Migrated from github.com)

I think this should either have its own index or be a sub-section as it comes up as a section now. Here: https://juaml.github.io/junifer/pr-preview/pr-174/using/index.html, you have the section entry and also have a separate outermost entry in ToC.

I think this should either have its own index or be a sub-section as it comes up as a section now. Here: https://juaml.github.io/junifer/pr-preview/pr-174/using/index.html, you have the section entry and also have a separate outermost entry in ToC.
fraimondo commented 2023-02-23 11:35:55 +00:00 (Migrated from github.com)

should be a subheading

should be a subheading
-------------------------------
The following sections explains common components of junifer that can be used across many steps of the pipeline.
.. toctree::
:maxdepth: 2
:caption: Contents:
masks

68
docs/using/masks.rst Normal file
View file

@ -0,0 +1,68 @@
.. include:: ../links.inc
LeSasse commented 2023-01-27 08:21:46 +00:00 (Migrated from github.com)

typo: constraint -> constrain

typo: constraint -> constrain
LeSasse commented 2023-01-27 08:42:06 +00:00 (Migrated from github.com)

Overall really cool, I think it will also be good to state here that in BOLD preprocessing of large 4D NIfTI files, the masks can have a beneficial effect on memory usage (just because when i started doing BOLD processing with nilearn without a mask, I would get the occasional memory error).

Overall really cool, I think it will also be good to state here that in BOLD preprocessing of large 4D NIfTI files, the masks can have a beneficial effect on memory usage (just because when i started doing BOLD processing with nilearn without a mask, I would get the occasional memory error).
fraimondo commented 2023-01-27 08:49:27 +00:00 (Migrated from github.com)

I am no expert in fMRI analysis, that's why I don't want to add much. I think this section could be improved.

Also, if we "mask" at the preprocessing, then all the markers must set the mask to "inherit". Otherwise, non-clean voxels might be used.

I am no expert in fMRI analysis, that's why I don't want to add much. I think this section could be improved. Also, if we "mask" at the preprocessing, then all the markers _must_ set the mask to "inherit". Otherwise, non-clean voxels might be used.
synchon commented 2023-01-30 10:04:07 +00:00 (Migrated from github.com)

... a mask ...

... a mask ...
synchon commented 2023-01-30 10:05:06 +00:00 (Migrated from github.com)

... ratio of gray matter to white matter / cerebrospinal fluid ...

... ratio of gray matter to white matter / cerebrospinal fluid ...
synchon commented 2023-01-30 10:06:30 +00:00 (Migrated from github.com)

... ``masks`` ...

```... ``masks`` ...```
synchon commented 2023-01-30 10:07:27 +00:00 (Migrated from github.com)

... **only** ...

```... **only** ...```
synchon commented 2023-01-30 10:08:45 +00:00 (Migrated from github.com)

... specifies the ``GM_prob0.2``...

```... specifies the ``GM_prob0.2``...```
synchon commented 2023-01-30 10:10:02 +00:00 (Migrated from github.com)

... ``compute_brain_mask`` ...

```... ``compute_brain_mask`` ...```
synchon commented 2023-01-30 10:11:12 +00:00 (Migrated from github.com)

... allows you to ...

... allows you to ...
synchon commented 2023-01-30 10:12:00 +00:00 (Migrated from github.com)

... ``GM_prob0.2`` and ``compute_brain_mask`` ...

```... ``GM_prob0.2`` and ``compute_brain_mask`` ...```
.. _using_masks:
Masks
=====
Masks are essentially boolean arrays that are used to constrain the extraction of features to voxels that are
meaningful. For example, in an fMRI imaging study, a mask can be used to constrain the extraction of features to
voxels that contain a certain ratio of gray matter to white matter / cerebrospinal fluid, ensuring that the features
are not extracted from voxels that contain mostly white matter or cerebrospinal fluid, which could add noise to the
BOLD signal.
Junifer provides a number of built-in masks, which can be listed using the :func:`junifer.data.masks.list_masks`. Some
masks are images, while other masks can be computed using :ref:`nilearn` functions.
For markers and steps that accept ``masks`` as an argument, the mask can be specified as a string, which will be the
name of a built-in mask, or as a dictionary in which the **only** key is the built-in mask name and the value is a
dictionary of keyword arguments to pass to the mask function.
For example, the following is a valid mask specification that specified the ``GM_prob0.2`` mask.
.. code-block:: yaml
masks: GM_prob0.2
The following is a valid mask specification that specifies the ``compute_brain_mask`` mask (function from nilearn),
with a threshold of 0.5.
.. code-block:: yaml
masks:
compute_brain_mask:
threshold: 0.5
Furthermore, junifer allows you to combine several masks using :func:`nilearn.masking.intersect_masks`. This is done by
specifying a list of masks, where each mask is a string or dictionary as described above. For example, the following
is a valid mask specification that specifies the intersection of the ``GM_prob0.2`` and ``compute_brain_mask`` masks.
.. code-block:: yaml
masks:
- GM_prob0.2
- compute_brain_mask:
threshold: 0.5
We can also specify the arguments of :func:`nilearn.masking.intersect_masks` (``threshold`` and ``connected``). The
following example combines the same masks as the previous one, but computing the full intersection.
.. code-block:: yaml
masks:
- GM_prob0.2
- compute_brain_mask:
threshold: 0.5
- threshold: 1 # intersection
Alternatively, we can also compute the union, even if the voxels do not form a connected component:
.. code-block:: yaml
masks:
- GM_prob0.2
- compute_brain_mask:
threshold: 0.5
- threshold: 0 # union
- connected: False # keep disconnected components

View file

@ -23,6 +23,7 @@ from nilearn.masking import (
compute_background_mask, compute_background_mask,
compute_brain_mask, compute_brain_mask,
compute_epi_mask, compute_epi_mask,
intersect_masks,
) )
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
@ -36,7 +37,10 @@ if TYPE_CHECKING:
_masks_path = Path(__file__).parent / "masks" _masks_path = Path(__file__).parent / "masks"
synchon commented 2023-01-30 10:27:31 +00:00 (Migrated from github.com)
  • ... Default is None ... => ... (default None).
  • Why have the extra_dict parameter if it's not used?
- ... Default is None ... => ... (default None). - Why have the `extra_dict` parameter if it's not used?
fraimondo commented 2023-02-23 11:40:16 +00:00 (Migrated from github.com)

it's used.

it's used.
fraimondo commented 2023-02-23 11:42:50 +00:00 (Migrated from github.com)

sorry, not used there, it was left by mistake.

sorry, not used there, it was left by mistake.
def _fetch_icbm152_brain_gm_mask(target_img: "Nifti1Image", **kwargs): def _fetch_icbm152_brain_gm_mask(
target_img: "Nifti1Image",
**kwargs,
):
"""Fetch ICBM152 brain mask and resample. """Fetch ICBM152 brain mask and resample.
Parameters Parameters
@ -65,6 +69,9 @@ data.
The built-in masks are files that are shipped with the package in the 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. data/masks directory. The user can also register their own masks.
Callable masks should be functions that take at least one parameter:
* `target_img`: the image to which the mask will be applied.
""" """
_available_masks: Dict[str, Dict[str, Any]] = { _available_masks: Dict[str, Dict[str, Any]] = {
"GM_prob0.2": {"family": "Vickery-Patil"}, "GM_prob0.2": {"family": "Vickery-Patil"},
@ -146,19 +153,23 @@ def list_masks() -> List[str]:
def get_mask( def get_mask(
mask: Union[str, Dict], masks: Union[str, Dict, List[Union[Dict, str]]],
target_data: Dict[str, Any], target_data: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None,
) -> "Nifti1Image": ) -> "Nifti1Image":
synchon commented 2023-01-30 10:31:37 +00:00 (Migrated from github.com)

... Default is None ... => ... (default None).

... Default is None ... => ... (default None).
LeSasse commented 2023-02-02 09:28:58 +00:00 (Migrated from github.com)

I think this should actually be called extra_input to be consistent with the extra_input parameter in other markers: https://github.com/juaml/junifer/blob/main/junifer/markers/parcel_aggregation.py#L103

I think this should actually be called `extra_input` to be consistent with the `extra_input` parameter in other markers: https://github.com/juaml/junifer/blob/main/junifer/markers/parcel_aggregation.py#L103
fraimondo commented 2023-02-23 11:42:57 +00:00 (Migrated from github.com)

good point

good point
"""Get mask, tailored for the target image. """Get mask, tailored for the target image.
Parameters Parameters
---------- ----------
masks : str or dict masks : str, dict or list of dict or str
The name of the mask, or the name of a callable mask and the parameters The name of the mask, or the name of a callable mask and the parameters
of the mask. of the mask as a dictionary. Several masks can be passed as a list.
target_data : dict target_data : dict
The corresponding item of the data object to which the mask will be The corresponding item of the data object to which the mask will be
applied. applied.
extra_input : dict, optional
The other fields in the data object. Useful for accessing other data
kinds that needs to be used in the computation of masks (default None).
Returns Returns
------- -------
@ -167,20 +178,68 @@ def get_mask(
""" """
# Get the min of the voxels sizes and use it as the resolution # Get the min of the voxels sizes and use it as the resolution
target_img = target_data["data"] target_img = target_data["data"]
inherited_mask_item = target_data.get("mask_item", None)
resolution = np.min(target_img.header.get_zooms()[:3]) resolution = np.min(target_img.header.get_zooms()[:3])
if isinstance(mask, dict): if not isinstance(masks, list):
if len(mask) != 1: masks = [masks]
# Check that dicts have only one key
invalid_elements = [
x for x in masks if isinstance(x, dict) and len(x) != 1
]
if len(invalid_elements) > 0:
raise_error( raise_error(
"The mask dictionary must have only one key, " "Each of the masks dictionary must have only one key, "
"the name of the mask." "the name of the mask. The following dictionaries are invalid: "
f"{invalid_elements}"
) )
mask_name = list(mask.keys())[0]
mask_params = mask[mask_name] # Check params for the intersection function
intersect_params = {}
true_masks = []
for t_mask in masks:
if isinstance(t_mask, dict):
if "threshold" in t_mask:
intersect_params["threshold"] = t_mask["threshold"]
continue
elif "connected" in t_mask:
intersect_params["connected"] = t_mask["connected"]
continue
# All the other elements are masks
true_masks.append(t_mask)
if len(true_masks) == 0:
raise_error("No mask was passed. At least one mask is required.")
# Get all the masks
all_masks = []
for t_mask in true_masks:
if isinstance(t_mask, dict):
mask_name = list(t_mask.keys())[0]
mask_params = t_mask[mask_name]
else: else:
mask_name = mask mask_name = t_mask
mask_params = None mask_params = None
if mask_name == "inherit":
if extra_input is None:
raise_error(
"Cannot inherit mask from another data item "
"because no extra data was passed."
)
if inherited_mask_item is None:
raise_error(
"Cannot inherit mask from another data item "
"because no mask item was specified "
"(missing `mask_item` key in the data object)."
)
if inherited_mask_item not in extra_input:
raise_error(
"Cannot inherit mask from another data item "
f"because the item ({inherited_mask_item}) does not exist."
)
mask_img = extra_input[inherited_mask_item]["data"]
else:
mask_object, _ = load_mask( mask_object, _ = load_mask(
mask_name, path_only=False, resolution=resolution mask_name, path_only=False, resolution=resolution
) )
@ -190,13 +249,26 @@ def get_mask(
mask_img = mask_object(target_img, **mask_params) mask_img = mask_object(target_img, **mask_params)
else: # Mask is a Nifti1Image else: # Mask is a Nifti1Image
if mask_params is not None: if mask_params is not None:
raise_error("Cannot pass callable params to a non-callable mask.") raise_error(
"Cannot pass callable params to a non-callable mask."
)
mask_img = resample_to_img( mask_img = resample_to_img(
mask_object, mask_object,
target_img, target_img,
interpolation="nearest", interpolation="nearest",
copy=True, copy=True,
) )
all_masks.append(mask_img)
if len(all_masks) > 1:
mask_img = intersect_masks(all_masks, **intersect_params)
else:
if len(intersect_params) > 0:
# Yes, I'm this strict!
raise_error(
"Cannot pass parameters to the intersection function "
"when there is only one mask."
)
mask_img = all_masks[0]
return mask_img return mask_img

View file

@ -6,8 +6,9 @@
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
from typing import Callable, Dict, Union from typing import Callable, Dict, List, Union
import numpy as np
import pytest import pytest
from nilearn.datasets import fetch_icbm152_brain_gm_mask from nilearn.datasets import fetch_icbm152_brain_gm_mask
from nilearn.image import resample_to_img from nilearn.image import resample_to_img
@ -15,6 +16,7 @@ from nilearn.masking import (
compute_background_mask, compute_background_mask,
compute_brain_mask, compute_brain_mask,
compute_epi_mask, compute_epi_mask,
intersect_masks,
) )
from numpy.testing import assert_array_almost_equal, assert_array_equal from numpy.testing import assert_array_almost_equal, assert_array_equal
@ -56,7 +58,9 @@ def test_register_mask_already_registered() -> None:
name="testmask", name="testmask",
mask_path="testmask.nii.gz", mask_path="testmask.nii.gz",
) )
assert load_mask("testmask", path_only=True)[1].name == "testmask.nii.gz" out = load_mask("testmask", path_only=True)
assert out[1] is not None
assert out[1].name == "testmask.nii.gz"
# Try registering again # Try registering again
with pytest.raises(ValueError, match=r"already registered."): with pytest.raises(ValueError, match=r"already registered."):
@ -70,7 +74,9 @@ def test_register_mask_already_registered() -> None:
overwrite=True, overwrite=True,
) )
assert load_mask("testmask", path_only=True)[1].name == "testmask2.nii.gz" out = load_mask("testmask", path_only=True)
assert out[1] is not None
assert out[1].name == "testmask2.nii.gz"
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -110,6 +116,7 @@ def test_register_mask(
# Load registered mask # Load registered mask
_, fname = load_mask(name=name, path_only=True) _, fname = load_mask(name=name, path_only=True)
# Check values for registered mask # Check values for registered mask
assert fname is not None
assert fname.name == f"{name}.nii.gz" assert fname.name == f"{name}.nii.gz"
@ -146,6 +153,7 @@ def test_vickery_patil() -> None:
mask.header["pixdim"][1:4], [1.5, 1.5, 1.5] # type: ignore mask.header["pixdim"][1:4], [1.5, 1.5, 1.5] # type: ignore
) )
assert fname is not None
assert fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz" assert fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz"
mask, fname = load_mask("GM_prob0.2", resolution=3) mask, fname = load_mask("GM_prob0.2", resolution=3)
@ -153,6 +161,7 @@ def test_vickery_patil() -> None:
mask.header["pixdim"][1:4], [3.0, 3.0, 3.0] # type: ignore mask.header["pixdim"][1:4], [3.0, 3.0, 3.0] # type: ignore
) )
assert fname is not None
assert ( assert (
fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz" fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz"
) )
@ -162,6 +171,7 @@ def test_vickery_patil() -> None:
mask.header["pixdim"][1:4], [3.0, 3.0, 3.0] # type: ignore mask.header["pixdim"][1:4], [3.0, 3.0, 3.0] # type: ignore
) )
assert fname is not None
assert fname.name == "GMprob0.2_cortex_3mm_NA_rm.nii.gz" assert fname.name == "GMprob0.2_cortex_3mm_NA_rm.nii.gz"
with pytest.raises(ValueError, match=r"find a Vickery-Patil mask "): with pytest.raises(ValueError, match=r"find a Vickery-Patil mask "):
@ -176,7 +186,7 @@ def test_get_mask() -> None:
input = reader.fit_transform(input) input = reader.fit_transform(input)
vbm_gm = input["VBM_GM"] vbm_gm = input["VBM_GM"]
vbm_gm_img = vbm_gm["data"] vbm_gm_img = vbm_gm["data"]
mask = get_mask(mask="GM_prob0.2", target_data=vbm_gm) mask = get_mask(masks="GM_prob0.2", target_data=vbm_gm)
assert mask.shape == vbm_gm_img.shape assert mask.shape == vbm_gm_img.shape
assert_array_equal(mask.affine, vbm_gm_img.affine) assert_array_equal(mask.affine, vbm_gm_img.affine)
@ -204,7 +214,7 @@ def test_mask_callable() -> None:
input = reader.fit_transform(input) input = reader.fit_transform(input)
vbm_gm = input["VBM_GM"] vbm_gm = input["VBM_GM"]
vbm_gm_img = vbm_gm["data"] vbm_gm_img = vbm_gm["data"]
mask = get_mask(mask="identity", target_data=vbm_gm) mask = get_mask(masks="identity", target_data=vbm_gm)
assert_array_equal(mask.get_fdata(), vbm_gm_img.get_fdata()) assert_array_equal(mask.get_fdata(), vbm_gm_img.get_fdata())
@ -218,11 +228,48 @@ def test_get_mask_errors() -> None:
input = dg["sub-01"] input = dg["sub-01"]
input = reader.fit_transform(input) input = reader.fit_transform(input)
vbm_gm = input["VBM_GM"] vbm_gm = input["VBM_GM"]
# Test wrong masks definitions (more than one key per dict)
with pytest.raises(ValueError, match=r"only one key"): with pytest.raises(ValueError, match=r"only one key"):
get_mask(mask={"GM_prob0.2": {}, "Other": {}}, target_data=vbm_gm) get_mask(masks={"GM_prob0.2": {}, "Other": {}}, target_data=vbm_gm)
# Test wrong masks definitions (pass paramaeters to non-callable mask)
with pytest.raises(ValueError, match=r"callable params"): with pytest.raises(ValueError, match=r"callable params"):
get_mask(mask={"GM_prob0.2": {"param": 1}}, target_data=vbm_gm) get_mask(masks={"GM_prob0.2": {"param": 1}}, target_data=vbm_gm)
# Pass only parametesr to the intersection function
with pytest.raises(
ValueError, match=r" At least one mask is required."
):
get_mask(masks={"threshold": 1}, target_data=vbm_gm)
# Pass parameters to the intersection function when only one mask
with pytest.raises(
ValueError, match=r"parameters to the intersection"
):
get_mask(
masks=["GM_prob0.2", {"threshold": 1}], target_data=vbm_gm
)
# Test "inherited" masks errors
# 1) No extra_data parameter
with pytest.raises(ValueError, match=r"no extra data was passed"):
get_mask(masks="inherit", target_data=vbm_gm)
extra_input = {"VBM_MASK": {}}
# 2) No mask_item key in target_data
with pytest.raises(ValueError, match=r"no mask item was specified"):
get_mask(
masks="inherit", target_data=vbm_gm, extra_input=extra_input
)
# 3) mask_item not in extra data
with pytest.raises(ValueError, match=r"does not exist"):
vbm_gm["mask_item"] = "wrong"
get_mask(
masks="inherit", target_data=vbm_gm, extra_input=extra_input
)
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -271,7 +318,7 @@ def test_nilearn_compute_masks(
else: else:
mask_spec = {mask_name: params} mask_spec = {mask_name: params}
mask = get_mask(mask=mask_spec, target_data=bold) mask = get_mask(masks=mask_spec, target_data=bold)
assert_array_equal(mask.affine, bold_img.affine) assert_array_equal(mask.affine, bold_img.affine)
@ -287,3 +334,114 @@ def test_nilearn_compute_masks(
copy=True, copy=True,
) )
assert_array_equal(mask.get_fdata(), ni_mask.get_fdata()) assert_array_equal(mask.get_fdata(), ni_mask.get_fdata())
def test_get_mask_inherit() -> None:
synchon commented 2023-02-24 10:00:24 +00:00 (Migrated from github.com)

Missing Parameters section in the docstring.

Missing Parameters section in the docstring.
"""Test using the inherit mask functionality."""
reader = DefaultDataReader()
with SPMAuditoryTestingDatagrabber() as dg:
input = dg["sub001"]
input = reader.fit_transform(input)
# Compute brain mask using nilearn
gm_mask = compute_brain_mask(input["BOLD"]["data"], threshold=0.2)
# Get mask using the compute_brain_mask function
mask1 = get_mask(
masks={"compute_brain_mask": {"threshold": 0.2}},
target_data=input["BOLD"],
)
# Now get the mask using the inherit functionality, passing the
# computed mask as extra data
extra_input = {"BOLD_MASK": {"data": gm_mask}}
input["BOLD"]["mask_item"] = "BOLD_MASK"
mask2 = get_mask(
masks="inherit", target_data=input["BOLD"], extra_input=extra_input
)
# Both masks should be equal
assert_array_equal(mask1.get_fdata(), mask2.get_fdata())
@pytest.mark.parametrize(
"masks,params",
[
(["GM_prob0.2", "compute_brain_mask"], {}),
(
["GM_prob0.2", "compute_brain_mask"],
{"threshold": 0.2},
),
(
[
"GM_prob0.2",
"compute_brain_mask",
"fetch_icbm152_brain_gm_mask",
],
{"threshold": 1, "connected": True},
),
],
)
def test_get_mask_multiple(
masks: Union[str, Dict, List[Union[Dict, str]]], params: Dict
) -> None:
"""Test getting multiple masks.
Parameters
----------
masks : str, dict, list of str or dict
Masks to get, junifer style.
params : dict
Parameters to pass to the intersect_masks function.
"""
reader = DefaultDataReader()
with SPMAuditoryTestingDatagrabber() as dg:
input = dg["sub001"]
input = reader.fit_transform(input)
if not isinstance(masks, list):
junifer_masks = [masks]
else:
junifer_masks = masks.copy()
if len(params) > 0:
# Convert params to junifer style (one dict per param)
junifer_params = [{k: params[k]} for k in params.keys()]
junifer_masks.extend(junifer_params)
target_img = input["BOLD"]["data"]
resolution = np.min(target_img.header.get_zooms()[:3])
computed = get_mask(masks=junifer_masks, target_data=input["BOLD"])
masks_names = [
list(x.keys())[0] if isinstance(x, dict) else x for x in masks
]
mask_funcs = [
x
for x in masks_names
if _available_masks[x]["family"] == "Callable"
]
mask_files = [
x
for x in masks_names
if _available_masks[x]["family"] != "Callable"
]
mask_imgs = [
load_mask(t_mask, path_only=False, resolution=resolution)[0]
for t_mask in mask_files
]
for t_func in mask_funcs:
mask_imgs.append(_available_masks[t_func]["func"](target_img))
mask_imgs = [
resample_to_img(
t_mask,
target_img,
interpolation="nearest",
copy=True,
)
for t_mask in mask_imgs
]
expected = intersect_masks(mask_imgs, **params)
assert_array_equal(computed.get_fdata(), expected.get_fdata())

View file

@ -32,10 +32,10 @@ class RSSETSMarker(BaseMarker):
agg_method_params : dict, optional agg_method_params : dict, optional
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` (default None). :func:`junifer.stats.get_aggfunc_by_name` (default None).
mask : str, optional masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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).
@ -49,13 +49,13 @@ class RSSETSMarker(BaseMarker):
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
mask: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = 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 = agg_method_params self.agg_method_params = agg_method_params
self.mask = mask self.masks = masks
super().__init__(name=name) super().__init__(name=name)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
@ -126,11 +126,11 @@ class RSSETSMarker(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, masks=self.masks,
) )
# 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)
edge_ts = _ets(out["data"]) edge_ts, _ = _ets(out["data"])
# Compute the RSS # Compute the RSS
out["data"] = np.sum(edge_ts**2, 1) ** 0.5 out["data"] = np.sum(edge_ts**2, 1) ** 0.5
# Set correct column label # Set correct column label

View file

@ -35,10 +35,10 @@ class AmplitudeLowFrequencyFluctuationParcels(
use_afni : bool, optional use_afni : bool, optional
Whether to use AFNI for computing. If None, will use AFNI only Whether to use AFNI for computing. If None, will use AFNI only
if available (default None). if available (default None).
mask : str, optional masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (default None).
method : str, optional 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").
@ -70,13 +70,13 @@ class AmplitudeLowFrequencyFluctuationParcels(
lowpass: float = 0.1, lowpass: float = 0.1,
tr: Optional[float] = None, tr: Optional[float] = None,
use_afni: Optional[bool] = None, use_afni: Optional[bool] = None,
mask: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
method: str = "mean", method: str = "mean",
method_params: Optional[Dict] = None, method_params: Optional[Dict] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation self.parcellation = parcellation
self.mask = mask self.masks = masks
self.method = method self.method = method
self.method_params = method_params self.method_params = method_params
super().__init__( super().__init__(
@ -116,7 +116,7 @@ class AmplitudeLowFrequencyFluctuationParcels(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.method, method=self.method,
method_params=self.method_params, method_params=self.method_params,
mask=self.mask, masks=self.masks,
on="fALFF", on="fALFF",
) )

View file

@ -5,7 +5,7 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Dict, Optional, Union from typing import Dict, List, Optional, Union
from ...api.decorators import register_marker from ...api.decorators import register_marker
from .. import SphereAggregation from .. import SphereAggregation
@ -39,10 +39,10 @@ class AmplitudeLowFrequencyFluctuationSpheres(
use_afni : bool, optional use_afni : bool, optional
Whether to use AFNI for computing. If None, will use AFNI only Whether to use AFNI for computing. If None, will use AFNI only
if available (default None). if available (default None).
mask : str, optional masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (default None).
method : str, optional 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").
@ -75,14 +75,14 @@ class AmplitudeLowFrequencyFluctuationSpheres(
lowpass: float = 0.1, lowpass: float = 0.1,
tr: Optional[float] = None, tr: Optional[float] = None,
use_afni: Optional[bool] = None, use_afni: Optional[bool] = None,
mask: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
method: str = "mean", method: str = "mean",
method_params: Optional[Dict] = None, method_params: Optional[Dict] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.mask = mask self.masks = masks
self.method = method self.method = method
self.method_params = method_params self.method_params = method_params
super().__init__( super().__init__(
@ -123,7 +123,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
radius=self.radius, radius=self.radius,
method=self.method, method=self.method,
method_params=self.method_params, method_params=self.method_params,
mask=self.mask, masks=self.masks,
on="fALFF", on="fALFF",
) )

View file

@ -30,10 +30,10 @@ class CrossParcellationFC(BaseMarker):
correlation_method : str, optional correlation_method : str, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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).
@ -47,7 +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: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
if parcellation_one == parcellation_two: if parcellation_one == parcellation_two:
@ -58,7 +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 self.masks = masks
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]:
@ -129,12 +129,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, masks=self.masks,
).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, masks=self.masks,
).compute(input) ).compute(input)
parcellated_ts_one = parcellation_one_dict["data"] parcellated_ts_one = parcellation_one_dict["data"]

View file

@ -35,10 +35,10 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
cor_method_params : dict, optional cor_method_params : dict, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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).
@ -59,7 +59,7 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
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, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation self.parcellation = parcellation
@ -68,7 +68,7 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
agg_method_params=agg_method_params, agg_method_params=agg_method_params,
cor_method=cor_method, cor_method=cor_method,
cor_method_params=cor_method_params, cor_method_params=cor_method_params,
mask=mask, masks=masks,
name=name, name=name,
) )
@ -78,7 +78,7 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
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, masks=self.masks,
on="BOLD", on="BOLD",
) )

View file

@ -4,7 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any, Dict, Optional from typing import Any, Dict, List, Optional, Union
from ...api.decorators import register_marker from ...api.decorators import register_marker
from ..sphere_aggregation import SphereAggregation from ..sphere_aggregation import SphereAggregation
@ -37,10 +37,10 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
cor_method_params : dict, optional cor_method_params : dict, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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_EdgeCentricFCSpheres where KIND is the kind of data it KIND_EdgeCentricFCSpheres where KIND is the kind of data it
@ -63,7 +63,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
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, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.coords = coords self.coords = coords
@ -75,7 +75,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
agg_method_params=agg_method_params, agg_method_params=agg_method_params,
cor_method=cor_method, cor_method=cor_method,
cor_method_params=cor_method_params, cor_method_params=cor_method_params,
mask=mask, masks=masks,
name=name, name=name,
) )
@ -86,7 +86,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
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, masks=self.masks,
on="BOLD", on="BOLD",
) )
bold_aggregated = sphere_aggregation.compute(input) bold_aggregated = sphere_aggregation.compute(input)

View file

@ -32,10 +32,10 @@ class FunctionalConnectivityBase(BaseMarker):
cor_method_params : dict, optional cor_method_params : dict, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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).
@ -50,7 +50,7 @@ class FunctionalConnectivityBase(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: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.agg_method = agg_method self.agg_method = agg_method
@ -62,7 +62,7 @@ class FunctionalConnectivityBase(BaseMarker):
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 self.masks = masks
super().__init__(on="BOLD", name=name) super().__init__(on="BOLD", name=name)
@abstractmethod @abstractmethod
@ -134,7 +134,7 @@ class FunctionalConnectivityBase(BaseMarker):
# Compute correlation # Compute correlation
if self.cor_method_params["empirical"]: if self.cor_method_params["empirical"]:
connectivity = ConnectivityMeasure( connectivity = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), cov_estimator=EmpiricalCovariance(), # type: ignore
kind=self.cor_method, kind=self.cor_method,
) )
else: else:

View file

@ -34,10 +34,10 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
cor_method_params : dict, optional cor_method_params : dict, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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).
@ -51,7 +51,7 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
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: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation self.parcellation = parcellation
@ -60,7 +60,7 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
agg_method_params=agg_method_params, agg_method_params=agg_method_params,
cor_method=cor_method, cor_method=cor_method,
cor_method_params=cor_method_params, cor_method_params=cor_method_params,
mask=mask, masks=masks,
name=name, name=name,
) )
@ -70,7 +70,7 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
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, masks=self.masks,
on="BOLD", on="BOLD",
) )
# Return the 2D timeseries after parcel aggregation # Return the 2D timeseries after parcel aggregation

View file

@ -5,7 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any, Dict, Optional, Union from typing import Any, Dict, List, Optional, Union
from ...api.decorators import register_marker from ...api.decorators import register_marker
from ..sphere_aggregation import SphereAggregation from ..sphere_aggregation import SphereAggregation
@ -38,10 +38,10 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
cor_method_params : dict, optional cor_method_params : dict, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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
@ -57,7 +57,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
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: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.coords = coords self.coords = coords
@ -69,7 +69,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
agg_method_params=agg_method_params, agg_method_params=agg_method_params,
cor_method=cor_method, cor_method=cor_method,
cor_method_params=cor_method_params, cor_method_params=cor_method_params,
mask=mask, masks=masks,
name=name, name=name,
) )
@ -80,7 +80,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
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, masks=self.masks,
on="BOLD", on="BOLD",
) )
# Return the 2D timeseries after sphere aggregation # Return the 2D timeseries after sphere aggregation

View file

@ -6,10 +6,12 @@
import pytest import pytest
# done to keep line length 79 # done to keep line length 79
import junifer.markers.functional_connectivity as fc from junifer.markers.functional_connectivity import (
functional_connectivity_base as fcb,
)
def test_base_functional_connectivity_marker_abstractness() -> None: def test_base_functional_connectivity_marker_abstractness() -> None:
"""Test FunctionalConnectivityBase is an abstract base class.""" """Test FunctionalConnectivityBase is an abstract base class."""
with pytest.raises(TypeError, match="abstract"): with pytest.raises(TypeError, match="abstract"):
fc.functional_connectivity_base.FunctionalConnectivityBase() fcb.FunctionalConnectivityBase() # type: ignore

View file

@ -32,10 +32,10 @@ class ParcelAggregation(BaseMarker):
method_params : dict, optional method_params : dict, optional
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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
@ -52,7 +52,7 @@ class ParcelAggregation(BaseMarker):
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
method: str, method: str,
method_params: Optional[Dict[str, Any]] = None, method_params: Optional[Dict[str, Any]] = None,
mask: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = 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:
@ -61,7 +61,7 @@ class ParcelAggregation(BaseMarker):
self.parcellation = parcellation self.parcellation = parcellation
self.method = method self.method = method
self.method_params = method_params or {} self.method_params = method_params or {}
self.mask = mask self.masks = masks
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]:
@ -184,9 +184,11 @@ class ParcelAggregation(BaseMarker):
img=parcellation_img_res, img=parcellation_img_res,
LeSasse commented 2023-02-02 12:44:19 +00:00 (Migrated from github.com)

this needs to hand over the extra_input to get_mask

this needs to hand over the `extra_input` to get_mask
fraimondo commented 2023-02-23 11:44:10 +00:00 (Migrated from github.com)

good catch!

good catch!
) )
if self.mask is not None: if self.masks is not None:
logger.debug(f"Masking with {self.mask}") logger.debug(f"Masking with {self.masks}")
mask_img = get_mask(mask=self.mask, target_data=input) mask_img = get_mask(
masks=self.masks, target_data=input, extra_input=extra_input
)
parcellation_bin = math_img( parcellation_bin = math_img(
"np.logical_and(img, mask)", "np.logical_and(img, mask)",

View file

@ -4,7 +4,7 @@
# License: AGPL # License: AGPL
from typing import Any, Dict, Optional, Union from typing import Any, Dict, List, Optional, Union
import numpy as np import numpy as np
@ -74,10 +74,10 @@ class ReHoParcels(ReHoBase):
agg_method_params : dict, optional agg_method_params : dict, optional
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` (default None). :func:`junifer.stats.get_aggfunc_by_name` (default None).
mask : str, optional masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (default None).
name : str, optional name : str, optional
The name of the marker. If None, it will use the class name The name of the marker. If None, it will use the class name
(default None). (default None).
@ -91,14 +91,14 @@ class ReHoParcels(ReHoBase):
reho_params: Optional[Dict] = None, reho_params: Optional[Dict] = None,
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
mask: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation self.parcellation = parcellation
self.reho_params = reho_params self.reho_params = reho_params
self.agg_method = agg_method self.agg_method = agg_method
self.agg_method_params = agg_method_params self.agg_method_params = agg_method_params
self.mask = mask self.masks = masks
super().__init__(use_afni=use_afni, name=name) super().__init__(use_afni=use_afni, name=name)
def compute( def compute(
@ -136,7 +136,7 @@ class ReHoParcels(ReHoBase):
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, masks=self.masks,
on="BOLD", on="BOLD",
) )
# Perform aggregation on reho map # Perform aggregation on reho map

View file

@ -4,7 +4,7 @@
# License: AGPL # License: AGPL
from typing import Any, Dict, Optional, Union from typing import Any, Dict, List, Optional, Union
import numpy as np import numpy as np
@ -79,10 +79,10 @@ class ReHoSpheres(ReHoBase):
(default None). (default None).
agg_method_params : dict, optional agg_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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (default None).
name : str, optional name : str, optional
The name of the marker. If None, it will use the class name The name of the marker. If None, it will use the class name
(default None). (default None).
@ -97,7 +97,7 @@ class ReHoSpheres(ReHoBase):
reho_params: Optional[Dict] = None, reho_params: Optional[Dict] = None,
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
mask: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.coords = coords self.coords = coords
@ -105,7 +105,7 @@ class ReHoSpheres(ReHoBase):
self.reho_params = reho_params self.reho_params = reho_params
self.agg_method = agg_method self.agg_method = agg_method
self.agg_method_params = agg_method_params self.agg_method_params = agg_method_params
self.mask = mask self.masks = masks
super().__init__(use_afni=use_afni, name=name) super().__init__(use_afni=use_afni, name=name)
def compute( def compute(
@ -144,7 +144,7 @@ class ReHoSpheres(ReHoBase):
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, masks=self.masks,
on="BOLD", on="BOLD",
) )
# Perform aggregation on reho map # Perform aggregation on reho map

View file

@ -34,10 +34,10 @@ class SphereAggregation(BaseMarker):
(default "mean"). (default "mean").
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 masks : str, dict or list of dict or str, optional
The name of the mask to apply to regions before extracting signals. The specification of the masks to apply to regions before extracting
Check valid options by calling :func:`junifer.data.masks.list_masks` signals. Check :ref:`Using Masks <using_masks>` for more details.
(default None). If None, will not apply any mask (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
@ -56,7 +56,7 @@ 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: Union[str, Dict, None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = 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:
@ -64,7 +64,7 @@ class SphereAggregation(BaseMarker):
self.radius = radius self.radius = radius
self.method = method self.method = method
self.method_params = method_params or {} self.method_params = method_params or {}
self.mask = mask self.masks = masks
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]:
@ -137,9 +137,11 @@ class SphereAggregation(BaseMarker):
) )
# Load mask # Load mask
mask_img = None mask_img = None
if self.mask is not None: if self.masks is not None:
logger.debug(f"Masking with {self.mask}") logger.debug(f"Masking with {self.masks}")
mask_img = get_mask(mask=self.mask, target_data=input) mask_img = get_mask(
masks=self.masks, target_data=input, extra_input=extra_input
)
# 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(

View file

@ -39,7 +39,7 @@ def test_marker_collection_incorrect_markers() -> None:
), ),
] ]
with pytest.raises(ValueError, match=r"must have different names"): with pytest.raises(ValueError, match=r"must have different names"):
MarkerCollection(wrong_markers) MarkerCollection(wrong_markers) # type: ignore
def test_marker_collection() -> None: def test_marker_collection() -> None:
@ -62,7 +62,7 @@ def test_marker_collection() -> None:
name="gmd_schaefer100x7_trim_mean90", name="gmd_schaefer100x7_trim_mean90",
), ),
] ]
mc = MarkerCollection(markers=markers) mc = MarkerCollection(markers=markers) # type: ignore
assert mc._markers == markers assert mc._markers == markers
assert mc._preprocessing is None assert mc._preprocessing is None
assert mc._storage is None assert mc._storage is None
@ -96,7 +96,7 @@ def test_marker_collection() -> None:
return input return input
mc2 = MarkerCollection( mc2 = MarkerCollection(
markers=markers, markers=markers, # type: ignore
preprocessing=BypassPreprocessing(), preprocessing=BypassPreprocessing(),
datareader=DefaultDataReader(), datareader=DefaultDataReader(),
) )
@ -127,7 +127,7 @@ def test_marker_collection_with_preprocessing() -> None:
), ),
] ]
mc = MarkerCollection( mc = MarkerCollection(
markers=markers, markers=markers, # type: ignore
preprocessing=fMRIPrepConfoundRemover(), preprocessing=fMRIPrepConfoundRemover(),
) )
assert mc._markers == markers assert mc._markers == markers
@ -173,7 +173,7 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
uri = tmp_path / "test_marker_collection_storage.sqlite" uri = tmp_path / "test_marker_collection_storage.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
mc = MarkerCollection( mc = MarkerCollection(
markers=markers, markers=markers, # type: ignore
storage=storage, storage=storage,
datareader=DefaultDataReader(), datareader=DefaultDataReader(),
) )
@ -185,7 +185,9 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
out = mc.fit(input) out = mc.fit(input)
assert out is None assert out is None
mc2 = MarkerCollection(markers=markers, datareader=DefaultDataReader()) mc2 = MarkerCollection(
markers=markers, datareader=DefaultDataReader() # type: ignore
)
mc2.validate(dg) mc2.validate(dg)
assert mc2._storage is None assert mc2._storage is None

View file

@ -27,13 +27,14 @@ def test_ets() -> None:
n_edges = int(n_rois * (n_rois - 1) / 2) n_edges = int(n_rois * (n_rois - 1) / 2)
# test without labels # test without labels
edge_ts = _ets(bold_ts) edge_ts, _ = _ets(bold_ts)
assert edge_ts.shape == (n_time, n_edges) assert edge_ts.shape == (n_time, n_edges)
# test with labels # test with labels
roi_labels = [f"Label_{x}" for x in range(n_rois)] roi_labels = [f"Label_{x}" for x in range(n_rois)]
edge_ts, edge_labels = _ets(bold_ts, roi_labels) edge_ts, edge_labels = _ets(bold_ts, roi_labels)
assert edge_ts.shape == (n_time, n_edges) assert edge_ts.shape == (n_time, n_edges)
assert edge_labels is not None
assert len(edge_labels) == n_edges assert len(edge_labels) == n_edges

View file

@ -229,7 +229,7 @@ def test_ParcelAggregation_3D_mask() -> None:
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="Schaefer100x7",
method="mean", method="mean",
mask="GM_prob0.2", masks="GM_prob0.2",
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
@ -274,7 +274,7 @@ def test_ParcelAggregation_3D_mask_computed() -> None:
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="Schaefer100x7",
method="mean", method="mean",
mask={"compute_brain_mask": {"threshold": 0.2}}, masks={"compute_brain_mask": {"threshold": 0.2}},
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument

View file

@ -159,7 +159,7 @@ def test_SphereAggregation_3D_mask() -> None:
method="mean", method="mean",
radius=RADIUS, radius=RADIUS,
on="VBM_GM", on="VBM_GM",
mask="GM_prob0.2", masks="GM_prob0.2",
) )
input = {"VBM_GM": {"data": img, "meta": {}}} input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"] jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]

View file

@ -7,7 +7,7 @@
# Federico Raimondo <f.raimondo@fz-juelich.de> # Federico Raimondo <f.raimondo@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any, Callable, Dict, List, Type, Union from typing import Any, Callable, Dict, List, Optional, Tuple, Type, Union
import numpy as np import numpy as np
import pandas as pd import pandas as pd
@ -56,8 +56,9 @@ def singleton(cls: Type) -> Type:
def _ets( def _ets(
bold_ts: np.ndarray, roi_names: Union[None, List[str]] = None bold_ts: np.ndarray,
) -> np.ndarray: roi_names: Union[None, List[str]] = None,
) -> Tuple[np.ndarray, Optional[List[str]]]:
"""Compute the edge-wise time series based on BOLD time series. """Compute the edge-wise time series based on BOLD time series.
Take a timeseries of brain areas, and calculate timeseries for each Take a timeseries of brain areas, and calculate timeseries for each
@ -80,9 +81,8 @@ def _ets(
edge-wise time series, i.e. estimate of functional connectivity at each edge-wise time series, i.e. estimate of functional connectivity at each
time point. time point.
edge_names : List[str] edge_names : List[str]
List of edge names corresponding to columns in the List of edge names corresponding to columns in the edge-wise time
edge-wise time series. This is only returned if the roi_names series. If roi_names are not specified, this is None.
are specified.
References References
---------- ----------
@ -102,18 +102,18 @@ def _ets(
ets = timeseries[:, u] * timeseries[:, v] ets = timeseries[:, u] * timeseries[:, v]
# Obtain the corresponding edge labels if specified else return # Obtain the corresponding edge labels if specified else return
if roi_names is None: if roi_names is None:
return ets return ets, None
else: else:
if len(roi_names) != n_roi: if len(roi_names) != n_roi:
raise_error( raise_error(
"List of roi names does not correspond " "List of roi names does not correspond "
"to the number of ROIs in the timeseries!" "to the number of ROIs in the timeseries!"
) )
roi_names = np.array(roi_names) _roi_names = np.array(roi_names)
edge_names = [ edge_names = [
"~".join([x, y]) for x, y in zip(roi_names[u], roi_names[v]) "~".join([x, y]) for x, y in zip(_roi_names[u], _roi_names[v])
] ]
return ets, list(edge_names) return ets, edge_names
def _correlate_dataframes( def _correlate_dataframes(

View file

@ -115,12 +115,22 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
if type_ in input.keys(): if type_ in input.keys():
logger.info(f"Computing {type_}") logger.info(f"Computing {type_}")
t_input = input[type_] t_input = input[type_]
extra_input = input.copy()
# Pass the other data types as extra input, removing
# the current type
extra_input = input
extra_input.pop(type_) extra_input.pop(type_)
key, t_out = self.preprocess( key, t_out = self.preprocess(
input=t_input, extra_input=extra_input input=t_input, extra_input=extra_input
) )
# Add the output to the Junifer Data object
out[key] = t_out out[key] = t_out
# In case we are creating a new type, re-add the original input
if key != type_:
out[type_] = t_input
self.update_meta(out[key], "preprocess") self.update_meta(out[key], "preprocess")
return out return out

View file

@ -11,9 +11,9 @@ import numpy as np
import pandas as pd import pandas as pd
from nilearn._utils.niimg_conversions import check_niimg_4d from nilearn._utils.niimg_conversions import check_niimg_4d
from nilearn.image import clean_img from nilearn.image import clean_img
from nilearn.masking import compute_brain_mask
from ...api.decorators import register_preprocessor from ...api.decorators import register_preprocessor
from ...data import get_mask
from ...utils import logger, raise_error from ...utils import logger, raise_error
from ..base import BasePreprocessor from ..base import BasePreprocessor
@ -134,11 +134,10 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
t_r : float, optional t_r : float, optional
Repetition time, in second (sampling period). Repetition time, in second (sampling period).
If None, it will use t_r from nifti header (default None). If None, it will use t_r from nifti header (default None).
mask_img: Niimg-like object, optional masks : str, dict or list of dict or str, optional
If provided, signal is only cleaned from voxels inside the mask. The specification of the masks to apply to regions before extracting
If mask is provided, it should have same shape and affine as imgs. signals. Check :ref:`Using Masks <using_masks>` for more details.
If not provided, a mask is computed using If None, will not apply any mask (default None).
:func:`nilearn.masking.compute_brain_mask` (default None).
""" """
@ -153,7 +152,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
low_pass: Optional[float] = None, low_pass: Optional[float] = None,
high_pass: Optional[float] = None, high_pass: Optional[float] = None,
t_r: Optional[float] = None, t_r: Optional[float] = None,
mask_img: Optional["Nifti1Image"] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
) -> None: ) -> None:
"""Initialise the class.""" """Initialise the class."""
if strategy is None: if strategy is None:
@ -169,7 +168,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
self.low_pass = low_pass self.low_pass = low_pass
self.high_pass = high_pass self.high_pass = high_pass
self.t_r = t_r self.t_r = t_r
self.mask_img = mask_img self.masks = masks
self._valid_components = ["motion", "wm_csf", "global_signal"] self._valid_components = ["motion", "wm_csf", "global_signal"]
self._valid_confounds = ["basic", "power2", "derivatives", "full"] self._valid_confounds = ["basic", "power2", "derivatives", "full"]
@ -521,19 +520,20 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
raise ValueError(f"Invalid confounds format {t_format}") raise ValueError(f"Invalid confounds format {t_format}")
def _remove_confounds( def _remove_confounds(
self, bold_img: "Nifti1Image", confounds_df: pd.DataFrame self,
input: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None,
) -> Union["Nifti1Image", "Nifti2Image", "MGHImage", List]: ) -> Union["Nifti1Image", "Nifti2Image", "MGHImage", List]:
"""Remove confounds from the BOLD image. """Remove confounds from the BOLD image.
Parameters Parameters
---------- ----------
bold_img : Niimg-like object input : dict
4D image. The signals in the last dimension are filtered Dictionary containing the ``BOLD`` value from the
(see http://nilearn.github.io/manipulating_images/input_output.html Junifer Data object.
for a detailed description of the valid input types). extra_input : dict, optional
confounds_df : pd.DataFrame Dictionary containing the rest of the Junifer Data object. Must
Dataframe containing confounds to remove. Number of rows should include the ``BOLD_confounds`` key.
correspond to number of volumes in the BOLD image.
Returns Returns
-------- --------
@ -541,8 +541,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
Input image with confounds removed. Input image with confounds removed.
""" """
assert extra_input is not None # Not the case, data is validated
confounds_df = self._pick_confounds(extra_input["BOLD_confounds"])
confounds_array = confounds_df.values confounds_array = confounds_df.values
bold_img = input["data"]
t_r = self.t_r t_r = self.t_r
if t_r is None: if t_r is None:
logger.info("No `t_r` specified, using t_r from nifti header") logger.info("No `t_r` specified, using t_r from nifti header")
@ -552,10 +555,17 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
f"Read t_r from nifti header: {t_r}", f"Read t_r from nifti header: {t_r}",
) )
mask_img = self.mask_img mask_img = None
if mask_img is None: if self.masks is not None:
logger.info("Computing brain mask from image") logger.debug(f"Masking with {self.masks}")
mask_img = compute_brain_mask(bold_img) mask_img = get_mask(
masks=self.masks, target_data=input, extra_input=extra_input
)
# Save the mask in the extra input and link it to the bold data
# this allows to use "inherit" down the pipeline
if extra_input is not None:
extra_input["BOLD_mask"] = {"data": mask_img}
input["mask_item"] = "BOLD_mask"
logger.info("Cleaning image") logger.info("Cleaning image")
logger.debug(f"\tdetrend: {self.detrend}") logger.debug(f"\tdetrend: {self.detrend}")
@ -600,8 +610,5 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
""" """
self._validate_data(input, extra_input) self._validate_data(input, extra_input)
assert extra_input is not None input["data"] = self._remove_confounds(input, extra_input=extra_input)
bold_img = input["data"]
confounds_df = self._pick_confounds(extra_input["BOLD_confounds"])
input["data"] = self._remove_confounds(bold_img, confounds_df)
return "BOLD", input return "BOLD", input

View file

@ -452,10 +452,10 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None:
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"] input = dg["sub-01"]
input = reader.fit_transform(input) input = reader.fit_transform(input)
confounds = confound_remover._pick_confounds(input["BOLD_confounds"])
raw_bold = input["BOLD"]["data"] raw_bold = input["BOLD"]["data"]
extra_input = {k: v for k, v in input.items() if k != "BOLD"}
clean_bold = confound_remover._remove_confounds( clean_bold = confound_remover._remove_confounds(
bold_img=raw_bold, confounds_df=confounds input=input["BOLD"], extra_input=extra_input
) )
clean_bold = typing.cast(nib.Nifti1Image, clean_bold) clean_bold = typing.cast(nib.Nifti1Image, clean_bold)
# TODO: Find a better way to test functionality here # TODO: Find a better way to test functionality here
@ -533,7 +533,63 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None:
assert t_meta["low_pass"] is None assert t_meta["low_pass"] is None
assert t_meta["high_pass"] is None assert t_meta["high_pass"] is None
assert t_meta["t_r"] is None assert t_meta["t_r"] is None
assert t_meta["mask_img"] is None assert t_meta["masks"] is None
assert "dependencies" in output["BOLD"]["meta"]
dependencies = output["BOLD"]["meta"]["dependencies"]
assert dependencies == {"numpy", "nilearn"}
def test_fMRIPrepConfoundRemover_fit_transform_masks() -> None:
"""Test fMRIPrepConfoundRemover with all confounds present."""
# need reader for the data
reader = DefaultDataReader()
# All strategies full, no spike
confound_remover = fMRIPrepConfoundRemover(
masks={"compute_brain_mask": {"threshold": 0.2}}
)
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"]
input = reader.fit_transform(input)
orig_bold = input["BOLD"]["data"].get_fdata().copy()
output = confound_remover.fit_transform(input)
trans_bold = output["BOLD"]["data"].get_fdata()
# Transformation is in place
assert_array_equal(trans_bold, input["BOLD"]["data"].get_fdata())
# Data should have the same shape
assert orig_bold.shape == trans_bold.shape
# but be different
assert_raises(
AssertionError, assert_array_equal, orig_bold, trans_bold
)
assert "meta" in output["BOLD"]
assert "preprocess" in output["BOLD"]["meta"]
t_meta = output["BOLD"]["meta"]["preprocess"]
assert t_meta["class"] == "fMRIPrepConfoundRemover"
# It should have all the default parameters
assert t_meta["strategy"] == confound_remover.strategy
assert t_meta["spike"] is None
assert t_meta["detrend"] is True
assert t_meta["standardize"] is True
assert t_meta["low_pass"] is None
assert t_meta["high_pass"] is None
assert t_meta["t_r"] is None
assert isinstance(t_meta["masks"], dict)
assert t_meta["masks"] is not None
assert len(t_meta["masks"]) == 1
assert "compute_brain_mask" in t_meta["masks"]
assert len(t_meta["masks"]["compute_brain_mask"]) == 1
assert "threshold" in t_meta["masks"]["compute_brain_mask"]
assert t_meta["masks"]["compute_brain_mask"]["threshold"] == 0.2
assert "BOLD_mask" in output
assert "mask_item" in output["BOLD"]
assert output["BOLD"]["mask_item"] == "BOLD_mask"
assert "dependencies" in output["BOLD"]["meta"] assert "dependencies" in output["BOLD"]["meta"]
dependencies = output["BOLD"]["meta"]["dependencies"] dependencies = output["BOLD"]["meta"]["dependencies"]