[ENH]: Adopt Pydantic for schema validation #364

Merged
synchon merged 138 commits from feat/pydantic-validation into main 2026-05-26 11:07:28 +00:00
196 changed files with 5115 additions and 5268 deletions

View file

@ -5,5 +5,5 @@ On-the-fly
:members: :members:
:imported-members: :imported-members:
.. automodule:: junifer.onthefly.brainprint .. automodule:: junifer.onthefly._brainprint
:members: :members:

View file

@ -158,16 +158,16 @@ Available
| (subject-native or other template spaces) | (subject-native or other template spaces)
- Done - Done
- 0.0.4 - 0.0.4
* - ``Smoothing`` * - :class:`.Smoothing`
- | Apply smoothing to data, particularly useful when dealing with - | Apply smoothing to data, particularly useful when dealing with
| ``fMRIPrep``-ed data | ``fMRIPrep``-ed data
- In Progress - In Progress
- :gh:`161` - :gh:`161`
* - ``TemporalSlicer`` * - :class:`.TemporalSlicer`
- Slice ``BOLD`` data temporally - Slice ``BOLD`` data temporally
- | Done - | Done
- :gh:`443` - :gh:`443`
* - ``TemporalFilter`` * - :class:`.TemporalFilter`
- Filter (clean) ``BOLD`` data temporally - Filter (clean) ``BOLD`` data temporally
- | Done - | Done
- :gh:`432` - :gh:`432`

View file

@ -0,0 +1 @@
Add documentation on adding confounds format by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Adopt Pydantic for user-facing core objects to perform automatic validation using typing annotations by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Enable confounds format to be added by introducing :func:`.register_confounds_format` by `Synchon Mandal`_

View file

@ -65,6 +65,8 @@ extensions = [
"sphinx_copybutton", # copy button for code blocks "sphinx_copybutton", # copy button for code blocks
"sphinxcontrib.mermaid", # mermaid support "sphinxcontrib.mermaid", # mermaid support
"sphinxcontrib.towncrier.ext", # towncrier fragment support "sphinxcontrib.towncrier.ext", # towncrier fragment support
"sphinxcontrib.autodoc_pydantic", # autodoc support for pydantic models
"enum_tools.autoenum", # enum support
] ]
if use_multiversion: if use_multiversion:
@ -97,6 +99,15 @@ nitpick_ignore_regex = [
("py:class", "pipeline.Pipeline"), # nilearn ("py:class", "pipeline.Pipeline"), # nilearn
("py:obj", "neurokit2.*"), # ignore neurokit2 ("py:obj", "neurokit2.*"), # ignore neurokit2
("py:obj", "datalad.*"), # ignore datalad ("py:obj", "datalad.*"), # ignore datalad
("py:obj", "junifer.*"), # ignore junifer internal
("py:class", "annotated_types.*"), # ignore pydantic annotated types
("py:obj", "variants"), # ignore variants
("py:obj", "them"), # ignore them
("py:class", "junifer.utils.helpers.ensure_list"), # ignore ensure_list
("py:class", "junifer.utils.helpers.ensure_list_or_none"), # ignore ensure_list_or_none
("py:class", "PydanticUndefined"), # ignore PydanticUndefined
("py:class", "FieldInfo"), # ignore FieldInfo
("py:class", "NoneType"), # ignore NoneType
] ]
# -- Options for HTML output ------------------------------------------------- # -- Options for HTML output -------------------------------------------------
@ -154,6 +165,7 @@ intersphinx_mapping = {
"pandas": ("https://pandas.pydata.org/pandas-docs/dev", None), "pandas": ("https://pandas.pydata.org/pandas-docs/dev", None),
# "sqlalchemy": ("https://docs.sqlalchemy.org/en/20/", None), # "sqlalchemy": ("https://docs.sqlalchemy.org/en/20/", None),
"scipy": ("https://docs.scipy.org/doc/scipy/", None), "scipy": ("https://docs.scipy.org/doc/scipy/", None),
"pydantic": ("https://docs.pydantic.dev/latest/", None),
} }
# -- sphinx.ext.extlinks configuration --------------------------------------- # -- sphinx.ext.extlinks configuration ---------------------------------------

View file

@ -0,0 +1,22 @@
.. include:: ../links.inc
.. _adding_confounds_format:
Adding Confounds Format
=======================
#. Check :ref:`extending junifer <extending_extension>` on how to create a
*junifer extension* if you have not done so.
#. Register the confounds format before defining / using a DataGrabber like so:
.. code-block:: python
from junifer.datagrabber import register_confounds_format
...
# registers the confounds format as "confounds" and accessible as
# ``ConfoundsFormat.Confounds``
register_confounds_format(name="Confounds", alias="confounds")
...

View file

@ -140,29 +140,24 @@ With the variables defined above, we can create our DataGrabber and name it
from pathlib import Path from pathlib import Path
from junifer.datagrabber import PatternDataGrabber from junifer.datagrabber import PatternDataGrabber, DataType
from junifer.typing import DataGrabberPatterns
class ExampleBIDSDataGrabber(PatternDataGrabber): class ExampleBIDSDataGrabber(PatternDataGrabber):
def __init__(self, datadir: str | Path) -> None:
types = ["T1w", "BOLD"] types: list[DataType] = [DataType.T1w, DataType.BOLD]
patterns = { patterns: DataGrabberPatterns = {
"T1w": { "T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", "pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native", "space": "native",
}, },
"BOLD": { "BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz", "pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym", "space": "MNI152NLin6Asym",
}, },
} }
replacements = ["subject", "session"] replacements: list[str] = ["subject", "session"]
super().__init__(
datadir=datadir,
types=types,
patterns=patterns,
replacements=replacements,
)
Our DataGrabber is ready to be used by ``junifer``. However, it is still unknown Our DataGrabber is ready to be used by ``junifer``. However, it is still unknown
to the library. We need to register it in the library. To do so, we need to to the library. We need to register it in the library. To do so, we need to
@ -175,29 +170,24 @@ use the :func:`.register_datagrabber` decorator.
from junifer.api.decorators import register_datagrabber from junifer.api.decorators import register_datagrabber
from junifer.datagrabber import PatternDataGrabber from junifer.datagrabber import PatternDataGrabber
from junifer.typing import DataGrabberPatterns
@register_datagrabber @register_datagrabber
class ExampleBIDSDataGrabber(PatternDataGrabber): class ExampleBIDSDataGrabber(PatternDataGrabber):
def __init__(self, datadir: str | Path) -> None:
types = ["T1w", "BOLD"] types: list[DataType] = [DataType.T1w, DataType.BOLD]
patterns = { patterns: DataGrabberPatterns = {
"T1w": { "T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", "pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native", "space": "native",
}, },
"BOLD": { "BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz", "pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym", "space": "MNI152NLin6Asym",
}, },
} }
replacements = ["subject", "session"] replacements: list[str] = ["subject", "session"]
super().__init__(
datadir=datadir,
types=types,
patterns=patterns,
replacements=replacements,
)
Now, we can use our DataGrabber in ``junifer``, by setting the ``datagrabber`` Now, we can use our DataGrabber in ``junifer``, by setting the ``datagrabber``
@ -259,35 +249,30 @@ And we can create our DataGrabber:
.. code-block:: python .. code-block:: python
from pathlib import Path
from junifer.api.decorators import register_datagrabber from junifer.api.decorators import register_datagrabber
from junifer.datagrabber import PatternDataladDataGrabber from junifer.datagrabber import PatternDataladDataGrabber
from pydantic import AnyUrl
@register_datagrabber @register_datagrabber
class ExampleBIDSDataGrabber(PatternDataladDataGrabber): class ExampleBIDSDataGrabber(PatternDataladDataGrabber):
def __init__(self) -> None:
types = ["T1w", "BOLD"] uri: AnyUrl = "https://gin.g-node.org/juaml/datalad-example-bids"
patterns = { types: list[DataType] = ["T1w", "BOLD"]
"T1w": { patterns: DataGrabberPatterns = {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", "T1w": {
"space": "native", "pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
}, "space": "native",
"BOLD": { },
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz", "BOLD": {
"space": "MNI152NLin6Asym", "pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
}, "space": "MNI152NLin6Asym",
} },
replacements = ["subject", "session"] }
uri = "https://gin.g-node.org/juaml/datalad-example-bids" replacements: list[str] = ["subject", "session"]
rootdir = "example_bids_ses" rootdir: Path = "example_bids_ses"
super().__init__(
datadir=None,
uri=uri,
rootdir=rootdir,
types=types,
patterns=patterns,
replacements=replacements,
)
This approach can be used directly from the YAML, like so: This approach can be used directly from the YAML, like so:
@ -376,8 +361,8 @@ need to implement the following methods:
.. note:: .. note::
The ``__init__`` method could also be implemented, but it is not mandatory. If the DataGrabber requires any extra parameter, they could be defined as
This is required if the DataGrabber requires any extra parameter. class attributes.
We will now implement our BIDS example with this method. We will now implement our BIDS example with this method.
@ -494,8 +479,8 @@ more information about the format of the confounds file. Thus, the
``BOLD.confounds`` element is a dictionary with the following keys: ``BOLD.confounds`` element is a dictionary with the following keys:
- ``path``: the path to the confounds file. - ``path``: the path to the confounds file.
- ``format``: the format of the confounds file. Currently, this can be either - ``format``: the format of the confounds file. Check :enum:`.ConfoundsFormat`
``fmriprep`` or ``adhoc``. for options.
The ``fmriprep`` format corresponds to the format of the confounds files The ``fmriprep`` format corresponds to the format of the confounds files
generated by `fMRIPrep`_. The ``adhoc`` format corresponds to a format that is generated by `fMRIPrep`_. The ``adhoc`` format corresponds to a format that is

View file

@ -46,7 +46,7 @@ by having a class attribute like so:
_EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [ _EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [
{ {
"name": "afni", "name": ExtDep.AFNI,
"commands": ["3dReHo", "3dAFNItoNIFTI"], "commands": ["3dReHo", "3dAFNItoNIFTI"],
}, },
] ]
@ -55,7 +55,7 @@ The above example is taken from the class which computes regional homogeneity
(ReHo) using AFNI. The general pattern is that you need to have the value of (ReHo) using AFNI. The general pattern is that you need to have the value of
``_EXT_DEPENDENCIES`` as a list of dictionary with two keys: ``_EXT_DEPENDENCIES`` as a list of dictionary with two keys:
* ``name`` (str) : lowercased name of the toolbox * ``name`` (:enum:`.ExtDep`) : name of the toolbox
* ``commands`` (list of str) : actual names of the commands you need to use * ``commands`` (list of str) : actual names of the commands you need to use
This is simple but powerful as we will see in the following sub-sections. This is simple but powerful as we will see in the following sub-sections.
@ -81,11 +81,11 @@ that it shows the problem a bit better and how we solve it:
_CONDITIONAL_DEPENDENCIES: ClassVar[ConditionalDependencies] = [ _CONDITIONAL_DEPENDENCIES: ClassVar[ConditionalDependencies] = [
{ {
"using": "fsl", "using": "fsl",
"depends_on": FSLWarper, "depends_on": [FSLWarper],
}, },
{ {
"using": "ants", "using": "ants",
"depends_on": ANTsWarper, "depends_on": [ANTSWarper],
}, },
{ {
"using": "auto", "using": "auto",
@ -93,18 +93,16 @@ that it shows the problem a bit better and how we solve it:
}, },
] ]
def __init__( using: str
self, using: str, reference: str, on: Union[List[str], str] reference: str
) -> None: on: List[DataType]
# validation and setting up
...
Here, you see a new class attribute ``_CONDITIONAL_DEPENDENCIES`` which is a Here, you see a new class attribute ``_CONDITIONAL_DEPENDENCIES`` which is a
list of dictionaries with two keys: list of dictionaries with two keys:
* ``using`` (str) : lowercased name of the toolbox * ``using`` (str) : lowercased name of the toolbox
* ``depends_on`` (object or list of objects) : a class or list of classes which \ * ``depends_on`` (list of objects) : list of classes which \
implements the particular tool's use implements the particular tool's use
It is mandatory to have the ``using`` positional argument in the constructor in It is mandatory to have the ``using`` positional argument in the constructor in
@ -128,7 +126,7 @@ similar. ``FSLWarper`` looks like this (only the relevant part is shown here):
_EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [ _EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [
{ {
"name": "fsl", "name": ExtDep.FSL,
"commands": ["flirt", "applywarp"], "commands": ["flirt", "applywarp"],
}, },
] ]

View file

@ -34,4 +34,5 @@ DataGrabbers, Preprocessors, Markers, etc., following the *junifer* way.
plugins plugins
data_registries data_registries
data_types data_types
confounds_format
data_dump_asset data_dump_asset

View file

@ -14,7 +14,7 @@ Most of the functionality of a ``junifer`` Marker has been taken care by the
:class:`.BaseMarker` class. Thus, only a few methods and class attributes are :class:`.BaseMarker` class. Thus, only a few methods and class attributes are
required: required:
#. ``__init__``: The initialisation method, where the Marker is configured. #. (optional) ``validate_marker_params``: The method to perform logical validation of parameters (if required).
#. ``compute``: The method that given the data, computes the Marker. #. ``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
@ -29,8 +29,8 @@ 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. The output of the Marker depends on ``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 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: inputs, it will be ``Vector``. Thus, we have a class attribute like so:
.. code-block:: python .. code-block:: python
@ -38,14 +38,14 @@ inputs, it will be ``vector``. Thus, we have a class attribute like so:
# You can have multiple features for one data type, # You can have multiple features for one data type,
# each feature having same or different storage type # each feature having same or different storage type
_MARKER_INOUT_MAPPINGS = { _MARKER_INOUT_MAPPINGS = {
"BOLD": { DataType.BOLD: {
"parcel_mean": "timeseries", "parcel_mean": StorageType.Timeseries,
}, },
"VBM_WM": { DataType.VBM_WM: {
"parcel_mean": "vector", "parcel_mean": StorageType.Vector,
}, },
"VBM_GM": { DataType.VBM_GM: {
"parcel_mean": "vector", "parcel_mean": StorageType.Vector,
}, },
} }
@ -57,13 +57,11 @@ Step 2: Initialise the Marker
In this step we need to define the parameters of the Marker the user can provide In this step we need to define the parameters of the Marker the user can provide
to configure how the Marker will behave. to configure how the Marker will behave.
The parameters of the Marker are defined in the ``__init__`` method. The The parameters of the Marker are defined as class attributes. The
:class:`.BaseMarker` class requires two optional parameters: :class:`.BaseMarker` class defines two optional parameters:
1. ``name``: the name of the Marker. This is used to identify the Marker in the 1. ``name``: the name of the Marker. This is used to identify the Marker in the configuration file.
configuration file. 2. ``on``: a list of :enum:`.DataType` with the data types that the Marker will be applied to.
2. ``on``: a list or string with the data types that the Marker will be applied
to.
.. attention:: .. attention::
@ -72,18 +70,11 @@ The parameters of the Marker are defined in the ``__init__`` method. The
JSON format, and JSON only supports these types. JSON format, and JSON only supports these types.
In this example, only parameter required for the computation is the name of the In this example, only parameter required for the computation is the name of the
parcellation to use. Thus, we can define the ``__init__`` method as follows: parcellation to use. Thus, we can define as follows:
.. code-block:: python .. code-block:: python
def __init__( parcellation: str
self,
parcellation: str,
on: str | list[str] | None = None,
name: str | None = None,
) -> None:
self.parcellation = parcellation
super().__init__(on=on, name=name)
.. caution:: .. caution::
@ -121,14 +112,14 @@ and the values would be a dictionary of storage type specific key-value pairs.
To simplify the ``store`` method, define keys of the dictionary based on the To simplify the ``store`` method, define keys of the dictionary based on the
corresponding store functions in the :ref:`storage types <storage_types>`. corresponding store functions in the :ref:`storage types <storage_types>`.
For example, if the output is a ``vector``, the keys of the dictionary should For example, if the output is a ``Vector``, the keys of the dictionary should
be ``data`` and ``col_names``. be ``data`` and ``col_names``.
.. code-block:: python .. code-block:: python
from typing import Any from typing import Any
from junifer.data import get_parcellation from junifer.data import get_data
from nilearn.maskers import NiftiLabelsMasker from nilearn.maskers import NiftiLabelsMasker
@ -141,8 +132,9 @@ and the values would be a dictionary of storage type specific key-value pairs.
data = input["data"] data = input["data"]
# Get the parcellation tailored for the target # Get the parcellation tailored for the target
t_parcellation, t_labels, _ = get_parcellation( t_parcellation, t_labels, _ = get_data(
name=self.parcellation_name, kind="parcellation",
name=[self.parcellation],
target_data=input, target_data=input,
extra_input=extra_input, extra_input=extra_input,
) )
@ -194,8 +186,10 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
from typing import Any, ClassVar 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_data
from junifer.datagrabber import DataType
from junifer.markers import BaseMarker from junifer.markers import BaseMarker
from junifer.storage import StorageType
from junifer.typing import Dependencies, MarkerInOutMappings from junifer.typing import Dependencies, MarkerInOutMappings
from nilearn.maskers import NiftiLabelsMasker from nilearn.maskers import NiftiLabelsMasker
@ -206,25 +200,18 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
_DEPENDENCIES: ClassVar[Dependencies] = {"nilearn", "numpy"} _DEPENDENCIES: ClassVar[Dependencies] = {"nilearn", "numpy"}
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = { _MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"BOLD": { DataType.BOLD: {
"parcel_mean": "timeseries", "parcel_mean": StorageType.Timeseries,
}, },
"VBM_WM": { DataType.VBM_WM: {
"parcel_mean": "vector", "parcel_mean": StorageType.Vector,
}, },
"VBM_GM": { DataType.VBM_GM: {
"parcel_mean": "vector", "parcel_mean": StorageType.Vector,
}, },
} }
def __init__( parcellation: str
self,
parcellation: str,
on: str | list[str] | None = None,
name: str | None = None,
) -> None:
self.parcellation = parcellation
super().__init__(on=on, name=name)
def compute( def compute(
self, self,
@ -235,8 +222,9 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
data = input["data"] data = input["data"]
# Get the parcellation tailored for the target # Get the parcellation tailored for the target
t_parcellation, t_labels, _ = get_parcellation( t_parcellation, t_labels, _ = get_data(
name=self.parcellation_name, kind="parcellation",
name=[self.parcellation],
target_data=input, target_data=input,
extra_input=extra_input, extra_input=extra_input,
) )
@ -280,9 +268,13 @@ Template for a custom Marker
# TODO: add the input-output mappings # TODO: add the input-output mappings
_MARKER_INOUT_MAPPINGS = {} _MARKER_INOUT_MAPPINGS = {}
def __init__(self, on=None, name=None): # TODO: define marker-specific parameters
# TODO: add marker-specific parameters
super().__init__(on=on, name=name) # optional
def validate_marker_params(self):
# TODO: add validation logic for marker parameters
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
return {}

View file

@ -124,7 +124,7 @@ parcellation when registering it. For example, we can add a
markers: markers:
fraimondo commented 2025-11-26 10:46:28 +00:00 (Migrated from github.com)

Same here: str or list[str]

Same here: str or list[str]
- name: CustomParcellation_mean - name: CustomParcellation_mean
kind: ParcelAggregation kind: ParcelAggregation
parcellation: my_custom_parcellation parcellation: <my_custom_parcellation>
method: mean method: mean
Now, you can simply use this YAML file to run your pipeline. Now, you can simply use this YAML file to run your pipeline.

View file

@ -16,8 +16,7 @@ own Preprocessor.
While implementing your own Preprocessor, you need to always inherit from While implementing your own Preprocessor, you need to always inherit from
:class:`.BasePreprocessor` and implement a few methods and class attributes: :class:`.BasePreprocessor` and implement a few methods and class attributes:
#. ``__init__``: The initialisation method, where the Preprocessor is #. (optional) ``validate_preprocessor_params``: The method to perform logical validation of parameters (if required).
configured.
#. ``preprocess``: The method that given the data, preprocesses the data. #. ``preprocess``: The method that given the data, preprocesses the data.
As an example, we will develop a ``NilearnSmoothing`` Preprocessor, which As an example, we will develop a ``NilearnSmoothing`` Preprocessor, which
@ -43,8 +42,8 @@ For input we can accept ``T1w``, ``T2w`` and ``BOLD``
Step 2: Initialise the Preprocessor Step 2: Initialise the Preprocessor
----------------------------------- -----------------------------------
Now we need to define our Preprocessor class' constructor which is also how Now we need to define our Preprocessor class' parameters as class attributes.
you configure it. Our class will have the following arguments: Our class will have the following:
1. ``fwhm``: The smoothing strength as a full-width at half maximum 1. ``fwhm``: The smoothing strength as a full-width at half maximum
(in millimetres). Since we depend on :func:`nilearn.image.smooth_img`, we (in millimetres). Since we depend on :func:`nilearn.image.smooth_img`, we
@ -59,6 +58,8 @@ you configure it. Our class will have the following arguments:
are allowed as parameters. This is because the parameters are stored in are allowed as parameters. This is because the parameters are stored in
JSON format, and JSON only supports these types. JSON format, and JSON only supports these types.
As :class:`.BasePreprocessor` already defines ``on``, we can define the other:
.. code-block:: python .. code-block:: python
from typing import Literal from typing import Literal
@ -68,15 +69,7 @@ you configure it. Our class will have the following arguments:
... ...
fwhm: int | float | ArrayLike | Literal["fast"] | None
def __init__(
self,
fwhm: int | float | ArrayLike | Literal["fast"] | None,
on: str | list[str] | None = None,
) -> None:
self.fwhm = fwhm
super().__init__(on=on)
... ...
@ -165,15 +158,9 @@ decorator and our final code should look like this:
_DEPENDENCIES = {"nilearn"} _DEPENDENCIES = {"nilearn"}
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["T1w", "T2w", "BOLD"] _VALID_DATA_TYPES: ClassVar[Sequence[DataType]] = ["T1w", "T2w", "BOLD"]
def __init__( fwhm: int | float | ArrayLike | Literal["fast"] | None
self,
fwhm: int | float | ArrayLike | Literal["fast"] | None,
on: str | list[str] | None = None,
) -> None:
self.fwhm = fwhm
super().__init__(on=on)
def preprocess( def preprocess(
self, self,
@ -191,7 +178,11 @@ Template for a custom Preprocessor
.. code-block:: python .. code-block:: python
from collections.abc import Sequence
from typing import ClassVar
from junifer.api.decorators import register_preprocessor from junifer.api.decorators import register_preprocessor
from junifer.datagrabber import DataType
from junifer.preprocess import BasePreprocessor from junifer.preprocess import BasePreprocessor
@ -202,11 +193,14 @@ Template for a custom Preprocessor
_DEPENDENCIES = {} _DEPENDENCIES = {}
# TODO: add the inputs # TODO: add the inputs
fraimondo commented 2026-05-26 08:12:48 +00:00 (Migrated from github.com)

We can already leave the typing in the template here

We can already leave the typing in the template here
synchon commented 2026-05-26 08:44:51 +00:00 (Migrated from github.com)

What do you mean?

What do you mean?
fraimondo commented 2026-05-26 08:54:51 +00:00 (Migrated from github.com)

This is supposed to be a template to create your own preprocessor:

_VALID_DATA_TYPES: ClassVar[Sequence[DataType]] = []
This is supposed to be a template to create your own preprocessor: ``` _VALID_DATA_TYPES: ClassVar[Sequence[DataType]] = [] ```
synchon commented 2026-05-26 09:05:40 +00:00 (Migrated from github.com)

Addressed in latest commit.

Addressed in latest commit.
_VALID_DATA_TYPES = [] _VALID_DATA_TYPES: ClassVar[Sequence[DataType]] = []
def __init__(self, on=None): # TODO: define preprocessor-specific parameters
# TODO: add preprocessor-specific parameters
super().__init__(on=on) # optional
def validate_preprocessor_params(self):
# TODO: add validation logic for preprocessor parameters
pass
def preprocess(self, input, extra_input): def preprocess(self, input, extra_input):
# TODO: add the preprocessor logic # TODO: add the preprocessor logic

View file

@ -265,7 +265,7 @@ Features
^^^^^^^^ ^^^^^^^^
- Introduce :func:`.normalize` and :func:`.reweight` functions for downstream - Introduce :func:`.normalize` and :func:`.reweight` functions for downstream
BrainPrint analysis in :mod:`.onthefly.brainprint` by `Synchon Mandal`_ BrainPrint analysis in :mod:`.onthefly._brainprint` by `Synchon Mandal`_
(:gh:`354`) (:gh:`354`)
- Introduce :class:`junifer.pipeline.PipelineComponentRegistry` to centralise - Introduce :class:`junifer.pipeline.PipelineComponentRegistry` to centralise
pipeline component management by `Synchon Mandal`_ (:gh:`362`) pipeline component management by `Synchon Mandal`_ (:gh:`362`)

View file

@ -15,8 +15,10 @@ from junifer.testing.datagrabbers import (
OasisVBMTestingDataGrabber, OasisVBMTestingDataGrabber,
SPMAuditoryTestingDataGrabber, SPMAuditoryTestingDataGrabber,
) )
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers import ParcelAggregation from junifer.markers import ParcelAggregation
from junifer.stats import AggFunc
from junifer.utils import configure_logging from junifer.utils import configure_logging
@ -32,7 +34,10 @@ with OasisVBMTestingDataGrabber() as dg:
# Read the element # Read the element
element_data = DefaultDataReader().fit_transform(dg[element]) element_data = DefaultDataReader().fit_transform(dg[element])
# Initialize marker # Initialize marker
marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean") marker = ParcelAggregation(
parcellation="Schaefer100x7",
method=AggFunc.Mean,
)
# Compute feature # Compute feature
feature = marker.fit_transform(element_data) feature = marker.fit_transform(element_data)
# Print the output # Print the output
@ -48,7 +53,9 @@ with SPMAuditoryTestingDataGrabber() as dg:
element_data = DefaultDataReader().fit_transform(dg[element]) element_data = DefaultDataReader().fit_transform(dg[element])
# Initialize marker # Initialize marker
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="BOLD" parcellation="Schaefer100x7",
method=AggFunc.Mean,
on=[DataType.BOLD],
) )
# Compute feature # Compute feature
feature = marker.fit_transform(element_data) feature = marker.fit_transform(element_data)

View file

