[ENH]: Enable Markers to extract multiple features #349

Merged
synchon merged 18 commits from refactor/marker-multi-feature-output into main 2024-07-18 10:24:13 +00:00
53 changed files with 1026 additions and 1088 deletions

View file

@ -0,0 +1 @@
``fractional`` parameter for ``ALFFBase``, :class:`.ALFFParcels` and :class:`.ALFFSpheres` have been removed in favour of returning both ALFF and fALFF by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Enable Markers to output multiple features by `Synchon Mandal`_

View file

@ -5,24 +5,17 @@
Creating Markers Creating Markers
================ ================
Computing a marker (a.k.a. *feature*) is the main goal of ``junifer``. While we Computing a marker (a.k.a. *feature(s)*) is the main goal of ``junifer``. While
aim to provide as many Markers as possible, it might be the case that the Marker we aim to provide as many Markers as possible, it might be the case that the
you are looking for is not available. In this case, you can create your own Marker Marker you are looking for is not available. In this case, you can create your
by following this tutorial. own Marker by following this tutorial.
Most of the functionality of a ``junifer`` Marker has been taken care by the Most of the functionality of a ``junifer`` Marker has been taken care by the
:class:`.BaseMarker` class. Thus, only a few methods are required: :class:`.BaseMarker` class. Thus, only a few methods and class attributes are
required:
#. ``get_valid_inputs``: The method to obtain the list of valid inputs for the
Marker. This is used to check that the inputs provided by the user are
valid. This method should return a list of strings, representing
:ref:`data types <data_types>`.
#. ``get_output_type``: The method to obtain the output type of the Marker.
This is used to check that the output of the Marker is compatible with the
storage. This method should return a string, representing
:ref:`storage types <storage_types>`.
#. ``compute``: The method that given the data, computes the Marker.
#. ``__init__``: The initialisation method, where the Marker is configured. #. ``__init__``: The initialisation method, where the Marker is configured.
#. ``compute``: The method that given the data, computes the Marker.
As an example, we will develop a ``ParcelMean`` Marker, a Marker that first As an example, we will develop a ``ParcelMean`` Marker, a Marker that first
applies a parcellation and then computes the mean of the data in each parcel. applies a parcellation and then computes the mean of the data in each parcel.
@ -35,24 +28,26 @@ Step 1: Configure input and output
This step is quite simple: we need to define the input and output of the Marker. This step is quite simple: we need to define the input and output of the Marker.
Based on the current :ref:`data types <data_types>`, we can have ``BOLD``, Based on the current :ref:`data types <data_types>`, we can have ``BOLD``,
``VBM_WM`` and ``VBM_GM`` as valid inputs. ``VBM_WM`` and ``VBM_GM`` as valid inputs. The output of the Marker depends on
the input. For ``BOLD``, it will be ``timeseries``, while for the rest of the
inputs, it will be ``vector``. Thus, we have a class attribute like so:
.. code-block:: python .. code-block:: python
def get_valid_inputs(self) -> list[str]: # NOTE: data type -> feature -> storage type
return ["BOLD", "VBM_WM", "VBM_GM"] # You can have multiple features for one data type,
# each feature having same or different storage type
The output of the Marker depends on the input. For ``BOLD``, it will be _MARKER_INOUT_MAPPINGS = {
``timeseries``, while for the rest of the inputs, it will be ``vector``. Thus, "BOLD": {
we can define the output as: "parcel_mean": "timeseries",
},
.. code-block:: python "VBM_WM": {
"parcel_mean": "vector",
def get_output_type(self, input_type: str) -> str: },
if input_type == "BOLD": "VBM_GM": {
return "timeseries" "parcel_mean": "vector",
else: },
return "vector" }
.. _extending_markers_init: .. _extending_markers_init:
@ -119,7 +114,8 @@ arguments:
Following the example, we will compute the mean of the data in each parcel using Following the example, we will compute the mean of the data in each parcel using
:class:`nilearn.maskers.NiftiLabelsMasker`. Importantly, the output of the :class:`nilearn.maskers.NiftiLabelsMasker`. Importantly, the output of the
compute function must be a dictionary. This dictionary will later be passed onto compute function must be a dictionary. This dictionary will later be passed onto
the ``store`` method. the ``store`` method. The dictionary's first level of keys would the feature name
and the values would be a dictionary of storage type specific key-value pairs.
.. hint:: .. hint::
@ -162,11 +158,13 @@ the ``store`` method.
# mask the data # mask the data
out_values = masker.fit_transform([data]) out_values = masker.fit_transform([data])
# Create the output dictionary # Create and return the output dictionary
out = {"data": out_values, "col_names": t_labels} return {
"parcel_mean": {
return out "data": out_values,
"col_names": t_labels,
},
}
.. _extending_markers_finalize: .. _extending_markers_finalize:
@ -193,11 +191,11 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
.. code-block:: python .. code-block:: python
from typing import Any from typing import Any, ClassVar
from junifer.api.decorators import register_marker from junifer.api.decorators import register_marker
from junifer.data import get_parcellation from junifer.data import get_parcellation
from junifer.markers.base import BaseMarker from junifer.markers import BaseMarker
from nilearn.maskers import NiftiLabelsMasker from nilearn.maskers import NiftiLabelsMasker
@ -206,6 +204,18 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
_DEPENDENCIES = {"nilearn", "numpy"} _DEPENDENCIES = {"nilearn", "numpy"}
_MARKER_INOUT_MAPPINGS: ClassVar[dict[str, dict[str, str]]] = {
"BOLD": {
"parcel_mean": "timeseries",
},
"VBM_WM": {
"parcel_mean": "vector",
},
"VBM_GM": {
"parcel_mean": "vector",
},
}
def __init__( def __init__(
self, self,
parcellation: str, parcellation: str,
@ -215,15 +225,6 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
self.parcellation = parcellation self.parcellation = parcellation
super().__init__(on=on, name=name) super().__init__(on=on, name=name)
def get_valid_inputs(self) -> list[str]:
return ["BOLD", "VBM_WM", "VBM_GM"]
def get_output_type(self, input_type: str) -> str:
if input_type == "BOLD":
return "timeseries"
else:
return "vector"
def compute( def compute(
self, self,
input: dict[str, Any], input: dict[str, Any],
@ -250,11 +251,13 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
# mask the data # mask the data
out_values = masker.fit_transform([data]) out_values = masker.fit_transform([data])
# Create the output dictionary # Create and return the output dictionary
out = {"data": out_values, "col_names": t_labels} return {
"parcel_mean": {
return out "data": out_values,
"col_names": t_labels,
},
}
.. _extending_markers_template: .. _extending_markers_template:
@ -269,22 +272,16 @@ Template for a custom Marker
@register_marker @register_marker
class TemplateMarker(BaseMarker): class TemplateMarker(BaseMarker):
# TODO: add the dependencies
_DEPENDENCIES = {}
# TODO: add the input-output mappings
_MARKER_INOUT_MAPPINGS = {}
def __init__(self, on=None, name=None): def __init__(self, on=None, name=None):
# TODO: add marker-specific parameters # TODO: add marker-specific parameters
super().__init__(on=on, name=name) super().__init__(on=on, name=name)
def get_valid_inputs(self):
# TODO: Complete with the valid inputs
valid = []
return valid
def get_output_type(self, input_type):
# TODO: Return the valid output type for each input type
pass
def compute(self, input, extra_input): def compute(self, input, extra_input):
# TODO: compute the marker and create the output dictionary # TODO: compute the marker and create the output dictionary
# Create the output dictionary
out = {"data": None, "col_names": None}
return out

View file

@ -26,7 +26,5 @@ on them outside the context as long as the actual data is in the memory and the
Python runtime has not garbage-collected it. Python runtime has not garbage-collected it.
If you are interested in using already provided Markers, please go to If you are interested in using already provided Markers, please go to
:doc:`../builtin`. And, if you want to implement your own Marker, you need to :doc:`../builtin`. And, if you want to implement your own Marker, please check
provide concrete implementation of :class:`.BaseMarker`. Specifically, you out :doc:`../extending/marker`.
need to override ``get_valid_inputs``, ``get_output_type`` and ``compute``
methods.

View file

@ -15,13 +15,13 @@ from junifer.datagrabber import DataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
"commit": "522dfb203afcd2cd55799bf347f9b211919a7338", "commit": "b87897cbe51bf0ee5514becaa5c7dd76491db5ad",
"id": "fec92475-d9c0-4409-92ba-f041b6a12c40", "id": "8fddff30-6993-420a-9d1e-b5b028c59468",
}, },
"example_bids_ses": { "example_bids_ses": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses",
"commit": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", "commit": "6b163aa98af76a9eac0272273c27e14127850181",
"id": "c83500d0-532f-45be-baf1-0dab703bdc2a", "id": "715c17cf-a1b9-42d6-9af8-9f74c1a4a724",
}, },
} }

View file

@ -15,13 +15,13 @@ from junifer.datagrabber import PatternDataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
"commit": "522dfb203afcd2cd55799bf347f9b211919a7338", "commit": "b87897cbe51bf0ee5514becaa5c7dd76491db5ad",
"id": "fec92475-d9c0-4409-92ba-f041b6a12c40", "id": "8fddff30-6993-420a-9d1e-b5b028c59468",
}, },
"example_bids_ses": { "example_bids_ses": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses",
"commit": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", "commit": "6b163aa98af76a9eac0272273c27e14127850181",
"id": "c83500d0-532f-45be-baf1-0dab703bdc2a", "id": "715c17cf-a1b9-42d6-9af8-9f74c1a4a724",
}, },
} }

View file

