[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:
:imported-members:
.. automodule:: junifer.onthefly.brainprint
.. automodule:: junifer.onthefly._brainprint
:members:

View file

@ -158,16 +158,16 @@ Available
| (subject-native or other template spaces)
- Done
- 0.0.4
* - ``Smoothing``
* - :class:`.Smoothing`
- | Apply smoothing to data, particularly useful when dealing with
| ``fMRIPrep``-ed data
- In Progress
- :gh:`161`
* - ``TemporalSlicer``
* - :class:`.TemporalSlicer`
- Slice ``BOLD`` data temporally
- | Done
- :gh:`443`
* - ``TemporalFilter``
* - :class:`.TemporalFilter`
- Filter (clean) ``BOLD`` data temporally
- | Done
- :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
"sphinxcontrib.mermaid", # mermaid support
"sphinxcontrib.towncrier.ext", # towncrier fragment support
"sphinxcontrib.autodoc_pydantic", # autodoc support for pydantic models
"enum_tools.autoenum", # enum support
]
if use_multiversion:
@ -97,6 +99,15 @@ nitpick_ignore_regex = [
("py:class", "pipeline.Pipeline"), # nilearn
("py:obj", "neurokit2.*"), # ignore neurokit2
("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 -------------------------------------------------
@ -154,6 +165,7 @@ intersphinx_mapping = {
"pandas": ("https://pandas.pydata.org/pandas-docs/dev", None),
# "sqlalchemy": ("https://docs.sqlalchemy.org/en/20/", None),
"scipy": ("https://docs.scipy.org/doc/scipy/", None),
"pydantic": ("https://docs.pydantic.dev/latest/", None),
}
# -- 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 junifer.datagrabber import PatternDataGrabber
from junifer.datagrabber import PatternDataGrabber, DataType
from junifer.typing import DataGrabberPatterns
class ExampleBIDSDataGrabber(PatternDataGrabber):
def __init__(self, datadir: str | Path) -> None:
types = ["T1w", "BOLD"]
patterns = {
"T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native",
},
"BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
}
replacements = ["subject", "session"]
super().__init__(
datadir=datadir,
types=types,
patterns=patterns,
replacements=replacements,
)
types: list[DataType] = [DataType.T1w, DataType.BOLD]
patterns: DataGrabberPatterns = {
"T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native",
},
"BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
}
replacements: list[str] = ["subject", "session"]
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
@ -175,29 +170,24 @@ use the :func:`.register_datagrabber` decorator.
from junifer.api.decorators import register_datagrabber
from junifer.datagrabber import PatternDataGrabber
from junifer.typing import DataGrabberPatterns
@register_datagrabber
class ExampleBIDSDataGrabber(PatternDataGrabber):
def __init__(self, datadir: str | Path) -> None:
types = ["T1w", "BOLD"]
patterns = {
"T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native",
},
"BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
}
replacements = ["subject", "session"]
super().__init__(
datadir=datadir,
types=types,
patterns=patterns,
replacements=replacements,
)
types: list[DataType] = [DataType.T1w, DataType.BOLD]
patterns: DataGrabberPatterns = {
"T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native",
},
"BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
}
replacements: list[str] = ["subject", "session"]
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
from pathlib import Path
from junifer.api.decorators import register_datagrabber
from junifer.datagrabber import PatternDataladDataGrabber
from pydantic import AnyUrl
@register_datagrabber
class ExampleBIDSDataGrabber(PatternDataladDataGrabber):
def __init__(self) -> None:
types = ["T1w", "BOLD"]
patterns = {
"T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native",
},
"BOLD": {
"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"
rootdir = "example_bids_ses"
super().__init__(
datadir=None,
uri=uri,
rootdir=rootdir,
types=types,
patterns=patterns,
replacements=replacements,
)
uri: AnyUrl = "https://gin.g-node.org/juaml/datalad-example-bids"
types: list[DataType] = ["T1w", "BOLD"]
patterns: DataGrabberPatterns = {
"T1w": {
"pattern": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz",
"space": "native",
},
"BOLD": {
"pattern": "{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
},
}
replacements: list[str] = ["subject", "session"]
rootdir: Path = "example_bids_ses"
This approach can be used directly from the YAML, like so:
@ -376,8 +361,8 @@ need to implement the following methods:
.. note::
The ``__init__`` method could also be implemented, but it is not mandatory.
This is required if the DataGrabber requires any extra parameter.
If the DataGrabber requires any extra parameter, they could be defined as
class attributes.
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:
- ``path``: the path to the confounds file.
- ``format``: the format of the confounds file. Currently, this can be either
``fmriprep`` or ``adhoc``.
- ``format``: the format of the confounds file. Check :enum:`.ConfoundsFormat`
for options.
The ``fmriprep`` format corresponds to the format of the confounds files
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] = [
{
"name": "afni",
"name": ExtDep.AFNI,
"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
``_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
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] = [
{
"using": "fsl",
"depends_on": FSLWarper,
"depends_on": [FSLWarper],
},
{
"using": "ants",
"depends_on": ANTsWarper,
"depends_on": [ANTSWarper],
},
{
"using": "auto",
@ -93,18 +93,16 @@ that it shows the problem a bit better and how we solve it:
},
]
def __init__(
self, using: str, reference: str, on: Union[List[str], str]
) -> None:
# validation and setting up
...
using: str
reference: str
on: List[DataType]
Here, you see a new class attribute ``_CONDITIONAL_DEPENDENCIES`` which is a
list of dictionaries with two keys:
* ``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
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] = [
{
"name": "fsl",
"name": ExtDep.FSL,
"commands": ["flirt", "applywarp"],
},
]

View file