@ -10,7 +10,7 @@ Authors: Federico Raimondo
License: BSD 3 clause License: BSD 3 clause
""" """
from junifer.datagrabber import PatternDataladDataGrabber from junifer.datagrabber import DataType, PatternDataladDataGrabber
from junifer.utils import configure_logging from junifer.utils import configure_logging
@ -23,7 +23,7 @@ configure_logging(level="INFO")
# The BIDS DataGrabber requires three parameters: the types of data we want, # The BIDS DataGrabber requires three parameters: the types of data we want,
# the specific pattern that matches each type, and the variables that will be # the specific pattern that matches each type, and the variables that will be
# replaced in the patterns. # replaced in the patterns.
types = ["T1w", "BOLD"] types = [DataType.T1w, DataType.BOLD]
patterns = { patterns = {
"T1w": { "T1w": {
"pattern": "{subject}/anat/{subject}_T1w.nii.gz", "pattern": "{subject}/anat/{subject}_T1w.nii.gz",

View file

@ -20,7 +20,7 @@ from julearn import run_cross_validation, PipelineCreator
import junifer.testing.registry # noqa: F401 import junifer.testing.registry # noqa: F401
from junifer.api import collect, run from junifer.api import collect, run
from junifer.storage.sqlite import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.utils import configure_logging from junifer.utils import configure_logging

View file

@ -5,6 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import atexit
import os import os
import shutil import shutil
from pathlib import Path from pathlib import Path
@ -182,6 +183,7 @@ def run(
"elements will be processed" "elements will be processed"
) )
WorkDirManager(**workdir) WorkDirManager(**workdir)
atexit.register(WorkDirManager()._cleanup)
# Get datagrabber to use # Get datagrabber to use
datagrabber_object = _get_datagrabber(datagrabber.copy()) datagrabber_object = _get_datagrabber(datagrabber.copy())

View file

@ -1,5 +1,13 @@
__all__ = ["QueueContextAdapter", "HTCondorAdapter", "GnuParallelLocalAdapter"] __all__ = [
"QueueContextAdapter",
"EnvKind",
"EnvShell",
"QueueContextEnv",
"HTCondorAdapter",
"HTCondorCollect",
"GnuParallelLocalAdapter",
]
from .queue_context_adapter import QueueContextAdapter from .queue_context_adapter import QueueContextAdapter, EnvKind, EnvShell, QueueContextEnv
from .htcondor_adapter import HTCondorAdapter from .htcondor_adapter import HTCondorAdapter, HTCondorCollect
from .gnu_parallel_local_adapter import GnuParallelLocalAdapter from .gnu_parallel_local_adapter import GnuParallelLocalAdapter

View file

@ -6,10 +6,16 @@
import shutil import shutil
import textwrap import textwrap
from pathlib import Path from pathlib import Path
from typing import Any
from ...typing import Elements from ...typing import Elements
from ...utils import logger, make_executable, raise_error, run_ext_cmd from ...utils import logger, make_executable, raise_error, run_ext_cmd
from .queue_context_adapter import QueueContextAdapter from .queue_context_adapter import (
EnvKind,
EnvShell,
QueueContextAdapter,
QueueContextEnv,
)
__all__ = ["GnuParallelLocalAdapter"] __all__ = ["GnuParallelLocalAdapter"]
@ -26,14 +32,14 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
The path to the job directory. The path to the job directory.
yaml_config_path : pathlib.Path yaml_config_path : pathlib.Path
The path to the YAML config file. The path to the YAML config file.
elements : list of str or tuple elements : ``Elements``
Element(s) to process. Will be used to index the DataGrabber. Element(s) to process. Will be used to index the DataGrabber.
pre_run : str or None, optional pre_run_cmds : str or None, optional
Extra shell commands to source before the run (default None). Extra shell commands to source before the run (default None).
pre_collect : str or None, optional pre_collect_cmds : str or None, optional
Extra bash commands to source before the collect (default None). Extra shell commands to source before the collect (default None).
env : dict, optional env : :class:`.QueueContextEnv` or None, optional
The Python environment configuration. If None, will run without a The environment configuration. If None, will run without a
virtual environment of any kind (default None). virtual environment of any kind (default None).
verbose : str, optional verbose : str, optional
The level of verbosity (default "info"). The level of verbosity (default "info").
@ -43,12 +49,6 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
submit : bool, optional submit : bool, optional
Whether to submit the jobs (default False). Whether to submit the jobs (default False).
Raises
------
ValueError
If ``env.kind`` is invalid or
if ``env.shell`` is invalid.
See Also See Also
-------- --------
QueueContextAdapter : QueueContextAdapter :
@ -58,87 +58,44 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
""" """
def __init__( job_name: str
self, job_dir: Path
job_name: str, yaml_config_path: Path
job_dir: Path, elements: Elements
yaml_config_path: Path, pre_run_cmds: str | None = None
elements: Elements, pre_collect_cmds: str | None = None
pre_run: str | None = None, env: QueueContextEnv | None = None
pre_collect: str | None = None, verbose: str = "info"
env: dict[str, str] | None = None, verbose_datalad: str | None = None
verbose: str = "info", submit: bool = False
verbose_datalad: str | None = None,
submit: bool = False,
) -> None:
"""Initialize the class."""
self._job_name = job_name
self._job_dir = job_dir
self._yaml_config_path = yaml_config_path
self._elements = elements
self._pre_run = pre_run
self._pre_collect = pre_collect
self._check_env(env)
self._verbose = verbose
self._verbose_datalad = verbose_datalad
self._submit = submit
self._log_dir = self._job_dir / "logs" def model_post_init(self, context: Any): # noqa: D102
self._pre_run_path = self._job_dir / "pre_run.sh" if self.env is None:
self._pre_collect_path = self._job_dir / "pre_collect.sh" self.env = QueueContextEnv(
self._run_path = self._job_dir / f"run_{self._job_name}.sh" kind=EnvKind.Local, shell=EnvShell.Bash, name=""
self._collect_path = self._job_dir / f"collect_{self._job_name}.sh"
self._run_joblog_path = self._job_dir / f"run_{self._job_name}_joblog"
self._elements_file_path = self._job_dir / "elements"
def _check_env(self, env: dict[str, str] | None) -> None:
"""Check value of env parameter on init.
Parameters
----------
env : dict or None
The value of env parameter.
Raises
------
ValueError
If ``env.kind`` is invalid.
"""
# Set env related variables
if env is None:
env = {"kind": "local"}
# Check env kind
valid_env_kinds = ["conda", "venv", "local"]
if env["kind"] not in valid_env_kinds:
raise_error(
f"Invalid value for `env.kind`: {env['kind']}, "
f"must be one of {valid_env_kinds}"
) )
if self.env["kind"] == EnvKind.Local:
# No virtual environment
self._executable = "junifer"
self._arguments = ""
else: else:
# Check shell if self.env["name"] is None:
shell = env.get("shell", "bash") raise_error("`env.name` is required")
valid_shells = ["bash", "zsh"] self._executable = f"run_{self.env['kind']}.{self.env['shell']}"
if shell not in valid_shells: self._arguments = f"{self.env['name']} junifer"
raise_error( self._exec_path = self.job_dir / self._executable
f"Invalid value for `env.shell`: {shell}, " self._log_dir = self.job_dir / "logs"
f"must be one of {valid_shells}" self._pre_run_path = self.job_dir / "pre_run.sh"
) self._pre_collect_path = self.job_dir / "pre_collect.sh"
self._shell = shell self._run_path = self.job_dir / f"run_{self.job_name}.sh"
# Set variables self._collect_path = self.job_dir / f"collect_{self.job_name}.sh"
if env["kind"] == "local": self._run_joblog_path = self.job_dir / f"run_{self.job_name}_joblog"
# No virtual environment self._elements_file_path = self.job_dir / "elements"
self._executable = "junifer"
self._arguments = ""
else:
self._executable = f"run_{env['kind']}.{self._shell}"
self._arguments = f"{env['name']} junifer"
self._exec_path = self._job_dir / self._executable
def elements(self) -> str: def elements_to_run(self) -> str:
"""Return elements to run.""" """Return elements to run."""
elements_to_run = [] elements_to_run = []
for element in self._elements: for element in self.elements:
# Stringify elements if tuple for operation # Stringify elements if tuple for operation
str_element = ( str_element = (
",".join(element) if isinstance(element, tuple) else element ",".join(element) if isinstance(element, tuple) else element
@ -150,23 +107,23 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
def pre_run(self) -> str: def pre_run(self) -> str:
"""Return pre-run commands.""" """Return pre-run commands."""
fixed = ( fixed = (
f"#!/usr/bin/env {self._shell}\n\n" f"#!/usr/bin/env {self.env['shell']}\n\n"
"# This script is auto-generated by junifer.\n\n" "# This script is auto-generated by junifer.\n\n"
"# Force datalad to run in non-interactive mode\n" "# Force datalad to run in non-interactive mode\n"
"DATALAD_UI_INTERACTIVE=false\n" "DATALAD_UI_INTERACTIVE=false\n"
) )
var = self._pre_run or "" var = self.pre_run_cmds or ""
return fixed + "\n" + var return fixed + "\n" + var
def run(self) -> str: def run(self) -> str:
"""Return run commands.""" """Return run commands."""
verbose_args = f"--verbose {self._verbose}" verbose_args = f"--verbose {self.verbose}"
if self._verbose_datalad: if self.verbose_datalad:
verbose_args = ( verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad}" f"{verbose_args} --verbose-datalad {self.verbose_datalad}"
) )
return ( return (
f"#!/usr/bin/env {self._shell}\n\n" f"#!/usr/bin/env {self.env['shell']}\n\n"
"# This script is auto-generated by junifer.\n\n" "# This script is auto-generated by junifer.\n\n"
"# Run pre_run.sh\n" "# Run pre_run.sh\n"
f"sh {self._pre_run_path.resolve()!s}\n\n" f"sh {self._pre_run_path.resolve()!s}\n\n"
@ -176,9 +133,9 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
"--delay 60 " # wait 1 min before next job is spawned "--delay 60 " # wait 1 min before next job is spawned
f"--results {self._log_dir} " f"--results {self._log_dir} "
f"--arg-file {self._elements_file_path.resolve()!s} " f"--arg-file {self._elements_file_path.resolve()!s} "
f"{self._job_dir.resolve()!s}/{self._executable} " f"{self.job_dir.resolve()!s}/{self._executable} "
f"{self._arguments} run " f"{self._arguments} run "
f"{self._yaml_config_path.resolve()!s} " f"{self.yaml_config_path.resolve()!s} "
f"{verbose_args} " f"{verbose_args} "
f"--element" f"--element"
) )
@ -186,28 +143,28 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
def pre_collect(self) -> str: def pre_collect(self) -> str:
"""Return pre-collect commands.""" """Return pre-collect commands."""
fixed = ( fixed = (
f"#!/usr/bin/env {self._shell}\n\n" f"#!/usr/bin/env {self.env['shell']}\n\n"
"# This script is auto-generated by junifer.\n" "# This script is auto-generated by junifer.\n"
) )
var = self._pre_collect or "" var = self.pre_collect_cmds or ""
return fixed + "\n" + var return fixed + "\n" + var
def collect(self) -> str: def collect(self) -> str:
"""Return collect commands.""" """Return collect commands."""
verbose_args = f"--verbose {self._verbose}" verbose_args = f"--verbose {self.verbose}"
if self._verbose_datalad: if self.verbose_datalad:
verbose_args = ( verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad}" f"{verbose_args} --verbose-datalad {self.verbose_datalad}"
) )
return ( return (
f"#!/usr/bin/env {self._shell}\n\n" f"#!/usr/bin/env {self.env['shell']}\n\n"
"# This script is auto-generated by junifer.\n\n" "# This script is auto-generated by junifer.\n\n"
"# Run pre_collect.sh\n" "# Run pre_collect.sh\n"
f"sh {self._pre_collect_path.resolve()!s}\n\n" f"sh {self._pre_collect_path.resolve()!s}\n\n"
"# Run `junifer collect`\n" "# Run `junifer collect`\n"
f"{self._job_dir.resolve()!s}/{self._executable} " f"{self.job_dir.resolve()!s}/{self._executable} "
f"{self._arguments} collect " f"{self._arguments} collect "
f"{self._yaml_config_path.resolve()!s} " f"{self.yaml_config_path.resolve()!s} "
f"{verbose_args}" f"{verbose_args}"
) )
@ -230,17 +187,19 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
f"{self._elements_file_path.resolve()!s}" f"{self._elements_file_path.resolve()!s}"
) )
self._elements_file_path.touch() self._elements_file_path.touch()
self._elements_file_path.write_text(textwrap.dedent(self.elements())) self._elements_file_path.write_text(
textwrap.dedent(self.elements_to_run())
)
# Create pre run # Create pre run
logger.info( logger.info(
f"Writing {self._pre_run_path.name} to {self._job_dir.resolve()!s}" f"Writing {self._pre_run_path.name} to {self.job_dir.resolve()!s}"
) )
self._pre_run_path.touch() self._pre_run_path.touch()
self._pre_run_path.write_text(textwrap.dedent(self.pre_run())) self._pre_run_path.write_text(textwrap.dedent(self.pre_run()))
make_executable(self._pre_run_path) make_executable(self._pre_run_path)
# Create run # Create run
logger.info( logger.info(
f"Writing {self._run_path.name} to {self._job_dir.resolve()!s}" f"Writing {self._run_path.name} to {self.job_dir.resolve()!s}"
) )
self._run_path.touch() self._run_path.touch()
self._run_path.write_text(textwrap.dedent(self.run())) self._run_path.write_text(textwrap.dedent(self.run()))
@ -248,14 +207,14 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
# Create pre collect # Create pre collect
logger.info( logger.info(
f"Writing {self._pre_collect_path.name} to " f"Writing {self._pre_collect_path.name} to "
f"{self._job_dir.resolve()!s}" f"{self.job_dir.resolve()!s}"
) )
self._pre_collect_path.touch() self._pre_collect_path.touch()
self._pre_collect_path.write_text(textwrap.dedent(self.pre_collect())) self._pre_collect_path.write_text(textwrap.dedent(self.pre_collect()))
make_executable(self._pre_collect_path) make_executable(self._pre_collect_path)
# Create collect # Create collect
logger.info( logger.info(
f"Writing {self._collect_path.name} to {self._job_dir.resolve()!s}" f"Writing {self._collect_path.name} to {self.job_dir.resolve()!s}"
) )
self._collect_path.touch() self._collect_path.touch()
self._collect_path.write_text(textwrap.dedent(self.collect())) self._collect_path.write_text(textwrap.dedent(self.collect()))
@ -263,7 +222,7 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
# Submit if required # Submit if required
run_cmd = f"sh {self._run_path.resolve()!s}" run_cmd = f"sh {self._run_path.resolve()!s}"
collect_cmd = f"sh {self._collect_path.resolve()!s}" collect_cmd = f"sh {self._collect_path.resolve()!s}"
if self._submit: if self.submit:
logger.info( logger.info(
"Shell scripts created, the following will be run:\n" "Shell scripts created, the following will be run:\n"
f"{run_cmd}\n" f"{run_cmd}\n"

View file

@ -5,14 +5,37 @@
import shutil import shutil
import textwrap import textwrap
from enum import Enum
from pathlib import Path from pathlib import Path
from typing import Any
from ...typing import Elements from ...typing import Elements
from ...utils import logger, make_executable, raise_error, run_ext_cmd from ...utils import logger, make_executable, raise_error, run_ext_cmd
from .queue_context_adapter import QueueContextAdapter from .queue_context_adapter import (
EnvKind,
EnvShell,
QueueContextAdapter,
QueueContextEnv,
)
__all__ = ["HTCondorAdapter"] __all__ = ["HTCondorAdapter", "HTCondorCollect"]
class HTCondorCollect(str, Enum):
"""Accepted HTCondor collect commands.
* ``"yes"``: Submit "collect" task and run even if some of the jobs
fail.
* ``"on_success_only"``: Submit "collect" task and run only if all jobs
succeed.
* ``"no"``: Do not submit "collect" task.
"""
Yes = "yes"
No = "no"
OnSuccessOnly = "on_success_only"
class HTCondorAdapter(QueueContextAdapter): class HTCondorAdapter(QueueContextAdapter):
@ -26,14 +49,14 @@ class HTCondorAdapter(QueueContextAdapter):
The path to the job directory. The path to the job directory.
yaml_config_path : pathlib.Path yaml_config_path : pathlib.Path
The path to the YAML config file. The path to the YAML config file.
elements : list of str or tuple elements : ``Elements``
Element(s) to process. Will be used to index the DataGrabber. Element(s) to process. Will be used to index the DataGrabber.
pre_run : str or None, optional pre_run_cmds : str or None, optional
Extra bash commands to source before the run (default None). Extra shell commands to source before the run (default None).
pre_collect : str or None, optional pre_collect_cmds : str or None, optional
Extra bash commands to source before the collect (default None). Extra shell commands to source before the collect (default None).
env : dict, optional env : :class:`.QueueContextEnv` or None, optional
The Python environment configuration. If None, will run without a The environment configuration. If None, will run without a
virtual environment of any kind (default None). virtual environment of any kind (default None).
verbose : str, optional verbose : str, optional
The level of verbosity (default "info"). The level of verbosity (default "info").
@ -48,25 +71,12 @@ class HTCondorAdapter(QueueContextAdapter):
The size of disk (HDD or SSD) to use (default "1G"). The size of disk (HDD or SSD) to use (default "1G").
extra_preamble : str or None, optional extra_preamble : str or None, optional
Extra commands to pass to HTCondor (default None). Extra commands to pass to HTCondor (default None).
collect : {"yes", "on_success_only", "no"}, optional collect_task : :class:`.HTCondorCollect`, optional
Whether to submit "collect" task for junifer (default "yes"). Whether to submit "collect" task for junifer (default "yes").
Valid options are:
* "yes": Submit "collect" task and run even if some of the jobs
fail.
* "on_success_only": Submit "collect" task and run only if all jobs
succeed.
* "no": Do not submit "collect" task.
submit : bool, optional submit : bool, optional
Whether to submit the jobs. In any case, .dag files will be created Whether to submit the jobs. In any case, .dag files will be created
for submission (default False). for submission (default False).
Raises
------
ValueError
If ``collect`` is invalid or if ``env`` is invalid.
See Also See Also
-------- --------
QueueContextAdapter : QueueContextAdapter :
@ -76,144 +86,67 @@ class HTCondorAdapter(QueueContextAdapter):
""" """
def __init__( job_name: str
self, job_dir: Path
job_name: str, yaml_config_path: Path
job_dir: Path, elements: Elements
yaml_config_path: Path, pre_run_cmds: str | None = None
elements: Elements, pre_collect_cmds: str | None = None
pre_run: str | None = None, env: QueueContextEnv | None = None
pre_collect: str | None = None, verbose: str = "info"
env: dict[str, str] | None = None, verbose_datalad: str | None = None
verbose: str = "info", cpus: int = 1
verbose_datalad: str | None = None, mem: str = "8G"
cpus: int = 1, disk: str = "1G"
mem: str = "8G", extra_preamble: str | None = None
disk: str = "1G", collect_task: HTCondorCollect = HTCondorCollect.Yes
extra_preamble: str | None = None, submit: bool = False
collect: str = "yes",
submit: bool = False,
) -> None:
"""Initialize the class."""
self._job_name = job_name
self._job_dir = job_dir
self._yaml_config_path = yaml_config_path
self._elements = elements
self._pre_run = pre_run
self._pre_collect = pre_collect
self._check_env(env)
self._verbose = verbose
self._verbose_datalad = verbose_datalad
self._cpus = cpus
self._mem = mem
self._disk = disk
self._extra_preamble = extra_preamble
self._collect = self._check_collect(collect)
self._submit = submit
self._log_dir = self._job_dir / "logs" def model_post_init(self, context: Any): # noqa: D102
self._pre_run_path = self._job_dir / "pre_run.sh" if self.env is None:
self._pre_collect_path = self._job_dir / "pre_collect.sh" self.env = QueueContextEnv(
self._submit_run_path = self._job_dir / f"run_{self._job_name}.submit" kind=EnvKind.Local, shell=EnvShell.Bash, name=""
)
if self.env["kind"] == EnvKind.Local:
# No virtual environment
self._executable = "junifer"
self._arguments = ""
else:
if self.env["name"] is None:
raise_error("`env.name` is required")
self._executable = f"run_{self.env['kind']}.{self.env['shell']}"
self._arguments = f"{self.env['name']} junifer"
self._exec_path = self.job_dir / self._executable
self._log_dir = self.job_dir / "logs"
self._pre_run_path = self.job_dir / "pre_run.sh"
self._pre_collect_path = self.job_dir / "pre_collect.sh"
self._submit_run_path = self.job_dir / f"run_{self.job_name}.submit"
self._submit_collect_path = ( self._submit_collect_path = (
self._job_dir / f"collect_{self._job_name}.submit" self.job_dir / f"collect_{self.job_name}.submit"
) )
self._dag_path = self._job_dir / f"{self._job_name}.dag" self._dag_path = self.job_dir / f"{self.job_name}.dag"
def _check_env(self, env: dict[str, str] | None) -> None:
"""Check value of env parameter on init.
Parameters
----------
env : dict or None
The value of env parameter.
Raises
------
ValueError
If ``env.kind`` is invalid or
if ``env.shell`` is invalid.
"""
# Set env related variables
if env is None:
env = {"kind": "local"}
# Check env kind
valid_env_kinds = ["conda", "venv", "local"]
if env["kind"] not in valid_env_kinds:
raise_error(
f"Invalid value for `env.kind`: {env['kind']}, "
f"must be one of {valid_env_kinds}"
)
else:
# Check shell
shell = env.get("shell", "bash")
valid_shells = ["bash", "zsh"]
if shell not in valid_shells:
raise_error(
f"Invalid value for `env.shell`: {shell}, "
f"must be one of {valid_shells}"
)
self._shell = shell
# Set variables
if env["kind"] == "local":
# No virtual environment
self._executable = "junifer"
self._arguments = ""
else:
self._executable = f"run_{env['kind']}.{self._shell}"
self._arguments = f"{env['name']} junifer"
self._exec_path = self._job_dir / self._executable
def _check_collect(self, collect: str) -> str:
"""Check value of collect parameter on init.
Parameters
----------
collect : str
The value of collect parameter.
Returns
-------
str
The checked value of collect parameter.
Raises
------
ValueError
If ``collect`` is invalid.
"""
valid_options = ["yes", "no", "on_success_only"]
if collect not in valid_options:
raise_error(
f"Invalid value for `collect`: {collect}, "
f"must be one of {valid_options}"
)
else:
return collect
def pre_run(self) -> str: def pre_run(self) -> str:
"""Return pre-run commands.""" """Return pre-run commands."""
fixed = ( fixed = (
f"#!/usr/bin/env {self._shell}\n\n" f"#!/usr/bin/env {self.env['shell']}\n\n"
"# This script is auto-generated by junifer.\n\n" "# This script is auto-generated by junifer.\n\n"
"# Force datalad to run in non-interactive mode\n" "# Force datalad to run in non-interactive mode\n"
"DATALAD_UI_INTERACTIVE=false\n" "DATALAD_UI_INTERACTIVE=false\n"
) )
var = self._pre_run or "" var = self.pre_run_cmds or ""
return fixed + "\n" + var return fixed + "\n" + var
def run(self) -> str: def run(self) -> str:
"""Return run commands.""" """Return run commands."""
verbose_args = f"--verbose {self._verbose} " verbose_args = f"--verbose {self.verbose} "
if self._verbose_datalad is not None: if self.verbose_datalad is not None:
verbose_args = ( verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad} " f"{verbose_args} --verbose-datalad {self.verbose_datalad} "
) )
junifer_run_args = ( junifer_run_args = (
"run " "run "
f"{self._yaml_config_path.resolve()!s} " f"{self.yaml_config_path.resolve()!s} "
f"{verbose_args}" f"{verbose_args}"
"--element $(element)" "--element $(element)"
) )
@ -226,11 +159,11 @@ class HTCondorAdapter(QueueContextAdapter):
"universe = vanilla\n" "universe = vanilla\n"
"getenv = True\n\n" "getenv = True\n\n"
"# Resources\n" "# Resources\n"
f"request_cpus = {self._cpus}\n" f"request_cpus = {self.cpus}\n"
f"request_memory = {self._mem}\n" f"request_memory = {self.mem}\n"
f"request_disk = {self._disk}\n\n" f"request_disk = {self.disk}\n\n"
"# Executable\n" "# Executable\n"
f"initial_dir = {self._job_dir.resolve()!s}\n" f"initial_dir = {self.job_dir.resolve()!s}\n"
f"executable = $(initial_dir)/{self._executable}\n" f"executable = $(initial_dir)/{self._executable}\n"
f"transfer_executable = False\n\n" f"transfer_executable = False\n\n"
f"arguments = {self._arguments} {junifer_run_args}\n\n" f"arguments = {self._arguments} {junifer_run_args}\n\n"
@ -239,31 +172,31 @@ class HTCondorAdapter(QueueContextAdapter):
f"output = {log_dir_prefix}.out\n" f"output = {log_dir_prefix}.out\n"
f"error = {log_dir_prefix}.err\n" f"error = {log_dir_prefix}.err\n"
) )
var = self._extra_preamble or "" var = self.extra_preamble or ""
return fixed + "\n" + var + "\n" + "queue" return fixed + "\n" + var + "\n" + "queue"
def pre_collect(self) -> str: def pre_collect(self) -> str:
"""Return pre-collect commands.""" """Return pre-collect commands."""
fixed = ( fixed = (
f"#!/usr/bin/env {self._shell}\n\n" f"#!/usr/bin/env {self.env['shell']}\n\n"
"# This script is auto-generated by junifer.\n" "# This script is auto-generated by junifer.\n"
) )
var = self._pre_collect or "" var = self.pre_collect_cmds or ""
# Add commands if collect="yes" # Add commands if collect="yes"
if self._collect == "yes": if self.collect_task == "yes":
var += 'if [ "${1}" == "4" ]; then\n exit 1\nfi\n' var += 'if [ "${1}" == "4" ]; then\n exit 1\nfi\n'
return fixed + "\n" + var return fixed + "\n" + var
def collect(self) -> str: def collect(self) -> str:
"""Return collect commands.""" """Return collect commands."""
verbose_args = f"--verbose {self._verbose} " verbose_args = f"--verbose {self.verbose} "
if self._verbose_datalad is not None: if self.verbose_datalad is not None:
verbose_args = ( verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad} " f"{verbose_args} --verbose-datalad {self.verbose_datalad} "
) )
junifer_collect_args = ( junifer_collect_args = (
f"collect {self._yaml_config_path.resolve()!s} {verbose_args}" f"collect {self.yaml_config_path.resolve()!s} {verbose_args}"
) )
log_dir_prefix = f"{self._log_dir.resolve()!s}/junifer_collect" log_dir_prefix = f"{self._log_dir.resolve()!s}/junifer_collect"
fixed = ( fixed = (
@ -272,11 +205,11 @@ class HTCondorAdapter(QueueContextAdapter):
"universe = vanilla\n" "universe = vanilla\n"
"getenv = True\n\n" "getenv = True\n\n"
"# Resources\n" "# Resources\n"
f"request_cpus = {self._cpus}\n" f"request_cpus = {self.cpus}\n"
f"request_memory = {self._mem}\n" f"request_memory = {self.mem}\n"
f"request_disk = {self._disk}\n\n" f"request_disk = {self.disk}\n\n"
"# Executable\n" "# Executable\n"
f"initial_dir = {self._job_dir.resolve()!s}\n" f"initial_dir = {self.job_dir.resolve()!s}\n"
f"executable = $(initial_dir)/{self._executable}\n" f"executable = $(initial_dir)/{self._executable}\n"
"transfer_executable = False\n\n" "transfer_executable = False\n\n"
f"arguments = {self._arguments} {junifer_collect_args}\n\n" f"arguments = {self._arguments} {junifer_collect_args}\n\n"
@ -285,13 +218,13 @@ class HTCondorAdapter(QueueContextAdapter):
f"output = {log_dir_prefix}.out\n" f"output = {log_dir_prefix}.out\n"
f"error = {log_dir_prefix}.err\n" f"error = {log_dir_prefix}.err\n"
) )
var = self._extra_preamble or "" var = self.extra_preamble or ""
return fixed + "\n" + var + "\n" + "queue" return fixed + "\n" + var + "\n" + "queue"
def dag(self) -> str: def dag(self) -> str:
"""Return HTCondor DAG commands.""" """Return HTCondor DAG commands."""
fixed = "" fixed = ""
for idx, element in enumerate(self._elements): for idx, element in enumerate(self.elements):
# Stringify elements if tuple for operation # Stringify elements if tuple for operation
str_element = ( str_element = (
",".join(element) if isinstance(element, tuple) else element ",".join(element) if isinstance(element, tuple) else element
@ -306,15 +239,15 @@ class HTCondorAdapter(QueueContextAdapter):
f'log_element="{log_element}"\n\n' # double quoted f'log_element="{log_element}"\n\n' # double quoted
) )
var = "" var = ""
if self._collect == "yes": if self.collect_task == "yes":
var += ( var += (
f"FINAL collect {self._submit_collect_path}\n" f"FINAL collect {self._submit_collect_path}\n"
f"SCRIPT PRE collect {self._pre_collect_path.as_posix()} " f"SCRIPT PRE collect {self._pre_collect_path.as_posix()} "
"$DAG_STATUS\n" "$DAG_STATUS\n"
) )
elif self._collect == "on_success_only": elif self.collect_task == "on_success_only":
var += f"JOB collect {self._submit_collect_path}\nPARENT " var += f"JOB collect {self._submit_collect_path}\nPARENT "
for idx, _ in enumerate(self._elements): for idx, _ in enumerate(self.elements):
var += f"run{idx} " var += f"run{idx} "
var += "CHILD collect\n" var += "CHILD collect\n"
@ -325,7 +258,7 @@ class HTCondorAdapter(QueueContextAdapter):
logger.info("Creating HTCondor job") logger.info("Creating HTCondor job")
# Create logs # Create logs
logger.info( logger.info(
f"Creating logs directory under {self._job_dir.resolve()!s}" f"Creating logs directory under {self.job_dir.resolve()!s}"
) )
self._log_dir.mkdir(exist_ok=True, parents=True) self._log_dir.mkdir(exist_ok=True, parents=True)
# Copy executable if not local # Copy executable if not local
@ -340,7 +273,7 @@ class HTCondorAdapter(QueueContextAdapter):
make_executable(self._exec_path) make_executable(self._exec_path)
# Create pre run # Create pre run
logger.info( logger.info(
f"Writing {self._pre_run_path.name} to {self._job_dir.resolve()!s}" f"Writing {self._pre_run_path.name} to {self.job_dir.resolve()!s}"
) )
self._pre_run_path.touch() self._pre_run_path.touch()
self._pre_run_path.write_text(textwrap.dedent(self.pre_run())) self._pre_run_path.write_text(textwrap.dedent(self.pre_run()))
@ -348,14 +281,14 @@ class HTCondorAdapter(QueueContextAdapter):
# Create run # Create run
logger.debug( logger.debug(
f"Writing {self._submit_run_path.name} to " f"Writing {self._submit_run_path.name} to "
f"{self._job_dir.resolve()!s}" f"{self.job_dir.resolve()!s}"
) )
self._submit_run_path.touch() self._submit_run_path.touch()
self._submit_run_path.write_text(textwrap.dedent(self.run())) self._submit_run_path.write_text(textwrap.dedent(self.run()))
# Create pre collect # Create pre collect
logger.info( logger.info(
f"Writing {self._pre_collect_path.name} to " f"Writing {self._pre_collect_path.name} to "
f"{self._job_dir.resolve()!s}" f"{self.job_dir.resolve()!s}"
) )
self._pre_collect_path.touch() self._pre_collect_path.touch()
self._pre_collect_path.write_text(textwrap.dedent(self.pre_collect())) self._pre_collect_path.write_text(textwrap.dedent(self.pre_collect()))
@ -363,13 +296,13 @@ class HTCondorAdapter(QueueContextAdapter):
# Create collect # Create collect
logger.debug( logger.debug(
f"Writing {self._submit_collect_path.name} to " f"Writing {self._submit_collect_path.name} to "
f"{self._job_dir.resolve()!s}" f"{self.job_dir.resolve()!s}"
) )
self._submit_collect_path.touch() self._submit_collect_path.touch()
self._submit_collect_path.write_text(textwrap.dedent(self.collect())) self._submit_collect_path.write_text(textwrap.dedent(self.collect()))
# Create DAG # Create DAG
logger.debug( logger.debug(
f"Writing {self._dag_path.name} to {self._job_dir.resolve()!s}" f"Writing {self._dag_path.name} to {self.job_dir.resolve()!s}"
) )
self._dag_path.touch() self._dag_path.touch()
self._dag_path.write_text(textwrap.dedent(self.dag())) self._dag_path.write_text(textwrap.dedent(self.dag()))
@ -379,7 +312,7 @@ class HTCondorAdapter(QueueContextAdapter):
"-include_env HOME", "-include_env HOME",
f"{self._dag_path.resolve()!s}", f"{self._dag_path.resolve()!s}",
] ]
if self._submit: if self.submit:
run_ext_cmd(name="condor_submit_dag", cmd=condor_submit_dag_cmd) run_ext_cmd(name="condor_submit_dag", cmd=condor_submit_dag_cmd)
else: else:
logger.info( logger.info(

View file

@ -3,22 +3,66 @@
# 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, 12): # pragma: no cover
from typing_extensions import TypedDict
else:
from typing import TypedDict
if sys.version_info < (3, 11): # pragma: no cover
from typing_extensions import Required
else:
from typing import Required
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from enum import Enum
from pydantic import BaseModel, ConfigDict
from ...utils import raise_error from ...utils import raise_error
__all__ = ["QueueContextAdapter"] __all__ = ["EnvKind", "EnvShell", "QueueContextAdapter", "QueueContextEnv"]
class QueueContextAdapter(ABC): class EnvKind(str, Enum):
"""Accepted Python environment kind."""
Venv = "venv"
Conda = "conda"
Local = "local"
class EnvShell(str, Enum):
"""Accepted environment shell."""
Bash = "bash"
Zsh = "zsh"
class QueueContextEnv(TypedDict, total=False):
"""Accepted environment configuration for queue context."""
kind: Required[EnvKind]
name: str
shell: Required[EnvShell]
class QueueContextAdapter(BaseModel, ABC):
"""Abstract base class for queue context adapter. """Abstract base class for queue context adapter.
For every interface that is required, one needs to provide a concrete For every queue context, one needs to provide a concrete
implementation of this abstract class. implementation of this abstract class.
""" """
model_config = ConfigDict(
use_enum_values=True,
extra="allow",
)
@abstractmethod @abstractmethod
def pre_run(self) -> str: def pre_run(self) -> str:
"""Return pre-run commands.""" """Return pre-run commands."""

View file

@ -11,30 +11,6 @@ import pytest
from junifer.api.queue_context import GnuParallelLocalAdapter from junifer.api.queue_context import GnuParallelLocalAdapter
def test_GnuParallelLocalAdapter_env_kind_error() -> None:
"""Test error for invalid env kind."""
with pytest.raises(ValueError, match=r"Invalid value for `env.kind`"):
GnuParallelLocalAdapter(
job_name="check_env_kind",
job_dir=Path("."),
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "jambalaya"},
)
def test_GnuParallelLocalAdapter_env_shell_error() -> None:
"""Test error for invalid env shell."""
with pytest.raises(ValueError, match=r"Invalid value for `env.shell`"):
GnuParallelLocalAdapter(
job_name="check_env_shell",
job_dir=Path("."),
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "conda", "shell": "fish"},
)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"elements, expected_text", "elements, expected_text",
[ [
@ -62,7 +38,7 @@ def test_GnuParallelLocalAdapter_elements(
yaml_config_path=Path("."), yaml_config_path=Path("."),
elements=elements, elements=elements,
) )
assert expected_text in adapter.elements() assert expected_text in adapter.elements_to_run()
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -97,7 +73,7 @@ def test_GnuParallelLocalAdapter_pre_run(
yaml_config_path=Path("."), yaml_config_path=Path("."),
elements=["sub01"], elements=["sub01"],
env={"kind": "conda", "name": "junifer", "shell": shell}, env={"kind": "conda", "name": "junifer", "shell": shell},
pre_run=pre_run, pre_run_cmds=pre_run,
) )
assert shell in adapter.pre_run() assert shell in adapter.pre_run()
assert expected_text in adapter.pre_run() assert expected_text in adapter.pre_run()
@ -135,7 +111,7 @@ def test_GnuParallelLocalAdapter_pre_collect(
yaml_config_path=Path("."), yaml_config_path=Path("."),
elements=["sub01"], elements=["sub01"],
env={"kind": "venv", "name": "junifer", "shell": shell}, env={"kind": "venv", "name": "junifer", "shell": shell},
pre_collect=pre_collect, pre_collect_cmds=pre_collect,
) )
assert shell in adapter.pre_collect() assert shell in adapter.pre_collect()
assert expected_text in adapter.pre_collect() assert expected_text in adapter.pre_collect()

View file

@ -11,42 +11,6 @@ import pytest
from junifer.api.queue_context import HTCondorAdapter from junifer.api.queue_context import HTCondorAdapter
def test_HTCondorAdapter_env_kind_error() -> None:
"""Test error for invalid env kind."""
with pytest.raises(ValueError, match=r"Invalid value for `env.kind`"):
HTCondorAdapter(
job_name="check_env_kind",
job_dir=Path("."),
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "jambalaya"},
)
def test_HTCondorAdapter_env_shell_error() -> None:
"""Test error for invalid env shell."""
with pytest.raises(ValueError, match=r"Invalid value for `env.shell`"):
HTCondorAdapter(
job_name="check_env_shell",
job_dir=Path("."),
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "conda", "shell": "fish"},
)
def test_HTCondorAdapter_collect_error() -> None:
"""Test error for invalid collect option."""
with pytest.raises(ValueError, match=r"Invalid value for `collect`"):
HTCondorAdapter(
job_name="check_collect",
job_dir=Path("."),
yaml_config_path=Path("."),
elements=["sub01"],
collect="off",
)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"pre_run, expected_text, shell", "pre_run, expected_text, shell",
[ [
@ -79,7 +43,7 @@ def test_HTCondorAdapter_pre_run(
yaml_config_path=Path("."), yaml_config_path=Path("."),
elements=["sub01"], elements=["sub01"],
env={"kind": "conda", "name": "junifer", "shell": shell}, env={"kind": "conda", "name": "junifer", "shell": shell},
pre_run=pre_run, pre_run_cmds=pre_run,
) )
assert shell in adapter.pre_run() assert shell in adapter.pre_run()
assert expected_text in adapter.pre_run() assert expected_text in adapter.pre_run()
@ -124,8 +88,8 @@ def test_HTCondorAdapter_pre_collect(
yaml_config_path=Path("."), yaml_config_path=Path("."),
elements=["sub01"], elements=["sub01"],
env={"kind": "venv", "name": "junifer", "shell": shell}, env={"kind": "venv", "name": "junifer", "shell": shell},
pre_collect=pre_collect, pre_collect_cmds=pre_collect,
collect=collect, collect_task=collect,
) )
assert shell in adapter.pre_collect() assert shell in adapter.pre_collect()
assert expected_text in adapter.pre_collect() assert expected_text in adapter.pre_collect()
@ -199,7 +163,7 @@ def test_HTCondor_dag(
job_dir=Path("."), job_dir=Path("."),
yaml_config_path=Path("."), yaml_config_path=Path("."),
elements=elements, elements=elements,
collect=collect, collect_task=collect,
) )
assert expected_text in adapter.dag() assert expected_text in adapter.dag()

View file

