[ENH]: Add support for on-the-fly template space transformation #299

Merged
synchon merged 47 commits from feat/multi-mni-support into main 2024-04-03 16:02:20 +00:00
34 changed files with 2052 additions and 1485 deletions

View file

@ -44,6 +44,20 @@ jobs:
run: | run: |
git config --global user.email "runner@github.com" git config --global user.email "runner@github.com"
git config --global user.name "GITHUB CI Runner" git config --global user.name "GITHUB CI Runner"
- name: Install ANTs
run: |
echo "++ Add universe repo"
sudo add-apt-repository -y universe
echo "++ Update package manager info"
sudo apt-get update -qq
echo "++ Downloading ANTs"
curl -fsSL -o ants.zip https://github.com/ANTsX/ANTs/releases/download/v2.5.1/ants-2.5.1-ubuntu-22.04-X64-gcc.zip
unzip ants.zip -d /opt
mv /opt/ants-2.5.1/bin/* /opt/ants-2.5.1
rm ants.zip
echo "/opt/ants-2.5.1" >> $GITHUB_PATH
- name: Test build docs - name: Test build docs
run: | run: |
BUILDDIR=_build/main make -C docs/ local BUILDDIR=_build/main make -C docs/ local

View file

@ -0,0 +1 @@
Add ``template_type`` parameter to :func:`.get_template` by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Rewrite :func:`.compute_brain_mask` to allow variable template fetching via templateflow, according to target data by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Add support for on-the-fly template space transformation in :func:`.get_parcellation` and :func:`.get_mask` to allow parcellation and mask in different template spaces to work with a ``DataGrabber``'s data in a specified template space.

View file

@ -0,0 +1,16 @@
workdir: /tmp
with: junifer.testing.registry
datagrabber:
kind: PartlyCloudyTestingDataGrabber
markers:
- kind: ParcelAggregation
parcellation: TianxS1x3TxMNInonlinear2009cAsym
method: mean
name: tian-s1-3T_mean
storage:
kind: HDF5FeatureStorage
uri: /tmp/partly_cloudy_agg_mean_tian.hdf5

View file

@ -0,0 +1,16 @@
workdir: /tmp
with: junifer.testing.registry
datagrabber:
kind: PartlyCloudyTestingDataGrabber
markers:
- kind: ParcelAggregation
parcellation: TianxS1x3TxMNInonlinear2009cAsym
method: mean
name: tian-s1-3T_mean
storage:
kind: HDF5FeatureStorage
uri: /tmp/partly_cloudy_agg_mean_tian.hdf5

View file

@ -54,16 +54,13 @@ def test_run_and_collect_commands(
""" """
# Get test config # Get test config
infile = Path(__file__).parent / "data" / "gmd_mean.yaml" infile = Path(__file__).parent / "data" / "partly_cloudy_agg_mean_tian.yml"
# Read test config # Read test config
contents = yaml.load(infile) contents = yaml.load(infile)
# Working directory # Working directory
workdir = tmp_path / "workdir" contents["workdir"] = str(tmp_path.resolve())
contents["workdir"] = str(workdir.resolve())
# Output directory
outdir = tmp_path / "outdir"
# Storage # Storage
contents["storage"]["uri"] = str(outdir.resolve()) contents["storage"]["uri"] = str((tmp_path / "out.hdf5").resolve())
# Write new test config # Write new test config
outfile = tmp_path / "in.yaml" outfile = tmp_path / "in.yaml"
yaml.dump(contents, stream=outfile) yaml.dump(contents, stream=outfile)
@ -117,16 +114,13 @@ def test_run_using_element_file(tmp_path: Path, elements: str) -> None:
f.write(elements) f.write(elements)
# Get test config # Get test config
infile = Path(__file__).parent / "data" / "gmd_mean.yaml" infile = Path(__file__).parent / "data" / "partly_cloudy_agg_mean_tian.yml"
# Read test config # Read test config
contents = yaml.load(infile) contents = yaml.load(infile)
# Working directory # Working directory
workdir = tmp_path / "workdir" contents["workdir"] = str(tmp_path.resolve())
contents["workdir"] = str(workdir.resolve())
# Output directory
outdir = tmp_path / "outdir"
# Storage # Storage
contents["storage"]["uri"] = str(outdir.resolve()) contents["storage"]["uri"] = str((tmp_path / "out.hdf5").resolve())
# Write new test config # Write new test config
outfile = tmp_path / "in.yaml" outfile = tmp_path / "in.yaml"
yaml.dump(contents, stream=outfile) yaml.dump(contents, stream=outfile)
@ -228,7 +222,7 @@ def test_queue(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize( @pytest.mark.parametrize(
"action, action_file", "action, action_file",
[ [
(run, "gmd_mean.yaml"), (run, "partly_cloudy_agg_mean_tian.yml"),
(queue, "gmd_mean_htcondor.yaml"), (queue, "gmd_mean_htcondor.yaml"),
], ],
) )

View file

@ -7,7 +7,7 @@
import logging import logging
from pathlib import Path from pathlib import Path
from typing import List, Tuple, Union from typing import Dict, List, Tuple, Union
import pytest import pytest
from ruamel.yaml import YAML from ruamel.yaml import YAML
@ -24,97 +24,104 @@ yaml.default_flow_style = False
yaml.allow_unicode = True yaml.allow_unicode = True
yaml.indent(mapping=2, sequence=4, offset=2) yaml.indent(mapping=2, sequence=4, offset=2)
# Define datagrabber
datagrabber = {
"kind": "OasisVBMTestingDataGrabber",
}
# Define markers @pytest.fixture
markers = [ def datagrabber() -> Dict[str, str]:
"""Return a datagrabber as a dictionary."""
return {
"kind": "PartlyCloudyTestingDataGrabber",
}
@pytest.fixture
def markers() -> List[Dict[str, str]]:
"""Return markers as a list of dictionary."""
return [
{ {
"name": "Schaefer1000x7_Mean", "name": "tian-s1-3T_mean",
"kind": "ParcelAggregation", "kind": "ParcelAggregation",
"parcellation": "Schaefer1000x7", "parcellation": "TianxS1x3TxMNInonlinear2009cAsym",
"method": "mean", "method": "mean",
}, },
{ {
"name": "Schaefer1000x7_Std", "name": "tian-s1-3T_std",
"kind": "ParcelAggregation", "kind": "ParcelAggregation",
"parcellation": "Schaefer1000x7", "parcellation": "TianxS1x3TxMNInonlinear2009cAsym",
"method": "std", "method": "std",
}, },
] ]
# Define storage
storage = { @pytest.fixture
def storage() -> Dict[str, str]:
"""Return a storage as a dictionary."""
return {
"kind": "SQLiteFeatureStorage", "kind": "SQLiteFeatureStorage",
} }
def test_run_single_element(tmp_path: Path) -> None: def test_run_single_element(
tmp_path: Path,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None:
"""Test run function with single element. """Test run function with single element.
Parameters Parameters
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
# Create working directory # Set storage
workdir = tmp_path / "workdir_single" storage["uri"] = str((tmp_path / "out.sqlite").resolve())
workdir.mkdir()
# Create output directory
outdir = workdir / "out"
outdir.mkdir()
# Create storage
uri = outdir / "test.sqlite"
storage["uri"] = uri # type: ignore
# Run operations # Run operations
run( run(
workdir=workdir, workdir=tmp_path,
datagrabber=datagrabber, datagrabber=datagrabber,
markers=markers, markers=markers,
storage=storage, storage=storage,
elements=["sub-01"], elements=["sub-01"],
) )
# Check files # Check files
files = list(outdir.glob("*.sqlite")) files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 1 assert len(files) == 1
def test_run_single_element_with_preprocessing(tmp_path: Path) -> None: def test_run_single_element_with_preprocessing(
tmp_path: Path,
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None:
"""Test run function with single element and pre-processing. """Test run function with single element and pre-processing.
Parameters Parameters
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
# Create working directory # Set storage
workdir = tmp_path / "workdir_single_with_preprocess" storage["uri"] = str((tmp_path / "out.sqlite").resolve())
workdir.mkdir()
# Create output directory
outdir = workdir / "out"
outdir.mkdir()
# Create storage
uri = outdir / "test.sqlite"
storage["uri"] = uri # type: ignore
# Run operations # Run operations
run( run(
workdir=workdir, workdir=tmp_path,
datagrabber={ datagrabber={
"kind": "PartlyCloudyTestingDataGrabber", "kind": "PartlyCloudyTestingDataGrabber",
"reduce_confounds": False, "reduce_confounds": False,
}, },
markers=[ markers=markers,
{
"name": "Schaefer100x17_mean_FC",
"kind": "FunctionalConnectivityParcels",
"parcellation": "Schaefer100x17",
"agg_method": "mean",
}
],
storage=storage, storage=storage,
preprocessors=[ preprocessors=[
{ {
@ -124,97 +131,110 @@ def test_run_single_element_with_preprocessing(tmp_path: Path) -> None:
elements=["sub-01"], elements=["sub-01"],
) )
# Check files # Check files
files = list(outdir.glob("*.sqlite")) files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 1 assert len(files) == 1
def test_run_multi_element(tmp_path: Path) -> None: def test_run_multi_element_multi_output(
"""Test run function with multi element. tmp_path: Path,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None:
"""Test run function with multi element and multi output.
Parameters Parameters
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
# Create working directory # Set storage
workdir = tmp_path / "workdir_multi" storage["uri"] = str((tmp_path / "out.sqlite").resolve())
workdir.mkdir()
# Create output directory
outdir = workdir / "out"
outdir.mkdir()
# Create storage
uri = outdir / "test.sqlite"
storage["uri"] = uri # type: ignore
storage["single_output"] = False # type: ignore storage["single_output"] = False # type: ignore
# Run operations # Run operations
run( run(
workdir=workdir, workdir=tmp_path,
datagrabber=datagrabber, datagrabber=datagrabber,
markers=markers, markers=markers,
storage=storage, storage=storage,
elements=["sub-01", "sub-03"], elements=["sub-01", "sub-03"],
) )
# Check files # Check files
files = list(outdir.glob("*.sqlite")) files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 2 assert len(files) == 2
def test_run_multi_element_single_output(tmp_path: Path) -> None: def test_run_multi_element_single_output(
"""Test run function with multi element. tmp_path: Path,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None:
"""Test run function with multi element and single output.
Parameters Parameters
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
# Create working directory # Set storage
workdir = tmp_path / "workdir_multi" storage["uri"] = str((tmp_path / "out.sqlite").resolve())
workdir.mkdir()
# Create output directory
outdir = workdir / "out"
outdir.mkdir()
# Create storage
uri = outdir / "test.sqlite"
storage["uri"] = uri # type: ignore
storage["single_output"] = True # type: ignore storage["single_output"] = True # type: ignore
# Run operations # Run operations
run( run(
workdir=workdir, workdir=tmp_path,
datagrabber=datagrabber, datagrabber=datagrabber,
markers=markers, markers=markers,
storage=storage, storage=storage,
elements=["sub-01", "sub-03"], elements=["sub-01", "sub-03"],
) )
# Check files # Check files
files = list(outdir.glob("*.sqlite")) files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 1 assert len(files) == 1
assert files[0].name == "test.sqlite" assert files[0].name == "out.sqlite"
def test_run_and_collect(tmp_path: Path) -> None: def test_run_and_collect(
tmp_path: Path,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None:
"""Test run and collect functions. """Test run and collect functions.
Parameters Parameters
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
# Create working directory # Set storage
workdir = tmp_path / "workdir" uri = tmp_path / "out.sqlite"
workdir.mkdir() storage["uri"] = str(uri.resolve())
# Create output directory
outdir = workdir / "out"
outdir.mkdir()
# Create storage
uri = outdir / "test.sqlite"
storage["uri"] = uri # type: ignore
storage["single_output"] = False # type: ignore storage["single_output"] = False # type: ignore
# Run operations # Run operations
run( run(
workdir=workdir, workdir=tmp_path,
datagrabber=datagrabber, datagrabber=datagrabber,
markers=markers, markers=markers,
storage=storage, storage=storage,
@ -225,7 +245,7 @@ def test_run_and_collect(tmp_path: Path) -> None:
) )
elements = dg.get_elements() # type: ignore elements = dg.get_elements() # type: ignore
# This should create 10 files # This should create 10 files
files = list(outdir.glob("*.sqlite")) files = list(tmp_path.glob("*.sqlite"))
assert len(files) == len(elements) assert len(files) == len(elements)
# But the test.sqlite file should not exist # But the test.sqlite file should not exist
assert not uri.exists() assert not uri.exists()
@ -239,6 +259,9 @@ def test_queue_correct_yaml_config(
tmp_path: Path, tmp_path: Path,
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture, caplog: pytest.LogCaptureFixture,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None: ) -> None:
"""Test proper YAML config generation for queueing. """Test proper YAML config generation for queueing.
@ -250,6 +273,12 @@ def test_queue_correct_yaml_config(
The pytest.MonkeyPatch object. The pytest.MonkeyPatch object.
caplog : pytest.LogCaptureFixture caplog : pytest.LogCaptureFixture
The pytest.LogCaptureFixture object. The pytest.LogCaptureFixture object.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
with monkeypatch.context() as m: with monkeypatch.context() as m:
@ -261,7 +290,7 @@ def test_queue_correct_yaml_config(
"workdir": str(tmp_path.resolve()), "workdir": str(tmp_path.resolve()),
"datagrabber": datagrabber, "datagrabber": datagrabber,
"markers": markers, "markers": markers,
"storage": {"kind": "SQLiteFeatureStorage"}, "storage": storage,
"env": { "env": {
"kind": "conda", "kind": "conda",
"name": "junifer", "name": "junifer",
@ -479,6 +508,7 @@ def test_queue_without_elements(
tmp_path: Path, tmp_path: Path,
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture, caplog: pytest.LogCaptureFixture,
datagrabber: Dict[str, str],
) -> None: ) -> None:
"""Test queue without elements. """Test queue without elements.
@ -490,6 +520,8 @@ def test_queue_without_elements(
The pytest.MonkeyPatch object. The pytest.MonkeyPatch object.
caplog : pytest.LogCaptureFixture caplog : pytest.LogCaptureFixture
The pytest.LogCaptureFixture object. The pytest.LogCaptureFixture object.
datagrabber : dict
Testing datagrabber as dictionary.
""" """
with monkeypatch.context() as m: with monkeypatch.context() as m:
@ -502,13 +534,24 @@ def test_queue_without_elements(
assert "Queue done" in caplog.text assert "Queue done" in caplog.text
def test_reset_run(tmp_path: Path) -> None: def test_reset_run(
tmp_path: Path,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
) -> None:
"""Test reset function for run. """Test reset function for run.
Parameters Parameters
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
""" """
# Create storage # Create storage
@ -535,7 +578,12 @@ def test_reset_run(tmp_path: Path) -> None:
), ),
) )
def test_reset_queue( def test_reset_queue(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, job_name: str tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
datagrabber: Dict[str, str],
markers: List[Dict[str, str]],
storage: Dict[str, str],
job_name: str,
) -> None: ) -> None:
"""Test reset function for queue. """Test reset function for queue.
@ -545,6 +593,12 @@ def test_reset_queue(
The path to the test directory. The path to the test directory.
monkeypatch : pytest.MonkeyPatch monkeypatch : pytest.MonkeyPatch
The pytest.MonkeyPatch object. The pytest.MonkeyPatch object.
datagrabber : dict
Testing datagrabber as dictionary.
markers : list of dict
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
job_name : str job_name : str
The parametrized job name. The parametrized job name.

View file

@ -138,7 +138,7 @@ def register_coordinates(
voi_names : list of str voi_names : list of str
The names of the VOIs. The names of the VOIs.
space : str space : str
The space of the coordinates. The space of the coordinates, for e.g., "MNI".
overwrite : bool, optional overwrite : bool, optional
If True, overwrite an existing list of coordinates with the same name. If True, overwrite an existing list of coordinates with the same name.
Does not apply to built-in coordinates (default False). Does not apply to built-in coordinates (default False).

View file

@ -20,16 +20,16 @@ from typing import (
import nibabel as nib import nibabel as nib
import numpy as np import numpy as np
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 get_data, new_img_like, resample_to_img
from nilearn.masking import ( from nilearn.masking import (
compute_background_mask, compute_background_mask,
compute_brain_mask,
compute_epi_mask, compute_epi_mask,
intersect_masks, intersect_masks,
) )
from ..pipeline import WorkDirManager from ..pipeline import WorkDirManager
from ..utils import logger, raise_error, run_ext_cmd from ..utils import logger, raise_error, run_ext_cmd, warn_with_log
from .template_spaces import get_template, get_xfm
from .utils import closest_resolution from .utils import closest_resolution
@ -40,10 +40,91 @@ if TYPE_CHECKING:
_masks_path = Path(__file__).parent / "masks" _masks_path = Path(__file__).parent / "masks"
def compute_brain_mask(
target_data: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None,
mask_type: str = "brain",
threshold: float = 0.5,
) -> "Nifti1Image":
"""Compute the whole-brain, grey-matter or white-matter mask.
This mask is calculated using the template space and resolution as found
in the ``target_data``.
Parameters
----------
target_data : dict
The corresponding item of the data object for which mask will be
loaded.
extra_input : dict, optional
The other fields in the data object. Useful for accessing other data
types (default None).
mask_type : {"brain", "gm", "wm"}, optional
Type of mask to be computed:
* "brain" : whole-brain mask
* "gm" : grey-matter mask
* "wm" : white-matter mask
(default "brain").
threshold : float, optional
The value under which the template is cut off (default 0.5).
Returns
-------
Nifti1Image
The mask (3D image).
Raises
------
ValueError
If ``mask_type`` is invalid or
if ``extra_input`` is None when ``target_data``'s space is native.
"""
logger.debug(f"Computing {mask_type} mask")
if mask_type not in ["brain", "gm", "wm"]:
raise_error(f"Unknown mask type: {mask_type}")
# Check pre-requirements for space manipulation
target_space = target_data["space"]
# Set target standard space to target space
target_std_space = target_space
# Extra data type requirement check if target space is native
if target_space == "native":
# Check for extra inputs
if extra_input is None:
raise_error(
"No extra input provided, requires `Warp` "
"data type to infer target template space."
)
# Set target standard space to warp file space source
target_std_space = extra_input["Warp"]["src"]
# Fetch template in closest resolution
template = get_template(
space=target_std_space,
target_data=target_data,
extra_input=extra_input,
template_type=mask_type if mask_type in ["gm", "wm"] else "T1w",
)
# Resample template to target image
target_img = target_data["data"]
resampled_template = resample_to_img(
source_img=template, target_img=target_img
)
# Threshold and get mask
mask = (get_data(resampled_template) >= threshold).astype("int8")
return new_img_like(target_img, mask) # type: ignore
def _fetch_icbm152_brain_gm_mask( def _fetch_icbm152_brain_gm_mask(
target_img: "Nifti1Image", target_img: "Nifti1Image",
**kwargs, **kwargs,
): ) -> "Nifti1Image":
"""Fetch ICBM152 brain mask and resample. """Fetch ICBM152 brain mask and resample.
Parameters Parameters
@ -59,7 +140,20 @@ def _fetch_icbm152_brain_gm_mask(
nibabel.Nifti1Image nibabel.Nifti1Image
The resampled mask. The resampled mask.
Warns
-----
DeprecationWarning
If this function is used.
""" """
warn_with_log(
msg=(
"It is recommended to use ``compute_brain_mask`` with "
"``mask_type='gm'``. This function will be removed in the next "
"release. For now, it's available for backward compatibility."
),
category=DeprecationWarning,
)
mask = fetch_icbm152_brain_gm_mask(**kwargs) mask = fetch_icbm152_brain_gm_mask(**kwargs)
mask = resample_to_img( mask = resample_to_img(
mask, target_img, interpolation="nearest", copy=True mask, target_img, interpolation="nearest", copy=True
@ -123,7 +217,7 @@ def register_mask(
mask_path : str or pathlib.Path mask_path : str or pathlib.Path
The path to the mask file. The path to the mask file.
space : str space : str
The space of the mask. The space of the mask, for e.g., "MNI152NLin6Asym".
overwrite : bool, optional overwrite : bool, optional
If True, overwrite an existing mask with the same name. If True, overwrite an existing mask with the same name.
Does not apply to built-in mask (default False). Does not apply to built-in mask (default False).
@ -198,30 +292,45 @@ def get_mask( # noqa: C901
Raises Raises
------ ------
RuntimeError RuntimeError
If masks are in different spaces and they need to be intersected / If warp / transformation file extension is not ".mat" or ".h5" or
unionized or if fetch_icbm152_brain_gm_mask is used and requires warping to
if warp / transformation file extension is not ".mat" or ".h5". other template space.
ValueError ValueError
If extra key is provided in addition to mask name in ``masks`` or If extra key is provided in addition to mask name in ``masks`` or
if no mask is provided or if no mask is provided or
if ``masks = "inherit"`` but ``extra_input`` is None or ``mask_item`` if ``masks = "inherit"`` but ``extra_input`` is None or ``mask_item``
is None or ``mask_items``'s value is not in ``extra_input`` or is None or ``mask_items``'s value is not in ``extra_input`` or
if callable parameters are passed to non-callable mask or if callable parameters are passed to non-callable mask or
if multiple masks are provided and their spaces do not match or
if parameters are passed to :func:`nilearn.masking.intersect_masks` if parameters are passed to :func:`nilearn.masking.intersect_masks`
when there is only one mask or when there is only one mask or
if ``extra_input`` is None when ``target_data``'s space is native. if ``extra_input`` is None when ``target_data``'s space is native.
""" """
# Check pre-requirements for space manipulation
target_space = target_data["space"]
# Set target standard space to target space
target_std_space = target_space
# Extra data type requirement check if target space is native
if target_space == "native":
# Check for extra inputs
if extra_input is None:
raise_error(
"No extra input provided, requires `Warp` and `T1w` "
"data types in particular for transformation to "
f"{target_data['space']} space for further computation."
)
# Set target standard space to warp file space source
target_std_space = extra_input["Warp"]["src"]
# 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])
# Convert masks to list if not already
if not isinstance(masks, list): if not isinstance(masks, list):
masks = [masks] masks = [masks]
# Check that dicts have only one key # Check that masks passed as dicts have only one key
invalid_elements = [ invalid_elements = [
x for x in masks if isinstance(x, dict) and len(x) != 1 x for x in masks if isinstance(x, dict) and len(x) != 1
] ]
@ -248,9 +357,19 @@ def get_mask( # noqa: C901
if len(true_masks) == 0: if len(true_masks) == 0:
raise_error("No mask was passed. At least one mask is required.") raise_error("No mask was passed. At least one mask is required.")
# Get the data type for the input data type's mask
inherited_mask_item = target_data.get("mask_item", None)
# Create component-scoped tempdir
tempdir = WorkDirManager().get_tempdir(prefix="masks")
# Create element-scoped tempdir so that warped mask is
# available later as nibabel stores file path reference for
# loading on computation
element_tempdir = WorkDirManager().get_element_tempdir(prefix="masks")
# Get all the masks # Get all the masks
all_masks = [] all_masks = []
all_spaces = []
for t_mask in true_masks: for t_mask in true_masks:
if isinstance(t_mask, dict): if isinstance(t_mask, dict):
mask_name = next(iter(t_mask.keys())) mask_name = next(iter(t_mask.keys()))
@ -281,21 +400,40 @@ def get_mask( # noqa: C901
f"because the item ({inherited_mask_item}) does not exist." f"because the item ({inherited_mask_item}) does not exist."
) )
mask_img = extra_input[inherited_mask_item]["data"] mask_img = extra_input[inherited_mask_item]["data"]
mask_space = target_data["space"]
# Starting with new mask # Starting with new mask
else: else:
# Restrict fetch_icbm152_brain_gm_mask if target std space doesn't
# match
if (
mask_name == "fetch_icbm152_brain_gm_mask"
and target_std_space != "MNI152NLin2009aAsym"
):
raise_error(
(
"``fetch_icbm152_brain_gm_mask`` is deprecated and "
"space transformation to any other template space is "
"prohibited as it will lead to unforeseen errors. "
"``compute_brain_mask`` is a better alternative."
),
klass=RuntimeError,
)
# Load mask # Load mask
mask_object, _, mask_space = load_mask( mask_object, _, mask_space = load_mask(
mask_name, path_only=False, resolution=resolution mask_name, path_only=False, resolution=resolution
) )
# Replace mask space with target space if mask's space is inherit # Replace mask space with target space if mask's space is inherit
if mask_space == "inherit": if mask_space == "inherit":
mask_space = target_data["space"] mask_space = target_std_space
# If mask is callable like from nilearn # If mask is callable like from nilearn
if callable(mask_object): if callable(mask_object):
if mask_params is None: if mask_params is None:
mask_params = {} mask_params = {}
# From nilearn
if mask_name != "compute_brain_mask":
mask_img = mask_object(target_img, **mask_params) mask_img = mask_object(target_img, **mask_params)
# Not from nilearn
else:
mask_img = mask_object(target_data, **mask_params)
# Mask is a Nifti1Image # Mask is a Nifti1Image
else: else:
# Mask params provided # Mask params provided
@ -306,31 +444,69 @@ def get_mask( # noqa: C901
) )
# Resample mask to target image # Resample mask to target image
mask_img = resample_to_img( mask_img = resample_to_img(
mask_object, source_img=mask_object,
target_img, target_img=target_img,
interpolation="nearest", interpolation="nearest",
copy=True, copy=True,
) )
all_spaces.append(mask_space) # Convert mask space if required
if mask_space != target_std_space:
# Get xfm file
xfm_file_path = get_xfm(src=mask_space, dst=target_std_space)
# Get target standard space template
target_std_space_template_img = get_template(
space=target_std_space,
target_data=target_data,
extra_input=extra_input,
)
# Save mask image to a component-scoped tempfile
mask_path = tempdir / f"{mask_name}.nii.gz"
nib.save(mask_img, mask_path)
# Save template
target_std_space_template_path = (
tempdir / f"{target_std_space}_T1w_{resolution}.nii.gz"
)
nib.save(
target_std_space_template_img,
target_std_space_template_path,
)
# Set warped mask path
warped_mask_path = element_tempdir / (
f"{mask_name}_warped_from_{mask_space}_to_"
f"{target_std_space}.nii.gz"
)
logger.debug(
f"Using ANTs to warp {mask_name} "
f"from {mask_space} to {target_std_space}"
)
# Set antsApplyTransforms command
apply_transforms_cmd = [
"antsApplyTransforms",
"-d 3",
"-e 3",
"-n 'GenericLabel[NearestNeighbor]'",
f"-i {mask_path.resolve()}",
f"-r {target_std_space_template_path.resolve()}",
f"-t {xfm_file_path.resolve()}",
f"-o {warped_mask_path.resolve()}",
]
# Call antsApplyTransforms
run_ext_cmd(
name="antsApplyTransforms", cmd=apply_transforms_cmd
)
mask_img = nib.load(warped_mask_path)
all_masks.append(mask_img) all_masks.append(mask_img)
# Multiple masks, need intersection / union # Multiple masks, need intersection / union
if len(all_masks) > 1: if len(all_masks) > 1:
# Make a set of unique spaces # Intersect / union of masks
unique_spaces = set(all_spaces)
# Intersect / union of masks only if all masks are in the same space
if len(unique_spaces) == 1:
mask_img = intersect_masks(all_masks, **intersect_params) mask_img = intersect_masks(all_masks, **intersect_params)
# Store the mask space for further checks
mask_space = next(iter(unique_spaces))
else:
raise_error(
msg=(
f"Masks are in different spaces: {unique_spaces}, "
"unable to merge."
),
klass=RuntimeError,
)
# Single mask # Single mask
else: else:
if len(intersect_params) > 0: if len(intersect_params) > 0:
@ -340,30 +516,13 @@ def get_mask( # noqa: C901
"when there is only one mask." "when there is only one mask."
) )
mask_img = all_masks[0] mask_img = all_masks[0]
mask_space = all_spaces[0]
# Warp mask if target data is native and mask space is not native
if target_data["space"] == "native" and target_data["space"] != mask_space:
# Check for extra inputs
if extra_input is None:
raise_error(
"No extra input provided, requires `Warp` and `T1w` "
"data types in particular for transformation to "
f"{target_data['space']} space for further computation."
)
# Create component-scoped tempdir
tempdir = WorkDirManager().get_tempdir(prefix="masks")
# Warp mask if target data is native
if target_space == "native":
# Save mask image to a component-scoped tempfile # Save mask image to a component-scoped tempfile
prewarp_mask_path = tempdir / "prewarp_mask.nii.gz" prewarp_mask_path = tempdir / "prewarp_mask.nii.gz"
nib.save(mask_img, prewarp_mask_path) nib.save(mask_img, prewarp_mask_path)
# Create element-scoped tempdir so that warped mask is
# available later as nibabel stores file path reference for
# loading on computation
element_tempdir = WorkDirManager().get_element_tempdir(prefix="masks")
# Create an element-scoped tempfile for warped output # Create an element-scoped tempfile for warped output
warped_mask_path = element_tempdir / "mask_warped.nii.gz" warped_mask_path = element_tempdir / "mask_warped.nii.gz"

View file

@ -22,6 +22,7 @@ from nilearn import datasets, image
from ..pipeline import WorkDirManager from ..pipeline import WorkDirManager
from ..utils import logger, raise_error, run_ext_cmd, warn_with_log from ..utils import logger, raise_error, run_ext_cmd, warn_with_log
from .template_spaces import get_template, get_xfm
from .utils import closest_resolution from .utils import closest_resolution
@ -154,7 +155,7 @@ def register_parcellation(
parcels_labels : list of str parcels_labels : list of str
The list of labels for the parcellation. The list of labels for the parcellation.
space : str space : str
The space of the parcellation. The template space of the parcellation, for e.g., "MNI152NLin6Asym".
overwrite : bool, optional overwrite : bool, optional
If True, overwrite an existing parcellation with the same name. If True, overwrite an existing parcellation with the same name.
Does not apply to built-in parcellations (default False). Does not apply to built-in parcellations (default False).
@ -236,57 +237,17 @@ def get_parcellation(
Raises Raises
------ ------
RuntimeError RuntimeError
If parcellations are in different spaces and they need to be merged or If warp / transformation file extension is not ".mat" or ".h5".
if warp / transformation file extension is not ".mat" or ".h5".
ValueError ValueError
If ``extra_input`` is None when ``target_data``'s space is native. If ``extra_input`` is None when ``target_data``'s space is native.
""" """
# Get the min of the voxels sizes and use it as the resolution # Check pre-requirements for space manipulation
target_img = target_data["data"] target_space = target_data["space"]
resolution = np.min(target_img.header.get_zooms()[:3]) # Set target standard space to target space
target_std_space = target_space
# Load the parcellations # Extra data type requirement check if target space is native
all_parcellations = [] if target_space == "native":
all_labels = []
all_spaces = []
for name in parcellation:
img, labels, _, space = load_parcellation(
name=name,
resolution=resolution,
)
# Resample all of them to the image
resampled_img = image.resample_to_img(
source_img=img,
target_img=target_img,
interpolation="nearest",
copy=True,
)
all_parcellations.append(resampled_img)
all_labels.append(labels)
all_spaces.append(space)
# Avoid merging if there is only one parcellation
if len(all_parcellations) == 1:
resampled_parcellation_img = all_parcellations[0]
labels = all_labels[0]
else:
# Merge the parcellations only if all parcellations are in the same
# space
if len(set(all_spaces)) == 1:
resampled_parcellation_img, labels = merge_parcellations(
parcellations_list=all_parcellations,
parcellations_names=parcellation,
labels_lists=all_labels,
)
else:
raise_error(
msg="Parcellations are in different spaces, unable to merge.",
klass=RuntimeError,
)
# Warp parcellation if target data is native
if target_data["space"] == "native":
# Check for extra inputs # Check for extra inputs
if extra_input is None: if extra_input is None:
raise_error( raise_error(
@ -294,20 +255,108 @@ def get_parcellation(
"data types in particular for transformation to " "data types in particular for transformation to "
f"{target_data['space']} space for further computation." f"{target_data['space']} space for further computation."
) )
# Set target standard space to warp file space source
target_std_space = extra_input["Warp"]["src"]
# Get the min of the voxels sizes and use it as the resolution
target_img = target_data["data"]
resolution = np.min(target_img.header.get_zooms()[:3])
# Create component-scoped tempdir # Create component-scoped tempdir
tempdir = WorkDirManager().get_tempdir(prefix="parcellations") tempdir = WorkDirManager().get_tempdir(prefix="parcellations")
# Save parcellation image to a component-scoped tempfile
prewarp_parcellation_path = tempdir / "prewarp_parcellation.nii.gz"
nib.save(resampled_parcellation_img, prewarp_parcellation_path)
# Create element-scoped tempdir so that warped parcellation is # Create element-scoped tempdir so that warped parcellation is
# available later as nibabel stores file path reference for # available later as nibabel stores file path reference for
# loading on computation # loading on computation
element_tempdir = WorkDirManager().get_element_tempdir( element_tempdir = WorkDirManager().get_element_tempdir(
prefix="parcellations" prefix="parcellations"
) )
# Load the parcellations
all_parcellations = []
all_labels = []
for name in parcellation:
img, labels, _, space = load_parcellation(
name=name,
resolution=resolution,
)
# Convert parcellation spaces if required
if space != target_std_space:
# Get xfm file
xfm_file_path = get_xfm(src=space, dst=target_std_space)
# Get target standard space template
target_std_space_template_img = get_template(
space=target_std_space,
target_data=target_data,
extra_input=extra_input,
)
# Save parcellation image to a component-scoped tempfile
parcellation_path = tempdir / f"{name}.nii.gz"
nib.save(img, parcellation_path)
# Save template
target_std_space_template_path = (
tempdir / f"{target_std_space}_T1w_{resolution}.nii.gz"
)
nib.save(
target_std_space_template_img, target_std_space_template_path
)
# Set warped parcellation path
warped_parcellation_path = element_tempdir / (
f"{name}_warped_from_{space}_to_" f"{target_std_space}.nii.gz"
)
logger.debug(
f"Using ANTs to warp {name} "
f"from {space} to {target_std_space}"
)
# Set antsApplyTransforms command
apply_transforms_cmd = [
"antsApplyTransforms",
"-d 3",
"-e 3",
"-n 'GenericLabel[NearestNeighbor]'",
f"-i {parcellation_path.resolve()}",
f"-r {target_std_space_template_path.resolve()}",
f"-t {xfm_file_path.resolve()}",
f"-o {warped_parcellation_path.resolve()}",
]
# Call antsApplyTransforms
run_ext_cmd(name="antsApplyTransforms", cmd=apply_transforms_cmd)
img = nib.load(warped_parcellation_path)
# Resample parcellation to target image
img_to_merge = image.resample_to_img(
source_img=img,
target_img=target_img,
interpolation="nearest",
copy=True,
)
all_parcellations.append(img_to_merge)
all_labels.append(labels)
# Avoid merging if there is only one parcellation
if len(all_parcellations) == 1:
resampled_parcellation_img = all_parcellations[0]
labels = all_labels[0]
# Parcellations are already transformed to target standard space
else:
resampled_parcellation_img, labels = merge_parcellations(
parcellations_list=all_parcellations,
parcellations_names=parcellation,
labels_lists=all_labels,
)
# Warp parcellation if target space is native
if target_space == "native":
# Save parcellation image to a component-scoped tempfile
prewarp_parcellation_path = tempdir / "prewarp_parcellation.nii.gz"
nib.save(resampled_parcellation_img, prewarp_parcellation_path)
# Create an element-scoped tempfile for warped output # Create an element-scoped tempfile for warped output
warped_parcellation_path = ( warped_parcellation_path = (
element_tempdir / "parcellation_warped.nii.gz" element_tempdir / "parcellation_warped.nii.gz"

View file

@ -99,6 +99,7 @@ def get_template(
space: str, space: str,
target_data: Dict[str, Any], target_data: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None, extra_input: Optional[Dict[str, Any]] = None,
template_type: str = "T1w",
) -> nib.Nifti1Image: ) -> nib.Nifti1Image:
"""Get template for the space, tailored for the target image. """Get template for the space, tailored for the target image.
@ -112,6 +113,8 @@ def get_template(
extra_input : dict, optional extra_input : dict, optional
The other fields in the data object. Useful for accessing other data The other fields in the data object. Useful for accessing other data
types (default None). types (default None).
template_type : {"T1w", "brain", "gm", "wm", "csf"}, optional
The template type to retrieve (default "T1w").
Returns Returns
------- -------
@ -121,15 +124,19 @@ def get_template(
Raises Raises
fraimondo commented 2024-03-27 15:00:50 +00:00 (Migrated from github.com)

What do you mean with "whole". Does it come from any toolbox? I know that usually it is named "brain" mask to reference the GM, WM and CSF masks.

What do you mean with `"whole"`. Does it come from any toolbox? I know that usually it is named "brain" mask to reference the GM, WM and CSF masks.
synchon commented 2024-03-27 15:15:25 +00:00 (Migrated from github.com)

It means "whole brain", but since I couldn't find an alternative for it, I used "whole".

It means "whole brain", but since I couldn't find an alternative for it, I used "whole".
fraimondo commented 2024-03-28 09:24:21 +00:00 (Migrated from github.com)

it should be "brain" then. That's the brain mask.

it should be "brain" then. That's the brain mask.
synchon commented 2024-03-28 09:38:13 +00:00 (Migrated from github.com)

Addressed.

Addressed.
------ ------
ValueError ValueError
If ``space`` is invalid. If ``space`` or ``template_type`` is invalid.
RuntimeError RuntimeError
If template in the required resolution is not found. If required template is not found.
""" """
# Check for invalid space; early check to raise proper error # Check for invalid space; early check to raise proper error
if space not in tflow.templates(): if space not in tflow.templates():
raise_error(f"Unknown template space: {space}") raise_error(f"Unknown template space: {space}")
# Check for template type
if template_type not in ["T1w", "brain", "gm", "wm", "csf"]:
raise_error(f"Unknown template type: {template_type}")
# 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"]
resolution = np.min(target_img.header.get_zooms()[:3]).astype(int) resolution = np.min(target_img.header.get_zooms()[:3]).astype(int)
@ -145,18 +152,38 @@ def get_template(
logger.info(f"Downloading template {space} in resolution {resolution}") logger.info(f"Downloading template {space} in resolution {resolution}")
# Retrieve template # Retrieve template
try: try:
suffix = None
desc = None
label = None
if template_type == "T1w":
suffix = template_type
desc = None
label = None
elif template_type == "brain":
suffix = "mask"
desc = "brain"
label = None
elif template_type in ["gm", "wm", "csf"]:
suffix = "probseg"
desc = None
label = template_type.upper()
# Set kwargs for fetching
kwargs = {
"suffix": suffix,
"desc": desc,
"label": label,
}
template_path = tflow.get( template_path = tflow.get(
space, space,
raise_empty=True, raise_empty=True,
resolution=resolution, resolution=resolution,
suffix="T1w",
desc=None,
extension="nii.gz", extension="nii.gz",
**kwargs,
) )
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
raise_error( raise_error(
f"Template {space} not found in the required resolution " f"Template {space} ({template_type}) with resolution {resolution} "
f"{resolution}", "not found",
klass=RuntimeError, klass=RuntimeError,
) )
else: else:

View file

@ -5,16 +5,16 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import socket
from pathlib import Path from pathlib import Path
from typing import Callable, Dict, List, Optional, Union from typing import Callable, Dict, List, Optional, Union
import nibabel as nib
import numpy as np import numpy as np
import pytest import pytest
from nilearn.datasets import fetch_icbm152_brain_gm_mask
from nilearn.image import resample_to_img from nilearn.image import resample_to_img
from nilearn.masking import ( from nilearn.masking import (
compute_background_mask, compute_background_mask,
compute_brain_mask,
compute_epi_mask, compute_epi_mask,
intersect_masks, intersect_masks,
) )
@ -23,18 +23,96 @@ from numpy.testing import assert_array_almost_equal, assert_array_equal
from junifer.data.masks import ( from junifer.data.masks import (
_available_masks, _available_masks,
_load_vickery_patil_mask, _load_vickery_patil_mask,
compute_brain_mask,
get_mask, get_mask,
list_masks, list_masks,
load_mask, load_mask,
register_mask, register_mask,
) )
from junifer.datagrabber import DMCC13Benchmark
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.testing.datagrabbers import ( from junifer.testing.datagrabbers import (
OasisVBMTestingDataGrabber, OasisVBMTestingDataGrabber,
PartlyCloudyTestingDataGrabber,
SPMAuditoryTestingDataGrabber, SPMAuditoryTestingDataGrabber,
) )
@pytest.mark.parametrize(
"mask_type, threshold",
[
("brain", 0.2),
("brain", 0.5),
("brain", 0.8),
("gm", 0.2),
("gm", 0.5),
("gm", 0.8),
("wm", 0.2),
("wm", 0.5),
("wm", 0.8),
],
)
def test_compute_brain_mask(mask_type: str, threshold: float) -> None:
"""Test compute_brain_mask().
Parameters
----------
mask_type : str
The parametrized mask type.
threshold : float
The parametrized threshold.
"""
with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
mask = compute_brain_mask(
target_data=element_data["BOLD"],
extra_input=None,
mask_type=mask_type,
)
assert isinstance(mask, nib.Nifti1Image)
@pytest.mark.skipif(
socket.gethostname() != "juseless",
reason="only for juseless",
)
@pytest.mark.parametrize(
"mask_type",
[
"brain",
"gm",
"wm",
],
)
def test_compute_brain_mask_for_native(mask_type: str) -> None:
"""Test compute_brain_mask().
Parameters
----------
mask_type : str
The parametrized mask type.
"""
with DMCC13Benchmark(
types=["BOLD"],
sessions=["wave1bas"],
tasks=["Rest"],
phase_encodings=["AP"],
runs=["1"],
native_t1w=True,
) as dg:
element_data = DefaultDataReader().fit_transform(
dg[("f1031ax", "wave1bas", "Rest", "AP", "1")]
)
mask = compute_brain_mask(
target_data=element_data["BOLD"],
extra_input=None,
mask_type=mask_type,
)
assert isinstance(mask, nib.Nifti1Image)
def test_register_mask_built_in_check() -> None: def test_register_mask_built_in_check() -> None:
"""Test mask registration check for built-in masks.""" """Test mask registration check for built-in masks."""
with pytest.raises(ValueError, match=r"built-in mask"): with pytest.raises(ValueError, match=r"built-in mask"):
@ -215,18 +293,19 @@ def test_vickery_patil_error() -> None:
def test_get_mask() -> None: def test_get_mask() -> None:
"""Test the get_mask function.""" """Test the get_mask function."""
reader = DefaultDataReader()
with OasisVBMTestingDataGrabber() as dg: with OasisVBMTestingDataGrabber() as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) vbm_gm = element_data["VBM_GM"]
vbm_gm = input["VBM_GM"]
vbm_gm_img = vbm_gm["data"] vbm_gm_img = vbm_gm["data"]
mask = get_mask(masks="GM_prob0.2", target_data=vbm_gm) mask = get_mask(masks="compute_brain_mask", 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)
raw_mask_img, _, _ = load_mask("GM_prob0.2", resolution=1.5) raw_mask_callable, _, _ = load_mask(
"compute_brain_mask", resolution=1.5
)
raw_mask_img = raw_mask_callable(vbm_gm) # type: ignore
res_mask_img = resample_to_img( res_mask_img = resample_to_img(
raw_mask_img, raw_mask_img,
vbm_gm_img, vbm_gm_img,
@ -245,13 +324,11 @@ def test_mask_callable() -> None:
_available_masks["identity"] = { _available_masks["identity"] = {
"family": "Callable", "family": "Callable",
"func": ident, "func": ident,
"space": "MNI", "space": "MNI152Lin",
} }
reader = DefaultDataReader()
with OasisVBMTestingDataGrabber() as dg: with OasisVBMTestingDataGrabber() as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) vbm_gm = element_data["VBM_GM"]
vbm_gm = input["VBM_GM"]
vbm_gm_img = vbm_gm["data"] vbm_gm_img = vbm_gm["data"]
mask = get_mask(masks="identity", target_data=vbm_gm) mask = get_mask(masks="identity", target_data=vbm_gm)
@ -262,11 +339,9 @@ def test_mask_callable() -> None:
def test_get_mask_errors() -> None: def test_get_mask_errors() -> None:
"""Test passing wrong parameters to get_mask.""" """Test passing wrong parameters to get_mask."""
reader = DefaultDataReader()
with OasisVBMTestingDataGrabber() as dg: with OasisVBMTestingDataGrabber() as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) vbm_gm = element_data["VBM_GM"]
vbm_gm = input["VBM_GM"]
# Test wrong masks definitions (more than one key per dict) # 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(masks={"GM_prob0.2": {}, "Other": {}}, target_data=vbm_gm) get_mask(masks={"GM_prob0.2": {}, "Other": {}}, target_data=vbm_gm)
@ -286,7 +361,8 @@ def test_get_mask_errors() -> None:
ValueError, match=r"parameters to the intersection" ValueError, match=r"parameters to the intersection"
): ):
get_mask( get_mask(
masks=["GM_prob0.2", {"threshold": 1}], target_data=vbm_gm masks=["compute_brain_mask", {"threshold": 1}],
target_data=vbm_gm,
) )
# Test "inherited" masks errors # Test "inherited" masks errors
@ -310,19 +386,20 @@ def test_get_mask_errors() -> None:
masks="inherit", target_data=vbm_gm, extra_input=extra_input masks="inherit", target_data=vbm_gm, extra_input=extra_input
) )
# Block fetch_icbm152_brain_gm_mask space transformation
with pytest.raises(RuntimeError, match="prohibited"):
get_mask(
masks="fetch_icbm152_brain_gm_mask",
target_data=vbm_gm,
extra_input=extra_input,
)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"mask_name,function,params,resample", "mask_name,function,params,resample",
[ [
("compute_brain_mask", compute_brain_mask, {"threshold": 0.2}, False),
("compute_background_mask", compute_background_mask, None, False), ("compute_background_mask", compute_background_mask, None, False),
("compute_epi_mask", compute_epi_mask, None, False), ("compute_epi_mask", compute_epi_mask, None, False),
(
"fetch_icbm152_brain_gm_mask",
fetch_icbm152_brain_gm_mask,
None,
True,
),
], ],
) )
def test_nilearn_compute_masks( def test_nilearn_compute_masks(
@ -345,11 +422,9 @@ def test_nilearn_compute_masks(
Whether to resample the mask to the target data. Whether to resample the mask to the target data.
""" """
reader = DefaultDataReader()
with SPMAuditoryTestingDataGrabber() as dg: with SPMAuditoryTestingDataGrabber() as dg:
input = dg["sub001"] element_data = DefaultDataReader().fit_transform(dg["sub001"])
input = reader.fit_transform(input) bold = element_data["BOLD"]
bold = input["BOLD"]
bold_img = bold["data"] bold_img = bold["data"]
if params is None: if params is None:
@ -378,27 +453,30 @@ def test_nilearn_compute_masks(
def test_get_mask_inherit() -> None: def test_get_mask_inherit() -> None:
"""Test using the inherit mask functionality.""" """Test using the inherit mask functionality."""
reader = DefaultDataReader()
with SPMAuditoryTestingDataGrabber() as dg: with SPMAuditoryTestingDataGrabber() as dg:
input = dg["sub001"] element_data = DefaultDataReader().fit_transform(dg["sub001"])
input = reader.fit_transform(input)
# Compute brain mask using nilearn # Compute brain mask using nilearn
gm_mask = compute_brain_mask(input["BOLD"]["data"], threshold=0.2) gm_mask = compute_brain_mask(element_data["BOLD"], threshold=0.2)
# Get mask using the compute_brain_mask function # Get mask using the compute_brain_mask function
mask1 = get_mask( mask1 = get_mask(
masks={"compute_brain_mask": {"threshold": 0.2}}, masks={"compute_brain_mask": {"threshold": 0.2}},
target_data=input["BOLD"], target_data=element_data["BOLD"],
) )
# Now get the mask using the inherit functionality, passing the # Now get the mask using the inherit functionality, passing the
# computed mask as extra data # computed mask as extra data
extra_input = { extra_input = {
"BOLD_MASK": {"data": gm_mask, "space": input["BOLD"]["space"]} "BOLD_MASK": {
"data": gm_mask,
"space": element_data["BOLD"]["space"],
} }
input["BOLD"]["mask_item"] = "BOLD_MASK" }
element_data["BOLD"]["mask_item"] = "BOLD_MASK"
mask2 = get_mask( mask2 = get_mask(
masks="inherit", target_data=input["BOLD"], extra_input=extra_input masks="inherit",
target_data=element_data["BOLD"],
extra_input=extra_input,
) )
# Both masks should be equal # Both masks should be equal
@ -408,7 +486,6 @@ def test_get_mask_inherit() -> None:
@pytest.mark.parametrize( @pytest.mark.parametrize(
"masks,params", "masks,params",
[ [
(["GM_prob0.2", "GM_prob0.2_cortex"], {}),
(["compute_brain_mask", "compute_background_mask"], {}), (["compute_brain_mask", "compute_background_mask"], {}),
(["compute_brain_mask", "compute_epi_mask"], {}), (["compute_brain_mask", "compute_epi_mask"], {}),
], ],
@ -426,10 +503,8 @@ def test_get_mask_multiple(
Parameters to pass to the intersect_masks function. Parameters to pass to the intersect_masks function.
""" """
reader = DefaultDataReader()
with SPMAuditoryTestingDataGrabber() as dg: with SPMAuditoryTestingDataGrabber() as dg:
input = dg["sub001"] element_data = DefaultDataReader().fit_transform(dg["sub001"])
input = reader.fit_transform(input)
if not isinstance(masks, list): if not isinstance(masks, list):
junifer_masks = [masks] junifer_masks = [masks]
else: else:
@ -438,10 +513,12 @@ def test_get_mask_multiple(
# Convert params to junifer style (one dict per param) # Convert params to junifer style (one dict per param)
junifer_params = [{k: params[k]} for k in params.keys()] junifer_params = [{k: params[k]} for k in params.keys()]
junifer_masks.extend(junifer_params) junifer_masks.extend(junifer_params)
target_img = input["BOLD"]["data"] target_img = element_data["BOLD"]["data"]
resolution = np.min(target_img.header.get_zooms()[:3]) resolution = np.min(target_img.header.get_zooms()[:3])
computed = get_mask(masks=junifer_masks, target_data=input["BOLD"]) computed = get_mask(
masks=junifer_masks, target_data=element_data["BOLD"]
)
masks_names = [ masks_names = [
next(iter(x.keys())) if isinstance(x, dict) else x for x in masks next(iter(x.keys())) if isinstance(x, dict) else x for x in masks
@ -464,6 +541,12 @@ def test_get_mask_multiple(
] ]
for t_func in mask_funcs: for t_func in mask_funcs:
# Bypass for custom mask
if t_func == "compute_brain_mask":
mask_imgs.append(
_available_masks[t_func]["func"](element_data["BOLD"])
)
else:
mask_imgs.append(_available_masks[t_func]["func"](target_img)) mask_imgs.append(_available_masks[t_func]["func"](target_img))
mask_imgs = [ mask_imgs = [
@ -478,21 +561,3 @@ def test_get_mask_multiple(
expected = intersect_masks(mask_imgs, **params) expected = intersect_masks(mask_imgs, **params)
assert_array_equal(computed.get_fdata(), expected.get_fdata()) assert_array_equal(computed.get_fdata(), expected.get_fdata())
def test_get_mask_multiple_incorrect_space() -> None:
"""Test incorrect space error for getting multiple masks."""
reader = DefaultDataReader()
with SPMAuditoryTestingDataGrabber() as dg:
input = dg["sub001"]
input = reader.fit_transform(input)
with pytest.raises(RuntimeError, match="unable to merge."):
get_mask(
masks=[
"GM_prob0.2",
"compute_brain_mask",
"fetch_icbm152_brain_gm_mask",
],
target_data=input["BOLD"],
)

View file

@ -30,7 +30,11 @@ from junifer.data.parcellations import (
register_parcellation, register_parcellation,
) )
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.testing.datagrabbers import OasisVBMTestingDataGrabber from junifer.pipeline.utils import _check_ants
from junifer.testing.datagrabbers import (
OasisVBMTestingDataGrabber,
PartlyCloudyTestingDataGrabber,
)
def test_register_parcellation_built_in_check() -> None: def test_register_parcellation_built_in_check() -> None:
@ -58,7 +62,7 @@ def test_register_parcellation_already_registered() -> None:
name="testparc", name="testparc",
parcellation_path="testparc.nii.gz", parcellation_path="testparc.nii.gz",
parcels_labels=["1", "2", "3"], parcels_labels=["1", "2", "3"],
space="MNI", space="MNI152Lin",
) )
assert ( assert (
load_parcellation("testparc", path_only=True)[2].name load_parcellation("testparc", path_only=True)[2].name
@ -71,13 +75,13 @@ def test_register_parcellation_already_registered() -> None:
name="testparc", name="testparc",
parcellation_path="testparc.nii.gz", parcellation_path="testparc.nii.gz",
parcels_labels=["1", "2", "3"], parcels_labels=["1", "2", "3"],
space="MNI", space="MNI152Lin",
) )
register_parcellation( register_parcellation(
name="testparc", name="testparc",
parcellation_path="testparc2.nii.gz", parcellation_path="testparc2.nii.gz",
parcels_labels=["1", "2", "3"], parcels_labels=["1", "2", "3"],
space="MNI", space="MNI152Lin",
overwrite=True, overwrite=True,
) )
@ -100,14 +104,16 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None:
assert schaefer is not None assert schaefer is not None
# Test wrong number of labels # Test wrong number of labels
register_parcellation("WrongLabels", schaefer_path, labels[:10], "MNI") register_parcellation(
"WrongLabels", schaefer_path, labels[:10], "MNI152Lin"
)
with pytest.raises(ValueError, match=r"has 100 parcels but 10"): with pytest.raises(ValueError, match=r"has 100 parcels but 10"):
load_parcellation("WrongLabels") load_parcellation("WrongLabels")
# Test wrong number of labels # Test wrong number of labels
register_parcellation( register_parcellation(
"WrongLabels2", schaefer_path, [*labels, "wrong"], "MNI" "WrongLabels2", schaefer_path, [*labels, "wrong"], "MNI152Lin"
) )
with pytest.raises(ValueError, match=r"has 100 parcels but 101"): with pytest.raises(ValueError, match=r"has 100 parcels but 101"):
@ -119,7 +125,9 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None:
new_schaefer_img = new_img_like(schaefer, schaefer_data) new_schaefer_img = new_img_like(schaefer, schaefer_data)
nib.save(new_schaefer_img, new_schaefer_path) nib.save(new_schaefer_img, new_schaefer_path)
register_parcellation("WrongValues", new_schaefer_path, labels[:-1], "MNI") register_parcellation(
"WrongValues", new_schaefer_path, labels[:-1], "MNI152Lin"
)
with pytest.raises(ValueError, match=r"the range [0, 99]"): with pytest.raises(ValueError, match=r"the range [0, 99]"):
load_parcellation("WrongValues") load_parcellation("WrongValues")
@ -129,7 +137,9 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None:
new_schaefer_img = new_img_like(schaefer, schaefer_data) new_schaefer_img = new_img_like(schaefer, schaefer_data)
nib.save(new_schaefer_img, new_schaefer_path) nib.save(new_schaefer_img, new_schaefer_path)
register_parcellation("WrongValues2", new_schaefer_path, labels, "MNI") register_parcellation(
"WrongValues2", new_schaefer_path, labels, "MNI152Lin"
)
with pytest.raises(ValueError, match=r"the range [0, 100]"): with pytest.raises(ValueError, match=r"the range [0, 100]"):
load_parcellation("WrongValues2") load_parcellation("WrongValues2")
@ -137,13 +147,25 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None:
@pytest.mark.parametrize( @pytest.mark.parametrize(
"name, parcellation_path, parcels_labels, space, overwrite", "name, parcellation_path, parcels_labels, space, overwrite",
[ [
("testparc_1", "testparc_1.nii.gz", ["1", "2", "3"], "MNI", True), (
("testparc_2", "testparc_2.nii.gz", ["1", "2", "6"], "MNI", True), "testparc_1",
"testparc_1.nii.gz",
["1", "2", "3"],
"MNI152Lin",
True,
),
(
"testparc_2",
"testparc_2.nii.gz",
["1", "2", "6"],
"MNI152Lin",
True,
),
( (
"testparc_3", "testparc_3",
Path("testparc_3.nii.gz"), Path("testparc_3.nii.gz"),
["1", "2", "6"], ["1", "2", "6"],
"MNI", "MNI152Lin",
True, True,
), ),
], ],
@ -1172,28 +1194,26 @@ def test_merge_parcellations_3D_multiple_duplicated_labels() -> None:
def test_get_parcellation_single() -> None: def test_get_parcellation_single() -> None:
"""Test tailored single parcellation fetch.""" """Test tailored single parcellation fetch."""
reader = DefaultDataReader() with PartlyCloudyTestingDataGrabber() as dg:
with OasisVBMTestingDataGrabber() as dg: element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element = dg["sub-01"] bold = element_data["BOLD"]
element_data = reader.fit_transform(element) bold_img = bold["data"]
vbm_gm = element_data["VBM_GM"]
vbm_gm_img = vbm_gm["data"]
# Get tailored parcellation # Get tailored parcellation
tailored_parcellation, tailored_labels = get_parcellation( tailored_parcellation, tailored_labels = get_parcellation(
parcellation=["Schaefer100x7"], parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
target_data=vbm_gm, target_data=bold,
) )
# Check shape and affine with original element data # Check shape and affine with original element data
assert tailored_parcellation.shape == vbm_gm_img.shape assert tailored_parcellation.shape == bold_img.shape[:3]
assert_array_equal(tailored_parcellation.affine, vbm_gm_img.affine) assert_array_equal(tailored_parcellation.affine, bold_img.affine)
# Get raw parcellation # Get raw parcellation
raw_parcellation, raw_labels, _, _ = load_parcellation( raw_parcellation, raw_labels, _, _ = load_parcellation(
"Schaefer100x7", "TianxS1x3TxMNInonlinear2009cAsym",
resolution=1.5, resolution=1.5,
) )
resampled_raw_parcellation = resample_to_img( resampled_raw_parcellation = resample_to_img(
source_img=raw_parcellation, source_img=raw_parcellation,
target_img=vbm_gm_img, target_img=bold_img,
interpolation="nearest", interpolation="nearest",
copy=True, copy=True,
) )
@ -1207,36 +1227,34 @@ def test_get_parcellation_single() -> None:
def test_get_parcellation_multi_same_space() -> None: def test_get_parcellation_multi_same_space() -> None:
"""Test tailored multi parcellation fetch in same space.""" """Test tailored multi parcellation fetch in same space."""
reader = DefaultDataReader() with PartlyCloudyTestingDataGrabber() as dg:
with OasisVBMTestingDataGrabber() as dg: element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element = dg["sub-01"] bold = element_data["BOLD"]
element_data = reader.fit_transform(element) bold_img = bold["data"]
vbm_gm = element_data["VBM_GM"]
vbm_gm_img = vbm_gm["data"]
# Get tailored parcellation # Get tailored parcellation
tailored_parcellation, tailored_labels = get_parcellation( tailored_parcellation, tailored_labels = get_parcellation(
parcellation=[ parcellation=[
"Schaefer100x7", "Shen_2015_268",
"TianxS2x3TxMNI6thgeneration", "TianxS1x3TxMNInonlinear2009cAsym",
], ],
target_data=vbm_gm, target_data=bold,
) )
# Check shape and affine with original element data # Check shape and affine with original element data
assert tailored_parcellation.shape == vbm_gm_img.shape assert tailored_parcellation.shape == bold_img.shape[:3]
assert_array_equal(tailored_parcellation.affine, vbm_gm_img.affine) assert_array_equal(tailored_parcellation.affine, bold_img.affine)
# Get raw parcellations # Get raw parcellations
raw_parcellations = [] raw_parcellations = []
raw_labels = [] raw_labels = []
parcellations_names = [ parcellations_names = [
"Schaefer100x7", "Shen_2015_268",
"TianxS2x3TxMNI6thgeneration", "TianxS1x3TxMNInonlinear2009cAsym",
] ]
for name in parcellations_names: for name in parcellations_names:
img, labels, _, _ = load_parcellation(name=name, resolution=1.5) img, labels, _, _ = load_parcellation(name=name, resolution=1.5)
# Resample raw parcellations # Resample raw parcellations
resampled_img = resample_to_img( resampled_img = resample_to_img(
source_img=img, source_img=img,
target_img=vbm_gm_img, target_img=bold_img,
interpolation="nearest", interpolation="nearest",
copy=True, copy=True,
) )
@ -1256,19 +1274,18 @@ def test_get_parcellation_multi_same_space() -> None:
assert tailored_labels == merged_labels assert tailored_labels == merged_labels
@pytest.mark.skipif(
_check_ants() is False, reason="requires ANTs to be in PATH"
)
def test_get_parcellation_multi_different_space() -> None: def test_get_parcellation_multi_different_space() -> None:
"""Test tailored multi parcellation fetch in different space.""" """Test tailored multi parcellation fetch in different space."""
reader = DefaultDataReader()
with OasisVBMTestingDataGrabber() as dg: with OasisVBMTestingDataGrabber() as dg:
element = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element_data = reader.fit_transform(element)
vbm_gm = element_data["VBM_GM"]
# Get tailored parcellation # Get tailored parcellation
with pytest.raises(RuntimeError, match="unable to merge."):
get_parcellation( get_parcellation(
parcellation=[ parcellation=[
"Schaefer100x7", "Schaefer100x7",
"SUITxSUIT", "TianxS1x3TxMNInonlinear2009cAsym",
], ],
target_data=vbm_gm, target_data=element_data["VBM_GM"],
) )

View file

@ -11,7 +11,10 @@ import pytest
from junifer.data import get_template, get_xfm from junifer.data import get_template, get_xfm
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.testing.datagrabbers import OasisVBMTestingDataGrabber from junifer.testing.datagrabbers import (
OasisVBMTestingDataGrabber,
PartlyCloudyTestingDataGrabber,
)
@pytest.mark.skipif( @pytest.mark.skipif(
@ -33,15 +36,32 @@ def test_get_xfm(tmp_path: Path) -> None:
assert isinstance(xfm_path, Path) assert isinstance(xfm_path, Path)
def test_get_template() -> None: @pytest.mark.parametrize(
"""Test tailored template image fetch.""" "template_type",
with OasisVBMTestingDataGrabber() as dg: [
"T1w",
"brain",
"gm",
"wm",
"csf",
],
)
def test_get_template(template_type: str) -> None:
"""Test tailored template image fetch.
Parameters
----------
template_type : str
The parametrized template type.
"""
with PartlyCloudyTestingDataGrabber() as dg:
element = dg["sub-01"] element = dg["sub-01"]
element_data = DefaultDataReader().fit_transform(element) element_data = DefaultDataReader().fit_transform(element)
vbm_gm = element_data["VBM_GM"] bold = element_data["BOLD"]
# Get tailored parcellation # Get tailored parcellation
tailored_template = get_template( tailored_template = get_template(
space=vbm_gm["space"], target_data=vbm_gm space=bold["space"], target_data=bold, template_type=template_type
) )
assert isinstance(tailored_template, nib.Nifti1Image) assert isinstance(tailored_template, nib.Nifti1Image)
@ -54,7 +74,22 @@ def test_get_template_invalid_space() -> None:
vbm_gm = element_data["VBM_GM"] vbm_gm = element_data["VBM_GM"]
# Get tailored parcellation # Get tailored parcellation
with pytest.raises(ValueError, match="Unknown template space:"): with pytest.raises(ValueError, match="Unknown template space:"):
_ = get_template(space="andromeda", target_data=vbm_gm) get_template(space="andromeda", target_data=vbm_gm)
def test_get_template_invalid_template_type() -> None:
"""Test invalid template type check for template fetch."""
with OasisVBMTestingDataGrabber() as dg:
element = dg["sub-01"]
element_data = DefaultDataReader().fit_transform(element)
vbm_gm = element_data["VBM_GM"]
# Get tailored parcellation
with pytest.raises(ValueError, match="Unknown template type:"):
get_template(
space=vbm_gm["space"],
target_data=vbm_gm,
template_type="xenon",
)
def test_get_template_closest_resolution() -> None: def test_get_template_closest_resolution() -> None:

View file

@ -154,4 +154,7 @@ class DataladAOMICID1000(PatternDataladDataGrabber):
out["T1w"].update({"space": "native"}) out["T1w"].update({"space": "native"})
else: else:
out["T1w"].update({"space": "MNI152NLin2009cAsym"}) out["T1w"].update({"space": "MNI152NLin2009cAsym"})
if out.get("Warp"):
# Add source space information
out["Warp"].update({"src": "MNI152NLin2009cAsym"})
return out return out

View file

@ -204,6 +204,9 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber):
out["T1w"].update({"space": "native"}) out["T1w"].update({"space": "native"})
else: else:
out["T1w"].update({"space": "MNI152NLin2009cAsym"}) out["T1w"].update({"space": "MNI152NLin2009cAsym"})
if out.get("Warp"):
# Add source space information
out["Warp"].update({"src": "MNI152NLin2009cAsym"})
return out return out
def get_elements(self) -> List: def get_elements(self) -> List:

View file

@ -204,4 +204,7 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber):
out["T1w"].update({"space": "native"}) out["T1w"].update({"space": "native"})
else: else:
out["T1w"].update({"space": "MNI152NLin2009cAsym"}) out["T1w"].update({"space": "MNI152NLin2009cAsym"})
if out.get("Warp"):
# Add source space information
out["Warp"].update({"src": "MNI152NLin2009cAsym"})
return out return out

View file

@ -271,6 +271,9 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
out["T1w"].update({"space": "native"}) out["T1w"].update({"space": "native"})
else: else:
out["T1w"].update({"space": "MNI152NLin2009cAsym"}) out["T1w"].update({"space": "MNI152NLin2009cAsym"})
if out.get("Warp"):
# Add source space information
out["Warp"].update({"src": "MNI152NLin2009cAsym"})
return out return out
def get_elements(self) -> List: def get_elements(self) -> List:

View file

@ -159,6 +159,9 @@ class HCP1200(PatternDataGrabber):
# Add space for T1w data type # Add space for T1w data type
if "T1w" in out: if "T1w" in out:
out["T1w"].update({"space": "native"}) out["T1w"].update({"space": "native"})
# Add source space for Warp data type
if "Warp" in out:
out["Warp"].update({"src": "MNI152NLin6Asym"})
return out return out
def get_elements(self) -> List: def get_elements(self) -> List:

View file

@ -18,7 +18,7 @@ from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
PARCELLATION = "Schaefer100x7" PARCELLATION = "TianxS1x3TxMNInonlinear2009cAsym"
def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None: def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@ -59,7 +59,7 @@ def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
output_bold_data = output_bold["data"] output_bold_data = output_bold["data"]
# Assert BOLD output data dimension # Assert BOLD output data dimension
assert output_bold_data.ndim == 2 assert output_bold_data.ndim == 2
assert output_bold_data.shape == (1, 100) assert output_bold_data.shape == (1, 16)
# Reset log capture # Reset log capture
caplog.clear() caplog.clear()
@ -123,4 +123,4 @@ def test_ALFFParcels_comparison(tmp_path: Path, fractional: bool) -> None:
junifer_output_bold["data"][0], junifer_output_bold["data"][0],
afni_output_bold["data"][0], afni_output_bold["data"][0],
) )
assert r > 0.99 assert r > 0.97

View file

@ -82,7 +82,7 @@ def test_ALFFSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
"fractional", [True, False], ids=["fractional", "non-fractional"] "fractional", [True, False], ids=["fractional", "non-fractional"]
) )
def test_ALFFSpheres_comparison(tmp_path: Path, fractional: bool) -> None: def test_ALFFSpheres_comparison(tmp_path: Path, fractional: bool) -> None:
"""Test ALFFSpheres using afni. """Test ALFFSpheres implementation comparison.
Parameters Parameters
---------- ----------

View file

@ -2,14 +2,17 @@
# Authors: Leonard Sasse <l.sasse@fz-juelich.de> # Authors: Leonard Sasse <l.sasse@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
import pytest import pytest
from nilearn import image
from junifer.datareader import DefaultDataReader
from junifer.markers.functional_connectivity import CrossParcellationFC from junifer.markers.functional_connectivity import CrossParcellationFC
from junifer.pipeline import WorkDirManager
from junifer.pipeline.utils import _check_ants
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
@ -18,32 +21,53 @@ parcellation_one = "Schaefer100x17"
parcellation_two = "Schaefer200x17" parcellation_two = "Schaefer200x17"
def test_compute() -> None: def test_init() -> None:
"""Test CrossParcellationFC compute().""" """Test CrossParcellationFC init()."""
with pytest.raises(ValueError, match="must be different"):
CrossParcellationFC(
parcellation_one="a",
parcellation_two="a",
correlation_method="pearson",
)
def test_get_output_type() -> None:
"""Test CrossParcellationFC get_output_type()."""
crossparcellation = CrossParcellationFC(
parcellation_one=parcellation_one, parcellation_two=parcellation_two
)
assert "matrix" == crossparcellation.get_output_type("BOLD")
@pytest.mark.skipif(
_check_ants() is False, reason="requires ANTs to be in PATH"
)
def test_compute(tmp_path: Path) -> None:
"""Test CrossParcellationFC compute().
Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
with SPMAuditoryTestingDataGrabber() as dg: with SPMAuditoryTestingDataGrabber() as dg:
out = dg["sub001"] element_data = DefaultDataReader().fit_transform(dg["sub001"])
niimg = image.load_img(str(out["BOLD"]["path"].absolute())) WorkDirManager().workdir = tmp_path
input_dict = {
"BOLD": {
"data": niimg,
"path": out["BOLD"]["path"],
"meta": {"element": "sub001"},
"space": "MNI",
}
}
crossparcellation = CrossParcellationFC( crossparcellation = CrossParcellationFC(
parcellation_one=parcellation_one, parcellation_one=parcellation_one,
parcellation_two=parcellation_two, parcellation_two=parcellation_two,
correlation_method="spearman", correlation_method="spearman",
) )
out = crossparcellation.compute(input_dict["BOLD"]) out = crossparcellation.compute(element_data["BOLD"])
assert out["data"].shape == (200, 100) assert out["data"].shape == (200, 100)
assert len(out["col_names"]) == 100 assert len(out["col_names"]) == 100
assert len(out["row_names"]) == 200 assert len(out["row_names"]) == 200
@pytest.mark.skipif(
_check_ants() is False, reason="requires ANTs to be in PATH"
)
def test_store(tmp_path: Path) -> None: def test_store(tmp_path: Path) -> None:
"""Test CrossParcellationFC store(). """Test CrossParcellationFC store().
@ -53,43 +77,20 @@ def test_store(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
with SPMAuditoryTestingDataGrabber() as dg: with SPMAuditoryTestingDataGrabber() as dg:
input_dict = dg["sub001"] element_data = DefaultDataReader().fit_transform(dg["sub001"])
niimg = image.load_img(str(input_dict["BOLD"]["path"].absolute())) WorkDirManager().workdir = tmp_path
input_dict["BOLD"]["data"] = niimg
crossparcellation = CrossParcellationFC( crossparcellation = CrossParcellationFC(
parcellation_one=parcellation_one, parcellation_one=parcellation_one,
parcellation_two=parcellation_two, parcellation_two=parcellation_two,
correlation_method="spearman", correlation_method="spearman",
) )
uri = tmp_path / "test_crossparcellation.sqlite" storage = SQLiteFeatureStorage(
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") uri=tmp_path / "test_crossparcellation.sqlite", upsert="ignore"
crossparcellation.fit_transform(input_dict, storage=storage) )
# Fit transform marker on data with storage
crossparcellation.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_CrossParcellationFC" for x in features.values() x["name"] == "BOLD_CrossParcellationFC" for x in features.values()
) )
def test_get_output_type() -> None:
"""Test CrossParcellationFC get_output_type()."""
crossparcellation = CrossParcellationFC(
parcellation_one=parcellation_one, parcellation_two=parcellation_two
)
input_ = "BOLD"
output = crossparcellation.get_output_type(input_)
assert output == "matrix"
def test_init_() -> None:
"""Test CrossParcellationFC init()."""
with pytest.raises(ValueError, match="must be different"):
CrossParcellationFC(
parcellation_one="a",
parcellation_two="a",
correlation_method="pearson",
)

View file

@ -6,10 +6,10 @@
from pathlib import Path from pathlib import Path
from nilearn import datasets, image from junifer.datareader import DefaultDataReader
from junifer.markers.functional_connectivity import EdgeCentricFCParcels from junifer.markers.functional_connectivity import EdgeCentricFCParcels
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
def test_EdgeCentricFCParcels(tmp_path: Path) -> None: def test_EdgeCentricFCParcels(tmp_path: Path) -> None:
@ -21,41 +21,34 @@ def test_EdgeCentricFCParcels(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset with PartlyCloudyTestingDataGrabber() as dg:
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") element_data = DefaultDataReader().fit_transform(dg["sub-01"])
fmri_img = image.concat_imgs(ni_data.func) # type: ignore marker = EdgeCentricFCParcels(
# Check empirical correlation method parameters
efc = EdgeCentricFCParcels(
parcellation="TianxS1x3TxMNInonlinear2009cAsym", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
cor_method_params={"empirical": True}, cor_method_params={"empirical": True},
) )
all_out = efc.fit_transform( # Check correct output
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} assert marker.get_output_type("BOLD") == "matrix"
)
out = all_out["BOLD"] # Fit-transform the data
edge_fc = marker.fit_transform(element_data)
edge_fc_bold = edge_fc["BOLD"]
# for 16 ROIs we should get (16 * (16 -1) / 2) edges in the ETS # For 16 ROIs we should get (16 * (16 -1) / 2) edges in the ETS
n_edges = int(16 * (16 - 1) / 2) n_edges = int(16 * (16 - 1) / 2)
assert "data" in out assert "data" in edge_fc_bold
assert "row_names" in out assert "row_names" in edge_fc_bold
assert "col_names" in out assert "col_names" in edge_fc_bold
assert out["data"].shape[0] == n_edges assert edge_fc_bold["data"].shape[0] == n_edges
assert out["data"].shape[1] == n_edges assert edge_fc_bold["data"].shape[1] == n_edges
assert len(set(out["row_names"])) == n_edges assert len(set(edge_fc_bold["row_names"])) == n_edges
assert len(set(out["col_names"])) == n_edges assert len(set(edge_fc_bold["col_names"])) == n_edges
# check correct output
assert efc.get_output_type("BOLD") == "matrix"
uri = tmp_path / "test_fc_parcellation.sqlite"
# Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}}
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
all_out = efc.fit_transform(input, storage=storage)
# Store
storage = SQLiteFeatureStorage(
uri=tmp_path / "test_edge_fc_parcels.sqlite", upsert="ignore"
)
marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_EdgeCentricFCParcels" for x in features.values() x["name"] == "BOLD_EdgeCentricFCParcels" for x in features.values()

View file

@ -6,10 +6,10 @@
from pathlib import Path from pathlib import Path
from nilearn import datasets, image from junifer.datareader import DefaultDataReader
from junifer.markers.functional_connectivity import EdgeCentricFCSpheres from junifer.markers.functional_connectivity import EdgeCentricFCSpheres
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
def test_EdgeCentricFCSpheres(tmp_path: Path) -> None: def test_EdgeCentricFCSpheres(tmp_path: Path) -> None:
@ -21,56 +21,40 @@ def test_EdgeCentricFCSpheres(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset with SPMAuditoryTestingDataGrabber() as dg:
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") element_data = DefaultDataReader().fit_transform(dg["sub001"])
fmri_img = image.concat_imgs(ni_data.func) # type: ignore marker = EdgeCentricFCSpheres(
efc = EdgeCentricFCSpheres(
coords="DMNBuckner", radius=5.0, cor_method="correlation" coords="DMNBuckner", radius=5.0, cor_method="correlation"
) )
all_out = efc.fit_transform( # Check correct output
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} assert marker.get_output_type("BOLD") == "matrix"
)
out = all_out["BOLD"] # Fit-transform the data
edge_fc = marker.fit_transform(element_data)
edge_fc_bold = edge_fc["BOLD"]
# There are six DMNBuckner coordinates, so # There are six DMNBuckner coordinates, so
# for 6 ROIs we should get (6 * (6 -1) / 2) edges in the ETS # for 6 ROIs we should get (6 * (6 -1) / 2) edges in the ETS
n_edges = int(6 * (6 - 1) / 2) n_edges = int(6 * (6 - 1) / 2)
assert "data" in out assert "data" in edge_fc_bold
assert "row_names" in out assert "row_names" in edge_fc_bold
assert "col_names" in out assert "col_names" in edge_fc_bold
assert out["data"].shape[0] == n_edges assert edge_fc_bold["data"].shape == (n_edges, n_edges)
assert out["data"].shape[1] == n_edges assert len(set(edge_fc_bold["row_names"])) == n_edges
assert len(set(out["row_names"])) == n_edges assert len(set(edge_fc_bold["col_names"])) == n_edges
assert len(set(out["col_names"])) == n_edges
# check correct output
assert efc.get_output_type("BOLD") == "matrix"
# Check empirical correlation method parameters # Check empirical correlation method parameters
efc = EdgeCentricFCSpheres( marker = EdgeCentricFCSpheres(
coords="DMNBuckner", coords="DMNBuckner",
radius=5.0, radius=5.0,
cor_method="correlation", cor_method="correlation",
cor_method_params={"empirical": True}, cor_method_params={"empirical": True},
) )
# Store
meta = { storage = SQLiteFeatureStorage(
"element": {"subject": "sub001"}, uri=tmp_path / "test_edge_fc_spheres.sqlite", upsert="ignore"
"dependencies": {"nilearn"},
}
all_out = efc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
) )
marker.fit_transform(input=element_data, storage=storage)
uri = tmp_path / "test_fc_parcellation.sqlite"
# Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}}
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
all_out = efc.fit_transform(input, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_EdgeCentricFCSpheres" for x in features.values() x["name"] == "BOLD_EdgeCentricFCSpheres" for x in features.values()

View file

@ -2,20 +2,22 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
from nilearn import datasets, image
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from nilearn.maskers import NiftiLabelsMasker from nilearn.maskers import NiftiLabelsMasker
from numpy.testing import assert_array_almost_equal, assert_array_equal from numpy.testing import assert_array_almost_equal
from junifer.data import get_parcellation
from junifer.datareader import DefaultDataReader
from junifer.markers.functional_connectivity import ( from junifer.markers.functional_connectivity import (
FunctionalConnectivityParcels, FunctionalConnectivityParcels,
) )
from junifer.markers.parcel_aggregation import ParcelAggregation
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
def test_FunctionalConnectivityParcels(tmp_path: Path) -> None: def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
@ -27,72 +29,58 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset with PartlyCloudyTestingDataGrabber() as dg:
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") element_data = DefaultDataReader().fit_transform(dg["sub-01"])
fmri_img = image.concat_imgs(ni_data.func) # type: ignore marker = FunctionalConnectivityParcels(
parcellation="TianxS1x3TxMNInonlinear2009cAsym"
fc = FunctionalConnectivityParcels(parcellation="Schaefer100x7")
all_out = fc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
) )
# Check correct output
assert marker.get_output_type("BOLD") == "matrix"
out = all_out["BOLD"] # Fit-transform the data
fc = marker.fit_transform(element_data)
fc_bold = fc["BOLD"]
assert "data" in out assert "data" in fc_bold
assert "row_names" in out assert "row_names" in fc_bold
assert "col_names" in out assert "col_names" in fc_bold
assert out["data"].shape[0] == 100 assert fc_bold["data"].shape == (16, 16)
assert out["data"].shape[1] == 100 assert len(set(fc_bold["row_names"])) == 16
assert len(set(out["row_names"])) == 100 assert len(set(fc_bold["col_names"])) == 16
assert len(set(out["col_names"])) == 100
# get the timeseries using pa # Compare with nilearn
pa = ParcelAggregation( # Load testing parcellation for the target data
parcellation="Schaefer100x7", method="mean", on="BOLD" testing_parcellation, _ = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
target_data=element_data["BOLD"],
) )
meta = { # Extract timeseries
"element": {"subject": "sub001"}, nifti_labels_masker = NiftiLabelsMasker(
"dependencies": {"nilearn"}, labels_img=testing_parcellation, standardize=False
}
ts = pa.compute({"data": fmri_img, "meta": meta, "space": "MNI"})
# compare with nilearn
# Get the testing parcellation (for nilearn)
parcellation = datasets.fetch_atlas_schaefer_2018(
n_rois=100, yeo_networks=7, resolution_mm=2
) )
masker = NiftiLabelsMasker( extracted_timeseries = nifti_labels_masker.fit_transform(
labels_img=parcellation["maps"], standardize=False element_data["BOLD"]["data"]
) )
ts_ni = masker.fit_transform(fmri_img) # Compute the connectivity measure
connectivity_measure = ConnectivityMeasure(
# check the TS are almost equal kind="covariance"
assert_array_equal(ts_ni, ts["data"]) ).fit_transform([extracted_timeseries])[0]
# Check that FC are almost equal # Check that FC are almost equal
cm = ConnectivityMeasure(kind="covariance") assert_array_almost_equal(
out_ni = cm.fit_transform([ts_ni])[0] connectivity_measure, fc_bold["data"], decimal=3
assert_array_almost_equal(out_ni, out["data"], decimal=3) )
# check correct output
assert fc.get_output_type("BOLD") == "matrix"
# Check empirical correlation method parameters # Check empirical correlation method parameters
fc = FunctionalConnectivityParcels( marker = FunctionalConnectivityParcels(
parcellation="Schaefer100x7", cor_method_params={"empirical": True} parcellation="TianxS1x3TxMNInonlinear2009cAsym",
cor_method_params={"empirical": True},
) )
# Store
all_out = fc.fit_transform( storage = SQLiteFeatureStorage(
{"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}} uri=tmp_path / "test_fc_parcels.sqlite", upsert="ignore"
) )
marker.fit_transform(input=element_data, storage=storage)
uri = tmp_path / "test_fc_parcellation.sqlite"
# Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}}
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
all_out = fc.fit_transform(input, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_FunctionalConnectivityParcels" x["name"] == "BOLD_FunctionalConnectivityParcels"

View file

@ -3,21 +3,24 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# Federico Raimondo <f.raimondo@fz-juelich.de> # Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
import pytest import pytest
from nilearn import datasets, image
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from nilearn.maskers import NiftiSpheresMasker
from numpy.testing import assert_array_almost_equal from numpy.testing import assert_array_almost_equal
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
from junifer.data import get_coordinates
from junifer.datareader import DefaultDataReader
from junifer.markers.functional_connectivity import ( from junifer.markers.functional_connectivity import (
FunctionalConnectivitySpheres, FunctionalConnectivitySpheres,
) )
from junifer.markers.sphere_aggregation import SphereAggregation
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None: def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
@ -29,51 +32,52 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset with SPMAuditoryTestingDataGrabber() as dg:
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") element_data = DefaultDataReader().fit_transform(dg["sub001"])
fmri_img = image.concat_imgs(ni_data.func) # type: ignore marker = FunctionalConnectivitySpheres(
fc = FunctionalConnectivitySpheres(
coords="DMNBuckner", radius=5.0, cor_method="correlation" coords="DMNBuckner", radius=5.0, cor_method="correlation"
) )
all_out = fc.fit_transform( # Check correct output
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} assert marker.get_output_type("BOLD") == "matrix"
# Fit-transform the data
fc = marker.fit_transform(element_data)
fc_bold = fc["BOLD"]
assert "data" in fc_bold
assert "row_names" in fc_bold
assert "col_names" in fc_bold
assert fc_bold["data"].shape == (6, 6)
assert len(set(fc_bold["row_names"])) == 6
assert len(set(fc_bold["col_names"])) == 6
# Compare with nilearn
# Load testing coordinates for the target data
testing_coords, _ = get_coordinates(
coords="DMNBuckner", target_data=element_data["BOLD"]
)
# Extract timeseries
nifti_spheres_masker = NiftiSpheresMasker(
seeds=testing_coords, radius=5.0
)
extracted_timeseries = nifti_spheres_masker.fit_transform(
element_data["BOLD"]["data"]
)
# Compute the connectivity measure
connectivity_measure = ConnectivityMeasure(
kind="correlation"
).fit_transform([extracted_timeseries])[0]
# Check that FC are almost equal
assert_array_almost_equal(
connectivity_measure, fc_bold["data"], decimal=3
) )
out = all_out["BOLD"] # Store
storage = SQLiteFeatureStorage(
assert "data" in out uri=tmp_path / "test_fc_spheres.sqlite", upsert="ignore"
assert "row_names" in out
assert "col_names" in out
assert out["data"].shape[0] == 6
assert out["data"].shape[1] == 6
assert len(set(out["row_names"])) == 6
assert len(set(out["col_names"])) == 6
# get the timeseries using sa
sa = SphereAggregation(
coords="DMNBuckner", radius=5.0, method="mean", on="BOLD"
) )
ts = sa.compute({"data": fmri_img, "meta": {}, "space": "MNI"}) marker.fit_transform(input=element_data, storage=storage)
# Check that FC are almost equal when using nileran
cm = ConnectivityMeasure(kind="correlation")
out_ni = cm.fit_transform([ts["data"]])[0]
assert_array_almost_equal(out_ni, out["data"], decimal=3)
# check correct output
assert fc.get_output_type("BOLD") == "matrix"
uri = tmp_path / "test_fc_parcel.sqlite"
# Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {
"element": {"subject": "test"},
"dependencies": {"numpy", "nilearn"},
}
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
all_out = fc.fit_transform(input, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_FunctionalConnectivitySpheres" x["name"] == "BOLD_FunctionalConnectivitySpheres"
@ -90,43 +94,49 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
with SPMAuditoryTestingDataGrabber() as dg:
# get a dataset element_data = DefaultDataReader().fit_transform(dg["sub001"])
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") marker = FunctionalConnectivitySpheres(
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
fc = FunctionalConnectivitySpheres(
coords="DMNBuckner", coords="DMNBuckner",
radius=5.0, radius=5.0,
cor_method="correlation", cor_method="correlation",
cor_method_params={"empirical": True}, cor_method_params={"empirical": True},
) )
all_out = fc.fit_transform( # Check correct output
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} assert marker.get_output_type("BOLD") == "matrix"
# Fit-transform the data
fc = marker.fit_transform(element_data)
fc_bold = fc["BOLD"]
assert "data" in fc_bold
assert "row_names" in fc_bold
assert "col_names" in fc_bold
assert fc_bold["data"].shape == (6, 6)
assert len(set(fc_bold["row_names"])) == 6
assert len(set(fc_bold["col_names"])) == 6
# Compare with nilearn
# Load testing coordinates for the target data
testing_coords, _ = get_coordinates(
coords="DMNBuckner", target_data=element_data["BOLD"]
) )
# Extract timeseries
out = all_out["BOLD"] nifti_spheres_masker = NiftiSpheresMasker(
seeds=testing_coords, radius=5.0
assert "data" in out
assert "row_names" in out
assert "col_names" in out
assert out["data"].shape[0] == 6
assert out["data"].shape[1] == 6
assert len(set(out["row_names"])) == 6
assert len(set(out["col_names"])) == 6
# get the timeseries using sa
sa = SphereAggregation(
coords="DMNBuckner", radius=5.0, method="mean", on="BOLD"
) )
ts = sa.compute({"data": fmri_img, "space": "MNI"}) extracted_timeseries = nifti_spheres_masker.fit_transform(
element_data["BOLD"]["data"]
# Check that FC are almost equal when using nileran )
cm = ConnectivityMeasure( # Compute the connectivity measure
connectivity_measure = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
).fit_transform([extracted_timeseries])[0]
# Check that FC are almost equal
assert_array_almost_equal(
connectivity_measure, fc_bold["data"], decimal=3
) )
out_ni = cm.fit_transform([ts["data"]])[0]
assert_array_almost_equal(out_ni, out["data"], decimal=3)
def test_FunctionalConnectivitySpheres_error() -> None: def test_FunctionalConnectivitySpheres_error() -> None:

View file

@ -12,12 +12,12 @@ import scipy as sp
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers import ReHoParcels from junifer.markers import ReHoParcels
from junifer.pipeline import WorkDirManager from junifer.pipeline import WorkDirManager
from junifer.pipeline.utils import _check_afni from junifer.pipeline.utils import _check_afni, _check_ants
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber from junifer.testing.datagrabbers import (
PartlyCloudyTestingDataGrabber,
SPMAuditoryTestingDataGrabber,
PARCELLATION = "Schaefer100x7" )
def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None: def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@ -32,13 +32,16 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
""" """
with caplog.at_level(logging.DEBUG): with caplog.at_level(logging.DEBUG):
with SPMAuditoryTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub001"]) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Update workdir to current test's tmp_path # Update workdir to current test's tmp_path
WorkDirManager().workdir = tmp_path WorkDirManager().workdir = tmp_path
# Initialize marker # Initialize marker
marker = ReHoParcels(parcellation=PARCELLATION, using="junifer") marker = ReHoParcels(
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
using="junifer",
)
# Fit transform marker on data # Fit transform marker on data
output = marker.fit_transform(element_data) output = marker.fit_transform(element_data)
@ -72,6 +75,9 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
assert "Creating cache" not in caplog.text assert "Creating cache" not in caplog.text
@pytest.mark.skipif(
_check_ants() is False, reason="requires ANTs to be in PATH"
)
@pytest.mark.skipif( @pytest.mark.skipif(
_check_afni() is False, reason="requires AFNI to be in PATH" _check_afni() is False, reason="requires AFNI to be in PATH"
) )
@ -91,7 +97,7 @@ def test_ReHoParcels_comparison(tmp_path: Path) -> None:
# Initialize marker # Initialize marker
junifer_marker = ReHoParcels( junifer_marker = ReHoParcels(
parcellation=PARCELLATION, using="junifer" parcellation="Schaefer100x7", using="junifer"
) )
# Fit transform marker on data # Fit transform marker on data
junifer_output = junifer_marker.fit_transform(element_data) junifer_output = junifer_marker.fit_transform(element_data)
@ -99,7 +105,7 @@ def test_ReHoParcels_comparison(tmp_path: Path) -> None:
junifer_output_bold = junifer_output["BOLD"] junifer_output_bold = junifer_output["BOLD"]
# Initialize marker # Initialize marker
afni_marker = ReHoParcels(parcellation=PARCELLATION, using="afni") afni_marker = ReHoParcels(parcellation="Schaefer100x7", using="afni")
# Fit transform marker on data # Fit transform marker on data
afni_output = afni_marker.fit_transform(element_data) afni_output = afni_marker.fit_transform(element_data)
# Get BOLD output # Get BOLD output
@ -110,4 +116,4 @@ def test_ReHoParcels_comparison(tmp_path: Path) -> None:
junifer_output_bold["data"].flatten(), junifer_output_bold["data"].flatten(),
afni_output_bold["data"].flatten(), afni_output_bold["data"].flatten(),
) )
assert r >= 0.3 # this is very bad, but they differ... assert r >= 0.2 # this is very bad, but they differ...

View file

@ -1,18 +1,39 @@
"""Provide tests for temporal signal-to-noise ratio using parcellation.""" """Provide tests for temporal signal-to-noise ratio using parcellation."""
# Authors: Leonard Sasse <l.sasse@fz-juelich.de> # Authors: Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
from nilearn import datasets, image from junifer.datareader import DefaultDataReader
from junifer.markers.temporal_snr import TemporalSNRParcels from junifer.markers.temporal_snr import TemporalSNRParcels
from junifer.storage import SQLiteFeatureStorage from junifer.storage import HDF5FeatureStorage
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
def test_TemporalSNRParcels(tmp_path: Path) -> None: def test_TemporalSNRParcels_computation() -> None:
"""Test TemporalSNRParcels. """Test TemporalSNRParcels fit-transform."""
with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
marker = TemporalSNRParcels(
parcellation="TianxS1x3TxMNInonlinear2009cAsym"
)
# Check correct output
assert marker.get_output_type("BOLD") == "vector"
# Fit-transform the data
tsnr_parcels = marker.fit_transform(element_data)
tsnr_parcels_bold = tsnr_parcels["BOLD"]
assert "data" in tsnr_parcels_bold
assert "col_names" in tsnr_parcels_bold
assert tsnr_parcels_bold["data"].shape == (1, 16)
assert len(set(tsnr_parcels_bold["col_names"])) == 16
def test_TemporalSNRParcels_storage(tmp_path: Path) -> None:
"""Test TemporalSNRParcels store.
Parameters Parameters
---------- ----------
@ -20,34 +41,14 @@ def test_TemporalSNRParcels(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset with PartlyCloudyTestingDataGrabber() as dg:
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") element_data = DefaultDataReader().fit_transform(dg["sub-01"])
fmri_img = image.concat_imgs(ni_data.func) # type: ignore marker = TemporalSNRParcels(
parcellation="TianxS1x3TxMNInonlinear2009cAsym"
tsnr_parcels = TemporalSNRParcels(parcellation="Schaefer100x7")
all_out = tsnr_parcels.fit_transform(
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
) )
# Store
out = all_out["BOLD"] storage = HDF5FeatureStorage(tmp_path / "test_tsnr_parcels.hdf5")
marker.fit_transform(input=element_data, storage=storage)
assert "data" in out
assert "col_names" in out
assert out["data"].shape[0] == 1
assert out["data"].shape[1] == 100
assert len(set(out["col_names"])) == 100
# check correct output
assert tsnr_parcels.get_output_type("BOLD") == "vector"
uri = tmp_path / "test_tsnr_parcellation.sqlite"
# Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}}
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
all_out = tsnr_parcels.fit_transform(input, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_TemporalSNRParcels" for x in features.values() x["name"] == "BOLD_TemporalSNRParcels" for x in features.values()

View file

@ -1,19 +1,39 @@
"""Provide tests for temporal signal-to-noise spheres.""" """Provide tests for temporal signal-to-noise spheres."""
# Authors: Leonard Sasse <l.sasse@fz-juelich.de> # Authors: Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
import pytest import pytest
from nilearn import datasets, image
from junifer.datareader import DefaultDataReader
from junifer.markers.temporal_snr import TemporalSNRSpheres from junifer.markers.temporal_snr import TemporalSNRSpheres
from junifer.storage import SQLiteFeatureStorage from junifer.storage import HDF5FeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
def test_TemporalSNRSpheres(tmp_path: Path) -> None: def test_TemporalSNRSpheres_computation() -> None:
"""Test TemporalSNRSpheres. """Test TemporalSNRSpheres fit-transform."""
with SPMAuditoryTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub001"])
marker = TemporalSNRSpheres(coords="DMNBuckner", radius=5.0)
# Check correct output
assert marker.get_output_type("BOLD") == "vector"
# Fit-transform the data
tsnr_spheres = marker.fit_transform(element_data)
tsnr_spheres_bold = tsnr_spheres["BOLD"]
assert "data" in tsnr_spheres_bold
assert "col_names" in tsnr_spheres_bold
assert tsnr_spheres_bold["data"].shape == (1, 6)
assert len(set(tsnr_spheres_bold["col_names"])) == 6
def test_TemporalSNRSpheres_storage(tmp_path: Path) -> None:
"""Test TemporalSNRSpheres store.
Parameters Parameters
---------- ----------
@ -21,36 +41,12 @@ def test_TemporalSNRSpheres(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset with SPMAuditoryTestingDataGrabber() as dg:
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") element_data = DefaultDataReader().fit_transform(dg["sub001"])
fmri_img = image.concat_imgs(ni_data.func) # type: ignore marker = TemporalSNRSpheres(coords="DMNBuckner", radius=5.0)
# Store
tsnr_spheres = TemporalSNRSpheres(coords="DMNBuckner", radius=5.0) storage = HDF5FeatureStorage(tmp_path / "test_tsnr_spheres.hdf5")
all_out = tsnr_spheres.fit_transform( marker.fit_transform(input=element_data, storage=storage)
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
)
out = all_out["BOLD"]
assert "data" in out
assert "col_names" in out
assert out["data"].shape[0] == 1
assert out["data"].shape[1] == 6
assert len(set(out["col_names"])) == 6
# check correct output
assert tsnr_spheres.get_output_type("BOLD") == "vector"
uri = tmp_path / "test_tsnr_coords.sqlite"
# Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {
"element": {"subject": "test"},
"dependencies": {"numpy", "nilearn"},
}
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
all_out = tsnr_spheres.fit_transform(input, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_TemporalSNRSpheres" for x in features.values() x["name"] == "BOLD_TemporalSNRSpheres" for x in features.values()

View file

@ -19,7 +19,6 @@ from junifer.pipeline import PipelineStepMixin
from junifer.preprocess import fMRIPrepConfoundRemover from junifer.preprocess import fMRIPrepConfoundRemover
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import ( from junifer.testing.datagrabbers import (
OasisVBMTestingDataGrabber,
PartlyCloudyTestingDataGrabber, PartlyCloudyTestingDataGrabber,
) )
@ -46,20 +45,20 @@ def test_marker_collection() -> None:
"""Test MarkerCollection.""" """Test MarkerCollection."""
markers = [ markers = [
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS2x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
), ),
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS2x3TxMNInonlinear2009cAsym",
method="std", method="std",
name="gmd_schaefer100x7_std", name="tian_std",
), ),
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS2x3TxMNInonlinear2009cAsym",
method="trim_mean", method="trim_mean",
method_params={"proportiontocut": 0.1}, method_params={"proportiontocut": 0.1},
name="gmd_schaefer100x7_trim_mean90", name="tian_trim_mean90",
), ),
] ]
mc = MarkerCollection(markers=markers) # type: ignore mc = MarkerCollection(markers=markers) # type: ignore
@ -69,7 +68,7 @@ def test_marker_collection() -> None:
assert isinstance(mc._datareader, DefaultDataReader) assert isinstance(mc._datareader, DefaultDataReader)
# Create testing datagrabber # Create testing datagrabber
dg = OasisVBMTestingDataGrabber() dg = PartlyCloudyTestingDataGrabber()
mc.validate(dg) mc.validate(dg)
with dg: with dg:
@ -78,17 +77,17 @@ def test_marker_collection() -> None:
assert out is not None assert out is not None
assert isinstance(out, dict) assert isinstance(out, dict)
assert len(out) == 3 assert len(out) == 3
assert "gmd_schaefer100x7_mean" in out assert "tian_mean" in out
assert "gmd_schaefer100x7_std" in out assert "tian_std" in out
assert "gmd_schaefer100x7_trim_mean90" in out assert "tian_trim_mean90" in out
for t_marker in markers: for t_marker in markers:
t_name = t_marker.name t_name = t_marker.name
assert "VBM_GM" in out[t_name] assert "BOLD" in out[t_name]
t_vbm = out[t_name]["VBM_GM"] t_bold = out[t_name]["BOLD"]
assert "data" in t_vbm assert "data" in t_bold
assert "col_names" in t_vbm assert "col_names" in t_bold
assert "meta" in t_vbm assert "meta" in t_bold
# Test preprocessing # Test preprocessing
class BypassPreprocessing(PipelineStepMixin): class BypassPreprocessing(PipelineStepMixin):
@ -108,7 +107,7 @@ def test_marker_collection() -> None:
for t_marker in markers: for t_marker in markers:
t_name = t_marker.name t_name = t_marker.name
assert_array_equal( assert_array_equal(
out[t_name]["VBM_GM"]["data"], out2[t_name]["VBM_GM"]["data"] out[t_name]["BOLD"]["data"], out2[t_name]["BOLD"]["data"]
) )
@ -151,27 +150,28 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
""" """
markers = [ markers = [
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS2x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
), ),
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS2x3TxMNInonlinear2009cAsym",
method="std", method="std",
name="gmd_schaefer100x7_std", name="tian_std",
), ),
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS2x3TxMNInonlinear2009cAsym",
method="trim_mean", method="trim_mean",
method_params={"proportiontocut": 0.1}, method_params={"proportiontocut": 0.1},
name="gmd_schaefer100x7_trim_mean90", name="tian_trim_mean90",
), ),
] ]
# Test storage # Setup datagrabber
dg = OasisVBMTestingDataGrabber() dg = PartlyCloudyTestingDataGrabber()
# Setup storage
uri = tmp_path / "test_marker_collection_storage.sqlite" storage = SQLiteFeatureStorage(
storage = SQLiteFeatureStorage(uri=uri) tmp_path / "test_marker_collection_storage.sqlite"
)
mc = MarkerCollection( mc = MarkerCollection(
markers=markers, # type: ignore markers=markers, # type: ignore
storage=storage, storage=storage,
@ -197,23 +197,24 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
features = storage.list_features() features = storage.list_features()
assert len(features) == 3 assert len(features) == 3
feature_md5 = next(iter(features.keys())) feature_md5 = next(iter(features.keys()))
t_feature = storage.read_df(feature_md5=feature_md5) t_feature = storage.read_df(feature_md5=feature_md5)
fname = "gmd_schaefer100x7_mean" fname = "tian_mean"
t_data = out[fname]["VBM_GM"]["data"] # type: ignore t_data = out[fname]["BOLD"]["data"] # type: ignore
cols = out[fname]["VBM_GM"]["col_names"] # type: ignore cols = out[fname]["BOLD"]["col_names"] # type: ignore
assert_array_equal(t_feature[cols].values, t_data) # type: ignore assert_array_equal(t_feature[cols].values, t_data) # type: ignore
feature_md5 = list(features.keys())[1] feature_md5 = list(features.keys())[1]
t_feature = storage.read_df(feature_md5=feature_md5) t_feature = storage.read_df(feature_md5=feature_md5)
fname = "gmd_schaefer100x7_std" fname = "tian_std"
t_data = out[fname]["VBM_GM"]["data"] # type: ignore t_data = out[fname]["BOLD"]["data"] # type: ignore
cols = out[fname]["VBM_GM"]["col_names"] # type: ignore cols = out[fname]["BOLD"]["col_names"] # type: ignore
assert_array_equal(t_feature[cols].values, t_data) # type: ignore assert_array_equal(t_feature[cols].values, t_data) # type: ignore
feature_md5 = list(features.keys())[2] feature_md5 = list(features.keys())[2]
t_feature = storage.read_df(feature_md5=feature_md5) t_feature = storage.read_df(feature_md5=feature_md5)
fname = "gmd_schaefer100x7_trim_mean90" fname = "tian_trim_mean90"
t_data = out[fname]["VBM_GM"]["data"] # type: ignore t_data = out[fname]["BOLD"]["data"] # type: ignore
cols = out[fname]["VBM_GM"]["col_names"] # type: ignore cols = out[fname]["BOLD"]["col_names"] # type: ignore
assert_array_equal(t_feature[cols].values, t_data) # type: ignore assert_array_equal(t_feature[cols].values, t_data) # type: ignore

View file

@ -8,52 +8,47 @@
from pathlib import Path from pathlib import Path
from nilearn import image
from nilearn.maskers import NiftiLabelsMasker from nilearn.maskers import NiftiLabelsMasker
from junifer.data import load_parcellation from junifer.data import get_parcellation
from junifer.datareader import DefaultDataReader
from junifer.markers.ets_rss import RSSETSMarker from junifer.markers.ets_rss import RSSETSMarker
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
# Set parcellation # Set parcellation
PARCELLATION = "Schaefer100x17" PARCELLATION = "TianxS1x3TxMNInonlinear2009cAsym"
def test_compute() -> None: def test_compute() -> None:
"""Test RSS ETS compute().""" """Test RSS ETS compute()."""
with SPMAuditoryTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
# Fetch element element_data = DefaultDataReader().fit_transform(dg["sub-01"])
out = dg["sub001"]
# Load BOLD image
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
# Create input data
input_dict = {
"data": niimg,
"path": out["BOLD"]["path"],
"space": "MNI",
}
# Compute the RSSETSMarker # Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) marker = RSSETSMarker(parcellation=PARCELLATION)
new_out = ets_rss_marker.compute(input_dict) rss_ets = marker.compute(element_data["BOLD"])
# Load parcellation # Compare with nilearn
test_parcellation, _, _, _ = load_parcellation(PARCELLATION) # Load testing parcellation
# Compute the NiftiLabelsMasker test_parcellation, _ = get_parcellation(
test_masker = NiftiLabelsMasker(test_parcellation) parcellation=[PARCELLATION],
test_ts = test_masker.fit_transform(niimg) target_data=element_data["BOLD"],
)
# Extract timeseries
nifti_labels_masker = NiftiLabelsMasker(labels_img=test_parcellation)
extacted_timeseries = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"]
)
# Assert the dimension of timeseries # Assert the dimension of timeseries
n_time, _ = test_ts.shape assert extacted_timeseries.shape[0] == len(rss_ets["data"])
assert n_time == len(new_out["data"])
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test RSS ETS get_output_type().""" """Test RSS ETS get_output_type()."""
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) assert "timeseries" == RSSETSMarker(
input_ = "BOLD" parcellation=PARCELLATION
output = ets_rss_marker.get_output_type(input_) ).get_output_type("BOLD")
assert output == "timeseries"
def test_store(tmp_path: Path) -> None: def test_store(tmp_path: Path) -> None:
@ -65,20 +60,13 @@ def test_store(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
with SPMAuditoryTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
# Fetch element element_data = DefaultDataReader().fit_transform(dg["sub-01"])
elem = dg["sub001"]
# Load BOLD image
niimg = image.load_img(str(elem["BOLD"]["path"].absolute()))
elem["BOLD"]["data"] = niimg
# Compute the RSSETSMarker # Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) marker = RSSETSMarker(parcellation=PARCELLATION)
# Create storage # Create storage
storage = SQLiteFeatureStorage( storage = SQLiteFeatureStorage(tmp_path / "test_rss_ets.sqlite")
uri=str((tmp_path / "test.sqlite").absolute())
)
# Store # Store
ets_rss_marker.fit_transform(input=elem, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any(x["name"] == "BOLD_RSSETSMarker" for x in features.values()) assert any(x["name"] == "BOLD_RSSETSMarker" for x in features.values())

View file

@ -10,16 +10,17 @@ from pathlib import Path
import nibabel as nib import nibabel as nib
import numpy as np import numpy as np
import pytest import pytest
from nilearn import datasets from nilearn.image import math_img, new_img_like
from nilearn.image import concat_imgs, math_img, new_img_like, resample_to_img
from nilearn.maskers import NiftiLabelsMasker, NiftiMasker from nilearn.maskers import NiftiLabelsMasker, NiftiMasker
from nilearn.masking import compute_brain_mask from nilearn.masking import compute_brain_mask
from numpy.testing import assert_array_almost_equal, assert_array_equal from numpy.testing import assert_array_almost_equal, assert_array_equal
from scipy.stats import trim_mean from scipy.stats import trim_mean
from junifer.data import load_mask, load_parcellation, register_parcellation from junifer.data import get_mask, get_parcellation, register_parcellation
from junifer.datareader import DefaultDataReader
from junifer.markers.parcel_aggregation import ParcelAggregation from junifer.markers.parcel_aggregation import ParcelAggregation
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
def test_ParcelAggregation_input_output() -> None: def test_ParcelAggregation_input_output() -> None:
@ -36,31 +37,38 @@ def test_ParcelAggregation_input_output() -> None:
def test_ParcelAggregation_3D() -> None: def test_ParcelAggregation_3D() -> None:
"""Test ParcelAggregation object on 3D images.""" """Test ParcelAggregation object on 3D images."""
# Get the testing parcellation (for nilearn) with PartlyCloudyTestingDataGrabber() as dg:
parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Create ParcelAggregation object
# Get the oasis VBM data marker = ParcelAggregation(
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) parcellation="TianxS1x3TxMNInonlinear2009cAsym",
vbm = oasis_dataset.gray_matter_maps[0] method="mean",
img = nib.load(vbm) on="BOLD",
# Mask parcellation manually
parcellation_img_res = resample_to_img(
parcellation.maps,
img,
interpolation="nearest",
) )
parcellation_bin = math_img( element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
# Compare with nilearn
# Load testing parcellation
testing_parcellation, _ = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
target_data=element_data["BOLD"],
)
# Binarize parcellation
testing_parcellation_bin = math_img(
"img != 0", "img != 0",
img=parcellation_img_res, img=testing_parcellation,
) )
# Create NiftiMasker # Create NiftiMasker
masker = NiftiMasker(parcellation_bin, target_affine=img.affine) masker = NiftiMasker(
data = masker.fit_transform(img) testing_parcellation_bin,
parcellation_values = masker.transform(parcellation_img_res) target_affine=element_data["BOLD"]["data"].affine,
parcellation_values = np.squeeze(parcellation_values).astype(int) )
data = masker.fit_transform(element_data["BOLD"]["data"])
parcellation_values = np.squeeze(
masker.transform(testing_parcellation)
).astype(int)
# Compute the mean manually # Compute the mean manually
manual = [] manual = []
for t_v in sorted(np.unique(parcellation_values)): for t_v in sorted(np.unique(parcellation_values)):
@ -69,41 +77,47 @@ def test_ParcelAggregation_3D() -> None:
manual = np.array(manual)[np.newaxis, :] manual = np.array(manual)[np.newaxis, :]
# Create NiftiLabelsMasker # Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps) nifti_labels_masker = NiftiLabelsMasker(
auto = nifti_masker.fit_transform(img) labels_img=testing_parcellation
)
nifti_labels_masked_bold = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"].slicer[..., 0:1]
)
parcel_agg_mean_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
# Check that arrays are almost equal # Check that arrays are almost equal
assert_array_almost_equal(auto, manual) assert_array_equal(parcel_agg_mean_bold_data, manual)
assert_array_almost_equal(nifti_labels_masked_bold, manual)
# Use the ParcelAggregation object # Check further
marker = ParcelAggregation( assert parcel_agg_mean_bold_data.ndim == 2
parcellation="Schaefer100x7", assert parcel_agg_mean_bold_data.shape[0] == 1
method="mean", assert_array_equal(
name="gmd_schaefer100x7_mean", nifti_labels_masked_bold.shape, parcel_agg_mean_bold_data.shape
on="VBM_GM", )
) # Test passing "on" as a keyword argument assert_array_equal(nifti_labels_masked_bold, parcel_agg_mean_bold_data)
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}}
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values3d_mean.ndim == 2 # Compute std manually
assert jun_values3d_mean.shape[0] == 1
assert_array_equal(manual, jun_values3d_mean)
# Test using another function (std)
manual = [] manual = []
for t_v in sorted(np.unique(parcellation_values)): for t_v in sorted(np.unique(parcellation_values)):
t_values = np.std(data[:, parcellation_values == t_v]) t_values = np.std(data[:, parcellation_values == t_v])
manual.append(t_values) manual.append(t_values)
manual = np.array(manual)[np.newaxis, :] manual = np.array(manual)[np.newaxis, :]
# Use the ParcelAggregation object # Create ParcelAggregation object
marker = ParcelAggregation(parcellation="Schaefer100x7", method="std") marker = ParcelAggregation(
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} parcellation="TianxS1x3TxMNInonlinear2009cAsym",
jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"] method="std",
on="BOLD",
assert jun_values3d_std.ndim == 2 )
assert jun_values3d_std.shape[0] == 1 parcel_agg_std_bold_data = marker.fit_transform(element_data)["BOLD"][
assert_array_equal(manual, jun_values3d_std) "data"
]
assert parcel_agg_std_bold_data.ndim == 2
assert parcel_agg_std_bold_data.shape[0] == 1
assert_array_equal(parcel_agg_std_bold_data, manual)
# Test using another function with parameters # Test using another function with parameters
manual = [] manual = []
@ -116,43 +130,52 @@ def test_ParcelAggregation_3D() -> None:
manual.append(t_values) manual.append(t_values)
manual = np.array(manual)[np.newaxis, :] manual = np.array(manual)[np.newaxis, :]
# Use the ParcelAggregation object # Create ParcelAggregation object
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="trim_mean", method="trim_mean",
method_params={"proportiontocut": 0.1}, method_params={"proportiontocut": 0.1},
on="BOLD",
) )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} parcel_agg_trim_mean_bold_data = marker.fit_transform(element_data)[
jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"] "BOLD"
]["data"]
assert jun_values3d_tm.ndim == 2 assert parcel_agg_trim_mean_bold_data.ndim == 2
assert jun_values3d_tm.shape[0] == 1 assert parcel_agg_trim_mean_bold_data.shape[0] == 1
assert_array_equal(manual, jun_values3d_tm) assert_array_equal(parcel_agg_trim_mean_bold_data, manual)
def test_ParcelAggregation_4D(): def test_ParcelAggregation_4D():
"""Test ParcelAggregation object on 4D images.""" """Test ParcelAggregation object on 4D images."""
# Get the testing parcellation (for nilearn) with PartlyCloudyTestingDataGrabber() as dg:
parcellation = datasets.fetch_atlas_schaefer_2018( element_data = DefaultDataReader().fit_transform(dg["sub-01"])
n_rois=100, yeo_networks=7, resolution_mm=2 # Create ParcelAggregation object
marker = ParcelAggregation(
parcellation="TianxS1x3TxMNInonlinear2009cAsym", method="mean"
)
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
# Compare with nilearn
# Load testing parcellation
testing_parcellation, _ = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
target_data=element_data["BOLD"],
)
# Extract data
nifti_labels_masker = NiftiLabelsMasker(
labels_img=testing_parcellation
)
nifti_labels_masked_bold = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"]
) )
# Get the SPM auditory data: assert parcel_agg_bold_data.ndim == 2
subject_data = datasets.fetch_spm_auditory() assert_array_equal(
fmri_img = concat_imgs(subject_data.func) # type: ignore nifti_labels_masked_bold.shape, parcel_agg_bold_data.shape
)
# Create NiftiLabelsMasker assert_array_equal(nifti_labels_masked_bold, parcel_agg_bold_data)
nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps)
auto4d = nifti_masker.fit_transform(fmri_img)
# Create ParcelAggregation object
marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d)
def test_ParcelAggregation_storage(tmp_path: Path) -> None: def test_ParcelAggregation_storage(tmp_path: Path) -> None:
@ -164,42 +187,38 @@ def test_ParcelAggregation_storage(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# Get the oasis VBM data # Store 3D
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) with PartlyCloudyTestingDataGrabber() as dg:
vbm = oasis_dataset.gray_matter_maps[0] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
img = nib.load(vbm) storage = SQLiteFeatureStorage(
uri = tmp_path / "test_sphere_storage_3D.sqlite" uri=tmp_path / "test_parcel_storage_3D.sqlite", upsert="ignore"
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {
"element": {"subject": "sub-01", "session": "ses-01"},
"dependencies": {"nilearn", "nibabel"},
}
input = {"VBM_GM": {"data": img, "meta": meta, "space": "MNI"}}
marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="VBM_GM"
) )
marker = ParcelAggregation(
marker.fit_transform(input, storage=storage) parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean",
on="BOLD",
)
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "VBM_GM_ParcelAggregation" for x in features.values() x["name"] == "BOLD_ParcelAggregation" for x in features.values()
) )
meta = { # Store 4D
"element": {"subject": "sub-01", "session": "ses-01"}, with PartlyCloudyTestingDataGrabber() as dg:
"dependencies": {"nilearn", "nibabel"}, element_data = DefaultDataReader().fit_transform(dg["sub-01"])
} storage = SQLiteFeatureStorage(
# Get the SPM auditory data uri=tmp_path / "test_parcel_storage_4D.sqlite", upsert="ignore"
subject_data = datasets.fetch_spm_auditory() )
fmri_img = concat_imgs(subject_data.func) # type: ignore
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="BOLD" parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean",
on="BOLD",
) )
marker.fit_transform(input=element_data, storage=storage)
marker.fit_transform(input, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_ParcelAggregation" for x in features.values() x["name"] == "BOLD_ParcelAggregation" for x in features.values()
@ -208,86 +227,108 @@ def test_ParcelAggregation_storage(tmp_path: Path) -> None:
def test_ParcelAggregation_3D_mask() -> None: def test_ParcelAggregation_3D_mask() -> None:
"""Test ParcelAggregation object on 3D images with mask.""" """Test ParcelAggregation object on 3D images with mask."""
with PartlyCloudyTestingDataGrabber() as dg:
# Get the testing parcellation (for nilearn) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) # Create ParcelAggregation object
# Get one mask
mask_img, _, _ = load_mask("GM_prob0.2")
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker(
labels_img=parcellation.maps, mask_img=mask_img
)
auto = nifti_masker.fit_transform(img)
# Use the ParcelAggregation object
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
masks="GM_prob0.2", name="tian_mean",
name="gmd_schaefer100x7_mean", on="BOLD",
on="VBM_GM", masks="compute_brain_mask",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] ..., 0:1
]
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
assert jun_values3d_mean.ndim == 2 # Compare with nilearn
assert jun_values3d_mean.shape[0] == 1 # Load testing parcellation
assert_array_almost_equal(auto, jun_values3d_mean) testing_parcellation, _ = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
target_data=element_data["BOLD"],
)
# Load mask
mask_img = get_mask(
"compute_brain_mask", target_data=element_data["BOLD"]
)
# Extract data
nifti_labels_masker = NiftiLabelsMasker(
labels_img=testing_parcellation, mask_img=mask_img
)
nifti_labels_masked_bold = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"].slicer[..., 0:1]
)
assert parcel_agg_bold_data.ndim == 2
assert_array_equal(
nifti_labels_masked_bold.shape, parcel_agg_bold_data.shape
)
assert_array_equal(nifti_labels_masked_bold, parcel_agg_bold_data)
def test_ParcelAggregation_3D_mask_computed() -> None: def test_ParcelAggregation_3D_mask_computed() -> None:
"""Test ParcelAggregation object on 3D images with computed masks.""" """Test ParcelAggregation object on 3D images with computed masks."""
with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
# Get the testing parcellation (for nilearn) # Compare with nilearn
parcellation = datasets.fetch_atlas_schaefer_2018(n_rois=100) # Load testing parcellation
testing_parcellation, _ = get_parcellation(
# Get the oasis VBM data parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) target_data=element_data["BOLD"],
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Get one mask
mask_img = compute_brain_mask(img, threshold=0.2)
# Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker(
labels_img=parcellation.maps, mask_img=mask_img
) )
auto = nifti_masker.fit_transform(img) # Get a mask
mask_img = compute_brain_mask(
# Get one mask element_data["BOLD"]["data"], threshold=0.2
mask_img = compute_brain_mask(img, threshold=0.5) )
# Create NiftiLabelsMasker
# Create NiftiLabelsMasker nifti_labels_masker = NiftiLabelsMasker(
nifti_masker = NiftiLabelsMasker( labels_img=testing_parcellation, mask_img=mask_img
labels_img=parcellation.maps, mask_img=mask_img )
nifti_labels_masked_bold_good = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"]
)
# Get another mask
mask_img = compute_brain_mask(
element_data["BOLD"]["data"], threshold=0.5
)
# Create NiftiLabelsMasker
nifti_labels_masker = NiftiLabelsMasker(
labels_img=testing_parcellation, mask_img=mask_img
)
nifti_labels_masked_bold_bad = nifti_labels_masker.fit_transform(
mask_img
) )
auto_bad = nifti_masker.fit_transform(img)
# Use the ParcelAggregation object # Use the ParcelAggregation object
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
masks={"compute_brain_mask": {"threshold": 0.2}}, masks={"compute_brain_mask": {"threshold": 0.2}},
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} parcel_agg_mean_bold_data = marker.fit_transform(element_data)["BOLD"][
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] "data"
]
assert jun_values3d_mean.ndim == 2 assert parcel_agg_mean_bold_data.ndim == 2
assert jun_values3d_mean.shape[0] == 1 assert parcel_agg_mean_bold_data.shape[0] == 1
assert_array_almost_equal(auto, jun_values3d_mean) assert_array_almost_equal(
nifti_labels_masked_bold_good, parcel_agg_mean_bold_data
)
with pytest.raises(AssertionError): with pytest.raises(AssertionError):
assert_array_almost_equal(jun_values3d_mean, auto_bad) assert_array_almost_equal(
parcel_agg_mean_bold_data, nifti_labels_masked_bold_bad
)
def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None: def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
@ -299,29 +340,34 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
# Get the testing parcellation # Load testing parcellation
parcellation, labels, _, _ = load_parcellation("Schaefer100x7") testing_parcellation, labels = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
assert parcellation is not None target_data=element_data["BOLD"],
)
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create two parcellations from it # Create two parcellations from it
parcellation_data = parcellation.get_fdata() parcellation_data = testing_parcellation.get_fdata()
parcellation1_data = parcellation_data.copy() parcellation1_data = parcellation_data.copy()
parcellation1_data[parcellation1_data > 50] = 0 parcellation1_data[parcellation1_data > 8] = 0
parcellation2_data = parcellation_data.copy() parcellation2_data = parcellation_data.copy()
parcellation2_data[parcellation2_data <= 50] = 0 parcellation2_data[parcellation2_data <= 8] = 0
parcellation2_data[parcellation2_data > 0] -= 50 parcellation2_data[parcellation2_data > 0] -= 8
labels1 = labels[:50] labels1 = labels[:8]
labels2 = labels[50:] labels2 = labels[8:]
parcellation1_img = new_img_like(parcellation, parcellation1_data) parcellation1_img = new_img_like(
parcellation2_img = new_img_like(parcellation, parcellation2_data) testing_parcellation, parcellation1_data
)
parcellation2_img = new_img_like(
testing_parcellation, parcellation2_data
)
parcellation1_path = tmp_path / "parcellation1.nii.gz" parcellation1_path = tmp_path / "parcellation1.nii.gz"
parcellation2_path = tmp_path / "parcellation2.nii.gz" parcellation2_path = tmp_path / "parcellation2.nii.gz"
@ -330,54 +376,53 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
nib.save(parcellation2_img, parcellation2_path) nib.save(parcellation2_img, parcellation2_path)
register_parcellation( register_parcellation(
name="Schaefer100x7_low", name="TianxS1x3TxMNInonlinear2009cAsym_low",
parcellation_path=parcellation1_path, parcellation_path=parcellation1_path,
parcels_labels=labels1, parcels_labels=labels1,
space="MNI", space="MNI152NLin2009cAsym",
overwrite=True, overwrite=True,
) )
register_parcellation( register_parcellation(
name="Schaefer100x7_high", name="TianxS1x3TxMNInonlinear2009cAsym_high",
parcellation_path=parcellation2_path, parcellation_path=parcellation2_path,
parcels_labels=labels2, parcels_labels=labels2,
space="MNI", space="MNI152NLin2009cAsym",
overwrite=True, overwrite=True,
) )
# Use the ParcelAggregation object on the original parcellation # Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation( marker_original = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} orig_mean = marker_original.fit_transform(element_data)["BOLD"]
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2 assert orig_mean_data.ndim == 2
assert orig_mean_data.shape[0] == 1 assert orig_mean_data.shape == (1, 16)
assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean)
# Use the ParcelAggregation object on the two parcellations # Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation( marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"], parcellation=[
"TianxS1x3TxMNInonlinear2009cAsym_low",
"TianxS1x3TxMNInonlinear2009cAsym_high",
],
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}}
# No warnings should be raised # No warnings should be raised
with warnings.catch_warnings(): with warnings.catch_warnings():
warnings.simplefilter("error", category=UserWarning) warnings.simplefilter("error", category=UserWarning)
split_mean = marker_split.fit_transform(input)["VBM_GM"] split_mean = marker_split.fit_transform(element_data)["BOLD"]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2 assert split_mean_data.ndim == 2
assert split_mean_data.shape[0] == 1 assert split_mean_data.shape == (1, 16)
assert split_mean_data.shape[1] == 100
# Data and labels should be the same # Data and labels should be the same
assert_array_equal(orig_mean_data, split_mean_data) assert_array_equal(orig_mean_data, split_mean_data)
@ -393,31 +438,36 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
# Get the testing parcellation # Load testing parcellation
parcellation, labels, _, _ = load_parcellation("Schaefer100x7") testing_parcellation, labels = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
assert parcellation is not None target_data=element_data["BOLD"],
)
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create two parcellations from it # Create two parcellations from it
parcellation_data = parcellation.get_fdata() parcellation_data = testing_parcellation.get_fdata()
parcellation1_data = parcellation_data.copy() parcellation1_data = parcellation_data.copy()
parcellation1_data[parcellation1_data > 50] = 0 parcellation1_data[parcellation1_data > 8] = 0
parcellation2_data = parcellation_data.copy() parcellation2_data = parcellation_data.copy()
# Make the second parcellation overlap with the first # Make the second parcellation overlap with the first
parcellation2_data[parcellation2_data <= 45] = 0 parcellation2_data[parcellation2_data <= 6] = 0
parcellation2_data[parcellation2_data > 0] -= 45 parcellation2_data[parcellation2_data > 0] -= 6
labels1 = [f"low_{x}" for x in labels[:50]] # Change the labels labels1 = [f"low_{x}" for x in labels[:8]] # Change the labels
labels2 = [f"high_{x}" for x in labels[45:]] # Change the labels labels2 = [f"high_{x}" for x in labels[6:]] # Change the labels
parcellation1_img = new_img_like(parcellation, parcellation1_data) parcellation1_img = new_img_like(
parcellation2_img = new_img_like(parcellation, parcellation2_data) testing_parcellation, parcellation1_data
)
parcellation2_img = new_img_like(
testing_parcellation, parcellation2_data
)
parcellation1_path = tmp_path / "parcellation1.nii.gz" parcellation1_path = tmp_path / "parcellation1.nii.gz"
parcellation2_path = tmp_path / "parcellation2.nii.gz" parcellation2_path = tmp_path / "parcellation2.nii.gz"
@ -426,62 +476,62 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
nib.save(parcellation2_img, parcellation2_path) nib.save(parcellation2_img, parcellation2_path)
register_parcellation( register_parcellation(
name="Schaefer100x7_low2", name="TianxS1x3TxMNInonlinear2009cAsym_low",
parcellation_path=parcellation1_path, parcellation_path=parcellation1_path,
parcels_labels=labels1, parcels_labels=labels1,
space="MNI", space="MNI152NLin2009cAsym",
overwrite=True, overwrite=True,
) )
register_parcellation( register_parcellation(
name="Schaefer100x7_high2", name="TianxS1x3TxMNInonlinear2009cAsym_high",
parcellation_path=parcellation2_path, parcellation_path=parcellation2_path,
parcels_labels=labels2, parcels_labels=labels2,
space="MNI", space="MNI152NLin2009cAsym",
overwrite=True, overwrite=True,
) )
# Use the ParcelAggregation object on the original parcellation # Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation( marker_original = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} orig_mean = marker_original.fit_transform(element_data)["BOLD"]
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2 assert orig_mean_data.ndim == 2
assert orig_mean_data.shape[0] == 1 assert orig_mean_data.shape == (1, 16)
assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean)
# Use the ParcelAggregation object on the two parcellations # Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation( marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"], parcellation=[
"TianxS1x3TxMNInonlinear2009cAsym_low",
"TianxS1x3TxMNInonlinear2009cAsym_high",
],
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} # Warning should be raised
with pytest.warns(RuntimeWarning, match="overlapping voxels"): with pytest.warns(RuntimeWarning, match="overlapping voxels"):
split_mean = marker_split.fit_transform(input)["VBM_GM"] split_mean = marker_split.fit_transform(element_data)["BOLD"]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2 assert split_mean_data.ndim == 2
assert split_mean_data.shape[0] == 1 assert split_mean_data.shape == (1, 18)
assert split_mean_data.shape[1] == 105
# Overlapping voxels should be NaN # Overlapping voxels should be NaN
assert np.isnan(split_mean_data[:, 50:55]).all() assert np.isnan(split_mean_data[:, 8:10]).all()
non_nan = split_mean_data[~np.isnan(split_mean_data)] non_nan = split_mean_data[~np.isnan(split_mean_data)]
# Data should be the same # Data should be the same
assert_array_equal(orig_mean_data, non_nan[None, :]) assert_array_equal(orig_mean_data, non_nan[None, :])
# Labels should be "low" for the first 50 and "high" for the second 50 # Labels should be "low" for the first 8 and "high" for the second 8
assert all(x.startswith("low") for x in split_mean["col_names"][:50]) assert all(x.startswith("low") for x in split_mean["col_names"][:8])
assert all(x.startswith("high") for x in split_mean["col_names"][50:]) assert all(x.startswith("high") for x in split_mean["col_names"][8:])
def test_ParcelAggregation_3D_multiple_duplicated_labels( def test_ParcelAggregation_3D_multiple_duplicated_labels(
@ -495,29 +545,34 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
The path to the test directory. The path to the test directory.
""" """
with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
# Get the testing parcellation # Load testing parcellation
parcellation, labels, _, _ = load_parcellation("Schaefer100x7") testing_parcellation, labels = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
assert parcellation is not None target_data=element_data["BOLD"],
)
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create two parcellations from it # Create two parcellations from it
parcellation_data = parcellation.get_fdata() parcellation_data = testing_parcellation.get_fdata()
parcellation1_data = parcellation_data.copy() parcellation1_data = parcellation_data.copy()
parcellation1_data[parcellation1_data > 50] = 0 parcellation1_data[parcellation1_data > 8] = 0
parcellation2_data = parcellation_data.copy() parcellation2_data = parcellation_data.copy()
parcellation2_data[parcellation2_data <= 50] = 0 parcellation2_data[parcellation2_data <= 8] = 0
parcellation2_data[parcellation2_data > 0] -= 50 parcellation2_data[parcellation2_data > 0] -= 8
labels1 = labels[:50] labels1 = labels[:8]
labels2 = labels[49:-1] # One label is duplicated labels2 = labels[7:-1] # One label is duplicated
parcellation1_img = new_img_like(parcellation, parcellation1_data) parcellation1_img = new_img_like(
parcellation2_img = new_img_like(parcellation, parcellation2_data) testing_parcellation, parcellation1_data
)
parcellation2_img = new_img_like(
testing_parcellation, parcellation2_data
)
parcellation1_path = tmp_path / "parcellation1.nii.gz" parcellation1_path = tmp_path / "parcellation1.nii.gz"
parcellation2_path = tmp_path / "parcellation2.nii.gz" parcellation2_path = tmp_path / "parcellation2.nii.gz"
@ -526,104 +581,128 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
nib.save(parcellation2_img, parcellation2_path) nib.save(parcellation2_img, parcellation2_path)
register_parcellation( register_parcellation(
name="Schaefer100x7_low", name="TianxS1x3TxMNInonlinear2009cAsym_low",
parcellation_path=parcellation1_path, parcellation_path=parcellation1_path,
parcels_labels=labels1, parcels_labels=labels1,
space="MNI", space="MNI152NLin2009cAsym",
overwrite=True, overwrite=True,
) )
register_parcellation( register_parcellation(
name="Schaefer100x7_high", name="TianxS1x3TxMNInonlinear2009cAsym_high",
parcellation_path=parcellation2_path, parcellation_path=parcellation2_path,
parcels_labels=labels2, parcels_labels=labels2,
space="MNI", space="MNI152NLin2009cAsym",
overwrite=True, overwrite=True,
) )
# Use the ParcelAggregation object on the original parcellation # Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation( marker_original = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} orig_mean = marker_original.fit_transform(element_data)["BOLD"]
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2 assert orig_mean_data.ndim == 2
assert orig_mean_data.shape[0] == 1 assert orig_mean_data.shape == (1, 16)
assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean)
# Use the ParcelAggregation object on the two parcellations # Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation( marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"], parcellation=[
"TianxS1x3TxMNInonlinear2009cAsym_low",
"TianxS1x3TxMNInonlinear2009cAsym_high",
],
method="mean", method="mean",
name="gmd_schaefer100x7_mean", name="tian_mean",
on="VBM_GM", on="BOLD",
) # Test passing "on" as a keyword argument )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}}
# Warning should be raised
with pytest.warns(RuntimeWarning, match="duplicated labels."): with pytest.warns(RuntimeWarning, match="duplicated labels."):
split_mean = marker_split.fit_transform(input)["VBM_GM"] split_mean = marker_split.fit_transform(element_data)["BOLD"]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2 assert split_mean_data.ndim == 2
assert split_mean_data.shape[0] == 1 assert split_mean_data.shape == (1, 16)
assert split_mean_data.shape[1] == 100
# Data should be the same # Data should be the same
assert_array_equal(orig_mean_data, split_mean_data) assert_array_equal(orig_mean_data, split_mean_data)
# Labels should be prefixed with the parcellation name # Labels should be prefixed with the parcellation name
col_names = [f"Schaefer100x7_low_{x}" for x in labels1] col_names = [
col_names += [f"Schaefer100x7_high_{x}" for x in labels2] f"TianxS1x3TxMNInonlinear2009cAsym_low_{x}" for x in labels1
]
col_names += [
f"TianxS1x3TxMNInonlinear2009cAsym_high_{x}" for x in labels2
]
assert col_names == split_mean["col_names"] assert col_names == split_mean["col_names"]
def test_ParcelAggregation_4D_agg_time(): def test_ParcelAggregation_4D_agg_time():
"""Test ParcelAggregation object on 4D images, aggregating time.""" """Test ParcelAggregation object on 4D images, aggregating time."""
# Get the testing parcellation (for nilearn) with PartlyCloudyTestingDataGrabber() as dg:
parcellation = datasets.fetch_atlas_schaefer_2018( element_data = DefaultDataReader().fit_transform(dg["sub-01"])
n_rois=100, yeo_networks=7, resolution_mm=2
)
# Get the SPM auditory data:
subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore
# Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps)
auto4d = nifti_masker.fit_transform(fmri_img)
auto_mean = auto4d.mean(axis=0)
# Create ParcelAggregation object # Create ParcelAggregation object
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", time_method="mean" parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean",
time_method="mean",
on="BOLD",
) )
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
jun_values4d = marker.fit_transform(input)["BOLD"]["data"] "data"
]
assert jun_values4d.ndim == 1 # Compare with nilearn
assert_array_equal(auto_mean.shape, jun_values4d.shape) # Loading testing parcellation
assert_array_almost_equal(auto_mean, jun_values4d, decimal=2) testing_parcellation, _ = get_parcellation(
parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
target_data=element_data["BOLD"],
)
# Extract data
nifti_labels_masker = NiftiLabelsMasker(
labels_img=testing_parcellation
)
nifti_labels_masked_bold = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"]
)
nifti_labels_masked_bold_mean = nifti_labels_masked_bold.mean(axis=0)
auto_pick_0 = auto4d[:1, :] assert parcel_agg_bold_data.ndim == 1
assert_array_equal(
nifti_labels_masked_bold_mean.shape, parcel_agg_bold_data.shape
)
assert_array_almost_equal(
nifti_labels_masked_bold_mean, parcel_agg_bold_data, decimal=2
)
# Test picking first time point
nifti_labels_masked_bold_pick_0 = nifti_labels_masked_bold[:1, :]
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean", method="mean",
time_method="select", time_method="select",
time_method_params={"pick": [0]}, time_method_params={"pick": [0]},
on="BOLD",
)
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
assert parcel_agg_bold_data.ndim == 2
assert_array_equal(
nifti_labels_masked_bold_pick_0.shape, parcel_agg_bold_data.shape
)
assert_array_equal(
nifti_labels_masked_bold_pick_0, parcel_agg_bold_data
) )
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto_pick_0.shape, jun_values4d.shape)
assert_array_equal(auto_pick_0, jun_values4d)
def test_ParcelAggregation_errors() -> None:
"""Test errors for ParcelAggregation."""
with pytest.raises(ValueError, match="can only be used with BOLD data"): with pytest.raises(ValueError, match="can only be used with BOLD data"):
ParcelAggregation( ParcelAggregation(
parcellation="Schaefer100x7", parcellation="Schaefer100x7",
@ -643,12 +722,22 @@ def test_ParcelAggregation_4D_agg_time():
on="VBM_GM", on="VBM_GM",
) )
with pytest.warns(RuntimeWarning, match="No time dimension to aggregate"):
input = { def test_ParcelAggregation_warning() -> None:
"BOLD": { """Test warning for ParcelAggregation."""
"data": fmri_img.slicer[..., 0:1], with PartlyCloudyTestingDataGrabber() as dg:
"meta": {}, element_data = DefaultDataReader().fit_transform(dg["sub-01"])
"space": "MNI", with pytest.warns(
} RuntimeWarning, match="No time dimension to aggregate"
} ):
marker.fit_transform(input) marker = ParcelAggregation(
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
method="mean",
time_method="select",
time_method_params={"pick": [0]},
on="BOLD",
)
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
marker.fit_transform(element_data)

View file

@ -1,22 +1,23 @@
"""Provide tests for sphere aggregation.""" """Provide tests for sphere aggregation."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de> # Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import typing
from pathlib import Path from pathlib import Path
from typing import Dict
import nibabel as nib
import pytest import pytest
from nilearn import datasets
from nilearn.image import concat_imgs
from nilearn.maskers import NiftiSpheresMasker from nilearn.maskers import NiftiSpheresMasker
from numpy.testing import assert_array_equal from numpy.testing import assert_array_equal
from junifer.data import load_coordinates, load_mask from junifer.data import get_coordinates, get_mask
from junifer.datareader import DefaultDataReader
from junifer.markers.sphere_aggregation import SphereAggregation from junifer.markers.sphere_aggregation import SphereAggregation
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import (
OasisVBMTestingDataGrabber,
SPMAuditoryTestingDataGrabber,
)
# Define common variables # Define common variables
@ -36,51 +37,66 @@ def test_SphereAggregation_input_output() -> None:
def test_SphereAggregation_3D() -> None: def test_SphereAggregation_3D() -> None:
"""Test SphereAggregation object on 3D images.""" """Test SphereAggregation object on 3D images."""
# Get the testing coordinates (for nilearn) with OasisVBMTestingDataGrabber() as dg:
coordinates, _, _ = load_coordinates(COORDS) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS)
auto4d = nifti_masker.fit_transform(img)
# Create SphereAggregation object # Create SphereAggregation object
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
) )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"] "data"
]
assert jun_values4d.ndim == 2 # Compare with nilearn
assert_array_equal(auto4d.shape, jun_values4d.shape) # Load testing coordinates
assert_array_equal(auto4d, jun_values4d) testing_coords, _ = get_coordinates(
coords=COORDS, target_data=element_data["VBM_GM"]
)
# Extract data
nifti_spheres_masker = NiftiSpheresMasker(
seeds=testing_coords, radius=RADIUS
)
nifti_spheres_masked_vbm_gm = nifti_spheres_masker.fit_transform(
element_data["VBM_GM"]["data"]
)
assert sphere_agg_vbm_gm_data.ndim == 2
assert_array_equal(
nifti_spheres_masked_vbm_gm.shape, sphere_agg_vbm_gm_data.shape
)
assert_array_equal(nifti_spheres_masked_vbm_gm, sphere_agg_vbm_gm_data)
def test_SphereAggregation_4D() -> None: def test_SphereAggregation_4D() -> None:
"""Test SphereAggregation object on 4D images.""" """Test SphereAggregation object on 4D images."""
# Get the testing coordinates (for nilearn) with SPMAuditoryTestingDataGrabber() as dg:
coordinates, _, _ = load_coordinates(COORDS) element_data = DefaultDataReader().fit_transform(dg["sub001"])
# Get the SPM auditory data
subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore
# Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS)
auto4d = nifti_masker.fit_transform(fmri_img)
# Create SphereAggregation object # Create SphereAggregation object
marker = SphereAggregation(coords=COORDS, method="mean", radius=RADIUS) marker = SphereAggregation(
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
jun_values4d = marker.fit_transform(input)["BOLD"]["data"] )
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
assert jun_values4d.ndim == 2 # Compare with nilearn
assert_array_equal(auto4d.shape, jun_values4d.shape) # Load testing coordinates
assert_array_equal(auto4d, jun_values4d) testing_coords, _ = get_coordinates(
coords=COORDS, target_data=element_data["BOLD"]
)
# Extract data
nifti_spheres_masker = NiftiSpheresMasker(
seeds=testing_coords, radius=RADIUS
)
nifti_spheres_masked_bold = nifti_spheres_masker.fit_transform(
element_data["BOLD"]["data"]
)
assert sphere_agg_bold_data.ndim == 2
assert_array_equal(
nifti_spheres_masked_bold.shape, sphere_agg_bold_data.shape
)
assert_array_equal(nifti_spheres_masked_bold, sphere_agg_bold_data)
def test_SphereAggregation_storage(tmp_path: Path) -> None: def test_SphereAggregation_storage(tmp_path: Path) -> None:
@ -92,43 +108,32 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# Get the oasis VBM data # Store 3D
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) with OasisVBMTestingDataGrabber() as dg:
vbm = oasis_dataset.gray_matter_maps[0] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
img = nib.load(vbm) storage = SQLiteFeatureStorage(
uri = tmp_path / "test_sphere_storage_3D.sqlite" uri=tmp_path / "test_sphere_storage_3D.sqlite", upsert="ignore"
)
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {
"element": {"subject": "sub-01", "session": "ses-01"},
"dependencies": {"nilearn", "nibabel"},
}
input = {"VBM_GM": {"data": img, "meta": meta, "space": "MNI"}}
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
) )
marker.fit_transform(input=element_data, storage=storage)
marker.fit_transform(input, storage=storage) features = storage.list_features()
features: Dict = typing.cast(Dict, storage.list_features())
assert any( assert any(
x["name"] == "VBM_GM_SphereAggregation" for x in features.values() x["name"] == "VBM_GM_SphereAggregation" for x in features.values()
) )
meta = { # Store 4D
"element": {"subject": "sub-01", "session": "ses-01"}, with SPMAuditoryTestingDataGrabber() as dg:
"dependencies": {"nilearn", "nibabel"}, element_data = DefaultDataReader().fit_transform(dg["sub001"])
} storage = SQLiteFeatureStorage(
# Get the SPM auditory data uri=tmp_path / "test_sphere_storage_4D.sqlite", upsert="ignore"
subject_data = datasets.fetch_spm_auditory() )
fmri_img = concat_imgs(subject_data.func) # type: ignore
input = {"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="BOLD" coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
) )
marker.fit_transform(input=element_data, storage=storage)
marker.fit_transform(input, storage=storage) features = storage.list_features()
features: Dict = typing.cast(Dict, storage.list_features())
assert any( assert any(
x["name"] == "BOLD_SphereAggregation" for x in features.values() x["name"] == "BOLD_SphereAggregation" for x in features.values()
) )
@ -136,80 +141,110 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
def test_SphereAggregation_3D_mask() -> None: def test_SphereAggregation_3D_mask() -> None:
"""Test SphereAggregation object on 3D images using mask.""" """Test SphereAggregation object on 3D images using mask."""
# Get the testing coordinates (for nilearn) with OasisVBMTestingDataGrabber() as dg:
coordinates, _, _ = load_coordinates(COORDS) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Get one mask
mask_img, _, _ = load_mask("GM_prob0.2")
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(
seeds=coordinates, radius=RADIUS, mask_img=mask_img
)
auto4d = nifti_masker.fit_transform(img)
# Create SphereAggregation object # Create SphereAggregation object
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, coords=COORDS,
method="mean", method="mean",
radius=RADIUS, radius=RADIUS,
on="VBM_GM", on="VBM_GM",
masks="GM_prob0.2", masks="compute_brain_mask",
) )
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}} sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"] "data"
]
assert jun_values4d.ndim == 2 # Compare with nilearn
assert_array_equal(auto4d.shape, jun_values4d.shape) # Load testing coordinates
assert_array_equal(auto4d, jun_values4d) testing_coords, _ = get_coordinates(
coords=COORDS, target_data=element_data["VBM_GM"]
)
# Load mask
mask_img = get_mask(
"compute_brain_mask", target_data=element_data["VBM_GM"]
)
# Extract data
nifti_spheres_masker = NiftiSpheresMasker(
seeds=testing_coords, radius=RADIUS, mask_img=mask_img
)
nifti_spheres_masked_vbm_agg = nifti_spheres_masker.fit_transform(
element_data["VBM_GM"]["data"]
)
assert sphere_agg_vbm_gm_data.ndim == 2
assert_array_equal(
nifti_spheres_masked_vbm_agg.shape,
nifti_spheres_masked_vbm_agg.shape,
)
assert_array_equal(
nifti_spheres_masked_vbm_agg, nifti_spheres_masked_vbm_agg
)
def test_SphereAggregation_4D_agg_time() -> None: def test_SphereAggregation_4D_agg_time() -> None:
"""Test SphereAggregation object on 4D images, aggregating time.""" """Test SphereAggregation object on 4D images, aggregating time."""
# Get the testing coordinates (for nilearn) with SPMAuditoryTestingDataGrabber() as dg:
coordinates, _, _ = load_coordinates(COORDS) element_data = DefaultDataReader().fit_transform(dg["sub001"])
# Get the SPM auditory data
subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore
# Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS)
auto4d = nifti_masker.fit_transform(fmri_img)
auto_mean = auto4d.mean(axis=0)
# Create SphereAggregation object # Create SphereAggregation object
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, time_method="mean" coords=COORDS,
method="mean",
radius=RADIUS,
time_method="mean",
on="BOLD",
) )
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
jun_values4d = marker.fit_transform(input)["BOLD"]["data"] "data"
]
assert jun_values4d.ndim == 1 # Compare with nilearn
assert_array_equal(auto_mean.shape, jun_values4d.shape) # Load testing coordinates
assert_array_equal(auto_mean, jun_values4d) testing_coords, _ = get_coordinates(
coords=COORDS, target_data=element_data["BOLD"]
)
# Extract data
nifti_spheres_masker = NiftiSpheresMasker(
seeds=testing_coords, radius=RADIUS
)
nifti_spheres_masked_bold = nifti_spheres_masker.fit_transform(
element_data["BOLD"]["data"]
)
nifti_spheres_masked_bold_mean = nifti_spheres_masked_bold.mean(axis=0)
auto_pick_0 = auto4d[:1, :] assert sphere_agg_bold_data.ndim == 1
assert_array_equal(
nifti_spheres_masked_bold_mean.shape, sphere_agg_bold_data.shape
)
assert_array_equal(
nifti_spheres_masked_bold_mean, sphere_agg_bold_data
)
# Test picking first time point
nifti_spheres_masked_bold_pick_0 = nifti_spheres_masked_bold[:1, :]
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, coords=COORDS,
method="mean", method="mean",
radius=RADIUS, radius=RADIUS,
time_method="select", time_method="select",
time_method_params={"pick": [0]}, time_method_params={"pick": [0]},
on="BOLD",
)
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
assert sphere_agg_bold_data.ndim == 2
assert_array_equal(
nifti_spheres_masked_bold_pick_0.shape, sphere_agg_bold_data.shape
)
assert_array_equal(
nifti_spheres_masked_bold_pick_0, sphere_agg_bold_data
) )
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto_pick_0.shape, jun_values4d.shape)
assert_array_equal(auto_pick_0, jun_values4d)
def test_SphereAggregation_errors() -> None:
"""Test errors for SphereAggregation."""
with pytest.raises(ValueError, match="can only be used with BOLD data"): with pytest.raises(ValueError, match="can only be used with BOLD data"):
SphereAggregation( SphereAggregation(
coords=COORDS, coords=COORDS,
@ -231,12 +266,23 @@ def test_SphereAggregation_4D_agg_time() -> None:
on="VBM_GM", on="VBM_GM",
) )
with pytest.warns(RuntimeWarning, match="No time dimension to aggregate"):
input = { def test_SphereAggregation_warning() -> None:
"BOLD": { """Test warning for SphereAggregation."""
"data": fmri_img.slicer[..., 0:1], with SPMAuditoryTestingDataGrabber() as dg:
"meta": {}, element_data = DefaultDataReader().fit_transform(dg["sub001"])
"space": "MNI", with pytest.warns(
} RuntimeWarning, match="No time dimension to aggregate"
} ):
marker.fit_transform(input) marker = SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
time_method="select",
time_method_params={"pick": [0]},
on="BOLD",
)
element_data["BOLD"]["data"] = element_data["BOLD"]["data"].slicer[
..., 0:1
]
marker.fit_transform(element_data)