[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."""
"name": "Schaefer1000x7_Mean", return {
"kind": "ParcelAggregation", "kind": "PartlyCloudyTestingDataGrabber",
"parcellation": "Schaefer1000x7", }
"method": "mean",
},
{
"name": "Schaefer1000x7_Std",
"kind": "ParcelAggregation",
"parcellation": "Schaefer1000x7",
"method": "std",
},
]
# Define storage
storage = {
"kind": "SQLiteFeatureStorage",
}
def test_run_single_element(tmp_path: Path) -> None: @pytest.fixture
def markers() -> List[Dict[str, str]]:
"""Return markers as a list of dictionary."""
return [
{
"name": "tian-s1-3T_mean",
"kind": "ParcelAggregation",
"parcellation": "TianxS1x3TxMNInonlinear2009cAsym",
"method": "mean",
},
{
"name": "tian-s1-3T_std",
"kind": "ParcelAggregation",
"parcellation": "TianxS1x3TxMNInonlinear2009cAsym",
"method": "std",
},
]
@pytest.fixture
def storage() -> Dict[str, str]:
"""Return a storage as a dictionary."""
return {
"kind": "SQLiteFeatureStorage",
}
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 = {}
mask_img = mask_object(target_img, **mask_params) # From nilearn
if mask_name != "compute_brain_mask":
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) mask_img = intersect_masks(all_masks, **intersect_params)
# 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)
# 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"
@ -413,8 +572,8 @@ def get_mask( # noqa: C901
# Load nifti # Load nifti
mask_img = nib.load(warped_mask_path) mask_img = nib.load(warped_mask_path)
# Delete tempdir # Delete tempdir
WorkDirManager().delete_tempdir(tempdir) WorkDirManager().delete_tempdir(tempdir)
return mask_img # type: ignore return mask_img # type: ignore

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"]
# Create component-scoped tempdir # Get the min of the voxels sizes and use it as the resolution
tempdir = WorkDirManager().get_tempdir(prefix="parcellations") target_img = target_data["data"]
resolution = np.min(target_img.header.get_zooms()[:3])
# Create component-scoped tempdir
tempdir = WorkDirManager().get_tempdir(prefix="parcellations")
# Create element-scoped tempdir so that warped parcellation is
# available later as nibabel stores file path reference for
# loading on computation
element_tempdir = WorkDirManager().get_element_tempdir(
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 # Save parcellation image to a component-scoped tempfile
prewarp_parcellation_path = tempdir / "prewarp_parcellation.nii.gz" prewarp_parcellation_path = tempdir / "prewarp_parcellation.nii.gz"
nib.save(resampled_parcellation_img, prewarp_parcellation_path) nib.save(resampled_parcellation_img, prewarp_parcellation_path)
# Create element-scoped tempdir so that warped parcellation is
# available later as nibabel stores file path reference for
# loading on computation
element_tempdir = WorkDirManager().get_element_tempdir(
prefix="parcellations"
)
# 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"
@ -359,8 +408,8 @@ def get_parcellation(
# Load nifti # Load nifti
resampled_parcellation_img = nib.load(warped_parcellation_path) resampled_parcellation_img = nib.load(warped_parcellation_path)
# Delete tempdir # Delete tempdir
WorkDirManager().delete_tempdir(tempdir) WorkDirManager().delete_tempdir(tempdir)
return resampled_parcellation_img, labels # type: ignore return resampled_parcellation_img, labels # type: ignore

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,7 +541,13 @@ def test_get_mask_multiple(
] ]
for t_func in mask_funcs: for t_func in mask_funcs:
mask_imgs.append(_available_masks[t_func]["func"](target_img)) # 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 = [ mask_imgs = [
resample_to_img( resample_to_img(
@ -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", "TianxS1x3TxMNInonlinear2009cAsym",
"SUITxSUIT", ],
], target_data=element_data["VBM_GM"],
target_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,14 +82,14 @@ 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
---------- ----------
tmp_path : pathlib.Path tmp_path : pathlib.Path
The path to the test directory. The path to the test directory.
fractional : bool fractional : bool
Whether to compute fractional ALFF or not. Whether to compute fractional ALFF or not.
""" """
with PartlyCloudyTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:

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,42 +21,35 @@ 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(
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
cor_method_params={"empirical": True},
)
# Check correct output
assert marker.get_output_type("BOLD") == "matrix"
# Check empirical correlation method parameters # Fit-transform the data
efc = EdgeCentricFCParcels( edge_fc = marker.fit_transform(element_data)
parcellation="TianxS1x3TxMNInonlinear2009cAsym", edge_fc_bold = edge_fc["BOLD"]
cor_method_params={"empirical": True},
)
all_out = efc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
)
out = all_out["BOLD"] # For 16 ROIs we should get (16 * (16 -1) / 2) edges in the ETS
n_edges = int(16 * (16 - 1) / 2)
assert "data" in edge_fc_bold
assert "row_names" in edge_fc_bold
assert "col_names" in edge_fc_bold
assert edge_fc_bold["data"].shape[0] == n_edges
assert edge_fc_bold["data"].shape[1] == n_edges
assert len(set(edge_fc_bold["row_names"])) == n_edges
assert len(set(edge_fc_bold["col_names"])) == n_edges
# for 16 ROIs we should get (16 * (16 -1) / 2) edges in the ETS # Store
n_edges = int(16 * (16 - 1) / 2) storage = SQLiteFeatureStorage(
assert "data" in out uri=tmp_path / "test_edge_fc_parcels.sqlite", upsert="ignore"
assert "row_names" in out )
assert "col_names" in out marker.fit_transform(input=element_data, storage=storage)
assert out["data"].shape[0] == n_edges features = storage.list_features()
assert out["data"].shape[1] == n_edges assert any(
assert len(set(out["row_names"])) == n_edges x["name"] == "BOLD_EdgeCentricFCParcels" for x in features.values()
assert len(set(out["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)
features = storage.list_features()
assert any(
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,57 +21,41 @@ 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(
coords="DMNBuckner", radius=5.0, cor_method="correlation"
)
# Check correct output
assert marker.get_output_type("BOLD") == "matrix"
efc = EdgeCentricFCSpheres( # Fit-transform the data
coords="DMNBuckner", radius=5.0, cor_method="correlation" edge_fc = marker.fit_transform(element_data)
) edge_fc_bold = edge_fc["BOLD"]
all_out = efc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
)
out = all_out["BOLD"] # There are six DMNBuckner coordinates, so
# for 6 ROIs we should get (6 * (6 -1) / 2) edges in the ETS
n_edges = int(6 * (6 - 1) / 2)
assert "data" in edge_fc_bold
assert "row_names" in edge_fc_bold
assert "col_names" in edge_fc_bold
assert edge_fc_bold["data"].shape == (n_edges, n_edges)
assert len(set(edge_fc_bold["row_names"])) == n_edges
assert len(set(edge_fc_bold["col_names"])) == n_edges
# There are six DMNBuckner coordinates, so # Check empirical correlation method parameters
# for 6 ROIs we should get (6 * (6 -1) / 2) edges in the ETS marker = EdgeCentricFCSpheres(
n_edges = int(6 * (6 - 1) / 2) coords="DMNBuckner",
assert "data" in out radius=5.0,
assert "row_names" in out cor_method="correlation",
assert "col_names" in out cor_method_params={"empirical": True},
assert out["data"].shape[0] == n_edges )
assert out["data"].shape[1] == n_edges # Store
assert len(set(out["row_names"])) == n_edges storage = SQLiteFeatureStorage(
assert len(set(out["col_names"])) == n_edges uri=tmp_path / "test_edge_fc_spheres.sqlite", upsert="ignore"
)
# check correct output marker.fit_transform(input=element_data, storage=storage)
assert efc.get_output_type("BOLD") == "matrix" features = storage.list_features()
assert any(
# Check empirical correlation method parameters x["name"] == "BOLD_EdgeCentricFCSpheres" for x in features.values()
efc = EdgeCentricFCSpheres( )
coords="DMNBuckner",
radius=5.0,
cor_method="correlation",
cor_method_params={"empirical": True},
)
meta = {
"element": {"subject": "sub001"},
"dependencies": {"nilearn"},
}
all_out = efc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
)
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()
assert any(
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,74 +29,60 @@ 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"
)
# Check correct output
assert marker.get_output_type("BOLD") == "matrix"
fc = FunctionalConnectivityParcels(parcellation="Schaefer100x7") # Fit-transform the data
all_out = fc.fit_transform( fc = marker.fit_transform(element_data)
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} fc_bold = fc["BOLD"]
)
out = all_out["BOLD"] assert "data" in fc_bold
assert "row_names" in fc_bold
assert "col_names" in fc_bold
assert fc_bold["data"].shape == (16, 16)
assert len(set(fc_bold["row_names"])) == 16
assert len(set(fc_bold["col_names"])) == 16
assert "data" in out # Compare with nilearn
assert "row_names" in out # Load testing parcellation for the target data
assert "col_names" in out testing_parcellation, _ = get_parcellation(
assert out["data"].shape[0] == 100 parcellation=["TianxS1x3TxMNInonlinear2009cAsym"],
assert out["data"].shape[1] == 100 target_data=element_data["BOLD"],
assert len(set(out["row_names"])) == 100 )
assert len(set(out["col_names"])) == 100 # Extract timeseries
nifti_labels_masker = NiftiLabelsMasker(
labels_img=testing_parcellation, standardize=False
)
extracted_timeseries = nifti_labels_masker.fit_transform(
element_data["BOLD"]["data"]
)
# Compute the connectivity measure
connectivity_measure = ConnectivityMeasure(
kind="covariance"
).fit_transform([extracted_timeseries])[0]
# get the timeseries using pa # Check that FC are almost equal
pa = ParcelAggregation( assert_array_almost_equal(
parcellation="Schaefer100x7", method="mean", on="BOLD" connectivity_measure, fc_bold["data"], decimal=3
) )
meta = {
"element": {"subject": "sub001"},
"dependencies": {"nilearn"},
}
ts = pa.compute({"data": fmri_img, "meta": meta, "space": "MNI"})
# compare with nilearn # Check empirical correlation method parameters
# Get the testing parcellation (for nilearn) marker = FunctionalConnectivityParcels(
parcellation = datasets.fetch_atlas_schaefer_2018( parcellation="TianxS1x3TxMNInonlinear2009cAsym",
n_rois=100, yeo_networks=7, resolution_mm=2 cor_method_params={"empirical": True},
) )
masker = NiftiLabelsMasker( # Store
labels_img=parcellation["maps"], standardize=False storage = SQLiteFeatureStorage(
) uri=tmp_path / "test_fc_parcels.sqlite", upsert="ignore"
ts_ni = masker.fit_transform(fmri_img) )
marker.fit_transform(input=element_data, storage=storage)
# check the TS are almost equal features = storage.list_features()
assert_array_equal(ts_ni, ts["data"]) assert any(
x["name"] == "BOLD_FunctionalConnectivityParcels"
# Check that FC are almost equal for x in features.values()
cm = ConnectivityMeasure(kind="covariance") )
out_ni = cm.fit_transform([ts_ni])[0]
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
fc = FunctionalConnectivityParcels(
parcellation="Schaefer100x7", cor_method_params={"empirical": True}
)
all_out = fc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": meta, "space": "MNI"}}
)
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()
assert any(
x["name"] == "BOLD_FunctionalConnectivityParcels"
for x in features.values()
)

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,56 +32,57 @@ 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(
coords="DMNBuckner", radius=5.0, cor_method="correlation"
)
# Check correct output
assert marker.get_output_type("BOLD") == "matrix"
fc = FunctionalConnectivitySpheres( # Fit-transform the data
coords="DMNBuckner", radius=5.0, cor_method="correlation" fc = marker.fit_transform(element_data)
) fc_bold = fc["BOLD"]
all_out = fc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
)
out = all_out["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
assert "data" in out # Compare with nilearn
assert "row_names" in out # Load testing coordinates for the target data
assert "col_names" in out testing_coords, _ = get_coordinates(
assert out["data"].shape[0] == 6 coords="DMNBuckner", target_data=element_data["BOLD"]
assert out["data"].shape[1] == 6 )
assert len(set(out["row_names"])) == 6 # Extract timeseries
assert len(set(out["col_names"])) == 6 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]
# get the timeseries using sa # Check that FC are almost equal
sa = SphereAggregation( assert_array_almost_equal(
coords="DMNBuckner", radius=5.0, method="mean", on="BOLD" connectivity_measure, fc_bold["data"], decimal=3
) )
ts = sa.compute({"data": fmri_img, "meta": {}, "space": "MNI"})
# Check that FC are almost equal when using nileran # Store
cm = ConnectivityMeasure(kind="correlation") storage = SQLiteFeatureStorage(
out_ni = cm.fit_transform([ts["data"]])[0] uri=tmp_path / "test_fc_spheres.sqlite", upsert="ignore"
assert_array_almost_equal(out_ni, out["data"], decimal=3) )
marker.fit_transform(input=element_data, storage=storage)
# check correct output features = storage.list_features()
assert fc.get_output_type("BOLD") == "matrix" assert any(
x["name"] == "BOLD_FunctionalConnectivitySpheres"
uri = tmp_path / "test_fc_parcel.sqlite" for x in features.values()
# 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()
assert any(
x["name"] == "BOLD_FunctionalConnectivitySpheres"
for x in features.values()
)
def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None: def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
@ -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:
element_data = DefaultDataReader().fit_transform(dg["sub001"])
marker = FunctionalConnectivitySpheres(
coords="DMNBuckner",
radius=5.0,
cor_method="correlation",
cor_method_params={"empirical": True},
)
# Check correct output
assert marker.get_output_type("BOLD") == "matrix"
# get a dataset # Fit-transform the data
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") fc = marker.fit_transform(element_data)
fmri_img = image.concat_imgs(ni_data.func) # type: ignore fc_bold = fc["BOLD"]
fc = FunctionalConnectivitySpheres( assert "data" in fc_bold
coords="DMNBuckner", assert "row_names" in fc_bold
radius=5.0, assert "col_names" in fc_bold
cor_method="correlation", assert fc_bold["data"].shape == (6, 6)
cor_method_params={"empirical": True}, assert len(set(fc_bold["row_names"])) == 6
) assert len(set(fc_bold["col_names"])) == 6
all_out = fc.fit_transform(
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}}
)
out = all_out["BOLD"] # 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(
cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
).fit_transform([extracted_timeseries])[0]
assert "data" in out # Check that FC are almost equal
assert "row_names" in out assert_array_almost_equal(
assert "col_names" in out connectivity_measure, fc_bold["data"], decimal=3
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"})
# Check that FC are almost equal when using nileran
cm = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
)
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,35 +41,15 @@ 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( # Store
{"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} storage = HDF5FeatureStorage(tmp_path / "test_tsnr_parcels.hdf5")
) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features()
out = all_out["BOLD"] assert any(
x["name"] == "BOLD_TemporalSNRParcels" for x in features.values()
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()
assert any(
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,40 +41,16 @@ 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"}} features = storage.list_features()
) assert any(
x["name"] == "BOLD_TemporalSNRSpheres" for x in features.values()
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()
assert any(
x["name"] == "BOLD_TemporalSNRSpheres" for x in features.values()
)
def test_TemporalSNRSpheres_error() -> None: def test_TemporalSNRSpheres_error() -> None:

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())

File diff suppressed because it is too large Load diff

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"])
# Create SphereAggregation object
marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
)
sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][
"data"
]
# Get the oasis VBM data # Compare with nilearn
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) # Load testing coordinates
vbm = oasis_dataset.gray_matter_maps[0] testing_coords, _ = get_coordinates(
img = nib.load(vbm) 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"]
)
# Create NiftSpheresMasker assert sphere_agg_vbm_gm_data.ndim == 2
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) assert_array_equal(
auto4d = nifti_masker.fit_transform(img) nifti_spheres_masked_vbm_gm.shape, sphere_agg_vbm_gm_data.shape
)
# Create SphereAggregation object assert_array_equal(nifti_spheres_masked_vbm_gm, sphere_agg_vbm_gm_data)
marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
)
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}}
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d)
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"])
# Create SphereAggregation object
marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
)
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
# Get the SPM auditory data # Compare with nilearn
subject_data = datasets.fetch_spm_auditory() # Load testing coordinates
fmri_img = concat_imgs(subject_data.func) # type: ignore 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"]
)
# Create NiftSpheresMasker assert sphere_agg_bold_data.ndim == 2
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) assert_array_equal(
auto4d = nifti_masker.fit_transform(fmri_img) nifti_spheres_masked_bold.shape, sphere_agg_bold_data.shape
)
# Create SphereAggregation object assert_array_equal(nifti_spheres_masked_bold, sphere_agg_bold_data)
marker = SphereAggregation(coords=COORDS, method="mean", radius=RADIUS)
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_SphereAggregation_storage(tmp_path: Path) -> None: def test_SphereAggregation_storage(tmp_path: Path) -> None:
@ -92,124 +108,143 @@ 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"
)
marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
)
marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features()
assert any(
x["name"] == "VBM_GM_SphereAggregation" for x in features.values()
)
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") # Store 4D
meta = { with SPMAuditoryTestingDataGrabber() as dg:
"element": {"subject": "sub-01", "session": "ses-01"}, element_data = DefaultDataReader().fit_transform(dg["sub001"])
"dependencies": {"nilearn", "nibabel"}, storage = SQLiteFeatureStorage(
} uri=tmp_path / "test_sphere_storage_4D.sqlite", upsert="ignore"
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="BOLD"
) )
marker.fit_transform(input=element_data, storage=storage)
marker.fit_transform(input, storage=storage) features = storage.list_features()
assert any(
features: Dict = typing.cast(Dict, storage.list_features()) x["name"] == "BOLD_SphereAggregation" for x in features.values()
assert any( )
x["name"] == "VBM_GM_SphereAggregation" for x in features.values()
)
meta = {
"element": {"subject": "sub-01", "session": "ses-01"},
"dependencies": {"nilearn", "nibabel"},
}
# Get the SPM auditory data
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(
coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
)
marker.fit_transform(input, storage=storage)
features: Dict = typing.cast(Dict, storage.list_features())
assert any(
x["name"] == "BOLD_SphereAggregation" for x in features.values()
)
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"])
# Create SphereAggregation object
marker = SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
on="VBM_GM",
masks="compute_brain_mask",
)
sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][
"data"
]
# Get one mask # Compare with nilearn
mask_img, _, _ = load_mask("GM_prob0.2") # Load testing coordinates
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"]
)
# Get the oasis VBM data assert sphere_agg_vbm_gm_data.ndim == 2
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) assert_array_equal(
vbm = oasis_dataset.gray_matter_maps[0] nifti_spheres_masked_vbm_agg.shape,
img = nib.load(vbm) nifti_spheres_masked_vbm_agg.shape,
)
# Create NiftSpheresMasker assert_array_equal(
nifti_masker = NiftiSpheresMasker( nifti_spheres_masked_vbm_agg, nifti_spheres_masked_vbm_agg
seeds=coordinates, radius=RADIUS, mask_img=mask_img )
)
auto4d = nifti_masker.fit_transform(img)
# Create SphereAggregation object
marker = SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
on="VBM_GM",
masks="GM_prob0.2",
)
input = {"VBM_GM": {"data": img, "meta": {}, "space": "MNI"}}
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d)
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"])
# Create SphereAggregation object
marker = SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
time_method="mean",
on="BOLD",
)
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
# Get the SPM auditory data # Compare with nilearn
subject_data = datasets.fetch_spm_auditory() # Load testing coordinates
fmri_img = concat_imgs(subject_data.func) # type: ignore 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)
# Create NiftSpheresMasker assert sphere_agg_bold_data.ndim == 1
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS) assert_array_equal(
auto4d = nifti_masker.fit_transform(fmri_img) nifti_spheres_masked_bold_mean.shape, sphere_agg_bold_data.shape
auto_mean = auto4d.mean(axis=0) )
assert_array_equal(
nifti_spheres_masked_bold_mean, sphere_agg_bold_data
)
# Create SphereAggregation object # Test picking first time point
marker = SphereAggregation( nifti_spheres_masked_bold_pick_0 = nifti_spheres_masked_bold[:1, :]
coords=COORDS, method="mean", radius=RADIUS, time_method="mean" marker = SphereAggregation(
) coords=COORDS,
input = {"BOLD": {"data": fmri_img, "meta": {}, "space": "MNI"}} method="mean",
jun_values4d = marker.fit_transform(input)["BOLD"]["data"] radius=RADIUS,
time_method="select",
time_method_params={"pick": [0]},
on="BOLD",
)
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data"
]
assert jun_values4d.ndim == 1 assert sphere_agg_bold_data.ndim == 2
assert_array_equal(auto_mean.shape, jun_values4d.shape) assert_array_equal(
assert_array_equal(auto_mean, jun_values4d) 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
)
auto_pick_0 = auto4d[:1, :]
marker = SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
time_method="select",
time_method_params={"pick": [0]},
)
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)