@ -62,7 +62,7 @@ def datagrabber() -> dict[str, str]:
@pytest.fixture @pytest.fixture
def markers() -> list[dict[str, str]]: def markers() -> list[dict[str, list[str] | str]]:
"""Return markers as a list of dictionary.""" """Return markers as a list of dictionary."""
return [ return [
{ {

View file

@ -17,4 +17,5 @@ queue:
env: env:
kind: conda kind: conda
name: junifer name: junifer
shell: bash
fraimondo commented 2026-05-26 08:15:23 +00:00 (Migrated from github.com)

We don't have bash as default here?

We don't have bash as default here?
synchon commented 2026-05-26 08:46:13 +00:00 (Migrated from github.com)

Yes the default is bash, this is to be explicit.

Yes the default is bash, this is to be explicit.
fraimondo commented 2026-05-26 08:55:12 +00:00 (Migrated from github.com)

ok

ok
mem: 8G mem: 8G

View file

@ -34,6 +34,7 @@ def test_get_dependency_information_short() -> None:
"""Test short version of _get_dependency_information().""" """Test short version of _get_dependency_information()."""
dependency_information = _get_dependency_information(long_=False) dependency_information = _get_dependency_information(long_=False)
dependency_list = [ dependency_list = [
"aenum",
"click", "click",
"numpy", "numpy",
"scipy", "scipy",
@ -50,6 +51,8 @@ def test_get_dependency_information_short() -> None:
"looseversion", "looseversion",
"junifer_data", "junifer_data",
"structlog", "structlog",
"pydantic",
"typing_extensions",
] ]
if sys.version_info < (3, 11): if sys.version_info < (3, 11):

View file

@ -2,3 +2,8 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de> # Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# License: AGPL # License: AGPL
import lazy_loader as lazy
__getattr__, __dir__, __all__ = lazy.attach_stub(__name__, __file__)

View file

View file

@ -2,12 +2,14 @@ __all__ = [
"JuselessDataladAOMICID1000VBM", "JuselessDataladAOMICID1000VBM",
"JuselessDataladCamCANVBM", "JuselessDataladCamCANVBM",
"JuselessDataladIXIVBM", "JuselessDataladIXIVBM",
"IXISite",
"JuselessUCLA", "JuselessUCLA",
"UCLATask",
"JuselessDataladUKBVBM", "JuselessDataladUKBVBM",
] ]
from .aomic_id1000_vbm import JuselessDataladAOMICID1000VBM from .aomic_id1000_vbm import JuselessDataladAOMICID1000VBM
from .camcan_vbm import JuselessDataladCamCANVBM from .camcan_vbm import JuselessDataladCamCANVBM
from .ixi_vbm import JuselessDataladIXIVBM from .ixi_vbm import JuselessDataladIXIVBM, IXISite
from .ucla import JuselessUCLA from .ucla import JuselessUCLA, UCLATask
from .ukb_vbm import JuselessDataladUKBVBM from .ukb_vbm import JuselessDataladUKBVBM

View file

@ -4,10 +4,13 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from typing import Literal
from pydantic import AnyUrl
from ....api.decorators import register_datagrabber from ....api.decorators import register_datagrabber
from ....datagrabber import PatternDataladDataGrabber from ....datagrabber import DataType, PatternDataladDataGrabber
from ....typing import DataGrabberPatterns
__all__ = ["JuselessDataladAOMICID1000VBM"] __all__ = ["JuselessDataladAOMICID1000VBM"]
@ -21,27 +24,19 @@ class JuselessDataladAOMICID1000VBM(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or pathlib.Path or None, optional datadir : pathlib.Path, optional
The directory where the datalad dataset will be cloned. If None, That path where the datalad dataset will be cloned.
the datalad dataset will be cloned into a temporary directory If not specified, the datalad dataset will be cloned into a temporary
(default None). directory.
""" """
def __init__(self, datadir: str | Path | None = None) -> None: uri: AnyUrl = AnyUrl("https://gin.g-node.org/felixh/ds003097_ReproVBM")
uri = "https://gin.g-node.org/felixh/ds003097_ReproVBM" types: list[Literal[DataType.VBM_GM]] = [DataType.VBM_GM] # noqa: RUF012
types = ["VBM_GM"] patterns: DataGrabberPatterns = { # noqa: RUF012
replacements = ["subject"] "VBM_GM": {
patterns = { "pattern": ("{subject}/mri/mwp1{subject}_run-2_T1w.nii.gz"),
"VBM_GM": { "space": "IXI549Space",
"pattern": ("{subject}/mri/mwp1{subject}_run-2_T1w.nii.gz"), },
"space": "IXI549Space", }
}, replacements: list[str] = ["subject"] # noqa: RUF012
}
super().__init__(
types=types,
datadir=datadir,
uri=uri,
replacements=replacements,
patterns=patterns,
)

View file

@ -5,10 +5,13 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from typing import Literal
from pydantic import AnyUrl
from ....api.decorators import register_datagrabber from ....api.decorators import register_datagrabber
from ....datagrabber import PatternDataladDataGrabber from ....datagrabber import DataType, PatternDataladDataGrabber
from ....typing import DataGrabberPatterns
__all__ = ["JuselessDataladCamCANVBM"] __all__ = ["JuselessDataladCamCANVBM"]
@ -22,30 +25,21 @@ class JuselessDataladCamCANVBM(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or pathlib.Path or None, optional datadir : pathlib.Path, optional
The directory where the datalad dataset will be cloned. If None, That path where the datalad dataset will be cloned.
the datalad dataset will be cloned into a temporary directory If not specified, the datalad dataset will be cloned into a temporary
(default None). directory.
""" """
def __init__(self, datadir: str | Path | None = None) -> None: uri: AnyUrl = AnyUrl(
uri = ( "ria+http://cat_12.5.ds.inm7.de#a139b26a-8406-11ea-8f94-a0369f287950"
"ria+http://cat_12.5.ds.inm7.de" )
"#a139b26a-8406-11ea-8f94-a0369f287950" types: list[Literal[DataType.VBM_GM]] = [DataType.VBM_GM] # noqa: RUF012
) patterns: DataGrabberPatterns = { # noqa: RUF012
types = ["VBM_GM"] "VBM_GM": {
replacements = ["subject"] "pattern": "{subject}/mri/m0wp1{subject}.nii.gz",
patterns = { "space": "IXI549Space",
"VBM_GM": { },
"pattern": "{subject}/mri/m0wp1{subject}.nii.gz", }
"space": "IXI549Space", replacements: list[str] = ["subject"] # noqa: RUF012
},
}
super().__init__(
types=types,
datadir=datadir,
uri=uri,
replacements=replacements,
patterns=patterns,
)

View file

@ -5,14 +5,29 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from enum import Enum
from typing import Annotated, ClassVar, Literal
from pydantic import AnyUrl, BeforeValidator
from ....api.decorators import register_datagrabber from ....api.decorators import register_datagrabber
from ....datagrabber import PatternDataladDataGrabber from ....datagrabber import DataType, PatternDataladDataGrabber
from ....utils import raise_error from ....typing import DataGrabberPatterns
from ....utils import ensure_list
__all__ = ["JuselessDataladIXIVBM"] __all__ = ["IXISite", "JuselessDataladIXIVBM"]
class IXISite(str, Enum):
"""Accepted IXI sites."""
Guys = "Guys"
HH = "HH"
IOP = "IOP"
_sites = Literal[IXISite.Guys, IXISite.HH, IXISite.IOP]
@register_datagrabber @register_datagrabber
@ -23,53 +38,31 @@ class JuselessDataladIXIVBM(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or pathlib.Path or None, optional datadir : pathlib.Path, optional
The directory where the datalad dataset will be cloned. If None, That path where the datalad dataset will be cloned.
the datalad dataset will be cloned into a temporary directory If not specified, the datalad dataset will be cloned into a temporary
(default None). directory.
sites : {"Guys", "HH", "IOP"} or list of the options or None, optional sites : {"Guys", "HH", "IOP"} or list of the options, optional
Which sites to access data from. If None, all available sites are IXI sites.
selected (default None). By default, all available sites are selected.
""" """
def __init__( uri: AnyUrl = AnyUrl(
self, "ria+http://cat_12.5.ds.inm7.de#b7107c52-8408-11ea-89c6-a0369f287950"
datadir: str | Path | None = None, )
sites: str | list[str] | None = None, types: list[Literal[DataType.VBM_GM]] = [DataType.VBM_GM] # noqa: RUF012
) -> None: sites: ClassVar[
uri = ( Annotated[IXISite | list[IXISite], BeforeValidator(ensure_list)]
"ria+http://cat_12.5.ds.inm7.de" ] = [
"#b7107c52-8408-11ea-89c6-a0369f287950" IXISite.Guys,
) IXISite.HH,
types = ["VBM_GM"] IXISite.IOP,
replacements = ["site", "subject"] ]
patterns = { patterns: DataGrabberPatterns = { # noqa: RUF012
"VBM_GM": { "VBM_GM": {
"pattern": ("{site}/{subject}/mri/m0wp1{subject}.nii.gz"), "pattern": ("{site}/{subject}/mri/m0wp1{subject}.nii.gz"),
"space": "IXI549Space", "space": "IXI549Space",
}, },
} }
replacements: list[str] = ["site", "subject"] # noqa: RUF012
# validate and/or transform 'site' input
all_sites = ["HH", "Guys", "IOP"]
if sites is None:
sites = all_sites
if isinstance(sites, str):
sites = [sites]
for s in sites:
if s not in all_sites:
raise_error(
f"{s} not a valid site in IXI VBM dataset!"
f"Available sites are {all_sites}"
)
self.sites = sites
super().__init__(
types=types,
datadir=datadir,
uri=uri,
replacements=replacements,
patterns=patterns,
)

View file

@ -31,10 +31,3 @@ def test_JuselessDataladIXIVBM() -> None:
out["VBM_GM"]["path"].name == f"m0wp1sub-{test_element[1]}.nii.gz" out["VBM_GM"]["path"].name == f"m0wp1sub-{test_element[1]}.nii.gz"
) )
assert out["VBM_GM"]["path"].exists() assert out["VBM_GM"]["path"].exists()
def test_JuselessDataladIXIVBM_invalid_site() -> None:
"""Test JuselessDataladIXIVBM with invalid site."""
with pytest.raises(ValueError, match="notavalidsite not a valid site"):
with JuselessDataladIXIVBM(sites="notavalidsite"):
pass

View file

@ -79,14 +79,6 @@ def test_JuselessUCLA_partial_data_access(
assert types in out assert types in out
def test_JuselessUCLA_incorrect_data_type() -> None:
"""Test JuselessUCLA DataGrabber incorrect data type."""
with pytest.raises(
ValueError, match="`patterns` must contain all `types`"
):
_ = JuselessUCLA(types="Eunomia")
@pytest.mark.parametrize( @pytest.mark.parametrize(
"tasks", "tasks",
[None, "rest", ["rest", "stopsignal"]], [None, "rest", ["rest", "stopsignal"]],
@ -124,11 +116,3 @@ def test_JuselessUCLA_task_params(tasks: str | None) -> None:
else: else:
for el in all_elements: for el in all_elements:
assert el[1] in ["rest", "stopsignal"] assert el[1] in ["rest", "stopsignal"]
def test_JuselessUCLA_invalid_tasks() -> None:
"""Test JuselessUCLA with invalid task parameters."""
with pytest.raises(
ValueError, match="invalid is not a valid task in the UCLA"
):
JuselessUCLA(tasks="invalid")

View file

@ -4,14 +4,52 @@
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# License: AGPL # License: AGPL
from enum import Enum
from pathlib import Path from pathlib import Path
from typing import Annotated, Literal
from pydantic import BeforeValidator
from ....api.decorators import register_datagrabber from ....api.decorators import register_datagrabber
from ....datagrabber import PatternDataGrabber from ....datagrabber import ConfoundsFormat, DataType, PatternDataGrabber
from ....utils import raise_error from ....typing import DataGrabberPatterns
from ....utils import ensure_list
__all__ = ["JuselessUCLA"] __all__ = ["JuselessUCLA", "UCLATask"]
class UCLATask(str, Enum):
"""Accepted UCLA tasks."""
REST = "rest"
BART = "bart"
BHT = "bht"
PAMENC = "pamenc"
PAMRET = "pamret"
SCAP = "scap"
TASKSWITCH = "taskswitch"
STOPSIGNAL = "stopsignal"
_types = Literal[
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
]
_tasks = Literal[
UCLATask.REST,
UCLATask.BART,
UCLATask.BHT,
UCLATask.PAMENC,
UCLATask.PAMRET,
UCLATask.SCAP,
UCLATask.TASKSWITCH,
UCLATask.STOPSIGNAL,
]
@register_datagrabber @register_datagrabber
@ -22,116 +60,88 @@ class JuselessUCLA(PatternDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or Path, optional datadir : Path, optional
The directory where the dataset is stored. The directory where the dataset is stored.
(default "/data/project/psychosis_thalamus/data/fmriprep"). (default "/data/project/psychosis_thalamus/data/fmriprep").
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM"} or \ types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM"} or \
list of the options, optional list of the options, optional
UCLA data types. If None, all available data types are selected. The data type(s) to grab.
(default None).
tasks : {"rest", "bart", "bht", "pamenc", "pamret", \ tasks : {"rest", "bart", "bht", "pamenc", "pamret", \
"scap", "taskswitch", "stopsignal"} or \ "scap", "taskswitch", "stopsignal"} or \
list of the options or None, optional list of the options, optional
UCLA task sessions. If None, all available task sessions are UCLA task sessions.
selected (default None). By default, all available task are selected.
""" """
def __init__( # the commented out uri leads to new open neuro dataset which does
self, # NOT have preprocessed data
datadir: str | Path = "/data/project/psychosis_thalamus/data/fmriprep", # uri = "https://github.com/OpenNeuroDatasets/ds000030.git"
types: str | list[str] | None = None, datadir: Path = Path("/data/project/psychosis_thalamus/data/fmriprep")
tasks: str | list[str] | None = None, types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
) -> None: DataType.BOLD,
# Declare all tasks DataType.T1w,
all_tasks = [ DataType.VBM_CSF,
"rest", DataType.VBM_GM,
"bart", DataType.VBM_WM,
"bht", ]
"pamenc", tasks: Annotated[
"pamret", UCLATask | list[UCLATask], BeforeValidator(ensure_list)
"scap", ] = [ # noqa: RUF012
"taskswitch", UCLATask.REST,
"stopsignal", UCLATask.BART,
] UCLATask.BHT,
# Set default tasks UCLATask.PAMENC,
if tasks is None: UCLATask.PAMRET,
tasks = all_tasks UCLATask.SCAP,
else: UCLATask.TASKSWITCH,
# Convert single task into list UCLATask.STOPSIGNAL,
if isinstance(tasks, str): ]
tasks = [tasks] patterns: DataGrabberPatterns = { # noqa: RUF012
# Verify valid tasks "BOLD": {
for t in tasks: "pattern": (
if t not in all_tasks: "{subject}/func/{subject}_task-{task}_bold_space-"
raise_error( "MNI152NLin2009cAsym_preproc.nii.gz"
f"{t} is not a valid task in the UCLA dataset!" ),
) "space": "MNI152NLin2009cAsym",
self.tasks = tasks "confounds": {
# The patterns
patterns = {
"BOLD": {
"pattern": ( "pattern": (
"{subject}/func/{subject}_task-{task}_bold_space-" "{subject}/func/{subject}_task-{task}_bold_confounds.tsv"
"MNI152NLin2009cAsym_preproc.nii.gz"
), ),
"space": "MNI152NLin2009cAsym", "space": "fmriprep",
"confounds": {
"pattern": (
"{subject}/func/{subject}_"
"task-{task}_bold_confounds.tsv"
),
"space": "fmriprep",
},
}, },
"T1w": { },
"pattern": ( "T1w": {
"{subject}/anat/{subject}_" "pattern": (
"T1w_space-MNI152NLin2009cAsym_preproc.nii.gz" "{subject}/anat/{subject}_"
), "T1w_space-MNI152NLin2009cAsym_preproc.nii.gz"
"space": "MNI152NLin2009cAsym", ),
}, "space": "MNI152NLin2009cAsym",
"VBM_CSF": { },
"pattern": ( "VBM_CSF": {
"{subject}/anat/{subject}_T1w_space-" "pattern": (
"MNI152NLin2009cAsym_class-CSF_probtissue.nii.gz" "{subject}/anat/{subject}_T1w_space-"
), "MNI152NLin2009cAsym_class-CSF_probtissue.nii.gz"
"space": "MNI152NLin2009cAsym", ),
}, "space": "MNI152NLin2009cAsym",
"VBM_GM": { },
"pattern": ( "VBM_GM": {
"{subject}/anat/{subject}_T1w_space-" "pattern": (
"MNI152NLin2009cAsym_class-GM_probtissue.nii.gz" "{subject}/anat/{subject}_T1w_space-"
), "MNI152NLin2009cAsym_class-GM_probtissue.nii.gz"
"space": "MNI152NLin2009cAsym", ),
}, "space": "MNI152NLin2009cAsym",
"VBM_WM": { },
"pattern": ( "VBM_WM": {
"{subject}/anat/{subject}_T1w_space" "pattern": (
"-MNI152NLin2009cAsym_class-WM_probtissue.nii.gz" "{subject}/anat/{subject}_T1w_space"
), "-MNI152NLin2009cAsym_class-WM_probtissue.nii.gz"
"space": "MNI152NLin2009cAsym", ),
}, "space": "MNI152NLin2009cAsym",
} },
# Set default types }
if types is None: replacements: list[str] = ["subject", "task"] # noqa: RUF012
types = list(patterns.keys()) confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
# Convert single type into list
else:
if not isinstance(types, list):
types = [types]
# The replacements
replacements = ["subject", "task"]
# the commented out uri leads to new open neuro dataset which does
# NOT have preprocessed data
# uri = "https://github.com/OpenNeuroDatasets/ds000030.git"
super().__init__(
types=types,
datadir=datadir,
patterns=patterns,
replacements=replacements,
confounds_format="fmriprep",
)
def get_elements(self) -> list: def get_elements(self) -> list:
"""Implement fetching list of elements in the dataset. """Implement fetching list of elements in the dataset.

View file

@ -6,9 +6,13 @@
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
from typing import Literal
from pydantic import AnyUrl
from ....api.decorators import register_datagrabber from ....api.decorators import register_datagrabber
from ....datagrabber import PatternDataladDataGrabber from ....datagrabber import DataType, PatternDataladDataGrabber
from ....typing import DataGrabberPatterns
__all__ = ["JuselessDataladUKBVBM"] __all__ = ["JuselessDataladUKBVBM"]
@ -22,29 +26,20 @@ class JuselessDataladUKBVBM(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or pathlib.Path or None, optional datadir : pathlib.Path, optional
The directory where the datalad dataset will be cloned. If None, That path where the datalad dataset will be cloned.
the datalad dataset will be cloned into a temporary directory If not specified, the datalad dataset will be cloned into a temporary
(default None). directory.
""" """
def __init__(self, datadir: str | Path | None = None) -> None: uri: AnyUrl = AnyUrl("ria+http://ukb.ds.inm7.de#~cat_m0wp1")
uri = "ria+http://ukb.ds.inm7.de#~cat_m0wp1" rootdir: Path = Path("m0wp1")
rootdir = "m0wp1" types: list[Literal[DataType.VBM_GM]] = [DataType.VBM_GM] # noqa: RUF012
types = ["VBM_GM"] patterns: DataGrabberPatterns = { # noqa: RUF012
replacements = ["subject", "session"] "VBM_GM": {
patterns = { "pattern": "m0wp1{subject}_ses-{session}_T1w.nii.gz",
"VBM_GM": { "space": "IXI549Space",
"pattern": "m0wp1{subject}_ses-{session}_T1w.nii.gz", },
"space": "IXI549Space", }
}, replacements: list[str] = ["subject", "session"] # noqa: RUF012
}
super().__init__(
types=types,
datadir=datadir,
uri=uri,
rootdir=rootdir,
replacements=replacements,
patterns=patterns,
)

0
junifer/configs/py.typed Normal file
View file

View file

@ -7,8 +7,9 @@
from pathlib import Path from pathlib import Path
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber import PatternDataladDataGrabber from junifer.datagrabber import DataType, PatternDataladDataGrabber
from junifer.utils.singleton import Singleton from junifer.utils.singleton import Singleton
@ -40,8 +41,8 @@ def maps_datagrabber(tmp_path: Path) -> PatternDataladDataGrabber:
""" """
dg = PatternDataladDataGrabber( dg = PatternDataladDataGrabber(
uri="https://github.com/OpenNeuroDatasets/ds005226.git", uri=AnyUrl("https://github.com/OpenNeuroDatasets/ds005226.git"),
types=["BOLD"], types=DataType.BOLD,
patterns={ patterns={
"BOLD": { "BOLD": {
"pattern": ( "pattern": (

View file

@ -94,7 +94,7 @@ class DataDispatcher(MutableMapping):
# Update global # Update global
self._registries[key] = value self._registries[key] = value
def popitem(): def popitem(self):
"""Not implemented.""" """Not implemented."""
pass pass

View file

@ -98,11 +98,11 @@ def test_compute_brain_mask_for_native(mask_type: str) -> None:
""" """
with DMCC13Benchmark( with DMCC13Benchmark(
types=["BOLD"], types="BOLD",
sessions=["ses-wave1bas"], sessions="ses-wave1bas",
tasks=["Rest"], tasks="Rest",
phase_encodings=["AP"], phase_encodings="AP",
runs=["1"], runs="1",
native_t1w=True, native_t1w=True,
) as dg: ) as dg:
element_data = DefaultDataReader().fit_transform( element_data = DefaultDataReader().fit_transform(
@ -177,7 +177,7 @@ def test_register_already_registered() -> None:
) )
def test_register( def test_register(
name: str, name: str,
mask_path: str, mask_path: str | Path,
space: str, space: str,
overwrite: bool, overwrite: bool,
) -> None: ) -> None:

View file

@ -3,30 +3,62 @@ __all__ = [
"DataladDataGrabber", "DataladDataGrabber",
"PatternDataGrabber", "PatternDataGrabber",
"PatternDataladDataGrabber", "PatternDataladDataGrabber",
"AOMICSpace",
"AOMICTask",
"DataladAOMICID1000", "DataladAOMICID1000",
"DataladAOMICPIOP1", "DataladAOMICPIOP1",
"DataladAOMICPIOP2", "DataladAOMICPIOP2",
"HCP1200", "HCP1200",
"HCP1200Task",
"HCP1200PhaseEncoding",
"DataladHCP1200", "DataladHCP1200",
"MultipleDataGrabber", "MultipleDataGrabber",
"DMCC13Benchmark", "DMCC13Benchmark",
"DMCCSession",
"DMCCTask",
"DMCCPhaseEncoding",
"DMCCRun",
"DataTypeManager", "DataTypeManager",
"DataTypeSchema", "DataTypeSchema",
"OptionalTypeSchema", "OptionalTypeSchema",
"PatternValidationMixin", "PatternValidationMixin",
"register_data_type", "register_data_type",
"DataType",
"ConfoundsFormat",
"register_confounds_format",
] ]
# These 4 need to be in this order, otherwise it is a circular import # These 4 need to be in this order, otherwise it is a circular import
from .base import BaseDataGrabber from .base import BaseDataGrabber, DataType
from .datalad_base import DataladDataGrabber from .datalad_base import DataladDataGrabber
from .pattern import PatternDataGrabber from .pattern import (
PatternDataGrabber,
ConfoundsFormat,
register_confounds_format,
)
from .pattern_datalad import PatternDataladDataGrabber from .pattern_datalad import PatternDataladDataGrabber
from .aomic import DataladAOMICID1000, DataladAOMICPIOP1, DataladAOMICPIOP2 from .aomic import (
from .hcp1200 import HCP1200, DataladHCP1200 AOMICSpace,
AOMICTask,
DataladAOMICID1000,
DataladAOMICPIOP1,
DataladAOMICPIOP2,
)
from .hcp1200 import (
HCP1200,
HCP1200Task,
HCP1200PhaseEncoding,
DataladHCP1200,
)
from .multiple import MultipleDataGrabber from .multiple import MultipleDataGrabber
from .dmcc13_benchmark import DMCC13Benchmark from .dmcc13_benchmark import (
DMCC13Benchmark,
DMCCSession,
DMCCTask,
DMCCPhaseEncoding,
DMCCRun,
)
from .pattern_validation_mixin import ( from .pattern_validation_mixin import (
DataTypeManager, DataTypeManager,

View file

@ -1,5 +1,12 @@
__all__ = ["DataladAOMICID1000", "DataladAOMICPIOP1", "DataladAOMICPIOP2"] __all__ = [
"AOMICSpace",
"AOMICTask",
"DataladAOMICID1000",
"DataladAOMICPIOP1",
"DataladAOMICPIOP2",
]
from ._types import AOMICSpace, AOMICTask
from .id1000 import DataladAOMICID1000 from .id1000 import DataladAOMICID1000
from .piop1 import DataladAOMICPIOP1 from .piop1 import DataladAOMICPIOP1
from .piop2 import DataladAOMICPIOP2 from .piop2 import DataladAOMICPIOP2

View file

@ -0,0 +1,22 @@
"""Provide common types for AOMIC DataGrabbers."""
from enum import Enum
class AOMICSpace(str, Enum):
"""Accepted spaces for AOMIC."""
Native = "native"
MNI152NLin2009cAsym = "MNI152NLin2009cAsym"
class AOMICTask(str, Enum):
"""Accepted tasks for AOMIC."""
RestingState = "restingstate"
Anticipation = "anticipation"
EmoMatching = "emomatching"
Faces = "faces"
Gstroop = "gstroop"
WorkingMemory = "workingmemory"
StopSignal = "stopsignal"

View file

@ -7,15 +7,32 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path from typing import Annotated, Literal
from pydantic import AnyUrl, BeforeValidator
from ...api.decorators import register_datagrabber from ...api.decorators import register_datagrabber
from ...utils import raise_error from ...typing import DataGrabberPatterns
from ...utils import ensure_list
from ..base import DataType
from ..pattern import ConfoundsFormat
from ..pattern_datalad import PatternDataladDataGrabber from ..pattern_datalad import PatternDataladDataGrabber
from ._types import AOMICSpace
__all__ = ["DataladAOMICID1000"] __all__ = ["DataladAOMICID1000"]
_types = Literal[
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
fraimondo commented 2025-11-26 10:56:05 +00:00 (Migrated from github.com)

should be a DataType of a list of them

should be a DataType of a list of them
fraimondo commented 2025-11-26 10:56:30 +00:00 (Migrated from github.com)

Same for all the datagrabbers, markers, steps, etc where it can be one or many.

Same for all the datagrabbers, markers, steps, etc where it can be one or many.
DataType.VBM_GM,
DataType.VBM_WM,
DataType.DWI,
DataType.FreeSurfer,
DataType.Warp,
]
@register_datagrabber @register_datagrabber
class DataladAOMICID1000(PatternDataladDataGrabber): class DataladAOMICID1000(PatternDataladDataGrabber):
@ -23,212 +40,206 @@ class DataladAOMICID1000(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or Path or None, optional types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
The directory where the datalad dataset will be cloned. If None, "FreeSurfer", "Warp"} or list of the options, optional
the datalad dataset will be cloned into a temporary directory The data type(s) to grab.
(default None). datadir : pathlib.Path, optional
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \ That path where the datalad dataset will be cloned.
"FreeSurfer"} or list of the options, optional If not specified, the datalad dataset will be cloned into a temporary
AOMIC data types. If None, all available data types are selected. directory.
(default None).
space : {"native", "MNI152NLin2009cAsym"}, optional space : {"native", "MNI152NLin2009cAsym"}, optional
The space to use for the data (default "MNI152NLin2009cAsym"). AOMIC space (default ``AOMICSpace.MNI152NLin2009cAsym``).
Raises
------
ValueError
If invalid value is passed for:
* ``space``
""" """
def __init__( uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds003097.git")
self, types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
datadir: str | Path | None = None, DataType.BOLD,
types: str | list[str] | None = None, DataType.T1w,
space: str = "MNI152NLin2009cAsym", DataType.VBM_CSF,
) -> None: DataType.VBM_GM,
valid_spaces = ["native", "MNI152NLin2009cAsym"] DataType.VBM_WM,
if space not in ["native", "MNI152NLin2009cAsym"]: DataType.DWI,
raise_error( DataType.FreeSurfer,
f"Invalid space {space}. Must be one of {valid_spaces}" DataType.Warp,
) ]
space: AOMICSpace = AOMICSpace.MNI152NLin2009cAsym
# Descriptor for space in `anat` patterns: DataGrabberPatterns = { # noqa: RUF012
sp_anat_desc = ( "BOLD": {
"" if space == "native" else "space-MNI152NLin2009cAsym_" "pattern": (
) "derivatives/fmriprep/{subject}/func/"
# Descriptor for space in `func` "{subject}_task-moviewatching_"
sp_func_desc = ( "{sp_func_desc}"
"space-T1w_" if space == "native" else "space-MNI152NLin2009cAsym_" "desc-preproc_bold.nii.gz"
) ),
# The patterns "mask": {
patterns = {
"BOLD": {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/func/" "derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_" "{subject}_task-moviewatching_"
f"{sp_func_desc}" "{sp_func_desc}"
"desc-preproc_bold.nii.gz" "desc-brain_mask.nii.gz"
), ),
"space": space,
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
f"{sp_func_desc}"
"desc-brain_mask.nii.gz"
),
"space": space,
},
"confounds": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
"desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
"reference": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
f"{sp_func_desc}"
"boldref.nii.gz"
),
},
}, },
"T1w": { "confounds": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
"desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
"reference": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
"{sp_func_desc}"
"boldref.nii.gz"
),
},
},
"T1w": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"desc-preproc_T1w.nii.gz"
),
"mask": {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_"
f"{sp_anat_desc}" "{sp_anat_desc}"
"desc-preproc_T1w.nii.gz" "desc-brain_mask.nii.gz"
), ),
"space": space,
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"desc-brain_mask.nii.gz"
),
"space": space,
},
}, },
"VBM_CSF": { },
"VBM_CSF": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-CSF_probseg.nii.gz"
),
},
"VBM_GM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-GM_probseg.nii.gz"
),
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-WM_probseg.nii.gz"
),
},
"DWI": {
"pattern": (
"derivatives/dwipreproc/{subject}/dwi/"
"{subject}_desc-preproc_dwi.nii.gz"
),
},
"FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
"aseg": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/aseg.mg[z]"
)
},
"norm": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/norm.mg[z]"
)
},
"lh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.whit[e]"
)
},
"rh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.whit[e]"
)
},
"lh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.pia[l]"
)
},
"rh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.pia[l]"
)
},
},
"Warp": [
{
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_from-MNI152NLin2009cAsym_to-T1w_"
f"{sp_anat_desc}" "mode-image_xfm.h5"
"label-CSF_probseg.nii.gz"
), ),
"space": space, "src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
}, },
"VBM_GM": { {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_from-T1w_to-MNI152NLin2009cAsym_"
f"{sp_anat_desc}" "mode-image_xfm.h5"
"label-GM_probseg.nii.gz"
), ),
"space": space, "src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
}, },
"VBM_WM": { ],
"pattern": ( }
"derivatives/fmriprep/{subject}/anat/" replacements: list[str] = ["subject"] # noqa: RUF012
"{subject}_" confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
f"{sp_anat_desc}"
"label-WM_probseg.nii.gz"
),
"space": space,
},
"DWI": {
"pattern": (
"derivatives/dwipreproc/{subject}/dwi/"
"{subject}_desc-preproc_dwi.nii.gz"
),
},
"FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
"aseg": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/aseg.mg[z]"
)
},
"norm": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/norm.mg[z]"
)
},
"lh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.whit[e]"
)
},
"rh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.whit[e]"
)
},
"lh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.pia[l]"
)
},
"rh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.pia[l]"
)
},
},
"Warp": [
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_from-MNI152NLin2009cAsym_to-T1w_"
"mode-image_xfm.h5"
),
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_from-T1w_to-MNI152NLin2009cAsym_"
"mode-image_xfm.h5"
),
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
],
}
if space == "native":
patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym"
else: def validate_datagrabber_params(self) -> None:
patterns["BOLD"]["prewarp_space"] = "native" """Run extra logical validation for datagrabber."""
# Descriptor for space in `anat`
# Use native T1w assets sp_anat_desc = (
self.space = space "" if self.space == "native" else "space-MNI152NLin2009cAsym_"
# Set default types
if types is None:
types = list(patterns.keys())
# Convert single type into list
else:
if not isinstance(types, list):
types = [types]
# The replacements
replacements = ["subject"]
uri = "https://github.com/OpenNeuroDatasets/ds003097.git"
super().__init__(
types=types,
datadir=datadir,
uri=uri,
patterns=patterns,
replacements=replacements,
confounds_format="fmriprep",
) )
# Descriptor for space in `func`
sp_func_desc = (
"space-T1w_"
if self.space == "native"
else "space-MNI152NLin2009cAsym_"
)
self.patterns["BOLD"]["pattern"] = self.patterns["BOLD"][
"pattern"
].replace("{sp_func_desc}", sp_func_desc)
self.patterns["BOLD"]["mask"]["pattern"] = self.patterns["BOLD"][
"mask"
]["pattern"].replace("{sp_func_desc}", sp_func_desc)
self.patterns["BOLD"]["reference"]["pattern"] = self.patterns["BOLD"][
"reference"
]["pattern"].replace("{sp_func_desc}", sp_func_desc)
self.patterns["T1w"]["pattern"] = self.patterns["T1w"][
"pattern"
].replace("{sp_anat_desc}", sp_anat_desc)
self.patterns["T1w"]["mask"]["pattern"] = self.patterns["T1w"]["mask"][
"pattern"
].replace("{sp_anat_desc}", sp_anat_desc)
for t in ["VBM_CSF", "VBM_GM", "VBM_WM"]:
self.patterns[t]["pattern"] = self.patterns[t]["pattern"].replace(
"{sp_anat_desc}", sp_anat_desc
)
for t in ["BOLD", "T1w"]:
self.patterns[t]["space"] = self.space
self.patterns[t]["mask"]["space"] = self.space
for t in ["VBM_CSF", "VBM_GM", "VBM_WM"]:
self.patterns[t]["space"] = self.space
if self.space == "native":
self.patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym"
else:
self.patterns["BOLD"]["prewarp_space"] = "native"
super().validate_datagrabber_params()

View file

