update: improve marker and storage interfaces #149

Merged
fraimondo merged 28 commits from fix/meta into main 2022-11-29 12:08:23 +00:00
73 changed files with 1626 additions and 1337 deletions

View file

@ -54,7 +54,7 @@ Contributions are welcome and greatly appreciated. Please read the [guidelines](
junifer is released under the AGPL v3 license: junifer is released under the AGPL v3 license:
julearn, FZJuelich AML neuroimaging feature extraction library. junifer, FZJuelich AML neuroimaging feature extraction library.
Copyright (C) 2022, authors of junifer. Copyright (C) 2022, authors of junifer.
This program is free software: you can redistribute it and/or modify This program is free software: you can redistribute it and/or modify

View file

@ -15,12 +15,11 @@ Thus, only a few methods are required:
1. ``get_valid_inputs``: a method to obtain the list of valid inputs for the marker. This is used to check that the 1. ``get_valid_inputs``: a method to obtain the list of valid inputs for the marker. This is used to check that the
inputs provided by the user are valid. This method should return a list of strings, representing inputs provided by the user are valid. This method should return a list of strings, representing
:ref:`data types <data_types>` :ref:`data types <data_types>`
2. ``get_output_kind``: a method to obtain the kind of output of the marker. This is used to check that the output 2. ``get_output_type``: a method to obtain the kind of output of the marker. This is used to check that the output
of the marker is compatible with the storage. This method should return a string, representing of the marker is compatible with the storage. This method should return a string, representing
:ref:`storage types <storage_types>` :ref:`storage types <storage_types>`
3. ``compute``: the method that given the data, computes the marker. 3. ``compute``: the method that given the data, computes the marker.
4. ``store``: the method that stores the computed marker. 4. ``__init__``: the initialization method, where the marker is configured.
5. ``__init__``: the initialization method, where the marker is configured.
As an example, we will develop a Parcel Mean marker, that is, a marker that first applies a parcellation and As an example, we will develop a Parcel Mean marker, that is, a marker that first applies a parcellation and
then computes the mean of the data in each parcel. This is a very simple example, but it will show you how to create then computes the mean of the data in each parcel. This is a very simple example, but it will show you how to create
@ -44,7 +43,7 @@ it will be ``table``. Thus, we can define the output as:
.. code-block:: python .. code-block:: python
def get_output_kind(self, input_kind): def get_output_type(self, input_kind):
if input_kind == 'BOLD': if input_kind == 'BOLD':
return 'timeseries' return 'timeseries'
else: else:
@ -143,38 +142,24 @@ This dictionary will later be passed onto the ``store`` method.
return out return out
.. _extending_markers_store: .. _extending_markers_finalize:
Step 4: Store the marker Step 4: Finalize the marker
------------------------ ---------------------------
In this step, we will define the method that stores the marker. This method will be called by junifer when needed, Once all of the above steps are done, we just need to give our marker a name, state its *dependencies* and register it
using the data provided by the ``compute`` method. The method ``store`` has three arguments: using the ``@register_marker`` decorator.
* ``kind``: A string indicating the :ref:`data type <data_types>` that was used to compute the marker. The *dependencies* are the core packages that are required to compute the marker. This will be later used to keep track
* ``out``: The output of the ``compute`` method. of the versions of the packages used to compute the marker. To inform junifer about the dependencies of a marker,
* ``storage``: The storage object, that will be used to store the marker. we need to define a ``_DEPENDENCIES`` attribute in the class. This attribute must be a set, with the names of the
packages as strings. For example, the ``ParcelMean`` marker has the following dependencies:
.. code-block:: python .. code-block:: python
def store(self, kind, out, storage): _DEPENDENCIES = {"nilearn"}
if kind in ["VBM_GM", "VBM_WM"]:
storage.store(kind="table", **out)
elif kind in ["BOLD"]:
storage.store(kind="timeseries", **out)
.. hint:: Check the hint on :ref:`extending_markers_compute`. If the output of the ``compute`` method is a dictionary Finally, we need to register the marker using the ``@register_marker`` decorator. This decorator takes the name of the
with keys based on the :ref:`storage types <storage_types>`, the ``store`` method can simply call the right
storage function, based on the ``kind`` parameter, with ``**out``.
.. _extending_markers_finalize:
Step 5: Finalize the marker
---------------------------
Once all of the above steps are done, we just need to give our marker a name an register it using the
``@register_marker`` decorator:
.. code-block:: python .. code-block:: python
@ -186,6 +171,8 @@ Once all of the above steps are done, we just need to give our marker a name an
@register_marker @register_marker
class ParcelMean(BaseMarker): class ParcelMean(BaseMarker):
_DEPENDENCIES = {"nilearn", "numpy"}
def __init__(self, parcellation_name, on=None, name=None): def __init__(self, parcellation_name, on=None, name=None):
self.parcellation_name = parcellation_name self.parcellation_name = parcellation_name
super().__init__(on=on, name=name) super().__init__(on=on, name=name)
@ -193,7 +180,7 @@ Once all of the above steps are done, we just need to give our marker a name an
def get_valid_inputs(self): def get_valid_inputs(self):
return ['BOLD', 'VBM_WM', 'VBM_GM'] return ['BOLD', 'VBM_WM', 'VBM_GM']
def get_output_kind(self, input_kind): def get_output_type(self, input_kind):
if input_kind == 'BOLD': if input_kind == 'BOLD':
return 'timeseries' return 'timeseries'
else: else:
@ -231,12 +218,6 @@ Once all of the above steps are done, we just need to give our marker a name an
out["row_names"] = "scan" out["row_names"] = "scan"
return out return out
def store(self, kind, out, storage):
if kind in ["VBM_GM", "VBM_WM"]:
storage.store(kind="table", **out)
elif kind in ["BOLD"]:
storage.store(kind="timeseries", **out)
.. _extending_markers_template: .. _extending_markers_template:
@ -260,7 +241,7 @@ Template for a custom Marker
valid = [] valid = []
return valid return valid
def get_output_kind(self, input_kind): def get_output_type(self, input_kind):
# TODO: Return the valid output kind for each input kind # TODO: Return the valid output kind for each input kind
pass pass
@ -270,7 +251,3 @@ Template for a custom Marker
# Create the output dictionary # Create the output dictionary
out = {"data": None, "columns": None} out = {"data": None, "columns": None}
return out return out
def store(self, kind, out, storage):
# TODO: store out using the storage object, based on the kind of data
pass

View file

@ -20,4 +20,4 @@ as the actual data is in the memory and the Python runtime has not garbage-colle
If you are interested in using already provided markers, please go to :doc:`../builtin`. And, if you want to implement If you are interested in using already provided markers, please go to :doc:`../builtin`. And, if you want to implement
your own marker, you need to provide concrete implementation of :class:`junifer.markers.BaseMarker`. Specifically, you your own marker, you need to provide concrete implementation of :class:`junifer.markers.BaseMarker`. Specifically, you
need to override ``get_output_kind``, ``store`` and ``compute`` methods. need to override ``get_output_type``, ``store`` and ``compute`` methods.

View file

@ -42,7 +42,10 @@ marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
############################################################################### ###############################################################################
# Prepare the input # Prepare the input
input = {"BOLD": {"data": fmri_img}, "VBM_GM": {"data": vbm_img}} input = {
"BOLD": {"data": fmri_img, "meta": {"element": "subject1"}},
"VBM_GM": {"data": vbm_img, "meta": {"element": "subject1"}},
}
############################################################################### ###############################################################################
# Fit transform the data # Fit transform the data

View file

@ -66,6 +66,9 @@ with tempfile.TemporaryDirectory() as tmpdir:
collect(storage=storage) collect(storage=storage)
# Create storage object to read in extracted features # Create storage object to read in extracted features
db = SQLiteFeatureStorage(uri=storage["uri"]) db = SQLiteFeatureStorage(uri=storage["uri"])
# List all the features
print(db.list_features())
# Read extracted features # Read extracted features
df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS") df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS")

View file

@ -16,8 +16,8 @@ import yaml
from ..utils.logging import ( from ..utils.logging import (
configure_logging, configure_logging,
logger, logger,
warn_with_log,
raise_error, raise_error,
warn_with_log,
) )
from .functions import collect as api_collect from .functions import collect as api_collect
from .functions import queue as api_queue from .functions import queue as api_queue
@ -69,7 +69,8 @@ def _parse_elements(element: str, config: Dict) -> Union[List, None]:
raise_error( raise_error(
"The 'elements' key is set in the configuration, but its value" "The 'elements' key is set in the configuration, but its value"
" is 'None'. It is likely that there is an empty 'elements' " " is 'None'. It is likely that there is an empty 'elements' "
"section in the yaml configuration file.") "section in the yaml configuration file."
)
return elements return elements

View file

@ -7,11 +7,11 @@
import shutil import shutil
import subprocess import subprocess
import textwrap
import typing import typing
from pathlib import Path from pathlib import Path
from typing import Dict, List, Optional, Tuple, Union from typing import Dict, List, Optional, Tuple, Union
import textwrap
import yaml import yaml
from ..datagrabber.base import BaseDataGrabber from ..datagrabber.base import BaseDataGrabber

View file

@ -5,11 +5,11 @@
# License: AGPL # License: AGPL
import importlib import importlib
from pathlib import Path
from typing import Dict, Union
import importlib.util import importlib.util
import os import os
import sys import sys
from pathlib import Path
from typing import Dict, Union
import yaml import yaml
@ -45,7 +45,8 @@ def parse_yaml(filepath: Union[str, Path]) -> Dict:
if contents["elements"] is None: if contents["elements"] is None:
raise_error( raise_error(
"The elements key was defined but its content is empty. " "The elements key was defined but its content is empty. "
"Please define the elements to operate on or remove the key.") "Please define the elements to operate on or remove the key."
)
# load modules # load modules
if "with" in contents: if "with" in contents:
to_load = contents["with"] to_load = contents["with"]
@ -58,9 +59,11 @@ def parse_yaml(filepath: Union[str, Path]) -> Dict:
file_path = Path(os.getcwd()) / t_module file_path = Path(os.getcwd()) / t_module
if not file_path.exists(): if not file_path.exists():
raise_error( raise_error(
f"File in 'with' section does not exist: {file_path}") f"File in 'with' section does not exist: {file_path}"
)
spec = importlib.util.spec_from_file_location( spec = importlib.util.spec_from_file_location(
t_module, file_path) t_module, file_path
)
module = importlib.util.module_from_spec(spec) # type: ignore module = importlib.util.module_from_spec(spec) # type: ignore
sys.modules[t_module] = module sys.modules[t_module] = module
spec.loader.exec_module(module) # type: ignore spec.loader.exec_module(module) # type: ignore

View file

@ -13,6 +13,7 @@ from click.testing import CliRunner
from junifer.api.cli import collect, run, selftest, wtf from junifer.api.cli import collect, run, selftest, wtf
# Create click test runner # Create click test runner
runner = CliRunner() runner = CliRunner()

View file

@ -11,6 +11,7 @@ import pytest
from junifer.configs.juseless.datagrabbers import JuselessDataladAOMICID1000VBM from junifer.configs.juseless.datagrabbers import JuselessDataladAOMICID1000VBM
from junifer.utils.logging import configure_logging from junifer.utils.logging import configure_logging
# Check if the test is running on juseless # Check if the test is running on juseless
if socket.gethostname() != "juseless": if socket.gethostname() != "juseless":
pytest.skip("These tests are only for juseless", allow_module_level=True) pytest.skip("These tests are only for juseless", allow_module_level=True)

View file

@ -12,6 +12,7 @@ import pytest
from junifer.configs.juseless.datagrabbers import JuselessDataladCamCANVBM from junifer.configs.juseless.datagrabbers import JuselessDataladCamCANVBM
from junifer.utils.logging import configure_logging from junifer.utils.logging import configure_logging
# Check if the test is running on juseless # Check if the test is running on juseless
if socket.gethostname() != "juseless": if socket.gethostname() != "juseless":
pytest.skip("These tests are only for juseless", allow_module_level=True) pytest.skip("These tests are only for juseless", allow_module_level=True)

View file

@ -12,6 +12,7 @@ import pytest
from junifer.configs.juseless.datagrabbers import JuselessDataladIXIVBM from junifer.configs.juseless.datagrabbers import JuselessDataladIXIVBM
from junifer.utils.logging import configure_logging from junifer.utils.logging import configure_logging
# Check if the test is running on juseless # Check if the test is running on juseless
if socket.gethostname() != "juseless": if socket.gethostname() != "juseless":
pytest.skip("These tests are only for juseless", allow_module_level=True) pytest.skip("These tests are only for juseless", allow_module_level=True)

View file

@ -13,6 +13,7 @@ import pytest
from junifer.configs.juseless.datagrabbers import JuselessUCLA from junifer.configs.juseless.datagrabbers import JuselessUCLA
from junifer.utils.logging import configure_logging from junifer.utils.logging import configure_logging
# Check if the test is running on juseless # Check if the test is running on juseless
if socket.gethostname() != "juseless": if socket.gethostname() != "juseless":
pytest.skip("These tests are only for juseless", allow_module_level=True) pytest.skip("These tests are only for juseless", allow_module_level=True)

View file

@ -12,6 +12,7 @@ import pytest
from junifer.configs.juseless.datagrabbers import JuselessDataladUKBVBM from junifer.configs.juseless.datagrabbers import JuselessDataladUKBVBM
from junifer.utils.logging import configure_logging from junifer.utils.logging import configure_logging
# Check if the test is running on juseless # Check if the test is running on juseless
if socket.gethostname() != "juseless": if socket.gethostname() != "juseless":
pytest.skip("These tests are only for juseless", allow_module_level=True) pytest.skip("These tests are only for juseless", allow_module_level=True)

View file

@ -12,6 +12,7 @@ import pytest
from junifer.datagrabber.hcp import DataladHCP1200 from junifer.datagrabber.hcp import DataladHCP1200
from junifer.utils.logging import configure_logging from junifer.utils.logging import configure_logging
# Check if the test is running on juseless # Check if the test is running on juseless
if socket.gethostname() != "juseless": if socket.gethostname() != "juseless":
pytest.skip("These tests are only for juseless", allow_module_level=True) pytest.skip("These tests are only for juseless", allow_module_level=True)

View file

@ -21,4 +21,4 @@ from .masks import (
register_mask, register_mask,
) )
from . import utils from . import utils

View file

@ -13,6 +13,7 @@ from numpy.typing import ArrayLike
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
# Path to the VOIs # Path to the VOIs
_vois_path = Path(__file__).parent / "VOIs" _vois_path = Path(__file__).parent / "VOIs"

View file

