[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
running
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_brain_mask,
compute_epi_mask,
intersect_masks,
)
from ..utils.logging import logger, raise_error
@ -36,7 +37,10 @@ if TYPE_CHECKING:
_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.
Parameters
@ -65,6 +69,9 @@ 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.
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]] = {
"GM_prob0.2": {"family": "Vickery-Patil"},
@ -146,19 +153,23 @@ def list_masks() -> List[str]:
def get_mask(
mask: Union[str, Dict],
masks: Union[str, Dict, List[Union[Dict, str]]],
target_data: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None,
) -> "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.
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
of the mask.
of the mask as a dictionary. Several masks can be passed as a list.
target_data : dict
The corresponding item of the data object to which the mask will be
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
-------
@ -167,20 +178,68 @@ def get_mask(
"""
# Get the min of the voxels sizes and use it as the resolution
target_img = target_data["data"]
inherited_mask_item = target_data.get("mask_item", None)
resolution = np.min(target_img.header.get_zooms()[:3])
if isinstance(mask, dict):
if len(mask) != 1:
if not isinstance(masks, list):
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(
"The mask dictionary must have only one key, "
"the name of the mask."
"Each of the masks dictionary must have only one key, "
"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:
mask_name = mask
mask_name = t_mask
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_name, path_only=False, resolution=resolution
)
@ -190,13 +249,26 @@ def get_mask(
mask_img = mask_object(target_img, **mask_params)
else: # Mask is a Nifti1Image
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_object,
target_img,
interpolation="nearest",
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

View file

@ -6,8 +6,9 @@
# License: AGPL
from pathlib import Path
from typing import Callable, Dict, Union
from typing import Callable, Dict, List, Union
import numpy as np
import pytest
from nilearn.datasets import fetch_icbm152_brain_gm_mask
from nilearn.image import resample_to_img
@ -15,6 +16,7 @@ from nilearn.masking import (
compute_background_mask,
compute_brain_mask,
compute_epi_mask,
intersect_masks,
)
from numpy.testing import assert_array_almost_equal, assert_array_equal
@ -56,7 +58,9 @@ def test_register_mask_already_registered() -> None:
name="testmask",
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
with pytest.raises(ValueError, match=r"already registered."):
@ -70,7 +74,9 @@ def test_register_mask_already_registered() -> None:
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(
@ -110,6 +116,7 @@ def test_register_mask(
# Load registered mask
_, fname = load_mask(name=name, path_only=True)
# Check values for registered mask
assert fname is not None
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
)
assert fname is not None
assert fname.name == "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz"
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
)
assert fname is not None
assert (
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
)
assert fname is not None
assert fname.name == "GMprob0.2_cortex_3mm_NA_rm.nii.gz"
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)
vbm_gm = input["VBM_GM"]
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_array_equal(mask.affine, vbm_gm_img.affine)
@ -204,7 +214,7 @@ def test_mask_callable() -> None:
input = reader.fit_transform(input)
vbm_gm = input["VBM_GM"]
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())
@ -218,11 +228,48 @@ def test_get_mask_errors() -> None:
input = dg["sub-01"]
input = reader.fit_transform(input)
vbm_gm = input["VBM_GM"]
# Test wrong masks definitions (more than one key per dict)
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"):
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(
@ -271,7 +318,7 @@ def test_nilearn_compute_masks(
else:
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)
@ -287,3 +334,114 @@ def test_nilearn_compute_masks(
copy=True,
)
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
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).
masks : str, dict or list of dict or str, optional
The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None).
name : str, optional
The name of the marker. If None, will use the class name (default
None).
@ -49,13 +49,13 @@ class RSSETSMarker(BaseMarker):
parcellation: Union[str, List[str]],
agg_method: str = "mean",
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,
) -> None:
self.parcellation = parcellation
self.agg_method = agg_method
self.agg_method_params = agg_method_params
self.mask = mask
self.masks = masks
super().__init__(name=name)
def get_valid_inputs(self) -> List[str]:
@ -126,11 +126,11 @@ class RSSETSMarker(BaseMarker):
parcellation=self.parcellation,
method=self.agg_method,
method_params=self.agg_method_params,
mask=self.mask,
masks=self.masks,
)
# Compute the parcel aggregation
out = parcel_aggregation.compute(input=input, extra_input=extra_input)
edge_ts = _ets(out["data"])
edge_ts, _ = _ets(out["data"])
# Compute the RSS
out["data"] = np.sum(edge_ts**2, 1) ** 0.5
# Set correct column label

View file

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

View file

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

View file

@ -30,10 +30,10 @@ class CrossParcellationFC(BaseMarker):
correlation_method : str, optional
Any method that can be passed to
:any:`pandas.DataFrame.corr` (default "pearson").
mask : str, optional
The name of the mask to apply to regions before extracting signals.
Check valid options by calling :func:`junifer.data.masks.list_masks`
(default None).
masks : str, dict or list of dict or str, optional
The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None).
name : str, optional
The name of the marker. If None, will use the class name
(default None).
@ -47,7 +47,7 @@ class CrossParcellationFC(BaseMarker):
parcellation_two: str,
aggregation_method: str = "mean",
correlation_method: str = "pearson",
mask: Union[str, Dict, None] = None,
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
name: Optional[str] = None,
) -> None:
if parcellation_one == parcellation_two:
@ -58,7 +58,7 @@ class CrossParcellationFC(BaseMarker):
self.parcellation_two = parcellation_two
self.aggregation_method = aggregation_method
self.correlation_method = correlation_method
self.mask = mask
self.masks = masks
super().__init__(on=["BOLD"], name=name)
def get_valid_inputs(self) -> List[str]:
@ -129,12 +129,12 @@ class CrossParcellationFC(BaseMarker):
parcellation_one_dict = ParcelAggregation(
parcellation=self.parcellation_one,
method=self.aggregation_method,
mask=self.mask,
masks=self.masks,
).compute(input)
parcellation_two_dict = ParcelAggregation(
parcellation=self.parcellation_two,
method=self.aggregation_method,
mask=self.mask,
masks=self.masks,
).compute(input)
parcellated_ts_one = parcellation_one_dict["data"]

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -6,10 +6,12 @@
import pytest
# 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:
"""Test FunctionalConnectivityBase is an abstract base class."""
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
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).
masks : str, dict or list of dict or str, optional
The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (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
@ -52,7 +52,7 @@ class ParcelAggregation(BaseMarker):
parcellation: Union[str, List[str]],
method: str,
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,
name: Optional[str] = None,
) -> None:
@ -61,7 +61,7 @@ class ParcelAggregation(BaseMarker):
self.parcellation = parcellation
self.method = method
self.method_params = method_params or {}
self.mask = mask
self.masks = masks
super().__init__(on=on, name=name)
def get_valid_inputs(self) -> List[str]:
@ -184,9 +184,11 @@ class ParcelAggregation(BaseMarker):
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:
logger.debug(f"Masking with {self.mask}")
mask_img = get_mask(mask=self.mask, target_data=input)
if self.masks is not None:
logger.debug(f"Masking with {self.masks}")
mask_img = get_mask(
masks=self.masks, target_data=input, extra_input=extra_input
)
parcellation_bin = math_img(
"np.logical_and(img, mask)",

View file

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

View file

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

View file

@ -34,10 +34,10 @@ class SphereAggregation(BaseMarker):
(default "mean").
method_params : dict, optional
The parameters to pass to the aggregation method (default None).
mask : str, optional
The name of the mask to apply to regions before extracting signals.
Check valid options by calling :func:`junifer.data.masks.list_masks`
(default None).
masks : str, dict or list of dict or str, optional
The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (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
@ -56,7 +56,7 @@ class SphereAggregation(BaseMarker):
radius: Optional[float] = None,
method: str = "mean",
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,
name: Optional[str] = None,
) -> None:
@ -64,7 +64,7 @@ class SphereAggregation(BaseMarker):
self.radius = radius
self.method = method
self.method_params = method_params or {}
self.mask = mask
self.masks = masks
super().__init__(on=on, name=name)
def get_valid_inputs(self) -> List[str]:
@ -137,9 +137,11 @@ class SphereAggregation(BaseMarker):
)
# Load mask
mask_img = None
if self.mask is not None:
logger.debug(f"Masking with {self.mask}")
mask_img = get_mask(mask=self.mask, target_data=input)
if self.masks is not None:
logger.debug(f"Masking with {self.masks}")
mask_img = get_mask(
masks=self.masks, target_data=input, extra_input=extra_input
)
# Get seeds and labels
coords, out_labels = load_coordinates(name=self.coords)
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"):
MarkerCollection(wrong_markers)
MarkerCollection(wrong_markers) # type: ignore
def test_marker_collection() -> None:
@ -62,7 +62,7 @@ def test_marker_collection() -> None:
name="gmd_schaefer100x7_trim_mean90",
),
]
mc = MarkerCollection(markers=markers)
mc = MarkerCollection(markers=markers) # type: ignore
assert mc._markers == markers
assert mc._preprocessing is None
assert mc._storage is None
@ -96,7 +96,7 @@ def test_marker_collection() -> None:
return input
mc2 = MarkerCollection(
markers=markers,
markers=markers, # type: ignore
preprocessing=BypassPreprocessing(),
datareader=DefaultDataReader(),
)
@ -127,7 +127,7 @@ def test_marker_collection_with_preprocessing() -> None:
),
]
mc = MarkerCollection(
markers=markers,
markers=markers, # type: ignore
preprocessing=fMRIPrepConfoundRemover(),
)
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"
storage = SQLiteFeatureStorage(uri=uri)
mc = MarkerCollection(
markers=markers,
markers=markers, # type: ignore
storage=storage,
datareader=DefaultDataReader(),
)
@ -185,7 +185,9 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
out = mc.fit(input)
assert out is None
mc2 = MarkerCollection(markers=markers, datareader=DefaultDataReader())
mc2 = MarkerCollection(
markers=markers, datareader=DefaultDataReader() # type: ignore
)
mc2.validate(dg)
assert mc2._storage is None

View file

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

View file

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

View file

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

View file

@ -7,7 +7,7 @@
# Federico Raimondo <f.raimondo@fz-juelich.de>
# 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 pandas as pd
@ -56,8 +56,9 @@ def singleton(cls: Type) -> Type:
def _ets(
bold_ts: np.ndarray, roi_names: Union[None, List[str]] = None
) -> np.ndarray:
bold_ts: 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.
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
time point.
edge_names : List[str]
List of edge names corresponding to columns in the
edge-wise time series. This is only returned if the roi_names
are specified.
List of edge names corresponding to columns in the edge-wise time
series. If roi_names are not specified, this is None.
References
----------
@ -102,18 +102,18 @@ def _ets(
ets = timeseries[:, u] * timeseries[:, v]
# Obtain the corresponding edge labels if specified else return
if roi_names is None:
return ets
return ets, None
else:
if len(roi_names) != n_roi:
raise_error(
"List of roi names does not correspond "
"to the number of ROIs in the timeseries!"
)
roi_names = np.array(roi_names)
_roi_names = np.array(roi_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(

View file

@ -115,12 +115,22 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
if type_ in input.keys():
logger.info(f"Computing {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_)
key, t_out = self.preprocess(
input=t_input, extra_input=extra_input
)
# Add the output to the Junifer Data object
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")
return out

View file

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

View file

@ -452,10 +452,10 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None:
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"]
input = reader.fit_transform(input)
confounds = confound_remover._pick_confounds(input["BOLD_confounds"])
raw_bold = input["BOLD"]["data"]
extra_input = {k: v for k, v in input.items() if k != "BOLD"}
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)
# 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["high_pass"] 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"]
dependencies = output["BOLD"]["meta"]["dependencies"]