@ -8,15 +8,41 @@
# License: AGPL # License: AGPL
from itertools import product from itertools import product
from pathlib import Path from typing import Annotated, Literal
from pydantic import AnyUrl, BeforeValidator
from ...api.decorators import register_datagrabber from ...api.decorators import register_datagrabber
from ...utils import raise_error from ...typing import DataGrabberPatterns
from ...utils import ensure_list
from ..base import DataType
from ..pattern import ConfoundsFormat
from ..pattern_datalad import PatternDataladDataGrabber from ..pattern_datalad import PatternDataladDataGrabber
from ._types import AOMICSpace, AOMICTask
__all__ = ["DataladAOMICPIOP1"] __all__ = ["DataladAOMICPIOP1"]
_types = Literal[
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.DWI,
DataType.FreeSurfer,
DataType.Warp,
]
_tasks = Literal[
AOMICTask.RestingState,
AOMICTask.Anticipation,
AOMICTask.EmoMatching,
AOMICTask.Faces,
AOMICTask.Gstroop,
AOMICTask.WorkingMemory,
]
@register_datagrabber @register_datagrabber
class DataladAOMICPIOP1(PatternDataladDataGrabber): class DataladAOMICPIOP1(PatternDataladDataGrabber):
@ -24,246 +50,224 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or Path or None, optional types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
The directory where the datalad dataset will be cloned. If None, "FreeSurfer", "Warp"} or list of the options, optional
the datalad dataset will be cloned into a temporary directory The data type(s) to grab.
(default None). datadir : pathlib.Path, optional
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \ That path where the datalad dataset will be cloned.
"FreeSurfer"} or list of the options, optional If not specified, the datalad dataset will be cloned into a temporary
AOMIC data types. If None, all available data types are selected. directory.
(default None).
tasks : {"restingstate", "anticipation", "emomatching", "faces", \ tasks : {"restingstate", "anticipation", "emomatching", "faces", \
"gstroop", "workingmemory"} or list of the options, optional "gstroop", "workingmemory"} or list of the options, optional
AOMIC PIOP1 task sessions. If None, all available task sessions are AOMIC PIOP1 task sessions.
selected (default None). By default, all available task sessions are selected.
space : {"native", "MNI152NLin2009cAsym"}, optional space : {"native", "MNI152NLin2009cAsym"}, optional
The space to use for the data (default "MNI152NLin2009cAsym"). AOMIC space (default ``AOMICSpace.MNI152NLin2009cAsym``).
Raises
------
ValueError
If invalid value is passed for:
* ``tasks``
* ``space``
""" """
def __init__( uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds002785")
self, types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
datadir: str | Path | None = None, DataType.BOLD,
types: str | list[str] | None = None, DataType.T1w,
tasks: str | list[str] | None = None, DataType.VBM_CSF,
space: str = "MNI152NLin2009cAsym", DataType.VBM_GM,
) -> None: DataType.VBM_WM,
valid_spaces = ["native", "MNI152NLin2009cAsym"] DataType.DWI,
if space not in ["native", "MNI152NLin2009cAsym"]: DataType.FreeSurfer,
raise_error( DataType.Warp,
f"Invalid space {space}. Must be one of {valid_spaces}" ]
) tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
# Declare all tasks AOMICTask.RestingState,
all_tasks = [ AOMICTask.Anticipation,
"restingstate", AOMICTask.EmoMatching,
"anticipation", AOMICTask.Faces,
"emomatching", AOMICTask.Gstroop,
"faces", AOMICTask.WorkingMemory,
"gstroop", ]
"workingmemory", space: AOMICSpace = AOMICSpace.MNI152NLin2009cAsym
] patterns: DataGrabberPatterns = { # noqa: RUF012
# Set default tasks "BOLD": {
if tasks is None: "pattern": (
tasks = all_tasks "derivatives/fmriprep/{subject}/func/"
else: "{subject}_task-{task}_"
# Convert single task into list "{sp_func_desc}"
if isinstance(tasks, str): "desc-preproc_bold.nii.gz"
tasks = [tasks] ),
# Verify valid tasks "mask": {
for t in tasks:
if t not in all_tasks:
raise_error(
f"{t} is not a valid task in the AOMIC PIOP1 dataset!"
)
self.tasks = tasks
# Descriptor for space in `anat`
sp_anat_desc = (
"" if space == "native" else "space-MNI152NLin2009cAsym_"
)
# Descriptor for space in `func`
sp_func_desc = (
"space-T1w_" if space == "native" else "space-MNI152NLin2009cAsym_"
)
# The patterns
patterns = {
"BOLD": {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/func/" "derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_" "{subject}_task-{task}_"
f"{sp_func_desc}" "{sp_func_desc}"
"desc-preproc_bold.nii.gz" "desc-brain_mask.nii.gz"
), ),
"space": space,
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
f"{sp_func_desc}"
"desc-brain_mask.nii.gz"
),
"space": space,
},
"confounds": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
"reference": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
f"{sp_func_desc}"
"boldref.nii.gz"
),
},
}, },
"T1w": { "confounds": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
"reference": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"{sp_func_desc}"
"boldref.nii.gz"
),
},
},
"T1w": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"desc-preproc_T1w.nii.gz"
),
"mask": {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_"
f"{sp_anat_desc}" "{sp_anat_desc}"
"desc-preproc_T1w.nii.gz" "desc-brain_mask.nii.gz"
), ),
"space": space,
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"desc-brain_mask.nii.gz"
),
"space": space,
},
}, },
"VBM_CSF": { },
"VBM_CSF": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-CSF_probseg.nii.gz"
),
},
"VBM_GM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-GM_probseg.nii.gz"
),
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-WM_probseg.nii.gz"
),
},
"DWI": {
"pattern": (
"derivatives/dwipreproc/{subject}/dwi/"
"{subject}_desc-preproc_dwi.nii.gz"
),
},
"FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
"aseg": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/aseg.mg[z]"
)
},
"norm": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/norm.mg[z]"
)
},
"lh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.whit[e]"
)
},
"rh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.whit[e]"
)
},
"lh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.pia[l]"
)
},
"rh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.pia[l]"
)
},
},
"Warp": [
{
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_from-MNI152NLin2009cAsym_to-T1w_"
f"{sp_anat_desc}" "mode-image_xfm.h5"
"label-CSF_probseg.nii.gz"
), ),
"space": space, "src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
}, },
"VBM_GM": { {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_from-T1w_to-MNI152NLin2009cAsym_"
f"{sp_anat_desc}" "mode-image_xfm.h5"
"label-GM_probseg.nii.gz"
), ),
"space": space, "src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
}, },
"VBM_WM": { ],
"pattern": ( }
"derivatives/fmriprep/{subject}/anat/" replacements: list[str] = ["subject", "task"] # noqa: RUF012
"{subject}_" confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
f"{sp_anat_desc}"
"label-WM_probseg.nii.gz"
),
"space": space,
},
"DWI": {
"pattern": (
"derivatives/dwipreproc/{subject}/dwi/"
"{subject}_desc-preproc_dwi.nii.gz"
),
},
"FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
"aseg": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/aseg.mg[z]"
)
},
"norm": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/norm.mg[z]"
)
},
"lh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.whit[e]"
)
},
"rh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.whit[e]"
)
},
"lh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.pia[l]"
)
},
"rh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.pia[l]"
)
},
},
"Warp": [
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_from-MNI152NLin2009cAsym_to-T1w_"
"mode-image_xfm.h5"
),
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_from-T1w_to-MNI152NLin2009cAsym_"
"mode-image_xfm.h5"
),
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
],
}
if space == "native": def validate_datagrabber_params(self) -> None:
patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym" """Run extra logical validation for datagrabber."""
else: # Descriptor for space in `anat`
patterns["BOLD"]["prewarp_space"] = "native" sp_anat_desc = (
"" if self.space == "native" else "space-MNI152NLin2009cAsym_"
# Use native T1w assets
self.space = space
# Set default types
if types is None:
types = list(patterns.keys())
# Convert single type into list
else:
if not isinstance(types, list):
types = [types]
# The replacements
replacements = ["subject", "task"]
uri = "https://github.com/OpenNeuroDatasets/ds002785"
super().__init__(
types=types,
datadir=datadir,
uri=uri,
patterns=patterns,
replacements=replacements,
confounds_format="fmriprep",
) )
# Descriptor for space in `func`
sp_func_desc = (
"space-T1w_"
if self.space == "native"
else "space-MNI152NLin2009cAsym_"
)
self.patterns["BOLD"]["pattern"] = self.patterns["BOLD"][
"pattern"
].replace("{sp_func_desc}", sp_func_desc)
self.patterns["BOLD"]["mask"]["pattern"] = self.patterns["BOLD"][
"mask"
]["pattern"].replace("{sp_func_desc}", sp_func_desc)
self.patterns["BOLD"]["reference"]["pattern"] = self.patterns["BOLD"][
"reference"
]["pattern"].replace("{sp_func_desc}", sp_func_desc)
self.patterns["T1w"]["pattern"] = self.patterns["T1w"][
"pattern"
].replace("{sp_anat_desc}", sp_anat_desc)
self.patterns["T1w"]["mask"]["pattern"] = self.patterns["T1w"]["mask"][
"pattern"
].replace("{sp_anat_desc}", sp_anat_desc)
for t in ["VBM_CSF", "VBM_GM", "VBM_WM"]:
self.patterns[t]["pattern"] = self.patterns[t]["pattern"].replace(
"{sp_anat_desc}", sp_anat_desc
)
for t in ["BOLD", "T1w"]:
self.patterns[t]["space"] = self.space
self.patterns[t]["mask"]["space"] = self.space
for t in ["VBM_CSF", "VBM_GM", "VBM_WM"]:
self.patterns[t]["space"] = self.space
if self.space == "native":
self.patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym"
else:
self.patterns["BOLD"]["prewarp_space"] = "native"
super().validate_datagrabber_params()
def get_item(self, subject: str, task: str) -> dict: def get_item(self, subject: str, task: str) -> dict:
"""Index one element in the dataset. """Get the specified item from the dataset.
Parameters Parameters
---------- ----------

View file

@ -8,15 +8,39 @@
# License: AGPL # License: AGPL
from itertools import product from itertools import product
from pathlib import Path from typing import Annotated, Literal
from pydantic import AnyUrl, BeforeValidator
from ...api.decorators import register_datagrabber from ...api.decorators import register_datagrabber
from ...utils import raise_error from ...typing import DataGrabberPatterns
from ...utils import ensure_list
from ..base import DataType
from ..pattern import ConfoundsFormat
from ..pattern_datalad import PatternDataladDataGrabber from ..pattern_datalad import PatternDataladDataGrabber
from ._types import AOMICSpace, AOMICTask
__all__ = ["DataladAOMICPIOP2"] __all__ = ["DataladAOMICPIOP2"]
_types = Literal[
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.DWI,
DataType.FreeSurfer,
DataType.Warp,
]
_tasks = Literal[
AOMICTask.RestingState,
AOMICTask.StopSignal,
AOMICTask.EmoMatching,
AOMICTask.WorkingMemory,
]
@register_datagrabber @register_datagrabber
class DataladAOMICPIOP2(PatternDataladDataGrabber): class DataladAOMICPIOP2(PatternDataladDataGrabber):
@ -24,241 +48,219 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or Path or None, optional types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
The directory where the datalad dataset will be cloned. If None, "FreeSurfer", "Warp"} or list of the options, optional
the datalad dataset will be cloned into a temporary directory The data type(s) to grab.
(default None). datadir : pathlib.Path, optional
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \ That path where the datalad dataset will be cloned.
"FreeSurfer"} or list of the options, optional If not specified, the datalad dataset will be cloned into a temporary
AOMIC data types. If None, all available data types are selected. directory.
(default None).
tasks : {"restingstate", "stopsignal", "workingmemory", "emomatching"} or \ tasks : {"restingstate", "stopsignal", "workingmemory", "emomatching"} or \
list of the options, optional list of the options, optional
AOMIC PIOP2 task sessions. If None, all available task sessions are AOMIC PIOP2 task sessions.
selected (default None). By default, all available task sessions are selected.
space : {"native", "MNI152NLin2009cAsym"}, optional space : {"native", "MNI152NLin2009cAsym"}, optional
The space to use for the data (default "MNI152NLin2009cAsym"). AOMIC space (default ``AOMICSpace.MNI152NLin2009cAsym``).
Raises
------
ValueError
If invalid value is passed for:
* ``tasks``
* ``space``
""" """
def __init__( uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds002790")
self, types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
datadir: str | Path | None = None, DataType.BOLD,
types: str | list[str] | None = None, DataType.T1w,
tasks: str | list[str] | None = None, DataType.VBM_CSF,
space: str = "MNI152NLin2009cAsym", DataType.VBM_GM,
) -> None: DataType.VBM_WM,
valid_spaces = ["native", "MNI152NLin2009cAsym"] DataType.DWI,
if space not in ["native", "MNI152NLin2009cAsym"]: DataType.FreeSurfer,
raise_error( DataType.Warp,
f"Invalid space {space}. Must be one of {valid_spaces}" ]
) tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
# Declare all tasks AOMICTask.RestingState,
all_tasks = [ AOMICTask.StopSignal,
"restingstate", AOMICTask.EmoMatching,
"stopsignal", AOMICTask.WorkingMemory,
"workingmemory", ]
"emomatching", space: AOMICSpace = AOMICSpace.MNI152NLin2009cAsym
] patterns: DataGrabberPatterns = { # noqa: RUF012
# Set default tasks "BOLD": {
if tasks is None: "pattern": (
tasks = all_tasks "derivatives/fmriprep/{subject}/func/"
else: "{subject}_task-{task}_"
# Convert single task into list "{sp_func_desc}"
if isinstance(tasks, str): "desc-preproc_bold.nii.gz"
tasks = [tasks] ),
# Verify valid tasks "mask": {
for t in tasks:
if t not in all_tasks:
raise_error(
f"{t} is not a valid task in the AOMIC PIOP2 dataset!"
)
self.tasks = tasks
# Descriptor for space in `anat`
sp_anat_desc = (
"" if space == "native" else "space-MNI152NLin2009cAsym_"
)
# Descriptor for space in `func`
sp_func_desc = (
"space-T1w_" if space == "native" else "space-MNI152NLin2009cAsym_"
)
# The patterns
patterns = {
"BOLD": {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/func/" "derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_" "{subject}_task-{task}_"
f"{sp_func_desc}" "{sp_func_desc}"
"desc-preproc_bold.nii.gz" "desc-brain_mask.nii.gz"
), ),
"space": space,
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
f"{sp_func_desc}"
"desc-brain_mask.nii.gz"
),
"space": space,
},
"confounds": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
"reference": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
f"{sp_func_desc}"
"boldref.nii.gz"
),
},
}, },
"T1w": { "confounds": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
"reference": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"{sp_func_desc}"
"boldref.nii.gz"
),
},
},
"T1w": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"desc-preproc_T1w.nii.gz"
),
"mask": {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_"
f"{sp_anat_desc}" "{sp_anat_desc}"
"desc-preproc_T1w.nii.gz" "desc-brain_mask.nii.gz"
), ),
"space": space,
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"desc-brain_mask.nii.gz"
),
"space": space,
},
}, },
"VBM_CSF": { },
"VBM_CSF": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-CSF_probseg.nii.gz"
),
},
"VBM_GM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-GM_probseg.nii.gz"
),
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
"{sp_anat_desc}"
"label-WM_probseg.nii.gz"
),
},
"DWI": {
"pattern": (
"derivatives/dwipreproc/{subject}/dwi/"
"{subject}_desc-preproc_dwi.nii.gz"
),
},
"FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
"aseg": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/aseg.mg[z]"
)
},
"norm": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/norm.mg[z]"
)
},
"lh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.whit[e]"
)
},
"rh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.whit[e]"
)
},
"lh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.pia[l]"
)
},
"rh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.pia[l]"
)
},
},
"Warp": [
{
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_from-MNI152NLin2009cAsym_to-T1w_"
f"{sp_anat_desc}" "mode-image_xfm.h5"
"label-CSF_probseg.nii.gz"
), ),
"space": space, "src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
}, },
"VBM_GM": { {
"pattern": ( "pattern": (
"derivatives/fmriprep/{subject}/anat/" "derivatives/fmriprep/{subject}/anat/"
"{subject}_" "{subject}_from-T1w_to-MNI152NLin2009cAsym_"
f"{sp_anat_desc}" "mode-image_xfm.h5"
"label-GM_probseg.nii.gz"
), ),
"space": space, "src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
}, },
"VBM_WM": { ],
"pattern": ( }
"derivatives/fmriprep/{subject}/anat/" replacements: list[str] = ["subject", "task"] # noqa: RUF012
"{subject}_" confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
f"{sp_anat_desc}"
"label-WM_probseg.nii.gz"
),
"space": space,
},
"DWI": {
"pattern": (
"derivatives/dwipreproc/{subject}/dwi/"
"{subject}_desc-preproc_dwi.nii.gz"
),
},
"FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
"aseg": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/aseg.mg[z]"
)
},
"norm": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/mri/norm.mg[z]"
)
},
"lh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.whit[e]"
)
},
"rh_white": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.whit[e]"
)
},
"lh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/lh.pia[l]"
)
},
"rh_pial": {
"pattern": (
"derivatives/freesurfer/[!f]{subject}/surf/rh.pia[l]"
)
},
},
"Warp": [
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_from-MNI152NLin2009cAsym_to-T1w_"
"mode-image_xfm.h5"
),
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_from-T1w_to-MNI152NLin2009cAsym_"
"mode-image_xfm.h5"
),
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
],
}
if space == "native": def validate_datagrabber_params(self) -> None:
patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym" """Run extra logical validation for datagrabber."""
else: # Descriptor for space in `anat`
patterns["BOLD"]["prewarp_space"] = "native" sp_anat_desc = (
"" if self.space == "native" else "space-MNI152NLin2009cAsym_"
# Use native T1w assets
self.space = space
# Set default types
if types is None:
types = list(patterns.keys())
# Convert single type into list
else:
if not isinstance(types, list):
types = [types]
# The replacements
replacements = ["subject", "task"]
uri = "https://github.com/OpenNeuroDatasets/ds002790"
super().__init__(
types=types,
datadir=datadir,
uri=uri,
patterns=patterns,
replacements=replacements,
confounds_format="fmriprep",
) )
# Descriptor for space in `func`
sp_func_desc = (
"space-T1w_"
if self.space == "native"
else "space-MNI152NLin2009cAsym_"
)
self.patterns["BOLD"]["pattern"] = self.patterns["BOLD"][
"pattern"
].replace("{sp_func_desc}", sp_func_desc)
self.patterns["BOLD"]["mask"]["pattern"] = self.patterns["BOLD"][
"mask"
]["pattern"].replace("{sp_func_desc}", sp_func_desc)
self.patterns["BOLD"]["reference"]["pattern"] = self.patterns["BOLD"][
"reference"
]["pattern"].replace("{sp_func_desc}", sp_func_desc)
self.patterns["T1w"]["pattern"] = self.patterns["T1w"][
"pattern"
].replace("{sp_anat_desc}", sp_anat_desc)
self.patterns["T1w"]["mask"]["pattern"] = self.patterns["T1w"]["mask"][
"pattern"
].replace("{sp_anat_desc}", sp_anat_desc)
for t in ["VBM_CSF", "VBM_GM", "VBM_WM"]:
self.patterns[t]["pattern"] = self.patterns[t]["pattern"].replace(
"{sp_anat_desc}", sp_anat_desc
)
for t in ["BOLD", "T1w"]:
self.patterns[t]["space"] = self.space
self.patterns[t]["mask"]["space"] = self.space
for t in ["VBM_CSF", "VBM_GM", "VBM_WM"]:
self.patterns[t]["space"] = self.space
if self.space == "native":
self.patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym"
else:
self.patterns["BOLD"]["prewarp_space"] = "native"
super().validate_datagrabber_params()
def get_elements(self) -> list: def get_elements(self) -> list:
"""Implement fetching list of elements in the dataset. """Implement fetching list of elements in the dataset.
@ -277,7 +279,7 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber):
return elems return elems
def get_item(self, subject: str, task: str) -> dict: def get_item(self, subject: str, task: str) -> dict:
"""Index one element in the dataset. """Get the specified item from the dataset.
Parameters Parameters
---------- ----------

View file

@ -8,31 +8,32 @@
# License: AGPL # License: AGPL
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber.aomic.id1000 import DataladAOMICID1000 from junifer.datagrabber.aomic.id1000 import DataladAOMICID1000
URI = "https://gin.g-node.org/juaml/datalad-example-aomic1000" URI = AnyUrl("https://gin.g-node.org/juaml/datalad-example-aomic1000")
@pytest.mark.parametrize( @pytest.mark.parametrize(
"type_, nested_types, space", "type_, nested_types, space",
[ [
("BOLD", ["confounds", "mask", "reference"], "MNI152NLin2009cAsym"), ("BOLD", ["confounds", "mask", "reference"], "MNI152NLin2009cAsym"),
("BOLD", ["confounds", "mask", "reference"], "native"), (["BOLD"], ["confounds", "mask", "reference"], "native"),
("T1w", ["mask"], "MNI152NLin2009cAsym"), ("T1w", ["mask"], "MNI152NLin2009cAsym"),
("T1w", ["mask"], "native"), (["T1w"], ["mask"], "native"),
("VBM_CSF", None, "MNI152NLin2009cAsym"), ("VBM_CSF", None, "MNI152NLin2009cAsym"),
("VBM_CSF", None, "native"), (["VBM_CSF"], None, "native"),
("VBM_GM", None, "MNI152NLin2009cAsym"), ("VBM_GM", None, "MNI152NLin2009cAsym"),
("VBM_GM", None, "native"), (["VBM_GM"], None, "native"),
("VBM_WM", None, "MNI152NLin2009cAsym"), ("VBM_WM", None, "MNI152NLin2009cAsym"),
("DWI", None, "MNI152NLin2009cAsym"), (["DWI"], None, "MNI152NLin2009cAsym"),
("FreeSurfer", None, "MNI152NLin2009cAsym"), (["FreeSurfer"], None, "MNI152NLin2009cAsym"),
], ],
) )
def test_DataladAOMICID1000( def test_DataladAOMICID1000(
type_: str, type_: str | list[str],
nested_types: list[str] | None, nested_types: list[str] | None,
space: str, space: str,
) -> None: ) -> None:
@ -40,7 +41,7 @@ def test_DataladAOMICID1000(
Parameters Parameters
---------- ----------
type_ : str type_ : str or list of str
The parametrized type. The parametrized type.
nested_types : list of str or None nested_types : list of str or None
The parametrized nested types. The parametrized nested types.
@ -48,32 +49,29 @@ def test_DataladAOMICID1000(
The parametrized space. The parametrized space.
""" """
dg = DataladAOMICID1000(types=type_, space=space) dg = DataladAOMICID1000(uri=URI, types=type_, space=space)
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element data
out = dg[test_element] out = dg[test_element]
# Assert data type # Assert data type
assert type_ in out if isinstance(type_, str):
assert out[type_]["path"].exists() type_ = [type_]
assert out[type_]["path"].is_file() for t in type_:
# Asserts data type metadata assert t in out
assert "meta" in out[type_] assert out[t]["path"].exists()
meta = out[type_]["meta"] assert out[t]["path"].is_file()
assert "element" in meta # Asserts data type metadata
assert "subject" in meta["element"] assert "meta" in out[t]
assert test_element == meta["element"]["subject"] meta = out[t]["meta"]
# Assert nested data type if not None assert "element" in meta
if nested_types is not None: assert "subject" in meta["element"]
for nested_type in nested_types: assert test_element == meta["element"]["subject"]
assert out[type_][nested_type]["path"].exists() # Assert nested data type if not None
assert out[type_][nested_type]["path"].is_file() if nested_types is not None:
for nested_type in nested_types:
assert out[t][nested_type]["path"].exists()
assert out[t][nested_type]["path"].is_file()
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -102,28 +100,13 @@ def test_DataladAOMICID1000_partial_data_access(
The parametrized types. The parametrized types.
""" """
dg = DataladAOMICID1000(types=types) dg = DataladAOMICID1000(uri=URI, types=types)
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element data
out = dg[test_element] out = dg[test_element]
# Assert data type # Assert data type
if isinstance(types, list): if isinstance(types, str):
for type_ in types: types = [types]
assert type_ in out for t in types:
else: assert t in out
assert types in out
def test_DataladAOMICID1000_incorrect_data_type() -> None:
"""Test DataladAOMICID1000 DataGrabber incorrect data type."""
with pytest.raises(
ValueError, match="`patterns` must contain all `types`"
):
_ = DataladAOMICID1000(types="Scooby-Doo")

View file

@ -8,11 +8,12 @@
# License: AGPL # License: AGPL
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber import DataladAOMICPIOP1 from junifer.datagrabber import DataladAOMICPIOP1
URI = "https://gin.g-node.org/juaml/datalad-example-aomicpiop1" URI = AnyUrl("https://gin.g-node.org/juaml/datalad-example-aomicpiop1")
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -21,18 +22,11 @@ URI = "https://gin.g-node.org/juaml/datalad-example-aomicpiop1"
( (
"BOLD", "BOLD",
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
None, "anticipation",
"MNI152NLin2009cAsym",
),
("BOLD", ["confounds", "mask", "reference"], None, "native"),
(
"BOLD",
["confounds", "mask", "reference"],
["anticipation"],
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
( (
"BOLD", ["BOLD"],
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
["emomatching", "faces"], ["emomatching", "faces"],
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
@ -40,97 +34,88 @@ URI = "https://gin.g-node.org/juaml/datalad-example-aomicpiop1"
( (
"BOLD", "BOLD",
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
["restingstate"], "restingstate",
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
( (
"BOLD", ["BOLD"],
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
["workingmemory", "gstroop"], ["workingmemory", "gstroop"],
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
( (
"BOLD", ["BOLD"],
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
["anticipation", "faces", "restingstate"], ["anticipation", "faces", "restingstate"],
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
("T1w", ["mask"], None, "MNI152NLin2009cAsym"), (["T1w"], ["mask"], "restingstate", "MNI152NLin2009cAsym"),
("T1w", ["mask"], None, "native"), ("T1w", ["mask"], ["restingstate"], "native"),
("VBM_CSF", None, None, "MNI152NLin2009cAsym"), (["VBM_CSF"], None, "restingstate", "MNI152NLin2009cAsym"),
("VBM_CSF", None, None, "native"), ("VBM_CSF", None, ["restingstate"], "native"),
("VBM_GM", None, None, "MNI152NLin2009cAsym"), (["VBM_GM"], None, "restingstate", "MNI152NLin2009cAsym"),
("VBM_GM", None, None, "native"), ("VBM_GM", None, ["restingstate"], "native"),
("VBM_WM", None, None, "MNI152NLin2009cAsym"), (["VBM_WM"], None, "restingstate", "MNI152NLin2009cAsym"),
("VBM_WM", None, None, "native"), ("VBM_WM", None, ["restingstate"], "native"),
("DWI", None, None, "MNI152NLin2009cAsym"), (["DWI"], None, "restingstate", "MNI152NLin2009cAsym"),
("FreeSurfer", None, None, "MNI152NLin2009cAsym"), (["FreeSurfer"], None, ["restingstate"], "MNI152NLin2009cAsym"),
], ],
) )
def test_DataladAOMICPIOP1( def test_DataladAOMICPIOP1(
type_: str, type_: str | list[str],
nested_types: list[str] | None, nested_types: list[str] | None,
tasks: list[str] | None, tasks: str | list[str],
space: str, space: str,
) -> None: ) -> None:
"""Test DataladAOMICPIOP1 DataGrabber. """Test DataladAOMICPIOP1 DataGrabber.
Parameters Parameters
---------- ----------
type_ : str type_ : str or list of str
The parametrized type. The parametrized type.
nested_types : list of str or None nested_types : list of str or None
The parametrized nested types. The parametrized nested types.
tasks : list of str or None tasks : str or list of str
The parametrized task values. The parametrized task values.
space: str space: str
The parametrized space. The parametrized space.
""" """
dg = DataladAOMICPIOP1(types=type_, tasks=tasks, space=space) dg = DataladAOMICPIOP1(uri=URI, types=type_, tasks=tasks, space=space)
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element data
out = dg[test_element]
# Get all elements
all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0]
# Get test element data
out = dg[test_element] out = dg[test_element]
# Assert data type # Assert data type
assert type_ in out if isinstance(type_, str):
# Check task name if BOLD type_ = [type_]
if type_ == "BOLD" and tasks is not None: for t in type_:
# Depending on task 'acquisition is different' assert t in out
task_acqs = { # Check task name if BOLD
"anticipation": "seq", if t == "BOLD":
"emomatching": "seq", # Depending on task 'acquisition is different'
"faces": "mb3", task_acqs = {
"gstroop": "seq", "anticipation": "seq",
"restingstate": "mb3", "emomatching": "seq",
"workingmemory": "seq", "faces": "mb3",
} "gstroop": "seq",
assert task_acqs[test_element[1]] in out[type_]["path"].name "restingstate": "mb3",
assert out[type_]["path"].exists() "workingmemory": "seq",
assert out[type_]["path"].is_file() }
# Asserts data type metadata assert task_acqs[test_element[1]] in out[t]["path"].name
assert "meta" in out[type_] assert out[t]["path"].exists()
meta = out[type_]["meta"] assert out[t]["path"].is_file()
assert "element" in meta # Asserts data type metadata
assert "subject" in meta["element"] assert "meta" in out[t]
assert test_element[0] == meta["element"]["subject"] meta = out[t]["meta"]
# Assert nested data type if not None assert "element" in meta
if nested_types is not None: assert "subject" in meta["element"]
for nested_type in nested_types: assert test_element[0] == meta["element"]["subject"]
assert out[type_][nested_type]["path"].exists() # Assert nested data type if not None
assert out[type_][nested_type]["path"].is_file() if nested_types is not None:
for nested_type in nested_types:
assert out[t][nested_type]["path"].exists()
assert out[t][nested_type]["path"].is_file()
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -159,40 +144,13 @@ def test_DataladAOMICPIOP1_partial_data_access(
The parametrized types. The parametrized types.
""" """
dg = DataladAOMICPIOP1(types=types) dg = DataladAOMICPIOP1(uri=URI, types=types)
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element data
out = dg[test_element] out = dg[test_element]
# Assert data type # Assert data type
if isinstance(types, list): if isinstance(types, str):
for type_ in types: types = [types]
assert type_ in out for t in types:
else: assert t in out
assert types in out
def test_DataladAOMICPIOP1_incorrect_data_type() -> None:
"""Test DataladAOMICPIOP1 DataGrabber incorrect data type."""
with pytest.raises(
ValueError, match="`patterns` must contain all `types`"
):
_ = DataladAOMICPIOP1(types="Ceres")
def test_DataladAOMICPIOP1_invalid_tasks():
"""Test DataladAOMICIDPIOP1 DataGrabber invalid tasks."""
with pytest.raises(
ValueError,
match=(
"thisisnotarealtask is not a valid task in "
"the AOMIC PIOP1 dataset!"
),
):
DataladAOMICPIOP1(tasks="thisisnotarealtask")

View file

@ -8,11 +8,12 @@
# License: AGPL # License: AGPL
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber import DataladAOMICPIOP2 from junifer.datagrabber import DataladAOMICPIOP2
URI = "https://gin.g-node.org/juaml/datalad-example-aomicpiop2" URI = AnyUrl("https://gin.g-node.org/juaml/datalad-example-aomicpiop2")
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -21,18 +22,11 @@ URI = "https://gin.g-node.org/juaml/datalad-example-aomicpiop2"
( (
"BOLD", "BOLD",
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
None, "restingstate",
"MNI152NLin2009cAsym",
),
("BOLD", ["confounds", "mask", "reference"], None, "native"),
(
"BOLD",
["confounds", "mask", "reference"],
["restingstate"],
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
( (
"BOLD", ["BOLD"],
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
["restingstate", "stopsignal"], ["restingstate", "stopsignal"],
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
@ -44,72 +38,69 @@ URI = "https://gin.g-node.org/juaml/datalad-example-aomicpiop2"
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
( (
"BOLD", ["BOLD"],
["confounds", "mask", "reference"], ["confounds", "mask", "reference"],
["workingmemory"], "workingmemory",
"MNI152NLin2009cAsym", "MNI152NLin2009cAsym",
), ),
("T1w", ["mask"], None, "MNI152NLin2009cAsym"), (["T1w"], ["mask"], "restingstate", "MNI152NLin2009cAsym"),
("T1w", ["mask"], None, "native"), ("T1w", ["mask"], ["restingstate"], "native"),
("VBM_CSF", None, None, "MNI152NLin2009cAsym"), (["VBM_CSF"], None, "restingstate", "MNI152NLin2009cAsym"),
("VBM_CSF", None, None, "native"), ("VBM_CSF", None, ["restingstate"], "native"),
("VBM_GM", None, None, "MNI152NLin2009cAsym"), (["VBM_GM"], None, "restingstate", "MNI152NLin2009cAsym"),
("VBM_GM", None, None, "native"), ("VBM_GM", None, ["restingstate"], "native"),
("VBM_WM", None, None, "MNI152NLin2009cAsym"), (["VBM_WM"], None, "restingstate", "MNI152NLin2009cAsym"),
("VBM_WM", None, None, "native"), ("VBM_WM", None, ["restingstate"], "native"),
("DWI", None, None, "MNI152NLin2009cAsym"), (["DWI"], None, "restingstate", "MNI152NLin2009cAsym"),
("FreeSurfer", None, None, "MNI152NLin2009cAsym"), (["FreeSurfer"], None, ["restingstate"], "MNI152NLin2009cAsym"),
], ],
) )
def test_DataladAOMICPIOP2( def test_DataladAOMICPIOP2(
type_: str, type_: str | list[str],
nested_types: list[str] | None, nested_types: list[str] | None,
tasks: list[str] | None, tasks: str | list[str],
space: str, space: str,
) -> None: ) -> None:
"""Test DataladAOMICPIOP2 DataGrabber. """Test DataladAOMICPIOP2 DataGrabber.
Parameters Parameters
---------- ----------
type_ : str type_ : str or list of str
The parametrized type. The parametrized type.
nested_types : list of str or None nested_types : list of str or None
The parametrized nested types. The parametrized nested types.
tasks : list of str or None tasks : str or list of str
The parametrized task values. The parametrized task values.
space: str space: str
The parametrized space. The parametrized space.
""" """
dg = DataladAOMICPIOP2(types=type_, tasks=tasks, space=space) dg = DataladAOMICPIOP2(uri=URI, types=type_, tasks=tasks, space=space)
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element data
out = dg[test_element] out = dg[test_element]
# Assert data type # Assert data type
assert type_ in out if isinstance(type_, str):
# Check task name if BOLD type_ = [type_]
if type_ == "BOLD" and tasks is not None: for t in type_:
assert test_element[1] in out[type_]["path"].name assert t in out
assert out[type_]["path"].exists() # Check task name if BOLD
assert out[type_]["path"].is_file() if t == "BOLD":
# Asserts data type metadata assert test_element[1] in out[t]["path"].name
assert "meta" in out[type_] assert out[t]["path"].exists()
meta = out[type_]["meta"] assert out[t]["path"].is_file()
assert "element" in meta # Asserts data type metadata
assert "subject" in meta["element"] assert "meta" in out[t]
assert test_element[0] == meta["element"]["subject"] meta = out[t]["meta"]
# Assert nested data type if not None assert "element" in meta
if nested_types is not None: assert "subject" in meta["element"]
for nested_type in nested_types: assert test_element[0] == meta["element"]["subject"]
assert out[type_][nested_type]["path"].exists() # Assert nested data type if not None
assert out[type_][nested_type]["path"].is_file() if nested_types is not None:
for nested_type in nested_types:
assert out[t][nested_type]["path"].exists()
assert out[t][nested_type]["path"].is_file()
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -138,40 +129,13 @@ def test_DataladAOMICPIOP2_partial_data_access(
The parametrized types. The parametrized types.
""" """
dg = DataladAOMICPIOP2(types=types) dg = DataladAOMICPIOP2(uri=URI, types=types)
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element data
out = dg[test_element] out = dg[test_element]
# Assert data type # Assert data type
if isinstance(types, list): if isinstance(types, str):
for type_ in types: types = [types]
assert type_ in out for t_ in types:
else: assert t_ in out
assert types in out
def test_DataladAOMICPIOP2_incorrect_data_type() -> None:
"""Test DataladAOMICPIOP2 DataGrabber incorrect data type."""
with pytest.raises(
ValueError, match="`patterns` must contain all `types`"
):
_ = DataladAOMICPIOP2(types="Vesta")
def test_DataladAOMICPIOP2_invalid_tasks():
"""Test DataladAOMICIDPIOP2 DataGrabber invalid tasks."""
with pytest.raises(
ValueError,
match=(
"thisisnotarealtask is not a valid task in "
"the AOMIC PIOP2 dataset!"
),
):
DataladAOMICPIOP2(tasks="thisisnotarealtask")

View file

@ -7,56 +7,77 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Iterator from collections.abc import Iterator
from enum import Enum
from pathlib import Path from pathlib import Path
from typing import Annotated, Any
from aenum import Enum as AEnum
from pydantic import BaseModel, BeforeValidator, ConfigDict, Field
from ..pipeline import UpdateMetaMixin from ..pipeline import UpdateMetaMixin
from ..typing import Element, Elements from ..typing import Element, Elements
from ..utils import logger, raise_error from ..utils import ensure_list, logger, raise_error
__all__ = ["BaseDataGrabber"] __all__ = ["BaseDataGrabber", "DataType"]
class BaseDataGrabber(ABC, UpdateMetaMixin): class DataType(str, AEnum):
"""Abstract base class for DataGrabber. """Accepted data type."""
fraimondo commented 2025-11-26 10:57:56 +00:00 (Migrated from github.com)

Now we hit an important issue here.

How can I create a junifer extension that allows to process a new DataType (e.g. EEG)?

Now we hit an important issue here. How can I create a junifer extension that allows to process a new DataType (e.g. EEG)?
synchon commented 2025-11-27 07:42:43 +00:00 (Migrated from github.com)

Same way as one does now.

Same way as one does now.
For every interface that is required, one needs to provide a concrete T1w = "T1w"
T2w = "T2w"
BOLD = "BOLD"
Warp = "Warp"
VBM_GM = "VBM_GM"
VBM_WM = "VBM_WM"
VBM_CSF = "VBM_CSF"
FALFF = "fALFF"
GCOR = "GCOR"
LCOR = "LCOR"
DWI = "DWI"
FreeSurfer = "FreeSurfer"
class BaseDataGrabber(BaseModel, ABC, UpdateMetaMixin):
"""Abstract base class for data fetcher.
For every datagrabber, one needs to provide a concrete
implementation of this abstract class. implementation of this abstract class.
Parameters Parameters
---------- ----------
types : list of str types : :enum:`.DataType` or list of variants
The types of data to be grabbed. The data type(s) to grab.
datadir : str or pathlib.Path datadir : pathlib.Path
The directory where the data is / will be stored. The path where the data is or will be stored.
Raises
------
TypeError
If ``types`` is not a list or if the values are not string.
""" """
def __init__(self, types: list[str], datadir: str | Path) -> None: model_config = ConfigDict(use_enum_values=True)
# Validate types
if not isinstance(types, list):
raise_error(msg="`types` must be a list", klass=TypeError)
if any(not isinstance(x, str) for x in types):
raise_error(
msg="`types` must be a list of strings", klass=TypeError
)
self.types = types
# Convert str to Path types: Annotated[
if not isinstance(datadir, Path): DataType | list[DataType],
datadir = Path(datadir) Field(frozen=True),
self._datadir = datadir BeforeValidator(ensure_list),
]
datadir: Path
def model_post_init(self, context: Any): # noqa: D102
logger.debug("Initializing BaseDataGrabber") logger.debug("Initializing BaseDataGrabber")
logger.debug(f"\t_datadir = {datadir}") logger.debug(f"\tdatadir = {self.datadir}")
logger.debug(f"\ttypes = {types}") logger.debug(f"\ttypes = {self.types}")
# Run extra validation for datagrabbers and fail early if needed
self.validate_datagrabber_params()
def __iter__(self) -> Iterator: def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber.
Subclasses can override to provide validation.
"""
pass
def __iter__(self) -> Iterator[Elements]:
"""Enable iterable support. """Enable iterable support.
Yields Yields
@ -72,7 +93,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
Parameters Parameters
---------- ----------
element : str or tuple of str element : `Element`
The element to be indexed. The element to be indexed.
Returns Returns
@ -82,10 +103,14 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
specified element. specified element.
""" """
# Convert element to tuple if not already and extract enum values if
# present
element = (
(element,)
if not isinstance(element, tuple)
else tuple(i.value if isinstance(i, Enum) else i for i in element)
)
logger.info(f"Getting element {element}") logger.info(f"Getting element {element}")
# Convert element to tuple if not already
if not isinstance(element, tuple):
element = (element,)
# Zip through element keys and actual values to construct element # Zip through element keys and actual values to construct element
# access dictionary # access dictionary
named_element: dict = dict( named_element: dict = dict(
@ -120,29 +145,30 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
Returns Returns
------- -------
list of str list of str
The types of data to be grabbed. The data type(s) to grab.
""" """
return self.types.copy() return [x.value if isinstance(x, Enum) else x for x in self.types]
@property @property
def datadir(self) -> Path: def fulldir(self) -> Path:
"""Get data directory path. """Get complete data directory path.
Returns Returns
------- -------
pathlib.Path pathlib.Path
Path to the data directory. Can be overridden by subclasses. Complete path to the data directory.
Can be overridden by subclasses.
""" """
return self._datadir return self.datadir
def filter(self, selection: Elements) -> Iterator: def filter(self, selection: Elements) -> Iterator:
"""Filter elements to be grabbed. """Filter elements to be grabbed.
Parameters Parameters
---------- ----------
selection : list selection : ``Elements``
The list of partial or complete element selectors to filter using. The list of partial or complete element selectors to filter using.
Yields Yields
@ -157,7 +183,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
Parameters Parameters
---------- ----------
element : str or tuple of str element : ``Elements``
The element to be filtered. The element to be filtered.
Returns Returns

View file

@ -9,12 +9,15 @@ import atexit
import os import os
import tempfile import tempfile
from pathlib import Path from pathlib import Path
from typing import Any, NoReturn
import datalad import datalad
import datalad.api as dl import datalad.api as dl
from datalad.support.exceptions import IncompleteResultsError from datalad.support.exceptions import IncompleteResultsError
from datalad.support.gitrepo import GitRepo from datalad.support.gitrepo import GitRepo
from pydantic import AnyUrl, Field, field_validator
from ..api.decorators import register_datagrabber
from ..pipeline import WorkDirManager from ..pipeline import WorkDirManager
from ..typing import Element from ..typing import Element
from ..utils import config, logger, raise_error, warn_with_log from ..utils import config, logger, raise_error, warn_with_log
@ -24,6 +27,26 @@ from .base import BaseDataGrabber
__all__ = ["DataladDataGrabber"] __all__ = ["DataladDataGrabber"]
def _create_datadir() -> Path:
"""Create a temporary directory for datalad dataset."""
datadir = WorkDirManager().get_tempdir(
prefix="datalad", suffix="juniferauto"
)
logger.info(
"Created a temporary directory for datalad dataset at: "
f"{datadir.resolve()!s}"
)
return datadir
def _remove_datadir(datadir: Path) -> None:
"""Remove temporary directory if it exists."""
if datadir.exists():
logger.debug(f"Removing temporary directory at: {datadir.resolve()!s}")
WorkDirManager().delete_tempdir(datadir)
@register_datagrabber
class DataladDataGrabber(BaseDataGrabber): class DataladDataGrabber(BaseDataGrabber):
"""Abstract base class for datalad-based data fetching. """Abstract base class for datalad-based data fetching.
@ -31,17 +54,15 @@ class DataladDataGrabber(BaseDataGrabber):
Parameters Parameters
---------- ----------
rootdir : str or pathlib.Path, optional uri : pydantic.AnyUrl
URI of the datalad sibling.
rootdir : pathlib.Path, optional
The path within the datalad dataset to the root directory The path within the datalad dataset to the root directory
(default "."). (default Path(".")).
datadir : str or pathlib.Path or None, optional datadir : pathlib.Path, optional
That directory where the datalad dataset will be cloned. If None, That path where the datalad dataset will be cloned.
the datalad dataset will be cloned into a temporary directory If not specified, the datalad dataset will be cloned into a temporary
(default None). directory.
uri : str or None, optional
URI of the datalad sibling (default None).
**kwargs
Keyword arguments passed to superclass.
Methods Methods
------- -------
@ -66,24 +87,62 @@ class DataladDataGrabber(BaseDataGrabber):
This class is intended to be used as a superclass of a subclass This class is intended to be used as a superclass of a subclass
with multiple inheritance. with multiple inheritance.
If the ``datadir`` is specified and has the stem prefix as ``"datalad"``
and the stem suffix as ``"juniferauto"``, it will be automatically
deleted after use.
""" """
def __init__( uri: AnyUrl = Field(frozen=True)
self, rootdir: Path = Field(frozen=True, default=Path("."))
rootdir: str | Path = ".", datadir: Path = Field(default_factory=lambda: _create_datadir())
datadir: str | Path | None = None, _repodir: Path = Path(".")
uri: str | None = None, # Flag to indicate if the dataset was cloned before and it might be
**kwargs, # dirty
): datalad_dirty: bool = False
if datadir is None: datalad_commit_id: str | None = None
logger.info("`datadir` is None, creating a temporary directory") datalad_id: str | None = None
# Create temporary directory _dataset: dl.Dataset | None = None
tmpdir = WorkDirManager().get_tempdir(prefix="datalad") _got_files: list[str] = [] # noqa: RUF012
self._tmpdir = tmpdir _was_cloned: bool = False
datadir = tmpdir / "datadir"
datadir.mkdir(parents=True, exist_ok=False) @field_validator("datadir", mode="after")
logger.info(f"`datadir` set to {datadir}") @classmethod
cache_dir = tmpdir / ".datalad_cache" def warn_existing_datalad_autodir(cls, value: Path) -> Path:
"""Warn if existing datalad autodir exists."""
if value.stem.startswith("datalad") and value.stem.endswith(
"juniferauto"
):
warn_with_log(
f"{value.resolve()!s} already exists and will reuse assets "
"from previous run."
)
return value
@field_validator("datalad_dirty", mode="before")
@classmethod
def disable_tag(cls, value: Any) -> NoReturn:
"""Disable setting datalad_dirty directly."""
raise_error(
msg="datalad_dirty cannot be set directly",
klass=ValueError,
)
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
logger.debug("Initializing DataladDataGrabber")
logger.debug(f"\turi = {self.uri}")
logger.debug(f"\trootdir = {self.rootdir}")
if self.datadir.stem.startswith(
"datalad"
) and self.datadir.stem.endswith("juniferauto"):
self._repodir = self.datadir / "dataset"
self._repodir.mkdir(parents=True, exist_ok=False)
logger.info(
"Datalad dataset installation path set to: "
f"{self._repodir.resolve()!s}"
fraimondo commented 2025-11-26 11:09:08 +00:00 (Migrated from github.com)

this is horrible, what if the datadir was set by the user to datalad_dataset_aomic?

this is horrible, what if the datadir was set by the user to `datalad_dataset_aomic`?
synchon commented 2025-11-27 08:15:16 +00:00 (Migrated from github.com)

Hmm that's a fair argument, I can do something like:

datadir = WorkDirManager().get_tempdir(
        prefix="datalad", suffix="juniferauto"
)

and then check with:

if self.datadir.stem.startswith(
    "datalad"
) and self.datadir.stem.endswith("juniferauto"):
Hmm that's a fair argument, I can do something like: ```python datadir = WorkDirManager().get_tempdir( prefix="datalad", suffix="juniferauto" ) ``` and then check with: ```python if self.datadir.stem.startswith( "datalad" ) and self.datadir.stem.endswith("juniferauto"): ```
fraimondo commented 2025-12-01 08:06:29 +00:00 (Migrated from github.com)

why can't we keep a private var _was_cloned and check that? What if we "run a test without cleaning the workdir" and then we use it?

The behaviour should be "clean whatever mess you made".

why can't we keep a private var `_was_cloned` and check that? What if we "run a test without cleaning the workdir" and then we use it? The behaviour should be "clean whatever mess you made".
synchon commented 2025-12-02 08:47:12 +00:00 (Migrated from github.com)

why can't we keep a private var _was_cloned and check that?

We already have _was_cloned for a similar purpose. I don't think I understand how you want it to be implemented.

> why can't we keep a private var `_was_cloned` and check that? We already have `_was_cloned` for a similar purpose. I don't think I understand how you want it to be implemented.
)
cache_dir = self.datadir / ".datalad_cache"
sockets_dir = cache_dir / "sockets" sockets_dir = cache_dir / "sockets"
locks_dir = cache_dir / "locks" locks_dir = cache_dir / "locks"
sockets_dir.mkdir(parents=True, exist_ok=False) sockets_dir.mkdir(parents=True, exist_ok=False)
@ -105,43 +164,29 @@ class DataladDataGrabber(BaseDataGrabber):
"Datalad locks set to " "Datalad locks set to "
f"{datalad.cfg.get('datalad.locations.locks')}" f"{datalad.cfg.get('datalad.locations.locks')}"
) )
atexit.register(self._rmtmpdir) atexit.register(_remove_datadir, self.datadir)
# TODO: uri can be converted to a positional argument else:
if uri is None: self._repodir = self.datadir
raise_error("`uri` must be provided") super().validate_datagrabber_params()
super().__init__(datadir=datadir, **kwargs)
logger.debug("Initializing DataladDataGrabber")
logger.debug(f"\turi = {uri}")
logger.debug(f"\t_rootdir = {rootdir}")
self.uri = uri
self._rootdir = rootdir
# Flag to indicate if the dataset was cloned before and it might be
# dirty
self.datalad_dirty = False
def __del__(self) -> None: def __del__(self) -> None:
"""Destructor.""" """Destructor."""
if hasattr(self, "_tmpdir"): if self.datadir.stem.startswith(
self._rmtmpdir() "datalad"
) and self.datadir.stem.endswith("juniferauto"):
def _rmtmpdir(self) -> None: _remove_datadir(self.datadir)
"""Remove temporary directory if it exists."""
if self._tmpdir.exists():
logger.debug("Removing temporary directory")
WorkDirManager().delete_tempdir(self._tmpdir)
@property @property
def datadir(self) -> Path: def fulldir(self) -> Path:
"""Get data directory path. """Get complete data directory path.
Returns Returns
------- -------
pathlib.Path pathlib.Path
Path to the data directory. Complete path to the data directory.
""" """
return super().datadir / self._rootdir return self._repodir / self.rootdir
def _get_dataset_id_remote(self) -> tuple[str, bool]: def _get_dataset_id_remote(self) -> tuple[str, bool]:
"""Get the dataset ID from the remote. """Get the dataset ID from the remote.
@ -164,8 +209,10 @@ class DataladDataGrabber(BaseDataGrabber):
with tempfile.TemporaryDirectory() as tmpdir: with tempfile.TemporaryDirectory() as tmpdir:
if not config.get("datagrabber.skipidcheck", False): if not config.get("datagrabber.skipidcheck", False):
logger.debug(f"Querying {self.uri} for dataset ID") logger.debug(f"Querying {self.uri} for dataset ID")
repo = GitRepo.clone( repo: GitRepo = GitRepo.clone(
self.uri, path=tmpdir, clone_options=["-n", "--depth=1"] str(self.uri),
path=tmpdir,
clone_options=["-n", "--depth=1"],
) )
repo.checkout(name=".datalad/config", options=["HEAD"]) repo.checkout(name=".datalad/config", options=["HEAD"])
remote_id = repo.config.get("datalad.dataset.id", None) remote_id = repo.config.get("datalad.dataset.id", None)
@ -178,10 +225,11 @@ class DataladDataGrabber(BaseDataGrabber):
is_dirty = False is_dirty = False
else: else:
logger.debug("Skipping dataset ID check") logger.debug("Skipping dataset ID check")
# Should be already set to the dataset
remote_id = self._dataset.id remote_id = self._dataset.id
is_dirty = False is_dirty = False
logger.debug( logger.debug(
f"Remote dataset is {'' if is_dirty else 'not'} dirty" f"Remote dataset is {'dirty' if is_dirty else 'not dirty'}"
) )
if remote_id is None: if remote_id is None:
raise_error("Could not get dataset ID from remote") raise_error("Could not get dataset ID from remote")
@ -251,7 +299,7 @@ class DataladDataGrabber(BaseDataGrabber):
return out return out
def install(self) -> None: def install(self) -> None:
"""Install the datalad dataset into the ``datadir``. """Installs the datalad dataset.
Raises Raises
------ ------
@ -261,12 +309,10 @@ class DataladDataGrabber(BaseDataGrabber):
If there is a datalad-related problem while cloning dataset. If there is a datalad-related problem while cloning dataset.
""" """
isinstalled = dl.Dataset(self._datadir).is_installed() is_installed = dl.Dataset(self._repodir).is_installed()
if isinstalled: if is_installed:
logger.debug("Dataset already installed") logger.debug("Dataset already installed")
self._got_files = [] self._dataset = dl.Dataset(self._repodir)
self._dataset: dl.Dataset = dl.Dataset(self._datadir)
# Check if dataset is already installed with a different ID # Check if dataset is already installed with a different ID
remote_id, is_dirty = self._get_dataset_id_remote() remote_id, is_dirty = self._get_dataset_id_remote()
if remote_id != self._dataset.id: if remote_id != self._dataset.id:
@ -274,7 +320,6 @@ class DataladDataGrabber(BaseDataGrabber):
"Dataset already installed but with a different " "Dataset already installed but with a different "
f"ID: {self._dataset.id} (local) != {remote_id} (remote)" f"ID: {self._dataset.id} (local) != {remote_id} (remote)"
) )
# Conditional reporting on dataset dirtiness # Conditional reporting on dataset dirtiness
self.datalad_dirty = is_dirty self.datalad_dirty = is_dirty
if self.datalad_dirty: if self.datalad_dirty:
@ -286,18 +331,18 @@ class DataladDataGrabber(BaseDataGrabber):
logger.debug(f"Dataset (id: {self._dataset.id}) is clean") logger.debug(f"Dataset (id: {self._dataset.id}) is clean")
else: else:
logger.debug(f"Installing dataset {self.uri} to {self._datadir}") logger.debug(f"Installing dataset {self.uri} to {self._repodir}")
try: try:
self._dataset: dl.Dataset = dl.clone( # type: ignore self._dataset = dl.clone(
self.uri, self._datadir, result_renderer="disabled" self.uri, self._repodir, result_renderer="disabled"
) )
except IncompleteResultsError as e: except IncompleteResultsError as e:
raise_error(f"Failed to clone dataset: {e.failed}") raise_error(f"Failed to clone dataset: {e.failed}")
logger.debug("Dataset installed") logger.debug("Dataset installed")
self._was_cloned = not isinstalled self._was_cloned = not is_installed
# Dataset should be set already
self.datalad_commit_id = self._dataset.repo.get_hexsha( # type: ignore self.datalad_commit_id = self._dataset.repo.get_hexsha(
self._dataset.repo.get_corresponding_branch() # type: ignore self._dataset.repo.get_corresponding_branch()
) )
self.datalad_id = self._dataset.id self.datalad_id = self._dataset.id
@ -320,7 +365,7 @@ class DataladDataGrabber(BaseDataGrabber):
Parameters Parameters
---------- ----------
element : str or tuple of str element : `Element`
The element to be indexed. If one string is provided, it is The element to be indexed. If one string is provided, it is
assumed to be a tuple with only one item. If a tuple is provided, assumed to be a tuple with only one item. If a tuple is provided,
each item in the tuple is the value for the replacement string each item in the tuple is the value for the replacement string