@ -34,4 +34,5 @@ DataGrabbers, Preprocessors, Markers, etc., following the *junifer* way.
plugins
data_registries
data_types
confounds_format
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
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.
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.
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
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:
the input. For ``BOLD``, it will be ``Timeseries``, while for the rest of the
inputs, it will be ``Vector``. Thus, we have a class attribute like so:
.. code-block:: python
@ -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,
# each feature having same or different storage type
_MARKER_INOUT_MAPPINGS = {
"BOLD": {
"parcel_mean": "timeseries",
DataType.BOLD: {
"parcel_mean": StorageType.Timeseries,
},
"VBM_WM": {
"parcel_mean": "vector",
DataType.VBM_WM: {
"parcel_mean": StorageType.Vector,
},
"VBM_GM": {
"parcel_mean": "vector",
DataType.VBM_GM: {
"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
to configure how the Marker will behave.
The parameters of the Marker are defined in the ``__init__`` method. The
:class:`.BaseMarker` class requires two optional parameters:
The parameters of the Marker are defined as class attributes. The
:class:`.BaseMarker` class defines two optional parameters:
1. ``name``: the name of the Marker. This is used to identify the Marker in the
configuration file.
2. ``on``: a list or string with the data types that the Marker will be applied
to.
1. ``name``: the name of the Marker. This is used to identify the Marker in the configuration file.
2. ``on``: a list of :enum:`.DataType` with the data types that the Marker will be applied to.
.. 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.
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
def __init__(
self,
parcellation: str,
on: str | list[str] | None = None,
name: str | None = None,
) -> None:
self.parcellation = parcellation
super().__init__(on=on, name=name)
parcellation: str
.. 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
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``.
.. code-block:: python
from typing import Any
from junifer.data import get_parcellation
from junifer.data import get_data
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"]
# Get the parcellation tailored for the target
t_parcellation, t_labels, _ = get_parcellation(
name=self.parcellation_name,
t_parcellation, t_labels, _ = get_data(
kind="parcellation",
name=[self.parcellation],
target_data=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 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.storage import StorageType
from junifer.typing import Dependencies, MarkerInOutMappings
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"}
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"BOLD": {
"parcel_mean": "timeseries",
DataType.BOLD: {
"parcel_mean": StorageType.Timeseries,
},
"VBM_WM": {
"parcel_mean": "vector",
DataType.VBM_WM: {
"parcel_mean": StorageType.Vector,
},
"VBM_GM": {
"parcel_mean": "vector",
DataType.VBM_GM: {
"parcel_mean": StorageType.Vector,
},
}
def __init__(
self,
parcellation: str,
on: str | list[str] | None = None,
name: str | None = None,
) -> None:
self.parcellation = parcellation
super().__init__(on=on, name=name)
parcellation: str
def compute(
self,
@ -235,8 +222,9 @@ Finally, we need to register the Marker using the ``@register_marker`` decorator
data = input["data"]
# Get the parcellation tailored for the target
t_parcellation, t_labels, _ = get_parcellation(
name=self.parcellation_name,
t_parcellation, t_labels, _ = get_data(
kind="parcellation",
name=[self.parcellation],
target_data=input,
extra_input=extra_input,
)
@ -280,9 +268,13 @@ Template for a custom Marker
# TODO: add the input-output mappings
_MARKER_INOUT_MAPPINGS = {}
def __init__(self, on=None, name=None):
# TODO: add marker-specific parameters
super().__init__(on=on, name=name)
# TODO: define marker-specific parameters
# optional
def validate_marker_params(self):
# TODO: add validation logic for marker parameters
pass
def compute(self, input, extra_input):
# 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:
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
kind: ParcelAggregation
parcellation: my_custom_parcellation
parcellation: <my_custom_parcellation>
method: mean
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
:class:`.BasePreprocessor` and implement a few methods and class attributes:
#. ``__init__``: The initialisation method, where the Preprocessor is
configured.
#. (optional) ``validate_preprocessor_params``: The method to perform logical validation of parameters (if required).
#. ``preprocess``: The method that given the data, preprocesses the data.
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
-----------------------------------
Now we need to define our Preprocessor class' constructor which is also how
you configure it. Our class will have the following arguments:
Now we need to define our Preprocessor class' parameters as class attributes.
Our class will have the following:
1. ``fwhm``: The smoothing strength as a full-width at half maximum
(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
JSON format, and JSON only supports these types.
As :class:`.BasePreprocessor` already defines ``on``, we can define the other:
.. code-block:: python
from typing import Literal
@ -68,15 +69,7 @@ you configure it. Our class will have the following arguments:
...
def __init__(
self,
fwhm: int | float | ArrayLike | Literal["fast"] | None,
on: str | list[str] | None = None,
) -> None:
self.fwhm = fwhm
super().__init__(on=on)
fwhm: int | float | ArrayLike | Literal["fast"] | None
...
@ -165,15 +158,9 @@ decorator and our final code should look like this:
_DEPENDENCIES = {"nilearn"}
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["T1w", "T2w", "BOLD"]
_VALID_DATA_TYPES: ClassVar[Sequence[DataType]] = ["T1w", "T2w", "BOLD"]
def __init__(
self,
fwhm: int | float | ArrayLike | Literal["fast"] | None,
on: str | list[str] | None = None,
) -> None:
self.fwhm = fwhm
super().__init__(on=on)
fwhm: int | float | ArrayLike | Literal["fast"] | None
def preprocess(
self,
@ -191,7 +178,11 @@ Template for a custom Preprocessor
.. code-block:: python
from collections.abc import Sequence
from typing import ClassVar
from junifer.api.decorators import register_preprocessor
from junifer.datagrabber import DataType
from junifer.preprocess import BasePreprocessor
@ -202,11 +193,14 @@ Template for a custom Preprocessor
_DEPENDENCIES = {}
# 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: add preprocessor-specific parameters
super().__init__(on=on)
# TODO: define preprocessor-specific parameters
# optional
def validate_preprocessor_params(self):
# TODO: add validation logic for preprocessor parameters
pass
def preprocess(self, input, extra_input):
# TODO: add the preprocessor logic

View file

@ -265,7 +265,7 @@ Features
^^^^^^^^
- 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`)
- Introduce :class:`junifer.pipeline.PipelineComponentRegistry` to centralise
pipeline component management by `Synchon Mandal`_ (:gh:`362`)

View file

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

View file

@ -10,7 +10,7 @@ Authors: Federico Raimondo
License: BSD 3 clause
"""
from junifer.datagrabber import PatternDataladDataGrabber
from junifer.datagrabber import DataType, PatternDataladDataGrabber
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 specific pattern that matches each type, and the variables that will be
# replaced in the patterns.
types = ["T1w", "BOLD"]
types = [DataType.T1w, DataType.BOLD]
patterns = {
"T1w": {
"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
from junifer.api import collect, run
from junifer.storage.sqlite import SQLiteFeatureStorage
from junifer.storage import SQLiteFeatureStorage
from junifer.utils import configure_logging

View file

@ -5,6 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
import atexit
import os
import shutil
from pathlib import Path
@ -182,6 +183,7 @@ def run(
"elements will be processed"
)
WorkDirManager(**workdir)
atexit.register(WorkDirManager()._cleanup)
# Get datagrabber to use
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 .htcondor_adapter import HTCondorAdapter
from .queue_context_adapter import QueueContextAdapter, EnvKind, EnvShell, QueueContextEnv
from .htcondor_adapter import HTCondorAdapter, HTCondorCollect
from .gnu_parallel_local_adapter import GnuParallelLocalAdapter

View file

@ -6,10 +6,16 @@
import shutil
import textwrap
from pathlib import Path
from typing import Any
from ...typing import Elements
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"]
@ -26,14 +32,14 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
The path to the job directory.
yaml_config_path : pathlib.Path
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.
pre_run : str or None, optional
pre_run_cmds : str or None, optional
Extra shell commands to source before the run (default None).
pre_collect : str or None, optional
Extra bash commands to source before the collect (default None).
env : dict, optional
The Python environment configuration. If None, will run without a
pre_collect_cmds : str or None, optional
Extra shell commands to source before the collect (default None).
env : :class:`.QueueContextEnv` or None, optional
The environment configuration. If None, will run without a
virtual environment of any kind (default None).
verbose : str, optional
The level of verbosity (default "info").
@ -43,12 +49,6 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
submit : bool, optional
Whether to submit the jobs (default False).
Raises
------
ValueError
If ``env.kind`` is invalid or
if ``env.shell`` is invalid.
See Also
--------
QueueContextAdapter :
@ -58,87 +58,44 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
"""
def __init__(
self,
job_name: str,
job_dir: Path,
yaml_config_path: Path,
elements: Elements,
pre_run: str | None = None,
pre_collect: str | None = None,
env: dict[str, str] | None = None,
verbose: str = "info",
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
job_name: str
job_dir: Path
yaml_config_path: Path
elements: Elements
pre_run_cmds: str | None = None
pre_collect_cmds: str | None = None
env: QueueContextEnv | None = None
verbose: str = "info"
verbose_datalad: str | None = None
submit: bool = False
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._run_path = self._job_dir / f"run_{self._job_name}.sh"
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}"
def model_post_init(self, context: Any): # noqa: D102
if self.env is None:
self.env = QueueContextEnv(
kind=EnvKind.Local, shell=EnvShell.Bash, name=""
)
if self.env["kind"] == EnvKind.Local:
# No virtual environment
self._executable = "junifer"
self._arguments = ""
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
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._run_path = self.job_dir / f"run_{self.job_name}.sh"
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 elements(self) -> str:
def elements_to_run(self) -> str:
"""Return elements to run."""
elements_to_run = []
for element in self._elements:
for element in self.elements:
# Stringify elements if tuple for operation
str_element = (
",".join(element) if isinstance(element, tuple) else element
@ -150,23 +107,23 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
def pre_run(self) -> str:
"""Return pre-run commands."""
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"
"# Force datalad to run in non-interactive mode\n"
"DATALAD_UI_INTERACTIVE=false\n"
)
var = self._pre_run or ""
var = self.pre_run_cmds or ""
return fixed + "\n" + var
def run(self) -> str:
"""Return run commands."""
verbose_args = f"--verbose {self._verbose}"
if self._verbose_datalad:
verbose_args = f"--verbose {self.verbose}"
if self.verbose_datalad:
verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad}"
f"{verbose_args} --verbose-datalad {self.verbose_datalad}"
)
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"
"# Run pre_run.sh\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
f"--results {self._log_dir} "
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._yaml_config_path.resolve()!s} "
f"{self.yaml_config_path.resolve()!s} "
f"{verbose_args} "
f"--element"
)
@ -186,28 +143,28 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
def pre_collect(self) -> str:
"""Return pre-collect commands."""
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"
)
var = self._pre_collect or ""
var = self.pre_collect_cmds or ""
return fixed + "\n" + var
def collect(self) -> str:
"""Return collect commands."""
verbose_args = f"--verbose {self._verbose}"
if self._verbose_datalad:
verbose_args = f"--verbose {self.verbose}"
if self.verbose_datalad:
verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad}"
f"{verbose_args} --verbose-datalad {self.verbose_datalad}"
)
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"
"# Run pre_collect.sh\n"
f"sh {self._pre_collect_path.resolve()!s}\n\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._yaml_config_path.resolve()!s} "
f"{self.yaml_config_path.resolve()!s} "
f"{verbose_args}"
)
@ -230,17 +187,19 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
f"{self._elements_file_path.resolve()!s}"
)
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
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.write_text(textwrap.dedent(self.pre_run()))
make_executable(self._pre_run_path)
# Create run
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.write_text(textwrap.dedent(self.run()))
@ -248,14 +207,14 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
# Create pre collect
logger.info(
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.write_text(textwrap.dedent(self.pre_collect()))
make_executable(self._pre_collect_path)
# Create collect
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.write_text(textwrap.dedent(self.collect()))
@ -263,7 +222,7 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
# Submit if required
run_cmd = f"sh {self._run_path.resolve()!s}"
collect_cmd = f"sh {self._collect_path.resolve()!s}"
if self._submit:
if self.submit:
logger.info(
"Shell scripts created, the following will be run:\n"
f"{run_cmd}\n"

View file

@ -5,14 +5,37 @@
import shutil
import textwrap
from enum import Enum
from pathlib import Path
from typing import Any
from ...typing import Elements
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):
@ -26,14 +49,14 @@ class HTCondorAdapter(QueueContextAdapter):
The path to the job directory.
yaml_config_path : pathlib.Path
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.
pre_run : str or None, optional
Extra bash commands to source before the run (default None).
pre_collect : str or None, optional
Extra bash commands to source before the collect (default None).
env : dict, optional
The Python environment configuration. If None, will run without a
pre_run_cmds : str or None, optional
Extra shell commands to source before the run (default None).
pre_collect_cmds : str or None, optional
Extra shell commands to source before the collect (default None).
env : :class:`.QueueContextEnv` or None, optional
The environment configuration. If None, will run without a
virtual environment of any kind (default None).
verbose : str, optional
The level of verbosity (default "info").
@ -48,25 +71,12 @@ class HTCondorAdapter(QueueContextAdapter):
The size of disk (HDD or SSD) to use (default "1G").
extra_preamble : str or None, optional
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").
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
Whether to submit the jobs. In any case, .dag files will be created
for submission (default False).
Raises
------
ValueError
If ``collect`` is invalid or if ``env`` is invalid.
See Also
--------
QueueContextAdapter :
@ -76,144 +86,67 @@ class HTCondorAdapter(QueueContextAdapter):
"""
def __init__(
self,
job_name: str,
job_dir: Path,
yaml_config_path: Path,
elements: Elements,
pre_run: str | None = None,
pre_collect: str | None = None,
env: dict[str, str] | None = None,
verbose: str = "info",
verbose_datalad: str | None = None,
cpus: int = 1,
mem: str = "8G",
disk: str = "1G",
extra_preamble: str | None = None,
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
job_name: str
job_dir: Path
yaml_config_path: Path
elements: Elements
pre_run_cmds: str | None = None
pre_collect_cmds: str | None = None
env: QueueContextEnv | None = None
verbose: str = "info"
verbose_datalad: str | None = None
cpus: int = 1
mem: str = "8G"
disk: str = "1G"
extra_preamble: str | None = None
collect_task: HTCondorCollect = HTCondorCollect.Yes
submit: bool = False
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"
def model_post_init(self, context: Any): # noqa: D102
if self.env is None:
self.env = QueueContextEnv(
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._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"
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
self._dag_path = self.job_dir / f"{self.job_name}.dag"
def pre_run(self) -> str:
"""Return pre-run commands."""
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"
"# Force datalad to run in non-interactive mode\n"
"DATALAD_UI_INTERACTIVE=false\n"
)
var = self._pre_run or ""
var = self.pre_run_cmds or ""
return fixed + "\n" + var
def run(self) -> str:
"""Return run commands."""
verbose_args = f"--verbose {self._verbose} "
if self._verbose_datalad is not None:
verbose_args = f"--verbose {self.verbose} "
if self.verbose_datalad is not None:
verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad} "
f"{verbose_args} --verbose-datalad {self.verbose_datalad} "
)
junifer_run_args = (
"run "
f"{self._yaml_config_path.resolve()!s} "
f"{self.yaml_config_path.resolve()!s} "
f"{verbose_args}"
"--element $(element)"
)
@ -226,11 +159,11 @@ class HTCondorAdapter(QueueContextAdapter):
"universe = vanilla\n"
"getenv = True\n\n"
"# Resources\n"
f"request_cpus = {self._cpus}\n"
f"request_memory = {self._mem}\n"
f"request_disk = {self._disk}\n\n"
f"request_cpus = {self.cpus}\n"
f"request_memory = {self.mem}\n"
f"request_disk = {self.disk}\n\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"transfer_executable = False\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"error = {log_dir_prefix}.err\n"
)
var = self._extra_preamble or ""
var = self.extra_preamble or ""
return fixed + "\n" + var + "\n" + "queue"
def pre_collect(self) -> str:
"""Return pre-collect commands."""
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"
)
var = self._pre_collect or ""
var = self.pre_collect_cmds or ""
# Add commands if collect="yes"
if self._collect == "yes":
if self.collect_task == "yes":
var += 'if [ "${1}" == "4" ]; then\n exit 1\nfi\n'
return fixed + "\n" + var
def collect(self) -> str:
"""Return collect commands."""
verbose_args = f"--verbose {self._verbose} "
if self._verbose_datalad is not None:
verbose_args = f"--verbose {self.verbose} "
if self.verbose_datalad is not None:
verbose_args = (
f"{verbose_args} --verbose-datalad {self._verbose_datalad} "
f"{verbose_args} --verbose-datalad {self.verbose_datalad} "
)
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"
fixed = (
@ -272,11 +205,11 @@ class HTCondorAdapter(QueueContextAdapter):
"universe = vanilla\n"
"getenv = True\n\n"
"# Resources\n"
f"request_cpus = {self._cpus}\n"
f"request_memory = {self._mem}\n"
f"request_disk = {self._disk}\n\n"
f"request_cpus = {self.cpus}\n"
f"request_memory = {self.mem}\n"
f"request_disk = {self.disk}\n\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"
"transfer_executable = False\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"error = {log_dir_prefix}.err\n"
)
var = self._extra_preamble or ""
var = self.extra_preamble or ""
return fixed + "\n" + var + "\n" + "queue"
def dag(self) -> str:
"""Return HTCondor DAG commands."""
fixed = ""
for idx, element in enumerate(self._elements):
for idx, element in enumerate(self.elements):
# Stringify elements if tuple for operation
str_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
)
var = ""
if self._collect == "yes":
if self.collect_task == "yes":
var += (
f"FINAL collect {self._submit_collect_path}\n"
f"SCRIPT PRE collect {self._pre_collect_path.as_posix()} "
"$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 "
for idx, _ in enumerate(self._elements):
for idx, _ in enumerate(self.elements):
var += f"run{idx} "
var += "CHILD collect\n"
@ -325,7 +258,7 @@ class HTCondorAdapter(QueueContextAdapter):
logger.info("Creating HTCondor job")
# Create logs
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)
# Copy executable if not local
@ -340,7 +273,7 @@ class HTCondorAdapter(QueueContextAdapter):
make_executable(self._exec_path)
# Create pre run
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.write_text(textwrap.dedent(self.pre_run()))
@ -348,14 +281,14 @@ class HTCondorAdapter(QueueContextAdapter):
# Create run
logger.debug(
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.write_text(textwrap.dedent(self.run()))
# Create pre collect
logger.info(
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.write_text(textwrap.dedent(self.pre_collect()))
@ -363,13 +296,13 @@ class HTCondorAdapter(QueueContextAdapter):
# Create collect
logger.debug(
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.write_text(textwrap.dedent(self.collect()))
# Create DAG
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.write_text(textwrap.dedent(self.dag()))
@ -379,7 +312,7 @@ class HTCondorAdapter(QueueContextAdapter):
"-include_env HOME",
f"{self._dag_path.resolve()!s}",
]
if self._submit:
if self.submit:
run_ext_cmd(name="condor_submit_dag", cmd=condor_submit_dag_cmd)
else:
logger.info(

View file

@ -3,22 +3,66 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# 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 enum import Enum
from pydantic import BaseModel, ConfigDict
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.
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.
"""
model_config = ConfigDict(
use_enum_values=True,
extra="allow",
)
@abstractmethod
def pre_run(self) -> str:
"""Return pre-run commands."""

View file

@ -11,30 +11,6 @@ import pytest
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(
"elements, expected_text",
[
@ -62,7 +38,7 @@ def test_GnuParallelLocalAdapter_elements(
yaml_config_path=Path("."),
elements=elements,
)
assert expected_text in adapter.elements()
assert expected_text in adapter.elements_to_run()
@pytest.mark.parametrize(
@ -97,7 +73,7 @@ def test_GnuParallelLocalAdapter_pre_run(
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "conda", "name": "junifer", "shell": shell},
pre_run=pre_run,
pre_run_cmds=pre_run,
)
assert shell in adapter.pre_run()
assert expected_text in adapter.pre_run()
@ -135,7 +111,7 @@ def test_GnuParallelLocalAdapter_pre_collect(
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "venv", "name": "junifer", "shell": shell},
pre_collect=pre_collect,
pre_collect_cmds=pre_collect,
)
assert shell 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
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(
"pre_run, expected_text, shell",
[
@ -79,7 +43,7 @@ def test_HTCondorAdapter_pre_run(
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "conda", "name": "junifer", "shell": shell},
pre_run=pre_run,
pre_run_cmds=pre_run,
)
assert shell in adapter.pre_run()
assert expected_text in adapter.pre_run()
@ -124,8 +88,8 @@ def test_HTCondorAdapter_pre_collect(
yaml_config_path=Path("."),
elements=["sub01"],
env={"kind": "venv", "name": "junifer", "shell": shell},
pre_collect=pre_collect,
collect=collect,
pre_collect_cmds=pre_collect,
collect_task=collect,
)
assert shell in adapter.pre_collect()
assert expected_text in adapter.pre_collect()
@ -199,7 +163,7 @@ def test_HTCondor_dag(
job_dir=Path("."),
yaml_config_path=Path("."),
elements=elements,
collect=collect,
collect_task=collect,
)
assert expected_text in adapter.dag()

View file

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

View file

@ -17,4 +17,5 @@ queue:
env:
kind: conda
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

View file

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

View file

@ -2,3 +2,8 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# 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",
"JuselessDataladCamCANVBM",
"JuselessDataladIXIVBM",
"IXISite",
"JuselessUCLA",
"UCLATask",
"JuselessDataladUKBVBM",
]
from .aomic_id1000_vbm import JuselessDataladAOMICID1000VBM
from .camcan_vbm import JuselessDataladCamCANVBM
from .ixi_vbm import JuselessDataladIXIVBM
from .ucla import JuselessUCLA
from .ixi_vbm import JuselessDataladIXIVBM, IXISite
from .ucla import JuselessUCLA, UCLATask
from .ukb_vbm import JuselessDataladUKBVBM

View file

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

View file

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

View file

@ -5,14 +5,29 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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 ....datagrabber import PatternDataladDataGrabber
from ....utils import raise_error
from ....datagrabber import DataType, PatternDataladDataGrabber
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
@ -23,53 +38,31 @@ class JuselessDataladIXIVBM(PatternDataladDataGrabber):
Parameters
----------
datadir : str or pathlib.Path or None, optional
The directory where the datalad dataset will be cloned. If None,
the datalad dataset will be cloned into a temporary directory
(default None).
sites : {"Guys", "HH", "IOP"} or list of the options or None, optional
Which sites to access data from. If None, all available sites are
selected (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.
sites : {"Guys", "HH", "IOP"} or list of the options, optional
IXI sites.
By default, all available sites are selected.
"""
def __init__(
self,
datadir: str | Path | None = None,
sites: str | list[str] | None = None,
) -> None:
uri = (
"ria+http://cat_12.5.ds.inm7.de"
"#b7107c52-8408-11ea-89c6-a0369f287950"
)
types = ["VBM_GM"]
replacements = ["site", "subject"]
patterns = {
"VBM_GM": {
"pattern": ("{site}/{subject}/mri/m0wp1{subject}.nii.gz"),
"space": "IXI549Space",
},
}
# 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,
)
uri: AnyUrl = AnyUrl(
"ria+http://cat_12.5.ds.inm7.de#b7107c52-8408-11ea-89c6-a0369f287950"
)
types: list[Literal[DataType.VBM_GM]] = [DataType.VBM_GM] # noqa: RUF012
sites: ClassVar[
Annotated[IXISite | list[IXISite], BeforeValidator(ensure_list)]
] = [
IXISite.Guys,
IXISite.HH,
IXISite.IOP,
]
patterns: DataGrabberPatterns = { # noqa: RUF012
"VBM_GM": {
"pattern": ("{site}/{subject}/mri/m0wp1{subject}.nii.gz"),
"space": "IXI549Space",
},
}
replacements: list[str] = ["site", "subject"] # noqa: RUF012

View file

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

View file

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

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

View file

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

View file

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

View file

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

View file

@ -3,30 +3,62 @@ __all__ = [
"DataladDataGrabber",
"PatternDataGrabber",
"PatternDataladDataGrabber",
"AOMICSpace",
"AOMICTask",
"DataladAOMICID1000",
"DataladAOMICPIOP1",
"DataladAOMICPIOP2",
"HCP1200",
"HCP1200Task",
"HCP1200PhaseEncoding",
"DataladHCP1200",
"MultipleDataGrabber",
"DMCC13Benchmark",
"DMCCSession",
"DMCCTask",
"DMCCPhaseEncoding",
"DMCCRun",
"DataTypeManager",
"DataTypeSchema",
"OptionalTypeSchema",
"PatternValidationMixin",
"register_data_type",
"DataType",
"ConfoundsFormat",
"register_confounds_format",
]
# 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 .pattern import PatternDataGrabber
from .pattern import (
PatternDataGrabber,
ConfoundsFormat,
register_confounds_format,
)
from .pattern_datalad import PatternDataladDataGrabber
from .aomic import DataladAOMICID1000, DataladAOMICPIOP1, DataladAOMICPIOP2
from .hcp1200 import HCP1200, DataladHCP1200
from .aomic import (
AOMICSpace,
AOMICTask,
DataladAOMICID1000,
DataladAOMICPIOP1,
DataladAOMICPIOP2,
)
from .hcp1200 import (
HCP1200,
HCP1200Task,
HCP1200PhaseEncoding,
DataladHCP1200,
)
from .multiple import MultipleDataGrabber
from .dmcc13_benchmark import DMCC13Benchmark
from .dmcc13_benchmark import (
DMCC13Benchmark,
DMCCSession,
DMCCTask,
DMCCPhaseEncoding,
DMCCRun,
)
from .pattern_validation_mixin import (
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 .piop1 import DataladAOMICPIOP1
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>
# License: AGPL
from pathlib import Path
from typing import Annotated, Literal
from pydantic import AnyUrl, BeforeValidator
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 ._types import AOMICSpace
__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
class DataladAOMICID1000(PatternDataladDataGrabber):
@ -23,212 +40,206 @@ class DataladAOMICID1000(PatternDataladDataGrabber):
Parameters
----------
datadir : str or Path or None, optional
The directory where the datalad dataset will be cloned. If None,
the datalad dataset will be cloned into a temporary directory
(default None).
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
"FreeSurfer"} or list of the options, optional
AOMIC data types. If None, all available data types are selected.
(default None).
types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
"FreeSurfer", "Warp"} or list of the options, optional
The data type(s) to grab.
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.
space : {"native", "MNI152NLin2009cAsym"}, optional
The space to use for the data (default "MNI152NLin2009cAsym").
Raises
------
ValueError
If invalid value is passed for:
* ``space``
AOMIC space (default ``AOMICSpace.MNI152NLin2009cAsym``).
"""
def __init__(
self,
datadir: str | Path | None = None,
types: str | list[str] | None = None,
space: str = "MNI152NLin2009cAsym",
) -> None:
valid_spaces = ["native", "MNI152NLin2009cAsym"]
if space not in ["native", "MNI152NLin2009cAsym"]:
raise_error(
f"Invalid space {space}. Must be one of {valid_spaces}"
)
# 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": {
uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds003097.git")
types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.DWI,
DataType.FreeSurfer,
DataType.Warp,
]
space: AOMICSpace = AOMICSpace.MNI152NLin2009cAsym
patterns: DataGrabberPatterns = { # noqa: RUF012
"BOLD": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
"{sp_func_desc}"
"desc-preproc_bold.nii.gz"
),
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-moviewatching_"
f"{sp_func_desc}"
"desc-preproc_bold.nii.gz"
"{sp_func_desc}"
"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": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"desc-preproc_T1w.nii.gz"
"{sp_anat_desc}"
"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": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"label-CSF_probseg.nii.gz"
"{subject}_from-MNI152NLin2009cAsym_to-T1w_"
"mode-image_xfm.h5"
),
"space": space,
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
"VBM_GM": {
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"label-GM_probseg.nii.gz"
"{subject}_from-T1w_to-MNI152NLin2009cAsym_"
"mode-image_xfm.h5"
),
"space": space,
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
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"
],
}
replacements: list[str] = ["subject"] # noqa: RUF012
confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
else:
patterns["BOLD"]["prewarp_space"] = "native"
# 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"]
uri = "https://github.com/OpenNeuroDatasets/ds003097.git"
super().__init__(
types=types,
datadir=datadir,
uri=uri,
patterns=patterns,
replacements=replacements,
confounds_format="fmriprep",
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
# Descriptor for space in `anat`
sp_anat_desc = (
"" if self.space == "native" else "space-MNI152NLin2009cAsym_"
)
# 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
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 ...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 ._types import AOMICSpace, AOMICTask
__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
class DataladAOMICPIOP1(PatternDataladDataGrabber):
@ -24,246 +50,224 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber):
Parameters
----------
datadir : str or Path or None, optional
The directory where the datalad dataset will be cloned. If None,
the datalad dataset will be cloned into a temporary directory
(default None).
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
"FreeSurfer"} or list of the options, optional
AOMIC data types. If None, all available data types are selected.
(default None).
types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
"FreeSurfer", "Warp"} or list of the options, optional
The data type(s) to grab.
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.
tasks : {"restingstate", "anticipation", "emomatching", "faces", \
"gstroop", "workingmemory"} or list of the options, optional
AOMIC PIOP1 task sessions. If None, all available task sessions are
selected (default None).
AOMIC PIOP1 task sessions.
By default, all available task sessions are selected.
space : {"native", "MNI152NLin2009cAsym"}, optional
The space to use for the data (default "MNI152NLin2009cAsym").
Raises
------
ValueError
If invalid value is passed for:
* ``tasks``
* ``space``
AOMIC space (default ``AOMICSpace.MNI152NLin2009cAsym``).
"""
def __init__(
self,
datadir: str | Path | None = None,
types: str | list[str] | None = None,
tasks: str | list[str] | None = None,
space: str = "MNI152NLin2009cAsym",
) -> None:
valid_spaces = ["native", "MNI152NLin2009cAsym"]
if space not in ["native", "MNI152NLin2009cAsym"]:
raise_error(
f"Invalid space {space}. Must be one of {valid_spaces}"
)
# Declare all tasks
all_tasks = [
"restingstate",
"anticipation",
"emomatching",
"faces",
"gstroop",
"workingmemory",
]
# Set default tasks
if tasks is None:
tasks = all_tasks
else:
# Convert single task into list
if isinstance(tasks, str):
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 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": {
uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds002785")
types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.DWI,
DataType.FreeSurfer,
DataType.Warp,
]
tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
AOMICTask.RestingState,
AOMICTask.Anticipation,
AOMICTask.EmoMatching,
AOMICTask.Faces,
AOMICTask.Gstroop,
AOMICTask.WorkingMemory,
]
space: AOMICSpace = AOMICSpace.MNI152NLin2009cAsym
patterns: DataGrabberPatterns = { # noqa: RUF012
"BOLD": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"{sp_func_desc}"
"desc-preproc_bold.nii.gz"
),
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
f"{sp_func_desc}"
"desc-preproc_bold.nii.gz"
"{sp_func_desc}"
"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": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"desc-preproc_T1w.nii.gz"
"{sp_anat_desc}"
"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": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"label-CSF_probseg.nii.gz"
"{subject}_from-MNI152NLin2009cAsym_to-T1w_"
"mode-image_xfm.h5"
),
"space": space,
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
"VBM_GM": {
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"label-GM_probseg.nii.gz"
"{subject}_from-T1w_to-MNI152NLin2009cAsym_"
"mode-image_xfm.h5"
),
"space": space,
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
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",
},
],
}
],
}
replacements: list[str] = ["subject", "task"] # noqa: RUF012
confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
if space == "native":
patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym"
else:
patterns["BOLD"]["prewarp_space"] = "native"
# 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",
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
# Descriptor for space in `anat`
sp_anat_desc = (
"" if self.space == "native" else "space-MNI152NLin2009cAsym_"
)
# 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:
"""Index one element in the dataset.
"""Get the specified item from the dataset.
Parameters
----------

View file

@ -8,15 +8,39 @@
# License: AGPL
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 ...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 ._types import AOMICSpace, AOMICTask
__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
class DataladAOMICPIOP2(PatternDataladDataGrabber):
@ -24,241 +48,219 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber):
Parameters
----------
datadir : str or Path or None, optional
The directory where the datalad dataset will be cloned. If None,
the datalad dataset will be cloned into a temporary directory
(default None).
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
"FreeSurfer"} or list of the options, optional
AOMIC data types. If None, all available data types are selected.
(default None).
types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "DWI", \
"FreeSurfer", "Warp"} or list of the options, optional
The data type(s) to grab.
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.
tasks : {"restingstate", "stopsignal", "workingmemory", "emomatching"} or \
list of the options, optional
AOMIC PIOP2 task sessions. If None, all available task sessions are
selected (default None).
AOMIC PIOP2 task sessions.
By default, all available task sessions are selected.
space : {"native", "MNI152NLin2009cAsym"}, optional
The space to use for the data (default "MNI152NLin2009cAsym").
Raises
------
ValueError
If invalid value is passed for:
* ``tasks``
* ``space``
AOMIC space (default ``AOMICSpace.MNI152NLin2009cAsym``).
"""
def __init__(
self,
datadir: str | Path | None = None,
types: str | list[str] | None = None,
tasks: str | list[str] | None = None,
space: str = "MNI152NLin2009cAsym",
) -> None:
valid_spaces = ["native", "MNI152NLin2009cAsym"]
if space not in ["native", "MNI152NLin2009cAsym"]:
raise_error(
f"Invalid space {space}. Must be one of {valid_spaces}"
)
# Declare all tasks
all_tasks = [
"restingstate",
"stopsignal",
"workingmemory",
"emomatching",
]
# Set default tasks
if tasks is None:
tasks = all_tasks
else:
# Convert single task into list
if isinstance(tasks, str):
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 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": {
uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds002790")
types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.DWI,
DataType.FreeSurfer,
DataType.Warp,
]
tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
AOMICTask.RestingState,
AOMICTask.StopSignal,
AOMICTask.EmoMatching,
AOMICTask.WorkingMemory,
]
space: AOMICSpace = AOMICSpace.MNI152NLin2009cAsym
patterns: DataGrabberPatterns = { # noqa: RUF012
"BOLD": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
"{sp_func_desc}"
"desc-preproc_bold.nii.gz"
),
"mask": {
"pattern": (
"derivatives/fmriprep/{subject}/func/"
"{subject}_task-{task}_"
f"{sp_func_desc}"
"desc-preproc_bold.nii.gz"
"{sp_func_desc}"
"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": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"desc-preproc_T1w.nii.gz"
"{sp_anat_desc}"
"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": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"label-CSF_probseg.nii.gz"
"{subject}_from-MNI152NLin2009cAsym_to-T1w_"
"mode-image_xfm.h5"
),
"space": space,
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
"VBM_GM": {
{
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
f"{sp_anat_desc}"
"label-GM_probseg.nii.gz"
"{subject}_from-T1w_to-MNI152NLin2009cAsym_"
"mode-image_xfm.h5"
),
"space": space,
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep/{subject}/anat/"
"{subject}_"
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",
},
],
}
],
}
replacements: list[str] = ["subject", "task"] # noqa: RUF012
confounds_format: ConfoundsFormat = ConfoundsFormat.FMRIPrep
if space == "native":
patterns["BOLD"]["prewarp_space"] = "MNI152NLin2009cAsym"
else:
patterns["BOLD"]["prewarp_space"] = "native"
# 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",
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
# Descriptor for space in `anat`
sp_anat_desc = (
"" if self.space == "native" else "space-MNI152NLin2009cAsym_"
)
# 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:
"""Implement fetching list of elements in the dataset.
@ -277,7 +279,7 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber):
return elems
def get_item(self, subject: str, task: str) -> dict:
"""Index one element in the dataset.
"""Get the specified item from the dataset.
Parameters
----------

View file

@ -8,31 +8,32 @@
# License: AGPL
import pytest
from pydantic import AnyUrl
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(
"type_, nested_types, space",
[
("BOLD", ["confounds", "mask", "reference"], "MNI152NLin2009cAsym"),
("BOLD", ["confounds", "mask", "reference"], "native"),
(["BOLD"], ["confounds", "mask", "reference"], "native"),
("T1w", ["mask"], "MNI152NLin2009cAsym"),
("T1w", ["mask"], "native"),
(["T1w"], ["mask"], "native"),
("VBM_CSF", None, "MNI152NLin2009cAsym"),
("VBM_CSF", None, "native"),
(["VBM_CSF"], None, "native"),
("VBM_GM", None, "MNI152NLin2009cAsym"),
("VBM_GM", None, "native"),
(["VBM_GM"], None, "native"),
("VBM_WM", None, "MNI152NLin2009cAsym"),
("DWI", None, "MNI152NLin2009cAsym"),
("FreeSurfer", None, "MNI152NLin2009cAsym"),
(["DWI"], None, "MNI152NLin2009cAsym"),
(["FreeSurfer"], None, "MNI152NLin2009cAsym"),
],
)
def test_DataladAOMICID1000(
type_: str,
type_: str | list[str],
nested_types: list[str] | None,
space: str,
) -> None:
@ -40,7 +41,7 @@ def test_DataladAOMICID1000(
Parameters
----------
type_ : str
type_ : str or list of str
The parametrized type.
nested_types : list of str or None
The parametrized nested types.
@ -48,32 +49,29 @@ def test_DataladAOMICID1000(
The parametrized space.
"""
dg = DataladAOMICID1000(types=type_, space=space)
# Set URI to Gin
dg.uri = URI
dg = DataladAOMICID1000(uri=URI, types=type_, space=space)
with dg:
# Get all elements
all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0]
# Get test element data
out = dg[test_element]
# Assert data type
assert type_ in out
assert out[type_]["path"].exists()
assert out[type_]["path"].is_file()
# Asserts data type metadata
assert "meta" in out[type_]
meta = out[type_]["meta"]
assert "element" in meta
assert "subject" in meta["element"]
assert test_element == meta["element"]["subject"]
# Assert nested data type if not None
if nested_types is not None:
for nested_type in nested_types:
assert out[type_][nested_type]["path"].exists()
assert out[type_][nested_type]["path"].is_file()
if isinstance(type_, str):
type_ = [type_]
for t in type_:
assert t in out
assert out[t]["path"].exists()
assert out[t]["path"].is_file()
# Asserts data type metadata
assert "meta" in out[t]
meta = out[t]["meta"]
assert "element" in meta
assert "subject" in meta["element"]
assert test_element == meta["element"]["subject"]
# Assert nested data type if not None
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(
@ -102,28 +100,13 @@ def test_DataladAOMICID1000_partial_data_access(
The parametrized types.
"""
dg = DataladAOMICID1000(types=types)
# Set URI to Gin
dg.uri = URI
dg = DataladAOMICID1000(uri=URI, types=types)
with dg:
# Get all elements
all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0]
# Get test element data
out = dg[test_element]
# Assert data type
if isinstance(types, list):
for type_ in types:
assert type_ in out
else:
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")
if isinstance(types, str):
types = [types]
for t in types:
assert t in out

View file

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

View file

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

View file

@ -7,56 +7,77 @@
from abc import ABC, abstractmethod
from collections.abc import Iterator
from enum import Enum
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 ..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):
"""Abstract base class for DataGrabber.
class DataType(str, AEnum):
"""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.
Parameters
----------
types : list of str
The types of data to be grabbed.
datadir : str or pathlib.Path
The directory where the data is / will be stored.
Raises
------
TypeError
If ``types`` is not a list or if the values are not string.
types : :enum:`.DataType` or list of variants
The data type(s) to grab.
datadir : pathlib.Path
The path where the data is or will be stored.
"""
def __init__(self, types: list[str], datadir: str | Path) -> None:
# 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
model_config = ConfigDict(use_enum_values=True)
# Convert str to Path
if not isinstance(datadir, Path):
datadir = Path(datadir)
self._datadir = datadir
types: Annotated[
DataType | list[DataType],
Field(frozen=True),
BeforeValidator(ensure_list),
]
datadir: Path
def model_post_init(self, context: Any): # noqa: D102
logger.debug("Initializing BaseDataGrabber")
logger.debug(f"\t_datadir = {datadir}")
logger.debug(f"\ttypes = {types}")
logger.debug(f"\tdatadir = {self.datadir}")
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.
Yields
@ -72,7 +93,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
Parameters
----------
element : str or tuple of str
element : `Element`
The element to be indexed.
Returns
@ -82,10 +103,14 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
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}")
# Convert element to tuple if not already
if not isinstance(element, tuple):
element = (element,)
# Zip through element keys and actual values to construct element
# access dictionary
named_element: dict = dict(
@ -120,29 +145,30 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
Returns
-------
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
def datadir(self) -> Path:
"""Get data directory path.
def fulldir(self) -> Path:
"""Get complete data directory path.
Returns
-------
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:
"""Filter elements to be grabbed.
Parameters
----------
selection : list
selection : ``Elements``
The list of partial or complete element selectors to filter using.
Yields
@ -157,7 +183,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
Parameters
----------
element : str or tuple of str
element : ``Elements``
The element to be filtered.
Returns

View file

@ -9,12 +9,15 @@ import atexit
import os
import tempfile
from pathlib import Path
from typing import Any, NoReturn
import datalad
import datalad.api as dl
from datalad.support.exceptions import IncompleteResultsError
from datalad.support.gitrepo import GitRepo
from pydantic import AnyUrl, Field, field_validator
from ..api.decorators import register_datagrabber
from ..pipeline import WorkDirManager
from ..typing import Element
from ..utils import config, logger, raise_error, warn_with_log
@ -24,6 +27,26 @@ from .base import BaseDataGrabber
__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):
"""Abstract base class for datalad-based data fetching.
@ -31,17 +54,15 @@ class DataladDataGrabber(BaseDataGrabber):
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
(default ".").
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).
uri : str or None, optional
URI of the datalad sibling (default None).
**kwargs
Keyword arguments passed to superclass.
(default Path(".")).
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.
Methods
-------
@ -66,24 +87,62 @@ class DataladDataGrabber(BaseDataGrabber):
This class is intended to be used as a superclass of a subclass
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__(
self,
rootdir: str | Path = ".",
datadir: str | Path | None = None,
uri: str | None = None,
**kwargs,
):
if datadir is None:
logger.info("`datadir` is None, creating a temporary directory")
# Create temporary directory
tmpdir = WorkDirManager().get_tempdir(prefix="datalad")
self._tmpdir = tmpdir
datadir = tmpdir / "datadir"
datadir.mkdir(parents=True, exist_ok=False)
logger.info(f"`datadir` set to {datadir}")
cache_dir = tmpdir / ".datalad_cache"
uri: AnyUrl = Field(frozen=True)
rootdir: Path = Field(frozen=True, default=Path("."))
datadir: Path = Field(default_factory=lambda: _create_datadir())
_repodir: Path = Path(".")
# Flag to indicate if the dataset was cloned before and it might be
# dirty
datalad_dirty: bool = False
datalad_commit_id: str | None = None
datalad_id: str | None = None
_dataset: dl.Dataset | None = None
_got_files: list[str] = [] # noqa: RUF012
_was_cloned: bool = False
@field_validator("datadir", mode="after")
@classmethod
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"
locks_dir = cache_dir / "locks"
sockets_dir.mkdir(parents=True, exist_ok=False)
@ -105,43 +164,29 @@ class DataladDataGrabber(BaseDataGrabber):
"Datalad locks set to "
f"{datalad.cfg.get('datalad.locations.locks')}"
)
atexit.register(self._rmtmpdir)
# TODO: uri can be converted to a positional argument
if uri is None:
raise_error("`uri` must be provided")
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
atexit.register(_remove_datadir, self.datadir)
else:
self._repodir = self.datadir
super().validate_datagrabber_params()
def __del__(self) -> None:
"""Destructor."""
if hasattr(self, "_tmpdir"):
self._rmtmpdir()
def _rmtmpdir(self) -> None:
"""Remove temporary directory if it exists."""
if self._tmpdir.exists():
logger.debug("Removing temporary directory")
WorkDirManager().delete_tempdir(self._tmpdir)
if self.datadir.stem.startswith(
"datalad"
) and self.datadir.stem.endswith("juniferauto"):
_remove_datadir(self.datadir)
@property
def datadir(self) -> Path:
"""Get data directory path.
def fulldir(self) -> Path:
"""Get complete data directory path.
Returns
-------
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]:
"""Get the dataset ID from the remote.
@ -164,8 +209,10 @@ class DataladDataGrabber(BaseDataGrabber):
with tempfile.TemporaryDirectory() as tmpdir:
if not config.get("datagrabber.skipidcheck", False):
logger.debug(f"Querying {self.uri} for dataset ID")
repo = GitRepo.clone(
self.uri, path=tmpdir, clone_options=["-n", "--depth=1"]
repo: GitRepo = GitRepo.clone(
str(self.uri),
path=tmpdir,
clone_options=["-n", "--depth=1"],
)
repo.checkout(name=".datalad/config", options=["HEAD"])
remote_id = repo.config.get("datalad.dataset.id", None)
@ -178,10 +225,11 @@ class DataladDataGrabber(BaseDataGrabber):
is_dirty = False
else:
logger.debug("Skipping dataset ID check")
# Should be already set to the dataset
remote_id = self._dataset.id
is_dirty = False
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:
raise_error("Could not get dataset ID from remote")
@ -251,7 +299,7 @@ class DataladDataGrabber(BaseDataGrabber):
return out
def install(self) -> None:
"""Install the datalad dataset into the ``datadir``.
"""Installs the datalad dataset.
Raises
------
@ -261,12 +309,10 @@ class DataladDataGrabber(BaseDataGrabber):
If there is a datalad-related problem while cloning dataset.
"""
isinstalled = dl.Dataset(self._datadir).is_installed()
if isinstalled:
is_installed = dl.Dataset(self._repodir).is_installed()
if is_installed:
logger.debug("Dataset already installed")
self._got_files = []
self._dataset: dl.Dataset = dl.Dataset(self._datadir)
self._dataset = dl.Dataset(self._repodir)
# Check if dataset is already installed with a different ID
remote_id, is_dirty = self._get_dataset_id_remote()
if remote_id != self._dataset.id:
@ -274,7 +320,6 @@ class DataladDataGrabber(BaseDataGrabber):
"Dataset already installed but with a different "
f"ID: {self._dataset.id} (local) != {remote_id} (remote)"
)
# Conditional reporting on dataset dirtiness
self.datalad_dirty = is_dirty
if self.datalad_dirty:
@ -286,18 +331,18 @@ class DataladDataGrabber(BaseDataGrabber):
logger.debug(f"Dataset (id: {self._dataset.id}) is clean")
else:
logger.debug(f"Installing dataset {self.uri} to {self._datadir}")
logger.debug(f"Installing dataset {self.uri} to {self._repodir}")
try:
self._dataset: dl.Dataset = dl.clone( # type: ignore
self.uri, self._datadir, result_renderer="disabled"
self._dataset = dl.clone(
self.uri, self._repodir, result_renderer="disabled"
)
except IncompleteResultsError as e:
raise_error(f"Failed to clone dataset: {e.failed}")
logger.debug("Dataset installed")
self._was_cloned = not isinstalled
self.datalad_commit_id = self._dataset.repo.get_hexsha( # type: ignore
self._dataset.repo.get_corresponding_branch() # type: ignore
self._was_cloned = not is_installed
# Dataset should be set already
self.datalad_commit_id = self._dataset.repo.get_hexsha(
self._dataset.repo.get_corresponding_branch()
)
self.datalad_id = self._dataset.id
@ -320,7 +365,7 @@ class DataladDataGrabber(BaseDataGrabber):
Parameters
----------
element : str or tuple of str
element : `Element`
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,
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>
# License: AGPL
from enum import Enum
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 ..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
__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
@ -20,191 +98,142 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
Parameters
----------
datadir : str or Path or None, optional
The directory where the datalad dataset will be cloned. If None,
the datalad dataset will be cloned into a temporary directory
(default None).
types: {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM"} or \
list of the options, optional
DMCC data types. If None, all available data types are selected.
(default None).
sessions: {"ses-wave1bas", "ses-wave1pro", "ses-wave1rea"} or \
list of the options, optional
DMCC sessions. If None, all available sessions are selected
(default None).
tasks: {"Rest", "Axcpt", "Cuedts", "Stern", "Stroop"} or \
list of the options, optional
DMCC task sessions. If None, all available task sessions are selected
(default None).
types : {"BOLD", "T1w", "VBM_CSF", "VBM_GM", "VBM_WM", "Warp"} or \
list of the options, optional
The data type(s) to grab.
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.
sessions : {"ses-wave1bas", "ses-wave1pro", "ses-wave1rea"} or \
list of the options, optional
DMCC sessions.
By default, all available sessions are selected.
tasks : {"Rest", "Axcpt", "Cuedts", "Stern", "Stroop"} or \
list of the options, optional
DMCC tasks.
By default, all available tasks are selected.
phase_encodings : {"AP", "PA"} or list of the options, optional
DMCC phase encoding directions. If None, all available phase encodings
are selected (default None).
DMCC phase encoding directions.
By default, all available phase encodings are selected.
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
Whether to use T1w in native space (default False).
Raises
------
ValueError
If invalid value is passed for:
* ``sessions``
* ``tasks``
* ``phase_encodings``
* ``runs``
"""
def __init__(
self,
datadir: str | Path | None = None,
types: str | list[str] | None = None,
sessions: str | list[str] | None = None,
tasks: str | list[str] | None = None,
phase_encodings: str | list[str] | None = None,
runs: str | list[str] | None = None,
native_t1w: bool = False,
) -> None:
# Declare all sessions
all_sessions = [
"ses-wave1bas",
"ses-wave1pro",
"ses-wave1rea",
]
# Set default sessions
if sessions is None:
sessions = all_sessions
else:
# Convert single session into list
if isinstance(sessions, str):
sessions = [sessions]
# Verify valid sessions
for s in sessions:
if s not in all_sessions:
raise_error(
f"{s} is not a valid session in the DMCC dataset"
)
self.sessions = sessions
# Declare all tasks
all_tasks = [
"Rest",
"Axcpt",
"Cuedts",
"Stern",
"Stroop",
]
# Set default tasks
if tasks is None:
tasks = all_tasks
else:
# Convert single task into list
if isinstance(tasks, str):
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": {
uri: AnyUrl = AnyUrl("https://github.com/OpenNeuroDatasets/ds003452.git")
types: Annotated[_types | list[_types], BeforeValidator(ensure_list)] = [ # noqa: RUF012
DataType.BOLD,
DataType.T1w,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
]
sessions: Annotated[
_sessions | list[_sessions], BeforeValidator(ensure_list)
] = [ # noqa: RUF012
DMCCSession.Wave1Bas,
DMCCSession.Wave1Pro,
DMCCSession.Wave1Rea,
]
tasks: Annotated[_tasks | list[_tasks], BeforeValidator(ensure_list)] = [ # noqa: RUF012
DMCCTask.Rest,
DMCCTask.Axcpt,
DMCCTask.Cuedts,
DMCCTask.Stern,
DMCCTask.Stroop,
]
phase_encodings: Annotated[
_phase_encodings | list[_phase_encodings],
BeforeValidator(ensure_list),
] = [ # noqa: RUF012
DMCCPhaseEncoding.AP,
DMCCPhaseEncoding.PA,
]
runs: Annotated[_runs | list[_runs], BeforeValidator(ensure_list)] = [ # noqa: RUF012
DMCCRun.One,
DMCCRun.Two,
]
native_t1w: bool = False
patterns: DataGrabberPatterns = { # noqa: RUF012
"BOLD": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/{session}/"
"func/{subject}_{session}_task-{task}_acq-mb4"
"{phase_encoding}_run-{run}_"
"space-MNI152NLin2009cAsym_desc-preproc_bold.nii.gz"
),
"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-preproc_bold.nii.gz"
"space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz"
),
"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": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_desc-preproc_T1w.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"
"{subject}_space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
"VBM_GM": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_label-GM_probseg.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_label-WM_probseg.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
}
# Use native T1w assets
self.native_t1w = False
if native_t1w:
self.native_t1w = True
patterns.update(
},
"VBM_CSF": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_label-CSF_probseg.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
"VBM_GM": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_label-GM_probseg.nii.gz"
),
"space": "MNI152NLin2009cAsym",
},
"VBM_WM": {
"pattern": (
"derivatives/fmriprep-1.3.2/{subject}/anat/"
"{subject}_space-MNI152NLin2009cAsym_label-WM_probseg.nii.gz"
),
"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": {
"pattern": (
@ -244,24 +273,8 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
],
}
)
# 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", "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",
)
self.types.append(DataType.Warp)
super().validate_datagrabber_params()
def get_item(
self,
@ -271,7 +284,7 @@ class DMCC13Benchmark(PatternDataladDataGrabber):
phase_encoding: str,
run: str,
) -> dict:
"""Index one element in the dataset.
"""Get the specified item from the dataset.
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

View file

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

View file

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

View file

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

View file

@ -5,10 +5,17 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from pathlib import Path
from typing import Annotated
from pydantic import BeforeValidator, ConfigDict
from ..api.decorators import register_datagrabber
from ..typing import DataGrabberLike
from ..utils import deep_update, raise_error
from .base import BaseDataGrabber
from ..typing import DataGrabberLike, Element
from ..utils import deep_update, ensure_list, raise_error
from .base import BaseDataGrabber, DataType
from .pattern import PatternDataGrabber
from .pattern_datalad import PatternDataladDataGrabber
__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 for same element keys
first_keys = datagrabbers[0].get_element_keys()
for dg in datagrabbers[1:]:
first_keys = self.datagrabbers[0].get_element_keys()
for dg in self.datagrabbers[1:]:
if dg.get_element_keys() != first_keys:
raise_error(
msg="DataGrabbers have different element keys",
klass=RuntimeError,
)
# 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 all(hasattr(dg, "patterns") for dg in datagrabbers):
first_patterns = datagrabbers[0].patterns
for dg in datagrabbers[1:]:
if all(hasattr(dg, "patterns") for dg in self.datagrabbers):
first_patterns = self.datagrabbers[0].patterns
for dg in self.datagrabbers[1:]:
for data_type in set(types):
dtype_pattern = dg.patterns.get(data_type)
if dtype_pattern is None:
@ -75,14 +93,13 @@ class MultipleDataGrabber(BaseDataGrabber):
msg="DataGrabbers have overlapping types",
klass=RuntimeError,
)
self._datagrabbers = datagrabbers
def __getitem__(self, element: str | tuple) -> dict:
def __getitem__(self, element: Element) -> dict:
"""Implement indexing.
Parameters
----------
element : str or tuple
element : `Element`
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,
each item in the tuple is the value for the replacement string
@ -98,7 +115,7 @@ class MultipleDataGrabber(BaseDataGrabber):
out = {}
metas = []
for dg in self._datagrabbers:
for dg in self.datagrabbers:
t_out = dg[element]
deep_update(out, t_out)
# Now get the meta for this datagrabber
@ -119,16 +136,15 @@ class MultipleDataGrabber(BaseDataGrabber):
def __enter__(self) -> "MultipleDataGrabber":
"""Implement context entry."""
for dg in self._datagrabbers:
for dg in self.datagrabbers:
dg.__enter__()
return self
def __exit__(self, exc_type, exc_value, exc_traceback) -> None:
"""Implement context exit."""
for dg in self._datagrabbers:
for dg in self.datagrabbers:
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]:
"""Get types.
@ -138,7 +154,7 @@ class MultipleDataGrabber(BaseDataGrabber):
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
def get_element_keys(self) -> list[str]:
@ -153,7 +169,7 @@ class MultipleDataGrabber(BaseDataGrabber):
The element keys.
"""
return self._datagrabbers[0].get_element_keys()
return self.datagrabbers[0].get_element_keys()
def get_elements(self) -> list:
"""Get elements.
@ -167,7 +183,7 @@ class MultipleDataGrabber(BaseDataGrabber):
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])
for s in all_elements[1:]:
elements.intersection_update(s)

View file

@ -10,6 +10,9 @@ from copy import deepcopy
from pathlib import Path
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 ..typing import DataGrabberPatterns, Elements
@ -18,11 +21,32 @@ from .base import BaseDataGrabber
from .pattern_validation_mixin import PatternValidationMixin
__all__ = ["PatternDataGrabber"]
__all__ = [
"ConfoundsFormat",
"PatternDataGrabber",
"register_confounds_format",
]
# Accepted formats for confounds specification
_CONFOUNDS_FORMATS = ("fmriprep", "adhoc")
class ConfoundsFormat(str, AEnum):
"""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
@ -33,125 +57,16 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
Parameters
----------
types : list of str
The types of data to be grabbed.
patterns : dict
Data type patterns as a dictionary. It has the following schema:
* ``"T1w"`` :
.. code-block:: none
{
"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
types : :enum:`.DataType` or list of variants
The data type(s) to grab.
datadir : pathlib.Path
The path where the data is stored.
patterns : ``DataGrabberPatterns``
The datagrabber patterns. Check :class:`.DataTypeSchema` for the \
schema.
replacements : list of str
All possible replacements in ``patterns.<data_type>.pattern``.
confounds_format : :enum:`.ConfoundsFormat` or None, optional
The format of the confounds for the dataset (default None).
partial_pattern_ok : bool, optional
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`
(default True).
Raises
------
ValueError
If ``confounds_format`` is invalid.
Attributes
----------
skip_file_check
"""
def __init__(
self,
types: list[str],
patterns: DataGrabberPatterns,
replacements: list[str] | str,
datadir: str | Path,
confounds_format: str | None = None,
partial_pattern_ok: bool = False,
) -> None:
# Convert replacements to list if not already
if not isinstance(replacements, list):
replacements = [replacements]
patterns: DataGrabberPatterns = Field(frozen=True)
replacements: list[str] = Field(frozen=True)
confounds_format: ConfoundsFormat | None = Field(None, frozen=True)
partial_pattern_ok: bool = Field(False, frozen=True)
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
# Validate patterns
self.validate_patterns(
types=types,
replacements=replacements,
patterns=patterns,
partial_pattern_ok=partial_pattern_ok,
types=self.types,
replacements=self.replacements,
patterns=self.patterns,
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(f"\tpatterns = {patterns}")
logger.debug(f"\treplacements = {replacements}")
logger.debug(f"\tconfounds_format = {confounds_format}")
logger.debug(f"\tpatterns = {self.patterns}")
logger.debug(f"\treplacements = {self.replacements}")
logger.debug(f"\tconfounds_format = {self.confounds_format}")
@property
def skip_file_check(self) -> bool:
@ -324,7 +217,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
resolved_pattern = self._replace_patterns_glob(element, pattern)
# Resolve path for wildcard
if "*" in resolved_pattern:
t_matches = list(self.datadir.absolute().glob(resolved_pattern))
t_matches = list(self.fulldir.absolute().glob(resolved_pattern))
# Multiple matches
if len(t_matches) > 1:
raise_error(
@ -340,7 +233,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
)
path = t_matches[0]
else:
path = self.datadir / resolved_pattern
path = self.fulldir / resolved_pattern
if not self.skip_file_check:
if not path.exists() and not path.is_symlink():
raise_error(
@ -367,7 +260,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
return self.replacements
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
replacing the ``patterns`` with actual values passed via ``**element``.
@ -514,8 +407,8 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
glob_pattern,
t_replacements,
) = self._replace_patterns_regex(pattern)
for fname in self.datadir.glob(glob_pattern):
suffix = fname.relative_to(self.datadir).as_posix()
for fname in self.fulldir.glob(glob_pattern):
suffix = fname.relative_to(self.fulldir).as_posix()
m = re.match(re_pattern, suffix)
if m is not None:
# Find the groups of replacements present in the

View file

@ -5,6 +5,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from pydantic import ConfigDict
from ..api.decorators import register_datagrabber
from ..utils import logger
from .datalad_base import DataladDataGrabber
@ -23,122 +25,24 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber):
Parameters
----------
types : list of str
The types of data to be grabbed.
patterns : dict
Data type patterns as a dictionary. It has the following schema:
* ``"T1w"`` :
.. code-block:: none
{
"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
uri : pydantic.AnyUrl
URI of the datalad sibling.
types : enum:`.DataType` or list of variants
The data type(s) to grab.
patterns : ``DataGrabberPatterns``
The datagrabber patterns. Check :class:`DataTypeSchema` for the schema.
replacements : list of str
All possible replacements in ``patterns.<data_type>.pattern``.
rootdir : pathlib.Path, optional
The path within the datalad dataset to the root directory
(default ".").
uri : str or None, optional
URI of the datalad sibling (default None).
(default Path(".")).
confounds_format : :enum:`.ConfoundsFormat` or None, optional
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
--------
@ -149,15 +53,11 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber):
"""
def __init__(
self,
**kwargs,
) -> None:
# TODO(synchon): needs to be reworked, DataladDataGrabber needs to be
# a mixin to avoid multiple inheritance wherever possible.
model_config = ConfigDict(extra="allow")
def validate_datagrabber_params(self) -> None:
"""Run extra logical validation for datagrabber."""
super().validate_datagrabber_params()
logger.debug("Initializing PatternDataladDataGrabber")
for key, val in kwargs.items():
for key, val in self.__pydantic_extra__.items():
logger.debug(f"\t{key} = {val}")
super().__init__(**kwargs)

View file

@ -3,9 +3,20 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from collections.abc import Iterator, MutableMapping
from typing import TypedDict
import sys
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 ..utils import logger, raise_error, warn_with_log
@ -191,37 +202,17 @@ def register_data_type(name: str, schema: DataTypeSchema) -> None:
----------
name : str
The data type name.
schema : DataTypeSchema
schema : ``DataTypeSchema``
The data type schema.
"""
DataTypeManager()[name] = schema
extend_enum(DataType, name, name)
class PatternValidationMixin:
"""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(
self,
replacements: list[str],
@ -234,15 +225,13 @@ class PatternValidationMixin:
----------
replacements : list of str
The replacements to validate.
patterns : dict
The patterns to validate replacements against.
patterns : ``DataGrabberPatterns``
The patterns to validate ``replacements`` against.
partial_pattern_ok : bool
Whether to raise error if partial pattern for a data type is found.
Raises
------
TypeError
If ``replacements`` is not a list or if the values are not string.
ValueError
If a value in ``replacements`` is not part of a data type pattern
and ``partial_pattern_ok=False`` or
@ -256,15 +245,6 @@ class PatternValidationMixin:
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
all_patterns = []
for dtype_val in patterns.values():
@ -390,7 +370,7 @@ class PatternValidationMixin:
def validate_patterns(
self,
types: list[str],
types: list[DataType],
replacements: list[str],
patterns: DataGrabberPatterns,
partial_pattern_ok: bool = False,
@ -399,11 +379,11 @@ class PatternValidationMixin:
Parameters
----------
types : list of str
The data types to check patterns of.
types : list of :enum:`.DataType`
The data type(s) to check patterns of.
replacements : list of str
The replacements to be replaced in the patterns.
patterns : dict
The replacements to be replaced in the ``patterns``.
patterns : ``DataGrabberPatterns``
The patterns to validate.
partial_pattern_ok : bool, optional
Whether to raise error if partial pattern for a data type is found.
@ -412,8 +392,6 @@ class PatternValidationMixin:
Raises
------
TypeError
If ``patterns`` is not a dictionary.
ValueError
If length of ``types`` and ``patterns`` are different or
if ``patterns`` is missing entries from ``types`` or
@ -421,12 +399,6 @@ class PatternValidationMixin:
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
if len(types) > len(patterns):
raise_error(

View file

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

View file

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

View file

@ -4,101 +4,95 @@
# License: AGPL
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(
"sessions, tasks, phase_encodings, runs, native_t1w",
[
(None, None, None, None, False),
("ses-wave1bas", "Rest", "AP", "1", False),
("ses-wave1bas", "Axcpt", "AP", "1", False),
("ses-wave1bas", "Cuedts", "AP", "1", False),
("ses-wave1bas", "Stern", "AP", "1", False),
("ses-wave1bas", "Stroop", "AP", "1", False),
("ses-wave1bas", "Rest", "PA", "2", False),
("ses-wave1bas", "Axcpt", "PA", "2", False),
("ses-wave1bas", "Cuedts", "PA", "2", False),
("ses-wave1bas", "Stern", "PA", "2", False),
("ses-wave1bas", "Stroop", "PA", "2", False),
("ses-wave1bas", "Rest", "AP", "1", True),
("ses-wave1bas", "Axcpt", "AP", "1", True),
("ses-wave1bas", "Cuedts", "AP", "1", True),
("ses-wave1bas", "Stern", "AP", "1", True),
("ses-wave1bas", "Stroop", "AP", "1", True),
("ses-wave1bas", "Rest", "PA", "2", True),
("ses-wave1bas", "Axcpt", "PA", "2", True),
("ses-wave1bas", "Cuedts", "PA", "2", True),
("ses-wave1bas", "Stern", "PA", "2", True),
("ses-wave1bas", "Stroop", "PA", "2", True),
("ses-wave1pro", "Rest", "AP", "1", False),
("ses-wave1pro", "Rest", "PA", "2", False),
("ses-wave1pro", "Rest", "AP", "1", True),
("ses-wave1pro", "Rest", "PA", "2", True),
("ses-wave1rea", "Rest", "AP", "1", False),
("ses-wave1rea", "Rest", "PA", "2", False),
("ses-wave1rea", "Rest", "AP", "1", True),
("ses-wave1rea", "Rest", "PA", "2", True),
(["ses-wave1bas"], ["Rest"], ["AP"], ["1"], False),
(["ses-wave1bas"], ["Axcpt"], ["AP"], ["1"], False),
(["ses-wave1bas"], ["Cuedts"], ["AP"], ["1"], False),
(["ses-wave1bas"], ["Stern"], ["AP"], ["1"], False),
(["ses-wave1bas"], ["Stroop"], ["AP"], ["1"], False),
(["ses-wave1bas"], ["Rest"], ["PA"], ["2"], False),
(["ses-wave1bas"], ["Axcpt"], ["PA"], ["2"], False),
(["ses-wave1bas"], ["Cuedts"], ["PA"], ["2"], False),
(["ses-wave1bas"], ["Stern"], ["PA"], ["2"], False),
(["ses-wave1bas"], ["Stroop"], ["PA"], ["2"], False),
(["ses-wave1bas"], ["Rest"], ["AP"], ["1"], True),
(["ses-wave1bas"], ["Axcpt"], ["AP"], ["1"], True),
(["ses-wave1bas"], ["Cuedts"], ["AP"], ["1"], True),
(["ses-wave1bas"], ["Stern"], ["AP"], ["1"], True),
(["ses-wave1bas"], ["Stroop"], ["AP"], ["1"], True),
(["ses-wave1bas"], ["Rest"], ["PA"], ["2"], True),
(["ses-wave1bas"], ["Axcpt"], ["PA"], ["2"], True),
(["ses-wave1bas"], ["Cuedts"], ["PA"], ["2"], True),
(["ses-wave1bas"], ["Stern"], ["PA"], ["2"], True),
(["ses-wave1bas"], ["Stroop"], ["PA"], ["2"], True),
(["ses-wave1pro"], ["Rest"], ["AP"], ["1"], False),
(["ses-wave1pro"], ["Rest"], ["PA"], ["2"], False),
(["ses-wave1pro"], ["Rest"], ["AP"], ["1"], True),
(["ses-wave1pro"], ["Rest"], ["PA"], ["2"], True),
(["ses-wave1rea"], ["Rest"], ["AP"], ["1"], False),
(["ses-wave1rea"], ["Rest"], ["PA"], ["2"], False),
(["ses-wave1rea"], ["Rest"], ["AP"], ["1"], True),
(["ses-wave1rea"], ["Rest"], ["PA"], ["2"], True),
],
)
def test_DMCC13Benchmark(
sessions: str | None,
tasks: str | None,
phase_encodings: str | None,
runs: str | None,
sessions: list[str],
tasks: list[str],
phase_encodings: list[str],
runs: list[str],
native_t1w: bool,
) -> None:
"""Test DMCC13Benchmark DataGrabber.
Parameters
----------
sessions : str or None
sessions : list of str
The parametrized session values.
tasks : str or None
tasks : list of str
The parametrized task values.
phase_encodings : str or None
phase_encodings : list of str
The parametrized phase encoding values.
runs : str or None
runs : list of str
The parametrized run values.
native_t1w : bool
The parametrized values for fetching native T1w.
"""
dg = DMCC13Benchmark(
uri=URI,
sessions=sessions,
tasks=tasks,
phase_encodings=phase_encodings,
runs=runs,
native_t1w=native_t1w,
)
# Set URI to Gin
dg.uri = URI
with dg:
# Get all elements
all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0]
# Get test element's access values
_, ses, task, phase, run = test_element
# Access data
out = dg[("sub-01", ses, task, phase, run)]
# Available data types
data_types = [
"BOLD",
"VBM_CSF",
"VBM_GM",
"VBM_WM",
"T1w",
DataType.BOLD,
DataType.VBM_CSF,
DataType.VBM_GM,
DataType.VBM_WM,
DataType.T1w,
]
# Add Warp if native T1w is accessed
if native_t1w:
data_types.append("Warp")
data_types.append(DataType.Warp)
# Data type file name formats
data_file_names = [
@ -131,9 +125,9 @@ def test_DMCC13Benchmark(
data_types, data_file_names, strict=False
):
# Assert data type
assert data_type in out
assert data_type in out.keys()
# Conditional for Warp
if data_type == "Warp":
if data_type is DataType.Warp:
for idx, fname in enumerate(data_file_name):
# Assert data file path exists
assert out[data_type][idx]["path"].exists()
@ -200,15 +194,15 @@ def test_DMCC13Benchmark(
@pytest.mark.parametrize(
"types, native_t1w",
[
("BOLD", True),
(["BOLD"], True),
("BOLD", False),
("T1w", True),
(["T1w"], True),
("T1w", False),
("VBM_CSF", True),
(["VBM_CSF"], True),
("VBM_CSF", False),
("VBM_GM", True),
(["VBM_GM"], True),
("VBM_GM", False),
("VBM_WM", True),
(["VBM_WM"], True),
("VBM_WM", False),
(["BOLD", "VBM_CSF"], True),
(["BOLD", "VBM_CSF"], False),
@ -232,66 +226,18 @@ def test_DMCC13Benchmark_partial_data_access(
The parametrized values for fetching native T1w.
"""
dg = DMCC13Benchmark(types=types, native_t1w=native_t1w)
# Set URI to Gin
dg.uri = URI
dg = DMCC13Benchmark(
uri=URI,
types=types,
native_t1w=native_t1w,
)
with dg:
# Get all elements
all_elements = dg.get_elements()
# Get test element
test_element = all_elements[0]
# Get test element's access values
_, ses, task, phase, run = test_element
# Access data
out = dg[("sub-01", ses, task, phase, run)]
# Assert data type
if isinstance(types, list):
for type_ in types:
assert type_ in out
else:
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")
if isinstance(types, str):
types = [types]
for type_ in types:
assert type_ in out

View file

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

View file

@ -10,7 +10,21 @@ from pathlib import Path
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:
@ -117,7 +131,7 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
"""
datagrabber_first = PatternDataGrabber(
datadir="/tmp/data",
datadir=Path("/tmp/data"),
types=["BOLD", "T1w"],
patterns={
"BOLD": {
@ -129,7 +143,7 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
"space": "native",
},
},
replacements="subject",
replacements=["subject"],
)
assert datagrabber_first.datadir == Path("/tmp/data")
assert set(datagrabber_first.types) == {"T1w", "BOLD"}
@ -181,7 +195,7 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
datagrabber_third = PatternDataGrabber(
datadir=tmpdir,
types=["T1w"],
types="T1w",
patterns={
"T1w": {
"pattern": "anat/{subject}_{session}.nii",
@ -258,7 +272,7 @@ def test_PatternDataGrabber_unix_path_expansion(tmp_path: Path) -> None:
# Create datagrabber
dg = PatternDataGrabber(
datadir=tmp_path,
types=["FreeSurfer"],
types="FreeSurfer",
patterns={
"FreeSurfer": {
"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
assert set(out["FreeSurfer"].keys()) == {"path", "aseg", "meta"}
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
import pytest
from pydantic import AnyUrl
from junifer.datagrabber import PatternDataladDataGrabber
from junifer.datagrabber import DataType, PatternDataladDataGrabber
_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:
"""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_uri = _testing_dataset["example_bids"]["uri"]
with PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri,
types=types,
patterns=patterns,
replacements=replacements,
uri=AnyUrl(repo_uri),
types=[DataType.T1w, DataType.BOLD],
patterns={
"T1w": {
"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:
subs = list(dg)
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]
assert "path" in t_sub["T1w"]
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 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"]
@ -88,7 +68,7 @@ def test_bids_PatternDataladDataGrabber() -> None:
assert "class" in dg_meta
assert dg_meta["class"] == "PatternDataladDataGrabber"
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 dg_meta["datalad_commit_id"] == repo_commit
@ -98,80 +78,64 @@ def test_bids_PatternDataladDataGrabber() -> None:
def test_bids_PatternDataladDataGrabber_datadir() -> None:
"""Test PatternDataladDataGrabber with a datadir set to a relative path."""
# Define patterns
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
datadir = Path("dataset") # use string and not absolute path
with PatternDataladDataGrabber(
uri=_testing_dataset["example_bids"]["uri"],
types=["T1w", "BOLD"],
patterns=patterns,
datadir=datadir,
rootdir="example_bids",
uri=AnyUrl(_testing_dataset["example_bids"]["uri"]),
types=[DataType.T1w, DataType.BOLD],
patterns={
"T1w": {
"pattern": "{subject}/anat/{subject}_T*w.nii.gz",
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_*.nii.gz",
"space": "MNI152NLin6Asym",
},
},
replacements=["subject"],
datadir=datadir,
rootdir=Path("example_bids"),
) as dg:
assert dg.datadir == Path(datadir) / "example_bids"
assert dg.fulldir == Path(datadir) / "example_bids"
for elem in dg:
t_sub = dg[elem]
assert "path" in t_sub["T1w"]
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 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():
"""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
repo_uri = _testing_dataset["example_bids_ses"]["uri"]
rootdir = "example_bids_ses"
repo_uri = AnyUrl(_testing_dataset["example_bids_ses"]["uri"])
rootdir = Path("example_bids_ses")
replacements = ["subject", "session"]
# With T1W and bold, only 2 sessions are available
with PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri,
types=types,
patterns=patterns,
types=[DataType.T1w, DataType.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=replacements,
rootdir=rootdir,
) as dg:
subs = list(dg.get_elements())
expected_subs = [
@ -182,21 +146,19 @@ def test_bids_PatternDataladDataGrabber_session():
assert set(subs) == set(expected_subs)
# 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(
rootdir=rootdir,
uri=repo_uri,
types=types,
patterns=patterns,
types=DataType.T1w,
patterns={
"T1w": {
"pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
),
"space": "MNI152NLin6Asym",
},
},
replacements=replacements,
rootdir=rootdir,
) as dg:
subs = list(dg)
expected_subs = [

View file

@ -8,7 +8,8 @@ from contextlib import AbstractContextManager, nullcontext
import pytest
from junifer.datagrabber.pattern_validation_mixin import (
from junifer.datagrabber import (
DataType,
DataTypeManager,
DataTypeSchema,
PatternValidationMixin,
@ -79,7 +80,7 @@ def test_dtype_mgr(dtype: DataTypeSchema) -> None:
Parameters
----------
dtype : DataTypeSchema
dtype : ``DataTypeSchema``
The parametrized schema.
"""
@ -110,33 +111,15 @@ def test_register_data_type() -> None:
)
assert "dtype" in DataTypeManager()
assert "dtype" in list(DataType)
_ = DataTypeManager().pop("dtype")
assert "dumb" not in DataTypeManager()
assert "dtype" not in DataTypeManager()
assert "dtype" in list(DataType)
@pytest.mark.parametrize(
"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"],
"",
@ -204,30 +187,6 @@ def test_register_data_type() -> None:
},
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"],
["subject", "session"],

View file

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

View file

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

View file

@ -2,3 +2,8 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# 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",
"EdgeCentricFCParcels",
"EdgeCentricFCSpheres",
"ReHoImpl",
"ReHoMaps",
"ReHoParcels",
"ReHoSpheres",
"ALFFImpl",
"ALFFMaps",
"ALFFParcels",
"ALFFSpheres",
@ -37,8 +39,8 @@ from .functional_connectivity import (
EdgeCentricFCParcels,
EdgeCentricFCSpheres,
)
from .reho import ReHoMaps, ReHoParcels, ReHoSpheres
from .falff import ALFFMaps, ALFFParcels, ALFFSpheres
from .reho import ReHoImpl, ReHoMaps, ReHoParcels, ReHoSpheres
from .falff import ALFFImpl, ALFFMaps, ALFFParcels, ALFFSpheres
from .temporal_snr import (
TemporalSNRMaps,
TemporalSNRParcels,

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -6,15 +6,21 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from enum import Enum
from pathlib import Path
from typing import (
TYPE_CHECKING,
Annotated,
Any,
ClassVar,
)
from pydantic import BeforeValidator, PositiveFloat
from ...datagrabber import DataType
from ...storage import StorageType
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 ._afni_falff import AFNIALFF
from ._junifer_falff import JuniferALFF
@ -24,7 +30,19 @@ if TYPE_CHECKING:
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):
@ -32,22 +50,28 @@ class ALFFBase(BaseMarker):
Parameters
----------
highpass : positive float
Highpass cutoff frequency.
lowpass : positive float
Lowpass cutoff frequency.
using : {"junifer", "afni"}
Implementation to use for computing ALFF:
* "junifer" : Use ``junifer``'s own ALFF implementation
* "afni" : Use AFNI's ``3dRSFC``
using : :enum:`.ALFFImpl`
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).
name : str, optional
The name of the marker. If None, it will use the class name
(default None).
The repetition time of the BOLD data.
If None, will extract the TR from NIfTI header (default None).
agg_method : str, optional
The aggregation function to use.
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
-----
@ -57,59 +81,36 @@ class ALFFBase(BaseMarker):
reported that some preprocessed data might not have the correct ``tr`` in
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] = [
{
"using": "afni",
"depends_on": AFNIALFF,
"using": ALFFImpl.afni,
"depends_on": [AFNIALFF],
},
{
"using": "junifer",
"depends_on": JuniferALFF,
"using": ALFFImpl.junifer,
"depends_on": [JuniferALFF],
},
]
_MARKER_INOUT_MAPPINGS: ClassVar[MarkerInOutMappings] = {
"BOLD": {
"alff": "vector",
"falff": "vector",
DataType.BOLD: {
"alff": StorageType.Vector,
"falff": StorageType.Vector,
},
}
def __init__(
self,
highpass: float,
lowpass: float,
using: str,
tr: float | None = None,
name: str | None = None,
) -> None:
if highpass < 0:
raise_error("Highpass must be positive or 0")
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)
using: ALFFImpl
highpass: PositiveFloat = 0.01
lowpass: PositiveFloat = 0.1
tr: PositiveFloat | None = None
agg_method: str = "mean"
agg_method_params: dict | None = None
masks: Annotated[
dict | str | list[dict | str] | None,
BeforeValidator(ensure_list_or_none),
] = None
def _compute(
self,

View file

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

View file

@ -6,10 +6,13 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Any
from typing import Annotated, Any
from pydantic import BeforeValidator
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 .falff_base import ALFFBase
@ -26,33 +29,27 @@ class ALFFParcels(ALFFBase):
parcellation : str or list of str
The name(s) of the parcellation(s) to use.
See :func:`.list_data` for options.
using : {"junifer", "afni"}
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).
using : :enum:`.ALFFImpl`
agg_method : str, optional
The method to perform aggregation using. Check valid options in
:func:`.get_aggfunc_by_name` (default "mean").
agg_method_params : dict, optional
Parameters to pass to the aggregation function. Check valid options in
:func:`.get_aggfunc_by_name` (default None).
masks : str, dict or list of dict or str, optional
The aggregation function to use.
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 valid options (default None).
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
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None).
name : str, optional
The name of the marker. If None, will use the class name (default
None).
name : str or None, optional
The name of the marker.
If None, will use the class name (default None).
Notes
-----
@ -68,30 +65,7 @@ class ALFFParcels(ALFFBase):
"""
def __init__(
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
parcellation: Annotated[str | list[str], BeforeValidator(ensure_list)]
def compute(
self,
@ -147,7 +121,7 @@ class ALFFParcels(ALFFBase):
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
on=DataType.BOLD,
).compute(
input=aggregation_alff_input,
extra_input=extra_input,
@ -159,7 +133,7 @@ class ALFFParcels(ALFFBase):
method=self.agg_method,
method_params=self.agg_method_params,
masks=self.masks,
on="BOLD",
on=DataType.BOLD,
).compute(
input=aggregation_falff_input,
extra_input=extra_input,

View file

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

View file

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

View file

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