@ -5,6 +5,7 @@
# License: AGPL # License: AGPL
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
@ -35,6 +36,8 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
Raises Raises
------ ------
AttributeError
If the marker does not have `_MARKER_INOUT_MAPPINGS` attribute.
ValueError ValueError
If required input data type(s) is(are) not found. If required input data type(s) is(are) not found.
@ -45,6 +48,12 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
on: Optional[Union[List[str], str]] = None, on: Optional[Union[List[str], str]] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
# Check for missing mapping attribute
if not hasattr(self, "_MARKER_INOUT_MAPPINGS"):
raise_error(
msg=("Missing `_MARKER_INOUT_MAPPINGS` for the marker"),
klass=AttributeError,
)
# Use all data types if not provided # Use all data types if not provided
if on is None: if on is None:
on = self.get_valid_inputs() on = self.get_valid_inputs()
@ -88,7 +97,6 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
) )
return [x for x in self._on if x in input] return [x for x in self._on if x in input]
@abstractmethod
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input. """Get valid data types for input.
@ -98,30 +106,25 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
The list of data types that can be used as input for this marker. The list of data types that can be used as input for this marker.
""" """
raise_error( return list(self._MARKER_INOUT_MAPPINGS.keys())
msg="Concrete classes need to implement get_valid_inputs().",
klass=NotImplementedError,
)
@abstractmethod def get_output_type(self, input_type: str, output_feature: str) -> str:
def get_output_type(self, input_type: str) -> str:
"""Get output type. """Get output type.
Parameters Parameters
---------- ----------
input_type : str input_type : str
The data type input to the marker. The data type input to the marker.
output_feature : str
The feature output of the marker.
Returns Returns
------- -------
str str
The storage type output by the marker. The storage type output of the marker.
""" """
raise_error( return self._MARKER_INOUT_MAPPINGS[input_type][output_feature]
msg="Concrete classes need to implement get_output_type().",
klass=NotImplementedError,
)
@abstractmethod @abstractmethod
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict: def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
@ -154,6 +157,7 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
def store( def store(
self, self,
type_: str, type_: str,
feature: str,
out: Dict[str, Any], out: Dict[str, Any],
storage: "BaseFeatureStorage", storage: "BaseFeatureStorage",
) -> None: ) -> None:
@ -163,13 +167,15 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
---------- ----------
type_ : str type_ : str
The data type to store. The data type to store.
feature : str
The feature to store.
out : dict out : dict
The computed result as a dictionary to store. The computed result as a dictionary to store.
storage : storage-like storage : storage-like
The storage class, for example, SQLiteFeatureStorage. The storage class, for example, SQLiteFeatureStorage.
""" """
output_type_ = self.get_output_type(type_) output_type_ = self.get_output_type(type_, feature)
logger.debug(f"Storing {output_type_} in {storage}") logger.debug(f"Storing {output_type_} in {storage}")
storage.store(kind=output_type_, **out) storage.store(kind=output_type_, **out)
@ -213,15 +219,35 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
t_meta["type"] = type_ t_meta["type"] = type_
# Compute marker # Compute marker
t_out = self.compute(input=t_input, extra_input=extra_input) t_out = self.compute(input=t_input, extra_input=extra_input)
t_out["meta"] = t_meta # Initialize empty dictionary if no storage object is provided
# Update metadata for step if storage is None:
self.update_meta(t_out, "marker") out[type_] = {}
# Check storage # Store individual features
if storage is not None: for feature_name, feature_data in t_out.items():
logger.info(f"Storing in {storage}") # Make deep copy of the feature data for manipulation
self.store(type_=type_, out=t_out, storage=storage) feature_data_copy = deepcopy(feature_data)
else: # Make deep copy of metadata and add to feature data
logger.info("No storage specified, returning dictionary") feature_data_copy["meta"] = deepcopy(t_meta)
out[type_] = t_out # Update metadata for the feature,
# feature data is not manipulated, only meta
self.update_meta(feature_data_copy, "marker")
# Update marker feature's metadata name
feature_data_copy["meta"]["marker"][
"name"
] += f"_{feature_name}"
if storage is not None:
logger.info(f"Storing in {storage}")
self.store(
type_=type_,
feature=feature_name,
out=feature_data_copy,
storage=storage,
)
else:
logger.info(
"No storage specified, returning dictionary"
)
out[type_][feature_name] = feature_data_copy
return out return out

View file

@ -3,21 +3,9 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import sys
if sys.version_info < (3, 11): # pragma: no cover
from importlib_metadata import packages_distributions
else:
from importlib.metadata import packages_distributions
import uuid import uuid
from copy import deepcopy
from importlib.util import find_spec
from itertools import chain
from pathlib import Path from pathlib import Path
from typing import ( from typing import (
TYPE_CHECKING,
Any, Any,
ClassVar, ClassVar,
Dict, Dict,
@ -37,15 +25,10 @@ from ..external.BrainPrint.brainprint.brainprint import (
) )
from ..external.BrainPrint.brainprint.surfaces import surf_to_vtk from ..external.BrainPrint.brainprint.surfaces import surf_to_vtk
from ..pipeline import WorkDirManager from ..pipeline import WorkDirManager
from ..pipeline.utils import check_ext_dependencies from ..utils import logger, run_ext_cmd
from ..utils import logger, raise_error, run_ext_cmd
from .base import BaseMarker from .base import BaseMarker
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
__all__ = ["BrainPrint"] __all__ = ["BrainPrint"]
@ -99,6 +82,15 @@ class BrainPrint(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"lapy", "numpy"} _DEPENDENCIES: ClassVar[Set[str]] = {"lapy", "numpy"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"FreeSurfer": {
"eigenvalues": "scalar_table",
"areas": "vector",
"volumes": "vector",
"distances": "vector",
}
}
def __init__( def __init__(
self, self,
num: int = 50, num: int = 50,
@ -121,117 +113,6 @@ class BrainPrint(BaseMarker):
self.use_cholmod = use_cholmod self.use_cholmod = use_cholmod
super().__init__(name=name, on="FreeSurfer") super().__init__(name=name, on="FreeSurfer")
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["FreeSurfer"]
# TODO: kept for making this class concrete; should be removed later
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "vector"
# TODO: overridden to allow multiple outputs from single data type; should
# be removed later
def validate(self, input: List[str]) -> List[str]:
"""Validate the the pipeline step.
Parameters
----------
input : list of str
The input to the pipeline step.
Returns
-------
list of str
The output of the pipeline step.
"""
def _check_dependencies(obj) -> None:
"""Check obj._DEPENDENCIES.
Parameters
----------
obj : object
Object to check _DEPENDENCIES of.
Raises
------
ImportError
If the pipeline step object is missing dependencies required
for its working.
"""
# Check if _DEPENDENCIES attribute is found;
# (markers and preprocessors will have them but not datareaders
# as of now)
dependencies_not_found = []
if hasattr(obj, "_DEPENDENCIES"):
# Check if dependencies are importable
for dependency in obj._DEPENDENCIES:
# First perform an easy check
if find_spec(dependency) is None:
# Then check mapped names
if dependency not in list(
chain.from_iterable(
packages_distributions().values()
)
):
dependencies_not_found.append(dependency)
# Raise error if any dependency is not found
if dependencies_not_found:
raise_error(
msg=(
f"{dependencies_not_found} are not installed but are "
f"required for using {obj.__class__.__name__}."
),
klass=ImportError,
)
def _check_ext_dependencies(obj) -> None:
"""Check obj._EXT_DEPENDENCIES.
Parameters
----------
obj : object
Object to check _EXT_DEPENDENCIES of.
"""
# Check if _EXT_DEPENDENCIES attribute is found;
# (some markers and preprocessors might have them)
if hasattr(obj, "_EXT_DEPENDENCIES"):
for dependency in obj._EXT_DEPENDENCIES:
check_ext_dependencies(**dependency)
# Check dependencies
_check_dependencies(self)
# Check external dependencies
# _check_ext_dependencies(self)
# Validate input
_ = self.validate_input(input=input)
# Validate output type
outputs = ["scalar_table", "vector"]
return outputs
def _create_aseg_surface( def _create_aseg_surface(
self, self,
aseg_path: Path, aseg_path: Path,
@ -426,6 +307,27 @@ class BrainPrint(BaseMarker):
), ),
} }
def _fix_nan(
self,
input_data: List[Union[float, str, npt.ArrayLike]],
) -> np.ndarray:
"""Convert BrainPrint output with string NaN to ``numpy.nan``.
Parameters
----------
input_data : list of str, float or numpy.ndarray-like
The data to convert.
Returns
-------
np.ndarray
The converted data as ``numpy.ndarray``.
"""
arr = np.asarray(input_data)
arr[arr == "NaN"] = np.nan
return arr.astype(np.float64)
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -443,16 +345,32 @@ class BrainPrint(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The dictionary has the following The computed result as dictionary. This will be either returned
keys: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``eigenvalues`` : dict of surface labels (str) and eigenvalues * ``eigenvalues`` : dictionary with the following keys:
(``np.ndarray``)
* ``eigenvectors`` : dict of surface labels (str) and eigenvectors - ``data`` : eigenvalues as ``np.ndarray``
(``np.ndarray``) if ``keep_eigenvectors=True`` - ``col_names`` : surface labels as list of str
else None - ``row_names`` : eigenvalue count labels as list of str
* ``distances`` : dict of ``{left_label}_{right_label}`` (str) and - ``row_header_col_name`` : "eigenvalue"
distance (float) if ``asymmetry=True`` else None ()
* ``areas`` : dictionary with the following keys:
- ``data`` : areas as ``np.ndarray``
- ``col_names`` : surface labels as list of str
* ``volumes`` : dictionary with the following keys:
- ``data`` : volumes as ``np.ndarray``
- ``col_names`` : surface labels as list of str
* ``distances`` : dictionary with the following keys
if ``asymmetry = True``:
- ``data`` : distances as ``np.ndarray``
- ``col_names`` : surface labels as list of str
References References
---------- ----------
@ -539,130 +457,3 @@ class BrainPrint(BaseMarker):
"col_names": list(distances.keys()), "col_names": list(distances.keys()),
} }
return output return output
def _fix_nan(
self,
input_data: List[Union[float, str, npt.ArrayLike]],
) -> np.ndarray:
"""Convert BrainPrint output with string NaN to ``numpy.nan``.
Parameters
----------
input_data : list of str, float or numpy.ndarray-like
The data to convert.
Returns
-------
np.ndarray
The converted data as ``numpy.ndarray``.
"""
arr = np.asarray(input_data)
arr[arr == "NaN"] = np.nan
return arr.astype(np.float64)
# TODO: overridden to allow storing multiple outputs from single input;
# should be removed later
def store(
self,
type_: str,
feature: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
type_ : str
The data type to store.
feature : {"eigenvalues", "distances", "areas", "volumes"}
The feature name to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
Raises
------
ValueError
If ``feature`` is invalid.
"""
if feature == "eigenvalues":
output_type = "scalar_table"
elif feature in ["distances", "areas", "volumes"]:
output_type = "vector"
else:
raise_error(f"Unknown feature: {feature}")
logger.debug(f"Storing {output_type} in {storage}")
storage.store(kind=output_type, **out)
# TODO: overridden to allow storing multiple outputs from single input;
# should be removed later
def _fit_transform(
self,
input: Dict[str, Dict],
storage: Optional["BaseFeatureStorage"] = None,
) -> Dict:
"""Fit and transform.
Parameters
----------
input : dict
The Junifer Data object.
storage : storage-like, optional
The storage class, for example, SQLiteFeatureStorage.
Returns
-------
dict
The processed output as a dictionary. If `storage` is provided,
empty dictionary is returned.
"""
out = {}
for type_ in self._on:
if type_ in input.keys():
logger.info(f"Computing {type_}")
t_input = input[type_]
extra_input = input.copy()
extra_input.pop(type_)
t_meta = t_input["meta"].copy()
t_meta["type"] = type_
# Returns multiple features
t_out = self.compute(input=t_input, extra_input=extra_input)
if storage is None:
out[type_] = {}
for feature_name, feature_data in t_out.items():
# Make deep copy of the feature data for manipulation
feature_data_copy = deepcopy(feature_data)
# Make deep copy of metadata and add to feature data
feature_data_copy["meta"] = deepcopy(t_meta)
# Update metadata for the feature,
# feature data is not manipulated, only meta
self.update_meta(feature_data_copy, "marker")
# Update marker feature's metadata name
feature_data_copy["meta"]["marker"][
"name"
] += f"_{feature_name}"
if storage is not None:
logger.info(f"Storing in {storage}")
self.store(
type_=type_,
feature=feature_name,
out=feature_data_copy,
storage=storage,
)
else:
logger.info(
"No storage specified, returning dictionary"
)
out[type_][feature_name] = feature_data_copy
return out

View file