@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import nibabel as nib import nibabel as nib
from .utils import closest_resolution
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
from .utils import closest_resolution
if TYPE_CHECKING: if TYPE_CHECKING:
@ -58,10 +58,9 @@ def register_mask(
if name in _available_masks: if name in _available_masks:
if overwrite is True: if overwrite is True:
logger.info(f"Overwriting {name} mask") logger.info(f"Overwriting {name} mask")
if (_available_masks[name]["family"] != "CustomUserMask"): if _available_masks[name]["family"] != "CustomUserMask":
raise_error( raise_error(
f"Cannot overwrite {name} mask. " f"Cannot overwrite {name} mask. " "It is a built-in mask."
"It is a built-in mask."
) )
else: else:
raise_error( raise_error(
@ -117,8 +116,7 @@ def load_mask(
""" """
if name not in _available_masks: if name not in _available_masks:
raise_error( raise_error(
f"Mask {name} not found. " f"Mask {name} not found. " f"Valid options are: {list_masks()}"
f"Valid options are: {list_masks()}"
) )
mask_definition = _available_masks[name].copy() mask_definition = _available_masks[name].copy()
@ -126,12 +124,10 @@ def load_mask(
if t_family == "CustomUserMask": if t_family == "CustomUserMask":
mask_fname = Path(mask_definition["path"]) mask_fname = Path(mask_definition["path"])
elif t_family == 'Vickery-Patil': elif t_family == "Vickery-Patil":
mask_fname = _load_vickery_patil_mask(name, resolution) mask_fname = _load_vickery_patil_mask(name, resolution)
else: else:
raise_error( raise_error(f"I don't know about the {t_family} mask family.")
f"I don't know about the {t_family} mask family."
)
logger.info(f"Loading mask {mask_fname.absolute()}") logger.info(f"Loading mask {mask_fname.absolute()}")
@ -167,8 +163,9 @@ def _load_vickery_patil_mask(
available_resolutions = [1.5, 3.0] available_resolutions = [1.5, 3.0]
to_load = closest_resolution(resolution, available_resolutions) to_load = closest_resolution(resolution, available_resolutions)
if to_load == 3.0: if to_load == 3.0:
mask_fname = \ mask_fname = (
"CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz" "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz"
)
elif to_load == 1.5: elif to_load == 1.5:
mask_fname = "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz" mask_fname = "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz"
else: else:
@ -178,9 +175,7 @@ def _load_vickery_patil_mask(
elif name == "GM_prob0.2_cortex": elif name == "GM_prob0.2_cortex":
mask_fname = "GMprob0.2_cortex_3mm_NA_rm.nii.gz" mask_fname = "GMprob0.2_cortex_3mm_NA_rm.nii.gz"
else: else:
raise_error( raise_error(f"Cannot find a Vickery-Patil mask called {name}")
f"Cannot find a Vickery-Patil mask called {name}"
)
mask_fname = _masks_path / "vickery-patil" / mask_fname mask_fname = _masks_path / "vickery-patil" / mask_fname
return mask_fname return mask_fname

View file

@ -18,8 +18,9 @@ import pandas as pd
import requests import requests
from nilearn import datasets from nilearn import datasets
from .utils import closest_resolution
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
from .utils import closest_resolution
if TYPE_CHECKING: if TYPE_CHECKING:
from nibabel import Nifti1Image from nibabel import Nifti1Image

View file

@ -5,8 +5,8 @@
from typing import List from typing import List
import pytest
import numpy as np import numpy as np
import pytest
from junifer.data.utils import closest_resolution from junifer.data.utils import closest_resolution

View file

@ -11,10 +11,10 @@ import pytest
from numpy.testing import assert_array_almost_equal from numpy.testing import assert_array_almost_equal
from junifer.data.masks import ( from junifer.data.masks import (
_load_vickery_patil_mask,
list_masks,
load_mask, load_mask,
register_mask, register_mask,
list_masks,
_load_vickery_patil_mask,
) )

View file

@ -8,11 +8,10 @@
from pathlib import Path from pathlib import Path
from typing import List from typing import List
import pytest
from numpy.testing import assert_array_almost_equal, assert_array_equal
import nibabel as nib import nibabel as nib
import pytest
from nilearn.image import new_img_like from nilearn.image import new_img_like
from numpy.testing import assert_array_almost_equal, assert_array_equal
from junifer.data.parcellations import ( from junifer.data.parcellations import (
_retrieve_parcellation, _retrieve_parcellation,

View file

@ -1,5 +1,5 @@
"""Provide utilities for data module.""" """Provide utilities for data module."""
from typing import Optional, Union, List from typing import List, Optional, Union
import numpy as np import numpy as np

View file

@ -111,8 +111,8 @@ def test_aomic1000_datagrabber() -> None:
assert out["DWI"]["path"].is_file() assert out["DWI"]["path"].is_file()
# asserts meta # asserts meta
assert "meta" in out assert "meta" in out["BOLD"]
meta = out["meta"] meta = out["BOLD"]["meta"]
assert "element" in meta assert "element" in meta
assert "subject" in meta["element"] assert "subject" in meta["element"]
assert test_element == meta["element"]["subject"] assert test_element == meta["element"]["subject"]

View file

@ -125,8 +125,8 @@ def test_aomic_piop1_datagrabber() -> None:
assert out["DWI"]["path"].is_file() assert out["DWI"]["path"].is_file()
# asserts meta # asserts meta
assert "meta" in out assert "meta" in out["BOLD"]
meta = out["meta"] meta = out["BOLD"]["meta"]
assert "element" in meta assert "element" in meta
assert "subject" in meta["element"] assert "subject" in meta["element"]
assert sub == meta["element"]["subject"] assert sub == meta["element"]["subject"]

View file

@ -121,8 +121,8 @@ def test_aomic_piop2_datagrabber() -> None:
assert out["DWI"]["path"].is_file() assert out["DWI"]["path"].is_file()
# asserts meta # asserts meta
assert "meta" in out assert "meta" in out["BOLD"]
meta = out["meta"] meta = out["BOLD"]["meta"]
assert "element" in meta assert "element" in meta
assert "subject" in meta["element"] assert "subject" in meta["element"]
assert sub == meta["element"]["subject"] assert sub == meta["element"]["subject"]

View file

@ -9,11 +9,12 @@ from abc import ABC, abstractmethod
from pathlib import Path from pathlib import Path
from typing import Dict, Iterator, List, Tuple, Union from typing import Dict, Iterator, List, Tuple, Union
from ..pipeline import UpdateMetaMixin
from ..utils import logger, raise_error from ..utils import logger, raise_error
from .utils import validate_types from .utils import validate_types
class BaseDataGrabber(ABC): class BaseDataGrabber(ABC, UpdateMetaMixin):
"""Abstract base class for datagrabber. """Abstract base class for datagrabber.
For every interface that is required, one needs to provide a concrete For every interface that is required, one needs to provide a concrete
@ -78,10 +79,11 @@ class BaseDataGrabber(ABC):
named_element = dict(zip(self.get_element_keys(), element)) named_element = dict(zip(self.get_element_keys(), element))
logger.debug(f"Named element: {named_element}") logger.debug(f"Named element: {named_element}")
out = self.get_item(**named_element) out = self.get_item(**named_element)
out["meta"] = {
"datagrabber": self.get_meta(), for _, t_val in out.items():
"element": named_element, self.update_meta(t_val, "datagrabber")
} t_val["meta"]["element"] = named_element
return out return out
def __enter__(self) -> "BaseDataGrabber": def __enter__(self) -> "BaseDataGrabber":
@ -103,22 +105,6 @@ class BaseDataGrabber(ABC):
""" """
return self.types.copy() return self.types.copy()
def get_meta(self) -> Dict:
"""Get metadata.
Returns
-------
dict
The metadata as dictionary.
"""
t_meta = {}
t_meta["class"] = self.__class__.__name__
for k, v in vars(self).items():
if not k.startswith("_"):
t_meta[k] = v
return t_meta
@property @property
def datadir(self) -> Path: def datadir(self) -> Path:
"""Get data directory path. """Get data directory path.

View file

@ -83,7 +83,7 @@ class DataladDataGrabber(BaseDataGrabber):
self._rootdir = rootdir self._rootdir = rootdir
# Flag to indicate if the dataset was cloned before and it might be # Flag to indicate if the dataset was cloned before and it might be
# dirty # dirty
self._dataset_dirty = False self.datalad_dirty = False
@property @property
def datadir(self) -> Path: def datadir(self) -> Path:
@ -175,7 +175,7 @@ class DataladDataGrabber(BaseDataGrabber):
# Check for dirty datasets: # Check for dirty datasets:
status = self._dataset.status() status = self._dataset.status()
if any([x["state"] != "clean" for x in status]): if any([x["state"] != "clean" for x in status]):
self._dataset_dirty = True self.datalad_dirty = True
warn_with_log( warn_with_log(
"At least one file is not clean, Junifer will " "At least one file is not clean, Junifer will "
"consider this dataset as dirty." "consider this dataset as dirty."
@ -191,11 +191,10 @@ class DataladDataGrabber(BaseDataGrabber):
logger.debug("Dataset installed") logger.debug("Dataset installed")
self._was_cloned = not isinstalled self._was_cloned = not isinstalled
self._datalad_commit_id = ( self.datalad_commit_id = self._dataset.repo.get_hexsha( # type: ignore
self._dataset.repo.get_hexsha( # type: ignore self._dataset.repo.get_corresponding_branch() # type: ignore
self._dataset.repo.get_corresponding_branch() # type: ignore
)
) )
self.datalad_id = self._dataset.id
def cleanup(self) -> None: def cleanup(self) -> None:
"""Cleanup the datalad dataset.""" """Cleanup the datalad dataset."""
@ -245,21 +244,3 @@ class DataladDataGrabber(BaseDataGrabber):
logger.debug("Cleaning up dataset") logger.debug("Cleaning up dataset")
self.cleanup() self.cleanup()
logger.debug("Dataset state restored") logger.debug("Dataset state restored")
def get_meta(self) -> Dict:
"""Get metadata.
Returns
-------
dict
The metadata as dictionary.
"""
t_meta = super().get_meta()
t_meta["datalad_commit_id"] = self._datalad_commit_id
t_meta["datalad_id"] = self._dataset.id
# Set a flag to indicate that the dataset was dirty
t_meta["datalad_dirty"] = self._dataset_dirty
return t_meta

View file

@ -5,7 +5,6 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path
from typing import Dict, List, Tuple, Union from typing import Dict, List, Tuple, Union
from .base import BaseDataGrabber from .base import BaseDataGrabber
@ -38,7 +37,7 @@ class MultipleDataGrabber(BaseDataGrabber):
raise ValueError("Datagrabbers have overlapping types.") raise ValueError("Datagrabbers have overlapping types.")
self._datagrabbers = datagrabbers self._datagrabbers = datagrabbers
def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: def __getitem__(self, element: Union[str, Tuple]) -> Dict:
"""Implement indexing. """Implement indexing.
Parameters Parameters
@ -58,9 +57,20 @@ class MultipleDataGrabber(BaseDataGrabber):
""" """
out = {} out = {}
metas = []
for dg in self._datagrabbers: for dg in self._datagrabbers:
t_out = dg[element] t_out = dg[element]
out.update(t_out) out.update(t_out)
# Now get the meta for this datagrabber
t_meta = {}
dg.update_meta(t_meta, "datagrabber")
# Store all the sub-datagrabbers meta
metas.append(t_meta["meta"]["datagrabber"])
# Update all the metas again
for kind in out:
self.update_meta(out[kind], "datagrabber")
out[kind]["meta"]["datagrabber"]["datagrabbers"] = metas
return out return out
def get_item(self, **element: Dict) -> Dict[str, Dict]: def get_item(self, **element: Dict) -> Dict[str, Dict]:
@ -136,17 +146,3 @@ class MultipleDataGrabber(BaseDataGrabber):
""" """
types = [x for dg in self._datagrabbers for x in dg.get_types()] types = [x for dg in self._datagrabbers for x in dg.get_types()]
return types return types
def get_meta(self) -> Dict:
"""Get metadata.
Returns
-------
dict
The metadata as dictionary.
"""
t_meta = {}
t_meta["class"] = self.__class__.__name__
t_meta["datagrabbers"] = [dg.get_meta() for dg in self._datagrabbers]
return t_meta

View file

@ -23,7 +23,7 @@ def test_BaseDataGrabber() -> None:
# Create concrete class. # Create concrete class.
class MyDataGrabber(BaseDataGrabber): class MyDataGrabber(BaseDataGrabber):
def get_item(self, subject): def get_item(self, subject):
return {} return {"BOLD": {}}
def get_elements(self): def get_elements(self):
return super().get_elements() return super().get_elements()
@ -31,22 +31,24 @@ def test_BaseDataGrabber() -> None:
def get_element_keys(self): def get_element_keys(self):
return ["subject"] return ["subject"]
dg = MyDataGrabber(datadir="/tmp", types=["func"]) dg = MyDataGrabber(datadir="/tmp", types=["BOLD"])
elem = dg["elem"] elem = dg["sub01"]
assert "meta" in elem assert "BOLD" in elem
assert "datagrabber" in elem["meta"] assert "meta" in elem["BOLD"]
assert "class" in elem["meta"]["datagrabber"] meta = elem["BOLD"]["meta"]
assert MyDataGrabber.__name__ in elem["meta"]["datagrabber"]["class"] assert "datagrabber" in meta
assert "element" in elem["meta"] assert "class" in meta["datagrabber"]
assert "subject" in elem["meta"]["element"] assert MyDataGrabber.__name__ in meta["datagrabber"]["class"]
assert "elem" in elem["meta"]["element"]["subject"] assert "element" in meta
assert "subject" in meta["element"]
assert "sub01" in meta["element"]["subject"]
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
dg.get_elements() dg.get_elements()
with dg: with dg:
assert dg.datadir == Path("/tmp") assert dg.datadir == Path("/tmp")
assert dg.types == ["func"] assert dg.types == ["BOLD"]
class MyDataGrabber2(BaseDataGrabber): class MyDataGrabber2(BaseDataGrabber):
def get_item(self, subject): def get_item(self, subject):
@ -58,7 +60,7 @@ def test_BaseDataGrabber() -> None:
def get_element_keys(self): def get_element_keys(self):
return super().get_element_keys() return super().get_element_keys()
dg = MyDataGrabber2(datadir="/tmp", types=["func"]) dg = MyDataGrabber2(datadir="/tmp", types=["BOLD"])
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
dg.get_element_keys() dg.get_element_keys()

View file

@ -11,6 +11,7 @@ import pytest
from junifer.datagrabber.datalad_base import DataladDataGrabber from junifer.datagrabber.datalad_base import DataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
@ -143,10 +144,11 @@ def test_datalad_clone_cleanup(
assert elem1_t1w.is_file() is False assert elem1_t1w.is_file() is False
assert elem1_t1w.is_symlink() is True assert elem1_t1w.is_symlink() is True
elem1 = dg["sub-01"] elem1 = dg["sub-01"]
assert "meta" in elem1 assert "meta" in elem1["BOLD"]
assert "datagrabber" in elem1["meta"] meta = elem1["BOLD"]["meta"]
assert "datalad_dirty" in elem1["meta"]["datagrabber"] assert "datagrabber" in meta
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False assert "datalad_dirty" in meta["datagrabber"]
assert meta["datagrabber"]["datalad_dirty"] is False
assert hasattr(dg, "_got_files") is False assert hasattr(dg, "_got_files") is False
assert datadir.exists() is True assert datadir.exists() is True
assert elem1_bold.is_file() is True assert elem1_bold.is_file() is True
@ -198,14 +200,15 @@ def test_datalad_previously_cloned(
assert datadir.exists() is True assert datadir.exists() is True
assert dg._was_cloned is False assert dg._was_cloned is False
elem1 = dg["sub-01"] elem1 = dg["sub-01"]
assert "meta" in elem1 assert "meta" in elem1["BOLD"]
assert "datagrabber" in elem1["meta"] meta = elem1["BOLD"]["meta"]
assert "datalad_dirty" in elem1["meta"]["datagrabber"] assert "datagrabber" in meta
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False assert "datalad_dirty" in meta["datagrabber"]
assert "datalad_commit_id" in elem1["meta"]["datagrabber"] assert meta["datagrabber"]["datalad_dirty"] is False
assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit assert "datalad_commit_id" in meta["datagrabber"]
assert "datalad_id" in elem1["meta"]["datagrabber"] assert meta["datagrabber"]["datalad_commit_id"] == commit
assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id assert "datalad_id" in meta["datagrabber"]
assert meta["datagrabber"]["datalad_id"] == remote_id
assert hasattr(dg, "_got_files") is True assert hasattr(dg, "_got_files") is True
# Files are there and symlinks are fixed # Files are there and symlinks are fixed
@ -264,7 +267,8 @@ def test_datalad_previously_cloned_and_get(
assert elem1_t1w.is_file() is False assert elem1_t1w.is_file() is False
dl.get( # type: ignore dl.get( # type: ignore
elem1_t1w, dataset=datadir, result_renderer="disabled") elem1_t1w, dataset=datadir, result_renderer="disabled"
)
assert elem1_bold.is_symlink() is True assert elem1_bold.is_symlink() is True
assert elem1_bold.is_file() is False assert elem1_bold.is_file() is False
@ -275,14 +279,15 @@ def test_datalad_previously_cloned_and_get(
assert datadir.exists() is True assert datadir.exists() is True
assert dg._was_cloned is False assert dg._was_cloned is False
elem1 = dg["sub-01"] elem1 = dg["sub-01"]
assert "meta" in elem1 assert "meta" in elem1["BOLD"]
assert "datagrabber" in elem1["meta"] meta = elem1["BOLD"]["meta"]
assert "datalad_dirty" in elem1["meta"]["datagrabber"] assert "datagrabber" in meta
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False assert "datalad_dirty" in meta["datagrabber"]
assert "datalad_commit_id" in elem1["meta"]["datagrabber"] assert meta["datagrabber"]["datalad_dirty"] is False
assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit assert "datalad_commit_id" in meta["datagrabber"]
assert "datalad_id" in elem1["meta"]["datagrabber"] assert meta["datagrabber"]["datalad_commit_id"] == commit
assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id assert "datalad_id" in meta["datagrabber"]
assert meta["datagrabber"]["datalad_id"] == remote_id
assert hasattr(dg, "_got_files") is True assert hasattr(dg, "_got_files") is True
# Files are there and symlinks are fixed # Files are there and symlinks are fixed
@ -344,7 +349,8 @@ def test_datalad_previously_cloned_and_get_dirty(
assert elem1_t1w.is_file() is False assert elem1_t1w.is_file() is False
dl.get( # type: ignore dl.get( # type: ignore
elem1_t1w, dataset=datadir, result_renderer="disabled") elem1_t1w, dataset=datadir, result_renderer="disabled"
)
assert elem1_bold.is_symlink() is True assert elem1_bold.is_symlink() is True
assert elem1_bold.is_file() is False assert elem1_bold.is_file() is False
@ -359,14 +365,15 @@ def test_datalad_previously_cloned_and_get_dirty(
assert datadir.exists() is True assert datadir.exists() is True
assert dg._was_cloned is False assert dg._was_cloned is False
elem1 = dg["sub-01"] elem1 = dg["sub-01"]
assert "meta" in elem1 assert "meta" in elem1["BOLD"]
assert "datagrabber" in elem1["meta"] meta = elem1["BOLD"]["meta"]
assert "datalad_dirty" in elem1["meta"]["datagrabber"] assert "datagrabber" in meta
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is True assert "datalad_dirty" in meta["datagrabber"]
assert "datalad_commit_id" in elem1["meta"]["datagrabber"] assert meta["datagrabber"]["datalad_dirty"] is True
assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit assert "datalad_commit_id" in meta["datagrabber"]
assert "datalad_id" in elem1["meta"]["datagrabber"] assert meta["datagrabber"]["datalad_commit_id"] == commit
assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id assert "datalad_id" in meta["datagrabber"]
assert meta["datagrabber"]["datalad_id"] == remote_id
assert hasattr(dg, "_got_files") is True assert hasattr(dg, "_got_files") is True
# Files are there and symlinks are fixed # Files are there and symlinks are fixed
@ -384,17 +391,18 @@ def test_datalad_previously_cloned_and_get_dirty(
assert datadir.exists() is True assert datadir.exists() is True
assert dg._was_cloned is False assert dg._was_cloned is False
elem2 = dg["sub-02"] elem2 = dg["sub-02"]
assert "meta" in elem2 assert "meta" in elem1["BOLD"]
assert "datagrabber" in elem2["meta"] meta = elem2["BOLD"]["meta"]
assert "datalad_dirty" in elem2["meta"]["datagrabber"] assert "datagrabber" in meta
assert "datalad_dirty" in meta["datagrabber"]
# Dataset is still dirty due to subject sub-01 # Dataset is still dirty due to subject sub-01
assert elem2["meta"]["datagrabber"]["datalad_dirty"] is True assert meta["datagrabber"]["datalad_dirty"] is True
assert "datalad_commit_id" in elem2["meta"]["datagrabber"] assert "datalad_commit_id" in meta["datagrabber"]
assert elem2["meta"]["datagrabber"]["datalad_commit_id"] == commit assert meta["datagrabber"]["datalad_commit_id"] == commit
assert "datalad_id" in elem2["meta"]["datagrabber"] assert "datalad_id" in meta["datagrabber"]
assert elem2["meta"]["datagrabber"]["datalad_id"] == remote_id assert meta["datagrabber"]["datalad_id"] == remote_id
assert hasattr(dg, "_got_files") is True assert hasattr(dg, "_got_files") is True
# Files are there and symlinks are fixed # Files are there and symlinks are fixed

View file

@ -10,6 +10,7 @@ import pytest
from junifer.datagrabber.hcp import DataladHCP1200 from junifer.datagrabber.hcp import DataladHCP1200
from junifer.utils import configure_logging from junifer.utils import configure_logging
URI = "https://gin.g-node.org/juaml/datalad-example-hcp1200" URI = "https://gin.g-node.org/juaml/datalad-example-hcp1200"
@ -79,8 +80,8 @@ def test_dataladhcp1200_datagrabber(
# Assert data file path is a file # Assert data file path is a file
assert out["BOLD"]["path"].is_file() assert out["BOLD"]["path"].is_file()
# Assert metadata # Assert metadata
assert "meta" in out assert "meta" in out["BOLD"]
meta = out["meta"] meta = out["BOLD"]["meta"]
assert "element" in meta assert "element" in meta
assert "subject" in meta["element"] assert "subject" in meta["element"]
assert test_element[0] == meta["element"]["subject"] assert test_element[0] == meta["element"]["subject"]

View file

@ -7,6 +7,7 @@ import pytest
from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
@ -63,17 +64,17 @@ def test_multiple() -> None:
subs = [x for x in dg] subs = [x for x in dg]
assert set(subs) == set(expected_subs) assert set(subs) == set(expected_subs)
data = dg[("sub-01", "ses-01")] elem = dg[("sub-01", "ses-01")]
assert "T1w" in data assert "T1w" in elem
assert "BOLD" in data assert "BOLD" in elem
assert "meta" in elem["BOLD"]
meta = dg.get_meta() meta = elem["BOLD"]["meta"]["datagrabber"]
assert "class" in meta assert "class" in meta
assert meta["class"] == "MultipleDataGrabber" assert meta["class"] == "MultipleDataGrabber"
assert "datagrabbers" in meta assert "datagrabbers" in meta
assert len(meta["datagrabbers"]) == 2 assert len(meta["datagrabbers"]) == 2
assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber" assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber"
assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber" assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber"
def test_multiple_no_intersection() -> None: def test_multiple_no_intersection() -> None:

View file

@ -256,6 +256,6 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
out1 = datagrabber[("sub000", "ses000", "task002")] out1 = datagrabber[("sub000", "ses000", "task002")]
out2 = datagrabber[("sub000", "ses000", "task003")] out2 = datagrabber[("sub000", "ses000", "task003")]
assert out1["func"] == out2["func"] assert out1["func"]["path"] == out2["func"]["path"]
assert out1["anat"] == out2["anat"] assert out1["anat"]["path"] == out2["anat"]["path"]
assert out1["vbm"] != out2["vbm"] assert out1["vbm"]["path"] != out2["vbm"]["path"]

View file

@ -11,6 +11,7 @@ import pytest
from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
@ -81,9 +82,10 @@ def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None:
dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz"
) )
assert "meta" in t_sub assert "meta" in t_sub["BOLD"]
assert "datagrabber" in t_sub["meta"] meta = t_sub["BOLD"]["meta"]
dg_meta = t_sub["meta"]["datagrabber"] assert "datagrabber" in meta
dg_meta = meta["datagrabber"]
assert "class" in dg_meta assert "class" in dg_meta
assert dg_meta["class"] == "PatternDataladDataGrabber" assert dg_meta["class"] == "PatternDataladDataGrabber"
assert "uri" in dg_meta assert "uri" in dg_meta

View file

@ -11,10 +11,11 @@ import nibabel as nib
import pandas as pd import pandas as pd
from ..api.decorators import register_datareader from ..api.decorators import register_datareader
from ..pipeline import PipelineStepMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
from ..utils.logging import logger, warn_with_log from ..utils.logging import logger, warn_with_log
# Map each file extension to a kind
# Map each file extension to a type
_extensions = { _extensions = {
".nii": "NIFTI", ".nii": "NIFTI",
".nii.gz": "NIFTI", ".nii.gz": "NIFTI",
@ -22,7 +23,7 @@ _extensions = {
".tsv": "TSV", ".tsv": "TSV",
} }
# Map each kind to a function and arguments # Map each type to a function and arguments
_readers = {} _readers = {}
_readers["NIFTI"] = {"func": nib.load, "params": None} _readers["NIFTI"] = {"func": nib.load, "params": None}
_readers["CSV"] = {"func": pd.read_csv, "params": None} _readers["CSV"] = {"func": pd.read_csv, "params": None}
@ -30,7 +31,7 @@ _readers["TSV"] = {"func": pd.read_csv, "params": {"sep": "\t"}}
@register_datareader @register_datareader
class DefaultDataReader(PipelineStepMixin): class DefaultDataReader(PipelineStepMixin, UpdateMetaMixin):
"""Mixin class for default data reader.""" """Mixin class for default data reader."""
def validate_input(self, input: List[str]) -> None: def validate_input(self, input: List[str]) -> None:
@ -46,8 +47,8 @@ class DefaultDataReader(PipelineStepMixin):
# Nothing to validate, any input is fine # Nothing to validate, any input is fine
pass pass
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input: List[str]) -> List[str]:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
@ -58,10 +59,10 @@ class DefaultDataReader(PipelineStepMixin):
Returns Returns
------- -------
list of str list of str
The updated list of output kinds, as reading possibilities. The updated list of output types, as reading possibilities.
""" """
# It will output the same kind of data as the input # It will output the same type of data as the input
return input return input
def fit_transform( def fit_transform(
@ -82,36 +83,33 @@ class DefaultDataReader(PipelineStepMixin):
------- -------
dict dict
The processed output as dictionary. The "data" key is added to The processed output as dictionary. The "data" key is added to
each data type dictionary except "meta". each data type dictionary.
""" """
# For each kind of data, try to read it # For each type of data, try to read it
out = input.copy() out = input.copy()
if params is None: if params is None:
params = {} params = {}
for kind in input.keys(): for type_ in input.keys():
if kind == "meta": if "path" not in input[type_]:
out["meta"] = input["meta"]
continue
if "path" not in input[kind]:
warn_with_log( warn_with_log(
f"Input kind {kind} does not provide a path. Skipping." f"Input type {type_} does not provide a path. Skipping."
) )
continue continue
t_path = input[kind]["path"] t_path = input[type_]["path"]
t_params = params.get(kind, {}) t_params = params.get(type_, {})
# Convert to Path if datareader is not well done # Convert to Path if datareader is not well done
if not isinstance(t_path, Path): if not isinstance(t_path, Path):
t_path = Path(t_path) t_path = Path(t_path)
out[kind]["path"] = t_path out[type_]["path"] = t_path
logger.info(f"Reading {kind} from {t_path.as_posix()}") logger.info(f"Reading {type_} from {t_path.as_posix()}")
fread = None fread = None
fname = t_path.name.lower() fname = t_path.name.lower()
for ext, ftype in _extensions.items(): for ext, ftype in _extensions.items():
if fname.endswith(ext): if fname.endswith(ext):
logger.info(f"{kind} is type {ftype}") logger.info(f"{type_} is type {ftype}")
reader_func = _readers[ftype]["func"] reader_func = _readers[ftype]["func"]
reader_params = _readers[ftype]["params"] reader_params = _readers[ftype]["params"]
if reader_params is not None: if reader_params is not None:
@ -123,8 +121,6 @@ class DefaultDataReader(PipelineStepMixin):
logger.info( logger.info(
f"Unknown file type {t_path.as_posix()}, skipping reading" f"Unknown file type {t_path.as_posix()}, skipping reading"
) )
out[kind]["data"] = fread out[type_]["data"] = fread
if "meta" not in out: self.update_meta(out[type_], "datareader")
out["meta"] = {}
out["meta"]["datareader"] = self.get_meta()
return out return out

View file

@ -17,37 +17,36 @@ from junifer.datareader import DefaultDataReader
@pytest.mark.parametrize( @pytest.mark.parametrize(
"kind", [["T1w", "BOLD", "T2", "dwi"], [], None, ["whatever"]] "type_", [["T1w", "BOLD", "T2", "dwi"], [], ["whatever"]]
) )
def test_validation(kind) -> None: def test_validation(type_) -> None:
"""Test validating input/output. """Test validating input/output.
Parameters Parameters
---------- ----------
kind : list of str or str or None type_ : list of str or str or None
The parametrized kind of data. The parametrized type_ of data.
""" """
reader = DefaultDataReader() reader = DefaultDataReader()
assert reader.validate_input(kind) is None assert reader.validate_input(type_) is None
assert reader.get_output_kind(kind) == kind assert reader.get_output_type(type_) == type_
assert reader.validate(kind) == kind assert reader.validate(type_) == type_
def test_meta() -> None: def test_meta() -> None:
"""Test reader metadata.""" """Test reader metadata."""
reader = DefaultDataReader() reader = DefaultDataReader()
t_meta = reader.get_meta()
assert t_meta["class"] == "DefaultDataReader"
nib_data_path = Path(nib_testing.data_path) nib_data_path = Path(nib_testing.data_path)
t_path = nib_data_path / "example4d.nii.gz" t_path = nib_data_path / "example4d.nii.gz"
input = {"BOLD": {"path": t_path}} input = {"BOLD": {"path": t_path}}
output = reader.fit_transform(input) output = reader.fit_transform(input)
assert "meta" in output assert "meta" in output["BOLD"]
assert "datareader" in output["meta"] meta = output["BOLD"]["meta"]
assert "class" in output["meta"]["datareader"] assert "datareader" in meta
assert output["meta"]["datareader"]["class"] == "DefaultDataReader" assert "class" in meta["datareader"]
assert meta["datareader"]["class"] == "DefaultDataReader"
@pytest.mark.parametrize( @pytest.mark.parametrize(

View file

@ -7,14 +7,15 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from ..pipeline import PipelineStepMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
from ..utils import logger, raise_error from ..utils import logger, raise_error
if TYPE_CHECKING: if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage from junifer.storage import BaseFeatureStorage
class BaseMarker(ABC, PipelineStepMixin): class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
"""Abstract base class for all markers. """Abstract base class for all markers.
Parameters Parameters
@ -57,27 +58,6 @@ class BaseMarker(ABC, PipelineStepMixin):
klass=NotImplementedError, klass=NotImplementedError,
) )
def get_meta(self, kind: str) -> Dict:
"""Get metadata.
Parameters
----------
kind : str
The kind of pipeline step.
Returns
-------
dict
The metadata as a dictionary with the only key 'marker'.
"""
s_meta = super().get_meta()
# same marker can be "fit"ted into different kinds, so the name
# is created from the kind and the name of the marker
s_meta["name"] = f"{kind}_{self.name}"
s_meta["kind"] = kind
return {"marker": s_meta}
def validate_input(self, input: List[str]) -> None: def validate_input(self, input: List[str]) -> None:
"""Validate input. """Validate input.
@ -101,23 +81,22 @@ class BaseMarker(ABC, PipelineStepMixin):
) )
@abstractmethod @abstractmethod
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the marker. The list must contain the The data type input to the marker.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of output kinds, as storage possibilities. The storage type output by the marker.
""" """
raise_error( raise_error(
msg="Concrete classes need to implement get_output_kind().", msg="Concrete classes need to implement get_output_type().",
klass=NotImplementedError, klass=NotImplementedError,
) )
@ -149,10 +128,9 @@ class BaseMarker(ABC, PipelineStepMixin):
klass=NotImplementedError, klass=NotImplementedError,
) )
@abstractmethod
def store( def store(
self, self,
kind: str, type_: str,
out: Dict[str, Any], out: Dict[str, Any],
storage: "BaseFeatureStorage", storage: "BaseFeatureStorage",
) -> None: ) -> None:
@ -160,18 +138,17 @@ class BaseMarker(ABC, PipelineStepMixin):
Parameters Parameters
---------- ----------
kind : str type_ : str
The data kind to store. The data type to store.
out : dict out : dict
The computed result as a dictionary to store. The computed result as a dictionary to store.
storage : storage-like storage : storage-like
The storage class. The storage class, for example, SQLiteFeatureStorage.
""" """
raise_error( output_type_ = self.get_output_type(type_)
msg="Concrete classes need to implement store().", logger.debug(f"Storing {output_type_} in {storage}")
klass=NotImplementedError, storage.store(kind=output_type_, **out)
)
def fit_transform( def fit_transform(
self, self,
@ -195,23 +172,25 @@ class BaseMarker(ABC, PipelineStepMixin):
""" """
out = {} out = {}
meta = input.get("meta", {}) for type_ in self._on:
for kind in self._on: if type_ in input.keys():
if kind in input.keys(): logger.info(f"Computing {type_}")
logger.info(f"Computing {kind}") t_input = input[type_]
t_input = input[kind]
extra_input = input.copy() extra_input = input.copy()
extra_input.pop(kind) extra_input.pop(type_)
t_meta = meta.copy() t_meta = t_input["meta"].copy()
t_meta.update(t_input.get("meta", {})) t_meta["type"] = type_
t_meta.update(self.get_meta(kind))
t_out = self.compute(input=t_input, extra_input=extra_input) t_out = self.compute(input=t_input, extra_input=extra_input)
t_out.update(meta=t_meta) t_out["meta"] = t_meta
self.update_meta(t_out, "marker")
if storage is not None: if storage is not None:
logger.info(f"Storing in {storage}") logger.info(f"Storing in {storage}")
self.store(kind=kind, out=t_out, storage=storage) self.store(type_=type_, out=t_out, storage=storage)
else: else:
logger.info("No storage specified, returning dictionary") logger.info("No storage specified, returning dictionary")
out[kind] = t_out out[type_] = t_out
return out return out

View file

@ -13,6 +13,7 @@ from ..pipeline import PipelineStepMixin
from ..storage.base import BaseFeatureStorage from ..storage.base import BaseFeatureStorage
from ..utils import logger from ..utils import logger
if TYPE_CHECKING: if TYPE_CHECKING:
from junifer.datagrabber import BaseDataGrabber from junifer.datagrabber import BaseDataGrabber

View file

@ -9,7 +9,6 @@ from typing import Any, Dict, List, Optional
import pandas as pd import pandas as pd
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..storage import BaseFeatureStorage
from ..utils import logger from ..utils import logger
from ..utils.logging import raise_error from ..utils.logging import raise_error
from .base import BaseMarker from .base import BaseMarker
@ -41,6 +40,8 @@ class CrossParcellationFC(BaseMarker):
(default None). (default None).
""" """
_DEPENDENCIES = {"nilearn"}
def __init__( def __init__(
self, self,
parcellation_one: str, parcellation_one: str,
@ -72,43 +73,21 @@ class CrossParcellationFC(BaseMarker):
""" """
return ["BOLD"] return ["BOLD"]
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the marker. The list must contain the The data type input to the marker.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of output kinds, as storage possibilities. The storage type output by the marker.
""" """
return ["matrix"] return "matrix"
def store(
self,
kind,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
kind : {"BOLD"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing BOLD-based marker in {storage}")
storage.store(kind="matrix", **out)
def compute( def compute(
self, self,

View file

@ -6,7 +6,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
import numpy as np import numpy as np
@ -16,9 +16,6 @@ from .base import BaseMarker
from .parcel_aggregation import ParcelAggregation from .parcel_aggregation import ParcelAggregation
from .utils import _ets from .utils import _ets
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker @register_marker
class RSSETSMarker(BaseMarker): class RSSETSMarker(BaseMarker):
@ -45,6 +42,8 @@ class RSSETSMarker(BaseMarker):
""" """
_DEPENDENCIES = {"nilearn"}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
@ -70,43 +69,21 @@ class RSSETSMarker(BaseMarker):
""" """
return ["BOLD"] return ["BOLD"]
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the marker. The list must contain the The data type input to the marker.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of output kinds, as storage possibilities. The storage type output by the marker.
""" """
return ["timeseries"] return "timeseries"
def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
kind : {"BOLD"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing BOLD in {storage}")
storage.store(kind="timeseries", **out)
def compute( def compute(
self, self,
@ -150,7 +127,7 @@ class RSSETSMarker(BaseMarker):
parcellation=self.parcellation, parcellation=self.parcellation,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
mask=self.mask mask=self.mask,
) )
# Compute the parcel aggregation # Compute the parcel aggregation
out = parcel_aggregation.compute(input=input, extra_input=extra_input) out = parcel_aggregation.compute(input=input, extra_input=extra_input)

View file

@ -4,19 +4,15 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..utils import logger
from .base import BaseMarker from .base import BaseMarker
from .parcel_aggregation import ParcelAggregation from .parcel_aggregation import ParcelAggregation
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker @register_marker
class FunctionalConnectivityParcels(BaseMarker): class FunctionalConnectivityParcels(BaseMarker):
@ -49,6 +45,8 @@ class FunctionalConnectivityParcels(BaseMarker):
None). None).
""" """
_DEPENDENCIES = {"nilearn", "scikit-learn"}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
@ -83,22 +81,21 @@ class FunctionalConnectivityParcels(BaseMarker):
""" """
return ["BOLD"] return ["BOLD"]
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the marker. The list must contain the The data type input to the marker.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of output kinds, as storage possibilities. The storage type output by the marker.
""" """
outputs = ["matrix"] return "matrix"
return outputs
def compute( def compute(
self, self,
@ -154,24 +151,3 @@ class FunctionalConnectivityParcels(BaseMarker):
out["col_names"] = ts["columns"] out["col_names"] = ts["columns"]
out["matrix_kind"] = "tril" out["matrix_kind"] = "tril"
return out return out
def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
kind : {"BOLD"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing {kind} in {storage}")
storage.store(kind="matrix", **out)

View file

@ -4,19 +4,16 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..utils import logger, raise_error from ..utils import raise_error
from .base import BaseMarker from .base import BaseMarker
from .sphere_aggregation import SphereAggregation from .sphere_aggregation import SphereAggregation
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker @register_marker
class FunctionalConnectivitySpheres(BaseMarker): class FunctionalConnectivitySpheres(BaseMarker):
@ -54,6 +51,8 @@ class FunctionalConnectivitySpheres(BaseMarker):
""" """
_DEPENDENCIES = {"nilearn", "scikit-learn"}
def __init__( def __init__(
self, self,
coords: str, coords: str,
@ -93,23 +92,21 @@ class FunctionalConnectivitySpheres(BaseMarker):
""" """
return ["BOLD"] return ["BOLD"]
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the marker. The list must contain the The data type input to the marker.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of output kinds, as storage possibilities. The storage type output by the marker.
""" """
outputs = ["matrix"] return "matrix"
return outputs
def compute( def compute(
self, self,
@ -166,24 +163,3 @@ class FunctionalConnectivitySpheres(BaseMarker):
out["col_names"] = ts["columns"] out["col_names"] = ts["columns"]
out["matrix_kind"] = "tril" out["matrix_kind"] = "tril"
return out return out
def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
kind : {"BOLD"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing {kind} in {storage}")
storage.store(kind="matrix", **out)

View file

@ -4,21 +4,18 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
import numpy as np import numpy as np
from nilearn.image import math_img, resample_to_img, new_img_like from nilearn.image import math_img, new_img_like, resample_to_img
from nilearn.maskers import NiftiMasker from nilearn.maskers import NiftiMasker
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..data import load_parcellation, load_mask from ..data import load_mask, load_parcellation
from ..stats import get_aggfunc_by_name from ..stats import get_aggfunc_by_name
from ..utils import logger from ..utils import logger
from .base import BaseMarker from .base import BaseMarker
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker @register_marker
class ParcelAggregation(BaseMarker): class ParcelAggregation(BaseMarker):
@ -48,6 +45,8 @@ class ParcelAggregation(BaseMarker):
None). None).
""" """
_DEPENDENCIES = {"nilearn", "numpy"}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]], parcellation: Union[str, List[str]],
@ -76,53 +75,27 @@ class ParcelAggregation(BaseMarker):
""" """
return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The kind of data to work on. The data type input to the marker.
Returns Returns
------- -------
list of str str
The list of storage kinds. The storage type output by the marker.
""" """
outputs = []
for t_input in input:
if t_input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
outputs.append("table")
elif t_input in ["BOLD"]:
outputs.append("timeseries")
else:
raise ValueError(f"Unknown input kind for {t_input}")
return outputs
def store( if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
self, return "table"
kind: str, elif input_type == "BOLD":
out: Dict[str, Any], return "timeseries"
storage: "BaseFeatureStorage", else:
) -> None: raise ValueError(f"Unknown input kind for {input_type}")
"""Store.
Parameters
----------
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing {kind} in {storage}")
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
storage.store(kind="table", **out)
elif kind in ["BOLD"]:
storage.store(kind="timeseries", **out)
def compute( def compute(
self, self,
@ -252,6 +225,4 @@ class ParcelAggregation(BaseMarker):
out_values = np.array(out_values).T out_values = np.array(out_values).T
out = {"data": out_values, "columns": out_labels} out = {"data": out_values, "columns": out_labels}
if out_values.shape[0] > 1:
out["row_names"] = "scan"
return out return out

View file

@ -4,7 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
from ..api.decorators import register_marker from ..api.decorators import register_marker
from ..data import load_coordinates, load_mask from ..data import load_coordinates, load_mask
@ -13,9 +13,6 @@ from ..stats import get_aggfunc_by_name
from ..utils import logger from ..utils import logger
from .base import BaseMarker from .base import BaseMarker
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker @register_marker
class SphereAggregation(BaseMarker): class SphereAggregation(BaseMarker):
@ -51,6 +48,8 @@ class SphereAggregation(BaseMarker):
""" """
_DEPENDENCIES = {"nilearn", "numpy"}
def __init__( def __init__(
self, self,
coords: str, coords: str,
@ -79,53 +78,27 @@ class SphereAggregation(BaseMarker):
""" """
return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The kind of data to work on. The data type input to the marker.
Returns Returns
------- -------
list of str str
The list of storage kinds. The storage type output by the marker.
""" """
outputs = []
for t_input in input:
if t_input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
outputs.append("table")
elif t_input in ["BOLD"]:
outputs.append("timeseries")
else:
raise ValueError(f"Unknown input kind for {t_input}")
return outputs
def store( if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
self, return "table"
kind: str, elif input_type == "BOLD":
out: Dict[str, Any], return "timeseries"
storage: "BaseFeatureStorage", else:
) -> None: raise ValueError(f"Unknown input kind for {input_type}")
"""Store.
Parameters
----------
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing {kind} in {storage}")
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
storage.store(kind="table", **out)
elif kind in ["BOLD"]:
storage.store(kind="timeseries", **out)
def compute( def compute(
self, self,
@ -180,6 +153,4 @@ class SphereAggregation(BaseMarker):
out_values = masker.fit_transform(t_input) out_values = masker.fit_transform(t_input)
# Format the output # Format the output
out = {"data": out_values, "columns": out_labels} out = {"data": out_values, "columns": out_labels}
if out_values.shape[0] > 1:
out["row_names"] = "scan"
return out return out

View file

@ -15,6 +15,7 @@ from junifer.markers.crossparcellation_functional_connectivity import (
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
parcellation_ONE = "Schaefer100x17" parcellation_ONE = "Schaefer100x17"
parcellation_TWO = "Schaefer200x17" parcellation_TWO = "Schaefer200x17"
@ -30,8 +31,7 @@ def test_compute() -> None:
"data": niimg, "data": niimg,
"path": out["BOLD"]["path"], "path": out["BOLD"]["path"],
"meta": {"element": "sub001"}, "meta": {"element": "sub001"},
}, }
"meta": {"element": "sub001"},
} }
crossparcellation = CrossParcellationFC( crossparcellation = CrossParcellationFC(
@ -43,12 +43,6 @@ def test_compute() -> None:
assert out["data"].shape == (200, 100) assert out["data"].shape == (200, 100)
assert len(out["col_names"]) == 100 assert len(out["col_names"]) == 100
assert len(out["row_names"]) == 200 assert len(out["row_names"]) == 200
meta = crossparcellation.get_meta("BOLD")["marker"]
assert meta["aggregation_method"] == "mean"
assert meta["class"] == "CrossParcellationFC"
assert meta["parcellation_one"] == "Schaefer100x17"
assert meta["parcellation_two"] == "Schaefer200x17"
assert meta["correlation_method"] == "spearman"
def test_store(tmp_path: Path) -> None: def test_store(tmp_path: Path) -> None:
@ -62,16 +56,10 @@ def test_store(tmp_path: Path) -> None:
""" """
with SPMAuditoryTestingDatagrabber() as dg: with SPMAuditoryTestingDatagrabber() as dg:
out = dg["sub001"] input_dict = dg["sub001"]
niimg = image.load_img(str(out["BOLD"]["path"].absolute())) niimg = image.load_img(str(input_dict["BOLD"]["path"].absolute()))
input_dict = {
"BOLD": { input_dict["BOLD"]["data"] = niimg
"data": niimg,
"path": out["BOLD"]["path"],
"meta": {"element": "sub001"},
},
"meta": {"element": "sub001"},
}
crossparcellation = CrossParcellationFC( crossparcellation = CrossParcellationFC(
parcellation_one=parcellation_ONE, parcellation_one=parcellation_ONE,
@ -80,19 +68,23 @@ def test_store(tmp_path: Path) -> None:
) )
uri = tmp_path / "test_crossparcellation.sqlite" uri = tmp_path / "test_crossparcellation.sqlite"
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
out = crossparcellation.fit_transform(input_dict, storage=storage) crossparcellation.fit_transform(input_dict, storage=storage)
features = storage.list_features()
assert any(
x["name"] == "BOLD_CrossParcellationFC"
for x in features.values()
)
def test_get_output_kind() -> None: def test_get_output_type() -> None:
"""Test CrossParcellationFC get_output_kind().""" """Test CrossParcellationFC get_output_type()."""
crossparcellation = CrossParcellationFC( crossparcellation = CrossParcellationFC(
parcellation_one=parcellation_ONE, parcellation_two=parcellation_TWO parcellation_one=parcellation_ONE, parcellation_two=parcellation_TWO
) )
input_list = ["BOLD"] input_ = "BOLD"
input_list = crossparcellation.get_output_kind(input_list) output = crossparcellation.get_output_type(input_)
assert len(input_list) == 1 assert output == "matrix"
assert input_list[0] in ["matrix"]
def test_init_() -> None: def test_init_() -> None:

View file

@ -16,6 +16,7 @@ from junifer.markers.ets_rss import RSSETSMarker
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
# Set parcellation # Set parcellation
PARCELLATION = "Schaefer100x17" PARCELLATION = "Schaefer100x17"
@ -42,21 +43,13 @@ def test_compute() -> None:
n_time, _ = test_ts.shape n_time, _ = test_ts.shape
assert n_time == len(new_out["data"]) assert n_time == len(new_out["data"])
# Assert the meta
meta = ets_rss_marker.get_meta("BOLD")["marker"]
assert meta["parcellation"] == "Schaefer100x17"
assert meta["agg_method"] == "mean"
assert meta["agg_method_params"] is None
assert meta["class"] == "RSSETSMarker"
def test_get_output_type() -> None:
def test_get_output_kind() -> None: """Test RSS ETS get_output_type()."""
"""Test RSS ETS get_output_kind()."""
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION)
input_list = ["BOLD"] input_ = "BOLD"
input_list = ets_rss_marker.get_output_kind(input_list) output = ets_rss_marker.get_output_type(input_)
assert len(input_list) == 1 assert output == "timeseries"
assert input_list[0] in ["timeseries"]
def test_store(tmp_path: Path) -> None: def test_store(tmp_path: Path) -> None:
@ -70,14 +63,18 @@ def test_store(tmp_path: Path) -> None:
""" """
with SPMAuditoryTestingDatagrabber() as dg: with SPMAuditoryTestingDatagrabber() as dg:
# Fetch element # Fetch element
out = dg["sub001"] elem = dg["sub001"]
# Load BOLD image # Load BOLD image
niimg = image.load_img(str(out["BOLD"]["path"].absolute())) niimg = image.load_img(str(elem["BOLD"]["path"].absolute()))
input_dict = {"data": niimg, "path": out["BOLD"]["path"]} elem["BOLD"]["data"] = niimg
# Compute the RSSETSMarker # Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION)
# Create storage # Create storage
storage = SQLiteFeatureStorage( storage = SQLiteFeatureStorage(
uri=str((tmp_path / "test.sqlite").absolute())) uri=str((tmp_path / "test.sqlite").absolute())
)
# Store # Store
ets_rss_marker.fit_transform(input=input_dict, storage=storage) ets_rss_marker.fit_transform(input=elem, storage=storage)
features = storage.list_features()
assert any(x["name"] == "BOLD_RSSETSMarker" for x in features.values())

View file

@ -32,7 +32,7 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
fmri_img = image.concat_imgs(ni_data.func) # type: ignore fmri_img = image.concat_imgs(ni_data.func) # type: ignore
fc = FunctionalConnectivityParcels(parcellation="Schaefer100x7") fc = FunctionalConnectivityParcels(parcellation="Schaefer100x7")
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}}) all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": {}}})
out = all_out["BOLD"] out = all_out["BOLD"]
@ -48,7 +48,11 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
pa = ParcelAggregation( pa = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="BOLD" parcellation="Schaefer100x7", method="mean", on="BOLD"
) )
ts = pa.compute({"data": fmri_img}) meta = {
"element": {"subject": "sub001"},
"dependencies": {"nilearn"},
}
ts = pa.compute({"data": fmri_img, "meta": meta})
# compare with nilearn # compare with nilearn
# Get the testing parcellation (for nilearn) # Get the testing parcellation (for nilearn)
@ -69,22 +73,24 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
assert_array_almost_equal(out_ni, out["data"], decimal=3) assert_array_almost_equal(out_ni, out["data"], decimal=3)
# check correct output # check correct output
assert fc.get_output_kind(["BOLD"]) == ["matrix"] assert fc.get_output_type("BOLD") == "matrix"
# Check empirical correlation method parameters # Check empirical correlation method parameters
fc = FunctionalConnectivityParcels( fc = FunctionalConnectivityParcels(
parcellation="Schaefer100x7", cor_method_params={"empirical": True} parcellation="Schaefer100x7", cor_method_params={"empirical": True}
) )
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}}) all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": meta}})
uri = tmp_path / "test_fc_parcellation.sqlite" uri = tmp_path / "test_fc_parcellation.sqlite"
# Single storage, must be the uri # Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = { meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}}
"element": "test", input = {"BOLD": {"data": fmri_img, "meta": meta}}
"version": "0.0.1",
"marker": {"name": "fcname"},
}
input = {"BOLD": {"data": fmri_img}, "meta": meta}
all_out = fc.fit_transform(input, storage=storage) all_out = fc.fit_transform(input, storage=storage)
features = storage.list_features()
assert any(
x["name"] == "BOLD_FunctionalConnectivityParcels"
for x in features.values()
)

View file

@ -36,7 +36,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
fc = FunctionalConnectivitySpheres( fc = FunctionalConnectivitySpheres(
coords="DMNBuckner", radius=5.0, cor_method="correlation" coords="DMNBuckner", radius=5.0, cor_method="correlation"
) )
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}}) all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": {}}})
out = all_out["BOLD"] out = all_out["BOLD"]
@ -52,7 +52,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
sa = SphereAggregation( sa = SphereAggregation(
coords="DMNBuckner", radius=5.0, method="mean", on="BOLD" coords="DMNBuckner", radius=5.0, method="mean", on="BOLD"
) )
ts = sa.compute({"data": fmri_img}) ts = sa.compute({"data": fmri_img, "meta": {}})
# Check that FC are almost equal when using nileran # Check that FC are almost equal when using nileran
cm = ConnectivityMeasure(kind="correlation") cm = ConnectivityMeasure(kind="correlation")
@ -60,19 +60,24 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
assert_array_almost_equal(out_ni, out["data"], decimal=3) assert_array_almost_equal(out_ni, out["data"], decimal=3)
# check correct output # check correct output
assert fc.get_output_kind(["BOLD"]) == ["matrix"] assert fc.get_output_type("BOLD") == "matrix"
uri = tmp_path / "test_fc_parcel.sqlite" uri = tmp_path / "test_fc_parcel.sqlite"
# Single storage, must be the uri # Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = { meta = {
"element": "test", "element": {"subject": "test"},
"version": "0.0.1", "dependencies": {"numpy", "nilearn"},
"marker": {"name": "fcname"},
} }
input = {"BOLD": {"data": fmri_img}, "meta": meta} input = {"BOLD": {"data": fmri_img, "meta": meta}}
all_out = fc.fit_transform(input, storage=storage) all_out = fc.fit_transform(input, storage=storage)
features = storage.list_features()
assert any(
x["name"] == "BOLD_FunctionalConnectivitySpheres"
for x in features.values()
)
def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None: def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
"""Test FunctionalConnectivitySpheres with empirical covariance. """Test FunctionalConnectivitySpheres with empirical covariance.
@ -94,7 +99,7 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
cor_method="correlation", cor_method="correlation",
cor_method_params={"empirical": True}, cor_method_params={"empirical": True},
) )
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}}) all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": {}}})
out = all_out["BOLD"] out = all_out["BOLD"]

View file

@ -19,10 +19,14 @@ def test_base_marker_subclassing() -> None:
"""Test proper subclassing of BaseMarker.""" """Test proper subclassing of BaseMarker."""
# Create concrete class # Create concrete class
class MyBaseMarker(BaseMarker): class MyBaseMarker(BaseMarker):
def __init__(self, on, name=None) -> None:
self.parameter = 1
super().__init__(on, name)
def get_valid_inputs(self): def get_valid_inputs(self):
return ["BOLD", "T1w"] return ["BOLD", "T1w"]
def get_output_kind(self, input): def get_output_type(self, input):
return ["timeseries"] return ["timeseries"]
def compute(self, input, extra_input): def compute(self, input, extra_input):
@ -32,18 +36,19 @@ def test_base_marker_subclassing() -> None:
"row_names": "row_names", "row_names": "row_names",
} }
def store(self, kind, out, storage):
return super().store(kind=kind, out=out, storage=storage)
with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"):
MyBaseMarker(on=["BOLD", "T2w"]) MyBaseMarker(on=["BOLD", "T2w"])
# Create input for marker # Create input for marker
input_ = { input_ = {
"meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"},
"BOLD": { "BOLD": {
"path": ".", "path": ".",
"data": "data", "data": "data",
"meta": {
"datagrabber": "dg",
"element": "elem",
"datareader": "dr",
},
}, },
} }
marker = MyBaseMarker(on=["BOLD"]) marker = MyBaseMarker(on=["BOLD"])
@ -57,14 +62,16 @@ def test_base_marker_subclassing() -> None:
assert "data" in output["BOLD"] assert "data" in output["BOLD"]
assert "columns" in output["BOLD"] assert "columns" in output["BOLD"]
assert "row_names" in output["BOLD"] assert "row_names" in output["BOLD"]
assert "meta" in output["BOLD"]
assert "datagrabber" in output["BOLD"]["meta"]
assert "element" in output["BOLD"]["meta"]
assert "datareader" in output["BOLD"]["meta"]
# Check no implementation check assert "meta" in output["BOLD"]
with pytest.raises(NotImplementedError): meta = output["BOLD"]["meta"]
marker.store(kind="kind", out="out", storage="storage") # type: ignore assert "datagrabber" in meta
assert "element" in meta
assert "datareader" in meta
assert "marker" in meta
assert "name" in meta["marker"]
assert "parameter" in meta["marker"]
assert meta["marker"]["parameter"] == 1
# Check attributes # Check attributes
assert marker.name == "MyBaseMarker" assert marker.name == "MyBaseMarker"

View file

@ -3,18 +3,21 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de> # Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path
import nibabel as nib import nibabel as nib
import numpy as np import numpy as np
import pytest import pytest
from pathlib import Path
from nilearn import datasets from nilearn import datasets
from nilearn.image import concat_imgs, math_img, resample_to_img, new_img_like from nilearn.image import concat_imgs, math_img, new_img_like, resample_to_img
from nilearn.maskers import NiftiLabelsMasker, NiftiMasker from nilearn.maskers import NiftiLabelsMasker, NiftiMasker
from numpy.testing import assert_array_almost_equal, assert_array_equal from numpy.testing import assert_array_almost_equal, assert_array_equal
from scipy.stats import trim_mean from scipy.stats import trim_mean
from junifer.data import load_mask, load_parcellation, register_parcellation from junifer.data import load_mask, load_parcellation, register_parcellation
from junifer.markers.parcel_aggregation import ParcelAggregation from junifer.markers.parcel_aggregation import ParcelAggregation
from junifer.storage import SQLiteFeatureStorage
def test_ParcelAggregation_input_output() -> None: def test_ParcelAggregation_input_output() -> None:
@ -22,12 +25,11 @@ def test_ParcelAggregation_input_output() -> None:
marker = ParcelAggregation( marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="VBM_GM" parcellation="Schaefer100x7", method="mean", on="VBM_GM"
) )
for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]:
output = marker.get_output_kind(["VBM_GM", "BOLD"]) assert marker.get_output_type(in_) == out_
assert output == ["table", "timeseries"]
with pytest.raises(ValueError, match="Unknown input"): with pytest.raises(ValueError, match="Unknown input"):
marker.get_output_kind(["VBM_GM", "BOLD", "unknown"]) marker.get_output_type("unknown")
def test_ParcelAggregation_3D() -> None: def test_ParcelAggregation_3D() -> None:
@ -78,22 +80,13 @@ def test_ParcelAggregation_3D() -> None:
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values3d_mean.ndim == 2 assert jun_values3d_mean.ndim == 2
assert jun_values3d_mean.shape[0] == 1 assert jun_values3d_mean.shape[0] == 1
assert_array_equal(manual, jun_values3d_mean) assert_array_equal(manual, jun_values3d_mean)
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Test using another function (std) # Test using another function (std)
manual = [] manual = []
for t_v in sorted(np.unique(parcellation_values)): for t_v in sorted(np.unique(parcellation_values)):
@ -103,22 +96,13 @@ def test_ParcelAggregation_3D() -> None:
# Use the ParcelAggregation object # Use the ParcelAggregation object
marker = ParcelAggregation(parcellation="Schaefer100x7", method="std") marker = ParcelAggregation(parcellation="Schaefer100x7", method="std")
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"] jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values3d_std.ndim == 2 assert jun_values3d_std.ndim == 2
assert jun_values3d_std.shape[0] == 1 assert jun_values3d_std.shape[0] == 1
assert_array_equal(manual, jun_values3d_std) assert_array_equal(manual, jun_values3d_std)
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "std"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Test using another function with parameters # Test using another function with parameters
manual = [] manual = []
for t_v in sorted(np.unique(parcellation_values)): for t_v in sorted(np.unique(parcellation_values)):
@ -136,22 +120,13 @@ def test_ParcelAggregation_3D() -> None:
method="trim_mean", method="trim_mean",
method_params={"proportiontocut": 0.1}, method_params={"proportiontocut": 0.1},
) )
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"] jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values3d_tm.ndim == 2 assert jun_values3d_tm.ndim == 2
assert jun_values3d_tm.shape[0] == 1 assert jun_values3d_tm.shape[0] == 1
assert_array_equal(manual, jun_values3d_tm) assert_array_equal(manual, jun_values3d_tm)
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "trim_mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_ParcelAggregation"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {"proportiontocut": 0.1}
def test_ParcelAggregation_4D(): def test_ParcelAggregation_4D():
"""Test ParcelAggregation object on 4D images.""" """Test ParcelAggregation object on 4D images."""
@ -170,21 +145,63 @@ def test_ParcelAggregation_4D():
# Create ParcelAggregation object # Create ParcelAggregation object
marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean") marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
input = dict(BOLD=dict(data=fmri_img)) input = {"BOLD": {"data": fmri_img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"] jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2 assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d) assert_array_equal(auto4d, jun_values4d)
meta = marker.get_meta("BOLD")["marker"]
assert meta["method"] == "mean" def test_ParcelAggregation_storage(tmp_path: Path) -> None:
assert meta["parcellation"] == ["Schaefer100x7"] """Test ParcelAggregation storage.
assert meta["mask"] is None
assert meta["name"] == "BOLD_ParcelAggregation" Parameters
assert meta["class"] == "ParcelAggregation" ----------
assert meta["kind"] == "BOLD" tmp_path : pathlib.Path
assert meta["method_params"] == {} The path to the test directory.
"""
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
uri = tmp_path / "test_sphere_storage_3D.sqlite"
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = {
"element": {"subject": "sub-01", "session": "ses-01"},
"dependencies": {"nilearn", "nibabel"},
}
input = {"VBM_GM": {"data": img, "meta": meta}}
marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="VBM_GM"
)
marker.fit_transform(input, storage=storage)
features = storage.list_features()
assert any(
x["name"] == "VBM_GM_ParcelAggregation" for x in features.values()
)
meta = {
"element": {"subject": "sub-01", "session": "ses-01"},
"dependencies": {"nilearn", "nibabel"},
}
# Get the SPM auditory data
subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore
input = {"BOLD": {"data": fmri_img, "meta": meta}}
marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="BOLD"
)
marker.fit_transform(input, storage=storage)
features = storage.list_features()
assert any(
x["name"] == "BOLD_ParcelAggregation" for x in features.values()
)
def test_ParcelAggregation_3D_mask() -> None: def test_ParcelAggregation_3D_mask() -> None:
@ -215,22 +232,13 @@ def test_ParcelAggregation_3D_mask() -> None:
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values3d_mean.ndim == 2 assert jun_values3d_mean.ndim == 2
assert jun_values3d_mean.shape[0] == 1 assert jun_values3d_mean.shape[0] == 1
assert_array_almost_equal(auto, jun_values3d_mean) assert_array_almost_equal(auto, jun_values3d_mean)
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] == "GM_prob0.2"
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None: def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
"""Test ParcelAggregation with multiple non-overlapping parcellations. """Test ParcelAggregation with multiple non-overlapping parcellations.
@ -281,7 +289,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
orig_mean = marker_original.fit_transform(input)["VBM_GM"] orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
@ -290,15 +298,6 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
assert orig_mean_data.shape[1] == 100 assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean) # assert_array_almost_equal(auto, jun_values3d_mean)
meta = marker_original.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Use the ParcelAggregation object on the two parcellations # Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation( marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"], parcellation=["Schaefer100x7_low", "Schaefer100x7_high"],
@ -306,7 +305,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
split_mean = marker_split.fit_transform(input)["VBM_GM"] split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
@ -314,15 +313,6 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
assert split_mean_data.shape[0] == 1 assert split_mean_data.shape[0] == 1
assert split_mean_data.shape[1] == 100 assert split_mean_data.shape[1] == 100
meta = marker_split.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7_low", "Schaefer100x7_high"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Data and labels should be the same # Data and labels should be the same
assert_array_equal(orig_mean_data, split_mean_data) assert_array_equal(orig_mean_data, split_mean_data)
assert orig_mean["columns"] == split_mean["columns"] assert orig_mean["columns"] == split_mean["columns"]
@ -379,7 +369,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
orig_mean = marker_original.fit_transform(input)["VBM_GM"] orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"] orig_mean_data = orig_mean["data"]
@ -388,15 +378,6 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
assert orig_mean_data.shape[1] == 100 assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean) # assert_array_almost_equal(auto, jun_values3d_mean)
meta = marker_original.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == ["Schaefer100x7"]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Use the ParcelAggregation object on the two parcellations # Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation( marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"], parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"],
@ -404,7 +385,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
name="gmd_schaefer100x7_mean", name="gmd_schaefer100x7_mean",
on="VBM_GM", on="VBM_GM",
) # Test passing "on" as a keyword argument ) # Test passing "on" as a keyword argument
input = dict(VBM_GM=dict(data=img)) input = {"VBM_GM": {"data": img, "meta": {}}}
split_mean = marker_split.fit_transform(input)["VBM_GM"] split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"] split_mean_data = split_mean["data"]
@ -412,18 +393,6 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
assert split_mean_data.shape[0] == 1 assert split_mean_data.shape[0] == 1
assert split_mean_data.shape[1] == 100 assert split_mean_data.shape[1] == 100
meta = marker_split.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["parcellation"] == [
"Schaefer100x7_low2",
"Schaefer100x7_high2",
]
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
assert meta["class"] == "ParcelAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
# Data should be the same # Data should be the same
assert_array_equal(orig_mean_data, split_mean_data) assert_array_equal(orig_mean_data, split_mean_data)