View file

@ -3,15 +3,93 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from enum import Enum
from itertools import product from itertools import product
from pathlib import Path from typing import Annotated, Literal
from pydantic import AnyUrl, BeforeValidator
from ..api.decorators import register_datagrabber from ..api.decorators import register_datagrabber
from ..utils import raise_error from ..typing import DataGrabberPatterns
from ..utils import ensure_list
from .base import DataType
from .pattern import ConfoundsFormat
from .pattern_datalad import PatternDataladDataGrabber from .pattern_datalad import PatternDataladDataGrabber
__all__ = ["DMCC13Benchmark"] __all__ = [
"DMCC13Benchmark",
"DMCCPhaseEncoding",
"DMCCRun",
"DMCCSession",
"DMCCTask",
]
class DMCCSession(str, Enum):
"""Accepted DMCC sessions."""
Wave1Bas = "ses-wave1bas"
Wave1Pro = "ses-wave1pro"
Wave1Rea = "ses-wave1rea"
class DMCCTask(str, Enum):
"""Accepted DMCC tasks."""
Rest = "Rest"
Axcpt = "Axcpt"
Cuedts = "Cuedts"
Stern = "Stern"
Stroop = "Stroop"
class DMCCPhaseEncoding(str, Enum):
"""Accepted DMCC phase encoding directions."""
AP = "AP"
PA = "PA"
class DMCCRun(str, Enum):
"""Accepted DMCC runs."""
One = "1"
Two = "2"
_types = Literal[
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.Warp,
]
_sessions = Literal[
DMCCSession.Wave1Bas,
DMCCSession.Wave1Pro,
DMCCSession.Wave1Rea,
]
_tasks = Literal[
DMCCTask.Rest,
DMCCTask.Axcpt,
DMCCTask.Cuedts,
DMCCTask.Stern,
DMCCTask.Stroop,
]
_phase_encodings = Literal[
DMCCPhaseEncoding.AP,
DMCCPhaseEncoding.PA,
]
_runs = Literal[
DMCCRun.One,
DMCCRun.Two,
]
@register_datagrabber @register_datagrabber
@ -20,191 +98,142 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or Path or None, optional types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "Warp"} or \
The directory where the datalad dataset will be cloned. If None, list of the options, optional
the datalad dataset will be cloned into a temporary directory The data type(s) to grab.
(default None). datadir : pathlib.Path, optional
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM"} or \ That path where the datalad dataset will be cloned.
list of the options, optional If not specified, the datalad dataset will be cloned into a temporary
DMCC data types. If None, all available data types are selected. directory.
(default None). sessions : {"ses-wave1bas", "ses-wave1pro", "ses-wave1rea"} or \
sessions: {"ses-wave1bas", "ses-wave1pro", "ses-wave1rea"} or \ list of the options, optional
list of the options, optional DMCC sessions.
DMCC sessions. If None, all available sessions are selected By default, all available sessions are selected.
(default None). tasks : {"Rest", "Axcpt", "Cuedts", "Stern", "Stroop"} or \
tasks: {"Rest", "Axcpt", "Cuedts", "Stern", "Stroop"} or \ list of the options, optional
list of the options, optional DMCC tasks.
DMCC task sessions. If None, all available task sessions are selected By default, all available tasks are selected.
(default None).
phase_encodings : {"AP", "PA"} or list of the options, optional phase_encodings : {"AP", "PA"} or list of the options, optional
DMCC phase encoding directions. If None, all available phase encodings DMCC phase encoding directions.
are selected (default None). By default, all available phase encodings are selected.
runs : {"1", "2"} or list of the options, optional runs : {"1", "2"} or list of the options, optional
DMCC runs. If None, all available runs are selected (default None). DMCC runs.
By default, all available runs are selected.
native_t1w : bool, optional native_t1w : bool, optional
Whether to use T1w in native space (default False). Whether to use T1w in native space (default False).
Raises
------
ValueError
If invalid value is passed for:
* ``sessions``
* ``tasks``
* ``phase_encodings``
* ``runs``
""" """
def __init__( uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds003452.git")
self, types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
datadir: str | Path | None = None, DataType.BOLD,
types: str | list[str] | None = None, DataType.T1w,
sessions: str | list[str] | None = None, DataType.VBM_CSF,
tasks: str | list[str] | None = None, DataType.VBM_GM,
phase_encodings: str | list[str] | None = None, DataType.VBM_WM,
runs: str | list[str] | None = None, ]
native_t1w: bool = False, sessions: Annotated[
) -> None: _sessions | list[_sessions], BeforeValidator(ensure_list)
# Declare all sessions ] = [ # noqa: RUF012
all_sessions = [ DMCCSession.Wave1Bas,
"ses-wave1bas", DMCCSession.Wave1Pro,
"ses-wave1pro", DMCCSession.Wave1Rea,
"ses-wave1rea", ]
] tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
# Set default sessions DMCCTask.Rest,
if sessions is None: DMCCTask.Axcpt,
sessions = all_sessions DMCCTask.Cuedts,
else: DMCCTask.Stern,
# Convert single session into list DMCCTask.Stroop,
if isinstance(sessions, str): ]
sessions = [sessions] phase_encodings: Annotated[
# Verify valid sessions _phase_encodings | list[_phase_encodings],
for s in sessions: BeforeValidator(ensure_list),
if s not in all_sessions: ] = [ # noqa: RUF012
raise_error( DMCCPhaseEncoding.AP,
f"{s} is not a valid session in the DMCC dataset" DMCCPhaseEncoding.PA,
) ]
self.sessions = sessions runs: Annotated[_runs | list[_runs], BeforeValidator(ensure_list)] = [ # noqa: RUF012
# Declare all tasks DMCCRun.One,
all_tasks = [ DMCCRun.Two,
"Rest", ]
"Axcpt", native_t1w: bool = False
"Cuedts", patterns: DataGrabberPatterns = { # noqa: RUF012
"Stern", "BOLD": {
"Stroop", "pattern": (
] "derivatives/fmriprep-1.3.2/{subject}/{session}/"
# Set default tasks "func/{subject}_{session}_task-{task}_acq-mb4"
if tasks is None: "{phase_encoding}_run-{run}_"
tasks = all_tasks "space-MNI152NLin2009cAsym_desc-preproc_bold.nii.gz"
else: ),
# Convert single task into list "space": "MNI152NLin2009cAsym",
if isinstance(tasks, str): "mask": {
tasks = [tasks]
# Verify valid tasks
for t in tasks:
if t not in all_tasks:
raise_error(f"{t} is not a valid task in the DMCC dataset")
self.tasks = tasks
# Declare all phase encodings
all_phase_encodings = ["AP", "PA"]
# Set default phase encodings
if phase_encodings is None:
phase_encodings = all_phase_encodings
else:
# Convert single phase encoding into list
if isinstance(phase_encodings, str):
phase_encodings = [phase_encodings]
# Verify valid phase encodings
for p in phase_encodings:
if p not in all_phase_encodings:
raise_error(
f"{p} is not a valid phase encoding in the DMCC "
"dataset"
)
self.phase_encodings = phase_encodings
# Declare all runs
all_runs = ["1", "2"]
# Set default runs
if runs is None:
runs = all_runs
else:
# Convert single run into list
if isinstance(runs, str):
runs = [runs]
# Verify valid runs
for r in runs:
if r not in all_runs:
raise_error(f"{r} is not a valid run in the DMCC dataset")
self.runs = runs
# The patterns
patterns = {
"BOLD": {
"pattern": ( "pattern": (
"derivatives/fmriprep-1.3.2/{subject}/{session}/" "derivatives/fmriprep-1.3.2/{subject}/{session}/"
"func/{subject}_{session}_task-{task}_acq-mb4" "func/{subject}_{session}_task-{task}_acq-mb4"
"{phase_encoding}_run-{run}_" "{phase_encoding}_run-{run}_"
"space-MNI152NLin2009cAsym_desc-preproc_bold.nii.gz" "space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz"
), ),
"space": "MNI152NLin2009cAsym", "space": "MNI152NLin2009cAsym",
"mask": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/{session}/"
"func/{subject}_{session}_task-{task}_acq-mb4"
"{phase_encoding}_run-{run}_"
"space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
"confounds": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/{session}/"
"func/{subject}_{session}_task-{task}_acq-mb4"
"{phase_encoding}_run-{run}_desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
}, },
"T1w": { "confounds": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/{session}/"
"func/{subject}_{session}_task-{task}_acq-mb4"
"{phase_encoding}_run-{run}_desc-confounds_regressors.tsv"
),
"format": "fmriprep",
},
},
"T1w": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_desc-preproc_T1w.nii.gz"
),
"space": "MNI152NLin2009cAsym",
"mask": {
"pattern": ( "pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/" "derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_desc-preproc_T1w.nii.gz" "{subject}_space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz"
),
"space": "MNI152NLin2009cAsym",
"mask": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
},
"VBM_CSF": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_label-CSF_probseg.nii.gz"
), ),
"space": "MNI152NLin2009cAsym", "space": "MNI152NLin2009cAsym",
}, },
"VBM_GM": { },
"pattern": ( "VBM_CSF": {
"derivatives/fmriprep-1.3.2/{subject}/anat/" "pattern": (
"{subject}_space-MNI152NLin2009cAsym_label-GM_probseg.nii.gz" "derivatives/fmriprep-1.3.2/{subject}/anat/"
), "{subject}_space-MNI152NLin2009cAsym_label-CSF_probseg.nii.gz"
"space": "MNI152NLin2009cAsym", ),
}, "space": "MNI152NLin2009cAsym",
"VBM_WM": { },
"pattern": ( "VBM_GM": {
"derivatives/fmriprep-1.3.2/{subject}/anat/" "pattern": (
"{subject}_space-MNI152NLin2009cAsym_label-WM_probseg.nii.gz" "derivatives/fmriprep-1.3.2/{subject}/anat/"
), "{subject}_space-MNI152NLin2009cAsym_label-GM_probseg.nii.gz"
"space": "MNI152NLin2009cAsym", ),
}, "space": "MNI152NLin2009cAsym",
} },
# Use native T1w assets "VBM_WM": {
self.native_t1w = False "pattern": (
if native_t1w: "derivatives/fmriprep-1.3.2/{subject}/anat/"
self.native_t1w = True "{subject}_space-MNI152NLin2009cAsym_label-WM_probseg.nii.gz"
patterns.update( ),
"space": "MNI152NLin2009cAsym",
},
}
replacements: list[str] = [ # noqa: RUF012
"subject",
"session",
"task",
"phase_encoding",
"run",
]
confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
if self.native_t1w:
self.patterns.update(
{ {
"T1w": { "T1w": {
"pattern": ( "pattern": (
@ -244,24 +273,8 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
], ],
} }
) )
# Set default types self.types.append(DataType.Warp)
if types is None: super().validate_datagrabber_params()
types = list(patterns.keys())
# Convert single type into list
else:
if not isinstance(types, list):
types = [types]
# The replacements
replacements = ["subject", "session", "task", "phase_encoding", "run"]
uri = "https://github.com/OpenNeuroDatasets/ds003452.git"
super().__init__(
types=types,
datadir=datadir,
uri=uri,
patterns=patterns,
replacements=replacements,
confounds_format="fmriprep",
)
def get_item( def get_item(
self, self,
@ -271,7 +284,7 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
phase_encoding: str, phase_encoding: str,
run: str, run: str,
) -> dict: ) -> dict:
"""Index one element in the dataset. """Get the specified item from the dataset.
Parameters Parameters
---------- ----------

View file

@ -1,4 +1,9 @@
__all__ = ["HCP1200", "DataladHCP1200"] __all__ = [
"HCP1200",
"HCP1200Task",
"HCP1200PhaseEncoding",
"DataladHCP1200",
]
from .hcp1200 import HCP1200 from .hcp1200 import HCP1200, HCP1200Task, HCP1200PhaseEncoding
from .datalad_hcp1200 import DataladHCP1200 from .datalad_hcp1200 import DataladHCP1200

View file

@ -6,66 +6,60 @@
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
from typing import Annotated, Literal
from junifer.datagrabber.datalad_base import DataladDataGrabber from pydantic import AnyUrl, BeforeValidator
from ...api.decorators import register_datagrabber from ...api.decorators import register_datagrabber
from ...utils import ensure_list
from ..base import DataType
from ..datalad_base import DataladDataGrabber
from .hcp1200 import HCP1200 from .hcp1200 import HCP1200
__all__ = ["DataladHCP1200"] __all__ = ["DataladHCP1200"]
_types = Literal[DataType.BOLD, DataType.T1w, DataType.Warp]
fraimondo commented 2025-11-26 11:12:19 +00:00 (Migrated from github.com)

will this expand? Otherwise I prefer the previous docstring as it tells you exactly what your options are.

will this expand? Otherwise I prefer the previous docstring as it tells you exactly what your options are.
synchon commented 2025-11-27 07:22:30 +00:00 (Migrated from github.com)
Here's the doc link: https://juaml.github.io/junifer/pr-preview/pr-364/api/datagrabbers.html#junifer.datagrabber.DataladHCP1200 and here's the doc link for the enum: https://juaml.github.io/junifer/pr-preview/pr-364/api/datagrabbers.html#junifer.datagrabber.HCP1200Task
fraimondo commented 2025-12-01 08:11:30 +00:00 (Migrated from github.com)

Ok, so now I found that we have an even bigger problem with the DOC.

This is the "user documentation" where we show what we have available: https://juaml.github.io/junifer/pr-preview/pr-364/builtin.html

When you click on DataladHCP1200, it goes to the API doc, which is then this: https://juaml.github.io/junifer/pr-preview/pr-364/api/datagrabbers.html#junifer.datagrabber.DataladHCP1200

This is terrible for users. types is not a string anymore, neither tasks nor phase_encodings which makes it cryptic. "What shall I put in the field?" The previous docstring was easier to understand E.g. "use the string "REST1" and it will work!"

Ok, so now I found that we have an even bigger problem with the DOC. This is the "user documentation" where we show what we have available: https://juaml.github.io/junifer/pr-preview/pr-364/builtin.html When you click on DataladHCP1200, it goes to the API doc, which is then this: https://juaml.github.io/junifer/pr-preview/pr-364/api/datagrabbers.html#junifer.datagrabber.DataladHCP1200 This is terrible for users. `types` is not a string anymore, neither `tasks` nor `phase_encodings` which makes it cryptic. "What shall I put in the field?" The previous docstring was easier to understand E.g. "use the string "REST1" and it will work!"
synchon commented 2025-12-02 09:15:23 +00:00 (Migrated from github.com)

Should be addressed with the latest commit.

Should be addressed with the latest commit.
@register_datagrabber @register_datagrabber
class DataladHCP1200(DataladDataGrabber, HCP1200): class DataladHCP1200(DataladDataGrabber, HCP1200):
"""Concrete implementation for datalad-based data fetching of HCP1200. """Concrete implementation for datalad-based data fetching of HCP1200.
Parameters Parameters
---------- ----------
datadir : str or Path or None, optional types : {"BOLD", "T1w", "Warp"} or list of the options, optional
The directory where the datalad dataset will be cloned. If None, The data type(s) to grab.
the datalad dataset will be cloned into a temporary directory datadir : pathlib.Path, optional
(default None). That path where the datalad dataset will be cloned.
If not specified, the datalad dataset will be cloned into a temporary
directory.
tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", \ tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", \
"LANGUAGE", "GAMBLING", "MOTOR"} or list of the options or None \ "LANGUAGE", "GAMBLING", "MOTOR"} or list of the options, optional
, optional HCP task sessions.
HCP task sessions. If None, all available task sessions are selected By default, all available task sessions are selected.
(default None). phase_encodings : {"LR", "RL"} or list of the options, optional
phase_encodings : {"LR", "RL"} or list of the options or None, optional HCP phase encoding directions.
HCP phase encoding directions. If None, both will be used By default, all are used.
(default None).
ica_fix : bool, optional ica_fix : bool, optional
Whether to retrieve data that was processed with ICA+FIX. Whether to retrieve data that was processed with ICA+FIX.
Only "REST1" and "REST2" tasks are available with ICA+FIX (default Only ``HCP1200Task.REST1`` and ``HCP1200Task.REST2`` tasks
False). are available with ICA+FIX
(default False).
Raises
------
ValueError
If invalid value is passed for ``tasks`` or ``phase_encodings``.
""" """
def __init__( uri: AnyUrl = AnyUrl(
self, "https://github.com/datalad-datasets/"
datadir: str | Path | None = None, "human-connectome-project-openaccess.git"
tasks: str | list[str] | None = None, )
phase_encodings: str | list[str] | None = None, types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
ica_fix: bool = False, DataType.BOLD,
) -> None: DataType.T1w,
uri = ( DataType.Warp,
"https://github.com/datalad-datasets/" ]
"human-connectome-project-openaccess.git" rootdir: Path = Path("HCP1200")
)
rootdir = "HCP1200"
super().__init__(
datadir=datadir,
tasks=tasks,
phase_encodings=phase_encodings,
uri=uri,
rootdir=rootdir,
ica_fix=ica_fix,
)
# Needed here as HCP1200's subjects are sub-datasets, so will not be # Needed here as HCP1200's subjects are sub-datasets, so will not be
# found when elements are checked. # found when elements are checked.

View file

@ -5,15 +5,61 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from enum import Enum
from itertools import product from itertools import product
from pathlib import Path from typing import Annotated, Literal
from pydantic import BeforeValidator
from ...api.decorators import register_datagrabber from ...api.decorators import register_datagrabber
from ...utils import raise_error from ...typing import DataGrabberPatterns
from ...utils import ensure_list, raise_error
from ..base import DataType
from ..pattern import PatternDataGrabber from ..pattern import PatternDataGrabber
__all__ = ["HCP1200"] __all__ = ["HCP1200", "HCP1200PhaseEncoding", "HCP1200Task"]
class HCP1200Task(str, Enum):
"""Accepted HCP1200 tasks."""
REST1 = "REST1"
REST2 = "REST2"
SOCIAL = "SOCIAL"
WM = "WM"
RELATIONAL = "RELATIONAL"
EMOTION = "EMOTION"
LANGUAGE = "LANGUAGE"
GAMBLING = "GAMBLING"
MOTOR = "MOTOR"
class HCP1200PhaseEncoding(str, Enum):
"""Accepted HCP1200 phase encoding directions."""
LR = "LR"
RL = "RL"
_types = Literal[DataType.BOLD, DataType.T1w, DataType.Warp]
_tasks = Literal[
HCP1200Task.REST1,
HCP1200Task.REST2,
HCP1200Task.SOCIAL,
HCP1200Task.WM,
HCP1200Task.RELATIONAL,
HCP1200Task.EMOTION,
HCP1200Task.LANGUAGE,
HCP1200Task.GAMBLING,
HCP1200Task.MOTOR,
]
_phase_encodings = Literal[
HCP1200PhaseEncoding.RL,
HCP1200PhaseEncoding.LR,
]
@register_datagrabber @register_datagrabber
@ -22,135 +68,99 @@ class HCP1200(PatternDataGrabber):
Parameters Parameters
---------- ----------
datadir : str or Path, optional types : {"BOLD", "T1w", "Warp"} or list of the options, optional
The directory where the data is / will be stored. The data type(s) to grab.
datadir : pathlib.Path
The path where the data is stored.
tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", \ tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", \
"LANGUAGE", "GAMBLING", "MOTOR"} or list of the options or None \ "LANGUAGE", "GAMBLING", "MOTOR"} or list of the options, optional
, optional HCP task sessions.
HCP task sessions. If None, all available task sessions are selected By default, all available task sessions are selected.
(default None). phase_encodings : {"LR", "RL"} or list of the options, optional
phase_encodings : {"LR", "RL"} or list of the options or None, optional HCP phase encoding directions.
HCP phase encoding directions. If None, both will be used By default, all are used.
(default None).
ica_fix : bool, optional ica_fix : bool, optional
Whether to retrieve data that was processed with ICA+FIX. Whether to retrieve data that was processed with ICA+FIX.
Only "REST1" and "REST2" tasks are available with ICA+FIX (default Only ``HCP1200Task.REST1`` and ``HCP1200Task.REST2`` tasks
False). are available with ICA+FIX
(default False).
Raises
------
ValueError
If invalid value is passed for ``tasks`` or ``phase_encodings``.
""" """
def __init__( types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
self, DataType.BOLD,
datadir: str | Path, DataType.T1w,
tasks: str | list[str] | None = None, DataType.Warp,
phase_encodings: str | list[str] | None = None, ]
ica_fix: bool = False, tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
) -> None: HCP1200Task.REST1,
# All tasks HCP1200Task.REST2,
all_tasks = [ HCP1200Task.SOCIAL,
"REST1", HCP1200Task.WM,
"REST2", HCP1200Task.RELATIONAL,
"SOCIAL", HCP1200Task.EMOTION,
"WM", HCP1200Task.LANGUAGE,
"RELATIONAL", HCP1200Task.GAMBLING,
"EMOTION", HCP1200Task.MOTOR,
"LANGUAGE", ]
"GAMBLING", phase_encodings: Annotated[
"MOTOR", _phase_encodings | list[_phase_encodings],
] BeforeValidator(ensure_list),
# Set default tasks ] = [ # noqa: RUF012
if tasks is None: HCP1200PhaseEncoding.RL,
self.tasks: list[str] = all_tasks HCP1200PhaseEncoding.LR,
# Convert single task into list ]
else: ica_fix: bool = False
if not isinstance(tasks, list): patterns: DataGrabberPatterns = { # noqa: RUF012
tasks = [tasks] "BOLD": {
# Check for invalid task(s) "pattern": (
for task in tasks: "{subject}/MNINonLinear/Results/"
if task not in all_tasks: "{task}_{phase_encoding}/"
raise_error( "{task}_{phase_encoding}"
f"'{task}' is not a valid HCP-YA fMRI task input. " "{suffix}.nii.gz"
f"Valid task values can be any or all of {all_tasks}." ),
) "space": "MNI152NLin6Asym",
self.tasks: list[str] = tasks },
"T1w": {
"pattern": "{subject}/T1w/T1w_acpc_dc_restore.nii.gz",
"space": "native",
},
"Warp": [
{
"pattern": (
"{subject}/MNINonLinear/xfms/standard2acpc_dc.nii.gz"
),
"src": "MNI152NLin6Asym",
"dst": "native",
"warper": "fsl",
},
{
"pattern": (
"{subject}/MNINonLinear/xfms/acpc_dc2standard.nii.gz"
),
"src": "native",
"dst": "MNI152NLin6Asym",
"warper": "fsl",
},
],
}
replacements: list[str] = ["subject", "task", "phase_encoding"] # noqa: RUF012
# All phase encodings def validate_datagrabber_params(self) -> None:
all_phase_encodings = ["LR", "RL"] """Run extra logical validation for datagrabber."""
# Set phase encodings if self.ica_fix:
if phase_encodings is None:
phase_encodings = all_phase_encodings
# Convert single phase encoding into list
if isinstance(phase_encodings, str):
phase_encodings = [phase_encodings]
# Check for invalid phase encoding(s)
for pe in phase_encodings:
if pe not in all_phase_encodings:
raise_error(
f"'{pe}' is not a valid HCP-YA phase encoding. "
"Valid phase encoding can be any or all of "
f"{all_phase_encodings}."
)
self.phase_encodings = phase_encodings
if ica_fix:
if not all(task in ["REST1", "REST2"] for task in self.tasks): if not all(task in ["REST1", "REST2"] for task in self.tasks):
raise_error( raise_error(
"ICA+FIX is only available for 'REST1' and 'REST2' tasks." "ICA+FIX is only available for 'REST1' and 'REST2' tasks."
) )
suffix = "_hp2000_clean" if ica_fix else "" suffix = "_hp2000_clean" if self.ica_fix else ""
self.patterns["BOLD"]["pattern"] = self.patterns["BOLD"][
# The types of data "pattern"
types = ["BOLD", "T1w", "Warp"] ].replace("{suffix}", suffix)
# The patterns super().validate_datagrabber_params()
patterns = {
"BOLD": {
"pattern": (
"{subject}/MNINonLinear/Results/"
"{task}_{phase_encoding}/"
"{task}_{phase_encoding}"
f"{suffix}.nii.gz"
),
"space": "MNI152NLin6Asym",
},
"T1w": {
"pattern": "{subject}/T1w/T1w_acpc_dc_restore.nii.gz",
"space": "native",
},
"Warp": [
{
"pattern": (
"{subject}/MNINonLinear/xfms/standard2acpc_dc.nii.gz"
),
"src": "MNI152NLin6Asym",
"dst": "native",
"warper": "fsl",
},
{
"pattern": (
"{subject}/MNINonLinear/xfms/acpc_dc2standard.nii.gz"
),
"src": "native",
"dst": "MNI152NLin6Asym",
"warper": "fsl",
},
],
}
# The replacements
replacements = ["subject", "task", "phase_encoding"]
super().__init__(
types=types,
datadir=datadir,
patterns=patterns,
replacements=replacements,
)
def get_item(self, subject: str, task: str, phase_encoding: str) -> dict: def get_item(self, subject: str, task: str, phase_encoding: str) -> dict:
"""Implement single element indexing in the database. """Get the specified item from the dataset.
Parameters Parameters
---------- ----------

View file