@ -53,6 +53,12 @@ class ComplexityBase(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "neurokit2"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "neurokit2"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"complexity": "vector",
},
}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
@ -78,33 +84,6 @@ class ComplexityBase(BaseMarker):
klass=NotImplementedError, klass=NotImplementedError,
) )
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "vector"
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -124,29 +103,30 @@ class ComplexityBase(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The following keys will be The computed result as dictionary. This will be either returned
included in the dictionary: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : ROI-wise complexity measures as ``numpy.ndarray`` * ``complexity`` : dictionary with the following keys:
* ``col_names`` : ROI labels for the complexity measures as list
- ``data`` : ROI-wise complexity measures as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
# Initialize a ParcelAggregation # Extract the 2D time series using ParcelAggregation
parcel_aggregation = ParcelAggregation( parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input=input, extra_input=extra_input)
# Extract the 2D time series using parcel aggregation
parcel_aggregation_map = parcel_aggregation.compute(
input=input, extra_input=extra_input
)
# Compute complexity measure # Compute complexity measure
parcel_aggregation_map["data"] = self.compute_complexity( return {
parcel_aggregation_map["data"] "complexity": {
) "data": self.compute_complexity(
parcel_aggregation["aggregation"]["data"]
return parcel_aggregation_map ),
"col_names": parcel_aggregation["aggregation"]["col_names"],
}
}

View file

@ -40,13 +40,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test HurstExponent get_output_type().""" """Test HurstExponent get_output_type()."""
marker = HurstExponent(parcellation=PARCELLATION) assert "vector" == HurstExponent(
assert marker.get_output_type("BOLD") == "vector" parcellation=PARCELLATION
).get_output_type(input_type="BOLD", output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -39,13 +39,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test MultiscaleEntropyAUC get_output_type().""" """Test MultiscaleEntropyAUC get_output_type()."""
marker = MultiscaleEntropyAUC(parcellation=PARCELLATION) assert "vector" == MultiscaleEntropyAUC(
assert marker.get_output_type("BOLD") == "vector" parcellation=PARCELLATION
).get_output_type(input_type="BOLD", output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -39,13 +39,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test PermEntropy get_output_type().""" """Test PermEntropy get_output_type()."""
marker = PermEntropy(parcellation=PARCELLATION) assert "vector" == PermEntropy(parcellation=PARCELLATION).get_output_type(
assert marker.get_output_type("BOLD") == "vector" input_type="BOLD", output_feature="complexity"
)
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -40,13 +40,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test RangeEntropy get_output_type().""" """Test RangeEntropy get_output_type()."""
marker = RangeEntropy(parcellation=PARCELLATION) assert "vector" == RangeEntropy(parcellation=PARCELLATION).get_output_type(
assert marker.get_output_type("BOLD") == "vector" input_type="BOLD", output_feature="complexity"
)
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -40,13 +40,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test RangeEntropyAUC get_output_type().""" """Test RangeEntropyAUC get_output_type()."""
marker = RangeEntropyAUC(parcellation=PARCELLATION) assert "vector" == RangeEntropyAUC(
assert marker.get_output_type("BOLD") == "vector" parcellation=PARCELLATION
).get_output_type(input_type="BOLD", output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -39,13 +39,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test SampleEntropy get_output_type().""" """Test SampleEntropy get_output_type()."""
marker = SampleEntropy(parcellation=PARCELLATION) assert "vector" == SampleEntropy(
assert marker.get_output_type("BOLD") == "vector" parcellation=PARCELLATION
).get_output_type(input_type="BOLD", output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -39,13 +39,14 @@ def test_compute() -> None:
# Compute the marker # Compute the marker
feature_map = marker.fit_transform(element_data) feature_map = marker.fit_transform(element_data)
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert feature_map["BOLD"]["data"].ndim == 2 assert feature_map["BOLD"]["complexity"]["data"].ndim == 2
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test WeightedPermEntropy get_output_type().""" """Test WeightedPermEntropy get_output_type()."""
marker = WeightedPermEntropy(parcellation=PARCELLATION) assert "vector" == WeightedPermEntropy(
assert marker.get_output_type("BOLD") == "vector" parcellation=PARCELLATION
).get_output_type(input_type="BOLD", output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -47,6 +47,12 @@ class RSSETSMarker(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"rss_ets": "timeseries",
},
}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
@ -61,33 +67,6 @@ class RSSETSMarker(BaseMarker):
self.masks = masks self.masks = masks
super().__init__(name=name) super().__init__(name=name)
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "timeseries"
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -109,8 +88,9 @@ class RSSETSMarker(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The dictionary has the following The computed result as dictionary. This will be either returned
keys: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``data`` : the actual computed values as a numpy.ndarray
* ``col_names`` : the column labels for the computed values as list * ``col_names`` : the column labels for the computed values as list
@ -124,20 +104,22 @@ class RSSETSMarker(BaseMarker):
""" """
logger.debug("Calculating root sum of squares of edgewise timeseries.") logger.debug("Calculating root sum of squares of edgewise timeseries.")
# Initialize a ParcelAggregation # Perform aggregation
parcel_aggregation = ParcelAggregation( aggregation = ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
) ).compute(input=input, extra_input=extra_input)
# Compute the parcel aggregation # Compute edgewise timeseries
out = parcel_aggregation.compute(input=input, extra_input=extra_input) edge_ts, _ = _ets(aggregation["aggregation"]["data"])
edge_ts, _ = _ets(out["data"]) # Compute the RSS of edgewise timeseries
# Compute the RSS rss = np.sum(edge_ts**2, 1) ** 0.5
out["data"] = np.sum(edge_ts**2, 1) ** 0.5
# Make it 2D return {
out["data"] = out["data"][:, np.newaxis] "rss_ets": {
# Set correct column label # Make it 2D
out["col_names"] = ["root_sum_of_squares_ets"] "data": rss[:, np.newaxis],
return out "col_names": ["root_sum_of_squares_ets"],
}
}

View file

@ -37,8 +37,6 @@ class ALFFBase(BaseMarker):
Parameters Parameters
---------- ----------
fractional : bool
Whether to compute fractional ALFF.
highpass : positive float highpass : positive float
Highpass cutoff frequency. Highpass cutoff frequency.
lowpass : positive float lowpass : positive float
@ -85,9 +83,15 @@ class ALFFBase(BaseMarker):
}, },
] ]
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"alff": "vector",
"falff": "vector",
},
}
def __init__( def __init__(
self, self,
fractional: bool,
highpass: float, highpass: float,
lowpass: float, lowpass: float,
using: str, using: str,
@ -110,45 +114,12 @@ class ALFFBase(BaseMarker):
) )
self.using = using self.using = using
self.tr = tr self.tr = tr
self.fractional = fractional
# Create a name based on the class name if none is provided
if name is None:
suffix = "_fractional" if fractional else ""
name = f"{self.__class__.__name__}{suffix}"
super().__init__(on="BOLD", name=name) super().__init__(on="BOLD", name=name)
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "vector"
def _compute( def _compute(
self, self,
input_data: Dict[str, Any], input_data: Dict[str, Any],
) -> Tuple["Nifti1Image", Path]: ) -> Tuple["Nifti1Image", "Nifti1Image", Path, Path]:
"""Compute ALFF and fALFF. """Compute ALFF and fALFF.
Parameters Parameters
@ -161,9 +132,13 @@ class ALFFBase(BaseMarker):
Returns Returns
------- -------
Niimg-like object Niimg-like object
The ALFF / fALFF as NIfTI. The ALFF as NIfTI.
Niimg-like object
The fALFF as NIfTI.
pathlib.Path pathlib.Path
The path to the ALFF / fALFF as NIfTI. The path to the ALFF as NIfTI.
pathlib.Path
The path to the fALFF as NIfTI.
""" """
logger.debug("Calculating ALFF and fALFF") logger.debug("Calculating ALFF and fALFF")
@ -186,11 +161,7 @@ class ALFFBase(BaseMarker):
# parcellation / coordinates to native space, else the # parcellation / coordinates to native space, else the
# path should be passed for use later if required. # path should be passed for use later if required.
# TODO(synchon): will be taken care in #292 # TODO(synchon): will be taken care in #292
if input_data["space"] == "native" and self.fractional: if input_data["space"] == "native":
return falff, input_data["path"] return alff, falff, input_data["path"], input_data["path"]
elif input_data["space"] == "native" and not self.fractional:
return alff, input_data["path"]
elif input_data["space"] != "native" and self.fractional:
return falff, falff_path
else: else:
return alff, alff_path return alff, falff, alff_path, falff_path

View file

@ -26,8 +26,6 @@ class ALFFParcels(ALFFBase):
parcellation : str or list of str parcellation : str or list of str
The name(s) of the parcellation(s). Check valid options by calling The name(s) of the parcellation(s). Check valid options by calling
:func:`.list_parcellations`. :func:`.list_parcellations`.
fractional : bool
Whether to compute fractional ALFF.
using : {"junifer", "afni"} using : {"junifer", "afni"}
Implementation to use for computing ALFF: Implementation to use for computing ALFF:
@ -73,7 +71,6 @@ class ALFFParcels(ALFFBase):
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
fractional: bool,
using: str, using: str,
highpass: float = 0.01, highpass: float = 0.01,
lowpass: float = 0.1, lowpass: float = 0.1,
@ -85,7 +82,6 @@ class ALFFParcels(ALFFBase):
) -> None: ) -> None:
# Superclass init first to validate `using` parameter # Superclass init first to validate `using` parameter
super().__init__( super().__init__(
fractional=fractional,
highpass=highpass, highpass=highpass,
lowpass=lowpass, lowpass=lowpass,
using=using, using=using,
@ -114,33 +110,63 @@ class ALFFParcels(ALFFBase):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The dictionary has the following The computed result as dictionary. This will be either returned
keys: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``alff`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
* ``falff`` : dictionary with the following keys:
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
logger.info("Calculating ALFF / fALFF for parcels") logger.info("Calculating ALFF / fALFF for parcels")
# Compute ALFF / fALFF # Compute ALFF + fALFF
output_data, output_file_path = self._compute(input_data=input) alff_output, falff_output, alff_output_path, falff_output_path = (
self._compute(input_data=input)
# Initialize parcel aggregation
parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation,
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
)
# Perform aggregation on ALFF / fALFF
parcel_aggregation_input = dict(input.items())
parcel_aggregation_input["data"] = output_data
parcel_aggregation_input["path"] = output_file_path
output = parcel_aggregation.compute(
input=parcel_aggregation_input,
extra_input=extra_input,
) )
return output # Perform aggregation on ALFF + fALFF
aggregation_alff_input = dict(input.items())
aggregation_falff_input = dict(input.items())
aggregation_alff_input["data"] = alff_output
aggregation_falff_input["data"] = falff_output
aggregation_alff_input["path"] = alff_output_path
aggregation_falff_input["path"] = falff_output_path
return {
"alff": {
**ParcelAggregation(
parcellation=self.parcellation,
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
).compute(
input=aggregation_alff_input,
extra_input=extra_input,
)[
"aggregation"
],
},
"falff": {
**ParcelAggregation(
parcellation=self.parcellation,
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
).compute(
input=aggregation_falff_input,
extra_input=extra_input,
)[
"aggregation"
],
},
}

View file

@ -26,8 +26,6 @@ class ALFFSpheres(ALFFBase):
coords : str coords : str
The name of the coordinates list to use. See The name of the coordinates list to use. See
:func:`.list_coordinates` for options. :func:`.list_coordinates` for options.
fractional : bool
Whether to compute fractional ALFF.
using : {"junifer", "afni"} using : {"junifer", "afni"}
Implementation to use for computing ALFF: Implementation to use for computing ALFF:
@ -80,7 +78,6 @@ class ALFFSpheres(ALFFBase):
def __init__( def __init__(
self, self,
coords: str, coords: str,
fractional: bool,
using: str, using: str,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False, allow_overlap: bool = False,
@ -94,7 +91,6 @@ class ALFFSpheres(ALFFBase):
) -> None: ) -> None:
# Superclass init first to validate `using` parameter # Superclass init first to validate `using` parameter
super().__init__( super().__init__(
fractional=fractional,
highpass=highpass, highpass=highpass,
lowpass=lowpass, lowpass=lowpass,
using=using, using=using,
@ -125,35 +121,67 @@ class ALFFSpheres(ALFFBase):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The dictionary has the following The computed result as dictionary. This will be either returned
keys: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``alff`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
* ``falff`` : dictionary with the following keys:
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
logger.info("Calculating ALFF / fALFF for spheres") logger.info("Calculating ALFF / fALFF for spheres")
# Compute ALFF / fALFF # Compute ALFF + fALFF
output_data, output_file_path = self._compute(input_data=input) alff_output, falff_output, alff_output_path, falff_output_path = (
self._compute(input_data=input)
# Initialize sphere aggregation
sphere_aggregation = SphereAggregation(
coords=self.coords,
radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
) )
# Perform aggregation on ALFF / fALFF # Perform aggregation on ALFF / fALFF
sphere_aggregation_input = dict(input.items()) aggregation_alff_input = dict(input.items())
sphere_aggregation_input["data"] = output_data aggregation_falff_input = dict(input.items())
sphere_aggregation_input["path"] = output_file_path aggregation_alff_input["data"] = alff_output
output = sphere_aggregation.compute( aggregation_falff_input["data"] = falff_output
input=sphere_aggregation_input, aggregation_alff_input["path"] = alff_output_path
extra_input=extra_input, aggregation_falff_input["path"] = falff_output_path
)
return output return {
"alff": {
**SphereAggregation(
coords=self.coords,
radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
).compute(
input=aggregation_alff_input,
extra_input=extra_input,
)[
"aggregation"
],
},
"falff": {
**SphereAggregation(
coords=self.coords,
radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
).compute(
input=aggregation_falff_input,
extra_input=extra_input,
)[
"aggregation"
],
},
}

View file

@ -21,6 +21,28 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
PARCELLATION = "TianxS1x3TxMNInonlinear2009cAsym" PARCELLATION = "TianxS1x3TxMNInonlinear2009cAsym"
@pytest.mark.parametrize(
"feature",
[
"alff",
"falff",
],
)
def test_ALFFParcels_get_output_type(feature: str) -> None:
"""Test ALFFParcels get_output_type().
Parameters
----------
feature : str
The parametrized feature name.
"""
assert "vector" == ALFFParcels(
parcellation=PARCELLATION,
using="junifer",
).get_output_type(input_type="BOLD", output_feature=feature)
def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None: def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
"""Test ALFFParcels. """Test ALFFParcels.
@ -41,7 +63,6 @@ def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Initialize marker # Initialize marker
marker = ALFFParcels( marker = ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
fractional=False,
using="junifer", using="junifer",
) )
# Fit transform marker on data # Fit transform marker on data
@ -51,15 +72,16 @@ def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Get BOLD output # Get BOLD output
assert "BOLD" in output assert "BOLD" in output
output_bold = output["BOLD"] for feature in output["BOLD"].keys():
# Assert BOLD output keys output_bold = output["BOLD"][feature]
assert "data" in output_bold # Assert BOLD output keys
assert "col_names" in output_bold assert "data" in output_bold
assert "col_names" in output_bold
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, 16) assert output_bold_data.shape == (1, 16)
# Reset log capture # Reset log capture
caplog.clear() caplog.clear()
@ -77,18 +99,13 @@ def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@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"
) )
@pytest.mark.parametrize( def test_ALFFParcels_comparison(tmp_path: Path) -> None:
"fractional", [True, False], ids=["fractional", "non-fractional"]
)
def test_ALFFParcels_comparison(tmp_path: Path, fractional: bool) -> None:
"""Test ALFFParcels implementation comparison. """Test ALFFParcels 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
Whether to compute fractional ALFF or not.
""" """
with PartlyCloudyTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
@ -99,7 +116,6 @@ def test_ALFFParcels_comparison(tmp_path: Path, fractional: bool) -> None:
# Initialize marker # Initialize marker
junifer_marker = ALFFParcels( junifer_marker = ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
fractional=fractional,
using="junifer", using="junifer",
) )
# Fit transform marker on data # Fit transform marker on data
@ -110,7 +126,6 @@ def test_ALFFParcels_comparison(tmp_path: Path, fractional: bool) -> None:
# Initialize marker # Initialize marker
afni_marker = ALFFParcels( afni_marker = ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
fractional=fractional,
using="afni", using="afni",
) )
# Fit transform marker on data # Fit transform marker on data
@ -118,9 +133,10 @@ def test_ALFFParcels_comparison(tmp_path: Path, fractional: bool) -> None:
# Get BOLD output # Get BOLD output
afni_output_bold = afni_output["BOLD"] afni_output_bold = afni_output["BOLD"]
# Check for Pearson correlation coefficient for feature in afni_output_bold.keys():
r, _ = sp.stats.pearsonr( # Check for Pearson correlation coefficient
junifer_output_bold["data"][0], r, _ = sp.stats.pearsonr(
afni_output_bold["data"][0], junifer_output_bold[feature]["data"][0],
) afni_output_bold[feature]["data"][0],
assert r > 0.97 )
assert r > 0.97

View file

@ -21,6 +21,28 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
COORDINATES = "DMNBuckner" COORDINATES = "DMNBuckner"
@pytest.mark.parametrize(
"feature",
[
"alff",
"falff",
],
)
def test_ALFFSpheres_get_output_type(feature: str) -> None:
"""Test ALFFSpheres get_output_type().
Parameters
----------
feature : str
The parametrized feature name.
"""
assert "vector" == ALFFSpheres(
coords=COORDINATES,
using="junifer",
).get_output_type(input_type="BOLD", output_feature=feature)
def test_ALFFSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None: def test_ALFFSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
"""Test ALFFSpheres. """Test ALFFSpheres.
@ -41,7 +63,6 @@ def test_ALFFSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Initialize marker # Initialize marker
marker = ALFFSpheres( marker = ALFFSpheres(
coords=COORDINATES, coords=COORDINATES,
fractional=False,
using="junifer", using="junifer",
radius=5.0, radius=5.0,
) )
@ -52,15 +73,16 @@ def test_ALFFSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Get BOLD output # Get BOLD output
assert "BOLD" in output assert "BOLD" in output
output_bold = output["BOLD"] for feature in output["BOLD"].keys():
# Assert BOLD output keys output_bold = output["BOLD"][feature]
assert "data" in output_bold # Assert BOLD output keys
assert "col_names" in output_bold assert "data" in output_bold
assert "col_names" in output_bold
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, 6) assert output_bold_data.shape == (1, 6)
# Reset log capture # Reset log capture
caplog.clear() caplog.clear()
@ -78,18 +100,13 @@ def test_ALFFSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@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"
) )
@pytest.mark.parametrize( def test_ALFFSpheres_comparison(tmp_path: Path) -> None:
"fractional", [True, False], ids=["fractional", "non-fractional"]
)
def test_ALFFSpheres_comparison(tmp_path: Path, fractional: bool) -> None:
"""Test ALFFSpheres implementation comparison. """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
Whether to compute fractional ALFF or not.
""" """
with PartlyCloudyTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
@ -100,7 +117,6 @@ def test_ALFFSpheres_comparison(tmp_path: Path, fractional: bool) -> None:
# Initialize marker # Initialize marker
junifer_marker = ALFFSpheres( junifer_marker = ALFFSpheres(
coords=COORDINATES, coords=COORDINATES,
fractional=fractional,
using="junifer", using="junifer",
radius=5.0, radius=5.0,
) )
@ -112,7 +128,6 @@ def test_ALFFSpheres_comparison(tmp_path: Path, fractional: bool) -> None:
# Initialize marker # Initialize marker
afni_marker = ALFFSpheres( afni_marker = ALFFSpheres(
coords=COORDINATES, coords=COORDINATES,
fractional=fractional,
using="afni", using="afni",
radius=5.0, radius=5.0,
) )
@ -121,9 +136,10 @@ def test_ALFFSpheres_comparison(tmp_path: Path, fractional: bool) -> None:
# Get BOLD output # Get BOLD output
afni_output_bold = afni_output["BOLD"] afni_output_bold = afni_output["BOLD"]
# Check for Pearson correlation coefficient for feature in afni_output_bold.keys():
r, _ = sp.stats.pearsonr( # Check for Pearson correlation coefficient
junifer_output_bold["data"][0], r, _ = sp.stats.pearsonr(
afni_output_bold["data"][0], junifer_output_bold[feature]["data"][0],
) afni_output_bold[feature]["data"][0],
assert r > 0.99 )
assert r > 0.99

View file

@ -45,6 +45,12 @@ class CrossParcellationFC(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"functional_connectivity": "matrix",
},
}
def __init__( def __init__(
self, self,
parcellation_one: str, parcellation_one: str,
@ -65,33 +71,6 @@ class CrossParcellationFC(BaseMarker):
self.masks = masks self.masks = masks
super().__init__(on=["BOLD"], name=name) super().__init__(on=["BOLD"], name=name)
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "matrix"
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -118,10 +97,14 @@ class CrossParcellationFC(BaseMarker):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the correlation values between the two parcellations * ``functional_connectivity`` : dictionary with the following keys:
as a numpy.ndarray
* ``col_names`` : the ROIs for first parcellation as a list - ``data`` : correlation between the two parcellations as
* ``row_names`` : the ROIs for second parcellation as a list ``numpy.ndarray``
- ``col_names`` : ROI labels for first parcellation as list of
str
- ``row_names`` : ROI labels for second parcellation as list of
str
""" """
logger.debug( logger.debug(
@ -129,31 +112,32 @@ class CrossParcellationFC(BaseMarker):
f" {self.parcellation_one} and " f" {self.parcellation_one} and "
f"{self.parcellation_two} parcellations." f"{self.parcellation_two} parcellations."
) )
# Initialize a ParcelAggregation # Perform aggregation using two parcellations
parcellation_one_dict = ParcelAggregation( aggregation_parcellation_one = ParcelAggregation(
parcellation=self.parcellation_one, parcellation=self.parcellation_one,
method=self.aggregation_method, method=self.aggregation_method,
masks=self.masks, masks=self.masks,
).compute(input, extra_input=extra_input) ).compute(input, extra_input=extra_input)
parcellation_two_dict = ParcelAggregation( aggregation_parcellation_two = ParcelAggregation(
parcellation=self.parcellation_two, parcellation=self.parcellation_two,
method=self.aggregation_method, method=self.aggregation_method,
masks=self.masks, masks=self.masks,
).compute(input, extra_input=extra_input) ).compute(input, extra_input=extra_input)
parcellated_ts_one = parcellation_one_dict["data"]
parcellated_ts_two = parcellation_two_dict["data"]
# columns should be named after parcellation 1
# rows should be named after parcellation 2
result = _correlate_dataframes(
pd.DataFrame(parcellated_ts_one),
pd.DataFrame(parcellated_ts_two),
method=self.correlation_method,
).values
return { return {
"data": result, "functional_connectivity": {
"col_names": parcellation_one_dict["col_names"], "data": _correlate_dataframes(
"row_names": parcellation_two_dict["col_names"], pd.DataFrame(
aggregation_parcellation_one["aggregation"]["data"]
),
pd.DataFrame(
aggregation_parcellation_two["aggregation"]["data"]
),
method=self.correlation_method,
).values,
# Columns should be named after parcellation 1
"col_names": aggregation_parcellation_one["col_names"],
# Rows should be named after parcellation 2
"row_names": aggregation_parcellation_two["col_names"],
},
} }

View file

@ -98,23 +98,29 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
parcel_aggregation = ParcelAggregation( # Perform aggregation
aggregation = ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input, extra_input=extra_input)
# Compute edgewise timeseries
bold_aggregated = parcel_aggregation.compute(
input, extra_input=extra_input
)
ets, edge_names = _ets( ets, edge_names = _ets(
bold_aggregated["data"], bold_aggregated["col_names"] bold_ts=aggregation["aggregation"]["data"],
roi_names=aggregation["aggregation"]["col_names"],
) )
return {"data": ets, "col_names": edge_names} return {
"aggregation": {
"data": ets,
"col_names": edge_names,
},
}

View file

@ -110,11 +110,14 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
sphere_aggregation = SphereAggregation( # Perform aggregation
aggregation = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap, allow_overlap=self.allow_overlap,
@ -122,12 +125,13 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input, extra_input=extra_input)
bold_aggregated = sphere_aggregation.compute( # Compute edgewise timeseries
input, extra_input=extra_input
)
ets, edge_names = _ets( ets, edge_names = _ets(
bold_aggregated["data"], bold_aggregated["col_names"] bold_ts=aggregation["aggregation"]["data"],
roi_names=aggregation["aggregation"]["col_names"],
) )
return {"data": ets, "col_names": edge_names} return {
"aggregation": {"data": ets, "col_names": edge_names},
}

View file

@ -47,6 +47,12 @@ class FunctionalConnectivityBase(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "scikit-learn"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "scikit-learn"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"functional_connectivity": "matrix",
},
}
def __init__( def __init__(
self, self,
agg_method: str = "mean", agg_method: str = "mean",
@ -80,33 +86,6 @@ class FunctionalConnectivityBase(BaseMarker):
klass=NotImplementedError, klass=NotImplementedError,
) )
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "matrix"
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -128,13 +107,16 @@ class FunctionalConnectivityBase(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The following keys will be The computed result as dictionary. This will be either returned
included in the dictionary: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : functional connectivity matrix as a ``numpy.ndarray``. * ``functional_connectivity`` : dictionary with the following keys:
* ``row_names`` : row names as a list
* ``col_names`` : column names as a list - ``data`` : functional connectivity matrix as ``numpy.ndarray``
* ``matrix_kind`` : the kind of matrix (tril, triu or full) - ``row_names`` : ROI labels as list of str
- ``col_names`` : ROI labels as list of str
- ``matrix_kind`` : the kind of matrix (tril, triu or full)
""" """
# Perform necessary aggregation # Perform necessary aggregation
@ -148,10 +130,14 @@ class FunctionalConnectivityBase(BaseMarker):
else: else:
connectivity = ConnectivityMeasure(kind=self.cor_method) connectivity = ConnectivityMeasure(kind=self.cor_method)
# Create dictionary for output # Create dictionary for output
out = {} return {
out["data"] = connectivity.fit_transform([aggregation["data"]])[0] "functional_connectivity": {
# Create column names "data": connectivity.fit_transform(
out["row_names"] = aggregation["col_names"] [aggregation["aggregation"]["data"]]
out["col_names"] = aggregation["col_names"] )[0],
out["matrix_kind"] = "tril" # Create column names
return out "row_names": aggregation["aggregation"]["col_names"],
"col_names": aggregation["aggregation"]["col_names"],
"matrix_kind": "tril",
},
}

View file

@ -90,16 +90,16 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
parcel_aggregation = ParcelAggregation( return ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input=input, extra_input=extra_input)
# Return the 2D timeseries after parcel aggregation
return parcel_aggregation.compute(input, extra_input=extra_input)

View file

@ -104,11 +104,13 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
sphere_aggregation = SphereAggregation( return SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap, allow_overlap=self.allow_overlap,
@ -116,6 +118,4 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input=input, extra_input=extra_input)
# Return the 2D timeseries after sphere aggregation
return sphere_aggregation.compute(input, extra_input=extra_input)

View file

@ -33,10 +33,11 @@ def test_init() -> None:
def test_get_output_type() -> None: def test_get_output_type() -> None:
"""Test CrossParcellationFC get_output_type().""" """Test CrossParcellationFC get_output_type()."""
crossparcellation = CrossParcellationFC( assert "matrix" == CrossParcellationFC(
parcellation_one=parcellation_one, parcellation_two=parcellation_two parcellation_one=parcellation_one, parcellation_two=parcellation_two
).get_output_type(
input_type="BOLD", output_feature="functional_connectivity"
) )
assert "matrix" == crossparcellation.get_output_type("BOLD")
@pytest.mark.skipif( @pytest.mark.skipif(
@ -59,7 +60,9 @@ def test_compute(tmp_path: Path) -> None:
parcellation_two=parcellation_two, parcellation_two=parcellation_two,
correlation_method="spearman", correlation_method="spearman",
) )
out = crossparcellation.compute(element_data["BOLD"]) out = crossparcellation.compute(element_data["BOLD"])[
"functional_connectivity"
]
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
@ -92,5 +95,6 @@ def test_store(tmp_path: Path) -> None:
crossparcellation.fit_transform(input=element_data, storage=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_functional_connectivity"
for x in features.values()
) )

View file

@ -28,11 +28,13 @@ def test_EdgeCentricFCParcels(tmp_path: Path) -> None:
cor_method_params={"empirical": True}, cor_method_params={"empirical": True},
) )
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "matrix" assert "matrix" == marker.get_output_type(
input_type="BOLD", output_feature="functional_connectivity"
)
# Fit-transform the data # Fit-transform the data
edge_fc = marker.fit_transform(element_data) edge_fc = marker.fit_transform(element_data)
edge_fc_bold = edge_fc["BOLD"] edge_fc_bold = edge_fc["BOLD"]["functional_connectivity"]
# For 16 ROIs we should get (16 * (16 -1) / 2) edges in the ETS # For 16 ROIs we should get (16 * (16 -1) / 2) edges in the ETS
n_edges = int(16 * (16 - 1) / 2) n_edges = int(16 * (16 - 1) / 2)
@ -51,5 +53,6 @@ def test_EdgeCentricFCParcels(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_EdgeCentricFCParcels" for x in features.values() x["name"] == "BOLD_EdgeCentricFCParcels_functional_connectivity"
for x in features.values()
) )

View file

@ -27,11 +27,13 @@ def test_EdgeCentricFCSpheres(tmp_path: Path) -> None:
coords="DMNBuckner", radius=5.0, cor_method="correlation" coords="DMNBuckner", radius=5.0, cor_method="correlation"
) )
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "matrix" assert "matrix" == marker.get_output_type(
input_type="BOLD", output_feature="functional_connectivity"
)
# Fit-transform the data # Fit-transform the data
edge_fc = marker.fit_transform(element_data) edge_fc = marker.fit_transform(element_data)
edge_fc_bold = edge_fc["BOLD"] edge_fc_bold = edge_fc["BOLD"]["functional_connectivity"]
# There are six DMNBuckner coordinates, so # There are six DMNBuckner coordinates, so
# for 6 ROIs we should get (6 * (6 -1) / 2) edges in the ETS # for 6 ROIs we should get (6 * (6 -1) / 2) edges in the ETS
@ -57,5 +59,6 @@ def test_EdgeCentricFCSpheres(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_EdgeCentricFCSpheres" for x in features.values() x["name"] == "BOLD_EdgeCentricFCSpheres_functional_connectivity"
for x in features.values()
) )

View file

@ -35,11 +35,13 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
parcellation="TianxS1x3TxMNInonlinear2009cAsym" parcellation="TianxS1x3TxMNInonlinear2009cAsym"
) )
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "matrix" assert "matrix" == marker.get_output_type(
input_type="BOLD", output_feature="functional_connectivity"
)
# Fit-transform the data # Fit-transform the data
fc = marker.fit_transform(element_data) fc = marker.fit_transform(element_data)
fc_bold = fc["BOLD"] fc_bold = fc["BOLD"]["functional_connectivity"]
assert "data" in fc_bold assert "data" in fc_bold
assert "row_names" in fc_bold assert "row_names" in fc_bold
@ -83,6 +85,7 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_FunctionalConnectivityParcels" x["name"]
== "BOLD_FunctionalConnectivityParcels_functional_connectivity"
for x in features.values() for x in features.values()
) )

View file

@ -38,11 +38,13 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
coords="DMNBuckner", radius=5.0, cor_method="correlation" coords="DMNBuckner", radius=5.0, cor_method="correlation"
) )
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "matrix" assert "matrix" == marker.get_output_type(
input_type="BOLD", output_feature="functional_connectivity"
)
# Fit-transform the data # Fit-transform the data
fc = marker.fit_transform(element_data) fc = marker.fit_transform(element_data)
fc_bold = fc["BOLD"] fc_bold = fc["BOLD"]["functional_connectivity"]
assert "data" in fc_bold assert "data" in fc_bold
assert "row_names" in fc_bold assert "row_names" in fc_bold
@ -80,7 +82,8 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_FunctionalConnectivitySpheres" x["name"]
== "BOLD_FunctionalConnectivitySpheres_functional_connectivity"
for x in features.values() for x in features.values()
) )
@ -103,11 +106,13 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
cor_method_params={"empirical": True}, cor_method_params={"empirical": True},
) )
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "matrix" assert "matrix" == marker.get_output_type(
input_type="BOLD", output_feature="functional_connectivity"
)
# Fit-transform the data # Fit-transform the data
fc = marker.fit_transform(element_data) fc = marker.fit_transform(element_data)
fc_bold = fc["BOLD"] fc_bold = fc["BOLD"]["functional_connectivity"]
assert "data" in fc_bold assert "data" in fc_bold
assert "row_names" in fc_bold assert "row_names" in fc_bold

View file

@ -63,6 +63,36 @@ class ParcelAggregation(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "numpy"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "numpy"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"T1w": {
"aggregation": "vector",
},
"T2w": {
"aggregation": "vector",
},
"BOLD": {
"aggregation": "timeseries",
},
"VBM_GM": {
"aggregation": "vector",
},
"VBM_WM": {
"aggregation": "vector",
},
"VBM_CSF": {
"aggregation": "vector",
},
"fALFF": {
"aggregation": "vector",
},
"GCOR": {
"aggregation": "vector",
},
"LCOR": {
"aggregation": "vector",
},
}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
@ -96,61 +126,6 @@ class ParcelAggregation(BaseMarker):
self.time_method = time_method self.time_method = time_method
self.time_method_params = time_method_params or {} self.time_method_params = time_method_params or {}
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return [
"T1w",
"T2w",
"BOLD",
"VBM_GM",
"VBM_WM",
"VBM_CSF",
"fALFF",
"GCOR",
"LCOR",
]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
Raises
------
ValueError
If the ``input_type`` is invalid.
"""
if input_type in [
"VBM_GM",
"VBM_WM",
"VBM_CSF",
"fALFF",
"GCOR",
"LCOR",
]:
return "vector"
elif input_type == "BOLD":
return "timeseries"
else:
raise_error(f"Unknown input kind for {input_type}")
def compute( def compute(
self, input: Dict[str, Any], extra_input: Optional[Dict] = None self, input: Dict[str, Any], extra_input: Optional[Dict] = None
) -> Dict: ) -> Dict:
@ -174,8 +149,10 @@ class ParcelAggregation(BaseMarker):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
Warns Warns
----- -----
@ -253,5 +230,9 @@ class ParcelAggregation(BaseMarker):
"available." "available."
) )
# Format the output # Format the output
out = {"data": out_values, "col_names": labels} return {
return out "aggregation": {
"data": out_values,
"col_names": labels,
},
}

View file

@ -62,6 +62,12 @@ class ReHoBase(BaseMarker):
}, },
] ]
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"reho": "vector",
},
}
def __init__( def __init__(
self, self,
using: str, using: str,
@ -76,33 +82,6 @@ class ReHoBase(BaseMarker):
self.using = using self.using = using
super().__init__(on="BOLD", name=name) super().__init__(on="BOLD", name=name)
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "vector"
def _compute( def _compute(
self, self,
input_data: Dict[str, Any], input_data: Dict[str, Any],

View file

@ -125,11 +125,14 @@ class ReHoParcels(ReHoBase):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The dictionary has the following The computed result as dictionary. This will be either returned
keys: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a 1D numpy.ndarray * ``reho`` : dictionary with the following keys:
* ``col_names`` : the column labels for the parcels as a list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
logger.info("Calculating ReHo for parcels") logger.info("Calculating ReHo for parcels")
@ -145,22 +148,27 @@ class ReHoParcels(ReHoBase):
else: else:
reho_map, reho_file_path = self._compute(input_data=input) reho_map, reho_file_path = self._compute(input_data=input)
# Initialize parcel aggregation # Perform aggregation on reho map
aggregation_input = dict(input.items())
aggregation_input["data"] = reho_map
aggregation_input["path"] = reho_file_path
parcel_aggregation = ParcelAggregation( parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(
# Perform aggregation on reho map input=aggregation_input,
parcel_aggregation_input = dict(input.items())
parcel_aggregation_input["data"] = reho_map
parcel_aggregation_input["path"] = reho_file_path
output = parcel_aggregation.compute(
input=parcel_aggregation_input,
extra_input=extra_input, extra_input=extra_input,
) )
# Only use the first row and expand row dimension
output["data"] = output["data"][0][np.newaxis, :] return {
return output "reho": {
# Only use the first row and expand row dimension
"data": parcel_aggregation["aggregation"]["data"][0][
np.newaxis, :
],
"col_names": parcel_aggregation["aggregation"]["col_names"],
}
}

View file

@ -140,11 +140,14 @@ class ReHoSpheres(ReHoBase):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The dictionary has the following The computed result as dictionary. This will be either returned
keys: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a 1D numpy.ndarray * ``reho`` : dictionary with the following keys:
* ``col_names`` : the column labels for the spheres as a list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
logger.info("Calculating ReHo for spheres") logger.info("Calculating ReHo for spheres")
@ -160,7 +163,10 @@ class ReHoSpheres(ReHoBase):
else: else:
reho_map, reho_file_path = self._compute(input_data=input) reho_map, reho_file_path = self._compute(input_data=input)
# Initialize sphere aggregation # Perform aggregation on reho map
aggregation_input = dict(input.items())
aggregation_input["data"] = reho_map
aggregation_input["path"] = reho_file_path
sphere_aggregation = SphereAggregation( sphere_aggregation = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
@ -169,14 +175,14 @@ class ReHoSpheres(ReHoBase):
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input=aggregation_input, extra_input=extra_input)
# Perform aggregation on reho map
sphere_aggregation_input = dict(input.items()) return {
sphere_aggregation_input["data"] = reho_map "reho": {
sphere_aggregation_input["path"] = reho_file_path # Only use the first row and expand row dimension
output = sphere_aggregation.compute( "data": sphere_aggregation["aggregation"]["data"][0][
input=sphere_aggregation_input, extra_input=extra_input np.newaxis, :
) ],
# Only use the first row and expand row dimension "col_names": sphere_aggregation["aggregation"]["col_names"],
output["data"] = output["data"][0][np.newaxis, :] }
return output }

View file

@ -42,6 +42,11 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
parcellation="TianxS1x3TxMNInonlinear2009cAsym", parcellation="TianxS1x3TxMNInonlinear2009cAsym",
using="junifer", using="junifer",
) )
# Check correct output
assert "vector" == marker.get_output_type(
input_type="BOLD", output_feature="reho"
)
# Fit transform marker on data # Fit transform marker on data
output = marker.fit_transform(element_data) output = marker.fit_transform(element_data)
@ -49,7 +54,7 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Get BOLD output # Get BOLD output
assert "BOLD" in output assert "BOLD" in output
output_bold = output["BOLD"] output_bold = output["BOLD"]["reho"]
# Assert BOLD output keys # Assert BOLD output keys
assert "data" in output_bold assert "data" in output_bold
assert "col_names" in output_bold assert "col_names" in output_bold
@ -102,14 +107,14 @@ def test_ReHoParcels_comparison(tmp_path: Path) -> None:
# 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)
# Get BOLD output # Get BOLD output
junifer_output_bold = junifer_output["BOLD"] junifer_output_bold = junifer_output["BOLD"]["reho"]
# Initialize marker # Initialize marker
afni_marker = ReHoParcels(parcellation="Schaefer100x7", 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
afni_output_bold = afni_output["BOLD"] afni_output_bold = afni_output["BOLD"]["reho"]
# Check for Pearson correlation coefficient # Check for Pearson correlation coefficient
r, _ = sp.stats.pearsonr( r, _ = sp.stats.pearsonr(

View file

@ -40,6 +40,11 @@ def test_ReHoSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
marker = ReHoSpheres( marker = ReHoSpheres(
coords=COORDINATES, using="junifer", radius=10.0 coords=COORDINATES, using="junifer", radius=10.0
) )
# Check correct output
assert "vector" == marker.get_output_type(
input_type="BOLD", output_feature="reho"
)
# Fit transform marker on data # Fit transform marker on data
output = marker.fit_transform(element_data) output = marker.fit_transform(element_data)
@ -47,7 +52,7 @@ def test_ReHoSpheres(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Get BOLD output # Get BOLD output
assert "BOLD" in output assert "BOLD" in output
output_bold = output["BOLD"] output_bold = output["BOLD"]["reho"]
# Assert BOLD output keys # Assert BOLD output keys
assert "data" in output_bold assert "data" in output_bold
assert "col_names" in output_bold assert "col_names" in output_bold
@ -99,7 +104,7 @@ def test_ReHoSpheres_comparison(tmp_path: Path) -> None:
# 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)
# Get BOLD output # Get BOLD output
junifer_output_bold = junifer_output["BOLD"] junifer_output_bold = junifer_output["BOLD"]["reho"]
# Initialize marker # Initialize marker
afni_marker = ReHoSpheres( afni_marker = ReHoSpheres(
@ -110,7 +115,7 @@ def test_ReHoSpheres_comparison(tmp_path: Path) -> None:
# 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
afni_output_bold = afni_output["BOLD"] afni_output_bold = afni_output["BOLD"]["reho"]
# Check for Pearson correlation coefficient # Check for Pearson correlation coefficient
r, _ = sp.stats.pearsonr( r, _ = sp.stats.pearsonr(

View file

@ -68,6 +68,36 @@ class SphereAggregation(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "numpy"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn", "numpy"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"T1w": {
"aggregation": "vector",
},
"T2w": {
"aggregation": "vector",
},
"BOLD": {
"aggregation": "timeseries",
},
"VBM_GM": {
"aggregation": "vector",
},
"VBM_WM": {
"aggregation": "vector",
},
"VBM_CSF": {
"aggregation": "vector",
},
"fALFF": {
"aggregation": "vector",
},
"GCOR": {
"aggregation": "vector",
},
"LCOR": {
"aggregation": "vector",
},
}
def __init__( def __init__(
self, self,
coords: str, coords: str,
@ -103,61 +133,6 @@ class SphereAggregation(BaseMarker):
self.time_method = time_method self.time_method = time_method
self.time_method_params = time_method_params or {} self.time_method_params = time_method_params or {}
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return [
"T1w",
"T2w",
"BOLD",
"VBM_GM",
"VBM_WM",
"VBM_CSF",
"fALFF",
"GCOR",
"LCOR",
]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
Raises
------
ValueError
If the ``input_type`` is invalid.
"""
if input_type in [
"VBM_GM",
"VBM_WM",
"VBM_CSF",
"fALFF",
"GCOR",
"LCOR",
]:
return "vector"
elif input_type == "BOLD":
return "timeseries"
else:
raise_error(f"Unknown input kind for {input_type}")
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -183,8 +158,10 @@ class SphereAggregation(BaseMarker):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : ROI values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
Warns Warns
----- -----
@ -241,5 +218,9 @@ class SphereAggregation(BaseMarker):
"available." "available."
) )
# Format the output # Format the output
out = {"data": out_values, "col_names": labels} return {
return out "aggregation": {
"data": out_values,
"col_names": labels,
},
}

View file

@ -39,6 +39,12 @@ class TemporalSNRBase(BaseMarker):
_DEPENDENCIES: ClassVar[Set[str]] = {"nilearn"} _DEPENDENCIES: ClassVar[Set[str]] = {"nilearn"}
_MARKER_INOUT_MAPPINGS: ClassVar[Dict[str, Dict[str, str]]] = {
"BOLD": {
"tsnr": "vector",
},
}
def __init__( def __init__(
self, self,
agg_method: str = "mean", agg_method: str = "mean",
@ -61,33 +67,6 @@ class TemporalSNRBase(BaseMarker):
klass=NotImplementedError, klass=NotImplementedError,
) )
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "vector"
def compute( def compute(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
@ -107,11 +86,14 @@ class TemporalSNRBase(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The following keys will be The computed result as dictionary. This will be either returned
included in the dictionary: to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the computed values as a ``numpy.ndarray`` * ``tsnr`` : dictionary with the following keys:
* ``col_names`` : the column labels for the computed values as list
- ``data`` : computed tSNR as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
# Calculate voxelwise temporal signal-to-noise ratio in an image # Calculate voxelwise temporal signal-to-noise ratio in an image
@ -129,4 +111,10 @@ class TemporalSNRBase(BaseMarker):
mask_img=mask_img, mask_img=mask_img,
) )
# Perform necessary aggregation and return # Perform necessary aggregation and return
return self.aggregate(input=input, extra_input=extra_input) return {
"tsnr": {
**self.aggregate(input=input, extra_input=extra_input)[
"aggregation"
]
}
}

View file

@ -77,16 +77,16 @@ class TemporalSNRParcels(TemporalSNRBase):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : ROI-wise temporal SNR as a ``numpy.ndarray`` * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the ROI labels for the computed values as list
- ``data`` : ROI-wise tSNR values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
parcel_aggregation = ParcelAggregation( return ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input=input, extra_input=extra_input)
# Return the 2D timeseries after parcel aggregation
return parcel_aggregation.compute(input=input, extra_input=extra_input)

View file

@ -92,11 +92,13 @@ class TemporalSNRSpheres(TemporalSNRBase):
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* ``data`` : VOI-wise temporal SNR as a ``numpy.ndarray`` * ``aggregation`` : dictionary with the following keys:
* ``col_names`` : the VOI labels for the computed values as list
- ``data`` : ROI-wise tSNR values as ``numpy.ndarray``
- ``col_names`` : ROI labels as list of str
""" """
sphere_aggregation = SphereAggregation( return SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap, allow_overlap=self.allow_overlap,
@ -104,6 +106,4 @@ class TemporalSNRSpheres(TemporalSNRBase):
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,
on="BOLD", on="BOLD",
) ).compute(input=input, extra_input=extra_input)
# Return the 2D timeseries after sphere aggregation
return sphere_aggregation.compute(input=input, extra_input=extra_input)

View file

@ -20,11 +20,13 @@ def test_TemporalSNRParcels_computation() -> None:
parcellation="TianxS1x3TxMNInonlinear2009cAsym" parcellation="TianxS1x3TxMNInonlinear2009cAsym"
) )
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "vector" assert "vector" == marker.get_output_type(
input_type="BOLD", output_feature="tsnr"
)
# Fit-transform the data # Fit-transform the data
tsnr_parcels = marker.fit_transform(element_data) tsnr_parcels = marker.fit_transform(element_data)
tsnr_parcels_bold = tsnr_parcels["BOLD"] tsnr_parcels_bold = tsnr_parcels["BOLD"]["tsnr"]
assert "data" in tsnr_parcels_bold assert "data" in tsnr_parcels_bold
assert "col_names" in tsnr_parcels_bold assert "col_names" in tsnr_parcels_bold
@ -51,5 +53,6 @@ def test_TemporalSNRParcels_storage(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_TemporalSNRParcels" for x in features.values() x["name"] == "BOLD_TemporalSNRParcels_tsnr"
for x in features.values()
) )

View file

@ -20,11 +20,13 @@ def test_TemporalSNRSpheres_computation() -> None:
element_data = DefaultDataReader().fit_transform(dg["sub001"]) element_data = DefaultDataReader().fit_transform(dg["sub001"])
marker = TemporalSNRSpheres(coords="DMNBuckner", radius=5.0) marker = TemporalSNRSpheres(coords="DMNBuckner", radius=5.0)
# Check correct output # Check correct output
assert marker.get_output_type("BOLD") == "vector" assert "vector" == marker.get_output_type(
input_type="BOLD", output_feature="tsnr"
)
# Fit-transform the data # Fit-transform the data
tsnr_spheres = marker.fit_transform(element_data) tsnr_spheres = marker.fit_transform(element_data)
tsnr_spheres_bold = tsnr_spheres["BOLD"] tsnr_spheres_bold = tsnr_spheres["BOLD"]["tsnr"]
assert "data" in tsnr_spheres_bold assert "data" in tsnr_spheres_bold
assert "col_names" in tsnr_spheres_bold assert "col_names" in tsnr_spheres_bold
@ -49,7 +51,8 @@ def test_TemporalSNRSpheres_storage(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_TemporalSNRSpheres" for x in features.values() x["name"] == "BOLD_TemporalSNRSpheres_tsnr"
for x in features.values()
) )

View file

@ -13,16 +13,29 @@ from junifer.markers import BrainPrint
from junifer.pipeline.utils import _check_freesurfer from junifer.pipeline.utils import _check_freesurfer
def test_get_output_type() -> None: @pytest.mark.parametrize(
"""Test BrainPrint get_output_type().""" "feature, storage_type",
marker = BrainPrint() [
assert marker.get_output_type("FreeSurfer") == "vector" ("eigenvalues", "scalar_table"),
("areas", "vector"),
("volumes", "vector"),
("distances", "vector"),
],
)
def test_get_output_type(feature: str, storage_type: str) -> None:
"""Test BrainPrint get_output_type().
Parameters
----------
feature : str
The parametrized feature name.
storage_type : str
The parametrized storage type.
def test_validate() -> None: """
"""Test BrainPrint validate().""" assert storage_type == BrainPrint().get_output_type(
marker = BrainPrint() input_type="FreeSurfer", output_feature=feature
assert set(marker.validate(["FreeSurfer"])) == {"scalar_table", "vector"} )
@pytest.mark.skipif( @pytest.mark.skipif(
@ -39,9 +52,7 @@ def test_compute() -> None:
element = dg["sub-0001"] element = dg["sub-0001"]
# Fetch element data # Fetch element data
element_data = DefaultDataReader().fit_transform(element) element_data = DefaultDataReader().fit_transform(element)
# Initialize the marker # Compute marker
marker = BrainPrint() feature_map = BrainPrint().fit_transform(element_data)
# Compute the marker
feature_map = marker.fit_transform(element_data)
# Assert the output keys # Assert the output keys
assert {"eigenvalues", "areas", "volumes"} == set(feature_map.keys()) assert {"eigenvalues", "areas", "volumes"} == set(feature_map.keys())

View file

@ -84,7 +84,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 "BOLD" in out[t_name] assert "BOLD" in out[t_name]
t_bold = out[t_name]["BOLD"] t_bold = out[t_name]["BOLD"]["aggregation"]
assert "data" in t_bold assert "data" in t_bold
assert "col_names" in t_bold assert "col_names" in t_bold
assert "meta" in t_bold assert "meta" in t_bold
@ -107,7 +107,8 @@ 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]["BOLD"]["data"], out2[t_name]["BOLD"]["data"] out[t_name]["BOLD"]["aggregation"]["data"],
out2[t_name]["BOLD"]["aggregation"]["data"],
) )
@ -201,20 +202,20 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
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 = "tian_mean" fname = "tian_mean"
t_data = out[fname]["BOLD"]["data"] # type: ignore t_data = out[fname]["BOLD"]["aggregation"]["data"] # type: ignore
cols = out[fname]["BOLD"]["col_names"] # type: ignore cols = out[fname]["BOLD"]["aggregation"]["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 = "tian_std" fname = "tian_std"
t_data = out[fname]["BOLD"]["data"] # type: ignore t_data = out[fname]["BOLD"]["aggregation"]["data"] # type: ignore
cols = out[fname]["BOLD"]["col_names"] # type: ignore cols = out[fname]["BOLD"]["aggregation"]["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 = "tian_trim_mean90" fname = "tian_trim_mean90"
t_data = out[fname]["BOLD"]["data"] # type: ignore t_data = out[fname]["BOLD"]["aggregation"]["data"] # type: ignore
cols = out[fname]["BOLD"]["col_names"] # type: ignore cols = out[fname]["BOLD"]["aggregation"]["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

@ -26,8 +26,9 @@ def test_compute() -> None:
with PartlyCloudyTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg["sub-01"]) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Compute the RSSETSMarker # Compute the RSSETSMarker
marker = RSSETSMarker(parcellation=PARCELLATION) rss_ets = RSSETSMarker(parcellation=PARCELLATION).compute(
rss_ets = marker.compute(element_data["BOLD"]) element_data["BOLD"]
)
# Compare with nilearn # Compare with nilearn
# Load testing parcellation # Load testing parcellation
@ -41,14 +42,14 @@ def test_compute() -> None:
element_data["BOLD"]["data"] element_data["BOLD"]["data"]
) )
# Assert the dimension of timeseries # Assert the dimension of timeseries
assert extacted_timeseries.shape[0] == len(rss_ets["data"]) assert extacted_timeseries.shape[0] == len(rss_ets["rss_ets"]["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()."""
assert "timeseries" == RSSETSMarker( assert "timeseries" == RSSETSMarker(
parcellation=PARCELLATION parcellation=PARCELLATION
).get_output_type("BOLD") ).get_output_type(input_type="BOLD", output_feature="rss_ets")
def test_store(tmp_path: Path) -> None: def test_store(tmp_path: Path) -> None:
@ -61,12 +62,17 @@ def test_store(tmp_path: Path) -> None:
""" """
with PartlyCloudyTestingDataGrabber() as dg: with PartlyCloudyTestingDataGrabber() as dg:
# Get element data
element_data = DefaultDataReader().fit_transform(dg["sub-01"]) element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Compute the RSSETSMarker
marker = RSSETSMarker(parcellation=PARCELLATION)
# Create storage # Create storage
storage = SQLiteFeatureStorage(tmp_path / "test_rss_ets.sqlite") storage = SQLiteFeatureStorage(tmp_path / "test_rss_ets.sqlite")
# Store # Compute the RSSETSMarker and store
marker.fit_transform(input=element_data, storage=storage) _ = RSSETSMarker(parcellation=PARCELLATION).fit_transform(
input=element_data, storage=storage
)
# Retrieve features
features = storage.list_features() features = storage.list_features()
assert any(x["name"] == "BOLD_RSSETSMarker" for x in features.values()) # Check marker name
assert any(
x["name"] == "BOLD_RSSETSMarker_rss_ets" for x in features.values()
)

View file

@ -20,27 +20,27 @@ def test_base_marker_subclassing() -> None:
# Create concrete class # Create concrete class
class MyBaseMarker(BaseMarker): class MyBaseMarker(BaseMarker):
_MARKER_INOUT_MAPPINGS = { # noqa: RUF012
"BOLD": {
"feat_1": "timeseries",
},
}
def __init__(self, on, name=None) -> None: def __init__(self, on, name=None) -> None:
self.parameter = 1 self.parameter = 1
super().__init__(on, name) super().__init__(on, name)
def get_valid_inputs(self):
return ["BOLD", "T1w"]
def get_output_type(self, input):
if input == "BOLD":
return "timeseries"
raise ValueError(f"Cannot compute output type for {input}")
def compute(self, input, extra_input): def compute(self, input, extra_input):
return { return {
"data": "data", "feat_1": {
"columns": "columns", "data": "data",
"row_names": "row_names", "col_names": ["columns"],
},
} }
with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): with pytest.raises(ValueError, match=r"cannot be computed on \['T1w'\]"):
MyBaseMarker(on=["BOLD", "T2w"]) MyBaseMarker(on=["BOLD", "T1w"])
# Create input for marker # Create input for marker
input_ = { input_ = {
@ -64,12 +64,11 @@ def test_base_marker_subclassing() -> None:
output = marker.fit_transform(input=input_) # process output = marker.fit_transform(input=input_) # process
# Check output # Check output
assert "BOLD" in output assert "BOLD" in output
assert "data" in output["BOLD"] assert "data" in output["BOLD"]["feat_1"]
assert "columns" in output["BOLD"] assert "col_names" in output["BOLD"]["feat_1"]
assert "row_names" in output["BOLD"]
assert "meta" in output["BOLD"] assert "meta" in output["BOLD"]["feat_1"]
meta = output["BOLD"]["meta"] meta = output["BOLD"]["feat_1"]["meta"]
assert "datagrabber" in meta assert "datagrabber" in meta
assert "element" in meta assert "element" in meta
assert "datareader" in meta assert "datareader" in meta

View file

@ -23,16 +23,63 @@ from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
def test_ParcelAggregation_input_output() -> None: @pytest.mark.parametrize(
"""Test ParcelAggregation input and output types.""" "input_type, storage_type",
marker = ParcelAggregation( [
parcellation="Schaefer100x7", method="mean", on="VBM_GM" (
) "T1w",
for in_, out_ in [("VBM_GM", "vector"), ("BOLD", "timeseries")]: "vector",
assert marker.get_output_type(in_) == out_ ),
(
"T2w",
"vector",
),
(
"BOLD",
"timeseries",
),
(
"VBM_GM",
"vector",
),
(
"VBM_WM",
"vector",
),
(
"VBM_CSF",
"vector",
),
(
"fALFF",
"vector",
),
(
"GCOR",
"vector",
),
(
"LCOR",
"vector",
),
],
)
def test_ParcelAggregation_input_output(
input_type: str, storage_type: str
) -> None:
"""Test ParcelAggregation input and output types.
with pytest.raises(ValueError, match="Unknown input"): Parameters
marker.get_output_type("unknown") ----------
input_type : str
The parametrized input type.
storage_type : str
The parametrized storage type.
"""
assert storage_type == ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on=input_type
).get_output_type(input_type=input_type, output_feature="aggregation")
def test_ParcelAggregation_3D() -> None: def test_ParcelAggregation_3D() -> None:
@ -85,8 +132,8 @@ def test_ParcelAggregation_3D() -> None:
) )
parcel_agg_mean_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_mean_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
# Check that arrays are almost equal # Check that arrays are almost equal
assert_array_equal(parcel_agg_mean_bold_data, manual) assert_array_equal(parcel_agg_mean_bold_data, manual)
assert_array_almost_equal(nifti_labels_masked_bold, manual) assert_array_almost_equal(nifti_labels_masked_bold, manual)
@ -113,8 +160,8 @@ def test_ParcelAggregation_3D() -> None:
on="BOLD", on="BOLD",
) )
parcel_agg_std_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_std_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
assert parcel_agg_std_bold_data.ndim == 2 assert parcel_agg_std_bold_data.ndim == 2
assert parcel_agg_std_bold_data.shape[0] == 1 assert parcel_agg_std_bold_data.shape[0] == 1
assert_array_equal(parcel_agg_std_bold_data, manual) assert_array_equal(parcel_agg_std_bold_data, manual)
@ -139,7 +186,7 @@ def test_ParcelAggregation_3D() -> None:
) )
parcel_agg_trim_mean_bold_data = marker.fit_transform(element_data)[ parcel_agg_trim_mean_bold_data = marker.fit_transform(element_data)[
"BOLD" "BOLD"
]["data"] ]["aggregation"]["data"]
assert parcel_agg_trim_mean_bold_data.ndim == 2 assert parcel_agg_trim_mean_bold_data.ndim == 2
assert parcel_agg_trim_mean_bold_data.shape[0] == 1 assert parcel_agg_trim_mean_bold_data.shape[0] == 1
assert_array_equal(parcel_agg_trim_mean_bold_data, manual) assert_array_equal(parcel_agg_trim_mean_bold_data, manual)
@ -154,8 +201,8 @@ def test_ParcelAggregation_4D():
parcellation="TianxS1x3TxMNInonlinear2009cAsym", method="mean" parcellation="TianxS1x3TxMNInonlinear2009cAsym", method="mean"
) )
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Load testing parcellation # Load testing parcellation
@ -204,7 +251,8 @@ def test_ParcelAggregation_storage(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_ParcelAggregation" for x in features.values() x["name"] == "BOLD_ParcelAggregation_aggregation"
for x in features.values()
) )
# Store 4D # Store 4D
@ -221,7 +269,8 @@ def test_ParcelAggregation_storage(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_ParcelAggregation" for x in features.values() x["name"] == "BOLD_ParcelAggregation_aggregation"
for x in features.values()
) )
@ -241,8 +290,8 @@ def test_ParcelAggregation_3D_mask() -> None:
..., 0:1 ..., 0:1
] ]
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Load testing parcellation # Load testing parcellation
@ -316,8 +365,8 @@ def test_ParcelAggregation_3D_mask_computed() -> None:
on="BOLD", on="BOLD",
) )
parcel_agg_mean_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_mean_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
assert parcel_agg_mean_bold_data.ndim == 2 assert parcel_agg_mean_bold_data.ndim == 2
assert parcel_agg_mean_bold_data.shape[0] == 1 assert parcel_agg_mean_bold_data.shape[0] == 1
@ -397,7 +446,9 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
name="tian_mean", name="tian_mean",
on="BOLD", on="BOLD",
) )
orig_mean = marker_original.fit_transform(element_data)["BOLD"] orig_mean = marker_original.fit_transform(element_data)["BOLD"][
"aggregation"
]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2 assert orig_mean_data.ndim == 2
@ -417,7 +468,9 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
# No warnings should be raised # No warnings should be raised
with warnings.catch_warnings(): with warnings.catch_warnings():
warnings.simplefilter("error", category=UserWarning) warnings.simplefilter("error", category=UserWarning)
split_mean = marker_split.fit_transform(element_data)["BOLD"] split_mean = marker_split.fit_transform(element_data)["BOLD"][
"aggregation"
]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
@ -497,7 +550,9 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
name="tian_mean", name="tian_mean",
on="BOLD", on="BOLD",
) )
orig_mean = marker_original.fit_transform(element_data)["BOLD"] orig_mean = marker_original.fit_transform(element_data)["BOLD"][
"aggregation"
]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2 assert orig_mean_data.ndim == 2
@ -515,7 +570,9 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
) )
# Warning should be raised # Warning should be raised
with pytest.warns(RuntimeWarning, match="overlapping voxels"): with pytest.warns(RuntimeWarning, match="overlapping voxels"):
split_mean = marker_split.fit_transform(element_data)["BOLD"] split_mean = marker_split.fit_transform(element_data)["BOLD"][
"aggregation"
]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
@ -602,7 +659,9 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
name="tian_mean", name="tian_mean",
on="BOLD", on="BOLD",
) )
orig_mean = marker_original.fit_transform(element_data)["BOLD"] orig_mean = marker_original.fit_transform(element_data)["BOLD"][
"aggregation"
]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2 assert orig_mean_data.ndim == 2
@ -621,7 +680,9 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
# Warning should be raised # Warning should be raised
with pytest.warns(RuntimeWarning, match="duplicated labels."): with pytest.warns(RuntimeWarning, match="duplicated labels."):
split_mean = marker_split.fit_transform(element_data)["BOLD"] split_mean = marker_split.fit_transform(element_data)["BOLD"][
"aggregation"
]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
@ -653,8 +714,8 @@ def test_ParcelAggregation_4D_agg_time():
on="BOLD", on="BOLD",
) )
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Loading testing parcellation # Loading testing parcellation
@ -689,8 +750,8 @@ def test_ParcelAggregation_4D_agg_time():
on="BOLD", on="BOLD",
) )
parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ parcel_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
assert parcel_agg_bold_data.ndim == 2 assert parcel_agg_bold_data.ndim == 2
assert_array_equal( assert_array_equal(

View file

@ -25,14 +25,65 @@ COORDS = "DMNBuckner"
RADIUS = 8 RADIUS = 8
def test_SphereAggregation_input_output() -> None: @pytest.mark.parametrize(
"""Test SphereAggregation input and output types.""" "input_type, storage_type",
marker = SphereAggregation(coords="DMNBuckner", method="mean", on="VBM_GM") [
for in_, out_ in [("VBM_GM", "vector"), ("BOLD", "timeseries")]: (
assert marker.get_output_type(in_) == out_ "T1w",
"vector",
),
(
"T2w",
"vector",
),
(
"BOLD",
"timeseries",
),
(
"VBM_GM",
"vector",
),
(
"VBM_WM",
"vector",
),
(
"VBM_CSF",
"vector",
),
(
"fALFF",
"vector",
),
(
"GCOR",
"vector",
),
(
"LCOR",
"vector",
),
],
)
def test_SphereAggregation_input_output(
input_type: str, storage_type: str
) -> None:
"""Test SphereAggregation input and output types.
with pytest.raises(ValueError, match="Unknown input"): Parameters
marker.get_output_type("unknown") ----------
input_type : str
The parametrized input type.
storage_type : str
The parametrized storage type.
"""
assert storage_type == SphereAggregation(
coords="DMNBuckner",
method="mean",
on=input_type,
).get_output_type(input_type=input_type, output_feature="aggregation")
def test_SphereAggregation_3D() -> None: def test_SphereAggregation_3D() -> None:
@ -44,8 +95,8 @@ def test_SphereAggregation_3D() -> None:
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
) )
sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][ sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Load testing coordinates # Load testing coordinates
@ -76,8 +127,8 @@ def test_SphereAggregation_4D() -> None:
coords=COORDS, method="mean", radius=RADIUS, on="BOLD" coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
) )
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Load testing coordinates # Load testing coordinates
@ -120,7 +171,8 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "VBM_GM_SphereAggregation" for x in features.values() x["name"] == "VBM_GM_SphereAggregation_aggregation"
for x in features.values()
) )
# Store 4D # Store 4D
@ -135,7 +187,8 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
marker.fit_transform(input=element_data, storage=storage) marker.fit_transform(input=element_data, storage=storage)
features = storage.list_features() features = storage.list_features()
assert any( assert any(
x["name"] == "BOLD_SphereAggregation" for x in features.values() x["name"] == "BOLD_SphereAggregation_aggregation"
for x in features.values()
) )
@ -152,8 +205,8 @@ def test_SphereAggregation_3D_mask() -> None:
masks="compute_brain_mask", masks="compute_brain_mask",
) )
sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][ sphere_agg_vbm_gm_data = marker.fit_transform(element_data)["VBM_GM"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Load testing coordinates # Load testing coordinates
@ -195,8 +248,8 @@ def test_SphereAggregation_4D_agg_time() -> None:
on="BOLD", on="BOLD",
) )
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
# Compare with nilearn # Compare with nilearn
# Load testing coordinates # Load testing coordinates
@ -231,8 +284,8 @@ def test_SphereAggregation_4D_agg_time() -> None:
on="BOLD", on="BOLD",
) )
sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][ sphere_agg_bold_data = marker.fit_transform(element_data)["BOLD"][
"data" "aggregation"
] ]["data"]
assert sphere_agg_bold_data.ndim == 2 assert sphere_agg_bold_data.ndim == 2
assert_array_equal( assert_array_equal(

View file

@ -210,7 +210,17 @@ class PipelineStepMixin:
# Validate input # Validate input
fit_input = self.validate_input(input=input) fit_input = self.validate_input(input=input)
# Validate output type # Validate output type
outputs = [self.get_output_type(t_input) for t_input in fit_input] # Nested output type for marker
if hasattr(self, "_MARKER_INOUT_MAPPINGS"):
outputs = list(
{
val
for t_input in fit_input
for val in self._MARKER_INOUT_MAPPINGS[t_input].values()
}
)
else:
outputs = [self.get_output_type(t_input) for t_input in fit_input]
return outputs return outputs
def fit_transform( def fit_transform(

View file

@ -101,7 +101,7 @@ def test_get_class():
register(step="datagrabber", name="bar", klass=str) register(step="datagrabber", name="bar", klass=str)
# Get class # Get class
obj = get_class(step="datagrabber", name="bar") obj = get_class(step="datagrabber", name="bar")
assert obj == str assert isinstance(obj, type(str))
# TODO: possible parametrization? # TODO: possible parametrization?