View file

@ -3,6 +3,8 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de> # Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# License: AGPL # License: AGPL
import typing
from typing import Dict
from pathlib import Path from pathlib import Path
import nibabel as nib import nibabel as nib
@ -16,6 +18,7 @@ from junifer.data import load_coordinates, load_mask
from junifer.markers.sphere_aggregation import SphereAggregation from junifer.markers.sphere_aggregation import SphereAggregation
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
# Define common variables # Define common variables
COORDS = "DMNBuckner" COORDS = "DMNBuckner"
RADIUS = 8 RADIUS = 8
@ -23,15 +26,12 @@ RADIUS = 8
def test_SphereAggregation_input_output() -> None: def test_SphereAggregation_input_output() -> None:
"""Test SphereAggregation input and output types.""" """Test SphereAggregation input and output types."""
marker = SphereAggregation( marker = SphereAggregation(coords="DMNBuckner", method="mean", on="VBM_GM")
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]:
) assert marker.get_output_type(in_) == out_
output = marker.get_output_kind(["VBM_GM", "BOLD"])
assert output == ["table", "timeseries"]
with pytest.raises(ValueError, match="Unknown input"): with pytest.raises(ValueError, match="Unknown input"):
marker.get_output_kind(["VBM_GM", "BOLD", "unknown"]) marker.get_output_type("unknown")
def test_SphereAggregation_3D() -> None: def test_SphereAggregation_3D() -> None:
@ -52,23 +52,13 @@ def test_SphereAggregation_3D() -> None:
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
) )
input = {"VBM_GM": {"data": img}} input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"] jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values4d.ndim == 2 assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d) assert_array_equal(auto4d, jun_values4d)
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["coords"] == COORDS
assert meta["radius"] == RADIUS
assert meta["mask"] is None
assert meta["name"] == "VBM_GM_SphereAggregation"
assert meta["class"] == "SphereAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}
def test_SphereAggregation_4D() -> None: def test_SphereAggregation_4D() -> None:
"""Test SphereAggregation object on 4D images.""" """Test SphereAggregation object on 4D images."""
@ -84,26 +74,14 @@ def test_SphereAggregation_4D() -> None:
auto4d = nifti_masker.fit_transform(fmri_img) auto4d = nifti_masker.fit_transform(fmri_img)
# Create SphereAggregation object # Create SphereAggregation object
marker = SphereAggregation( marker = SphereAggregation(coords=COORDS, method="mean", radius=RADIUS)
coords=COORDS, method="mean", radius=RADIUS input = {"BOLD": {"data": fmri_img, "meta": {}}}
)
input = {"BOLD": {"data": fmri_img}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"] jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2 assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d) assert_array_equal(auto4d, jun_values4d)
meta = marker.get_meta("BOLD")["marker"]
assert meta["method"] == "mean"
assert meta["coords"] == COORDS
assert meta["radius"] == RADIUS
assert meta["mask"] is None
assert meta["name"] == "BOLD_SphereAggregation"
assert meta["class"] == "SphereAggregation"
assert meta["kind"] == "BOLD"
assert meta["method_params"] == {}
def test_SphereAggregation_storage(tmp_path: Path) -> None: def test_SphereAggregation_storage(tmp_path: Path) -> None:
"""Test SphereAggregation storage. """Test SphereAggregation storage.
@ -122,31 +100,38 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
meta = { meta = {
"element": "test", "element": {"subject": "sub-01", "session": "ses-01"},
"version": "0.0.1", "dependencies": {"nilearn", "nibabel"},
"marker": {"name": "fcname"},
} }
input = {"VBM_GM": {"data": img}, "meta": meta} input = {"VBM_GM": {"data": img, "meta": meta}}
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
) )
marker.fit_transform(input, storage=storage) marker.fit_transform(input, storage=storage)
features: Dict = typing.cast(Dict, storage.list_features())
assert any(
x["name"] == "VBM_GM_SphereAggregation" for x in features.values()
)
meta = { meta = {
"element": "test", "element": {"subject": "sub-01", "session": "ses-01"},
"version": "0.0.1", "dependencies": {"nilearn", "nibabel"},
"marker": {"name": "BOLD_fcname"},
} }
# Get the SPM auditory data # Get the SPM auditory data
subject_data = datasets.fetch_spm_auditory() subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore fmri_img = concat_imgs(subject_data.func) # type: ignore
input = {"BOLD": {"data": fmri_img}, "meta": meta} input = {"BOLD": {"data": fmri_img, "meta": meta}}
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="BOLD" coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
) )
marker.fit_transform(input, storage=storage) marker.fit_transform(input, storage=storage)
features: Dict = typing.cast(Dict, storage.list_features())
assert any(
x["name"] == "BOLD_SphereAggregation" for x in features.values()
)
def test_SphereAggregation_3D_mask() -> None: def test_SphereAggregation_3D_mask() -> None:
@ -164,26 +149,21 @@ def test_SphereAggregation_3D_mask() -> None:
# Create NiftSpheresMasker # Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker( nifti_masker = NiftiSpheresMasker(
seeds=coordinates, radius=RADIUS, mask_img=mask_img) seeds=coordinates, radius=RADIUS, mask_img=mask_img
)
auto4d = nifti_masker.fit_transform(img) auto4d = nifti_masker.fit_transform(img)
# Create SphereAggregation object # Create SphereAggregation object
marker = SphereAggregation( marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM", coords=COORDS,
mask="GM_prob0.2" method="mean",
radius=RADIUS,
on="VBM_GM",
mask="GM_prob0.2",
) )
input = {"VBM_GM": {"data": img}} input = {"VBM_GM": {"data": img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"] jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
assert jun_values4d.ndim == 2 assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d) assert_array_equal(auto4d, jun_values4d)
meta = marker.get_meta("VBM_GM")["marker"]
assert meta["method"] == "mean"
assert meta["coords"] == COORDS
assert meta["radius"] == RADIUS
assert meta["name"] == "VBM_GM_SphereAggregation"
assert meta["class"] == "SphereAggregation"
assert meta["kind"] == "VBM_GM"
assert meta["method_params"] == {}

View file

@ -5,3 +5,4 @@
from . import registry from . import registry
from .pipeline_step_mixin import PipelineStepMixin from .pipeline_step_mixin import PipelineStepMixin
from .update_meta_mixin import UpdateMetaMixin

View file

@ -4,6 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from importlib.util import find_spec
from typing import Dict, List from typing import Dict, List
from ..utils import raise_error from ..utils import raise_error
@ -12,22 +13,6 @@ from ..utils import raise_error
class PipelineStepMixin: class PipelineStepMixin:
"""Mixin class for pipeline.""" """Mixin class for pipeline."""
def get_meta(self) -> Dict:
"""Get metadata.
Returns
-------
dict
The metadata as a dictionary.
"""
t_meta = {}
t_meta["class"] = self.__class__.__name__
for k, v in vars(self).items():
if not k.startswith("_"):
t_meta[k] = v
return t_meta
def validate_input(self, input: List[str]) -> None: def validate_input(self, input: List[str]) -> None:
"""Validate the input to the pipeline step. """Validate the input to the pipeline step.
@ -48,24 +33,22 @@ class PipelineStepMixin:
klass=NotImplementedError, klass=NotImplementedError,
) )
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get the kind of the pipeline step. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the pipeline step. The list must contain the The data type input to the marker.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of available Junifer Data dictionary keys after The storage type output by the marker.
the pipeline step.
""" """
raise_error( raise_error(
msg="Concrete classes need to implement get_output_kind().", msg="Concrete classes need to implement get_output_type().",
klass=NotImplementedError, klass=NotImplementedError,
) )
@ -85,11 +68,30 @@ class PipelineStepMixin:
Raises Raises
------ ------
ValueError ValueError
If the input does not have the required data. If the pipeline step object is missing dependencies required for
its working or if the input does not have the required data.
""" """
# Check if _DEPENDENCIES attribute is found;
# (markers and preprocessors will have them but not datareaders
# as of now)
dependencies_not_found = []
if hasattr(self, "_DEPENDENCIES"):
# Check if dependencies are importable
for dependency in self._DEPENDENCIES: # type: ignore
if find_spec(dependency) is None:
dependencies_not_found.append(dependency)
# Raise error if any dependency is not found
if dependencies_not_found:
raise_error(
msg=f"{dependencies_not_found} are not installed but are "
"required for using {self.name}.",
klass=ImportError,
)
self.validate_input(input=input) self.validate_input(input=input)
return self.get_output_kind(input=input) outputs = [self.get_output_type(t_input) for t_input in input]
return outputs
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]: def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
"""Fit and transform. """Fit and transform.

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Union
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
if TYPE_CHECKING: if TYPE_CHECKING:
from ..datagrabber import BaseDataGrabber from ..datagrabber import BaseDataGrabber
from ..storage import BaseFeatureStorage from ..storage import BaseFeatureStorage