@ -9,62 +9,61 @@ from collections.abc import Iterable
from pathlib import Path from pathlib import Path
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber import HCP1200, DataladHCP1200 from junifer.datagrabber import HCP1200, DataladHCP1200
from junifer.utils import configure_logging from junifer.utils import config, configure_logging
URI = "https://gin.g-node.org/juaml/datalad-example-hcp1200"
@pytest.fixture(scope="module") @pytest.fixture(scope="module")
def hcpdg() -> Iterable[DataladHCP1200]: def hcpdg() -> Iterable[DataladHCP1200]:
"""Return a HCP1200 DataGrabber.""" """Return a HCP1200 DataGrabber."""
tmpdir = Path(tempfile.gettempdir()) tmpdir = Path(tempfile.gettempdir())
dg = DataladHCP1200(datadir=tmpdir / "datadir") config.set(key="datagrabber.skipidcheck", val=True)
# Set URI to Gin dg = DataladHCP1200(
dg.uri = URI uri=AnyUrl("https://gin.g-node.org/juaml/datalad-example-hcp1200"),
# Set correct root directory datadir=tmpdir / "hcp1200_test",
dg._rootdir = "." rootdir=Path("."),
)
with dg: with dg:
for t_elem in dg.get_elements(): for t_elem in dg.get_elements():
dg[t_elem] dg[t_elem]
yield dg yield dg
shutil.rmtree(tmpdir / "datadir", ignore_errors=True) config.set(key="datagrabber.skipidcheck", val=False)
shutil.rmtree(tmpdir / "hcp1200_test", ignore_errors=True)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"tasks, phase_encodings, ica_fix, expected_path_name", "tasks, phase_encodings, ica_fix, expected_path_name",
[ [
(None, None, False, "rfMRI_REST1_LR.nii.gz"),
("REST1", "LR", False, "rfMRI_REST1_LR.nii.gz"), ("REST1", "LR", False, "rfMRI_REST1_LR.nii.gz"),
("REST1", "RL", False, "rfMRI_REST1_RL.nii.gz"), (["REST1"], ["RL"], False, "rfMRI_REST1_RL.nii.gz"),
("REST2", "LR", False, "rfMRI_REST2_LR.nii.gz"), ("REST2", "LR", False, "rfMRI_REST2_LR.nii.gz"),
("REST2", "RL", False, "rfMRI_REST2_RL.nii.gz"), (["REST2"], ["RL"], False, "rfMRI_REST2_RL.nii.gz"),
("SOCIAL", "LR", False, "tfMRI_SOCIAL_LR.nii.gz"), ("SOCIAL", "LR", False, "tfMRI_SOCIAL_LR.nii.gz"),
("SOCIAL", "RL", False, "tfMRI_SOCIAL_RL.nii.gz"), (["SOCIAL"], ["RL"], False, "tfMRI_SOCIAL_RL.nii.gz"),
("WM", "LR", False, "tfMRI_WM_LR.nii.gz"), ("WM", "LR", False, "tfMRI_WM_LR.nii.gz"),
("WM", "RL", False, "tfMRI_WM_RL.nii.gz"), (["WM"], ["RL"], False, "tfMRI_WM_RL.nii.gz"),
("RELATIONAL", "LR", False, "tfMRI_RELATIONAL_LR.nii.gz"), ("RELATIONAL", "LR", False, "tfMRI_RELATIONAL_LR.nii.gz"),
("RELATIONAL", "RL", False, "tfMRI_RELATIONAL_RL.nii.gz"), (["RELATIONAL"], ["RL"], False, "tfMRI_RELATIONAL_RL.nii.gz"),
("EMOTION", "LR", False, "tfMRI_EMOTION_LR.nii.gz"), ("EMOTION", "LR", False, "tfMRI_EMOTION_LR.nii.gz"),
("EMOTION", "RL", False, "tfMRI_EMOTION_RL.nii.gz"), (["EMOTION"], ["RL"], False, "tfMRI_EMOTION_RL.nii.gz"),
("LANGUAGE", "LR", False, "tfMRI_LANGUAGE_LR.nii.gz"), ("LANGUAGE", "LR", False, "tfMRI_LANGUAGE_LR.nii.gz"),
("LANGUAGE", "RL", False, "tfMRI_LANGUAGE_RL.nii.gz"), (["LANGUAGE"], ["RL"], False, "tfMRI_LANGUAGE_RL.nii.gz"),
("GAMBLING", "LR", False, "tfMRI_GAMBLING_LR.nii.gz"), ("GAMBLING", "LR", False, "tfMRI_GAMBLING_LR.nii.gz"),
("GAMBLING", "RL", False, "tfMRI_GAMBLING_RL.nii.gz"), (["GAMBLING"], ["RL"], False, "tfMRI_GAMBLING_RL.nii.gz"),
("MOTOR", "LR", False, "tfMRI_MOTOR_LR.nii.gz"), ("MOTOR", "LR", False, "tfMRI_MOTOR_LR.nii.gz"),
("MOTOR", "RL", False, "tfMRI_MOTOR_RL.nii.gz"), (["MOTOR"], ["RL"], False, "tfMRI_MOTOR_RL.nii.gz"),
("REST1", "LR", True, "rfMRI_REST1_LR_hp2000_clean.nii.gz"), ("REST1", "LR", True, "rfMRI_REST1_LR_hp2000_clean.nii.gz"),
("REST1", "RL", True, "rfMRI_REST1_RL_hp2000_clean.nii.gz"), (["REST1"], ["RL"], True, "rfMRI_REST1_RL_hp2000_clean.nii.gz"),
("REST2", "LR", True, "rfMRI_REST2_LR_hp2000_clean.nii.gz"), ("REST2", "LR", True, "rfMRI_REST2_LR_hp2000_clean.nii.gz"),
("REST2", "RL", True, "rfMRI_REST2_RL_hp2000_clean.nii.gz"), (["REST2"], ["RL"], True, "rfMRI_REST2_RL_hp2000_clean.nii.gz"),
], ],
) )
def test_HCP1200( def test_HCP1200(
hcpdg: DataladHCP1200, hcpdg: DataladHCP1200,
tasks: str | None, tasks: str | list[str],
phase_encodings: str | None, phase_encodings: str | list[str],
ica_fix: bool, ica_fix: bool,
expected_path_name: str, expected_path_name: str,
) -> None: ) -> None:
@ -75,9 +74,9 @@ def test_HCP1200(
hcpdg : DataladHCP1200 hcpdg : DataladHCP1200
The Datalad version of the DataGrabber with the first subject The Datalad version of the DataGrabber with the first subject
already cloned. already cloned.
tasks : str tasks : str or list of str
The parametrized tasks. The parametrized tasks.
phase_encodings : str phase_encodings : str or list of str
The parametrized phase encodings. The parametrized phase encodings.
ica_fix : bool ica_fix : bool
The parametrized ICA-FIX flag. The parametrized ICA-FIX flag.
@ -118,29 +117,29 @@ def test_HCP1200(
"tasks, phase_encodings", "tasks, phase_encodings",
[ [
("REST1", "LR"), ("REST1", "LR"),
("REST1", "RL"), (["REST1"], ["RL"]),
("REST2", "LR"), ("REST2", "LR"),
("REST2", "RL"), (["REST2"], ["RL"]),
("SOCIAL", "LR"), ("SOCIAL", "LR"),
("SOCIAL", "RL"), (["SOCIAL"], ["RL"]),
("WM", "LR"), ("WM", "LR"),
("WM", "RL"), (["WM"], ["RL"]),
("RELATIONAL", "LR"), ("RELATIONAL", "LR"),
("RELATIONAL", "RL"), (["RELATIONAL"], ["RL"]),
("EMOTION", "LR"), ("EMOTION", "LR"),
("EMOTION", "RL"), (["EMOTION"], ["RL"]),
("LANGUAGE", "LR"), ("LANGUAGE", "LR"),
("LANGUAGE", "RL"), (["LANGUAGE"], ["RL"]),
("GAMBLING", "LR"), ("GAMBLING", "LR"),
("GAMBLING", "RL"), (["GAMBLING"], ["RL"]),
("MOTOR", "LR"), ("MOTOR", "LR"),
("MOTOR", "RL"), (["MOTOR"], ["RL"]),
], ],
) )
def test_HCP1200_single_access( def test_HCP1200_single_access(
hcpdg: DataladHCP1200, hcpdg: DataladHCP1200,
tasks: str | None, tasks: str | list[str],
phase_encodings: str | None, phase_encodings: str | list[str],
) -> None: ) -> None:
"""Test HCP1200 DataGrabber single access. """Test HCP1200 DataGrabber single access.
@ -149,9 +148,9 @@ def test_HCP1200_single_access(
hcpdg : DataladHCP1200 hcpdg : DataladHCP1200
The Datalad version of the DataGrabber with the first subject The Datalad version of the DataGrabber with the first subject
already cloned. already cloned.
tasks : str tasks : str or list of str
The parametrized tasks. The parametrized tasks.
phase_encodings : str phase_encodings : str or list of str
The parametrized phase encodings. The parametrized phase encodings.
""" """
@ -166,21 +165,16 @@ def test_HCP1200_single_access(
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Check only specified task and phase encoding are found # Check only specified task and phase encoding are found
for element in all_elements: for element in all_elements:
assert element[1] == tasks assert element[1] == tasks if isinstance(tasks, str) else tasks[0]
assert element[2] == phase_encodings assert (
element[2] == phase_encodings
if isinstance(phase_encodings, str)
else phase_encodings[0]
)
@pytest.mark.parametrize(
"tasks, phase_encodings",
[
(["REST1", "REST2"], ["LR", "RL"]),
(["REST1", "REST2"], None),
],
)
def test_HCP1200_multi_access( def test_HCP1200_multi_access(
hcpdg: DataladHCP1200, hcpdg: DataladHCP1200,
tasks: str | None,
phase_encodings: str | None,
) -> None: ) -> None:
"""Test HCP1200 DataGrabber multiple access. """Test HCP1200 DataGrabber multiple access.
@ -189,17 +183,13 @@ def test_HCP1200_multi_access(
hcpdg : DataladHCP1200 hcpdg : DataladHCP1200
The Datalad version of the DataGrabber with the first subject The Datalad version of the DataGrabber with the first subject
already cloned. already cloned.
tasks : str
The parametrized tasks.
phase_encodings : str
The parametrized phase encodings.
""" """
configure_logging(level="DEBUG") configure_logging(level="DEBUG")
dg = HCP1200( dg = HCP1200(
datadir=hcpdg.datadir, datadir=hcpdg.datadir,
tasks=tasks, tasks=["REST1", "REST2"],
phase_encodings=phase_encodings, phase_encodings=["LR", "RL"],
) )
with dg: with dg:
# Get all elements # Get all elements
@ -264,70 +254,6 @@ def test_HCP1200_multi_access_phase_simple(
assert element[2] == "LR" assert element[2] == "LR"
@pytest.mark.parametrize(
"tasks, phase_encodings",
[
("FOO", ["LR", "RL"]),
("FOO", "RL"),
(["FOO", "BAR"], ["LR", "RL"]),
(["FOO", "BAR"], "LR"),
],
)
def test_HCP1200_incorrect_access_task(
tasks: str | None,
phase_encodings: str | None,
) -> None:
"""Test HCP1200 DataGrabber incorrect access for task.
Parameters
----------
tasks : str
The parametrized tasks.
phase_encodings : str
The parametrized phase encodings.
"""
configure_logging(level="DEBUG")
with pytest.raises(ValueError, match="not a valid HCP-YA fMRI task input"):
_ = HCP1200(
datadir=".",
tasks=tasks,
phase_encodings=phase_encodings,
)
@pytest.mark.parametrize(
"tasks, phase_encodings",
[
("REST1", ["FOO", "BAR"]),
("REST1", "FOO"),
(["REST1", "REST2"], ["FOO", "BAR"]),
(["REST1", "REST2"], "BAR"),
],
)
def test_HCP1200_incorrect_access_phase(
tasks: str | None,
phase_encodings: str | None,
) -> None:
"""Test HCP1200 DataGrabber incorrect access for phase.
Parameters
----------
tasks : str
The parametrized tasks.
phase_encodings : str
The parametrized phase encodings.
"""
configure_logging(level="DEBUG")
with pytest.raises(ValueError, match="not a valid HCP-YA phase encoding"):
_ = HCP1200(
datadir=".",
tasks=tasks,
phase_encodings=phase_encodings,
)
def test_HCP1200_elements( def test_HCP1200_elements(
hcpdg: DataladHCP1200, hcpdg: DataladHCP1200,
) -> None: ) -> None:
@ -363,23 +289,24 @@ def test_HCP1200_elements(
@pytest.mark.parametrize( @pytest.mark.parametrize(
"tasks, ica_fix", "tasks, ica_fix",
[ [
("SOCIAL", True), (["SOCIAL"], True),
("WM", True), ("WM", True),
("RELATIONAL", True), (["RELATIONAL"], True),
("EMOTION", True), ("EMOTION", True),
("LANGUAGE", True), (["LANGUAGE"], True),
("GAMBLING", True), ("GAMBLING", True),
("MOTOR", True), (["MOTOR"], True),
], ],
) )
def test_HCP1200_incorrect_access_icafix( def test_HCP1200_incorrect_access_icafix(
tasks: str | None, ica_fix: bool tasks: str | list[str],
ica_fix: bool,
) -> None: ) -> None:
"""Test HCP1200 DataGrabber incorrect access for icafix. """Test HCP1200 DataGrabber incorrect access for icafix.
Parameters Parameters
---------- ----------
tasks : str tasks : str or list of str
The parametrized tasks. The parametrized tasks.
ica_fix : bool ica_fix : bool
The parametrized ICA-FIX flag. The parametrized ICA-FIX flag.
@ -388,7 +315,7 @@ def test_HCP1200_incorrect_access_icafix(
configure_logging(level="DEBUG") configure_logging(level="DEBUG")
with pytest.raises(ValueError, match="is only available for"): with pytest.raises(ValueError, match="is only available for"):
_ = HCP1200( _ = HCP1200(
datadir=".", datadir=Path("."),
tasks=tasks, tasks=tasks,
ica_fix=ica_fix, ica_fix=ica_fix,
) )

View file

@ -5,10 +5,17 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path
from typing import Annotated
from pydantic import BeforeValidator, ConfigDict
from ..api.decorators import register_datagrabber from ..api.decorators import register_datagrabber
from ..typing import DataGrabberLike from ..typing import DataGrabberLike, Element
from ..utils import deep_update, raise_error from ..utils import deep_update, ensure_list, raise_error
from .base import BaseDataGrabber from .base import BaseDataGrabber, DataType
from .pattern import PatternDataGrabber
from .pattern_datalad import PatternDataladDataGrabber
__all__ = ["MultipleDataGrabber"] __all__ = ["MultipleDataGrabber"]
@ -36,22 +43,33 @@ class MultipleDataGrabber(BaseDataGrabber):
""" """
def __init__(self, datagrabbers: list[DataGrabberLike], **kwargs) -> None: model_config = ConfigDict(extra="allow")
datagrabbers: list[
DataGrabberLike | PatternDataGrabber | PatternDataladDataGrabber
]
types: Annotated[
DataType | list[DataType], BeforeValidator(ensure_list)
] = [] # noqa: RUF012
datadir: Path = Path(".")
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
# Check datagrabbers consistency # Check datagrabbers consistency
# Check for same element keys # Check for same element keys
first_keys = datagrabbers[0].get_element_keys() first_keys = self.datagrabbers[0].get_element_keys()
for dg in datagrabbers[1:]: for dg in self.datagrabbers[1:]:
if dg.get_element_keys() != first_keys: if dg.get_element_keys() != first_keys:
raise_error( raise_error(
msg="DataGrabbers have different element keys", msg="DataGrabbers have different element keys",
klass=RuntimeError, klass=RuntimeError,
) )
# Check for no overlapping types (and nested data types) # Check for no overlapping types (and nested data types)
types = [x for dg in datagrabbers for x in dg.get_types()] types = [x for dg in self.datagrabbers for x in dg.get_types()]
if len(types) != len(set(types)): if len(types) != len(set(types)):
if all(hasattr(dg, "patterns") for dg in datagrabbers): if all(hasattr(dg, "patterns") for dg in self.datagrabbers):
first_patterns = datagrabbers[0].patterns first_patterns = self.datagrabbers[0].patterns
for dg in datagrabbers[1:]: for dg in self.datagrabbers[1:]:
for data_type in set(types): for data_type in set(types):
dtype_pattern = dg.patterns.get(data_type) dtype_pattern = dg.patterns.get(data_type)
if dtype_pattern is None: if dtype_pattern is None:
@ -75,14 +93,13 @@ class MultipleDataGrabber(BaseDataGrabber):
msg="DataGrabbers have overlapping types", msg="DataGrabbers have overlapping types",
klass=RuntimeError, klass=RuntimeError,
) )
self._datagrabbers = datagrabbers
def __getitem__(self, element: str | tuple) -> dict: def __getitem__(self, element: Element) -> dict:
"""Implement indexing. """Implement indexing.
Parameters Parameters
---------- ----------
element : str or tuple element : `Element`
The element to be indexed. If one string is provided, it is The element to be indexed. If one string is provided, it is
assumed to be a tuple with only one item. If a tuple is provided, assumed to be a tuple with only one item. If a tuple is provided,
each item in the tuple is the value for the replacement string each item in the tuple is the value for the replacement string
@ -98,7 +115,7 @@ class MultipleDataGrabber(BaseDataGrabber):
out = {} out = {}
metas = [] metas = []
for dg in self._datagrabbers: for dg in self.datagrabbers:
t_out = dg[element] t_out = dg[element]
deep_update(out, t_out) deep_update(out, t_out)
# Now get the meta for this datagrabber # Now get the meta for this datagrabber
@ -119,16 +136,15 @@ class MultipleDataGrabber(BaseDataGrabber):
def __enter__(self) -> "MultipleDataGrabber": def __enter__(self) -> "MultipleDataGrabber":
"""Implement context entry.""" """Implement context entry."""
for dg in self._datagrabbers: for dg in self.datagrabbers:
dg.__enter__() dg.__enter__()
return self return self
def __exit__(self, exc_type, exc_value, exc_traceback) -> None: def __exit__(self, exc_type, exc_value, exc_traceback) -> None:
"""Implement context exit.""" """Implement context exit."""
for dg in self._datagrabbers: for dg in self.datagrabbers:
dg.__exit__(exc_type, exc_value, exc_traceback) dg.__exit__(exc_type, exc_value, exc_traceback)
# TODO: return type should be List[List[str]], but base type is List[str]
def get_types(self) -> list[str]: def get_types(self) -> list[str]:
"""Get types. """Get types.
@ -138,7 +154,7 @@ class MultipleDataGrabber(BaseDataGrabber):
The types of data to be grabbed. The types of data to be grabbed.
""" """
types = [x for dg in self._datagrabbers for x in dg.get_types()] types = [x for dg in self.datagrabbers for x in dg.get_types()]
return types return types
def get_element_keys(self) -> list[str]: def get_element_keys(self) -> list[str]:
@ -153,7 +169,7 @@ class MultipleDataGrabber(BaseDataGrabber):
The element keys. The element keys.
""" """
return self._datagrabbers[0].get_element_keys() return self.datagrabbers[0].get_element_keys()
def get_elements(self) -> list: def get_elements(self) -> list:
"""Get elements. """Get elements.
@ -167,7 +183,7 @@ class MultipleDataGrabber(BaseDataGrabber):
related DataGrabbers. related DataGrabbers.
""" """
all_elements = [dg.get_elements() for dg in self._datagrabbers] all_elements = [dg.get_elements() for dg in self.datagrabbers]
elements = set(all_elements[0]) elements = set(all_elements[0])
for s in all_elements[1:]: for s in all_elements[1:]:
elements.intersection_update(s) elements.intersection_update(s)

View file

@ -10,6 +10,9 @@ from copy import deepcopy
from pathlib import Path from pathlib import Path
import numpy as np import numpy as np
from aenum import Enum as AEnum
from aenum import extend_enum
from pydantic import Field
from ..api.decorators import register_datagrabber from ..api.decorators import register_datagrabber
from ..typing import DataGrabberPatterns, Elements from ..typing import DataGrabberPatterns, Elements
@ -18,11 +21,32 @@ from .base import BaseDataGrabber
from .pattern_validation_mixin import PatternValidationMixin from .pattern_validation_mixin import PatternValidationMixin
__all__ = ["PatternDataGrabber"] __all__ = [
"ConfoundsFormat",
"PatternDataGrabber",
"register_confounds_format",
]
# Accepted formats for confounds specification class ConfoundsFormat(str, AEnum):
_CONFOUNDS_FORMATS = ("fmriprep", "adhoc") """Accepted confounds format."""
FMRIPrep = "fmriprep"
AdHoc = "adhoc"
def register_confounds_format(name: str, alias: str) -> None:
"""Register custom confounds format.
Parameters
----------
name : str
The confounds format name to be referred.
alias : str
The confounds format alias for string representation.
"""
extend_enum(ConfoundsFormat, name, alias)
@register_datagrabber @register_datagrabber
@ -33,125 +57,16 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
Parameters Parameters
---------- ----------
types : list of str types : :enum:`.DataType` or list of variants
The types of data to be grabbed. The data type(s) to grab.
patterns : dict datadir : pathlib.Path
Data type patterns as a dictionary. It has the following schema: The path where the data is stored.
patterns : ``DataGrabberPatterns``
* ``"T1w"`` : The datagrabber patterns. Check :class:`.DataTypeSchema` for the \
schema.
.. code-block:: none replacements : list of str
All possible replacements in ``patterns.<data_type>.pattern``.
{ confounds_format : :enum:`.ConfoundsFormat` or None, optional
"mandatory": ["pattern", "space"],
"optional": {
"mask": {
"mandatory": ["pattern", "space"],
"optional": []
}
}
}
* ``"T2w"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": {
"mask": {
"mandatory": ["pattern", "space"],
"optional": []
}
}
}
* ``"BOLD"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": {
"mask": {
"mandatory": ["pattern", "space"],
"optional": []
}
"confounds": {
"mandatory": ["pattern", "format"],
"optional": []
}
}
}
* ``"Warp"`` :
.. code-block:: none
{
"mandatory": ["pattern", "src", "dst", "warper"],
"optional": []
}
* ``"VBM_GM"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": []
}
* ``"VBM_WM"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": []
}
Basically, for each data type, one needs to provide ``mandatory`` keys
and can choose to also provide ``optional`` keys. The value for each
key is a string. So, one needs to provide necessary data types as a
dictionary, for example:
.. code-block:: none
{
"BOLD": {
"pattern": "...",
"space": "...",
},
"T1w": {
"pattern": "...",
"space": "...",
},
}
except ``Warp``, which needs to be a list of dictionaries as there can
be multiple spaces to warp (for example, with fMRIPrep):
.. code-block:: none
{
"Warp": [
{
"pattern": "...",
"src": "...",
"dst": "...",
"warper": "...",
},
],
}
taken from :class:`.HCP1200`.
replacements : str or list of str
Replacements in the ``pattern`` key of each data type. The value needs
to be a list of all possible replacements.
datadir : str or pathlib.Path
The directory where the data is / will be stored.
confounds_format : {"fmriprep", "adhoc"} or None, optional
The format of the confounds for the dataset (default None). The format of the confounds for the dataset (default None).
partial_pattern_ok : bool, optional partial_pattern_ok : bool, optional
Whether to raise error if partial pattern for a data type is found. Whether to raise error if partial pattern for a data type is found.
@ -161,52 +76,30 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
powerful when used with :class:`.MultipleDataGrabber` powerful when used with :class:`.MultipleDataGrabber`
(default True). (default True).
Raises Attributes
------ ----------
ValueError skip_file_check
If ``confounds_format`` is invalid.
""" """
def __init__( patterns: DataGrabberPatterns = Field(frozen=True)
self, replacements: list[str] = Field(frozen=True)
types: list[str], confounds_format: ConfoundsFormat | None = Field(None, frozen=True)
patterns: DataGrabberPatterns, partial_pattern_ok: bool = Field(False, frozen=True)
replacements: list[str] | str,
datadir: str | Path, def validate_datagrabber_params(self) -> None:
confounds_format: str | None = None, """Run extra logical validation for datagrabber."""
partial_pattern_ok: bool = False,
) -> None:
# Convert replacements to list if not already
if not isinstance(replacements, list):
replacements = [replacements]
# Validate patterns # Validate patterns
self.validate_patterns( self.validate_patterns(
types=types, types=self.types,
replacements=replacements, replacements=self.replacements,
patterns=patterns, patterns=self.patterns,
partial_pattern_ok=partial_pattern_ok, partial_pattern_ok=self.partial_pattern_ok,
) )
self.replacements = replacements
self.patterns = patterns
self.partial_pattern_ok = partial_pattern_ok
# Validate confounds format
if (
confounds_format is not None
and confounds_format not in _CONFOUNDS_FORMATS
):
raise_error(
"Invalid value for `confounds_format`, should be one of "
f"{_CONFOUNDS_FORMATS}."
)
self.confounds_format = confounds_format
super().__init__(types=types, datadir=datadir)
logger.debug("Initializing PatternDataGrabber") logger.debug("Initializing PatternDataGrabber")
logger.debug(f"\tpatterns = {patterns}") logger.debug(f"\tpatterns = {self.patterns}")
logger.debug(f"\treplacements = {replacements}") logger.debug(f"\treplacements = {self.replacements}")
logger.debug(f"\tconfounds_format = {confounds_format}") logger.debug(f"\tconfounds_format = {self.confounds_format}")
@property @property
def skip_file_check(self) -> bool: def skip_file_check(self) -> bool:
@ -324,7 +217,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
resolved_pattern = self._replace_patterns_glob(element, pattern) resolved_pattern = self._replace_patterns_glob(element, pattern)
# Resolve path for wildcard # Resolve path for wildcard
if "*" in resolved_pattern: if "*" in resolved_pattern:
t_matches = list(self.datadir.absolute().glob(resolved_pattern)) t_matches = list(self.fulldir.absolute().glob(resolved_pattern))
# Multiple matches # Multiple matches
if len(t_matches) > 1: if len(t_matches) > 1:
raise_error( raise_error(
@ -340,7 +233,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
) )
path = t_matches[0] path = t_matches[0]
else: else:
path = self.datadir / resolved_pattern path = self.fulldir / resolved_pattern
if not self.skip_file_check: if not self.skip_file_check:
if not path.exists() and not path.is_symlink(): if not path.exists() and not path.is_symlink():
raise_error( raise_error(
@ -367,7 +260,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
return self.replacements return self.replacements
def get_item(self, **element: dict) -> dict[str, dict]: def get_item(self, **element: dict) -> dict[str, dict]:
"""Implement single element indexing for the datagrabber. """Get the specified item from the dataset.
This method constructs a real path to the requested item's data, by This method constructs a real path to the requested item's data, by
replacing the ``patterns`` with actual values passed via ``**element``. replacing the ``patterns`` with actual values passed via ``**element``.
@ -514,8 +407,8 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
glob_pattern, glob_pattern,
t_replacements, t_replacements,
) = self._replace_patterns_regex(pattern) ) = self._replace_patterns_regex(pattern)
for fname in self.datadir.glob(glob_pattern): for fname in self.fulldir.glob(glob_pattern):
suffix = fname.relative_to(self.datadir).as_posix() suffix = fname.relative_to(self.fulldir).as_posix()
m = re.match(re_pattern, suffix) m = re.match(re_pattern, suffix)
if m is not None: if m is not None:
# Find the groups of replacements present in the # Find the groups of replacements present in the

View file

@ -5,6 +5,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pydantic import ConfigDict
from ..api.decorators import register_datagrabber from ..api.decorators import register_datagrabber
from ..utils import logger from ..utils import logger
from .datalad_base import DataladDataGrabber from .datalad_base import DataladDataGrabber
@ -23,122 +25,24 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber):
Parameters Parameters
---------- ----------
types : list of str uri : pydantic.AnyUrl
The types of data to be grabbed. URI of the datalad sibling.
patterns : dict types : enum:`.DataType` or list of variants
Data type patterns as a dictionary. It has the following schema: The data type(s) to grab.
patterns : ``DataGrabberPatterns``
* ``"T1w"`` : The datagrabber patterns. Check :class:`DataTypeSchema` for the schema.
replacements : list of str
.. code-block:: none All possible replacements in ``patterns.<data_type>.pattern``.
rootdir : pathlib.Path, optional
{
"mandatory": ["pattern", "space"],
"optional": {
"mask": {
"mandatory": ["pattern", "space"],
"optional": []
}
}
}
* ``"T2w"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": {
"mask": {
"mandatory": ["pattern", "space"],
"optional": []
}
}
}
* ``"BOLD"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": {
"mask": {
"mandatory": ["pattern", "space"],
"optional": []
}
"confounds": {
"mandatory": ["pattern", "format"],
"optional": []
}
}
}
* ``"Warp"`` :
.. code-block:: none
{
"mandatory": ["pattern", "src", "dst"],
"optional": []
}
* ``"VBM_GM"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": []
}
* ``"VBM_WM"`` :
.. code-block:: none
{
"mandatory": ["pattern", "space"],
"optional": []
}
Basically, for each data type, one needs to provide ``mandatory`` keys
and can choose to also provide ``optional`` keys. The value for each
key is a string. So, one needs to provide necessary data types as a
dictionary, for example:
.. code-block:: none
{
"BOLD": {
"pattern": "...",
"space": "...",
},
"T1w": {
"pattern": "...",
"space": "...",
},
"Warp": {
"pattern": "...",
"src": "...",
"dst": "...",
}
}
taken from :class:`.HCP1200`.
replacements : str or list of str
Replacements in the ``pattern`` key of each data type. The value needs
to be a list of all possible replacements.
confounds_format : {"fmriprep", "adhoc"} or None, optional
The format of the confounds for the dataset (default None).
datadir : str or pathlib.Path or None, optional
That directory where the datalad dataset will be cloned. If None,
the datalad dataset will be cloned into a temporary directory
(default None).
rootdir : str or pathlib.Path, optional
The path within the datalad dataset to the root directory The path within the datalad dataset to the root directory
(default "."). (default Path(".")).
uri : str or None, optional confounds_format : :enum:`.ConfoundsFormat` or None, optional
URI of the datalad sibling (default None). The format of the confounds for the dataset (default None).
datadir : pathlib.Path, optional
That path where the datalad dataset will be cloned.
If not specified, the datalad dataset will be cloned into a temporary
directory.
See Also See Also
-------- --------
@ -149,15 +53,11 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber):
""" """
def __init__( model_config = ConfigDict(extra="allow")
self,
**kwargs,
) -> None:
# TODO(synchon): needs to be reworked, DataladDataGrabber needs to be
# a mixin to avoid multiple inheritance wherever possible.
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
super().validate_datagrabber_params()
logger.debug("Initializing PatternDataladDataGrabber") logger.debug("Initializing PatternDataladDataGrabber")
for key, val in kwargs.items(): for key, val in self.__pydantic_extra__.items():
logger.debug(f"\t{key} = {val}") logger.debug(f"\t{key} = {val}")
super().__init__(**kwargs)

View file

@ -3,9 +3,20 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Iterator, MutableMapping import sys
from typing import TypedDict
if sys.version_info < (3, 12): # pragma: no cover
from typing_extensions import TypedDict
else:
from typing import TypedDict
from collections.abc import Iterator, MutableMapping
from aenum import extend_enum
from ..datagrabber import DataType
from ..typing import DataGrabberPatterns from ..typing import DataGrabberPatterns
from ..utils import logger, raise_error, warn_with_log from ..utils import logger, raise_error, warn_with_log
@ -191,37 +202,17 @@ def register_data_type(name: str, schema: DataTypeSchema) -> None:
---------- ----------
name : str name : str
The data type name. The data type name.
schema : DataTypeSchema schema : ``DataTypeSchema``
The data type schema. The data type schema.
""" """
DataTypeManager()[name] = schema DataTypeManager()[name] = schema
extend_enum(DataType, name, name)
class PatternValidationMixin: class PatternValidationMixin:
"""Mixin class for pattern validation.""" """Mixin class for pattern validation."""
def _validate_types(self, types: list[str]) -> None:
"""Validate the types.
Parameters
----------
types : list of str
The data types to validate.
Raises
------
TypeError
If ``types`` is not a list or if the values are not string.
"""
if not isinstance(types, list):
raise_error(msg="`types` must be a list", klass=TypeError)
if any(not isinstance(x, str) for x in types):
raise_error(
msg="`types` must be a list of strings", klass=TypeError
)
def _validate_replacements( def _validate_replacements(
self, self,
replacements: list[str], replacements: list[str],
@ -234,15 +225,13 @@ class PatternValidationMixin:
---------- ----------
replacements : list of str replacements : list of str
The replacements to validate. The replacements to validate.
patterns : dict patterns : ``DataGrabberPatterns``
The patterns to validate replacements against. The patterns to validate ``replacements`` against.
partial_pattern_ok : bool partial_pattern_ok : bool
Whether to raise error if partial pattern for a data type is found. Whether to raise error if partial pattern for a data type is found.
Raises Raises
------ ------
TypeError
If ``replacements`` is not a list or if the values are not string.
ValueError ValueError
If a value in ``replacements`` is not part of a data type pattern If a value in ``replacements`` is not part of a data type pattern
and ``partial_pattern_ok=False`` or and ``partial_pattern_ok=False`` or
@ -256,15 +245,6 @@ class PatternValidationMixin:
and ``partial_pattern_ok=True``. and ``partial_pattern_ok=True``.
""" """
if not isinstance(replacements, list):
raise_error(msg="`replacements` must be a list.", klass=TypeError)
if any(not isinstance(x, str) for x in replacements):
raise_error(
msg="`replacements` must be a list of strings",
klass=TypeError,
)
# Make a list of all patterns recursively # Make a list of all patterns recursively
all_patterns = [] all_patterns = []
for dtype_val in patterns.values(): for dtype_val in patterns.values():
@ -390,7 +370,7 @@ class PatternValidationMixin:
def validate_patterns( def validate_patterns(
self, self,
types: list[str], types: list[DataType],
replacements: list[str], replacements: list[str],
patterns: DataGrabberPatterns, patterns: DataGrabberPatterns,
partial_pattern_ok: bool = False, partial_pattern_ok: bool = False,
@ -399,11 +379,11 @@ class PatternValidationMixin:
Parameters Parameters
---------- ----------
types : list of str types : list of :enum:`.DataType`
The data types to check patterns of. The data type(s) to check patterns of.
replacements : list of str replacements : list of str
The replacements to be replaced in the patterns. The replacements to be replaced in the ``patterns``.
patterns : dict patterns : ``DataGrabberPatterns``
The patterns to validate. The patterns to validate.
partial_pattern_ok : bool, optional partial_pattern_ok : bool, optional
Whether to raise error if partial pattern for a data type is found. Whether to raise error if partial pattern for a data type is found.
@ -412,8 +392,6 @@ class PatternValidationMixin:
Raises Raises
------ ------
TypeError
If ``patterns`` is not a dictionary.
ValueError ValueError
If length of ``types`` and ``patterns`` are different or If length of ``types`` and ``patterns`` are different or
if ``patterns`` is missing entries from ``types`` or if ``patterns`` is missing entries from ``types`` or
@ -421,12 +399,6 @@ class PatternValidationMixin:
if data type pattern key contains '*' as value. if data type pattern key contains '*' as value.
""" """
# Validate types
self._validate_types(types=types)
# Validate patterns
if not isinstance(patterns, dict):
raise_error(msg="`patterns` must be a dict", klass=TypeError)
# Unequal length of objects # Unequal length of objects
if len(types) > len(patterns): if len(types) > len(patterns):
raise_error( raise_error(

View file

@ -26,7 +26,7 @@ def test_BaseDataGrabber() -> None:
def get_element_keys(self): def get_element_keys(self):
return ["subject"] return ["subject"]
dg = MyDataGrabber(datadir="/tmp", types=["BOLD"]) dg = MyDataGrabber(datadir="/tmp", types="BOLD")
elem = dg["sub01"] elem = dg["sub01"]
assert "BOLD" in elem assert "BOLD" in elem
assert "meta" in elem["BOLD"] assert "meta" in elem["BOLD"]
@ -55,7 +55,7 @@ def test_BaseDataGrabber() -> None:
def get_element_keys(self): def get_element_keys(self):
return super().get_element_keys() return super().get_element_keys()
dg = MyDataGrabber2(datadir="/tmp", types=["BOLD"]) dg = MyDataGrabber2(datadir="/tmp", types="BOLD")
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
dg.get_element_keys() dg.get_element_keys()
@ -77,7 +77,7 @@ def test_BaseDataGrabber_filter_single() -> None:
def get_element_keys(self): def get_element_keys(self):
return ["subject"] return ["subject"]
dg = FilterDataGrabber(datadir="/tmp", types=["BOLD"]) dg = FilterDataGrabber(datadir="/tmp", types="BOLD")
with dg: with dg:
assert "sub01" in list(dg.filter(["sub01"])) assert "sub01" in list(dg.filter(["sub01"]))
assert "sub02" not in list(dg.filter(["sub01"])) assert "sub02" not in list(dg.filter(["sub01"]))
@ -104,7 +104,7 @@ def test_BaseDataGrabber_filter_multi() -> None:
def get_element_keys(self): def get_element_keys(self):
return ["subject", "task"] return ["subject", "task"]
dg = FilterDataGrabber(datadir="/tmp", types=["BOLD"]) dg = FilterDataGrabber(datadir="/tmp", types="BOLD")
with dg: with dg:
assert ("sub01", "rest") in list( assert ("sub01", "rest") in list(
dg.filter([("sub01", "rest")]) # type: ignore dg.filter([("sub01", "rest")]) # type: ignore

View file

@ -9,7 +9,7 @@ from pathlib import Path
import datalad.api as dl import datalad.api as dl
import pytest import pytest
from junifer.datagrabber import DataladDataGrabber from junifer.datagrabber import DataladDataGrabber, DataType
from junifer.utils import config from junifer.utils import config
@ -38,23 +38,18 @@ def concrete_datagrabber() -> type[DataladDataGrabber]:
""" """
class MyDataGrabber(DataladDataGrabber): # type: ignore class MyDataGrabber(DataladDataGrabber):
def __init__(self, datadir, uri): types: list[DataType] = [DataType.T1w, DataType.BOLD] # noqa: RUF012
super().__init__( rootdir: Path = Path("example_bids")
datadir=datadir,
rootdir="example_bids",
uri=uri,
types=["T1w", "BOLD"],
)
def get_item(self, subject): def get_item(self, subject):
out = { out = {
"T1w": { "T1w": {
"path": self.datadir "path": self.fulldir
/ f"{subject}/anat/{subject}_T1w.nii.gz" / f"{subject}/anat/{subject}_T1w.nii.gz"
}, },
"BOLD": { "BOLD": {
"path": self.datadir "path": self.fulldir
/ f"{subject}/func/{subject}_task-rest_bold.nii.gz" / f"{subject}/func/{subject}_task-rest_bold.nii.gz"
}, },
} }
@ -91,7 +86,7 @@ def test_DataladDataGrabber_install_errors(
# Files are not there # Files are not there
assert datadir.exists() is False assert datadir.exists() is False
# Clone dataset # Clone dataset
dl.clone(uri, datadir) # type: ignore dl.clone(uri, datadir)
dg = concrete_datagrabber(datadir=datadir, uri=uri2) dg = concrete_datagrabber(datadir=datadir, uri=uri2)
with pytest.raises(ValueError, match=r"different ID"): with pytest.raises(ValueError, match=r"different ID"):
with dg: with dg:
@ -160,7 +155,6 @@ def test_DataladDataGrabber_clone_cleanup(
assert "datagrabber" in meta assert "datagrabber" in meta
assert "datalad_dirty" in meta["datagrabber"] assert "datalad_dirty" in meta["datagrabber"]
assert meta["datagrabber"]["datalad_dirty"] is False assert meta["datagrabber"]["datalad_dirty"] is False
assert hasattr(dg, "_got_files") is False
assert datadir.exists() is True assert datadir.exists() is True
assert elem1_bold.is_file() is True assert elem1_bold.is_file() is True
assert elem1_bold.is_symlink() is True assert elem1_bold.is_symlink() is True
@ -185,8 +179,8 @@ def test_DataladDataGrabber_clone_create_cleanup(
# Clone whole dataset # Clone whole dataset
uri = _testing_dataset["example_bids"]["uri"] uri = _testing_dataset["example_bids"]["uri"]
with concrete_datagrabber(datadir=None, uri=uri) as dg: with concrete_datagrabber(uri=uri) as dg:
datadir = dg._tmpdir / "datadir" datadir = dg._repodir
elem1_bold = ( elem1_bold = (
datadir / "example_bids/sub-01/func/sub-01_task-rest_bold.nii.gz" datadir / "example_bids/sub-01/func/sub-01_task-rest_bold.nii.gz"
) )
@ -206,7 +200,6 @@ def test_DataladDataGrabber_clone_create_cleanup(
assert "datagrabber" in meta assert "datagrabber" in meta
assert "datalad_dirty" in meta["datagrabber"] assert "datalad_dirty" in meta["datagrabber"]
assert meta["datagrabber"]["datalad_dirty"] is False assert meta["datagrabber"]["datalad_dirty"] is False
assert hasattr(dg, "_got_files") is False
assert datadir.exists() is True assert datadir.exists() is True
assert elem1_bold.is_file() is True assert elem1_bold.is_file() is True
assert elem1_bold.is_symlink() is True assert elem1_bold.is_symlink() is True
@ -246,7 +239,7 @@ def test_DataladDataGrabber_previously_cloned(
assert elem1_t1w.exists() is False assert elem1_t1w.exists() is False
# Clone dataset # Clone dataset
dl.clone(uri, datadir, result_renderer="disabled") # type: ignore dl.clone(uri, datadir, result_renderer="disabled")
# Files are there, but are empty symbolic links # Files are there, but are empty symbolic links
assert datadir.exists() is True assert datadir.exists() is True
@ -316,7 +309,7 @@ def test_DataladDataGrabber_previously_cloned_and_get(
assert elem1_t1w.exists() is False assert elem1_t1w.exists() is False
# Clone dataset # Clone dataset
dl.clone(uri, datadir, result_renderer="disabled") # type: ignore dl.clone(uri, datadir, result_renderer="disabled")
# Files are there, but are empty symbolic links # Files are there, but are empty symbolic links
assert datadir.exists() is True assert datadir.exists() is True
@ -325,9 +318,7 @@ def test_DataladDataGrabber_previously_cloned_and_get(
assert elem1_t1w.is_symlink() is True assert elem1_t1w.is_symlink() is True
assert elem1_t1w.is_file() is False assert elem1_t1w.is_file() is False
dl.get( # type: ignore dl.get(elem1_t1w, dataset=datadir, result_renderer="disabled")
elem1_t1w, dataset=datadir, result_renderer="disabled"
)
assert elem1_bold.is_symlink() is True assert elem1_bold.is_symlink() is True
assert elem1_bold.is_file() is False assert elem1_bold.is_file() is False
@ -399,7 +390,7 @@ def test_DataladDataGrabber_previously_cloned_and_get_dirty(
assert elem1_t1w.exists() is False assert elem1_t1w.exists() is False
# Clone dataset # Clone dataset
dl.clone(uri, datadir, result_renderer="disabled") # type: ignore dl.clone(uri, datadir, result_renderer="disabled")
# Files are there, but are empty symbolic links # Files are there, but are empty symbolic links
assert datadir.exists() is True assert datadir.exists() is True
@ -408,9 +399,7 @@ def test_DataladDataGrabber_previously_cloned_and_get_dirty(
assert elem1_t1w.is_symlink() is True assert elem1_t1w.is_symlink() is True
assert elem1_t1w.is_file() is False assert elem1_t1w.is_file() is False
dl.get( # type: ignore dl.get(elem1_t1w, dataset=datadir, result_renderer="disabled")
elem1_t1w, dataset=datadir, result_renderer="disabled"
)
assert elem1_bold.is_symlink() is True assert elem1_bold.is_symlink() is True
assert elem1_bold.is_file() is False assert elem1_bold.is_file() is False

View file

@ -4,101 +4,95 @@
# License: AGPL # License: AGPL
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber import DMCC13Benchmark from junifer.datagrabber import DataType, DMCC13Benchmark
URI = "https://gin.g-node.org/synchon/datalad-example-dmcc13-benchmark" URI = AnyUrl("https://gin.g-node.org/synchon/datalad-example-dmcc13-benchmark")
@pytest.mark.parametrize( @pytest.mark.parametrize(
"sessions, tasks, phase_encodings, runs, native_t1w", "sessions, tasks, phase_encodings, runs, native_t1w",
[ [
(None, None, None, None, False), (["ses-wave1bas"], ["Rest"], ["AP"], ["1"], False),
("ses-wave1bas", "Rest", "AP", "1", False), (["ses-wave1bas"], ["Axcpt"], ["AP"], ["1"], False),
("ses-wave1bas", "Axcpt", "AP", "1", False), (["ses-wave1bas"], ["Cuedts"], ["AP"], ["1"], False),
("ses-wave1bas", "Cuedts", "AP", "1", False), (["ses-wave1bas"], ["Stern"], ["AP"], ["1"], False),
("ses-wave1bas", "Stern", "AP", "1", False), (["ses-wave1bas"], ["Stroop"], ["AP"], ["1"], False),
("ses-wave1bas", "Stroop", "AP", "1", False), (["ses-wave1bas"], ["Rest"], ["PA"], ["2"], False),
("ses-wave1bas", "Rest", "PA", "2", False), (["ses-wave1bas"], ["Axcpt"], ["PA"], ["2"], False),
("ses-wave1bas", "Axcpt", "PA", "2", False), (["ses-wave1bas"], ["Cuedts"], ["PA"], ["2"], False),
("ses-wave1bas", "Cuedts", "PA", "2", False), (["ses-wave1bas"], ["Stern"], ["PA"], ["2"], False),
("ses-wave1bas", "Stern", "PA", "2", False), (["ses-wave1bas"], ["Stroop"], ["PA"], ["2"], False),
("ses-wave1bas", "Stroop", "PA", "2", False), (["ses-wave1bas"], ["Rest"], ["AP"], ["1"], True),
("ses-wave1bas", "Rest", "AP", "1", True), (["ses-wave1bas"], ["Axcpt"], ["AP"], ["1"], True),
("ses-wave1bas", "Axcpt", "AP", "1", True), (["ses-wave1bas"], ["Cuedts"], ["AP"], ["1"], True),
("ses-wave1bas", "Cuedts", "AP", "1", True), (["ses-wave1bas"], ["Stern"], ["AP"], ["1"], True),
("ses-wave1bas", "Stern", "AP", "1", True), (["ses-wave1bas"], ["Stroop"], ["AP"], ["1"], True),
("ses-wave1bas", "Stroop", "AP", "1", True), (["ses-wave1bas"], ["Rest"], ["PA"], ["2"], True),
("ses-wave1bas", "Rest", "PA", "2", True), (["ses-wave1bas"], ["Axcpt"], ["PA"], ["2"], True),
("ses-wave1bas", "Axcpt", "PA", "2", True), (["ses-wave1bas"], ["Cuedts"], ["PA"], ["2"], True),
("ses-wave1bas", "Cuedts", "PA", "2", True), (["ses-wave1bas"], ["Stern"], ["PA"], ["2"], True),
("ses-wave1bas", "Stern", "PA", "2", True), (["ses-wave1bas"], ["Stroop"], ["PA"], ["2"], True),
("ses-wave1bas", "Stroop", "PA", "2", True), (["ses-wave1pro"], ["Rest"], ["AP"], ["1"], False),
("ses-wave1pro", "Rest", "AP", "1", False), (["ses-wave1pro"], ["Rest"], ["PA"], ["2"], False),
("ses-wave1pro", "Rest", "PA", "2", False), (["ses-wave1pro"], ["Rest"], ["AP"], ["1"], True),
("ses-wave1pro", "Rest", "AP", "1", True), (["ses-wave1pro"], ["Rest"], ["PA"], ["2"], True),
("ses-wave1pro", "Rest", "PA", "2", True), (["ses-wave1rea"], ["Rest"], ["AP"], ["1"], False),
("ses-wave1rea", "Rest", "AP", "1", False), (["ses-wave1rea"], ["Rest"], ["PA"], ["2"], False),
("ses-wave1rea", "Rest", "PA", "2", False), (["ses-wave1rea"], ["Rest"], ["AP"], ["1"], True),
("ses-wave1rea", "Rest", "AP", "1", True), (["ses-wave1rea"], ["Rest"], ["PA"], ["2"], True),
("ses-wave1rea", "Rest", "PA", "2", True),
], ],
) )
def test_DMCC13Benchmark( def test_DMCC13Benchmark(
sessions: str | None, sessions: list[str],
tasks: str | None, tasks: list[str],
phase_encodings: str | None, phase_encodings: list[str],
runs: str | None, runs: list[str],
native_t1w: bool, native_t1w: bool,
) -> None: ) -> None:
"""Test DMCC13Benchmark DataGrabber. """Test DMCC13Benchmark DataGrabber.
Parameters Parameters
---------- ----------
sessions : str or None sessions : list of str
The parametrized session values. The parametrized session values.
tasks : str or None tasks : list of str
The parametrized task values. The parametrized task values.
phase_encodings : str or None phase_encodings : list of str
The parametrized phase encoding values. The parametrized phase encoding values.
runs : str or None runs : list of str
The parametrized run values. The parametrized run values.
native_t1w : bool native_t1w : bool
The parametrized values for fetching native T1w. The parametrized values for fetching native T1w.
""" """
dg = DMCC13Benchmark( dg = DMCC13Benchmark(
uri=URI,
sessions=sessions, sessions=sessions,
tasks=tasks, tasks=tasks,
phase_encodings=phase_encodings, phase_encodings=phase_encodings,
runs=runs, runs=runs,
native_t1w=native_t1w, native_t1w=native_t1w,
) )
# Set URI to Gin
dg.uri = URI
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element's access values
_, ses, task, phase, run = test_element _, ses, task, phase, run = test_element
# Access data
out = dg[("sub-01", ses, task, phase, run)] out = dg[("sub-01", ses, task, phase, run)]
# Available data types # Available data types
data_types = [ data_types = [
"BOLD", DataType.BOLD,
"VBM_CSF", DataType.VBM_CSF,
"VBM_GM", DataType.VBM_GM,
"VBM_WM", DataType.VBM_WM,
"T1w", DataType.T1w,
] ]
# Add Warp if native T1w is accessed # Add Warp if native T1w is accessed
if native_t1w: if native_t1w:
data_types.append("Warp") data_types.append(DataType.Warp)
# Data type file name formats # Data type file name formats
data_file_names = [ data_file_names = [
@ -131,9 +125,9 @@ def test_DMCC13Benchmark(
data_types, data_file_names, strict=False data_types, data_file_names, strict=False
): ):
# Assert data type # Assert data type
assert data_type in out assert data_type in out.keys()
# Conditional for Warp # Conditional for Warp
if data_type == "Warp": if data_type is DataType.Warp:
for idx, fname in enumerate(data_file_name): for idx, fname in enumerate(data_file_name):
# Assert data file path exists # Assert data file path exists
assert out[data_type][idx]["path"].exists() assert out[data_type][idx]["path"].exists()
@ -200,15 +194,15 @@ def test_DMCC13Benchmark(
@pytest.mark.parametrize( @pytest.mark.parametrize(
"types, native_t1w", "types, native_t1w",
[ [
("BOLD", True), (["BOLD"], True),
("BOLD", False), ("BOLD", False),
("T1w", True), (["T1w"], True),
("T1w", False), ("T1w", False),
("VBM_CSF", True), (["VBM_CSF"], True),
("VBM_CSF", False), ("VBM_CSF", False),
("VBM_GM", True), (["VBM_GM"], True),
("VBM_GM", False), ("VBM_GM", False),
("VBM_WM", True), (["VBM_WM"], True),
("VBM_WM", False), ("VBM_WM", False),
(["BOLD", "VBM_CSF"], True), (["BOLD", "VBM_CSF"], True),
(["BOLD", "VBM_CSF"], False), (["BOLD", "VBM_CSF"], False),
@ -232,66 +226,18 @@ def test_DMCC13Benchmark_partial_data_access(
The parametrized values for fetching native T1w. The parametrized values for fetching native T1w.
""" """
dg = DMCC13Benchmark(types=types, native_t1w=native_t1w) dg = DMCC13Benchmark(
# Set URI to Gin uri=URI,
dg.uri = URI types=types,
native_t1w=native_t1w,
)
with dg: with dg:
# Get all elements
all_elements = dg.get_elements() all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0] test_element = all_elements[0]
# Get test element's access values
_, ses, task, phase, run = test_element _, ses, task, phase, run = test_element
# Access data
out = dg[("sub-01", ses, task, phase, run)] out = dg[("sub-01", ses, task, phase, run)]
# Assert data type # Assert data type
if isinstance(types, list): if isinstance(types, str):
for type_ in types: types = [types]
assert type_ in out for type_ in types:
else: assert type_ in out
assert types in out
def test_DMCC13Benchmark_incorrect_data_type() -> None:
"""Test DMCC13Benchmark DataGrabber incorrect data type."""
with pytest.raises(
ValueError, match="`patterns` must contain all `types`"
):
_ = DMCC13Benchmark(types="Orcus")
def test_DMCC13Benchmark_invalid_sessions():
"""Test DMCC13Benchmark DataGrabber invalid sessions."""
with pytest.raises(
ValueError,
match=("phonyses is not a valid session in the DMCC dataset"),
):
DMCC13Benchmark(sessions="phonyses")
def test_DMCC13Benchmark_invalid_tasks():
"""Test DMCC13Benchmark DataGrabber invalid tasks."""
with pytest.raises(
ValueError,
match=("thisisnotarealtask is not a valid task in the DMCC dataset"),
):
DMCC13Benchmark(tasks="thisisnotarealtask")
def test_DMCC13Benchmark_phase_encodings():
"""Test DMCC13Benchmark DataGrabber invalid phase encodings."""
with pytest.raises(
ValueError,
match=("moonphase is not a valid phase encoding in the DMCC dataset"),
):
DMCC13Benchmark(phase_encodings="moonphase")
def test_DMCC13Benchmark_runs():
"""Test DMCC13Benchmark DataGrabber invalid runs."""
with pytest.raises(
ValueError,
match=("cerebralrun is not a valid run in the DMCC dataset"),
):
DMCC13Benchmark(runs="cerebralrun")

View file

@ -3,7 +3,10 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de> # Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path
import pytest import pytest
from pydantic import AnyUrl
from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber
@ -22,8 +25,8 @@ _testing_dataset = {
def test_MultipleDataGrabber() -> None: def test_MultipleDataGrabber() -> None:
"""Test MultipleDataGrabber.""" """Test MultipleDataGrabber."""
repo_uri = _testing_dataset["example_bids_ses"]["uri"] repo_uri = AnyUrl(_testing_dataset["example_bids_ses"]["uri"])
rootdir = "example_bids_ses" rootdir = Path("example_bids_ses")
replacements = ["subject", "session"] replacements = ["subject", "session"]
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
@ -73,7 +76,7 @@ def test_MultipleDataGrabber() -> None:
dg2 = PatternDataladDataGrabber( dg2 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=repo_uri, uri=repo_uri,
types=["BOLD"], types="BOLD",
patterns={ patterns={
"BOLD": { "BOLD": {
"pattern": ( "pattern": (
@ -93,7 +96,7 @@ def test_MultipleDataGrabber() -> None:
replacements=replacements, replacements=replacements,
) )
dg = MultipleDataGrabber([dg1, dg2]) dg = MultipleDataGrabber(datagrabbers=[dg1, dg2])
types = dg.get_types() types = dg.get_types()
assert "T1w" in types assert "T1w" in types
@ -129,12 +132,12 @@ def test_MultipleDataGrabber() -> None:
def test_MultipleDataGrabber_no_intersection() -> None: def test_MultipleDataGrabber_no_intersection() -> None:
"""Test MultipleDataGrabber without intersection (0 elements).""" """Test MultipleDataGrabber without intersection (0 elements)."""
rootdir = "example_bids_ses" rootdir = Path("example_bids_ses")
replacements = ["subject", "session"] replacements = ["subject", "session"]
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=_testing_dataset["example_bids"]["uri"], uri=AnyUrl(_testing_dataset["example_bids"]["uri"]),
types=["T1w", "Warp"], types=["T1w", "Warp"],
patterns={ patterns={
"T1w": { "T1w": {
@ -171,8 +174,8 @@ def test_MultipleDataGrabber_no_intersection() -> None:
dg2 = PatternDataladDataGrabber( dg2 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=_testing_dataset["example_bids_ses"]["uri"], uri=AnyUrl(_testing_dataset["example_bids_ses"]["uri"]),
types=["BOLD"], types="BOLD",
patterns={ patterns={
"BOLD": { "BOLD": {
"pattern": ( "pattern": (
@ -185,7 +188,7 @@ def test_MultipleDataGrabber_no_intersection() -> None:
replacements=replacements, replacements=replacements,
) )
dg = MultipleDataGrabber([dg1, dg2]) dg = MultipleDataGrabber(datagrabbers=[dg1, dg2])
expected_subs = set() expected_subs = set()
with dg: with dg:
subs = list(dg) subs = list(dg)
@ -195,9 +198,9 @@ def test_MultipleDataGrabber_no_intersection() -> None:
def test_MultipleDataGrabber_get_item() -> None: def test_MultipleDataGrabber_get_item() -> None:
"""Test MultipleDataGrabber get_item() error.""" """Test MultipleDataGrabber get_item() error."""
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
rootdir="example_bids_ses", rootdir=Path("example_bids_ses"),
uri=_testing_dataset["example_bids"]["uri"], uri=AnyUrl(_testing_dataset["example_bids"]["uri"]),
types=["T1w"], types="T1w",
patterns={ patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
@ -209,19 +212,19 @@ def test_MultipleDataGrabber_get_item() -> None:
replacements=["subject", "session"], replacements=["subject", "session"],
) )
dg = MultipleDataGrabber([dg1]) dg = MultipleDataGrabber(datagrabbers=[dg1])
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
dg.get_item(subject="sub-01") # type: ignore dg.get_item(subject="sub-01")
def test_MultipleDataGrabber_validation() -> None: def test_MultipleDataGrabber_validation() -> None:
"""Test MultipleDataGrabber init validation.""" """Test MultipleDataGrabber init validation."""
rootdir = "example_bids_ses" rootdir = Path("example_bids_ses")
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=_testing_dataset["example_bids"]["uri"], uri=AnyUrl(_testing_dataset["example_bids"]["uri"]),
types=["T1w"], types="T1w",
patterns={ patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
@ -235,8 +238,8 @@ def test_MultipleDataGrabber_validation() -> None:
dg2 = PatternDataladDataGrabber( dg2 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=_testing_dataset["example_bids_ses"]["uri"], uri=AnyUrl(_testing_dataset["example_bids_ses"]["uri"]),
types=["BOLD"], types="BOLD",
patterns={ patterns={
"BOLD": { "BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz", "pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz",
@ -247,22 +250,22 @@ def test_MultipleDataGrabber_validation() -> None:
) )
with pytest.raises(RuntimeError, match="have different element keys"): with pytest.raises(RuntimeError, match="have different element keys"):
MultipleDataGrabber([dg1, dg2]) MultipleDataGrabber(datagrabbers=[dg1, dg2])
with pytest.raises(RuntimeError, match="have overlapping mandatory"): with pytest.raises(RuntimeError, match="have overlapping mandatory"):
MultipleDataGrabber([dg1, dg1]) MultipleDataGrabber(datagrabbers=[dg1, dg1])
def test_MultipleDataGrabber_partial_pattern() -> None: def test_MultipleDataGrabber_partial_pattern() -> None:
"""Test MultipleDataGrabber partial pattern.""" """Test MultipleDataGrabber partial pattern."""
repo_uri = _testing_dataset["example_bids_ses"]["uri"] repo_uri = AnyUrl(_testing_dataset["example_bids_ses"]["uri"])
rootdir = "example_bids_ses" rootdir = Path("example_bids_ses")
replacements = ["subject", "session"] replacements = ["subject", "session"]
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=repo_uri, uri=repo_uri,
types=["BOLD"], types="BOLD",
patterns={ patterns={
"BOLD": { "BOLD": {
"pattern": ( "pattern": (
@ -278,7 +281,7 @@ def test_MultipleDataGrabber_partial_pattern() -> None:
dg2 = PatternDataladDataGrabber( dg2 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=repo_uri, uri=repo_uri,
types=["BOLD"], types="BOLD",
patterns={ patterns={
"BOLD": { "BOLD": {
"confounds": { "confounds": {
@ -295,7 +298,7 @@ def test_MultipleDataGrabber_partial_pattern() -> None:
partial_pattern_ok=True, partial_pattern_ok=True,
) )
dg = MultipleDataGrabber([dg1, dg2]) dg = MultipleDataGrabber(datagrabbers=[dg1, dg2])
types = dg.get_types() types = dg.get_types()
assert "BOLD" in types assert "BOLD" in types

View file

@ -10,7 +10,21 @@ from pathlib import Path
import pytest import pytest
from junifer.datagrabber import PatternDataGrabber from junifer.datagrabber import (
ConfoundsFormat,
PatternDataGrabber,
register_confounds_format,
)
def test_register_confounds_format() -> None:
"""Test confounds format registration."""
register_confounds_format(
name="Confounds",
alias="confounds",
)
assert "confounds" in list(ConfoundsFormat)
def test_PatternDataGrabber_errors(tmp_path: Path) -> None: def test_PatternDataGrabber_errors(tmp_path: Path) -> None:
@ -117,7 +131,7 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
""" """
datagrabber_first = PatternDataGrabber( datagrabber_first = PatternDataGrabber(
datadir="/tmp/data", datadir=Path("/tmp/data"),
types=["BOLD", "T1w"], types=["BOLD", "T1w"],
patterns={ patterns={
"BOLD": { "BOLD": {
@ -129,7 +143,7 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
"space": "native", "space": "native",
}, },
}, },
replacements="subject", replacements=["subject"],
) )
assert datagrabber_first.datadir == Path("/tmp/data") assert datagrabber_first.datadir == Path("/tmp/data")
assert set(datagrabber_first.types) == {"T1w", "BOLD"} assert set(datagrabber_first.types) == {"T1w", "BOLD"}
@ -181,7 +195,7 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
datagrabber_third = PatternDataGrabber( datagrabber_third = PatternDataGrabber(
datadir=tmpdir, datadir=tmpdir,
types=["T1w"], types="T1w",
patterns={ patterns={
"T1w": { "T1w": {
"pattern": "anat/{subject}_{session}.nii", "pattern": "anat/{subject}_{session}.nii",
@ -258,7 +272,7 @@ def test_PatternDataGrabber_unix_path_expansion(tmp_path: Path) -> None:
# Create datagrabber # Create datagrabber
dg = PatternDataGrabber( dg = PatternDataGrabber(
datadir=tmp_path, datadir=tmp_path,
types=["FreeSurfer"], types="FreeSurfer",
patterns={ patterns={
"FreeSurfer": { "FreeSurfer": {
"pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]", "pattern": "derivatives/freesurfer/[!f]{subject}/mri/T1.mg[z]",
@ -279,22 +293,3 @@ def test_PatternDataGrabber_unix_path_expansion(tmp_path: Path) -> None:
# Check paths are found # Check paths are found
assert set(out["FreeSurfer"].keys()) == {"path", "aseg", "meta"} assert set(out["FreeSurfer"].keys()) == {"path", "aseg", "meta"}
assert list(out["FreeSurfer"]["aseg"].keys()) == ["path"] assert list(out["FreeSurfer"]["aseg"].keys()) == ["path"]
def test_PatternDataGrabber_confounds_format_error_on_init() -> None:
"""Test PatterDataGrabber confounds format error on initialisation."""
with pytest.raises(
ValueError, match="Invalid value for `confounds_format`"
):
PatternDataGrabber(
types=["BOLD"],
patterns={
"BOLD": {
"pattern": "func/{subject}.nii",
"space": "MNI152NLin6Asym",
},
},
replacements=["subject"],
datadir="/tmp",
confounds_format="foobar",
)

View file

@ -7,9 +7,9 @@
from pathlib import Path from pathlib import Path
import pytest from pydantic import AnyUrl
from junifer.datagrabber import PatternDataladDataGrabber from junifer.datagrabber import DataType, PatternDataladDataGrabber
_testing_dataset = { _testing_dataset = {
@ -26,45 +26,25 @@ _testing_dataset = {
} }
def test_bids_PatternDataladDataGrabber_missing_uri() -> None:
"""Test check of missing URI in PatternDataladDataGrabber."""
with pytest.raises(ValueError, match=r"`uri` must be provided"):
PatternDataladDataGrabber(
datadir=None,
types=[],
patterns={},
replacements=[],
)
def test_bids_PatternDataladDataGrabber() -> None: def test_bids_PatternDataladDataGrabber() -> None:
"""Test subject-based BIDS PatternDataladDataGrabber.""" """Test subject-based BIDS PatternDataladDataGrabber."""
# Define types
types = ["T1w", "BOLD"]
# Define patterns
patterns = {
"T1w": {
"pattern": "{subject}/anat/{subject}_T1w.nii.gz",
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
}
# Define replacements
replacements = ["subject"]
repo_uri = _testing_dataset["example_bids"]["uri"]
rootdir = "example_bids"
repo_commit = _testing_dataset["example_bids"]["commit"] repo_commit = _testing_dataset["example_bids"]["commit"]
repo_uri = _testing_dataset["example_bids"]["uri"]
with PatternDataladDataGrabber( with PatternDataladDataGrabber(
rootdir=rootdir, uri=AnyUrl(repo_uri),
uri=repo_uri, types=[DataType.T1w, DataType.BOLD],
types=types, patterns={
patterns=patterns, "T1w": {
replacements=replacements, "pattern": "{subject}/anat/{subject}_T1w.nii.gz",
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
},
replacements=["subject"],
rootdir=Path("example_bids"),
) as dg: ) as dg:
subs = list(dg) subs = list(dg)
expected_subs = [f"sub-{i:02d}" for i in range(1, 10)] expected_subs = [f"sub-{i:02d}" for i in range(1, 10)]
@ -74,11 +54,11 @@ def test_bids_PatternDataladDataGrabber() -> None:
t_sub = dg[elem] t_sub = dg[elem]
assert "path" in t_sub["T1w"] assert "path" in t_sub["T1w"]
assert t_sub["T1w"]["path"] == ( assert t_sub["T1w"]["path"] == (
dg.datadir / f"{elem}/anat/{elem}_T1w.nii.gz" dg.fulldir / f"{elem}/anat/{elem}_T1w.nii.gz"
) )
assert "path" in t_sub["BOLD"] assert "path" in t_sub["BOLD"]
assert t_sub["BOLD"]["path"] == ( assert t_sub["BOLD"]["path"] == (
dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" dg.fulldir / f"{elem}/func/{elem}_task-rest_bold.nii.gz"
) )
assert "meta" in t_sub["BOLD"] assert "meta" in t_sub["BOLD"]
@ -88,7 +68,7 @@ def test_bids_PatternDataladDataGrabber() -> None:
assert "class" in dg_meta assert "class" in dg_meta
assert dg_meta["class"] == "PatternDataladDataGrabber" assert dg_meta["class"] == "PatternDataladDataGrabber"
assert "uri" in dg_meta assert "uri" in dg_meta
assert dg_meta["uri"] == repo_uri assert str(dg_meta["uri"]) == repo_uri
assert "datalad_commit_id" in dg_meta assert "datalad_commit_id" in dg_meta
assert dg_meta["datalad_commit_id"] == repo_commit assert dg_meta["datalad_commit_id"] == repo_commit
@ -98,80 +78,64 @@ def test_bids_PatternDataladDataGrabber() -> None:
def test_bids_PatternDataladDataGrabber_datadir() -> None: def test_bids_PatternDataladDataGrabber_datadir() -> None:
"""Test PatternDataladDataGrabber with a datadir set to a relative path.""" """Test PatternDataladDataGrabber with a datadir set to a relative path."""
# Define patterns datadir = Path("dataset") # use string and not absolute path
patterns = {
"T1w": {
"pattern": "{subject}/anat/{subject}_T*w.nii.gz",
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_*.nii.gz",
"space": "MNI152NLin6Asym",
},
}
# Define datadir
datadir = "dataset" # use string and not absolute path
with PatternDataladDataGrabber( with PatternDataladDataGrabber(
uri=_testing_dataset["example_bids"]["uri"], uri=AnyUrl(_testing_dataset["example_bids"]["uri"]),
types=["T1w", "BOLD"], types=[DataType.T1w, DataType.BOLD],
patterns=patterns, patterns={
datadir=datadir, "T1w": {
rootdir="example_bids", "pattern": "{subject}/anat/{subject}_T*w.nii.gz",
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_*.nii.gz",
"space": "MNI152NLin6Asym",
},
},
replacements=["subject"], replacements=["subject"],
datadir=datadir,
rootdir=Path("example_bids"),
) as dg: ) as dg:
assert dg.datadir == Path(datadir) / "example_bids" assert dg.fulldir == Path(datadir) / "example_bids"
for elem in dg: for elem in dg:
t_sub = dg[elem] t_sub = dg[elem]
assert "path" in t_sub["T1w"] assert "path" in t_sub["T1w"]
assert t_sub["T1w"]["path"] == ( assert t_sub["T1w"]["path"] == (
dg.datadir / f"{elem}/anat/{elem}_T1w.nii.gz" dg.fulldir / f"{elem}/anat/{elem}_T1w.nii.gz"
) )
assert "path" in t_sub["BOLD"] assert "path" in t_sub["BOLD"]
assert t_sub["BOLD"]["path"] == ( assert t_sub["BOLD"]["path"] == (
dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" dg.fulldir / f"{elem}/func/{elem}_task-rest_bold.nii.gz"
) )
def test_bids_PatternDataladDataGrabber_session(): def test_bids_PatternDataladDataGrabber_session():
"""Test a subject and session-based BIDS PatternDataladDataGrabber.""" """Test a subject and session-based BIDS PatternDataladDataGrabber."""
types = ["T1w", "BOLD"]
patterns = {
"T1w": {
"pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
),
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": (
"{subject}/{session}/func/"
"{subject}_{session}_task-rest_bold.nii.gz"
),
"space": "MNI152NLin6Asym",
},
}
replacements = ["subject", "session"]
# Check error
with pytest.raises(ValueError, match=r"`uri` must be provided"):
PatternDataladDataGrabber(
datadir=None,
types=types,
patterns=patterns,
replacements=replacements,
)
# Set parameters # Set parameters
repo_uri = _testing_dataset["example_bids_ses"]["uri"] repo_uri = AnyUrl(_testing_dataset["example_bids_ses"]["uri"])
rootdir = "example_bids_ses" rootdir = Path("example_bids_ses")
replacements = ["subject", "session"]
# With T1W and bold, only 2 sessions are available # With T1W and bold, only 2 sessions are available
with PatternDataladDataGrabber( with PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri, uri=repo_uri,
types=types, types=[DataType.T1w, DataType.BOLD],
patterns=patterns, patterns={
"T1w": {
"pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
),
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": (
"{subject}/{session}/func/"
"{subject}_{session}_task-rest_bold.nii.gz"
),
"space": "MNI152NLin6Asym",
},
},
replacements=replacements, replacements=replacements,
rootdir=rootdir,
) as dg: ) as dg:
subs = list(dg.get_elements()) subs = list(dg.get_elements())
expected_subs = [ expected_subs = [
@ -182,21 +146,19 @@ def test_bids_PatternDataladDataGrabber_session():
assert set(subs) == set(expected_subs) assert set(subs) == set(expected_subs)
# Test with a different T1w only, it should have 3 sessions # Test with a different T1w only, it should have 3 sessions
types = ["T1w"]
patterns = {
"T1w": {
"pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
),
"space": "MNI152NLin6Asym",
},
}
with PatternDataladDataGrabber( with PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri, uri=repo_uri,
types=types, types=DataType.T1w,
patterns=patterns, patterns={
"T1w": {
"pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
),
"space": "MNI152NLin6Asym",
},
},
replacements=replacements, replacements=replacements,
rootdir=rootdir,
) as dg: ) as dg:
subs = list(dg) subs = list(dg)
expected_subs = [ expected_subs = [

View file

@ -8,7 +8,8 @@ from contextlib import AbstractContextManager, nullcontext
import pytest import pytest
from junifer.datagrabber.pattern_validation_mixin import ( from junifer.datagrabber import (
DataType,
DataTypeManager, DataTypeManager,
DataTypeSchema, DataTypeSchema,
PatternValidationMixin, PatternValidationMixin,
@ -79,7 +80,7 @@ def test_dtype_mgr(dtype: DataTypeSchema) -> None:
Parameters Parameters
---------- ----------
dtype : DataTypeSchema dtype : ``DataTypeSchema``
The parametrized schema. The parametrized schema.
""" """
@ -110,33 +111,15 @@ def test_register_data_type() -> None:
) )
assert "dtype" in DataTypeManager() assert "dtype" in DataTypeManager()
assert "dtype" in list(DataType)
_ = DataTypeManager().pop("dtype") _ = DataTypeManager().pop("dtype")
assert "dumb" not in DataTypeManager() assert "dtype" not in DataTypeManager()
assert "dtype" in list(DataType)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"types, replacements, patterns, expect", "types, replacements, patterns, expect",
[ [
(
"wrong",
[],
{},
pytest.raises(TypeError, match="`types` must be a list"),
),
(
[1],
[],
{},
pytest.raises(
TypeError, match="`types` must be a list of strings"
),
),
(
["BOLD"],
[],
"wrong",
pytest.raises(TypeError, match="`patterns` must be a dict"),
),
( (
["T1w", "BOLD"], ["T1w", "BOLD"],
"", "",
@ -204,30 +187,6 @@ def test_register_data_type() -> None:
}, },
pytest.raises(ValueError, match="following a replacement"), pytest.raises(ValueError, match="following a replacement"),
), ),
(
["T1w"],
"wrong",
{
"T1w": {
"pattern": "{subject}/anat/{subject}_T1w.nii",
"space": "native",
},
},
pytest.raises(TypeError, match="`replacements` must be a list"),
),
(
["T1w"],
[1],
{
"T1w": {
"pattern": "{subject}/anat/{subject}_T1w.nii",
"space": "native",
},
},
pytest.raises(
TypeError, match="`replacements` must be a list of strings"
),
),
( (
["T1w", "BOLD"], ["T1w", "BOLD"],
["subject", "session"], ["subject", "session"],

View file

@ -8,6 +8,7 @@ from pathlib import Path
import nibabel as nib import nibabel as nib
import pandas as pd import pandas as pd
from pydantic import BaseModel
from ..api.decorators import register_datareader from ..api.decorators import register_datareader
from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
@ -33,7 +34,7 @@ _readers["TSV"] = {"func": pd.read_csv, "params": {"sep": "\t"}}
@register_datareader @register_datareader
class DefaultDataReader(PipelineStepMixin, UpdateMetaMixin): class DefaultDataReader(BaseModel, PipelineStepMixin, UpdateMetaMixin):
"""Concrete implementation for common data reading.""" """Concrete implementation for common data reading."""
def validate_input(self, input: list[str]) -> list[str]: def validate_input(self, input: list[str]) -> list[str]:

View file

@ -30,7 +30,7 @@ def test_DefaultDataReader_validation(type_) -> None:
""" """
reader = DefaultDataReader() reader = DefaultDataReader()
assert reader.validate_input(type_) == type_ assert reader.validate_input(type_) == type_
assert reader.validate(type_) == type_ assert reader.validate_component(type_) == type_
def test_DefaultDataReader_meta() -> None: def test_DefaultDataReader_meta() -> None:

View file

@ -2,3 +2,8 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import lazy_loader as lazy
__getattr__, __dir__, __all__ = lazy.attach_stub(__name__, __file__)

0
junifer/external/__init__.pyi vendored Normal file
View file

0
junifer/external/py.typed vendored Normal file
View file

View file

@ -11,9 +11,11 @@ __all__ = [
"EdgeCentricFCMaps", "EdgeCentricFCMaps",
"EdgeCentricFCParcels", "EdgeCentricFCParcels",
"EdgeCentricFCSpheres", "EdgeCentricFCSpheres",
"ReHoImpl",
"ReHoMaps", "ReHoMaps",
"ReHoParcels", "ReHoParcels",
"ReHoSpheres", "ReHoSpheres",
"ALFFImpl",
"ALFFMaps", "ALFFMaps",
"ALFFParcels", "ALFFParcels",
"ALFFSpheres", "ALFFSpheres",
@ -37,8 +39,8 @@ from .functional_connectivity import (
EdgeCentricFCParcels, EdgeCentricFCParcels,
EdgeCentricFCSpheres, EdgeCentricFCSpheres,
) )
from .reho import ReHoMaps, ReHoParcels, ReHoSpheres from .reho import ReHoImpl, ReHoMaps, ReHoParcels, ReHoSpheres
from .falff import ALFFMaps, ALFFParcels, ALFFSpheres from .falff import ALFFImpl, ALFFMaps, ALFFParcels, ALFFSpheres
from .temporal_snr import ( from .temporal_snr import (
TemporalSNRMaps, TemporalSNRMaps,
TemporalSNRParcels, TemporalSNRParcels,

View file

@ -8,7 +8,11 @@ from abc import ABC, abstractmethod
from copy import deepcopy from copy import deepcopy
from typing import Any, ClassVar from typing import Any, ClassVar
from pydantic import BaseModel, ConfigDict
from ..datagrabber import DataType
from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
from ..storage import StorageType
from ..typing import MarkerInOutMappings, StorageLike from ..typing import MarkerInOutMappings, StorageLike
from ..utils import logger, raise_error from ..utils import logger, raise_error
@ -16,7 +20,7 @@ from ..utils import logger, raise_error
__all__ = ["BaseMarker"] __all__ = ["BaseMarker"]
class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin): class BaseMarker(BaseModel, ABC, PipelineStepMixin, UpdateMetaMixin):
"""Abstract base class for marker. """Abstract base class for marker.
For every marker, one needs to provide a concrete For every marker, one needs to provide a concrete
@ -24,12 +28,17 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
Parameters Parameters
---------- ----------
on : str or list of str or None, optional on : :enum:`.DataType` or list of variants or None, optional
The data type to apply the marker on. If None, The data type(s) to apply the marker on.
will work on all available data types (default None). If None, will work on all available data types.
name : str, optional Check :enum:`.DataType` for valid values (default None).
The name of the marker. If None, will use the class name as the name : str or None, optional
name of the marker (default None). The name of the marker.
If None, will use the class name (default None).
Attributes
----------
valid_inputs
Raises Raises
------ ------
@ -42,11 +51,12 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] _MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings]
def __init__( model_config = ConfigDict(use_enum_values=True)
self,
on: list[str] | str | None = None, on: list[DataType] | None = None
name: str | None = None, name: str | None = None
) -> None:
def model_post_init(self, context: Any): # noqa: D102
# Check for missing mapping attribute # Check for missing mapping attribute
if not hasattr(self, "_MARKER_INOUT_MAPPINGS"): if not hasattr(self, "_MARKER_INOUT_MAPPINGS"):
raise_error( raise_error(
@ -54,18 +64,48 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
klass=AttributeError, klass=AttributeError,
) )
# Use all data types if not provided # Use all data types if not provided
if on is None: if self.on is None:
on = self.get_valid_inputs() self.on = self.valid_inputs
# Convert data types to list else:
if not isinstance(on, list): # Convert to correct data type
on = [on] self.on = [
DataType(t) if isinstance(t, str) else t for t in self.on
]
# Check if required input data types are provided
if any(x not in self.valid_inputs for x in self.on):
wrong_on = [
x.value for x in self.on if x not in self.valid_inputs
]
raise_error(
f"{self.__class__.__name__} cannot be computed on "
f"{wrong_on}"
)
# Run extra validation for markers and fail early if needed
self.validate_marker_params()
# Set default name if not provided # Set default name if not provided
self.name = self.__class__.__name__ if name is None else name self.name = self.__class__.__name__ if self.name is None else self.name
# Check if required inputs are found
if any(x not in self.get_valid_inputs() for x in on): @property
wrong_on = [x for x in on if x not in self.get_valid_inputs()] def valid_inputs(self) -> list[DataType]:
raise_error(f"{self.name} cannot be computed on {wrong_on}") """Valid data types to operate on.
self._on = on
Returns
-------
list of :enum:`.DataType`
The list of data types that can be used as input for this marker.
"""
return [
DataType(x) if isinstance(x, str) else x
for x in self._MARKER_INOUT_MAPPINGS.keys()
]
def validate_marker_params(self) -> None:
"""Run extra logical validation for marker.
Subclasses can override to provide validation.
"""
pass
def validate_input(self, input: list[str]) -> list[str]: def validate_input(self, input: list[str]) -> list[str]:
"""Validate input. """Validate input.
@ -88,38 +128,29 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
If the input does not have the required data. If the input does not have the required data.
""" """
if not any(x in input for x in self._on): if not any(x in input for x in self.on):
raise_error( raise_error(
"Input does not have the required data." "Input does not have the required data.\n"
f"\t Input: {input}" f"\t Input: {input}\n"
f"\t Required (any of): {self._on}" f"\t Required (any of): {[t.value for t in self.on]}"
) )
return [x for x in self._on if x in input] return [x.value for x in self.on if x in input]
def get_valid_inputs(self) -> list[str]: def storage_type(
"""Get valid data types for input. self, input_type: DataType, output_feature: str
) -> StorageType:
Returns """Get :enum:`.StorageType` for a feature.
-------
list of str
The list of data types that can be used as input for this marker.
"""
return list(self._MARKER_INOUT_MAPPINGS.keys())
def storage_type(self, input_type: str, output_feature: str) -> str:
"""Get storage type for a feature.
Parameters Parameters
---------- ----------
input_type : str input_type : :enum:`.DataType`
The data type input to the marker. The data type input to the marker.
output_feature : str output_feature : str
The feature output of the marker. The feature output of the marker.
Returns Returns
------- -------
str :enum:`.StorageType`
The storage type output of the marker. The storage type output of the marker.
""" """
@ -155,28 +186,28 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
def store( def store(
self, self,
type_: str, data_type: DataType,
feature: str, feature: str,
out: dict[str, Any], output: dict[str, Any],
storage: StorageLike, storage: StorageLike,
) -> None: ) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
type_ : str data_type : :enum:`.DataType`
The data type to store. The data type to store.
feature : str feature : str
The feature to store. The feature to store.
out : dict output : 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.storage_type(type_, feature) s_type = self.storage_type(data_type, feature)
logger.debug(f"Storing {output_type_} in {storage}") logger.debug(f"Storing {s_type} in {storage}")
storage.store(kind=output_type_, **out) storage.store(kind=s_type, **output)
def _fit_transform( def _fit_transform(
self, self,
@ -195,58 +226,56 @@ class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
Returns Returns
------- -------
dict dict
The processed output as a dictionary. If `storage` is provided, The processed output as a dictionary. If ``storage`` is provided,
empty dictionary is returned. empty dictionary is returned.
""" """
out = {} out = {}
for type_ in self._on: for t in self.on:
if type_ in input.keys(): if t in input.keys():
logger.info(f"Computing {type_}") logger.info(f"Computing {t}")
# Get data dict for data type # Get data dict for data type
t_input = input[type_] t_input = input[t]
# Pass the other data types as extra input, removing # Pass the other data types as extra input, removing
# the current type # the current type
extra_input = input.copy() extra_input = input.copy()
extra_input.pop(type_) extra_input.pop(t)
logger.debug( logger.debug(
f"Extra data type for feature extraction: " f"Extra data type for feature extraction: "
f"{extra_input.keys()}" f"{extra_input.keys()}"
) )
# Copy metadata # Copy metadata
t_meta = t_input["meta"].copy() t_meta = t_input["meta"].copy()
t_meta["type"] = type_ t_meta["type"] = t.value
# 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)
# Initialize empty dictionary if no storage object is provided # Initialize empty dictionary if no storage object is provided
if storage is None: if storage is None:
out[type_] = {} out[t] = {}
# Store individual features # Store individual features
for feature_name, feature_data in t_out.items(): for f_name, f_data in t_out.items():
# Make deep copy of the feature data for manipulation # Make deep copy of the feature data for manipulation
feature_data_copy = deepcopy(feature_data) f_data_copy = deepcopy(f_data)
# Make deep copy of metadata and add to feature data # Make deep copy of metadata and add to feature data
feature_data_copy["meta"] = deepcopy(t_meta) f_data_copy["meta"] = deepcopy(t_meta)
# Update metadata for the feature, # Update metadata for the feature,
# feature data is not manipulated, only meta # feature data is not manipulated, only meta
self.update_meta(feature_data_copy, "marker") self.update_meta(f_data_copy, "marker")
# Update marker feature's metadata name # Update marker feature's metadata name
feature_data_copy["meta"]["marker"]["name"] += ( f_data_copy["meta"]["marker"]["name"] += f"_{f_name}"
f"_{feature_name}"
)
if storage is not None: if storage is not None:
logger.info(f"Storing in {storage}") logger.info(f"Storing in {storage}")
self.store( self.store(
type_=type_, data_type=t,
feature=feature_name, feature=f_name,
out=feature_data_copy, output=f_data_copy,
storage=storage, storage=storage,
) )
else: else:
logger.info( logger.info(
"No storage specified, returning dictionary" "No storage specified, returning dictionary"
) )
out[type_][feature_name] = feature_data_copy out[t][f_name] = f_data_copy
return out return out

View file

@ -12,14 +12,17 @@ from typing import (
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
from pydantic import PositiveInt
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..datagrabber import DataType
from ..external.BrainPrint.brainprint.brainprint import ( from ..external.BrainPrint.brainprint.brainprint import (
compute_asymmetry, compute_asymmetry,
compute_brainprint, compute_brainprint,
) )
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 ExtDep, WorkDirManager
from ..storage import StorageType
from ..typing import Dependencies, ExternalDependencies, MarkerInOutMappings from ..typing import Dependencies, ExternalDependencies, MarkerInOutMappings
from ..utils import logger, run_ext_cmd from ..utils import logger, run_ext_cmd
from .base import BaseMarker from .base import BaseMarker
@ -58,15 +61,15 @@ class BrainPrint(BaseMarker):
execution speed. Requires the ``scikit-sparse`` library. If it cannot execution speed. Requires the ``scikit-sparse`` library. If it cannot
be found, an error will be thrown. If False, will use slower LU be found, an error will be thrown. If False, will use slower LU
decomposition (default False). decomposition (default False).
name : str, optional name : str or None, optional
The name of the marker. If None, will use the class name (default The name of the marker.
None). If None, will use the class name (default None).
""" """
_EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [ _EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [
{ {
"name": "freesurfer", "name": ExtDep.FreeSurfer,
"commands": [ "commands": [
"mri_binarize", "mri_binarize",
"mri_pretess", "mri_pretess",
@ -79,35 +82,25 @@ class BrainPrint(BaseMarker):
_DEPENDENCIES: ClassVar[Dependencies] = {"lapy", "numpy"} _DEPENDENCIES: ClassVar[Dependencies] = {"lapy", "numpy"}
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = { _MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"FreeSurfer": { DataType.FreeSurfer: {
"eigenvalues": "scalar_table", "eigenvalues": StorageType.ScalarTable,
"areas": "vector", "areas": StorageType.Vector,
"volumes": "vector", "volumes": StorageType.Vector,
"distances": "vector", "distances": StorageType.Vector,
} }
} }
def __init__( num: PositiveInt = 50
self, skip_cortex: bool = False
num: int = 50, keep_eigenvectors: bool = False
skip_cortex=False, norm: str = "none"
keep_eigenvectors: bool = False, reweight: bool = False
norm: str = "none", asymmetry: bool = False
reweight: bool = False, asymmetry_distance: str = "euc"
asymmetry: bool = False, use_cholmod: bool = False
asymmetry_distance: str = "euc",
use_cholmod: bool = False, _tempdir = Path()
name: str | None = None, _element_tempdir = Path()
) -> None:
self.num = num
self.skip_cortex = skip_cortex
self.keep_eigenvectors = keep_eigenvectors
self.norm = norm
self.reweight = reweight
self.asymmetry = asymmetry
self.asymmetry_distance = asymmetry_distance
self.use_cholmod = use_cholmod
super().__init__(name=name, on="FreeSurfer")
def _create_aseg_surface( def _create_aseg_surface(
self, self,
@ -351,7 +344,7 @@ class BrainPrint(BaseMarker):
- ``col_names`` : surface labels as list of str - ``col_names`` : surface labels as list of str
- ``row_names`` : eigenvalue count labels as list of str - ``row_names`` : eigenvalue count labels as list of str
- ``row_header_col_name`` : "eigenvalue" - ``row_header_col_name`` : "eigenvalue"
()
* ``areas`` : dictionary with the following keys: * ``areas`` : dictionary with the following keys:
- ``data`` : areas as ``np.ndarray`` - ``data`` : areas as ``np.ndarray``
@ -362,7 +355,7 @@ class BrainPrint(BaseMarker):
- ``data`` : volumes as ``np.ndarray`` - ``data`` : volumes as ``np.ndarray``
- ``col_names`` : surface labels as list of str - ``col_names`` : surface labels as list of str
* ``distances`` : dictionary with the following keys * ``distances`` : dictionary with the following keys \
if ``asymmetry = True``: if ``asymmetry = True``:
- ``data`` : distances as ``np.ndarray`` - ``data`` : distances as ``np.ndarray``

View file

@ -1,17 +1,23 @@
"""Provide base class for complexity.""" """Provide base class for complexity."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from abc import abstractmethod from abc import abstractmethod
from typing import ( from typing import (
TYPE_CHECKING, TYPE_CHECKING,
Annotated,
Any, Any,
ClassVar, ClassVar,
) )
from pydantic import BeforeValidator
from ...datagrabber import DataType
from ...storage import StorageType
from ...typing import Dependencies, MarkerInOutMappings from ...typing import Dependencies, MarkerInOutMappings
from ...utils import raise_error from ...utils import ensure_list, ensure_list_or_none, raise_error
from ..base import BaseMarker from ..base import BaseMarker
from ..parcel_aggregation import ParcelAggregation from ..parcel_aggregation import ParcelAggregation
@ -24,50 +30,45 @@ __all__ = ["ComplexityBase"]
class ComplexityBase(BaseMarker): class ComplexityBase(BaseMarker):
"""Base class for complexity computation. """Abstract base class for complexity computation.
Parameters Parameters
---------- ----------
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) to use.
:func:`junifer.data.parcellations.list_parcellations`. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
name : str, optional name : str or None, optional
The name of the marker. If None, it will use the class name The name of the marker.
(default None). If None, will use the class name (default None).
""" """
_DEPENDENCIES: ClassVar[Dependencies] = {"nilearn", "neurokit2"} _DEPENDENCIES: ClassVar[Dependencies] = {"nilearn", "neurokit2"}
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = { _MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"BOLD": { DataType.BOLD: {
"complexity": "vector", "complexity": StorageType.Vector,
}, },
} }
def __init__( parcellation: Annotated[str | list[str], BeforeValidator(ensure_list)]
self, agg_method: str = "mean"
parcellation: str | list[str], agg_method_params: dict | None = None
agg_method: str = "mean", masks: Annotated[
agg_method_params: dict | None = None, dict | str | list[dict | str] | None,
masks: str | dict | list[dict | str] | None = None, BeforeValidator(ensure_list_or_none),
name: str | None = None, ] = None
) -> None:
self.parcellation = parcellation
self.agg_method = agg_method
self.agg_method_params = agg_method_params
self.masks = masks
super().__init__(on="BOLD", name=name)
@abstractmethod @abstractmethod
def compute_complexity( def compute_complexity(
@ -115,7 +116,7 @@ class ComplexityBase(BaseMarker):
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=DataType.BOLD,
).compute(input=input, extra_input=extra_input) ).compute(input=input, extra_input=extra_input)
# Compute complexity measure # Compute complexity measure
return { return {

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,22 +26,23 @@ class HurstExponent(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the Hurst exponent calculation function. For more The parameters to pass to the Hurst exponent calculation function.
information, check out ``junifer.markers.utils._hurst_exponent``. See ``junifer.markers.utils._hurst_exponent`` for more information.
If None, value is set to {"method": "dfa"} (default None). If None, value is set to ``{"method": "dfa"}`` (default None).
name : str, optional name : str or None, optional
The name of the marker. If None, it will use the class name The name of the marker.
(default None). If None, will use the class name (default None).
Warnings Warnings
-------- --------
@ -51,26 +53,12 @@ class HurstExponent(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"method": "dfa"} self.params = {"method": "dfa"}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,23 +26,25 @@ class MultiscaleEntropyAUC(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the AUC of multiscale entropy calculation The parameters to pass to the AUC of multiscale entropy calculation
function. For more information, check out function. See
``junifer.markers.utils._multiscale_entropy_auc``. If None, value ``junifer.markers.utils._multiscale_entropy_auc`` for more information.
is set to {"m": 2, "tol": 0.5, "scale": 10} (default None). If None, value is set to ``{"m": 2, "tol": 0.5, "scale": 10}``
name : str, optional
The name of the marker. If None, it will use the class name
(default None). (default None).
name : str or None, optional
The name of the marker.
If None, will use the class name (default None).
Warnings Warnings
-------- --------
@ -52,26 +55,12 @@ class MultiscaleEntropyAUC(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"m": 2, "tol": 0.5, "scale": 10} self.params = {"m": 2, "tol": 0.5, "scale": 10}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,23 +26,23 @@ class PermEntropy(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the permutation entropy calculation function. The parameters to pass to the permutation entropy calculation function.
For more information, check out See ``junifer.markers.utils._perm_entropy`` for more information.
``junifer.markers.utils._perm_entropy``. If None, value is set to If None, value is set to ``{"m": 2, "delay": 1}`` (default None).
{"m": 2, "delay": 1} (default None). name : str or None, optional
name : str, optional The name of the marker.
The name of the marker. If None, it will use the class name If None, will use the class name (default None).
(default None).
Warnings Warnings
-------- --------
@ -52,26 +53,12 @@ class PermEntropy(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"m": 4, "delay": 1} self.params = {"m": 4, "delay": 1}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,23 +26,24 @@ class RangeEntropy(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the range entropy calculation function. For more The parameters to pass to the range entropy calculation function.
information, check out ``junifer.markers.utils._range_entropy``. See ``junifer.markers.utils._range_entropy`` for more information.
If None, value is set to {"m": 2, "tol": 0.5, "delay": 1} If None, value is set to ``{"m": 2, "tol": 0.5, "delay": 1}``
(default None).
name : str, optional
The name of the marker. If None, it will use the class name
(default None). (default None).
name : str or None, optional
The name of the marker.
If None, will use the class name (default None).
Warnings Warnings
-------- --------
@ -52,26 +54,12 @@ class RangeEntropy(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"m": 2, "tol": 0.5, "delay": 1} self.params = {"m": 2, "tol": 0.5, "delay": 1}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,23 +26,24 @@ class RangeEntropyAUC(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the range entropy calculation function. For more The parameters to pass to the range entropy calculation function.
information, check out ``junifer.markers.utils._range_entropy``. See ``junifer.markers.utils._range_entropy`` for more information.
If None, value is set to {"m": 2, "delay": 1, "n_r": 10} If None, value is set to ``{"m": 2, "delay": 1, "n_r": 10}``
(default None).
name : str, optional
The name of the marker. If None, it will use the class name
(default None). (default None).
name : str or None, optional
The name of the marker.
If None, will use the class name (default None).
Warnings Warnings
-------- --------
@ -52,26 +54,12 @@ class RangeEntropyAUC(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"m": 2, "delay": 1, "n_r": 10} self.params = {"m": 2, "delay": 1, "n_r": 10}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,24 +26,24 @@ class SampleEntropy(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the sample entropy calculation function. The parameters to pass to the sample entropy calculation function.
For more information, check out See ``junifer.markers.utils._sample_entropy`` for more information.
``junifer.markers.utils._sample_entropy``. If None, value is set to ``{"m": 2, "delay": 1, "tol": 0.5}``
If None, value is set to
{"m": 2, "delay": 1, "tol": 0.5} (default None).
name : str, optional
The name of the marker. If None, it will use the class name
(default None). (default None).
name : str or None, optional
The name of the marker.
If None, will use the class name (default None).
Warnings Warnings
-------- --------
@ -53,26 +54,12 @@ class SampleEntropy(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"m": 4, "delay": 1, "tol": 0.5} self.params = {"m": 4, "delay": 1, "tol": 0.5}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -12,6 +12,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import HurstExponent from junifer.markers.complexity import HurstExponent
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -46,7 +47,7 @@ def test_compute() -> None:
def test_storage_type() -> None: def test_storage_type() -> None:
"""Test HurstExponent storage_type.""" """Test HurstExponent storage_type."""
assert "vector" == HurstExponent(parcellation=PARCELLATION).storage_type( assert "vector" == HurstExponent(parcellation=PARCELLATION).storage_type(
input_type="BOLD", output_feature="complexity" input_type=DataType.BOLD, output_feature="complexity"
) )

View file

@ -11,6 +11,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import MultiscaleEntropyAUC from junifer.markers.complexity import MultiscaleEntropyAUC
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -46,7 +47,7 @@ def test_storage_type() -> None:
"""Test MultiscaleEntropyAUC storage_type.""" """Test MultiscaleEntropyAUC storage_type."""
assert "vector" == MultiscaleEntropyAUC( assert "vector" == MultiscaleEntropyAUC(
parcellation=PARCELLATION parcellation=PARCELLATION
).storage_type(input_type="BOLD", output_feature="complexity") ).storage_type(input_type=DataType.BOLD, output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -11,6 +11,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import PermEntropy from junifer.markers.complexity import PermEntropy
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -45,7 +46,7 @@ def test_compute() -> None:
def test_storage_type() -> None: def test_storage_type() -> None:
"""Test PermEntropy storage_type.""" """Test PermEntropy storage_type."""
assert "vector" == PermEntropy(parcellation=PARCELLATION).storage_type( assert "vector" == PermEntropy(parcellation=PARCELLATION).storage_type(
input_type="BOLD", output_feature="complexity" input_type=DataType.BOLD, output_feature="complexity"
) )

View file

@ -12,6 +12,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import RangeEntropy from junifer.markers.complexity import RangeEntropy
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -46,7 +47,7 @@ def test_compute() -> None:
def test_storage_type() -> None: def test_storage_type() -> None:
"""Test RangeEntropy storage_type.""" """Test RangeEntropy storage_type."""
assert "vector" == RangeEntropy(parcellation=PARCELLATION).storage_type( assert "vector" == RangeEntropy(parcellation=PARCELLATION).storage_type(
input_type="BOLD", output_feature="complexity" input_type=DataType.BOLD, output_feature="complexity"
) )

View file

@ -12,6 +12,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import RangeEntropyAUC from junifer.markers.complexity import RangeEntropyAUC
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -46,7 +47,7 @@ def test_compute() -> None:
def test_storage_type() -> None: def test_storage_type() -> None:
"""Test RangeEntropyAUC storage_type.""" """Test RangeEntropyAUC storage_type."""
assert "vector" == RangeEntropyAUC(parcellation=PARCELLATION).storage_type( assert "vector" == RangeEntropyAUC(parcellation=PARCELLATION).storage_type(
input_type="BOLD", output_feature="complexity" input_type=DataType.BOLD, output_feature="complexity"
) )

View file

@ -11,6 +11,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import SampleEntropy from junifer.markers.complexity import SampleEntropy
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -45,7 +46,7 @@ def test_compute() -> None:
def test_storage_type() -> None: def test_storage_type() -> None:
"""Test SampleEntropy storage_type.""" """Test SampleEntropy storage_type."""
assert "vector" == SampleEntropy(parcellation=PARCELLATION).storage_type( assert "vector" == SampleEntropy(parcellation=PARCELLATION).storage_type(
input_type="BOLD", output_feature="complexity" input_type=DataType.BOLD, output_feature="complexity"
) )

View file

@ -11,6 +11,7 @@ import pytest
pytest.importorskip("neurokit2") pytest.importorskip("neurokit2")
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.complexity import WeightedPermEntropy from junifer.markers.complexity import WeightedPermEntropy
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
@ -46,7 +47,7 @@ def test_storage_type() -> None:
"""Test WeightedPermEntropy storage_type.""" """Test WeightedPermEntropy storage_type."""
assert "vector" == WeightedPermEntropy( assert "vector" == WeightedPermEntropy(
parcellation=PARCELLATION parcellation=PARCELLATION
).storage_type(input_type="BOLD", output_feature="complexity") ).storage_type(input_type=DataType.BOLD, output_feature="complexity")
@pytest.mark.skipif( @pytest.mark.skipif(

View file

@ -2,6 +2,7 @@
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import neurokit2 as nk import neurokit2 as nk
@ -25,24 +26,24 @@ class WeightedPermEntropy(ComplexityBase):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`junifer.stats.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
params : dict, optional params : dict or None, optional
Parameters to pass to the weighted permutation entropy calculation The parameters to pass to the weighted permutation entropy calculation
function. function. See ``junifer.markers.utils._weighted_perm_entropy`` for more
For more information, check out information. If None, value is set to ``{"m": 2, "delay": 1}``
``junifer.markers.utils._weighted_perm_entropy``. If None, value
is set to {"m": 2, "delay": 1} (default None).
name : str, optional
The name of the marker. If None, it will use the class name
(default None). (default None).
name : str or None, optional
The name of the marker.
If None, will use the class name (default None).
Warnings Warnings
-------- --------
@ -53,26 +54,12 @@ class WeightedPermEntropy(ComplexityBase):
""" """
def __init__( params: dict | None = None
self,
parcellation: str | list[str], def validate_marker_params(self) -> None:
agg_method: str = "mean", """Run extra logical validation for marker."""
agg_method_params: dict | None = None, if self.params is None:
masks: str | dict | list[dict | str] | None = None,
params: dict | None = None,
name: str | None = None,
) -> None:
super().__init__(
parcellation=parcellation,
agg_method=agg_method,
agg_method_params=agg_method_params,
masks=masks,
name=name,
)
if params is None:
self.params = {"m": 4, "delay": 1} self.params = {"m": 4, "delay": 1}
else:
self.params = params
def compute_complexity( def compute_complexity(
self, self,

View file

@ -6,13 +6,16 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any, ClassVar from typing import Annotated, Any, ClassVar
import numpy as np import numpy as np
from pydantic import BeforeValidator
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..datagrabber import DataType
from ..storage import StorageType
from ..typing import Dependencies, MarkerInOutMappings from ..typing import Dependencies, MarkerInOutMappings
from ..utils import logger from ..utils import ensure_list, ensure_list_or_none, logger
from .base import BaseMarker from .base import BaseMarker
from .parcel_aggregation import ParcelAggregation from .parcel_aggregation import ParcelAggregation
from .utils import _ets from .utils import _ets
@ -31,42 +34,37 @@ class RSSETSMarker(BaseMarker):
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options
agg_method_params : dict, optional (default "mean").
Parameters to pass to the aggregation function. Check valid options in agg_method_params : dict or None, optional
:func:`.get_aggfunc_by_name` (default None). The parameters to pass to the aggregation function.
masks : str, dict or list of dict or str, optional See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
name : str, optional name : str or None, optional
The name of the marker. If None, will use the class name (default The name of the marker.
None). If None, will use the class name (default None).
""" """
_DEPENDENCIES: ClassVar[Dependencies] = {"nilearn"} _DEPENDENCIES: ClassVar[Dependencies] = {"nilearn"}
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = { _MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"BOLD": { DataType.BOLD: {
"rss_ets": "timeseries", "rss_ets": StorageType.Timeseries,
}, },
} }
def __init__( parcellation: Annotated[str | list[str], BeforeValidator(ensure_list)]
self, agg_method: str = "mean"
parcellation: str | list[str], agg_method_params: dict | None = None
agg_method: str = "mean", masks: Annotated[
agg_method_params: dict | None = None, dict | str | list[dict | str] | None,
masks: str | dict | list[dict | str] | None = None, BeforeValidator(ensure_list_or_none),
name: str | None = None, ] = None
) -> None:
self.parcellation = parcellation
self.agg_method = agg_method
self.agg_method_params = agg_method_params
self.masks = masks
super().__init__(name=name)
def compute( def compute(
self, self,
@ -110,6 +108,7 @@ class RSSETSMarker(BaseMarker):
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
on=DataType.BOLD,
masks=self.masks, masks=self.masks,
).compute(input=input, extra_input=extra_input) ).compute(input=input, extra_input=extra_input)
# Compute edgewise timeseries # Compute edgewise timeseries

View file

@ -1,5 +1,11 @@
__all__ = ["ALFFMaps", "ALFFParcels", "ALFFSpheres"] __all__ = [
"ALFFImpl",
"ALFFMaps",
"ALFFParcels",
"ALFFSpheres",
]
from .falff_base import ALFFImpl
from .falff_maps import ALFFMaps from .falff_maps import ALFFMaps
from .falff_parcels import ALFFParcels from .falff_parcels import ALFFParcels
from .falff_spheres import ALFFSpheres from .falff_spheres import ALFFSpheres

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
import atexit
from functools import lru_cache from functools import lru_cache
from pathlib import Path from pathlib import Path
from typing import ( from typing import (
@ -12,7 +13,7 @@ from typing import (
import nibabel as nib import nibabel as nib
from ...pipeline import WorkDirManager from ...pipeline import ExtDep, WorkDirManager
from ...typing import ExternalDependencies from ...typing import ExternalDependencies
from ...utils import logger, run_ext_cmd from ...utils import logger, run_ext_cmd
from ...utils.singleton import Singleton from ...utils.singleton import Singleton
@ -35,13 +36,15 @@ class AFNIALFF(metaclass=Singleton):
_EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [ _EXT_DEPENDENCIES: ClassVar[ExternalDependencies] = [
{ {
"name": "afni", "name": ExtDep.AFNI,
"commands": ["3dRSFC", "3dAFNItoNIFTI"], "commands": ["3dRSFC", "3dAFNItoNIFTI"],
}, },
] ]
def __del__(self) -> None: def __init__(self) -> None:
"""Terminate the class.""" atexit.register(self._del)
def _del(self) -> None:
# Clear the computation cache # Clear the computation cache
logger.debug("Clearing cache for ALFF computation via AFNI") logger.debug("Clearing cache for ALFF computation via AFNI")
self.compute.cache_clear() self.compute.cache_clear()

View file

@ -6,15 +6,21 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from enum import Enum
from pathlib import Path from pathlib import Path
from typing import ( from typing import (
TYPE_CHECKING, TYPE_CHECKING,
Annotated,
Any, Any,
ClassVar, ClassVar,
) )
from pydantic import BeforeValidator, PositiveFloat
from ...datagrabber import DataType
from ...storage import StorageType
from ...typing import ConditionalDependencies, MarkerInOutMappings from ...typing import ConditionalDependencies, MarkerInOutMappings
from ...utils.logging import logger, raise_error from ...utils import ensure_list_or_none, logger
from ..base import BaseMarker from ..base import BaseMarker
from ._afni_falff import AFNIALFF from ._afni_falff import AFNIALFF
from ._junifer_falff import JuniferALFF from ._junifer_falff import JuniferALFF
@ -24,7 +30,19 @@ if TYPE_CHECKING:
from nibabel.nifti1 import Nifti1Image from nibabel.nifti1 import Nifti1Image
__all__ = ["ALFFBase"] __all__ = ["ALFFBase", "ALFFImpl"]
class ALFFImpl(str, Enum):
"""Accepted ALFF implementations.
* ``junifer`` : ``junifer``'s ALFF
* ``afni`` : AFNI's ``3dRSFC``
"""
junifer = "junifer"
afni = "afni"
class ALFFBase(BaseMarker): class ALFFBase(BaseMarker):
@ -32,22 +50,28 @@ class ALFFBase(BaseMarker):
Parameters Parameters
---------- ----------
highpass : positive float using : :enum:`.ALFFImpl`
Highpass cutoff frequency. highpass : positive float, optional
lowpass : positive float Highpass cutoff frequency (default 0.01).
Lowpass cutoff frequency. lowpass : positive float, optional
using : {"junifer", "afni"} Lowpass cutoff frequency (default 0.1).
Implementation to use for computing ALFF:
* "junifer" : Use ``junifer``'s own ALFF implementation
* "afni" : Use AFNI's ``3dRSFC``
tr : positive float, optional tr : positive float, optional
The Repetition Time of the BOLD data. If None, will extract The repetition time of the BOLD data.
the TR from NIfTI header (default None). If None, will extract the TR from NIfTI header (default None).
name : str, optional agg_method : str, optional
The name of the marker. If None, it will use the class name The aggregation function to use.
(default None). See :func:`.get_aggfunc_by_name` for options
(default "mean").
agg_method_params : dict or None, optional
The parameters to pass to the aggregation function.
See :func:`.get_aggfunc_by_name` for options (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None).
name : str or None, optional
The name of the marker.
If None, it will use the class name (default None).
Notes Notes
----- -----
@ -57,59 +81,36 @@ class ALFFBase(BaseMarker):
reported that some preprocessed data might not have the correct ``tr`` in reported that some preprocessed data might not have the correct ``tr`` in
the NIfTI header. the NIfTI header.
Raises
------
ValueError
If ``highpass`` is not positive or zero or
if ``lowpass`` is not positive or
if ``highpass`` is higher than ``lowpass`` or
if ``using`` is invalid.
""" """
_CONDITIONAL_DEPENDENCIES: ClassVar[ConditionalDependencies] = [ _CONDITIONAL_DEPENDENCIES: ClassVar[ConditionalDependencies] = [
{ {
"using": "afni", "using": ALFFImpl.afni,
"depends_on": AFNIALFF, "depends_on": [AFNIALFF],
}, },
{ {
"using": "junifer", "using": ALFFImpl.junifer,
"depends_on": JuniferALFF, "depends_on": [JuniferALFF],
}, },
] ]
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = { _MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"BOLD": { DataType.BOLD: {
"alff": "vector", "alff": StorageType.Vector,
"falff": "vector", "falff": StorageType.Vector,
}, },
} }
def __init__( using: ALFFImpl
self, highpass: PositiveFloat = 0.01
highpass: float, lowpass: PositiveFloat = 0.1
lowpass: float, tr: PositiveFloat | None = None
using: str, agg_method: str = "mean"
tr: float | None = None, agg_method_params: dict | None = None
name: str | None = None, masks: Annotated[
) -> None: dict | str | list[dict | str] | None,
if highpass < 0: BeforeValidator(ensure_list_or_none),
raise_error("Highpass must be positive or 0") ] = None
if lowpass <= 0:
raise_error("Lowpass must be positive")
if highpass >= lowpass:
raise_error("Highpass must be lower than lowpass")
self.highpass = highpass
self.lowpass = lowpass
# Validate `using` parameter
valid_using = [dep["using"] for dep in self._CONDITIONAL_DEPENDENCIES]
if using not in valid_using:
raise_error(
f"Invalid value for `using`, should be one of: {valid_using}"
)
self.using = using
self.tr = tr
super().__init__(on="BOLD", name=name)
def _compute( def _compute(
self, self,

View file

@ -6,6 +6,7 @@
from typing import Any from typing import Any
from ...api.decorators import register_marker from ...api.decorators import register_marker
from ...datagrabber import DataType
from ...utils import logger from ...utils import logger
from ..maps_aggregation import MapsAggregation from ..maps_aggregation import MapsAggregation
from .falff_base import ALFFBase from .falff_base import ALFFBase
@ -23,27 +24,21 @@ class ALFFMaps(ALFFBase):
maps : str maps : str
The name of the map(s) to use. The name of the map(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
using : {"junifer", "afni"} using : :enum:`.ALFFImpl`
Implementation to use for computing ALFF:
* "junifer" : Use ``junifer``'s own ALFF implementation
* "afni" : Use AFNI's ``3dRSFC``
highpass : positive float, optional highpass : positive float, optional
The highpass cutoff frequency for the bandpass filter. If 0, Highpass cutoff frequency (default 0.01).
it will not apply a highpass filter (default 0.01).
lowpass : positive float, optional lowpass : positive float, optional
The lowpass cutoff frequency for the bandpass filter (default 0.1). Lowpass cutoff frequency (default 0.1).
tr : positive float, optional tr : positive float, optional
The Repetition Time of the BOLD data. If None, will extract The repetition time of the BOLD data.
the TR from NIfTI header (default None). If None, will extract the TR from NIfTI header (default None).
masks : str, dict or list of dict or str, optional masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
name : str, optional name : str or None, optional
The name of the marker. If None, will use the class name (default The name of the marker.
None). If None, will use the class name (default None).
Notes Notes
----- -----
@ -59,26 +54,7 @@ class ALFFMaps(ALFFBase):
""" """
def __init__( maps: str
self,
maps: str,
using: str,
highpass: float = 0.01,
lowpass: float = 0.1,
tr: float | None = None,
masks: str | dict | list[dict | str] | None = None,
name: str | None = None,
) -> None:
# Superclass init first to validate `using` parameter
super().__init__(
highpass=highpass,
lowpass=lowpass,
using=using,
tr=tr,
name=name,
)
self.maps = maps
self.masks = masks
def compute( def compute(
self, self,
@ -132,7 +108,7 @@ class ALFFMaps(ALFFBase):
**MapsAggregation( **MapsAggregation(
maps=self.maps, maps=self.maps,
masks=self.masks, masks=self.masks,
on="BOLD", on=DataType.BOLD,
).compute( ).compute(
input=aggregation_alff_input, input=aggregation_alff_input,
extra_input=extra_input, extra_input=extra_input,
@ -142,7 +118,7 @@ class ALFFMaps(ALFFBase):
**MapsAggregation( **MapsAggregation(
maps=self.maps, maps=self.maps,
masks=self.masks, masks=self.masks,
on="BOLD", on=DataType.BOLD,
).compute( ).compute(
input=aggregation_falff_input, input=aggregation_falff_input,
extra_input=extra_input, extra_input=extra_input,

View file

@ -6,10 +6,13 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any from typing import Annotated, Any
from pydantic import BeforeValidator
from ...api.decorators import register_marker from ...api.decorators import register_marker
from ...utils import logger from ...datagrabber import DataType
from ...utils import ensure_list, logger
from ..parcel_aggregation import ParcelAggregation from ..parcel_aggregation import ParcelAggregation
from .falff_base import ALFFBase from .falff_base import ALFFBase
@ -26,33 +29,27 @@ class ALFFParcels(ALFFBase):
parcellation : str or list of str parcellation : str or list of str
The name(s) of the parcellation(s) to use. The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
using : {"junifer", "afni"} using : :enum:`.ALFFImpl`
Implementation to use for computing ALFF:
* "junifer" : Use ``junifer``'s own ALFF implementation
* "afni" : Use AFNI's ``3dRSFC``
highpass : positive float, optional
The highpass cutoff frequency for the bandpass filter. If 0,
it will not apply a highpass filter (default 0.01).
lowpass : positive float, optional
The lowpass cutoff frequency for the bandpass filter (default 0.1).
tr : positive float, optional
The Repetition Time of the BOLD data. If None, will extract
the TR from NIfTI header (default None).
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options (default "mean").
agg_method_params : dict, optional agg_method_params : dict or None, optional
Parameters to pass to the aggregation function. Check valid options in The parameters to pass to the aggregation function.
:func:`.get_aggfunc_by_name` (default None). See :func:`.get_aggfunc_by_name` for valid options (default None).
masks : str, dict or list of dict or str, optional highpass : positive float, optional
Highpass cutoff frequency (default 0.01).
lowpass : positive float, optional
Lowpass cutoff frequency (default 0.1).
tr : positive float, optional
The repetition time of the BOLD data.
If None, will extract the TR from NIfTI header (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
name : str, optional name : str or None, optional
The name of the marker. If None, will use the class name (default The name of the marker.
None). If None, will use the class name (default None).
Notes Notes
----- -----
@ -68,30 +65,7 @@ class ALFFParcels(ALFFBase):
""" """
def __init__( parcellation: Annotated[str | list[str], BeforeValidator(ensure_list)]
self,
parcellation: str | list[str],
using: str,
highpass: float = 0.01,
lowpass: float = 0.1,
tr: float | None = None,
agg_method: str = "mean",
agg_method_params: dict | None = None,
masks: str | dict | list[dict | str] | None = None,
name: str | None = None,
) -> None:
# Superclass init first to validate `using` parameter
super().__init__(
highpass=highpass,
lowpass=lowpass,
using=using,
tr=tr,
name=name,
)
self.parcellation = parcellation
self.agg_method = agg_method
self.agg_method_params = agg_method_params
self.masks = masks
def compute( def compute(
self, self,
@ -147,7 +121,7 @@ class ALFFParcels(ALFFBase):
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=DataType.BOLD,
).compute( ).compute(
input=aggregation_alff_input, input=aggregation_alff_input,
extra_input=extra_input, extra_input=extra_input,
@ -159,7 +133,7 @@ class ALFFParcels(ALFFBase):
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=DataType.BOLD,
).compute( ).compute(
input=aggregation_falff_input, input=aggregation_falff_input,
extra_input=extra_input, extra_input=extra_input,

View file

@ -6,9 +6,12 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any from typing import Any, Literal
from pydantic import PositiveFloat
from ...api.decorators import register_marker from ...api.decorators import register_marker
from ...datagrabber import DataType
from ...utils import logger from ...utils import logger
from ..sphere_aggregation import SphereAggregation from ..sphere_aggregation import SphereAggregation
from .falff_base import ALFFBase from .falff_base import ALFFBase
@ -26,40 +29,35 @@ class ALFFSpheres(ALFFBase):
coords : str coords : str
The name of the coordinates list to use. The name of the coordinates list to use.
See :func:`.list_data` for options. See :func:`.list_data` for options.
using : {"junifer", "afni"} using : :enum:`.ALFFImpl`
Implementation to use for computing ALFF: radius : ``zero`` or positive float or None, optional
The radius of the sphere in millimetres.
* "junifer" : Use ``junifer``'s own ALFF implementation If None, the signal will be extracted from a single voxel.
* "afni" : Use AFNI's ``3dRSFC`` See :class:`.JuniferNiftiSpheresMasker` for more information
(default None).
radius : float, optional
The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None).
allow_overlap : bool, optional allow_overlap : bool, optional
Whether to allow overlapping spheres. If False, an error is raised if Whether to allow overlapping spheres.
the spheres overlap (default is False). If False, an error is raised if the spheres overlap (default False).
highpass : positive float, optional
The highpass cutoff frequency for the bandpass filter. If 0,
it will not apply a highpass filter (default 0.01).
lowpass : positive float, optional
The lowpass cutoff frequency for the bandpass filter (default 0.1).
tr : positive float, optional
The Repetition Time of the BOLD data. If None, will extract
the TR from NIfTI header (default None).
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The aggregation function to use.
:func:`.get_aggfunc_by_name` (default "mean"). See :func:`.get_aggfunc_by_name` for options (default "mean").
agg_method_params : dict, optional agg_method_params : dict or None, optional
Parameters to pass to the aggregation function. Check valid options in The parameters to pass to the aggregation function.
:func:`.get_aggfunc_by_name`. See :func:`.get_aggfunc_by_name` for valid options (default None).
masks : str, dict or list of dict or str, optional highpass : positive float, optional
Highpass cutoff frequency (default 0.01).
lowpass : positive float, optional
Lowpass cutoff frequency (default 0.1).
tr : positive float, optional
The repetition time of the BOLD data.
If None, will extract the TR from NIfTI header (default None).
masks : str, dict, list of them or None, optional
The specification of the masks to apply to regions before extracting The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details. signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None). If None, will not apply any mask (default None).
name : str, optional name : str or None, optional
The name of the marker. If None, will use the class name (default The name of the marker.
None). If None, will use the class name (default None).
Notes Notes
----- -----
@ -75,34 +73,9 @@ class ALFFSpheres(ALFFBase):
""" """
def __init__( coords: str
self, radius: Literal[0] | PositiveFloat | None = None
coords: str, allow_overlap: bool = False
using: str,
radius: float | None = None,
allow_overlap: bool = False,
highpass: float = 0.01,
lowpass: float = 0.1,
tr: float | None = None,
agg_method: str = "mean",
agg_method_params: dict | None = None,
masks: str | dict | list[dict | str] | None = None,
name: str | None = None,
) -> None:
# Superclass init first to validate `using` parameter
super().__init__(
highpass=highpass,
lowpass=lowpass,
using=using,
tr=tr,
name=name,
)
self.coords = coords
self.radius = radius
self.allow_overlap = allow_overlap
self.agg_method = agg_method
self.agg_method_params = agg_method_params
self.masks = masks
def compute( def compute(
self, self,
@ -160,7 +133,7 @@ class ALFFSpheres(ALFFBase):
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=DataType.BOLD,
).compute( ).compute(
input=aggregation_alff_input, input=aggregation_alff_input,
extra_input=extra_input, extra_input=extra_input,
@ -174,7 +147,7 @@ class ALFFSpheres(ALFFBase):
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=DataType.BOLD,
).compute( ).compute(
input=aggregation_falff_input, input=aggregation_falff_input,
extra_input=extra_input, extra_input=extra_input,

View file

@ -7,11 +7,11 @@ import logging
from pathlib import Path from pathlib import Path
import pytest import pytest
import scipy as sp import scipy.stats as sps
from junifer.datagrabber import PatternDataladDataGrabber from junifer.datagrabber import DataType, PatternDataladDataGrabber
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers import ALFFMaps from junifer.markers import ALFFImpl, ALFFMaps
from junifer.pipeline import WorkDirManager from junifer.pipeline import WorkDirManager
from junifer.pipeline.utils import _check_afni from junifer.pipeline.utils import _check_afni
from junifer.storage import HDF5FeatureStorage from junifer.storage import HDF5FeatureStorage
@ -38,8 +38,8 @@ def test_ALFFMaps_storage_type(feature: str) -> None:
""" """
assert "vector" == ALFFMaps( assert "vector" == ALFFMaps(
maps=MAPS, maps=MAPS,
using="junifer", using=ALFFImpl.junifer,
).storage_type(input_type="BOLD", output_feature=feature) ).storage_type(input_type=DataType.BOLD, output_feature=feature)
def test_ALFFMaps( def test_ALFFMaps(
@ -70,12 +70,12 @@ def test_ALFFMaps(
# Initialize marker # Initialize marker
marker = ALFFMaps( marker = ALFFMaps(
maps=MAPS, maps=MAPS,
using="junifer", using=ALFFImpl.junifer,
) )
# Check correct output # Check correct output
for name in ["alff", "falff"]: for name in ["alff", "falff"]:
assert "vector" == marker.storage_type( assert "vector" == marker.storage_type(
input_type="BOLD", output_feature=name input_type=DataType.BOLD, output_feature=name
) )
# Fit transform marker on data # Fit transform marker on data
@ -99,7 +99,7 @@ def test_ALFFMaps(
# Reset log capture # Reset log capture
caplog.clear() caplog.clear()
# Initialize storage # Initialize storage
storage = HDF5FeatureStorage(tmp_path / "falff_maps.hdf5") storage = HDF5FeatureStorage(uri=tmp_path / "falff_maps.hdf5")
# Fit transform marker on data with storage # Fit transform marker on data with storage
marker.fit_transform( marker.fit_transform(
input=element_data, input=element_data,
@ -141,7 +141,7 @@ def test_ALFFMaps_comparison(
# Initialize marker # Initialize marker
junifer_marker = ALFFMaps( junifer_marker = ALFFMaps(
maps=MAPS, maps=MAPS,
using="junifer", using=ALFFImpl.junifer,
) )
# Fit transform marker on data # Fit transform marker on data
junifer_output = junifer_marker.fit_transform(element_data) junifer_output = junifer_marker.fit_transform(element_data)
@ -151,7 +151,7 @@ def test_ALFFMaps_comparison(
# Initialize marker # Initialize marker
afni_marker = ALFFMaps( afni_marker = ALFFMaps(
maps=MAPS, maps=MAPS,
using="afni", using=ALFFImpl.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)
@ -160,7 +160,7 @@ def test_ALFFMaps_comparison(
for feature in afni_output_bold.keys(): for feature in afni_output_bold.keys():
# Check for Pearson correlation coefficient # Check for Pearson correlation coefficient
r, _ = sp.stats.pearsonr( r, _ = sps.pearsonr(
junifer_output_bold[feature]["data"][0], junifer_output_bold[feature]["data"][0],
afni_output_bold[feature]["data"][0], afni_output_bold[feature]["data"][0],
) )

View file

@ -8,10 +8,11 @@ import logging
from pathlib import Path from pathlib import Path
import pytest import pytest
import scipy as sp import scipy.stats as sps
from junifer.datagrabber import DataType
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.markers.falff import ALFFParcels from junifer.markers import ALFFImpl, ALFFParcels
from junifer.pipeline import WorkDirManager from junifer.pipeline import WorkDirManager
from junifer.pipeline.utils import _check_afni from junifer.pipeline.utils import _check_afni
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
@ -39,8 +40,8 @@ def test_ALFFParcels_storage_type(feature: str) -> None:
""" """
assert "vector" == ALFFParcels( assert "vector" == ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
using="junifer", using=ALFFImpl.junifer,
).storage_type(input_type="BOLD", output_feature=feature) ).storage_type(input_type=DataType.BOLD, output_feature=feature)
def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None: def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@ -63,7 +64,7 @@ def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Initialize marker # Initialize marker
marker = ALFFParcels( marker = ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
using="junifer", using=ALFFImpl.junifer,
) )
# Fit transform marker on data # Fit transform marker on data
output = marker.fit_transform(element_data) output = marker.fit_transform(element_data)
@ -86,7 +87,9 @@ def test_ALFFParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
# Reset log capture # Reset log capture
caplog.clear() caplog.clear()
# Initialize storage # Initialize storage
storage = SQLiteFeatureStorage(tmp_path / "falff_parcels.sqlite") storage = SQLiteFeatureStorage(
uri=tmp_path / "falff_parcels.sqlite"
)
# Fit transform marker on data with storage # Fit transform marker on data with storage
marker.fit_transform( marker.fit_transform(
input=element_data, input=element_data,
@ -116,7 +119,7 @@ def test_ALFFParcels_comparison(tmp_path: Path) -> None:
# Initialize marker # Initialize marker
junifer_marker = ALFFParcels( junifer_marker = ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
using="junifer", using=ALFFImpl.junifer,
) )
# Fit transform marker on data # Fit transform marker on data
junifer_output = junifer_marker.fit_transform(element_data) junifer_output = junifer_marker.fit_transform(element_data)
@ -126,7 +129,7 @@ def test_ALFFParcels_comparison(tmp_path: Path) -> None:
# Initialize marker # Initialize marker
afni_marker = ALFFParcels( afni_marker = ALFFParcels(
parcellation=PARCELLATION, parcellation=PARCELLATION,
using="afni", using=ALFFImpl.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)
@ -135,7 +138,7 @@ def test_ALFFParcels_comparison(tmp_path: Path) -> None:
for feature in afni_output_bold.keys(): for feature in afni_output_bold.keys():
# Check for Pearson correlation coefficient # Check for Pearson correlation coefficient
r, _ = sp.stats.pearsonr( r, _ = sps.pearsonr(
junifer_output_bold[feature]["data"][0], junifer_output_bold[feature]["data"][0],
afni_output_bold[feature]["data"][0], afni_output_bold[feature]["data"][0],
) )

Some files were not shown because too many files have changed in this diff Show more