View file

@ -4,6 +4,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Dict, List
import pytest import pytest
from junifer.pipeline.pipeline_step_mixin import PipelineStepMixin from junifer.pipeline.pipeline_step_mixin import PipelineStepMixin
@ -15,13 +17,49 @@ def test_PipelineStepMixin() -> None:
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
mixin.validate_input([]) mixin.validate_input([])
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
mixin.get_output_kind([]) mixin.get_output_type("")
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
mixin.fit_transform({}) mixin.fit_transform({})
def test_pipeline_step_mixin_meta(): def test_pipeline_step_mixin_validate_correct_dependencies() -> None:
"""Test metadata for PipelineStepMixin.""" """Test validate with correct dependencies."""
pipemixin = PipelineStepMixin()
t_meta = pipemixin.get_meta() class CorrectMixer(PipelineStepMixin):
assert t_meta["class"] == "PipelineStepMixin" """Test class for validation."""
_DEPENDENCIES = {"setuptools"}
def validate_input(self, input: List[str]) -> None:
print(input)
def get_output_type(self, input_type: str) -> str:
return input_type
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
return {"input": input}
mixer = CorrectMixer()
mixer.validate([])
def test_pipeline_step_mixin_validate_incorrect_dependencies() -> None:
"""Test validate with incorrect dependencies."""
class IncorrectMixer(PipelineStepMixin):
"""Test class for validation."""
_DEPENDENCIES = {"foobar"}
def validate_input(self, input: List[str]) -> None:
print(input)
def get_output_type(self, input_type: str) -> str:
return input_type
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
return {"input": input}
mixer = IncorrectMixer()
with pytest.raises(ImportError, match="not installed"):
mixer.validate([])

View file

@ -0,0 +1,51 @@
"""Provide tests for update meta mixin."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Dict, List, Set, Union
import pytest
from junifer.pipeline.update_meta_mixin import UpdateMetaMixin
@pytest.mark.parametrize(
"input, step_name, dependencies, expected",
[
({}, "step_name", None, set()),
({}, "step_name", ["numpy"], set(["numpy"])),
({}, "step_name", "numpy", set(["numpy"])),
({}, "step_name", set(["numpy"]), set(["numpy"])),
],
)
def test_UpdateMetaMixin(
input: Dict,
step_name: str,
dependencies: Union[Set, List, str, None],
expected: Set,
) -> None:
"""Test UpdateMetaMixin.
Parameters
----------
input : dict
The data object to update.
step_name : str
The name of the pipeline step.
dependencies : set, list, None, str
The dependencies of the pipeline step.
expected : set
The expected dependencies.
"""
class TestUpdateMetaMixin(UpdateMetaMixin):
"""Test UpdateMetaMixin."""
_DEPENDENCIES = dependencies
obj = TestUpdateMetaMixin()
obj.update_meta(input, step_name=step_name)
assert input["meta"]["step_name"]["class"] == "TestUpdateMetaMixin"
assert input["meta"]["dependencies"] == expected

View file

@ -0,0 +1,43 @@
"""Provide mixin class for updating metadata."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Dict
class UpdateMetaMixin:
"""Mixin class for updating meta."""
def update_meta(
self,
input: Dict,
step_name: str,
) -> None:
"""Update metadata.
Parameters
----------
input : dict
The data object to update.
step_name : str
The name of the pipeline step.
"""
t_meta = {}
t_meta["class"] = self.__class__.__name__
for k, v in vars(self).items():
if not k.startswith("_"):
t_meta[k] = v
if "meta" not in input:
input["meta"] = {}
input["meta"][step_name] = t_meta
if "dependencies" not in input["meta"]:
input["meta"]["dependencies"] = set()
dependencies = getattr(self, "_DEPENDENCIES", set())
if dependencies is not None:
if not isinstance(dependencies, (set, list)):
dependencies = set([dependencies])
input["meta"]["dependencies"].update(dependencies)

View file

@ -7,11 +7,11 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any, Dict, List, Optional, Tuple, Union
from ..pipeline import PipelineStepMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
from ..utils import logger, raise_error from ..utils import logger, raise_error
class BasePreprocessor(ABC, PipelineStepMixin): class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
"""Provide abstract base class for all preprocessors. """Provide abstract base class for all preprocessors.
Parameters Parameters
@ -58,8 +58,8 @@ class BasePreprocessor(ABC, PipelineStepMixin):
) )
@abstractmethod @abstractmethod
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input: List[str]) -> List[str]:
"""Get output kind. """Get output type.
Parameters Parameters
---------- ----------
@ -75,7 +75,7 @@ class BasePreprocessor(ABC, PipelineStepMixin):
""" """
raise_error( raise_error(
msg="Concrete classes need to implement get_output_kind().", msg="Concrete classes need to implement get_output_type().",
klass=NotImplementedError, klass=NotImplementedError,
) )
@ -93,22 +93,6 @@ class BasePreprocessor(ABC, PipelineStepMixin):
klass=NotImplementedError, klass=NotImplementedError,
) )
def get_meta(self, kind: str) -> Dict:
"""Get metadata.
Parameters
----------
kind : str
The kind of pipeline step.
Returns
-------
dict
The metadata as a dictionary with the only key 'preprocess'.
"""
s_meta = super().get_meta()
return {"preprocess": s_meta}
def fit_transform( def fit_transform(
self, self,
input: Dict[str, Dict], input: Dict[str, Dict],
@ -127,19 +111,17 @@ class BasePreprocessor(ABC, PipelineStepMixin):
""" """
out = input out = input
for kind in self._on: for type_ in self._on:
if kind in input.keys(): if type_ in input.keys():
logger.info(f"Computing {kind}") logger.info(f"Computing {type_}")
t_input = input[kind] t_input = input[type_]
extra_input = input.copy() extra_input = input.copy()
extra_input.pop(kind) extra_input.pop(type_)
t_meta = t_input.get("meta", {}) # input kind meta
t_meta.update(self.get_meta(kind))
key, t_out = self.preprocess( key, t_out = self.preprocess(
input=t_input, extra_input=extra_input input=t_input, extra_input=extra_input
) )
t_out.update(meta=t_meta)
out[key] = t_out out[key] = t_out
self.update_meta(out[key], "preprocess")
return out return out
@abstractmethod @abstractmethod

View file

@ -17,6 +17,7 @@ from ...api.decorators import register_preprocessor
from ...utils import logger, raise_error from ...utils import logger, raise_error
from ..base import BasePreprocessor from ..base import BasePreprocessor
if TYPE_CHECKING: if TYPE_CHECKING:
from nibabel import MGHImage, Nifti1Image, Nifti2Image from nibabel import MGHImage, Nifti1Image, Nifti2Image
@ -141,6 +142,8 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
""" """
_DEPENDENCIES = {"numpy", "nilearn"}
def __init__( def __init__(
self, self,
strategy: Optional[Dict[str, str]] = None, strategy: Optional[Dict[str, str]] = None,
@ -222,7 +225,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
klass=ValueError, klass=ValueError,
) )
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_type(self, input: List[str]) -> List[str]:
"""Get the kind of the pipeline step. """Get the kind of the pipeline step.
Parameters Parameters

View file

@ -62,7 +62,7 @@ def test_fMRIPrepConfoundRemover_validate_input() -> None:
confound_remover.validate_input(input) confound_remover.validate_input(input)
def test_fMRIPrepConfoundRemover_get_output_kind() -> None: def test_fMRIPrepConfoundRemover_get_output_type() -> None:
"""Test fMRIPrepConfoundRemover validate_input.""" """Test fMRIPrepConfoundRemover validate_input."""
confound_remover = fMRIPrepConfoundRemover() confound_remover = fMRIPrepConfoundRemover()
inputs = [ inputs = [
@ -72,7 +72,7 @@ def test_fMRIPrepConfoundRemover_get_output_kind() -> None:
] ]
# Confound remover works in place # Confound remover works in place
for input in inputs: for input in inputs:
assert confound_remover.get_output_kind(input) == input assert confound_remover.get_output_type(input) == input
def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None: def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None:
@ -460,8 +460,8 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None:
clean_bold = typing.cast(nib.Nifti1Image, clean_bold) clean_bold = typing.cast(nib.Nifti1Image, clean_bold)
# TODO: Find a better way to test functionality here # TODO: Find a better way to test functionality here
assert ( assert (
clean_bold.header.get_zooms() == # type: ignore clean_bold.header.get_zooms() # type: ignore
raw_bold.header.get_zooms() == raw_bold.header.get_zooms() # type: ignore
) )
assert clean_bold.get_fdata().shape == raw_bold.get_fdata().shape assert clean_bold.get_fdata().shape == raw_bold.get_fdata().shape
# TODO: Test confound remover with mask, needs #79 to be implemented # TODO: Test confound remover with mask, needs #79 to be implemented
@ -521,7 +521,6 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None:
AssertionError, assert_array_equal, orig_bold, trans_bold AssertionError, assert_array_equal, orig_bold, trans_bold
) )
assert output["meta"] == input["meta"] # general meta does not change
assert "meta" in output["BOLD"] assert "meta" in output["BOLD"]
assert "preprocess" in output["BOLD"]["meta"] assert "preprocess" in output["BOLD"]["meta"]
t_meta = output["BOLD"]["meta"]["preprocess"] t_meta = output["BOLD"]["meta"]["preprocess"]
@ -535,3 +534,7 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None:
assert t_meta["high_pass"] is None assert t_meta["high_pass"] is None
assert t_meta["t_r"] is None assert t_meta["t_r"] is None
assert t_meta["mask_img"] is None assert t_meta["mask_img"] is None
assert "dependencies" in output["BOLD"]["meta"]
dependencies = output["BOLD"]["meta"]["dependencies"]
assert dependencies == {"numpy", "nilearn"}

View file

@ -19,10 +19,14 @@ def test_base_preprocessor_subclassing() -> None:
"""Test proper subclassing of BasePreprocessor.""" """Test proper subclassing of BasePreprocessor."""
# Create concrete class # Create concrete class
class MyBasePreprocessor(BasePreprocessor): class MyBasePreprocessor(BasePreprocessor):
def __init__(self, on):
self.parameter = 1
super().__init__(on=on)
def get_valid_inputs(self): def get_valid_inputs(self):
return ["BOLD", "T1w"] return ["BOLD", "T1w"]
def get_output_kind(self, input): def get_output_type(self, input):
return ["timeseries"] return ["timeseries"]
def preprocess(self, input, extra_input): def preprocess(self, input, extra_input):
@ -37,14 +41,23 @@ def test_base_preprocessor_subclassing() -> None:
# Create input for marker # Create input for marker
input_ = { input_ = {
"meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"},
"BOLD": { "BOLD": {
"path": ".", "path": ".",
"data": "data", "data": "data",
"meta": {
"datagrabber": "dg",
"element": "elem",
"datareader": "dr",
},
}, },
"T1w": { "T1w": {
"path": ".", "path": ".",
"data": "data", "data": "data",
"meta": {
"datagrabber": "dg",
"element": "elem",
"datareader": "dr",
},
}, },
} }
prep = MyBasePreprocessor(on=["BOLD"]) prep = MyBasePreprocessor(on=["BOLD"])
@ -60,6 +73,13 @@ def test_base_preprocessor_subclassing() -> None:
assert "path" in output["BOLD"] assert "path" in output["BOLD"]
assert "meta" in output["BOLD"] assert "meta" in output["BOLD"]
meta = output["BOLD"]["meta"]
assert "preprocess" in meta
assert "class" in meta["preprocess"]
assert "MyBasePreprocessor" == meta["preprocess"]["class"]
assert "parameter" in meta["preprocess"]
assert 1 == meta["preprocess"]["parameter"]
assert "T1w" in output assert "T1w" in output
assert "data" in output["T1w"] assert "data" in output["T1w"]
assert output["T1w"]["data"] == "data" assert output["T1w"]["data"] == "data"

View file

@ -6,12 +6,13 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from pathlib import Path from pathlib import Path
from typing import Dict, List, Optional, Union from typing import Dict, Iterable, List, Optional, Union
import numpy as np
import pandas as pd import pandas as pd
from .._version import __version__
from ..utils import raise_error from ..utils import raise_error
from .utils import process_meta
class BaseFeatureStorage(ABC): class BaseFeatureStorage(ABC):
@ -40,23 +41,29 @@ class BaseFeatureStorage(ABC):
self.uri = uri self.uri = uri
if not isinstance(storage_types, list): if not isinstance(storage_types, list):
storage_types = [storage_types] storage_types = [storage_types]
if any(x not in self.get_valid_inputs() for x in storage_types):
wrong_storage_types = [
x for x in storage_types if x not in self.get_valid_inputs()
]
raise ValueError(
f"{self.__class__.__name__} cannot store {wrong_storage_types}"
)
self._valid_inputs = storage_types self._valid_inputs = storage_types
self.single_output = single_output self.single_output = single_output
def get_meta(self) -> Dict: def get_valid_inputs(self) -> List[str]:
"""Get metadata. """Get valid storage types for input.
Returns Returns
------- -------
dict list of str
The metadata as a dictionary. The list of storage types that can be used as input for this "
"storage.
""" """
meta = {} raise_error(
meta["versions"] = { msg="Concrete classes need to implement get_valid_inputs().",
"junifer": __version__, klass=NotImplementedError,
} )
return meta
def validate(self, input_: List[str]) -> None: def validate(self, input_: List[str]) -> None:
"""Validate the input to the pipeline step. """Validate the input to the pipeline step.
@ -80,23 +87,15 @@ class BaseFeatureStorage(ABC):
) )
@abstractmethod @abstractmethod
def list_features( def list_features(self) -> Dict:
self, return_df: bool = False
) -> Union[Dict[str, Dict], pd.DataFrame]:
"""List the features in the storage. """List the features in the storage.
Parameters
----------
return_df : bool, optional
If True, returns a pandas DataFrame. If False, returns a
dictionary (default False).
Returns Returns
------- -------
dict or pandas.DataFrame dict
List of features in the storage. If dictionary is returned, the List of features in the storage. The keys are the feature names to
keys are the feature names to be used in read_features() and the be used in read_features() and the values are the metadata of each
values are the metadata of each feature. feature.
""" """
raise_error( raise_error(
@ -131,19 +130,17 @@ class BaseFeatureStorage(ABC):
) )
@abstractmethod @abstractmethod
def store_metadata(self, meta: Dict) -> str: def store_metadata(self, meta_md5: str, element: Dict, meta: Dict) -> None:
"""Store metadata. """Store metadata.
Parameters Parameters
---------- ----------
meta_md5 : str
The metadata MD5 hash.
element : dict
The element as a dictionary.
meta : dict meta : dict
The metadata as a dictionary. The metadata as a dictionary.
Returns
-------
str
The metadata column.
""" """
raise_error( raise_error(
msg="Concrete classes need to implement store_metadata().", msg="Concrete classes need to implement store_metadata().",
@ -166,65 +163,117 @@ class BaseFeatureStorage(ABC):
If ``kind`` is invalid. If ``kind`` is invalid.
""" """
# Do the check before calling the abstract methods, otherwise the
# meta might be stored even if the data is not stored.
if kind not in self._valid_inputs:
raise_error(
msg=f"I don't know how to store {kind}.",
klass=ValueError,
)
t_meta = kwargs.pop("meta")
meta_md5, t_meta, t_element = process_meta(t_meta)
self.store_metadata(meta_md5=meta_md5, element=t_element, meta=t_meta)
if kind == "matrix": if kind == "matrix":
self.store_matrix(**kwargs) self.store_matrix(meta_md5=meta_md5, element=t_element, **kwargs)
elif kind == "timeseries": elif kind == "timeseries":
self.store_timeseries(**kwargs) self.store_timeseries(
meta_md5=meta_md5, element=t_element, **kwargs
)
elif kind == "table": elif kind == "table":
self.store_table(**kwargs) self.store_table(meta_md5=meta_md5, element=t_element, **kwargs)
else:
raise ValueError(f"I don't know how to store {kind}")
def store_df(self, **kwargs) -> None: def store_matrix(
"""Store pandas DataFrame. self,
meta_md5: str,
Parameters element: Dict,
---------- data: np.ndarray,
**kwargs : dict col_names: Optional[Iterable[str]] = None,
The keyword arguments. row_names: Optional[Iterable[str]] = None,
matrix_kind: Optional[str] = "full",
""" diagonal: bool = True,
raise_error( ) -> None:
msg="Concrete classes need to implement store_df().",
klass=NotImplementedError,
)
def store_matrix(self, **kwargs) -> None:
"""Store matrix. """Store matrix.
Parameters Parameters
---------- ----------
**kwargs : dict meta_md5 : str
The keyword arguments. The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray
The matrix data to store.
col_names : list or tuple of str, optional
The column names (default None).
row_names : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
matrix_kind : str, optional
The kind of matrix:
* ``triu`` : store upper triangular only
* ``tril`` : store lower triangular
* ``full`` : full matrix
(default "full").
diagonal : bool, optional
Whether to store the diagonal. If `matrix_kind` is "full", setting
this to False will raise an error (default True).
""" """
raise_error( raise_error(
msg="Concrete classes need to implement store_matrix2d().", msg="Concrete classes need to implement store_matrix2d().",
klass=NotImplementedError, klass=NotImplementedError,
) )
def store_table(self, **kwargs) -> None: def store_table(
self,
meta_md5: str,
element: Dict,
data: Union[np.ndarray, List],
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Store table. """Store table.
Parameters Parameters
---------- ----------
**kwargs : dict meta_md5 : str
The keyword arguments. The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray or list
The table data to store.
columns : list or tuple of str, optional
The columns (default None).
rows_col_name : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
""" """
raise_error( raise_error(
msg="Concrete classes need to implement store_table().", msg="Concrete classes need to implement store_table().",
klass=NotImplementedError, klass=NotImplementedError,
) )
def store_timeseries(self, **kwargs) -> None: def store_timeseries(
"""Store timeseries. self,
meta_md5: str,
element: Dict,
data: np.ndarray,
columns: Optional[Iterable[str]] = None,
) -> None:
"""Implement timeseries storing.
Parameters Parameters
---------- ----------
**kwargs : dict meta_md5 : str
The keyword arguments. The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray
The timeseries data to store.
columns : list or tuple of str, optional
The column labels (default None).
""" """
raise_error( raise_error(
msg="Concrete classes need to implement store_timeseries().", msg="Concrete classes need to implement store_timeseries().",

View file

@ -6,8 +6,9 @@
import json import json
from pathlib import Path from pathlib import Path
from typing import Dict, Union from typing import Any, Dict, Iterable, List, Optional, Union
import numpy as np
import pandas as pd import pandas as pd
from .base import BaseFeatureStorage from .base import BaseFeatureStorage
@ -39,6 +40,17 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
) -> None: ) -> None:
super().__init__(uri=uri, single_output=single_output, **kwargs) super().__init__(uri=uri, single_output=single_output, **kwargs)
def get_valid_inputs(self) -> List[str]:
"""Get valid storage types for input.
Returns
-------
list of str
The list of storage types that can be used as input for this "
"storage.
"""
return ["matrix", "table", "timeseries"]
def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame: def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame:
"""Convert the metadata to a pandas DataFrame. """Convert the metadata to a pandas DataFrame.
@ -57,8 +69,166 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
data_df = {} data_df = {}
for k, v in meta.items(): for k, v in meta.items():
data_df[k] = json.dumps(v, sort_keys=True) data_df[k] = json.dumps(v, sort_keys=True)
if "marker" in meta:
data_df["name"] = meta["marker"]["name"]
df = pd.DataFrame(data_df, index=[meta_md5]) df = pd.DataFrame(data_df, index=[meta_md5])
df.index.name = "meta_md5" df.index.name = "meta_md5"
return df return df
@staticmethod
def element_to_index(
element: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None
) -> pd.MultiIndex:
"""Convert the element metadata to index.
Parameters
----------
element : dict
The element as a dictionary.
n_rows : int, optional
Number of rows to create (default 1).
rows_col_name: str, optional
The column name to use in case `n_rows` > 1. If None and
n_rows > 1, the name will be "idx" (default None).
Returns
-------
pandas.MultiIndex
The index of the dataframe to store.
Raises
------
ValueError
If `meta` does not contain the key "element".
"""
# Check rows_col_name
if rows_col_name is None:
rows_col_name = "idx"
elem_idx: Dict[Any, Any] = {
k: [v] * n_rows for k, v in element.items()
}
elem_idx[rows_col_name] = np.arange(n_rows)
# Create index
index = pd.MultiIndex.from_frame(
pd.DataFrame(elem_idx, index=range(n_rows))
)
return index
def store_df(
self, meta_md5: str, element: Dict, df: Union[pd.DataFrame, pd.Series]
) -> None:
"""Implement pandas DataFrame storing.
Parameters
----------
df : pandas.DataFrame or pandas.Series
The pandas DataFrame or Series to store.
meta : dict
The metadata as a dictionary.
Raises
------
ValueError
If the dataframe index has items that are not in the index
generated from the metadata.
"""
raise NotImplementedError("Implement in subclass.")
def _store_2d(
self,
meta_md5: str,
element: Dict,
data: Union[np.ndarray, List],
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Store 2D dataframe.
Parameters
----------
meta_md5 : str
The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray or List
The data to store.
columns : list or tuple of str, optional
The columns (default None).
rows_col_name : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
"""
n_rows = len(data)
# Convert element metadata to index
idx = self.element_to_index(
element=element, n_rows=n_rows, rows_col_name=rows_col_name
)
# Prepare new dataframe
data_df = pd.DataFrame( # type: ignore
data, columns=columns, index=idx # type: ignore
)
# Store dataframe
self.store_df(meta_md5=meta_md5, element=element, df=data_df)
def store_table(
self,
meta_md5: str,
element: Dict,
data: Union[np.ndarray, List],
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Implement table storing.
Parameters
----------
meta_md5 : str
The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray or List
The table data to store.
columns : list or tuple of str, optional
The columns (default None).
rows_col_name : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
"""
self._store_2d(
meta_md5=meta_md5,
element=element,
data=data,
columns=columns,
rows_col_name=rows_col_name,
)
def store_timeseries(
self,
meta_md5: str,
element: Dict,
data: np.ndarray,
columns: Optional[Iterable[str]] = None,
) -> None:
"""Implement timeseries storing.
Parameters
----------
meta_md5 : str
The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray
The timeseries data to store.
columns : list or tuple of str, optional
The column labels (default None).
"""
self._store_2d(
meta_md5=meta_md5,
element=element,
data=data,
columns=columns,
rows_col_name="timepoint",
)

View file

@ -5,8 +5,9 @@
# License: AGPL # License: AGPL
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Dict, Iterable, List, Optional, Union from typing import TYPE_CHECKING, Dict, List, Optional, Union
import json
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from pandas.core.base import NoNewAttributesMixin from pandas.core.base import NoNewAttributesMixin
@ -17,7 +18,8 @@ from tqdm import tqdm
from ..api.decorators import register_storage from ..api.decorators import register_storage
from ..utils import logger, raise_error, warn_with_log from ..utils import logger, raise_error, warn_with_log
from .pandas_base import PandasBaseFeatureStorage from .pandas_base import PandasBaseFeatureStorage
from .utils import element_to_index, element_to_prefix, process_meta from .utils import element_to_prefix
if TYPE_CHECKING: if TYPE_CHECKING:
from sqlalchemy.engine import Engine from sqlalchemy.engine import Engine
@ -86,7 +88,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
# Set upsert # Set upsert
self._upsert = upsert self._upsert = upsert
def get_engine(self, meta: Optional[Dict] = None) -> "Engine": def get_engine(self, element: Optional[Dict] = None) -> "Engine":
"""Get engine. """Get engine.
Parameters Parameters
@ -100,11 +102,6 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
The sqlalchemy engine. The sqlalchemy engine.
""" """
# Set metadata as empty dictionary if None
if meta is None:
meta = {}
# Retrieve element key from metadata
element = meta.get("element", None)
# Prefixed elements # Prefixed elements
prefix = "" prefix = ""
if self.single_output is False: if self.single_output is False:
@ -208,58 +205,15 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
msg=f"Invalid option {if_exists} for if_exists." msg=f"Invalid option {if_exists} for if_exists."
) )
def _store_2d( def list_features(self) -> Dict:
self, """List the features in the storage.
data: Dict,
meta: Dict,
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Store 2D dataframe.
Parameters
----------
data : dict
The data to store.
meta : dict
The metadata as a dictionary.
columns : list or tuple of str, optional
The columns (default None).
rows_col_name : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
"""
n_rows = len(data)
# Convert element metadata to index
idx = element_to_index(
meta=meta, n_rows=n_rows, rows_col_name=rows_col_name
)
# Prepare new dataframe
data_df = pd.DataFrame( # type: ignore
data, columns=columns, index=idx # type: ignore
)
# Store dataframe
self.store_df(df=data_df, meta=meta)
def list_features(
self, return_df: bool = False
) -> Union[Dict, pd.DataFrame]:
"""Implement features listing from the storage.
Parameters
----------
return_df : bool, optional
If True, returns a pandas DataFrame. If False, returns a
dictionary (default False).
Returns Returns
------- -------
dict or pandas.DataFrame dict
List of features in the storage. If dictionary is returned, the List of features in the storage. The keys are the feature names to
keys are the feature names to be used in read_features() and the be used in read_features() and the values are the metadata of each
values are the metadata of each feature. feature.
""" """
meta_df = pd.read_sql( meta_df = pd.read_sql(
@ -267,10 +221,11 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
con=self.get_engine(), con=self.get_engine(),
index_col="meta_md5", index_col="meta_md5",
) )
out = meta_df meta_df.index = meta_df.index.str.replace(r"meta_", "")
# Return dictionary out = meta_df.to_dict(orient="index") # type: ignore
if return_df is False: for md5, t_meta in out.items():
out = meta_df.to_dict(orient="index") # type: ignore for k, v in t_meta.items():
out[md5][k] = json.loads(v)
return out return out
def read_df( def read_df(
@ -327,7 +282,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
con=engine, con=engine,
index_col="meta_md5", index_col="meta_md5",
) )
t_df = meta_df.query(f"name == '{feature_name}'")
# Wrap in double quotes as the fields are in JSON format
t_df = meta_df.query(f"name == '\"{feature_name}\"'")
if len(t_df) == 0: if len(t_df) == 0:
raise_error(msg=f"Feature {feature_name} not found") raise_error(msg=f"Feature {feature_name} not found")
elif len(t_df) > 1: elif len(t_df) > 1:
@ -339,8 +296,11 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
) )
) )
table_name = f"meta_{t_df.index[0]}" table_name = f"meta_{t_df.index[0]}"
if table_name not in inspect(engine).get_table_names():
raise_error(msg=f"Feature MD5 {feature_md5} not found")
# Read metadata from table # Read metadata from table
df = pd.read_sql(sql=table_name, con=engine) df = pd.read_sql(sql=table_name, con=engine)
# Read the index # Read the index
query = ( query = (
"SELECT ii.name FROM sqlite_master AS m, " "SELECT ii.name FROM sqlite_master AS m, "
@ -356,36 +316,30 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
df = df.set_index(index_names) df = df.set_index(index_names)
return df return df
def store_metadata(self, meta: Dict) -> str: def store_metadata(self, meta_md5: str, element: Dict, meta: Dict) -> None:
r"""Implement metadata storing in the storage. """Implement metadata storing in the storage.
Parameters Parameters
---------- ----------
meta_md5 : str
The metadata MD5 hash.
element : dict
The element as a dictionary.
meta : dict meta : dict
The metadata as a dictionary. The metadata as a dictionary.
Returns
-------
str
The MD5 hash of the metadata prefixed with "meta\_".
""" """
# Copy metadata
t_meta = meta.copy()
# Update metadata
t_meta.update(self.get_meta())
# Process metadata
meta_md5, t_meta_row = process_meta(t_meta)
# Get sqlalchemy engine # Get sqlalchemy engine
engine = self.get_engine(meta=t_meta) engine = self.get_engine(element=element)
if meta_md5 not in inspect(engine).get_table_names(): table_name = f"meta_{meta_md5}"
if table_name not in inspect(engine).get_table_names():
# Convert metadata to dataframe # Convert metadata to dataframe
meta_df = self._meta_row(meta=t_meta_row, meta_md5=meta_md5) meta_df = self._meta_row(meta=meta, meta_md5=meta_md5)
# Save dataframe # Save dataframe
self._save_upsert(meta_df, "meta", engine) self._save_upsert(meta_df, "meta", engine)
return f"meta_{meta_md5}"
def store_df(self, df: Union[pd.DataFrame, pd.Series], meta: Dict) -> None: def store_df(
self, meta_md5: str, element: Dict, df: Union[pd.DataFrame, pd.Series]
) -> None:
"""Implement pandas DataFrame storing. """Implement pandas DataFrame storing.
Parameters Parameters
@ -405,7 +359,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
# TODO: Test this function # TODO: Test this function
# Check that the index generated by meta matches the one in # Check that the index generated by meta matches the one in
# the dataframe. # the dataframe.
idx = element_to_index(meta) idx = self.element_to_index(element)
# Given the meta, we might not know if there is an extra column added # Given the meta, we might not know if there is an extra column added
# when storing a timeseries or 2d elements. We need to check if the # when storing a timeseries or 2d elements. We need to check if the
# extra element is only one. # extra element is only one.
@ -418,24 +372,25 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
elif len(extra) == 1: elif len(extra) == 1:
# The df has one extra index item, this should be the new name # The df has one extra index item, this should be the new name
# of the missing element in the index # of the missing element in the index
idx = element_to_index(meta, rows_col_name=extra[0]) idx = self.element_to_index(element, rows_col_name=extra[0])
if any(x not in df.index.names for x in idx.names): if any(x not in df.index.names for x in idx.names):
raise_error( raise_error(
"The index of the dataframe is missing index items that are " "The index of the dataframe is missing index items that are "
"generated from the metadata." "generated from the metadata."
) )
# Get table name
table_name = self.store_metadata(meta) table_name = f"meta_{meta_md5}"
# Get sqlalchemy engine # Get sqlalchemy engine
engine = self.get_engine(meta) engine = self.get_engine(element)
# Save data # Save data
self._save_upsert(df, table_name, engine) self._save_upsert(df, table_name, engine)
def store_matrix( def store_matrix(
self, self,
meta_md5: str,
element: Dict,
data: np.ndarray, data: np.ndarray,
meta: Dict,
col_names: Optional[List[str]] = None, col_names: Optional[List[str]] = None,
row_names: Optional[List[str]] = None, row_names: Optional[List[str]] = None,
matrix_kind: Optional[str] = "full", matrix_kind: Optional[str] = "full",
@ -445,6 +400,10 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
Parameters Parameters
---------- ----------
meta_md5 : str
The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray data : numpy.ndarray
The matrix data to store. The matrix data to store.
meta : dict meta : dict
@ -465,7 +424,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
(default "full"). (default "full").
diagonal : bool, optional diagonal : bool, optional
Whether to store the diagonal. If `matrix_kind` is "full", setting Whether to store the diagonal. If `matrix_kind` is "full", setting
this to False will raise an error (default True).. this to False will raise an error (default True).
""" """
if diagonal is False and matrix_kind not in ["triu", "tril"]: if diagonal is False and matrix_kind not in ["triu", "tril"]:
@ -519,74 +478,27 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
# Convert element metadata to index # Convert element metadata to index
n_rows = 1 n_rows = 1
idx = element_to_index(meta=meta, n_rows=n_rows, rows_col_name=None) idx = self.element_to_index(
element=element, n_rows=n_rows, rows_col_name=None
)
# Prepare new dataframe # Prepare new dataframe
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx) data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
if len(columns) > 2000: # TODO: check SQLITE_MAX_COLUMN if len(columns) > 2000:
warn_with_log(
msg="The number of columns is greater than 2000. "
"The data will be stored in long format. "
"This will make it slower to collect the data. "
"Future versions of junifer will provide additional storage "
"options that will not raise this warning.",
)
data_df = data_df.stack() data_df = data_df.stack()
new_names = [x for x in data_df.index.names[:-1]] new_names = [x for x in data_df.index.names[:-1]]
new_names.append("pair") new_names.append("pair")
data_df.index.names = new_names data_df.index.names = new_names
# Store dataframe # Store dataframe
self.store_df(df=data_df, meta=meta) self.store_df(meta_md5=meta_md5, element=element, df=data_df)
def store_table(
self,
data: Dict,
meta: Dict,
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Implement table storing.
Parameters
----------
data : dict
The table data to store.
meta : dict
The metadata as a dictionary.
columns : list or tuple of str, optional
The columns (default None).
rows_col_name : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
"""
self._store_2d(
data=data, meta=meta, columns=columns, rows_col_name=rows_col_name
)
def store_timeseries(
self,
data: Dict,
meta: Dict,
columns: Optional[Iterable[str]] = None,
row_names: str = "timepoint",
) -> None:
"""Implement timeseries storing.
Parameters
----------
data : dict
The timeseries data to store.
meta : dict
The metadata as a dictionary.
columns : list or tuple of str, optional
The column labels (default None).
row_names : str, optional
The column name to use in case number of rows greater than 1
(default "timepoint").
"""
self._store_2d(
data=data,
meta=meta,
columns=columns,
rows_col_name="timepoint", # explicit so as to stop overriding
)
def collect(self) -> None: def collect(self) -> None:
"""Implement data collection. """Implement data collection.

View file

@ -0,0 +1,70 @@
"""Provide tests for pandas base feature storage."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from junifer.storage.pandas_base import PandasBaseFeatureStorage
def test_element_to_index() -> None:
"""Test element to index."""
element = {"foo": "bar"}
index = PandasBaseFeatureStorage.element_to_index(element)
assert index.names == ["foo", "idx"]
assert index.levels[0].name == "foo"
assert index.levels[0].values[0] == "bar"
assert all(x == "bar" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "idx"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (1,)
index = PandasBaseFeatureStorage.element_to_index(element, n_rows=10)
assert index.names == ["foo", "idx"]
assert index.levels[0].name == "foo"
assert all(x == "bar" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "idx"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (10,)
index = PandasBaseFeatureStorage.element_to_index(
element, n_rows=1, rows_col_name="scan"
)
assert index.names == ["foo", "scan"]
assert index.levels[0].name == "foo"
assert index.levels[0].values[0] == "bar"
assert all(x == "bar" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "scan"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (1,)
index = PandasBaseFeatureStorage.element_to_index(
element, n_rows=7, rows_col_name="scan"
)
assert index.names == ["foo", "scan"]
assert index.levels[0].name == "foo"
assert all(x == "bar" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "scan"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (7,)
element = {"subject": "sub-01", "session": "ses-01"}
index = PandasBaseFeatureStorage.element_to_index(element, n_rows=10)
assert index.levels[0].name == "subject"
assert all(x == "sub-01" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "session"
assert all(x == "ses-01" for x in index.levels[1].values)
assert index.levels[1].values.shape == (1,)
assert index.levels[2].name == "idx"
assert all(x == i for i, x in enumerate(index.levels[2].values))
assert index.levels[2].values.shape == (10,)

View file

@ -15,47 +15,44 @@ from pandas.testing import assert_frame_equal
from sqlalchemy import create_engine from sqlalchemy import create_engine
from junifer.storage.sqlite import SQLiteFeatureStorage from junifer.storage.sqlite import SQLiteFeatureStorage
from junifer.storage.utils import ( from junifer.storage.utils import element_to_prefix, process_meta
element_to_index,
element_to_prefix,
process_meta,
)
df1 = pd.DataFrame( df1 = pd.DataFrame(
{ {
"element": [1, 2, 3, 4, 5], "subject": [1, 2, 3, 4, 5],
"pk2": ["a", "b", "c", "d", "e"], "pk2": ["a", "b", "c", "d", "e"],
"col1": [11, 22, 33, 44, 55], "col1": [11, 22, 33, 44, 55],
"col2": [111, 222, 333, 444, 555], "col2": [111, 222, 333, 444, 555],
} }
).set_index(["element", "pk2"]) ).set_index(["subject", "pk2"])
df2 = pd.DataFrame( df2 = pd.DataFrame(
{ {
"element": [2, 5, 6], "subject": [2, 5, 6],
"pk2": ["b", "e", "f"], "pk2": ["b", "e", "f"],
"col1": [2222, 5555, 66], "col1": [2222, 5555, 66],
"col2": [22222, 55555, 666], "col2": [22222, 55555, 666],
} }
).set_index(["element", "pk2"]) ).set_index(["subject", "pk2"])
df_update = pd.DataFrame( df_update = pd.DataFrame(
{ {
"element": [1, 2, 3, 4, 5, 6], "subject": [1, 2, 3, 4, 5, 6],
"pk2": ["a", "b", "c", "d", "e", "f"], "pk2": ["a", "b", "c", "d", "e", "f"],
"col1": [11, 2222, 33, 44, 5555, 66], "col1": [11, 2222, 33, 44, 5555, 66],
"col2": [111, 22222, 333, 444, 55555, 666], "col2": [111, 22222, 333, 444, 55555, 666],
} }
).set_index(["element", "pk2"]) ).set_index(["subject", "pk2"])
df_ignore = pd.DataFrame( df_ignore = pd.DataFrame(
{ {
"element": [1, 2, 3, 4, 5, 6], "subject": [1, 2, 3, 4, 5, 6],
"pk2": ["a", "b", "c", "d", "e", "f"], "pk2": ["a", "b", "c", "d", "e", "f"],
"col1": [11, 22, 33, 44, 55, 66], "col1": [11, 22, 33, 44, 55, 66],
"col2": [111, 222, 333, 444, 555, 666], "col2": [111, 222, 333, 444, 555, 666],
} }
).set_index(["element", "pk2"]) ).set_index(["subject", "pk2"])
def _read_sql( def _read_sql(
@ -148,15 +145,14 @@ def test_upsert_replace(tmp_path: Path) -> None:
uri = tmp_path / "test_upsert_replace.sqlite" uri = tmp_path / "test_upsert_replace.sqlite"
# Single storage, must be the uri # Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
# Metadata to store
meta = {"element": "test", "version": "0.0.1"}
# Save to database # Save to database
storage.store_df(df=df1, meta=meta) storage.store_df(
# Store metadata meta_md5="table_name", element={"subject": "test"}, df=df1
table_name = storage.store_metadata(meta) )
table_name = "meta_table_name"
# Read stored table # Read stored table
c_df1 = _read_sql( c_df1 = _read_sql(
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(df1, c_df1) assert_frame_equal(df1, c_df1)
@ -164,7 +160,7 @@ def test_upsert_replace(tmp_path: Path) -> None:
storage._save_upsert(df=df2, name=table_name, if_exists="replace") storage._save_upsert(df=df2, name=table_name, if_exists="replace")
# Read stored table # Read stored table
c_df2 = _read_sql( c_df2 = _read_sql(
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(df2, c_df2) assert_frame_equal(df2, c_df2)
@ -182,24 +178,25 @@ def test_upsert_ignore(tmp_path: Path) -> None:
uri = tmp_path / "test_upsert_ignore.sqlite" uri = tmp_path / "test_upsert_ignore.sqlite"
# Single storage, must be the uri # Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
# Metadata to store
meta = {"element": "test", "version": "0.0.1"}
# Save to database # Save to database
storage.store_df(df=df1, meta=meta) storage.store_df(
# Store metadata meta_md5="table_name", element={"subject": "test"}, df=df1
table_name = storage.store_metadata(meta=meta) )
table_name = "meta_table_name"
# Read stored table # Read stored table
c_df1 = _read_sql( c_df1 = _read_sql(
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(df1, c_df1) assert_frame_equal(df1, c_df1)
# Check for warning # Check for warning
with pytest.warns(RuntimeWarning, match="are already present"): with pytest.warns(RuntimeWarning, match="are already present"):
storage.store_df(df2, meta) storage.store_df(
meta_md5="table_name", element={"subject": "test"}, df=df2
)
# Read stored table # Read stored table
c_dfignore = _read_sql( c_dfignore = _read_sql(
table_name, uri=uri.as_posix(), index_col=["element", "pk2"] table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(c_dfignore, df_ignore) assert_frame_equal(c_dfignore, df_ignore)
@ -219,23 +216,24 @@ def test_upsert_update(tmp_path: Path) -> None:
""" """
uri = tmp_path / "test_upsert_delete.sqlite" uri = tmp_path / "test_upsert_delete.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Metadata to store
meta = {"element": "test", "version": "0.0.1"}
# Save to database # Save to database
storage.store_df(df1, meta) storage.store_df(
# Store metadata meta_md5="table_name", element={"subject": "test"}, df=df1
table_name = storage.store_metadata(meta) )
table_name = "meta_table_name"
# Read stored table # Read stored table
c_df1 = _read_sql( c_df1 = _read_sql(
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(df1, c_df1) assert_frame_equal(df1, c_df1)
# Save to database # Save to database
storage.store_df(df2, meta) storage.store_df(
meta_md5="table_name", element={"subject": "test"}, df=df2
)
# Read stored table # Read stored table
c_dfupdate = _read_sql( c_dfupdate = _read_sql(
table_name, uri=uri.as_posix(), index_col=["element", "pk2"] table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(c_dfupdate, df_update) assert_frame_equal(c_dfupdate, df_update)
@ -268,50 +266,66 @@ def test_store_df_and_read_df(tmp_path: Path) -> None:
uri = tmp_path / "test_store_df_and_read_df.sqlite" uri = tmp_path / "test_store_df_and_read_df.sqlite"
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
# Metadata to store # Metadata to store
meta_md5 = "feature_md5"
element = {"subject": "test"}
meta = { meta = {
"element": "test", "element": element,
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fcname"}, "marker": {
"name": "markername",
},
"type": "BOLD",
} }
_, meta_to_store, element_to_store = process_meta(meta)
# Store metadata
storage.store_metadata(
meta_md5="feature_md5", element=element_to_store, meta=meta_to_store
)
# Columns to store # Columns to store
to_store = df1[["col1", "col2"]] to_store = df1[["col1", "col2"]]
# Check for error while storing # Check for error while storing
with pytest.raises(ValueError, match=r"missing index items"): with pytest.raises(ValueError, match=r"missing index items"):
storage.store_df(to_store.set_index("col1"), meta) storage.store_df(
meta_md5=meta_md5, element=element, df=to_store.set_index("col1")
)
# Set index # Set index
to_store = df1.reset_index().set_index(["element", "pk2", "col1"]) to_store = df1.reset_index().set_index(["subject", "pk2", "col1"])
# Check for error while storing # Check for error while storing
with pytest.raises(ValueError, match=r"extra items"): with pytest.raises(ValueError, match=r"extra items"):
storage.store_df(to_store, meta) storage.store_df(meta_md5=meta_md5, element=element, df=to_store)
# Convert element to index # Convert element to index
idx = element_to_index(meta, n_rows=len(to_store)) idx = storage.element_to_index(element=element, n_rows=len(to_store))
# Set index # Set index
to_store = to_store.set_index(idx) to_store = to_store.set_index(idx)
# Store dataframe # Store dataframe
storage.store_df(to_store, meta) storage.store_df(meta_md5=meta_md5, element=element, df=to_store)
# Store metadata
table_name = storage.store_metadata(meta)
# List stored features # List stored features
features = storage.list_features() features = storage.list_features()
# Check correct usage # Check correct usage
assert len(features) == 1 assert len(features) == 1
assert table_name.replace("meta_", "") in features assert "feature_md5" in features
# Check for missing feature # Check for missing feature
with pytest.raises(ValueError, match="not found"): with pytest.raises(ValueError, match="not found"):
storage.read_df("wrong_md5") storage.read_df(feature_md5="wrong_md5")
with pytest.raises(ValueError, match="not found"):
storage.read_df(feature_name="wrong_name")
# Check for missing feature to fetch # Check for missing feature to fetch
with pytest.raises(ValueError, match="least one"): with pytest.raises(ValueError, match="least one"):
storage.read_df() storage.read_df()
# Check for multiple features to fetch # Check for multiple features to fetch
with pytest.raises(ValueError, match="Only one"): with pytest.raises(ValueError, match="Only one"):
storage.read_df("wrong_md5", "wrong_name") storage.read_df(feature_name="wrong_name", feature_md5="wrong_md5")
# Get MD5 hash of features # Get MD5 hash of features
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "feature_md5" == feature_md5
# Check for key # Check for key
assert "fcname" == features[feature_md5]["name"] assert "BOLD_markername" == features[feature_md5]["name"]
# Read into dataframes # Read into dataframes
read_df1 = storage.read_df(feature_md5=feature_md5) read_df1 = storage.read_df(feature_md5=feature_md5)
read_df2 = storage.read_df(feature_name="fcname") read_df2 = storage.read_df(feature_name="BOLD_markername")
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(read_df1, read_df2) assert_frame_equal(read_df1, read_df2)
assert_frame_equal(read_df1, to_store) assert_frame_equal(read_df1, to_store)
@ -330,10 +344,21 @@ def test_store_metadata(tmp_path: Path) -> None:
# Single storage, must be the uri # Single storage, must be the uri
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
# Metadata to store # Metadata to store
meta = {"element": "test", "version": "0.0.1"} meta = {
"element": {"subject": "test"},
"dependencies": ["numpy"],
"marker": {"name": "test"},
"type": "BOLD",
}
meta_md5, meta_to_store, element_to_store = process_meta(meta)
# Store metadata # Store metadata
table_name = storage.store_metadata(meta) storage.store_metadata(
assert table_name.startswith("meta_") meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
features = storage.list_features()
feature_md5 = list(features.keys())[0]
assert meta_md5 == feature_md5
def test_store_table(tmp_path: Path) -> None: def test_store_table(tmp_path: Path) -> None:
@ -348,7 +373,21 @@ def test_store_table(tmp_path: Path) -> None:
uri = tmp_path / "test_store_table.sqlite" uri = tmp_path / "test_store_table.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Metadata to store # Metadata to store
meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}} element = {"subject": "test"}
dependencies = ["numpy"]
meta = {
"element": element,
"dependencies": dependencies,
"marker": {"name": "fc"},
"type": "BOLD",
}
meta_md5, meta_to_store, element_to_store = process_meta(meta)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
# Data to store # Data to store
data = [ data = [
[1, 10], [1, 10],
@ -358,18 +397,23 @@ def test_store_table(tmp_path: Path) -> None:
[5, 50], [5, 50],
] ]
# Convert element to index # Convert element to index
idx = element_to_index(meta, n_rows=5, rows_col_name="scan") idx = storage.element_to_index(element, n_rows=5, rows_col_name="scan")
# Create dataframe # Create dataframe
df = pd.DataFrame(data, columns=["f1", "f2"], index=idx) df = pd.DataFrame(data, columns=["f1", "f2"], index=idx)
# Store table # Store table
storage.store_table(data, meta, columns=["f1", "f2"], rows_col_name="scan") storage.store_table(
# Store metadata meta_md5=meta_md5,
table_name = storage.store_metadata(meta) element=element_to_store,
data=data,
columns=["f1", "f2"],
rows_col_name="scan",
)
# Read stored table # Read stored table
c_df = _read_sql( c_df = _read_sql(
table_name=table_name, table_name=f"meta_{meta_md5}",
uri=uri.as_posix(), uri=uri.as_posix(),
index_col=["element", "scan"], index_col=["subject", "scan"],
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(df, c_df) assert_frame_equal(df, c_df)
@ -377,19 +421,24 @@ def test_store_table(tmp_path: Path) -> None:
# New data to store # New data to store
data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]] data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]]
# Convert element to index # Convert element to index
idx_new = element_to_index(meta, n_rows=6, rows_col_name="scan") idx_new = storage.element_to_index(element, n_rows=6, rows_col_name="scan")
# Create dataframe # Create dataframe
df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new) df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new)
# Check warning # Check warning
with pytest.warns(RuntimeWarning, match=r"Some rows"): with pytest.warns(RuntimeWarning, match=r"Some rows"):
# Store table
storage.store_table( storage.store_table(
data_new, meta, columns=["f1", "f2"], rows_col_name="scan" meta_md5=meta_md5,
element=element_to_store,
data=data_new,
columns=["f1", "f2"],
rows_col_name="scan",
) )
# Read stored table # Read stored table
c_df_new = _read_sql( c_df_new = _read_sql(
table_name=table_name, table_name=f"meta_{meta_md5}",
uri=uri.as_posix(), uri=uri.as_posix(),
index_col=["element", "scan"], index_col=["subject", "scan"],
) )
# Check if dataframes are equal # Check if dataframes are equal
assert_frame_equal(df_new, c_df_new) assert_frame_equal(df_new, c_df_new)
@ -407,7 +456,20 @@ def test_store_matrix(tmp_path: Path) -> None:
uri = tmp_path / "test_store_table.sqlite" uri = tmp_path / "test_store_table.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Metadata to store # Metadata to store
meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}} element = {"subject": "test"}
dependencies = ["numpy"]
meta = {
"element": element,
"dependencies": dependencies,
"marker": {"name": "fc"},
"type": "BOLD",
}
meta_md5, meta_to_store, element_to_store = process_meta(meta)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
# Store 4 x 3 full matrix # Store 4 x 3 full matrix
data = np.array( data = np.array(
@ -418,17 +480,17 @@ def test_store_matrix(tmp_path: Path) -> None:
# Store table # Store table
storage.store_matrix( storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data, data=data,
meta=meta,
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
) )
stored_names = [f"{i}~{j}" for i in row_names for j in col_names] stored_names = [f"{i}~{j}" for i in row_names for j in col_names]
features = storage.list_features() features = storage.list_features()
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "fc" == features[feature_md5]["name"] assert "BOLD_fc" == features[feature_md5]["name"]
read_df = storage.read_df(feature_md5=feature_md5) read_df = storage.read_df(feature_md5=feature_md5)
assert read_df.shape == (1, 12) assert read_df.shape == (1, 12)
@ -437,7 +499,14 @@ def test_store_matrix(tmp_path: Path) -> None:
# Store without row and column names # Store without row and column names
uri = tmp_path / "test_store_table_nonames.sqlite" uri = tmp_path / "test_store_table_nonames.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
storage.store_matrix(data=data, meta=meta) # Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
# Store table
storage.store_matrix(
meta_md5=meta_md5, element=element_to_store, data=data
)
stored_names = [ stored_names = [
f"r{i}~c{j}" f"r{i}~c{j}"
for i in range(data.shape[0]) for i in range(data.shape[0])
@ -445,20 +514,37 @@ def test_store_matrix(tmp_path: Path) -> None:
] ]
features = storage.list_features() features = storage.list_features()
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "fc" == features[feature_md5]["name"] assert "BOLD_fc" == features[feature_md5]["name"]
read_df = storage.read_df(feature_md5=feature_md5) read_df = storage.read_df(feature_md5=feature_md5)
assert list(read_df.columns) == stored_names assert list(read_df.columns) == stored_names
with pytest.raises(ValueError, match="Invalid kind"): with pytest.raises(ValueError, match="Invalid kind"):
storage.store_matrix(data=data, meta=meta, matrix_kind="wrong") storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data,
row_names=row_names,
col_names=col_names,
matrix_kind="wrong",
)
with pytest.raises(ValueError, match="non-square"): with pytest.raises(ValueError, match="non-square"):
storage.store_matrix(data=data, meta=meta, matrix_kind="triu") storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data,
row_names=row_names,
col_names=col_names,
matrix_kind="triu",
)
with pytest.raises(ValueError, match="cannot be False"): with pytest.raises(ValueError, match="cannot be False"):
storage.store_matrix( storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data, data=data,
meta=meta, row_names=row_names,
col_names=col_names,
matrix_kind="full", matrix_kind="full",
diagonal=False, diagonal=False,
) )
@ -469,12 +555,17 @@ def test_store_matrix(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"] col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_triu.sqlite" uri = tmp_path / "test_store_table_triu.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
storage.store_matrix( storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data, data=data,
meta=meta,
matrix_kind="triu",
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
matrix_kind="triu",
) )
stored_names = [ stored_names = [
@ -488,7 +579,7 @@ def test_store_matrix(tmp_path: Path) -> None:
features = storage.list_features() features = storage.list_features()
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "fc" == features[feature_md5]["name"] assert "BOLD_fc" == features[feature_md5]["name"]
read_df = storage.read_df(feature_md5=feature_md5) read_df = storage.read_df(feature_md5=feature_md5)
assert list(read_df.columns) == stored_names assert list(read_df.columns) == stored_names
assert_array_equal( assert_array_equal(
@ -498,12 +589,17 @@ def test_store_matrix(tmp_path: Path) -> None:
# Store upper triangular matrix without diagonal # Store upper triangular matrix without diagonal
uri = tmp_path / "test_store_table_triu_nodiagonal.sqlite" uri = tmp_path / "test_store_table_triu_nodiagonal.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
storage.store_matrix( storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data, data=data,
meta=meta,
matrix_kind="triu",
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
matrix_kind="triu",
diagonal=False, diagonal=False,
) )
@ -515,7 +611,7 @@ def test_store_matrix(tmp_path: Path) -> None:
features = storage.list_features() features = storage.list_features()
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "fc" == features[feature_md5]["name"] assert "BOLD_fc" == features[feature_md5]["name"]
read_df = storage.read_df(feature_md5=feature_md5) read_df = storage.read_df(feature_md5=feature_md5)
assert list(read_df.columns) == stored_names assert list(read_df.columns) == stored_names
assert_array_equal( assert_array_equal(
@ -528,12 +624,17 @@ def test_store_matrix(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"] col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_tril.sqlite" uri = tmp_path / "test_store_table_tril.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
storage.store_matrix( storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
data=data, data=data,
meta=meta,
matrix_kind="tril",
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
matrix_kind="tril",
) )
stored_names = [ stored_names = [
@ -547,7 +648,7 @@ def test_store_matrix(tmp_path: Path) -> None:
features = storage.list_features() features = storage.list_features()
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "fc" == features[feature_md5]["name"] assert "BOLD_fc" == features[feature_md5]["name"]
read_df = storage.read_df(feature_md5=feature_md5) read_df = storage.read_df(feature_md5=feature_md5)
assert list(read_df.columns) == stored_names assert list(read_df.columns) == stored_names
assert_array_equal( assert_array_equal(
@ -557,15 +658,19 @@ def test_store_matrix(tmp_path: Path) -> None:
# Store lower triangular matrix without diagonal # Store lower triangular matrix without diagonal
uri = tmp_path / "test_store_table_tril_nodiagonal.sqlite" uri = tmp_path / "test_store_table_tril_nodiagonal.sqlite"
storage = SQLiteFeatureStorage(uri=uri) storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
storage.store_matrix( storage.store_matrix(
data, meta_md5=meta_md5,
meta, element=element_to_store,
matrix_kind="tril", data=data,
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
matrix_kind="tril",
diagonal=False, diagonal=False,
) )
stored_names = [ stored_names = [
"row2~col1", "row2~col1",
"row3~col1", "row3~col1",
@ -574,7 +679,7 @@ def test_store_matrix(tmp_path: Path) -> None:
features = storage.list_features() features = storage.list_features()
feature_md5 = list(features.keys())[0] feature_md5 = list(features.keys())[0]
assert "fc" == features[feature_md5]["name"] assert "BOLD_fc" == features[feature_md5]["name"]
read_df = storage.read_df(feature_md5=feature_md5) read_df = storage.read_df(feature_md5=feature_md5)
assert list(read_df.columns) == stored_names assert list(read_df.columns) == stored_names
assert_array_equal( assert_array_equal(
@ -597,18 +702,21 @@ def test_store_multiple_output(tmp_path: Path):
# Metadata to store # Metadata to store
meta1 = { meta1 = {
"element": {"subject": "test-01", "session": "ses-01"}, "element": {"subject": "test-01", "session": "ses-01"},
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fc"}, "marker": {"name": "fc"},
"type": "BOLD",
} }
meta2 = { meta2 = {
"element": {"subject": "test-02", "session": "ses-01"}, "element": {"subject": "test-02", "session": "ses-01"},
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fc"}, "marker": {"name": "fc"},
"type": "BOLD",
} }
meta3 = { meta3 = {
"element": {"subject": "test-01", "session": "ses-02"}, "element": {"subject": "test-01", "session": "ses-02"},
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fc"}, "marker": {"name": "fc"},
"type": "BOLD",
} }
# Data to store # Data to store
data1 = np.array( data1 = np.array(
@ -623,35 +731,62 @@ def test_store_multiple_output(tmp_path: Path):
data2 = data1 * 10 data2 = data1 * 10
data3 = data1 * 20 data3 = data1 * 20
# Process metadata for storage # Process metadata for storage
hash1, _ = process_meta(meta1) hash1, meta_to_store1, element_to_store1 = process_meta(meta1)
# Convert element to index # Convert element to index
idx1 = element_to_index(meta1, n_rows=5, rows_col_name="scan") idx1 = storage.element_to_index(
element_to_store1, n_rows=5, rows_col_name="scan"
)
# Create dataframe # Create dataframe
df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1) df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1)
# Process metadata for storage # Process metadata for storage
hash2, _ = process_meta(meta2) hash2, meta_to_store2, element_to_store2 = process_meta(meta2)
# Convert element to index # Convert element to index
idx2 = element_to_index(meta2, n_rows=5, rows_col_name="scan") idx2 = storage.element_to_index(
element_to_store2, n_rows=5, rows_col_name="scan"
)
# Create dataframe # Create dataframe
df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2) df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2)
# Process metadata for storage # Process metadata for storage
hash3, _ = process_meta(meta3) hash3, meta_to_store3, element_to_store3 = process_meta(meta3)
# Convert element to index # Convert element to index
idx3 = element_to_index(meta3, n_rows=5, rows_col_name="scan") idx3 = storage.element_to_index(
element_to_store3, n_rows=5, rows_col_name="scan"
)
# Create dataframe # Create dataframe
df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3) df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3)
# Check hash equality # Check hash equality
assert hash1 == hash2 assert hash1 == hash2
assert hash2 == hash3 assert hash2 == hash3
# Store tables # Store tables
storage.store_table( storage.store_metadata(
data1, meta1, columns=["f1", "f2"], rows_col_name="scan" meta_md5=hash1, element=element_to_store1, meta=meta_to_store1
)
storage.store_metadata(
meta_md5=hash2, element=element_to_store2, meta=meta_to_store2
)
storage.store_metadata(
meta_md5=hash3, element=element_to_store3, meta=meta_to_store3
) )
storage.store_table( storage.store_table(
data2, meta2, columns=["f1", "f2"], rows_col_name="scan" meta_md5=hash1,
element=element_to_store1,
data=data1,
columns=["f1", "f2"],
rows_col_name="scan",
) )
storage.store_table( storage.store_table(
data3, meta3, columns=["f1", "f2"], rows_col_name="scan" meta_md5=hash2,
element=element_to_store2,
data=data2,
columns=["f1", "f2"],
rows_col_name="scan",
)
storage.store_table(
meta_md5=hash3,
element=element_to_store3,
data=data3,
columns=["f1", "f2"],
rows_col_name="scan",
) )
# Check that URI does not exist yet # Check that URI does not exist yet
assert not uri.exists() assert not uri.exists()
@ -667,10 +802,9 @@ def test_store_multiple_output(tmp_path: Path):
assert uri1.exists() assert uri1.exists()
assert uri2.exists() assert uri2.exists()
assert uri3.exists() assert uri3.exists()
# Store metadata
table_name = storage.store_metadata(meta1)
# Set index columns # Set index columns
cols = ["subject", "session", "scan"] cols = ["subject", "session", "scan"]
table_name = f"meta_{hash1}"
# Read stored tables # Read stored tables
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)
cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols) cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols)
@ -696,18 +830,21 @@ def test_collect(tmp_path: Path) -> None:
# Metadata for storage # Metadata for storage
meta1 = { meta1 = {
"element": {"subject": "test-01", "session": "ses-01"}, "element": {"subject": "test-01", "session": "ses-01"},
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fc"}, "marker": {"name": "fc"},
"type": "BOLD",
} }
meta2 = { meta2 = {
"element": {"subject": "test-02", "session": "ses-01"}, "element": {"subject": "test-02", "session": "ses-01"},
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fc"}, "marker": {"name": "fc"},
"type": "BOLD",
} }
meta3 = { meta3 = {
"element": {"subject": "test-01", "session": "ses-02"}, "element": {"subject": "test-01", "session": "ses-02"},
"version": "0.0.1", "dependencies": ["numpy"],
"marker": {"name": "fc"}, "marker": {"name": "fc"},
"type": "BOLD",
} }
# Data for storage # Data for storage
data1 = np.array( data1 = np.array(
@ -722,14 +859,38 @@ def test_collect(tmp_path: Path) -> None:
data2 = data1 * 10 data2 = data1 * 10
data3 = data1 * 20 data3 = data1 * 20
# Store tables # Store tables
storage.store_table( hash1, meta_to_store1, element_to_store1 = process_meta(meta1)
data1, meta1, columns=["f1", "f2"], rows_col_name="scan" hash2, meta_to_store2, element_to_store2 = process_meta(meta2)
hash3, meta_to_store3, element_to_store3 = process_meta(meta3)
storage.store_metadata(
meta_md5=hash1, element=element_to_store1, meta=meta_to_store1
)
storage.store_metadata(
meta_md5=hash2, element=element_to_store2, meta=meta_to_store2
)
storage.store_metadata(
meta_md5=hash3, element=element_to_store3, meta=meta_to_store3
) )
storage.store_table( storage.store_table(
data2, meta2, columns=["f1", "f2"], rows_col_name="scan" meta_md5=hash1,
element=element_to_store1,
data=data1,
columns=["f1", "f2"],
rows_col_name="scan",
) )
storage.store_table( storage.store_table(
data3, meta3, columns=["f1", "f2"], rows_col_name="scan" meta_md5=hash2,
element=element_to_store2,
data=data2,
columns=["f1", "f2"],
rows_col_name="scan",
)
storage.store_table(
meta_md5=hash3,
element=element_to_store3,
data=data3,
columns=["f1", "f2"],
rows_col_name="scan",
) )
# Convert element to prefix # Convert element to prefix
prefix1 = element_to_prefix(meta1["element"]) prefix1 = element_to_prefix(meta1["element"])
@ -752,7 +913,7 @@ def test_collect(tmp_path: Path) -> None:
# Set index columns # Set index columns
cols = ["subject", "session", "scan"] cols = ["subject", "session", "scan"]
# Store metadata # Store metadata
table_name = storage.store_metadata(meta1) table_name = f"meta_{hash1}"
# Read stored tables # Read stored tables
all_df = _read_sql(table_name, uri.as_posix(), index_col=cols) all_df = _read_sql(table_name, uri.as_posix(), index_col=cols)
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)

View file

@ -12,7 +12,7 @@ from junifer.storage.base import BaseFeatureStorage
def test_BaseFeatureStorage_abstractness() -> None: def test_BaseFeatureStorage_abstractness() -> None:
"""Test BaseFeatureStorage is abstract base class.""" """Test BaseFeatureStorage is abstract base class."""
with pytest.raises(TypeError, match=r"abstract"): with pytest.raises(TypeError, match=r"abstract"):
BaseFeatureStorage(uri="/tmp", storage_types=["matrix"]) BaseFeatureStorage(uri="/tm", storage_types=["matrix"]) # type: ignore
def test_BaseFeatureStorage() -> None: def test_BaseFeatureStorage() -> None:
@ -22,13 +22,16 @@ def test_BaseFeatureStorage() -> None:
"""Implement concrete class.""" """Implement concrete class."""
def __init__(self, uri, single_output=False): def __init__(self, uri, single_output=False):
storage_types = ["matrix"] storage_types = ["matrix", "table", "timeseries"]
super().__init__( super().__init__(
uri=uri, uri=uri,
storage_types=storage_types, storage_types=storage_types,
single_output=single_output, single_output=single_output,
) )
def get_valid_inputs(self):
return ["matrix", "table", "timeseries"]
def list_features(self): def list_features(self):
super().list_features() super().list_features()
@ -38,8 +41,8 @@ def test_BaseFeatureStorage() -> None:
feature_md5=feature_md5, feature_md5=feature_md5,
) )
def store_metadata(self, metadata): def store_metadata(self, meta_md5, meta, element):
super().store_metadata(metadata) super().store_metadata(meta_md5, meta, element)
def collect(self): def collect(self):
return super().collect() return super().collect()
@ -55,7 +58,7 @@ def test_BaseFeatureStorage() -> None:
st.validate(input_=["matrix"]) st.validate(input_=["matrix"])
# Check validate with invalid argument # Check validate with invalid argument
with pytest.raises(ValueError): with pytest.raises(ValueError):
st.validate(input_=["table"]) st.validate(input_=["duck"])
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.list_features() st.list_features()
@ -63,22 +66,31 @@ def test_BaseFeatureStorage() -> None:
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.read_df(None) st.read_df(None)
element = {"subject": "test"}
dependencies = ["numpy"]
meta = {
"element": element,
"dependencies": dependencies,
"marker": {"name": "fc"},
"type": "BOLD",
}
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.store_metadata(None) st.store(kind="matrix", meta=meta)
with pytest.raises(NotImplementedError):
st.store_metadata("md5", meta=meta, element={})
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.collect() st.collect()
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.store(kind="matrix") st.store(kind="timeseries", meta=meta)
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.store(kind="timeseries") st.store(kind="table", meta=meta)
with pytest.raises(NotImplementedError):
st.store(kind="table")
with pytest.raises(ValueError): with pytest.raises(ValueError):
st.store(kind="lego") st.store(kind="lego", meta=meta)
assert st.uri == "/tmp" assert st.uri == "/tmp"

View file

@ -4,17 +4,51 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Dict, List, Tuple, Union from typing import Dict, List
import pytest import pytest
from junifer.storage.utils import ( from junifer.storage.utils import (
element_to_index,
element_to_prefix, element_to_prefix,
get_dependency_version,
process_meta, process_meta,
) )
@pytest.mark.parametrize(
"dependency, max_version",
[
("click", "8.2"),
("numpy", "1.24"),
("datalad", "0.18"),
("pandas", "1.6"),
("nibabel", "4.1"),
("nilearn", "1.0"),
("sqlalchemy", "1.5.0"),
("pyyaml", "7.0"),
],
)
def test_get_dependency_version(dependency: str, max_version: str) -> None:
"""Test dependency resolution for installed dependencies.
Parameters
----------
dependency : str
The parametrized dependency name.
max_version : str
The parametrized maximum version of the dependency.
"""
version = get_dependency_version(dependency)
assert version < max_version
def test_get_dependency_version_invalid() -> None:
"""Test invalid package name handling for dependency resolution."""
with pytest.raises(ValueError, match="Could not obtain"):
get_dependency_version("foobar")
def test_process_meta_invalid_metadata_type() -> None: def test_process_meta_invalid_metadata_type() -> None:
"""Test invalid metadata type check for metadata hash processing.""" """Test invalid metadata type check for metadata hash processing."""
meta = None meta = None
@ -25,58 +59,133 @@ def test_process_meta_invalid_metadata_type() -> None:
# TODO: parameterize # TODO: parameterize
def test_process_meta_hash() -> None: def test_process_meta_hash() -> None:
"""Test metadata hash processing.""" """Test metadata hash processing."""
meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]} meta = {
hash1, _ = process_meta(meta) "element": {"foo": "bar"},
"A": 1,
"B": [2, 3, 4, 5, 6],
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
}
hash1, _, element1 = process_meta(meta)
assert element1 == {"foo": "bar"}
meta = {"element": "foo", "B": [2, 3, 4, 5, 6], "A": 1} meta = {
hash2, _ = process_meta(meta) "element": {"foo": "baz"},
"B": [2, 3, 4, 5, 6],
"A": 1,
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
}
hash2, _, element2 = process_meta(meta)
assert hash1 == hash2 assert hash1 == hash2
assert element2 == {"foo": "baz"}
meta = {"element": "foo", "A": 1, "B": [2, 3, 1, 5, 6]} meta = {
hash3, _ = process_meta(meta) "element": {"foo": "bar"},
"A": 1,
"B": [2, 3, 1, 5, 6],
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
}
hash3, _, element3 = process_meta(meta)
assert hash1 != hash3 assert hash1 != hash3
assert element3 == element1
meta1 = { meta4 = {
"element": "foo", "element": {"foo": "bar"},
"B": { "B": {
"B2": [2, 3, 4, 5, 6], "B2": [2, 3, 4, 5, 6],
"B1": [9.22, 3.14, 1.41, 5.67, 6.28], "B1": [9.22, 3.14, 1.41, 5.67, 6.28],
"B3": (1, "car"), "B3": (1, "car"),
}, },
"A": 1, "A": 1,
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
} }
meta2 = { meta5 = {
"A": 1, "A": 1,
"B": { "B": {
"B3": (1, "car"), "B3": (1, "car"),
"B1": [9.22, 3.14, 1.41, 5.67, 6.28], "B1": [9.22, 3.14, 1.41, 5.67, 6.28],
"B2": [2, 3, 4, 5, 6], "B2": [2, 3, 4, 5, 6],
}, },
"element": "foo", "element": {"foo": "baz"},
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
} }
hash4, _ = process_meta(meta1) hash4, _, _ = process_meta(meta4)
hash5, _ = process_meta(meta2) hash5, _, _ = process_meta(meta5)
assert hash4 == hash5 assert hash4 == hash5
# Different element keys should give a different hash
meta6 = {
"A": 1,
"B": {
"B3": (1, "car"),
"B1": [9.22, 3.14, 1.41, 5.67, 6.28],
"B2": [2, 3, 4, 5, 6],
},
"element": {"bar": "baz"},
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
}
hash6, _, _ = process_meta(meta6)
assert hash4 != hash6
def test_process_meta_invalid_metadata_key() -> None: def test_process_meta_invalid_metadata_key() -> None:
"""Test invalid metadata key check for metadata hash processing.""" """Test invalid metadata key check for metadata hash processing."""
meta = {} meta = {}
with pytest.raises(ValueError, match=r"_element_keys"): with pytest.raises(ValueError, match=r"element"):
process_meta(meta)
meta = {"element": {}}
with pytest.raises(ValueError, match=r"marker"):
process_meta(meta)
meta = {"element": {}, "marker": {}}
with pytest.raises(ValueError, match=r"key 'name'"):
process_meta(meta)
meta = {"element": {}, "marker": {"name": "test"}}
with pytest.raises(ValueError, match=r"key 'type'"):
process_meta(meta)
meta = {"element": {}, "marker": {"name": "test"}, "type": "BOLD"}
with pytest.raises(ValueError, match=r"dependencies"):
process_meta(meta) process_meta(meta)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"meta,elements", "meta,elements",
[ [
({"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}, ["element"]), (
{
"element": {"foo": "bar"},
"A": 1,
"B": [2, 3, 4, 5, 6],
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
},
["foo"],
),
( (
{ {
"element": {"subject": "foo", "session": "bar"}, "element": {"subject": "foo", "session": "bar"},
"B": [2, 3, 4, 5, 6], "B": [2, 3, 4, 5, 6],
"A": 1, "A": 1,
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
}, },
["subject", "session"], ["subject", "session"],
), ),
@ -93,34 +202,30 @@ def test_process_meta_element(meta: Dict, elements: List[str]) -> None:
The parametrized elements to assert against. The parametrized elements to assert against.
""" """
hash1, processed_meta = process_meta(meta) hash1, processed_meta, _ = process_meta(meta)
assert "_element_keys" in processed_meta assert "_element_keys" in processed_meta
assert processed_meta["_element_keys"] == elements assert processed_meta["_element_keys"] == elements
assert "A" in processed_meta assert "A" in processed_meta
assert "B" in processed_meta assert "B" in processed_meta
assert "element" not in processed_meta assert "element" not in processed_meta
hash2, processed_meta2 = process_meta(processed_meta) assert isinstance(processed_meta["dependencies"], Dict)
assert all(
assert hash1, hash2 x in processed_meta["dependencies"] for x in meta["dependencies"]
assert processed_meta == processed_meta2 )
assert "name" in processed_meta
assert processed_meta["name"] == f"{meta['type']}_{meta['marker']['name']}"
@pytest.mark.parametrize( @pytest.mark.parametrize(
"element,prefix", "element,prefix",
[ [
("sub-01", "element_sub-01_"),
(1, "element_1_"),
({"subject": "sub-01"}, "element_sub-01_"), ({"subject": "sub-01"}, "element_sub-01_"),
({"subject": 1}, "element_1_"), ({"subject": 1}, "element_1_"),
({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"), ({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"),
({"subject": 1, "session": 2}, "element_1_2_"), ({"subject": 1, "session": 2}, "element_1_2_"),
(("sub-01", "ses-02"), "element_sub-01_ses-02_"),
((1, 2), "element_1_2_"),
], ],
) )
def test_element_to_prefix( def test_element_to_prefix(element: Dict, prefix: str) -> None:
element: Union[str, int, Dict, Tuple], prefix: str
) -> None:
"""Test converting element to prefix (for file naming). """Test converting element to prefix (for file naming).
Parameters Parameters
@ -138,75 +243,5 @@ def test_element_to_prefix(
def test_element_to_prefix_invalid_type() -> None: def test_element_to_prefix_invalid_type() -> None:
"""Test element to prefix type checking.""" """Test element to prefix type checking."""
element = 2.3 element = 2.3
with pytest.raises(ValueError, match=r"convert element of type"): with pytest.raises(ValueError, match=r"must be a dict"):
element_to_prefix(element) # type: ignore element_to_prefix(element) # type: ignore
def test_element_to_index_check_meta_invalid_key() -> None:
"""Test element to index metadata key checking."""
meta = {"noelement": "foo"}
with pytest.raises(ValueError, match=r"metadata must contain the key"):
element_to_index(meta)
def test_element_to_index() -> None:
"""Test element to index."""
meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}
index = element_to_index(meta)
assert index.names == ["element", "idx"]
assert index.levels[0].name == "element"
assert index.levels[0].values[0] == "foo"
assert all(x == "foo" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "idx"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (1,)
index = element_to_index(meta, n_rows=10)
assert index.names == ["element", "idx"]
assert index.levels[0].name == "element"
assert all(x == "foo" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "idx"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (10,)
index = element_to_index(meta, n_rows=1, rows_col_name="scan")
assert index.names == ["element", "scan"]
assert index.levels[0].name == "element"
assert index.levels[0].values[0] == "foo"
assert all(x == "foo" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "scan"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (1,)
index = element_to_index(meta, n_rows=7, rows_col_name="scan")
assert index.names == ["element", "scan"]
assert index.levels[0].name == "element"
assert all(x == "foo" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "scan"
assert all(x == i for i, x in enumerate(index.levels[1].values))
assert index.levels[1].values.shape == (7,)
meta = {
"element": {"subject": "sub-01", "session": "ses-01"},
"A": 1,
"B": [2, 3, 4, 5, 6],
}
index = element_to_index(meta, n_rows=10)
assert index.levels[0].name == "subject"
assert all(x == "sub-01" for x in index.levels[0].values)
assert index.levels[0].values.shape == (1,)
assert index.levels[1].name == "session"
assert all(x == "ses-01" for x in index.levels[1].values)
assert index.levels[1].values.shape == (1,)
assert index.levels[2].name == "idx"
assert all(x == i for i, x in enumerate(index.levels[2].values))
assert index.levels[2].values.shape == (10,)

View file

@ -6,14 +6,39 @@
import hashlib import hashlib
import json import json
from typing import Any, Dict, Optional, Tuple, Union from importlib.metadata import PackageNotFoundError, version
from typing import Dict, Tuple
import numpy as np
import pandas as pd
from ..utils.logging import logger, raise_error from ..utils.logging import logger, raise_error
def get_dependency_version(dependency: str) -> str:
"""Get dependency version.
Parameters
----------
dependency : str
The dependency to fetch version for.
Returns
-------
str
The version of the dependency.
"""
dep_version = ""
try:
dep_version = version(dependency)
except PackageNotFoundError as e:
raise_error(
f"Could not obtain the version of {dependency}. "
"Have you specified the DEPENDENCIES variable correctly?",
exception=e,
)
return dep_version
def _meta_hash(meta: Dict) -> str: def _meta_hash(meta: Dict) -> str:
"""Compute the MD5 hash of the metadata. """Compute the MD5 hash of the metadata.
@ -29,6 +54,12 @@ def _meta_hash(meta: Dict) -> str:
""" """
logger.debug(f"Hashing metadata: {meta}") logger.debug(f"Hashing metadata: {meta}")
if "dependencies" not in meta:
raise_error("The metadata must contain the key 'dependencies'")
# Convert dependencies set into {dependency: version} dictionary
meta["dependencies"] = {
dep: get_dependency_version(dep) for dep in meta["dependencies"]
}
meta_md5 = hashlib.md5( meta_md5 = hashlib.md5(
json.dumps(meta, sort_keys=True).encode("utf-8") json.dumps(meta, sort_keys=True).encode("utf-8")
).hexdigest() ).hexdigest()
@ -36,7 +67,7 @@ def _meta_hash(meta: Dict) -> str:
return meta_md5 return meta_md5
def process_meta(meta: Dict) -> Tuple[str, Dict]: def process_meta(meta: Dict) -> Tuple[str, Dict, Dict]:
"""Process the metadata for storage. """Process the metadata for storage.
It removes the key "element" and adds the "_element_keys" with the keys It removes the key "element" and adds the "_element_keys" with the keys
@ -53,12 +84,13 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]:
The MD5 hash of the metadata. The MD5 hash of the metadata.
dict dict
The processed metadata for storage. The processed metadata for storage.
tuple
The element.
Raises Raises
------ ------
ValueError ValueError
If `meta` is None or if it does not contain the key "element" or If `meta` is None or if it does not contain the key "element".
"_element_keys".
""" """
if meta is None: if meta is None:
@ -68,97 +100,42 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]:
# Remove key "element" # Remove key "element"
element = t_meta.pop("element", None) element = t_meta.pop("element", None)
if element is None: if element is None:
if "_element_keys" not in t_meta: raise_error(msg="`meta` must contain the key 'element'")
raise_error( if "marker" not in t_meta:
msg="`meta` must contain the key 'element' or '_element_keys'" raise_error(msg="`meta` must contain the key 'marker'")
) if "name" not in t_meta["marker"]:
else: raise_error(msg="`meta['marker']` must contain the key 'name'")
if isinstance(element, dict): if "type" not in t_meta:
t_meta["_element_keys"] = list(element.keys()) raise_error(msg="`meta` must contain the key 'type'")
else:
t_meta["_element_keys"] = ["element"] t_meta["_element_keys"] = list(element.keys())
type_ = t_meta["type"]
name = t_meta["marker"]["name"]
t_meta["name"] = f"{type_}_{name}"
# MD5 hash of the metadata # MD5 hash of the metadata
md5_hash = _meta_hash(t_meta) md5_hash = _meta_hash(t_meta)
return md5_hash, t_meta return md5_hash, t_meta, element
def element_to_prefix(element: Union[Tuple, Dict, str, int]) -> str: def element_to_prefix(element: Dict) -> str:
"""Convert the element metadata to prefix. """Convert the element metadata to prefix.
Parameters Parameters
---------- ----------
element : tuple, dict, str or int element : dict
The element to convert to prefix. The element to convert to prefix.
Returns Returns
------- -------
str str
The element converted to prefix. The element converted to prefix.
Raises
------
ValueError
If invalid type is passed for `element`.
""" """
logger.debug(f"Converting element {element} to prefix.") logger.debug(f"Converting element {element} to prefix.")
prefix = "element" prefix = "element"
if isinstance(element, tuple): if not isinstance(element, dict):
prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}" raise_error(msg="`element` must be a dict")
elif isinstance(element, dict):
prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}" prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}"
elif isinstance(element, (str, int)):
prefix = f"{prefix}_{element}"
else:
raise_error(
f"Cannot convert element of type {type(element)} to prefix. "
"Must be a str, int, tuple or dict."
)
logger.debug(f"Converted prefix: {prefix}") logger.debug(f"Converted prefix: {prefix}")
return f"{prefix}_" return f"{prefix}_"
def element_to_index(
meta: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None
) -> pd.MultiIndex:
"""Convert the element metadata to index.
Parameters
----------
meta : dict
The metadata as a dictionary. Must contain the key "element"."
n_rows : int, optional
Number of rows to create (default 1).
rows_col_name: str, optional
The column name to use in case `n_rows` > 1. If None and
n_rows > 1, the name will be "idx" (default None).
Returns
-------
pandas.MultiIndex
The index of the dataframe to store.
Raises
------
ValueError
If `meta` does not contain the key "element".
"""
if "element" not in meta:
raise_error(
msg="To create and index, metadata must contain the key 'element'."
)
# Get element
element = meta["element"]
if not isinstance(element, dict):
element = {"element": element}
# Check rows_col_name
if rows_col_name is None:
rows_col_name = "idx"
elem_idx: Dict[Any, Any] = {k: [v] * n_rows for k, v in element.items()}
elem_idx[rows_col_name] = np.arange(n_rows)
# Create index
index = pd.MultiIndex.from_frame(
pd.DataFrame(elem_idx, index=range(n_rows))
)
return index

View file

@ -10,6 +10,7 @@ from .datagrabbers import (
SPMAuditoryTestingDatagrabber, SPMAuditoryTestingDatagrabber,
) )
# Register testing datagrabber # Register testing datagrabber
register( register(
step="datagrabber", step="datagrabber",

View file

@ -14,6 +14,7 @@ from warnings import warn
import datalad import datalad
logger = logging.getLogger("JUNIFER") logger = logging.getLogger("JUNIFER")
# Set up datalad logger level to warning by default # Set up datalad logger level to warning by default
@ -262,7 +263,11 @@ def configure_logging(
log_versions() # log versions of installed packages log_versions() # log versions of installed packages
def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: def raise_error(
msg: str,
klass: Type[Exception] = ValueError,
exception: Optional[Exception] = None,
) -> NoReturn:
"""Raise error, but first log it. """Raise error, but first log it.
Parameters Parameters
@ -271,10 +276,15 @@ def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn:
The message for the exception. The message for the exception.
klass : subclass of Exception, optional klass : subclass of Exception, optional
The subclass of Exception to raise using (default ValueError). The subclass of Exception to raise using (default ValueError).
exception : Exception, optional
The original exception to follow up on (default None).
""" """
logger.error(msg) logger.error(msg)
raise klass(msg) if exception is not None:
raise klass(msg) from exception
else:
raise klass(msg)
def warn_with_log( def warn_with_log(

View file

@ -11,7 +11,7 @@ name = "junifer"
description = "JUelich NeuroImaging FEature extractoR" description = "JUelich NeuroImaging FEature extractoR"
readme = "README.md" readme = "README.md"
requires-python = ">=3.8" requires-python = ">=3.8"
license = {file = "LICENSE.md"} license = {text = "AGPL-3.0-only"}
authors = [ authors = [
{name = "Fede Raimondo", email = "f.raimondo@fz-juelich.de"}, {name = "Fede Raimondo", email = "f.raimondo@fz-juelich.de"},
{name = "Synchon Mandal", email = "s.mandal@fz-juelich.de"}, {name = "Synchon Mandal", email = "s.mandal@fz-juelich.de"},