diff --git a/README.md b/README.md index 18cf7a61c..b8079569b 100644 --- a/README.md +++ b/README.md @@ -54,7 +54,7 @@ Contributions are welcome and greatly appreciated. Please read the [guidelines]( 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. This program is free software: you can redistribute it and/or modify diff --git a/docs/extending/marker.rst b/docs/extending/marker.rst index 85c084e35..a9f79c1d5 100644 --- a/docs/extending/marker.rst +++ b/docs/extending/marker.rst @@ -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 inputs provided by the user are valid. This method should return a list of strings, representing :ref:`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 :ref:`storage types ` 3. ``compute``: the method that given the data, computes the marker. -4. ``store``: the method that stores the computed marker. -5. ``__init__``: the initialization method, where the marker is configured. +4. ``__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 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 - def get_output_kind(self, input_kind): + def get_output_type(self, input_kind): if input_kind == 'BOLD': return 'timeseries' else: @@ -143,38 +142,24 @@ This dictionary will later be passed onto the ``store`` method. 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, -using the data provided by the ``compute`` method. The method ``store`` has three arguments: +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 ``@register_marker`` decorator. -* ``kind``: A string indicating the :ref:`data type ` that was used to compute the marker. -* ``out``: The output of the ``compute`` method. -* ``storage``: The storage object, that will be used to store the marker. +The *dependencies* are the core packages that are required to compute the marker. This will be later used to keep track +of the versions of the packages used to compute the marker. To inform junifer about the dependencies of a 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 - 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) + _DEPENDENCIES = {"nilearn"} -.. hint:: Check the hint on :ref:`extending_markers_compute`. If the output of the ``compute`` method is a dictionary - with keys based on the :ref:`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: +Finally, we need to register the marker using the ``@register_marker`` decorator. This decorator takes the name of the .. 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 class ParcelMean(BaseMarker): + _DEPENDENCIES = {"nilearn", "numpy"} + def __init__(self, parcellation_name, on=None, name=None): self.parcellation_name = parcellation_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): return ['BOLD', 'VBM_WM', 'VBM_GM'] - def get_output_kind(self, input_kind): + def get_output_type(self, input_kind): if input_kind == 'BOLD': return 'timeseries' 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" 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: @@ -260,7 +241,7 @@ Template for a custom Marker 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 pass @@ -270,7 +251,3 @@ Template for a custom Marker # Create the output dictionary out = {"data": None, "columns": None} return out - - def store(self, kind, out, storage): - # TODO: store out using the storage object, based on the kind of data - pass \ No newline at end of file diff --git a/docs/understanding/marker.rst b/docs/understanding/marker.rst index e744ccfe3..1ccdfce12 100644 --- a/docs/understanding/marker.rst +++ b/docs/understanding/marker.rst @@ -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 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. diff --git a/examples/run_compute_parcel_mean.py b/examples/run_compute_parcel_mean.py index fa7df2e97..d4d00f046 100644 --- a/examples/run_compute_parcel_mean.py +++ b/examples/run_compute_parcel_mean.py @@ -42,7 +42,10 @@ marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean") ############################################################################### # 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 diff --git a/examples/run_ets_rss_marker.py b/examples/run_ets_rss_marker.py index ff6b5e86b..c364eec60 100644 --- a/examples/run_ets_rss_marker.py +++ b/examples/run_ets_rss_marker.py @@ -66,6 +66,9 @@ with tempfile.TemporaryDirectory() as tmpdir: collect(storage=storage) # Create storage object to read in extracted features db = SQLiteFeatureStorage(uri=storage["uri"]) + + # List all the features + print(db.list_features()) # Read extracted features df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS") diff --git a/junifer/api/cli.py b/junifer/api/cli.py index 9293dcc81..964a12029 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -16,8 +16,8 @@ import yaml from ..utils.logging import ( configure_logging, logger, - warn_with_log, raise_error, + warn_with_log, ) from .functions import collect as api_collect from .functions import queue as api_queue @@ -69,7 +69,8 @@ def _parse_elements(element: str, config: Dict) -> Union[List, None]: raise_error( "The 'elements' key is set in the configuration, but its value" " 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 diff --git a/junifer/api/functions.py b/junifer/api/functions.py index 21aef3850..44835cf97 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -7,11 +7,11 @@ import shutil import subprocess +import textwrap import typing from pathlib import Path from typing import Dict, List, Optional, Tuple, Union -import textwrap import yaml from ..datagrabber.base import BaseDataGrabber diff --git a/junifer/api/parser.py b/junifer/api/parser.py index fabdd1647..18d37552c 100644 --- a/junifer/api/parser.py +++ b/junifer/api/parser.py @@ -5,11 +5,11 @@ # License: AGPL import importlib -from pathlib import Path -from typing import Dict, Union import importlib.util import os import sys +from pathlib import Path +from typing import Dict, Union import yaml @@ -45,7 +45,8 @@ def parse_yaml(filepath: Union[str, Path]) -> Dict: if contents["elements"] is None: raise_error( "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 if "with" in contents: to_load = contents["with"] @@ -58,9 +59,11 @@ def parse_yaml(filepath: Union[str, Path]) -> Dict: file_path = Path(os.getcwd()) / t_module if not file_path.exists(): 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( - t_module, file_path) + t_module, file_path + ) module = importlib.util.module_from_spec(spec) # type: ignore sys.modules[t_module] = module spec.loader.exec_module(module) # type: ignore diff --git a/junifer/api/tests/test_cli.py b/junifer/api/tests/test_cli.py index 27935ad9b..8e90dc765 100644 --- a/junifer/api/tests/test_cli.py +++ b/junifer/api/tests/test_cli.py @@ -13,6 +13,7 @@ from click.testing import CliRunner from junifer.api.cli import collect, run, selftest, wtf + # Create click test runner runner = CliRunner() diff --git a/junifer/configs/juseless/datagrabbers/tests/test_aomic_id1000_vbm.py b/junifer/configs/juseless/datagrabbers/tests/test_aomic_id1000_vbm.py index fa3279e1c..fad38837f 100644 --- a/junifer/configs/juseless/datagrabbers/tests/test_aomic_id1000_vbm.py +++ b/junifer/configs/juseless/datagrabbers/tests/test_aomic_id1000_vbm.py @@ -11,6 +11,7 @@ import pytest from junifer.configs.juseless.datagrabbers import JuselessDataladAOMICID1000VBM from junifer.utils.logging import configure_logging + # Check if the test is running on juseless if socket.gethostname() != "juseless": pytest.skip("These tests are only for juseless", allow_module_level=True) diff --git a/junifer/configs/juseless/datagrabbers/tests/test_camcan_vbm.py b/junifer/configs/juseless/datagrabbers/tests/test_camcan_vbm.py index c80cf418f..efb149bd2 100644 --- a/junifer/configs/juseless/datagrabbers/tests/test_camcan_vbm.py +++ b/junifer/configs/juseless/datagrabbers/tests/test_camcan_vbm.py @@ -12,6 +12,7 @@ import pytest from junifer.configs.juseless.datagrabbers import JuselessDataladCamCANVBM from junifer.utils.logging import configure_logging + # Check if the test is running on juseless if socket.gethostname() != "juseless": pytest.skip("These tests are only for juseless", allow_module_level=True) diff --git a/junifer/configs/juseless/datagrabbers/tests/test_ixi_vbm.py b/junifer/configs/juseless/datagrabbers/tests/test_ixi_vbm.py index 73d655307..b9c78dd84 100644 --- a/junifer/configs/juseless/datagrabbers/tests/test_ixi_vbm.py +++ b/junifer/configs/juseless/datagrabbers/tests/test_ixi_vbm.py @@ -12,6 +12,7 @@ import pytest from junifer.configs.juseless.datagrabbers import JuselessDataladIXIVBM from junifer.utils.logging import configure_logging + # Check if the test is running on juseless if socket.gethostname() != "juseless": pytest.skip("These tests are only for juseless", allow_module_level=True) diff --git a/junifer/configs/juseless/datagrabbers/tests/test_ucla.py b/junifer/configs/juseless/datagrabbers/tests/test_ucla.py index eb7f224ff..aa3ffdb75 100644 --- a/junifer/configs/juseless/datagrabbers/tests/test_ucla.py +++ b/junifer/configs/juseless/datagrabbers/tests/test_ucla.py @@ -13,6 +13,7 @@ import pytest from junifer.configs.juseless.datagrabbers import JuselessUCLA from junifer.utils.logging import configure_logging + # Check if the test is running on juseless if socket.gethostname() != "juseless": pytest.skip("These tests are only for juseless", allow_module_level=True) diff --git a/junifer/configs/juseless/datagrabbers/tests/test_ukb_vbm.py b/junifer/configs/juseless/datagrabbers/tests/test_ukb_vbm.py index 1a4d2e053..3cda8b42c 100644 --- a/junifer/configs/juseless/datagrabbers/tests/test_ukb_vbm.py +++ b/junifer/configs/juseless/datagrabbers/tests/test_ukb_vbm.py @@ -12,6 +12,7 @@ import pytest from junifer.configs.juseless.datagrabbers import JuselessDataladUKBVBM from junifer.utils.logging import configure_logging + # Check if the test is running on juseless if socket.gethostname() != "juseless": pytest.skip("These tests are only for juseless", allow_module_level=True) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index ca74a53ee..3a7f3e74a 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -12,6 +12,7 @@ import pytest from junifer.datagrabber.hcp import DataladHCP1200 from junifer.utils.logging import configure_logging + # Check if the test is running on juseless if socket.gethostname() != "juseless": pytest.skip("These tests are only for juseless", allow_module_level=True) diff --git a/junifer/data/__init__.py b/junifer/data/__init__.py index f1508945b..f75220f6b 100644 --- a/junifer/data/__init__.py +++ b/junifer/data/__init__.py @@ -21,4 +21,4 @@ from .masks import ( register_mask, ) -from . import utils \ No newline at end of file +from . import utils diff --git a/junifer/data/coordinates.py b/junifer/data/coordinates.py index bb1500161..8d6ef224e 100644 --- a/junifer/data/coordinates.py +++ b/junifer/data/coordinates.py @@ -13,6 +13,7 @@ from numpy.typing import ArrayLike from ..utils.logging import logger, raise_error + # Path to the VOIs _vois_path = Path(__file__).parent / "VOIs" diff --git a/junifer/data/masks.py b/junifer/data/masks.py index 710ad0bde..611176c2e 100644 --- a/junifer/data/masks.py +++ b/junifer/data/masks.py @@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union import nibabel as nib -from .utils import closest_resolution from ..utils.logging import logger, raise_error +from .utils import closest_resolution if TYPE_CHECKING: @@ -58,10 +58,9 @@ def register_mask( if name in _available_masks: if overwrite is True: logger.info(f"Overwriting {name} mask") - if (_available_masks[name]["family"] != "CustomUserMask"): + if _available_masks[name]["family"] != "CustomUserMask": raise_error( - f"Cannot overwrite {name} mask. " - "It is a built-in mask." + f"Cannot overwrite {name} mask. " "It is a built-in mask." ) else: raise_error( @@ -117,8 +116,7 @@ def load_mask( """ if name not in _available_masks: raise_error( - f"Mask {name} not found. " - f"Valid options are: {list_masks()}" + f"Mask {name} not found. " f"Valid options are: {list_masks()}" ) mask_definition = _available_masks[name].copy() @@ -126,12 +124,10 @@ def load_mask( if t_family == "CustomUserMask": mask_fname = Path(mask_definition["path"]) - elif t_family == 'Vickery-Patil': + elif t_family == "Vickery-Patil": mask_fname = _load_vickery_patil_mask(name, resolution) else: - raise_error( - f"I don't know about the {t_family} mask family." - ) + raise_error(f"I don't know about the {t_family} mask family.") logger.info(f"Loading mask {mask_fname.absolute()}") @@ -167,8 +163,9 @@ def _load_vickery_patil_mask( available_resolutions = [1.5, 3.0] to_load = closest_resolution(resolution, available_resolutions) if to_load == 3.0: - mask_fname = \ + mask_fname = ( "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz" + ) elif to_load == 1.5: mask_fname = "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz" else: @@ -178,9 +175,7 @@ def _load_vickery_patil_mask( elif name == "GM_prob0.2_cortex": mask_fname = "GMprob0.2_cortex_3mm_NA_rm.nii.gz" else: - raise_error( - f"Cannot find a Vickery-Patil mask called {name}" - ) + raise_error(f"Cannot find a Vickery-Patil mask called {name}") mask_fname = _masks_path / "vickery-patil" / mask_fname return mask_fname diff --git a/junifer/data/parcellations.py b/junifer/data/parcellations.py index e9c7c5baf..f51aa0d41 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -18,8 +18,9 @@ import pandas as pd import requests from nilearn import datasets -from .utils import closest_resolution from ..utils.logging import logger, raise_error +from .utils import closest_resolution + if TYPE_CHECKING: from nibabel import Nifti1Image diff --git a/junifer/data/tests/test_data_utils.py b/junifer/data/tests/test_data_utils.py index 80b2adf66..614250b84 100644 --- a/junifer/data/tests/test_data_utils.py +++ b/junifer/data/tests/test_data_utils.py @@ -5,8 +5,8 @@ from typing import List -import pytest import numpy as np +import pytest from junifer.data.utils import closest_resolution diff --git a/junifer/data/tests/test_masks.py b/junifer/data/tests/test_masks.py index 13926774f..b00a76df6 100644 --- a/junifer/data/tests/test_masks.py +++ b/junifer/data/tests/test_masks.py @@ -11,10 +11,10 @@ import pytest from numpy.testing import assert_array_almost_equal from junifer.data.masks import ( + _load_vickery_patil_mask, + list_masks, load_mask, register_mask, - list_masks, - _load_vickery_patil_mask, ) diff --git a/junifer/data/tests/test_parcellations.py b/junifer/data/tests/test_parcellations.py index 1e04f10f4..fcc01c5e2 100644 --- a/junifer/data/tests/test_parcellations.py +++ b/junifer/data/tests/test_parcellations.py @@ -8,11 +8,10 @@ from pathlib import Path from typing import List -import pytest -from numpy.testing import assert_array_almost_equal, assert_array_equal - import nibabel as nib +import pytest from nilearn.image import new_img_like +from numpy.testing import assert_array_almost_equal, assert_array_equal from junifer.data.parcellations import ( _retrieve_parcellation, diff --git a/junifer/data/utils.py b/junifer/data/utils.py index 56450e0dd..97fb524d6 100644 --- a/junifer/data/utils.py +++ b/junifer/data/utils.py @@ -1,5 +1,5 @@ """Provide utilities for data module.""" -from typing import Optional, Union, List +from typing import List, Optional, Union import numpy as np diff --git a/junifer/datagrabber/aomic/tests/test_id1000.py b/junifer/datagrabber/aomic/tests/test_id1000.py index 532777218..925bbce66 100644 --- a/junifer/datagrabber/aomic/tests/test_id1000.py +++ b/junifer/datagrabber/aomic/tests/test_id1000.py @@ -111,8 +111,8 @@ def test_aomic1000_datagrabber() -> None: assert out["DWI"]["path"].is_file() # asserts meta - assert "meta" in out - meta = out["meta"] + assert "meta" in out["BOLD"] + meta = out["BOLD"]["meta"] assert "element" in meta assert "subject" in meta["element"] assert test_element == meta["element"]["subject"] diff --git a/junifer/datagrabber/aomic/tests/test_piop1.py b/junifer/datagrabber/aomic/tests/test_piop1.py index b3f44f542..1378c2dd1 100644 --- a/junifer/datagrabber/aomic/tests/test_piop1.py +++ b/junifer/datagrabber/aomic/tests/test_piop1.py @@ -125,8 +125,8 @@ def test_aomic_piop1_datagrabber() -> None: assert out["DWI"]["path"].is_file() # asserts meta - assert "meta" in out - meta = out["meta"] + assert "meta" in out["BOLD"] + meta = out["BOLD"]["meta"] assert "element" in meta assert "subject" in meta["element"] assert sub == meta["element"]["subject"] diff --git a/junifer/datagrabber/aomic/tests/test_piop2.py b/junifer/datagrabber/aomic/tests/test_piop2.py index a1b8850d0..19f45704d 100644 --- a/junifer/datagrabber/aomic/tests/test_piop2.py +++ b/junifer/datagrabber/aomic/tests/test_piop2.py @@ -121,8 +121,8 @@ def test_aomic_piop2_datagrabber() -> None: assert out["DWI"]["path"].is_file() # asserts meta - assert "meta" in out - meta = out["meta"] + assert "meta" in out["BOLD"] + meta = out["BOLD"]["meta"] assert "element" in meta assert "subject" in meta["element"] assert sub == meta["element"]["subject"] diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index ee6c9c22d..5931bfb2e 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -9,11 +9,12 @@ from abc import ABC, abstractmethod from pathlib import Path from typing import Dict, Iterator, List, Tuple, Union +from ..pipeline import UpdateMetaMixin from ..utils import logger, raise_error from .utils import validate_types -class BaseDataGrabber(ABC): +class BaseDataGrabber(ABC, UpdateMetaMixin): """Abstract base class for datagrabber. 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)) logger.debug(f"Named element: {named_element}") out = self.get_item(**named_element) - out["meta"] = { - "datagrabber": self.get_meta(), - "element": named_element, - } + + for _, t_val in out.items(): + self.update_meta(t_val, "datagrabber") + t_val["meta"]["element"] = named_element + return out def __enter__(self) -> "BaseDataGrabber": @@ -103,22 +105,6 @@ class BaseDataGrabber(ABC): """ 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 def datadir(self) -> Path: """Get data directory path. diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 3237f33a3..49ac30b8e 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -83,7 +83,7 @@ class DataladDataGrabber(BaseDataGrabber): self._rootdir = rootdir # Flag to indicate if the dataset was cloned before and it might be # dirty - self._dataset_dirty = False + self.datalad_dirty = False @property def datadir(self) -> Path: @@ -175,7 +175,7 @@ class DataladDataGrabber(BaseDataGrabber): # Check for dirty datasets: status = self._dataset.status() if any([x["state"] != "clean" for x in status]): - self._dataset_dirty = True + self.datalad_dirty = True warn_with_log( "At least one file is not clean, Junifer will " "consider this dataset as dirty." @@ -191,11 +191,10 @@ class DataladDataGrabber(BaseDataGrabber): logger.debug("Dataset installed") self._was_cloned = not isinstalled - self._datalad_commit_id = ( - self._dataset.repo.get_hexsha( # type: ignore - self._dataset.repo.get_corresponding_branch() # type: ignore - ) + self.datalad_commit_id = self._dataset.repo.get_hexsha( # type: ignore + self._dataset.repo.get_corresponding_branch() # type: ignore ) + self.datalad_id = self._dataset.id def cleanup(self) -> None: """Cleanup the datalad dataset.""" @@ -245,21 +244,3 @@ class DataladDataGrabber(BaseDataGrabber): logger.debug("Cleaning up dataset") self.cleanup() 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 diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py index 512765c72..ea8dd8eb2 100644 --- a/junifer/datagrabber/multiple.py +++ b/junifer/datagrabber/multiple.py @@ -5,7 +5,6 @@ # Synchon Mandal # License: AGPL -from pathlib import Path from typing import Dict, List, Tuple, Union from .base import BaseDataGrabber @@ -38,7 +37,7 @@ class MultipleDataGrabber(BaseDataGrabber): raise ValueError("Datagrabbers have overlapping types.") self._datagrabbers = datagrabbers - def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + def __getitem__(self, element: Union[str, Tuple]) -> Dict: """Implement indexing. Parameters @@ -58,9 +57,20 @@ class MultipleDataGrabber(BaseDataGrabber): """ out = {} + metas = [] for dg in self._datagrabbers: t_out = dg[element] 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 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()] 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 diff --git a/junifer/datagrabber/tests/test_base.py b/junifer/datagrabber/tests/test_base.py index 4eaa054f1..2d73b8031 100644 --- a/junifer/datagrabber/tests/test_base.py +++ b/junifer/datagrabber/tests/test_base.py @@ -23,7 +23,7 @@ def test_BaseDataGrabber() -> None: # Create concrete class. class MyDataGrabber(BaseDataGrabber): def get_item(self, subject): - return {} + return {"BOLD": {}} def get_elements(self): return super().get_elements() @@ -31,22 +31,24 @@ def test_BaseDataGrabber() -> None: def get_element_keys(self): return ["subject"] - dg = MyDataGrabber(datadir="/tmp", types=["func"]) - elem = dg["elem"] - assert "meta" in elem - assert "datagrabber" in elem["meta"] - assert "class" in elem["meta"]["datagrabber"] - assert MyDataGrabber.__name__ in elem["meta"]["datagrabber"]["class"] - assert "element" in elem["meta"] - assert "subject" in elem["meta"]["element"] - assert "elem" in elem["meta"]["element"]["subject"] + dg = MyDataGrabber(datadir="/tmp", types=["BOLD"]) + elem = dg["sub01"] + assert "BOLD" in elem + assert "meta" in elem["BOLD"] + meta = elem["BOLD"]["meta"] + assert "datagrabber" in meta + assert "class" in meta["datagrabber"] + assert MyDataGrabber.__name__ in meta["datagrabber"]["class"] + assert "element" in meta + assert "subject" in meta["element"] + assert "sub01" in meta["element"]["subject"] with pytest.raises(NotImplementedError): dg.get_elements() with dg: assert dg.datadir == Path("/tmp") - assert dg.types == ["func"] + assert dg.types == ["BOLD"] class MyDataGrabber2(BaseDataGrabber): def get_item(self, subject): @@ -58,7 +60,7 @@ def test_BaseDataGrabber() -> None: def get_element_keys(self): return super().get_element_keys() - dg = MyDataGrabber2(datadir="/tmp", types=["func"]) + dg = MyDataGrabber2(datadir="/tmp", types=["BOLD"]) with pytest.raises(NotImplementedError): dg.get_element_keys() diff --git a/junifer/datagrabber/tests/test_datalad_base.py b/junifer/datagrabber/tests/test_datalad_base.py index 50ebcbe76..4eee066b8 100644 --- a/junifer/datagrabber/tests/test_datalad_base.py +++ b/junifer/datagrabber/tests/test_datalad_base.py @@ -11,6 +11,7 @@ import pytest from junifer.datagrabber.datalad_base import DataladDataGrabber + _testing_dataset = { "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_symlink() is True elem1 = dg["sub-01"] - assert "meta" in elem1 - assert "datagrabber" in elem1["meta"] - assert "datalad_dirty" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False + assert "meta" in elem1["BOLD"] + meta = elem1["BOLD"]["meta"] + assert "datagrabber" in meta + assert "datalad_dirty" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_dirty"] is False assert hasattr(dg, "_got_files") is False assert datadir.exists() is True assert elem1_bold.is_file() is True @@ -198,14 +200,15 @@ def test_datalad_previously_cloned( assert datadir.exists() is True assert dg._was_cloned is False elem1 = dg["sub-01"] - assert "meta" in elem1 - assert "datagrabber" in elem1["meta"] - assert "datalad_dirty" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False - assert "datalad_commit_id" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit - assert "datalad_id" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id + assert "meta" in elem1["BOLD"] + meta = elem1["BOLD"]["meta"] + assert "datagrabber" in meta + assert "datalad_dirty" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_dirty"] is False + assert "datalad_commit_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_id"] == remote_id assert hasattr(dg, "_got_files") is True # 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 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_file() is False @@ -275,14 +279,15 @@ def test_datalad_previously_cloned_and_get( assert datadir.exists() is True assert dg._was_cloned is False elem1 = dg["sub-01"] - assert "meta" in elem1 - assert "datagrabber" in elem1["meta"] - assert "datalad_dirty" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False - assert "datalad_commit_id" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit - assert "datalad_id" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id + assert "meta" in elem1["BOLD"] + meta = elem1["BOLD"]["meta"] + assert "datagrabber" in meta + assert "datalad_dirty" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_dirty"] is False + assert "datalad_commit_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_id"] == remote_id assert hasattr(dg, "_got_files") is True # 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 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_file() is False @@ -359,14 +365,15 @@ def test_datalad_previously_cloned_and_get_dirty( assert datadir.exists() is True assert dg._was_cloned is False elem1 = dg["sub-01"] - assert "meta" in elem1 - assert "datagrabber" in elem1["meta"] - assert "datalad_dirty" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_dirty"] is True - assert "datalad_commit_id" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit - assert "datalad_id" in elem1["meta"]["datagrabber"] - assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id + assert "meta" in elem1["BOLD"] + meta = elem1["BOLD"]["meta"] + assert "datagrabber" in meta + assert "datalad_dirty" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_dirty"] is True + assert "datalad_commit_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_id"] == remote_id assert hasattr(dg, "_got_files") is True # 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 dg._was_cloned is False elem2 = dg["sub-02"] - assert "meta" in elem2 - assert "datagrabber" in elem2["meta"] - assert "datalad_dirty" in elem2["meta"]["datagrabber"] + assert "meta" in elem1["BOLD"] + meta = elem2["BOLD"]["meta"] + assert "datagrabber" in meta + assert "datalad_dirty" in meta["datagrabber"] # 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 elem2["meta"]["datagrabber"]["datalad_commit_id"] == commit - assert "datalad_id" in elem2["meta"]["datagrabber"] - assert elem2["meta"]["datagrabber"]["datalad_id"] == remote_id + assert "datalad_commit_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_commit_id"] == commit + assert "datalad_id" in meta["datagrabber"] + assert meta["datagrabber"]["datalad_id"] == remote_id assert hasattr(dg, "_got_files") is True # Files are there and symlinks are fixed diff --git a/junifer/datagrabber/tests/test_hcp.py b/junifer/datagrabber/tests/test_hcp.py index 6fb4a9434..bfbdf4cd2 100644 --- a/junifer/datagrabber/tests/test_hcp.py +++ b/junifer/datagrabber/tests/test_hcp.py @@ -10,6 +10,7 @@ import pytest from junifer.datagrabber.hcp import DataladHCP1200 from junifer.utils import configure_logging + 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 out["BOLD"]["path"].is_file() # Assert metadata - assert "meta" in out - meta = out["meta"] + assert "meta" in out["BOLD"] + meta = out["BOLD"]["meta"] assert "element" in meta assert "subject" in meta["element"] assert test_element[0] == meta["element"]["subject"] diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py index 47abdd525..7ad9c4dec 100644 --- a/junifer/datagrabber/tests/test_multiple.py +++ b/junifer/datagrabber/tests/test_multiple.py @@ -7,6 +7,7 @@ import pytest from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber + _testing_dataset = { "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] assert set(subs) == set(expected_subs) - data = dg[("sub-01", "ses-01")] - assert "T1w" in data - assert "BOLD" in data - - meta = dg.get_meta() - assert "class" in meta - assert meta["class"] == "MultipleDataGrabber" - assert "datagrabbers" in meta - assert len(meta["datagrabbers"]) == 2 - assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber" - assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber" + elem = dg[("sub-01", "ses-01")] + assert "T1w" in elem + assert "BOLD" in elem + assert "meta" in elem["BOLD"] + meta = elem["BOLD"]["meta"]["datagrabber"] + assert "class" in meta + assert meta["class"] == "MultipleDataGrabber" + assert "datagrabbers" in meta + assert len(meta["datagrabbers"]) == 2 + assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber" + assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber" def test_multiple_no_intersection() -> None: diff --git a/junifer/datagrabber/tests/test_pattern.py b/junifer/datagrabber/tests/test_pattern.py index 71ba60701..8d2ec5bc8 100644 --- a/junifer/datagrabber/tests/test_pattern.py +++ b/junifer/datagrabber/tests/test_pattern.py @@ -256,6 +256,6 @@ def test_PatternDataGrabber(tmp_path: Path) -> None: out1 = datagrabber[("sub000", "ses000", "task002")] out2 = datagrabber[("sub000", "ses000", "task003")] - assert out1["func"] == out2["func"] - assert out1["anat"] == out2["anat"] - assert out1["vbm"] != out2["vbm"] + assert out1["func"]["path"] == out2["func"]["path"] + assert out1["anat"]["path"] == out2["anat"]["path"] + assert out1["vbm"]["path"] != out2["vbm"]["path"] diff --git a/junifer/datagrabber/tests/test_pattern_datalad.py b/junifer/datagrabber/tests/test_pattern_datalad.py index 2deace6eb..e0986f1e7 100644 --- a/junifer/datagrabber/tests/test_pattern_datalad.py +++ b/junifer/datagrabber/tests/test_pattern_datalad.py @@ -11,6 +11,7 @@ import pytest from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber + _testing_dataset = { "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" ) - assert "meta" in t_sub - assert "datagrabber" in t_sub["meta"] - dg_meta = t_sub["meta"]["datagrabber"] + assert "meta" in t_sub["BOLD"] + meta = t_sub["BOLD"]["meta"] + assert "datagrabber" in meta + dg_meta = meta["datagrabber"] assert "class" in dg_meta assert dg_meta["class"] == "PatternDataladDataGrabber" assert "uri" in dg_meta diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index d1a74a8ef..7cfa43c2a 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -11,10 +11,11 @@ import nibabel as nib import pandas as pd from ..api.decorators import register_datareader -from ..pipeline import PipelineStepMixin +from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..utils.logging import logger, warn_with_log -# Map each file extension to a kind + +# Map each file extension to a type _extensions = { ".nii": "NIFTI", ".nii.gz": "NIFTI", @@ -22,7 +23,7 @@ _extensions = { ".tsv": "TSV", } -# Map each kind to a function and arguments +# Map each type to a function and arguments _readers = {} _readers["NIFTI"] = {"func": nib.load, "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 -class DefaultDataReader(PipelineStepMixin): +class DefaultDataReader(PipelineStepMixin, UpdateMetaMixin): """Mixin class for default data reader.""" def validate_input(self, input: List[str]) -> None: @@ -46,8 +47,8 @@ class DefaultDataReader(PipelineStepMixin): # Nothing to validate, any input is fine pass - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input: List[str]) -> List[str]: + """Get output type. Parameters ---------- @@ -58,10 +59,10 @@ class DefaultDataReader(PipelineStepMixin): Returns ------- 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 def fit_transform( @@ -82,36 +83,33 @@ class DefaultDataReader(PipelineStepMixin): ------- dict 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() if params is None: params = {} - for kind in input.keys(): - if kind == "meta": - out["meta"] = input["meta"] - continue - if "path" not in input[kind]: + for type_ in input.keys(): + if "path" not in input[type_]: warn_with_log( - f"Input kind {kind} does not provide a path. Skipping." + f"Input type {type_} does not provide a path. Skipping." ) continue - t_path = input[kind]["path"] - t_params = params.get(kind, {}) + t_path = input[type_]["path"] + t_params = params.get(type_, {}) # Convert to Path if datareader is not well done if not isinstance(t_path, Path): t_path = Path(t_path) - out[kind]["path"] = t_path - logger.info(f"Reading {kind} from {t_path.as_posix()}") + out[type_]["path"] = t_path + logger.info(f"Reading {type_} from {t_path.as_posix()}") fread = None fname = t_path.name.lower() for ext, ftype in _extensions.items(): if fname.endswith(ext): - logger.info(f"{kind} is type {ftype}") + logger.info(f"{type_} is type {ftype}") reader_func = _readers[ftype]["func"] reader_params = _readers[ftype]["params"] if reader_params is not None: @@ -123,8 +121,6 @@ class DefaultDataReader(PipelineStepMixin): logger.info( f"Unknown file type {t_path.as_posix()}, skipping reading" ) - out[kind]["data"] = fread - if "meta" not in out: - out["meta"] = {} - out["meta"]["datareader"] = self.get_meta() + out[type_]["data"] = fread + self.update_meta(out[type_], "datareader") return out diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index 4ebf61b34..61022472d 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -17,37 +17,36 @@ from junifer.datareader import DefaultDataReader @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. Parameters ---------- - kind : list of str or str or None - The parametrized kind of data. + type_ : list of str or str or None + The parametrized type_ of data. """ reader = DefaultDataReader() - assert reader.validate_input(kind) is None - assert reader.get_output_kind(kind) == kind - assert reader.validate(kind) == kind + assert reader.validate_input(type_) is None + assert reader.get_output_type(type_) == type_ + assert reader.validate(type_) == type_ def test_meta() -> None: """Test reader metadata.""" reader = DefaultDataReader() - t_meta = reader.get_meta() - assert t_meta["class"] == "DefaultDataReader" nib_data_path = Path(nib_testing.data_path) t_path = nib_data_path / "example4d.nii.gz" input = {"BOLD": {"path": t_path}} output = reader.fit_transform(input) - assert "meta" in output - assert "datareader" in output["meta"] - assert "class" in output["meta"]["datareader"] - assert output["meta"]["datareader"]["class"] == "DefaultDataReader" + assert "meta" in output["BOLD"] + meta = output["BOLD"]["meta"] + assert "datareader" in meta + assert "class" in meta["datareader"] + assert meta["datareader"]["class"] == "DefaultDataReader" @pytest.mark.parametrize( diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 04087ef9a..7711cb1a7 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -7,14 +7,15 @@ from abc import ABC, abstractmethod 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 + if TYPE_CHECKING: from junifer.storage import BaseFeatureStorage -class BaseMarker(ABC, PipelineStepMixin): +class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin): """Abstract base class for all markers. Parameters @@ -57,27 +58,6 @@ class BaseMarker(ABC, PipelineStepMixin): 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: """Validate input. @@ -101,23 +81,22 @@ class BaseMarker(ABC, PipelineStepMixin): ) @abstractmethod - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The input to the marker. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the marker. Returns ------- - list of str - The updated list of output kinds, as storage possibilities. + str + The storage type output by the marker. """ raise_error( - msg="Concrete classes need to implement get_output_kind().", + msg="Concrete classes need to implement get_output_type().", klass=NotImplementedError, ) @@ -149,10 +128,9 @@ class BaseMarker(ABC, PipelineStepMixin): klass=NotImplementedError, ) - @abstractmethod def store( self, - kind: str, + type_: str, out: Dict[str, Any], storage: "BaseFeatureStorage", ) -> None: @@ -160,18 +138,17 @@ class BaseMarker(ABC, PipelineStepMixin): Parameters ---------- - kind : str - The data kind to store. + type_ : str + The data type to store. out : dict The computed result as a dictionary to store. storage : storage-like - The storage class. + The storage class, for example, SQLiteFeatureStorage. """ - raise_error( - msg="Concrete classes need to implement store().", - klass=NotImplementedError, - ) + output_type_ = self.get_output_type(type_) + logger.debug(f"Storing {output_type_} in {storage}") + storage.store(kind=output_type_, **out) def fit_transform( self, @@ -195,23 +172,25 @@ class BaseMarker(ABC, PipelineStepMixin): """ out = {} - meta = input.get("meta", {}) - for kind in self._on: - if kind in input.keys(): - logger.info(f"Computing {kind}") - t_input = input[kind] + for type_ in self._on: + if type_ in input.keys(): + logger.info(f"Computing {type_}") + t_input = input[type_] extra_input = input.copy() - extra_input.pop(kind) - t_meta = meta.copy() - t_meta.update(t_input.get("meta", {})) - t_meta.update(self.get_meta(kind)) + extra_input.pop(type_) + t_meta = t_input["meta"].copy() + t_meta["type"] = type_ + 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: 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: logger.info("No storage specified, returning dictionary") - out[kind] = t_out + out[type_] = t_out return out diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index 2b94642e9..d2e1f3e61 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -13,6 +13,7 @@ from ..pipeline import PipelineStepMixin from ..storage.base import BaseFeatureStorage from ..utils import logger + if TYPE_CHECKING: from junifer.datagrabber import BaseDataGrabber diff --git a/junifer/markers/crossparcellation_functional_connectivity.py b/junifer/markers/crossparcellation_functional_connectivity.py index 4e8ccbad7..e57e9180e 100644 --- a/junifer/markers/crossparcellation_functional_connectivity.py +++ b/junifer/markers/crossparcellation_functional_connectivity.py @@ -9,7 +9,6 @@ from typing import Any, Dict, List, Optional import pandas as pd from ..api.decorators import register_marker -from ..storage import BaseFeatureStorage from ..utils import logger from ..utils.logging import raise_error from .base import BaseMarker @@ -41,6 +40,8 @@ class CrossParcellationFC(BaseMarker): (default None). """ + _DEPENDENCIES = {"nilearn"} + def __init__( self, parcellation_one: str, @@ -72,43 +73,21 @@ class CrossParcellationFC(BaseMarker): """ return ["BOLD"] - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The input to the marker. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the marker. Returns ------- - list of str - The updated list of output kinds, as storage possibilities. + str + The storage type output by the marker. """ - 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) + return "matrix" def compute( self, diff --git a/junifer/markers/ets_rss.py b/junifer/markers/ets_rss.py index a7e7f939a..5605fa28a 100644 --- a/junifer/markers/ets_rss.py +++ b/junifer/markers/ets_rss.py @@ -6,7 +6,7 @@ # Synchon Mandal # License: AGPL -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union import numpy as np @@ -16,9 +16,6 @@ from .base import BaseMarker from .parcel_aggregation import ParcelAggregation from .utils import _ets -if TYPE_CHECKING: - from junifer.storage import BaseFeatureStorage - @register_marker class RSSETSMarker(BaseMarker): @@ -45,6 +42,8 @@ class RSSETSMarker(BaseMarker): """ + _DEPENDENCIES = {"nilearn"} + def __init__( self, parcellation: Union[str, List[str]], @@ -70,43 +69,21 @@ class RSSETSMarker(BaseMarker): """ return ["BOLD"] - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The input to the marker. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the marker. Returns ------- - list of str - The updated list of output kinds, as storage possibilities. + str + The storage type output by the marker. """ - 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) + return "timeseries" def compute( self, @@ -150,7 +127,7 @@ class RSSETSMarker(BaseMarker): parcellation=self.parcellation, method=self.agg_method, method_params=self.agg_method_params, - mask=self.mask + mask=self.mask, ) # Compute the parcel aggregation out = parcel_aggregation.compute(input=input, extra_input=extra_input) diff --git a/junifer/markers/functional_connectivity_parcels.py b/junifer/markers/functional_connectivity_parcels.py index b42c1171d..649038e6d 100644 --- a/junifer/markers/functional_connectivity_parcels.py +++ b/junifer/markers/functional_connectivity_parcels.py @@ -4,19 +4,15 @@ # Kaustubh R. Patil # 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 sklearn.covariance import EmpiricalCovariance from ..api.decorators import register_marker -from ..utils import logger from .base import BaseMarker from .parcel_aggregation import ParcelAggregation -if TYPE_CHECKING: - from junifer.storage import BaseFeatureStorage - @register_marker class FunctionalConnectivityParcels(BaseMarker): @@ -49,6 +45,8 @@ class FunctionalConnectivityParcels(BaseMarker): None). """ + _DEPENDENCIES = {"nilearn", "scikit-learn"} + def __init__( self, parcellation: Union[str, List[str]], @@ -83,22 +81,21 @@ class FunctionalConnectivityParcels(BaseMarker): """ return ["BOLD"] - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The input to the marker. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the marker. Returns ------- - list of str - The updated list of output kinds, as storage possibilities. + str + The storage type output by the marker. + """ - outputs = ["matrix"] - return outputs + return "matrix" def compute( self, @@ -154,24 +151,3 @@ class FunctionalConnectivityParcels(BaseMarker): out["col_names"] = ts["columns"] out["matrix_kind"] = "tril" 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) diff --git a/junifer/markers/functional_connectivity_spheres.py b/junifer/markers/functional_connectivity_spheres.py index 193510220..97b485e7b 100644 --- a/junifer/markers/functional_connectivity_spheres.py +++ b/junifer/markers/functional_connectivity_spheres.py @@ -4,19 +4,16 @@ # Kaustubh R. Patil # License: AGPL -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import Any, Dict, List, Optional from nilearn.connectome import ConnectivityMeasure from sklearn.covariance import EmpiricalCovariance from ..api.decorators import register_marker -from ..utils import logger, raise_error +from ..utils import raise_error from .base import BaseMarker from .sphere_aggregation import SphereAggregation -if TYPE_CHECKING: - from junifer.storage import BaseFeatureStorage - @register_marker class FunctionalConnectivitySpheres(BaseMarker): @@ -54,6 +51,8 @@ class FunctionalConnectivitySpheres(BaseMarker): """ + _DEPENDENCIES = {"nilearn", "scikit-learn"} + def __init__( self, coords: str, @@ -93,23 +92,21 @@ class FunctionalConnectivitySpheres(BaseMarker): """ return ["BOLD"] - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The input to the marker. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the marker. Returns ------- - list of str - The updated list of output kinds, as storage possibilities. + str + The storage type output by the marker. """ - outputs = ["matrix"] - return outputs + return "matrix" def compute( self, @@ -166,24 +163,3 @@ class FunctionalConnectivitySpheres(BaseMarker): out["col_names"] = ts["columns"] out["matrix_kind"] = "tril" 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) diff --git a/junifer/markers/parcel_aggregation.py b/junifer/markers/parcel_aggregation.py index f76c6c8c4..0c4388629 100644 --- a/junifer/markers/parcel_aggregation.py +++ b/junifer/markers/parcel_aggregation.py @@ -4,21 +4,18 @@ # Synchon Mandal # License: AGPL -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union 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 ..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 ..utils import logger from .base import BaseMarker -if TYPE_CHECKING: - from junifer.storage import BaseFeatureStorage - @register_marker class ParcelAggregation(BaseMarker): @@ -48,6 +45,8 @@ class ParcelAggregation(BaseMarker): None). """ + _DEPENDENCIES = {"nilearn", "numpy"} + def __init__( self, parcellation: Union[str, List[str]], @@ -76,53 +75,27 @@ class ParcelAggregation(BaseMarker): """ return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The kind of data to work on. + input_type : str + The data type input to the marker. Returns ------- - list of str - The list of storage kinds. + str + 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( - self, - kind: str, - out: Dict[str, Any], - storage: "BaseFeatureStorage", - ) -> None: - """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) + if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: + return "table" + elif input_type == "BOLD": + return "timeseries" + else: + raise ValueError(f"Unknown input kind for {input_type}") def compute( self, @@ -252,6 +225,4 @@ class ParcelAggregation(BaseMarker): out_values = np.array(out_values).T out = {"data": out_values, "columns": out_labels} - if out_values.shape[0] > 1: - out["row_names"] = "scan" return out diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index 8fa648b4c..20c16ccfc 100644 --- a/junifer/markers/sphere_aggregation.py +++ b/junifer/markers/sphere_aggregation.py @@ -4,7 +4,7 @@ # Synchon Mandal # 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 ..data import load_coordinates, load_mask @@ -13,9 +13,6 @@ from ..stats import get_aggfunc_by_name from ..utils import logger from .base import BaseMarker -if TYPE_CHECKING: - from junifer.storage import BaseFeatureStorage - @register_marker class SphereAggregation(BaseMarker): @@ -51,6 +48,8 @@ class SphereAggregation(BaseMarker): """ + _DEPENDENCIES = {"nilearn", "numpy"} + def __init__( self, coords: str, @@ -79,53 +78,27 @@ class SphereAggregation(BaseMarker): """ return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The kind of data to work on. + input_type : str + The data type input to the marker. Returns ------- - list of str - The list of storage kinds. + str + 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( - self, - kind: str, - out: Dict[str, Any], - storage: "BaseFeatureStorage", - ) -> None: - """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) + if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: + return "table" + elif input_type == "BOLD": + return "timeseries" + else: + raise ValueError(f"Unknown input kind for {input_type}") def compute( self, @@ -180,6 +153,4 @@ class SphereAggregation(BaseMarker): out_values = masker.fit_transform(t_input) # Format the output out = {"data": out_values, "columns": out_labels} - if out_values.shape[0] > 1: - out["row_names"] = "scan" return out diff --git a/junifer/markers/tests/test_crossparcellation_functional_connectivity.py b/junifer/markers/tests/test_crossparcellation_functional_connectivity.py index 5a6dee709..8e9fd53da 100644 --- a/junifer/markers/tests/test_crossparcellation_functional_connectivity.py +++ b/junifer/markers/tests/test_crossparcellation_functional_connectivity.py @@ -15,6 +15,7 @@ from junifer.markers.crossparcellation_functional_connectivity import ( from junifer.storage import SQLiteFeatureStorage from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber + parcellation_ONE = "Schaefer100x17" parcellation_TWO = "Schaefer200x17" @@ -30,8 +31,7 @@ def test_compute() -> None: "data": niimg, "path": out["BOLD"]["path"], "meta": {"element": "sub001"}, - }, - "meta": {"element": "sub001"}, + } } crossparcellation = CrossParcellationFC( @@ -43,12 +43,6 @@ def test_compute() -> None: assert out["data"].shape == (200, 100) assert len(out["col_names"]) == 100 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: @@ -62,16 +56,10 @@ def test_store(tmp_path: Path) -> None: """ with SPMAuditoryTestingDatagrabber() as dg: - out = dg["sub001"] - niimg = image.load_img(str(out["BOLD"]["path"].absolute())) - input_dict = { - "BOLD": { - "data": niimg, - "path": out["BOLD"]["path"], - "meta": {"element": "sub001"}, - }, - "meta": {"element": "sub001"}, - } + input_dict = dg["sub001"] + niimg = image.load_img(str(input_dict["BOLD"]["path"].absolute())) + + input_dict["BOLD"]["data"] = niimg crossparcellation = CrossParcellationFC( parcellation_one=parcellation_ONE, @@ -80,19 +68,23 @@ def test_store(tmp_path: Path) -> None: ) uri = tmp_path / "test_crossparcellation.sqlite" 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: - """Test CrossParcellationFC get_output_kind().""" +def test_get_output_type() -> None: + """Test CrossParcellationFC get_output_type().""" crossparcellation = CrossParcellationFC( parcellation_one=parcellation_ONE, parcellation_two=parcellation_TWO ) - input_list = ["BOLD"] - input_list = crossparcellation.get_output_kind(input_list) - assert len(input_list) == 1 - assert input_list[0] in ["matrix"] + input_ = "BOLD" + output = crossparcellation.get_output_type(input_) + assert output == "matrix" def test_init_() -> None: diff --git a/junifer/markers/tests/test_ets_rss.py b/junifer/markers/tests/test_ets_rss.py index be390da88..d18ff4a81 100644 --- a/junifer/markers/tests/test_ets_rss.py +++ b/junifer/markers/tests/test_ets_rss.py @@ -16,6 +16,7 @@ from junifer.markers.ets_rss import RSSETSMarker from junifer.storage import SQLiteFeatureStorage from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber + # Set parcellation PARCELLATION = "Schaefer100x17" @@ -42,21 +43,13 @@ def test_compute() -> None: n_time, _ = test_ts.shape 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_kind() -> None: - """Test RSS ETS get_output_kind().""" +def test_get_output_type() -> None: + """Test RSS ETS get_output_type().""" ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) - input_list = ["BOLD"] - input_list = ets_rss_marker.get_output_kind(input_list) - assert len(input_list) == 1 - assert input_list[0] in ["timeseries"] + input_ = "BOLD" + output = ets_rss_marker.get_output_type(input_) + assert output == "timeseries" def test_store(tmp_path: Path) -> None: @@ -70,14 +63,18 @@ def test_store(tmp_path: Path) -> None: """ with SPMAuditoryTestingDatagrabber() as dg: # Fetch element - out = dg["sub001"] + elem = dg["sub001"] # Load BOLD image - niimg = image.load_img(str(out["BOLD"]["path"].absolute())) - input_dict = {"data": niimg, "path": out["BOLD"]["path"]} + niimg = image.load_img(str(elem["BOLD"]["path"].absolute())) + elem["BOLD"]["data"] = niimg # Compute the RSSETSMarker ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION) # Create storage storage = SQLiteFeatureStorage( - uri=str((tmp_path / "test.sqlite").absolute())) + uri=str((tmp_path / "test.sqlite").absolute()) + ) # 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()) diff --git a/junifer/markers/tests/test_functional_connectivity_parcels.py b/junifer/markers/tests/test_functional_connectivity_parcels.py index 667c8e246..5b5758970 100644 --- a/junifer/markers/tests/test_functional_connectivity_parcels.py +++ b/junifer/markers/tests/test_functional_connectivity_parcels.py @@ -32,7 +32,7 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None: fmri_img = image.concat_imgs(ni_data.func) # type: ignore 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"] @@ -48,7 +48,11 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None: pa = ParcelAggregation( 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 # 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) # check correct output - assert fc.get_output_kind(["BOLD"]) == ["matrix"] + assert fc.get_output_type("BOLD") == "matrix" # Check empirical correlation method parameters fc = FunctionalConnectivityParcels( parcellation="Schaefer100x7", cor_method_params={"empirical": True} ) - all_out = fc.fit_transform({"BOLD": {"data": fmri_img}}) + all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": meta}}) uri = tmp_path / "test_fc_parcellation.sqlite" # Single storage, must be the uri storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") - meta = { - "element": "test", - "version": "0.0.1", - "marker": {"name": "fcname"}, - } - input = {"BOLD": {"data": fmri_img}, "meta": meta} + meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}} + input = {"BOLD": {"data": fmri_img, "meta": meta}} all_out = fc.fit_transform(input, storage=storage) + + features = storage.list_features() + assert any( + x["name"] == "BOLD_FunctionalConnectivityParcels" + for x in features.values() + ) diff --git a/junifer/markers/tests/test_functional_connectivity_spheres.py b/junifer/markers/tests/test_functional_connectivity_spheres.py index c4eaf86fe..819d4a6ad 100644 --- a/junifer/markers/tests/test_functional_connectivity_spheres.py +++ b/junifer/markers/tests/test_functional_connectivity_spheres.py @@ -36,7 +36,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None: fc = FunctionalConnectivitySpheres( 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"] @@ -52,7 +52,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None: sa = SphereAggregation( 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 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) # check correct output - assert fc.get_output_kind(["BOLD"]) == ["matrix"] + assert fc.get_output_type("BOLD") == "matrix" uri = tmp_path / "test_fc_parcel.sqlite" # Single storage, must be the uri storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") meta = { - "element": "test", - "version": "0.0.1", - "marker": {"name": "fcname"}, + "element": {"subject": "test"}, + "dependencies": {"numpy", "nilearn"}, } - input = {"BOLD": {"data": fmri_img}, "meta": meta} + input = {"BOLD": {"data": fmri_img, "meta": meta}} 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: """Test FunctionalConnectivitySpheres with empirical covariance. @@ -94,7 +99,7 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None: cor_method="correlation", 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"] diff --git a/junifer/markers/tests/test_markers_base.py b/junifer/markers/tests/test_markers_base.py index 220d3dbfe..b8590b02f 100644 --- a/junifer/markers/tests/test_markers_base.py +++ b/junifer/markers/tests/test_markers_base.py @@ -19,10 +19,14 @@ def test_base_marker_subclassing() -> None: """Test proper subclassing of BaseMarker.""" # Create concrete class class MyBaseMarker(BaseMarker): + def __init__(self, on, name=None) -> None: + self.parameter = 1 + super().__init__(on, name) + def get_valid_inputs(self): return ["BOLD", "T1w"] - def get_output_kind(self, input): + def get_output_type(self, input): return ["timeseries"] def compute(self, input, extra_input): @@ -32,18 +36,19 @@ def test_base_marker_subclassing() -> None: "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'\]"): MyBaseMarker(on=["BOLD", "T2w"]) # Create input for marker input_ = { - "meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"}, "BOLD": { "path": ".", "data": "data", + "meta": { + "datagrabber": "dg", + "element": "elem", + "datareader": "dr", + }, }, } marker = MyBaseMarker(on=["BOLD"]) @@ -57,14 +62,16 @@ def test_base_marker_subclassing() -> None: assert "data" in output["BOLD"] assert "columns" 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 - with pytest.raises(NotImplementedError): - marker.store(kind="kind", out="out", storage="storage") # type: ignore + assert "meta" in output["BOLD"] + meta = output["BOLD"]["meta"] + 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 assert marker.name == "MyBaseMarker" diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index cace74bba..3274e0ad4 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -3,18 +3,21 @@ # Authors: Federico Raimondo # Synchon Mandal # License: AGPL + +from pathlib import Path + import nibabel as nib import numpy as np import pytest -from pathlib import Path 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 numpy.testing import assert_array_almost_equal, assert_array_equal from scipy.stats import trim_mean from junifer.data import load_mask, load_parcellation, register_parcellation from junifer.markers.parcel_aggregation import ParcelAggregation +from junifer.storage import SQLiteFeatureStorage def test_ParcelAggregation_input_output() -> None: @@ -22,12 +25,11 @@ def test_ParcelAggregation_input_output() -> None: marker = ParcelAggregation( parcellation="Schaefer100x7", method="mean", on="VBM_GM" ) - - output = marker.get_output_kind(["VBM_GM", "BOLD"]) - assert output == ["table", "timeseries"] + for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]: + assert marker.get_output_type(in_) == out_ 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: @@ -78,22 +80,13 @@ def test_ParcelAggregation_3D() -> None: name="gmd_schaefer100x7_mean", on="VBM_GM", ) # 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"] assert jun_values3d_mean.ndim == 2 assert jun_values3d_mean.shape[0] == 1 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) manual = [] for t_v in sorted(np.unique(parcellation_values)): @@ -103,22 +96,13 @@ def test_ParcelAggregation_3D() -> None: # Use the ParcelAggregation object 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"] assert jun_values3d_std.ndim == 2 assert jun_values3d_std.shape[0] == 1 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 manual = [] for t_v in sorted(np.unique(parcellation_values)): @@ -136,22 +120,13 @@ def test_ParcelAggregation_3D() -> None: method="trim_mean", 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"] assert jun_values3d_tm.ndim == 2 assert jun_values3d_tm.shape[0] == 1 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(): """Test ParcelAggregation object on 4D images.""" @@ -170,21 +145,63 @@ def test_ParcelAggregation_4D(): # Create ParcelAggregation object 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"] assert jun_values4d.ndim == 2 assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d, jun_values4d) - meta = marker.get_meta("BOLD")["marker"] - assert meta["method"] == "mean" - assert meta["parcellation"] == ["Schaefer100x7"] - assert meta["mask"] is None - assert meta["name"] == "BOLD_ParcelAggregation" - assert meta["class"] == "ParcelAggregation" - assert meta["kind"] == "BOLD" - assert meta["method_params"] == {} + +def test_ParcelAggregation_storage(tmp_path: Path) -> None: + """Test ParcelAggregation storage. + + Parameters + ---------- + tmp_path : pathlib.Path + 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: @@ -215,22 +232,13 @@ def test_ParcelAggregation_3D_mask() -> None: name="gmd_schaefer100x7_mean", on="VBM_GM", ) # 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"] assert jun_values3d_mean.ndim == 2 assert jun_values3d_mean.shape[0] == 1 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: """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", on="VBM_GM", ) # 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_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_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 marker_split = ParcelAggregation( parcellation=["Schaefer100x7_low", "Schaefer100x7_high"], @@ -306,7 +305,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None: name="gmd_schaefer100x7_mean", on="VBM_GM", ) # 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_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[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 assert_array_equal(orig_mean_data, split_mean_data) 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", on="VBM_GM", ) # 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_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_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 marker_split = ParcelAggregation( parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"], @@ -404,7 +385,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None: name="gmd_schaefer100x7_mean", on="VBM_GM", ) # 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_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[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 assert_array_equal(orig_mean_data, split_mean_data) diff --git a/junifer/markers/tests/test_sphere_aggregation.py b/junifer/markers/tests/test_sphere_aggregation.py index 5c3311387..5081fb324 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -3,6 +3,8 @@ # Authors: Federico Raimondo # License: AGPL +import typing +from typing import Dict from pathlib import Path 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.storage import SQLiteFeatureStorage + # Define common variables COORDS = "DMNBuckner" RADIUS = 8 @@ -23,15 +26,12 @@ RADIUS = 8 def test_SphereAggregation_input_output() -> None: """Test SphereAggregation input and output types.""" - marker = SphereAggregation( - coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" - ) - - output = marker.get_output_kind(["VBM_GM", "BOLD"]) - assert output == ["table", "timeseries"] + marker = SphereAggregation(coords="DMNBuckner", method="mean", on="VBM_GM") + for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]: + assert marker.get_output_type(in_) == out_ 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: @@ -52,23 +52,13 @@ def test_SphereAggregation_3D() -> None: marker = SphereAggregation( 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"] assert jun_values4d.ndim == 2 assert_array_equal(auto4d.shape, jun_values4d.shape) 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: """Test SphereAggregation object on 4D images.""" @@ -84,26 +74,14 @@ def test_SphereAggregation_4D() -> None: auto4d = nifti_masker.fit_transform(fmri_img) # Create SphereAggregation object - marker = SphereAggregation( - coords=COORDS, method="mean", radius=RADIUS - ) - input = {"BOLD": {"data": fmri_img}} + marker = SphereAggregation(coords=COORDS, method="mean", radius=RADIUS) + input = {"BOLD": {"data": fmri_img, "meta": {}}} jun_values4d = marker.fit_transform(input)["BOLD"]["data"] assert jun_values4d.ndim == 2 assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d, jun_values4d) - 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: """Test SphereAggregation storage. @@ -122,31 +100,38 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None: storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") meta = { - "element": "test", - "version": "0.0.1", - "marker": {"name": "fcname"}, + "element": {"subject": "sub-01", "session": "ses-01"}, + "dependencies": {"nilearn", "nibabel"}, } - input = {"VBM_GM": {"data": img}, "meta": meta} + input = {"VBM_GM": {"data": img, "meta": meta}} marker = SphereAggregation( coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM" ) 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 = { - "element": "test", - "version": "0.0.1", - "marker": {"name": "BOLD_fcname"}, + "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} + input = {"BOLD": {"data": fmri_img, "meta": meta}} marker = SphereAggregation( coords=COORDS, method="mean", radius=RADIUS, on="BOLD" ) marker.fit_transform(input, storage=storage) + features: Dict = typing.cast(Dict, storage.list_features()) + assert any( + x["name"] == "BOLD_SphereAggregation" for x in features.values() + ) def test_SphereAggregation_3D_mask() -> None: @@ -164,26 +149,21 @@ def test_SphereAggregation_3D_mask() -> None: # Create NiftSpheresMasker 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) # Create SphereAggregation object marker = SphereAggregation( - coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM", - mask="GM_prob0.2" + coords=COORDS, + 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"] assert jun_values4d.ndim == 2 assert_array_equal(auto4d.shape, jun_values4d.shape) 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"] == {} diff --git a/junifer/pipeline/__init__.py b/junifer/pipeline/__init__.py index f0ddfb38d..8d1681388 100644 --- a/junifer/pipeline/__init__.py +++ b/junifer/pipeline/__init__.py @@ -5,3 +5,4 @@ from . import registry from .pipeline_step_mixin import PipelineStepMixin +from .update_meta_mixin import UpdateMetaMixin diff --git a/junifer/pipeline/pipeline_step_mixin.py b/junifer/pipeline/pipeline_step_mixin.py index 915458810..acd9ca900 100644 --- a/junifer/pipeline/pipeline_step_mixin.py +++ b/junifer/pipeline/pipeline_step_mixin.py @@ -4,6 +4,7 @@ # Synchon Mandal # License: AGPL +from importlib.util import find_spec from typing import Dict, List from ..utils import raise_error @@ -12,22 +13,6 @@ from ..utils import raise_error class PipelineStepMixin: """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: """Validate the input to the pipeline step. @@ -48,24 +33,22 @@ class PipelineStepMixin: klass=NotImplementedError, ) - def get_output_kind(self, input: List[str]) -> List[str]: - """Get the kind of the pipeline step. + def get_output_type(self, input_type: str) -> str: + """Get output type. Parameters ---------- - input : list of str - The input to the pipeline step. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the marker. Returns ------- - list of str - The updated list of available Junifer Data dictionary keys after - the pipeline step. + str + The storage type output by the marker. """ raise_error( - msg="Concrete classes need to implement get_output_kind().", + msg="Concrete classes need to implement get_output_type().", klass=NotImplementedError, ) @@ -85,11 +68,30 @@ class PipelineStepMixin: Raises ------ 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) - 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]: """Fit and transform. diff --git a/junifer/pipeline/registry.py b/junifer/pipeline/registry.py index b38c46ab2..07cb40dc9 100644 --- a/junifer/pipeline/registry.py +++ b/junifer/pipeline/registry.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Union from ..utils.logging import logger, raise_error + if TYPE_CHECKING: from ..datagrabber import BaseDataGrabber from ..storage import BaseFeatureStorage diff --git a/junifer/pipeline/tests/test_pipeline_step_mixin.py b/junifer/pipeline/tests/test_pipeline_step_mixin.py index 160d57201..457bc8891 100644 --- a/junifer/pipeline/tests/test_pipeline_step_mixin.py +++ b/junifer/pipeline/tests/test_pipeline_step_mixin.py @@ -4,6 +4,8 @@ # Synchon Mandal # License: AGPL +from typing import Dict, List + import pytest from junifer.pipeline.pipeline_step_mixin import PipelineStepMixin @@ -15,13 +17,49 @@ def test_PipelineStepMixin() -> None: with pytest.raises(NotImplementedError): mixin.validate_input([]) with pytest.raises(NotImplementedError): - mixin.get_output_kind([]) + mixin.get_output_type("") with pytest.raises(NotImplementedError): mixin.fit_transform({}) -def test_pipeline_step_mixin_meta(): - """Test metadata for PipelineStepMixin.""" - pipemixin = PipelineStepMixin() - t_meta = pipemixin.get_meta() - assert t_meta["class"] == "PipelineStepMixin" +def test_pipeline_step_mixin_validate_correct_dependencies() -> None: + """Test validate with correct dependencies.""" + + class CorrectMixer(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([]) diff --git a/junifer/pipeline/tests/test_update_meta_mixin.py b/junifer/pipeline/tests/test_update_meta_mixin.py new file mode 100644 index 000000000..5b857506f --- /dev/null +++ b/junifer/pipeline/tests/test_update_meta_mixin.py @@ -0,0 +1,51 @@ +"""Provide tests for update meta mixin.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# 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 diff --git a/junifer/pipeline/update_meta_mixin.py b/junifer/pipeline/update_meta_mixin.py new file mode 100644 index 000000000..7b3b4a460 --- /dev/null +++ b/junifer/pipeline/update_meta_mixin.py @@ -0,0 +1,43 @@ +"""Provide mixin class for updating metadata.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# 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) diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index f1e181035..f0ad1ad94 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -7,11 +7,11 @@ from abc import ABC, abstractmethod from typing import Any, Dict, List, Optional, Tuple, Union -from ..pipeline import PipelineStepMixin +from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..utils import logger, raise_error -class BasePreprocessor(ABC, PipelineStepMixin): +class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): """Provide abstract base class for all preprocessors. Parameters @@ -58,8 +58,8 @@ class BasePreprocessor(ABC, PipelineStepMixin): ) @abstractmethod - def get_output_kind(self, input: List[str]) -> List[str]: - """Get output kind. + def get_output_type(self, input: List[str]) -> List[str]: + """Get output type. Parameters ---------- @@ -75,7 +75,7 @@ class BasePreprocessor(ABC, PipelineStepMixin): """ raise_error( - msg="Concrete classes need to implement get_output_kind().", + msg="Concrete classes need to implement get_output_type().", klass=NotImplementedError, ) @@ -93,22 +93,6 @@ class BasePreprocessor(ABC, PipelineStepMixin): 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( self, input: Dict[str, Dict], @@ -127,19 +111,17 @@ class BasePreprocessor(ABC, PipelineStepMixin): """ out = input - for kind in self._on: - if kind in input.keys(): - logger.info(f"Computing {kind}") - t_input = input[kind] + for type_ in self._on: + if type_ in input.keys(): + logger.info(f"Computing {type_}") + t_input = input[type_] extra_input = input.copy() - extra_input.pop(kind) - t_meta = t_input.get("meta", {}) # input kind meta - t_meta.update(self.get_meta(kind)) + extra_input.pop(type_) key, t_out = self.preprocess( input=t_input, extra_input=extra_input ) - t_out.update(meta=t_meta) out[key] = t_out + self.update_meta(out[key], "preprocess") return out @abstractmethod diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index 2e5be6807..8018c0448 100644 --- a/junifer/preprocess/confounds/fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/fmriprep_confound_remover.py @@ -17,6 +17,7 @@ from ...api.decorators import register_preprocessor from ...utils import logger, raise_error from ..base import BasePreprocessor + if TYPE_CHECKING: from nibabel import MGHImage, Nifti1Image, Nifti2Image @@ -141,6 +142,8 @@ class fMRIPrepConfoundRemover(BasePreprocessor): """ + _DEPENDENCIES = {"numpy", "nilearn"} + def __init__( self, strategy: Optional[Dict[str, str]] = None, @@ -222,7 +225,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor): 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. Parameters diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index 98e028d86..5d52a24e2 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -62,7 +62,7 @@ def test_fMRIPrepConfoundRemover_validate_input() -> None: confound_remover.validate_input(input) -def test_fMRIPrepConfoundRemover_get_output_kind() -> None: +def test_fMRIPrepConfoundRemover_get_output_type() -> None: """Test fMRIPrepConfoundRemover validate_input.""" confound_remover = fMRIPrepConfoundRemover() inputs = [ @@ -72,7 +72,7 @@ def test_fMRIPrepConfoundRemover_get_output_kind() -> None: ] # Confound remover works in place 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: @@ -460,8 +460,8 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None: clean_bold = typing.cast(nib.Nifti1Image, clean_bold) # TODO: Find a better way to test functionality here assert ( - clean_bold.header.get_zooms() == # type: ignore - raw_bold.header.get_zooms() + clean_bold.header.get_zooms() # type: ignore + == raw_bold.header.get_zooms() # type: ignore ) assert clean_bold.get_fdata().shape == raw_bold.get_fdata().shape # 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 ) - assert output["meta"] == input["meta"] # general meta does not change assert "meta" in output["BOLD"] assert "preprocess" in output["BOLD"]["meta"] 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["t_r"] is None assert t_meta["mask_img"] is None + + assert "dependencies" in output["BOLD"]["meta"] + dependencies = output["BOLD"]["meta"]["dependencies"] + assert dependencies == {"numpy", "nilearn"} diff --git a/junifer/preprocess/tests/test_preprocess_base.py b/junifer/preprocess/tests/test_preprocess_base.py index 90d3c89ef..05a1a3f80 100644 --- a/junifer/preprocess/tests/test_preprocess_base.py +++ b/junifer/preprocess/tests/test_preprocess_base.py @@ -19,10 +19,14 @@ def test_base_preprocessor_subclassing() -> None: """Test proper subclassing of BasePreprocessor.""" # Create concrete class class MyBasePreprocessor(BasePreprocessor): + def __init__(self, on): + self.parameter = 1 + super().__init__(on=on) + def get_valid_inputs(self): return ["BOLD", "T1w"] - def get_output_kind(self, input): + def get_output_type(self, input): return ["timeseries"] def preprocess(self, input, extra_input): @@ -37,14 +41,23 @@ def test_base_preprocessor_subclassing() -> None: # Create input for marker input_ = { - "meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"}, "BOLD": { "path": ".", "data": "data", + "meta": { + "datagrabber": "dg", + "element": "elem", + "datareader": "dr", + }, }, "T1w": { "path": ".", "data": "data", + "meta": { + "datagrabber": "dg", + "element": "elem", + "datareader": "dr", + }, }, } prep = MyBasePreprocessor(on=["BOLD"]) @@ -60,6 +73,13 @@ def test_base_preprocessor_subclassing() -> None: assert "path" 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 "data" in output["T1w"] assert output["T1w"]["data"] == "data" diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 624a1bab7..38b2f5a5b 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -6,12 +6,13 @@ from abc import ABC, abstractmethod 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 -from .._version import __version__ from ..utils import raise_error +from .utils import process_meta class BaseFeatureStorage(ABC): @@ -40,23 +41,29 @@ class BaseFeatureStorage(ABC): self.uri = uri if not isinstance(storage_types, list): 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.single_output = single_output - def get_meta(self) -> Dict: - """Get metadata. + def get_valid_inputs(self) -> List[str]: + """Get valid storage types for input. Returns ------- - dict - The metadata as a dictionary. - + list of str + The list of storage types that can be used as input for this " + "storage. """ - meta = {} - meta["versions"] = { - "junifer": __version__, - } - return meta + raise_error( + msg="Concrete classes need to implement get_valid_inputs().", + klass=NotImplementedError, + ) def validate(self, input_: List[str]) -> None: """Validate the input to the pipeline step. @@ -80,23 +87,15 @@ class BaseFeatureStorage(ABC): ) @abstractmethod - def list_features( - self, return_df: bool = False - ) -> Union[Dict[str, Dict], pd.DataFrame]: + def list_features(self) -> Dict: """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 ------- - dict or pandas.DataFrame - List of features in the storage. If dictionary is returned, the - keys are the feature names to be used in read_features() and the - values are the metadata of each feature. + dict + List of features in the storage. The keys are the feature names to + be used in read_features() and the values are the metadata of each + feature. """ raise_error( @@ -131,19 +130,17 @@ class BaseFeatureStorage(ABC): ) @abstractmethod - def store_metadata(self, meta: Dict) -> str: + def store_metadata(self, meta_md5: str, element: Dict, meta: Dict) -> None: """Store metadata. Parameters ---------- + meta_md5 : str + The metadata MD5 hash. + element : dict + The element as a dictionary. meta : dict The metadata as a dictionary. - - Returns - ------- - str - The metadata column. - """ raise_error( msg="Concrete classes need to implement store_metadata().", @@ -166,65 +163,117 @@ class BaseFeatureStorage(ABC): 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": - self.store_matrix(**kwargs) + self.store_matrix(meta_md5=meta_md5, element=t_element, **kwargs) elif kind == "timeseries": - self.store_timeseries(**kwargs) + self.store_timeseries( + meta_md5=meta_md5, element=t_element, **kwargs + ) elif kind == "table": - self.store_table(**kwargs) - else: - raise ValueError(f"I don't know how to store {kind}") + self.store_table(meta_md5=meta_md5, element=t_element, **kwargs) - def store_df(self, **kwargs) -> None: - """Store pandas DataFrame. - - Parameters - ---------- - **kwargs : dict - The keyword arguments. - - """ - raise_error( - msg="Concrete classes need to implement store_df().", - klass=NotImplementedError, - ) - - def store_matrix(self, **kwargs) -> None: + def store_matrix( + self, + meta_md5: str, + element: Dict, + data: np.ndarray, + col_names: Optional[Iterable[str]] = None, + row_names: Optional[Iterable[str]] = None, + matrix_kind: Optional[str] = "full", + diagonal: bool = True, + ) -> None: """Store matrix. Parameters ---------- - **kwargs : dict - The keyword arguments. + meta_md5 : str + 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( msg="Concrete classes need to implement store_matrix2d().", 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. Parameters ---------- - **kwargs : dict - The keyword arguments. - + 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). """ raise_error( msg="Concrete classes need to implement store_table().", klass=NotImplementedError, ) - def store_timeseries(self, **kwargs) -> None: - """Store timeseries. + def store_timeseries( + self, + meta_md5: str, + element: Dict, + data: np.ndarray, + columns: Optional[Iterable[str]] = None, + ) -> None: + """Implement timeseries storing. Parameters ---------- - **kwargs : dict - The keyword arguments. - + 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). """ raise_error( msg="Concrete classes need to implement store_timeseries().", diff --git a/junifer/storage/pandas_base.py b/junifer/storage/pandas_base.py index 4ee40b712..bc6b5983c 100644 --- a/junifer/storage/pandas_base.py +++ b/junifer/storage/pandas_base.py @@ -6,8 +6,9 @@ import json 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 from .base import BaseFeatureStorage @@ -39,6 +40,17 @@ class PandasBaseFeatureStorage(BaseFeatureStorage): ) -> None: 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: """Convert the metadata to a pandas DataFrame. @@ -57,8 +69,166 @@ class PandasBaseFeatureStorage(BaseFeatureStorage): data_df = {} for k, v in meta.items(): 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.index.name = "meta_md5" 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", + ) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 8256db1cf..3d9ed05c6 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -5,8 +5,9 @@ # License: AGPL 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 pandas as pd from pandas.core.base import NoNewAttributesMixin @@ -17,7 +18,8 @@ from tqdm import tqdm from ..api.decorators import register_storage from ..utils import logger, raise_error, warn_with_log 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: from sqlalchemy.engine import Engine @@ -86,7 +88,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): # Set upsert self._upsert = upsert - def get_engine(self, meta: Optional[Dict] = None) -> "Engine": + def get_engine(self, element: Optional[Dict] = None) -> "Engine": """Get engine. Parameters @@ -100,11 +102,6 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): 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 prefix = "" if self.single_output is False: @@ -208,58 +205,15 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): msg=f"Invalid option {if_exists} for if_exists." ) - def _store_2d( - self, - 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). + def list_features(self) -> Dict: + """List the features in the storage. Returns ------- - dict or pandas.DataFrame - List of features in the storage. If dictionary is returned, the - keys are the feature names to be used in read_features() and the - values are the metadata of each feature. + dict + List of features in the storage. The keys are the feature names to + be used in read_features() and the values are the metadata of each + feature. """ meta_df = pd.read_sql( @@ -267,10 +221,11 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): con=self.get_engine(), index_col="meta_md5", ) - out = meta_df - # Return dictionary - if return_df is False: - out = meta_df.to_dict(orient="index") # type: ignore + meta_df.index = meta_df.index.str.replace(r"meta_", "") + out = meta_df.to_dict(orient="index") # type: ignore + for md5, t_meta in out.items(): + for k, v in t_meta.items(): + out[md5][k] = json.loads(v) return out def read_df( @@ -327,7 +282,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): con=engine, 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: raise_error(msg=f"Feature {feature_name} not found") elif len(t_df) > 1: @@ -339,8 +296,11 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): ) ) 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 df = pd.read_sql(sql=table_name, con=engine) + # Read the index query = ( "SELECT ii.name FROM sqlite_master AS m, " @@ -356,36 +316,30 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): df = df.set_index(index_names) return df - def store_metadata(self, meta: Dict) -> str: - r"""Implement metadata storing in the storage. + def store_metadata(self, meta_md5: str, element: Dict, meta: Dict) -> None: + """Implement metadata storing in the storage. Parameters ---------- + meta_md5 : str + The metadata MD5 hash. + element : dict + The element as a dictionary. meta : dict 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 - engine = self.get_engine(meta=t_meta) - if meta_md5 not in inspect(engine).get_table_names(): + engine = self.get_engine(element=element) + table_name = f"meta_{meta_md5}" + if table_name not in inspect(engine).get_table_names(): # 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 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. Parameters @@ -405,7 +359,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): # TODO: Test this function # Check that the index generated by meta matches the one in # 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 # when storing a timeseries or 2d elements. We need to check if the # extra element is only one. @@ -418,24 +372,25 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): elif len(extra) == 1: # The df has one extra index item, this should be the new name # 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): raise_error( "The index of the dataframe is missing index items that are " "generated from the metadata." ) - # Get table name - table_name = self.store_metadata(meta) + + table_name = f"meta_{meta_md5}" # Get sqlalchemy engine - engine = self.get_engine(meta) + engine = self.get_engine(element) # Save data self._save_upsert(df, table_name, engine) def store_matrix( self, + meta_md5: str, + element: Dict, data: np.ndarray, - meta: Dict, col_names: Optional[List[str]] = None, row_names: Optional[List[str]] = None, matrix_kind: Optional[str] = "full", @@ -445,6 +400,10 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): Parameters ---------- + meta_md5 : str + The metadata MD5 hash. + element : dict + The element as a dictionary. data : numpy.ndarray The matrix data to store. meta : dict @@ -465,7 +424,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): (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).. + this to False will raise an error (default True). """ if diagonal is False and matrix_kind not in ["triu", "tril"]: @@ -519,74 +478,27 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): # Convert element metadata to index 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 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() new_names = [x for x in data_df.index.names[:-1]] new_names.append("pair") data_df.index.names = new_names # Store dataframe - self.store_df(df=data_df, meta=meta) - - 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 - ) + self.store_df(meta_md5=meta_md5, element=element, df=data_df) def collect(self) -> None: """Implement data collection. diff --git a/junifer/storage/tests/test_pandas_base.py b/junifer/storage/tests/test_pandas_base.py new file mode 100644 index 000000000..f81dd4247 --- /dev/null +++ b/junifer/storage/tests/test_pandas_base.py @@ -0,0 +1,70 @@ +"""Provide tests for pandas base feature storage.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# 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,) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 8124aa987..65b3e6ffc 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -15,47 +15,44 @@ from pandas.testing import assert_frame_equal from sqlalchemy import create_engine from junifer.storage.sqlite import SQLiteFeatureStorage -from junifer.storage.utils import ( - element_to_index, - element_to_prefix, - process_meta, -) +from junifer.storage.utils import element_to_prefix, process_meta + df1 = pd.DataFrame( { - "element": [1, 2, 3, 4, 5], + "subject": [1, 2, 3, 4, 5], "pk2": ["a", "b", "c", "d", "e"], "col1": [11, 22, 33, 44, 55], "col2": [111, 222, 333, 444, 555], } -).set_index(["element", "pk2"]) +).set_index(["subject", "pk2"]) df2 = pd.DataFrame( { - "element": [2, 5, 6], + "subject": [2, 5, 6], "pk2": ["b", "e", "f"], "col1": [2222, 5555, 66], "col2": [22222, 55555, 666], } -).set_index(["element", "pk2"]) +).set_index(["subject", "pk2"]) 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"], "col1": [11, 2222, 33, 44, 5555, 66], "col2": [111, 22222, 333, 444, 55555, 666], } -).set_index(["element", "pk2"]) +).set_index(["subject", "pk2"]) 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"], "col1": [11, 22, 33, 44, 55, 66], "col2": [111, 222, 333, 444, 555, 666], } -).set_index(["element", "pk2"]) +).set_index(["subject", "pk2"]) def _read_sql( @@ -148,15 +145,14 @@ def test_upsert_replace(tmp_path: Path) -> None: uri = tmp_path / "test_upsert_replace.sqlite" # Single storage, must be the uri storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") - # Metadata to store - meta = {"element": "test", "version": "0.0.1"} # Save to database - storage.store_df(df=df1, meta=meta) - # Store metadata - table_name = storage.store_metadata(meta) + storage.store_df( + meta_md5="table_name", element={"subject": "test"}, df=df1 + ) + table_name = "meta_table_name" # Read stored table 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 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") # Read stored table 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 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" # Single storage, must be the uri storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") - # Metadata to store - meta = {"element": "test", "version": "0.0.1"} # Save to database - storage.store_df(df=df1, meta=meta) - # Store metadata - table_name = storage.store_metadata(meta=meta) + storage.store_df( + meta_md5="table_name", element={"subject": "test"}, df=df1 + ) + table_name = "meta_table_name" # Read stored table 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 assert_frame_equal(df1, c_df1) # Check for warning 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 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 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" storage = SQLiteFeatureStorage(uri=uri) - # Metadata to store - meta = {"element": "test", "version": "0.0.1"} # Save to database - storage.store_df(df1, meta) - # Store metadata - table_name = storage.store_metadata(meta) + storage.store_df( + meta_md5="table_name", element={"subject": "test"}, df=df1 + ) + table_name = "meta_table_name" # Read stored table 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 assert_frame_equal(df1, c_df1) # Save to database - storage.store_df(df2, meta) + storage.store_df( + meta_md5="table_name", element={"subject": "test"}, df=df2 + ) # Read stored table 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 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" storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") # Metadata to store + meta_md5 = "feature_md5" + element = {"subject": "test"} meta = { - "element": "test", - "version": "0.0.1", - "marker": {"name": "fcname"}, + "element": element, + "dependencies": ["numpy"], + "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 to_store = df1[["col1", "col2"]] # Check for error while storing 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 - 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 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 - idx = element_to_index(meta, n_rows=len(to_store)) + idx = storage.element_to_index(element=element, n_rows=len(to_store)) # Set index to_store = to_store.set_index(idx) # Store dataframe - storage.store_df(to_store, meta) - # Store metadata - table_name = storage.store_metadata(meta) + storage.store_df(meta_md5=meta_md5, element=element, df=to_store) + # List stored features features = storage.list_features() # Check correct usage assert len(features) == 1 - assert table_name.replace("meta_", "") in features + assert "feature_md5" in features # Check for missing feature 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 with pytest.raises(ValueError, match="least one"): storage.read_df() # Check for multiple features to fetch 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 feature_md5 = list(features.keys())[0] + assert "feature_md5" == feature_md5 # Check for key - assert "fcname" == features[feature_md5]["name"] + assert "BOLD_markername" == features[feature_md5]["name"] # Read into dataframes 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 assert_frame_equal(read_df1, read_df2) 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 storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") # 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 - table_name = storage.store_metadata(meta) - assert table_name.startswith("meta_") + storage.store_metadata( + 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: @@ -348,7 +373,21 @@ def test_store_table(tmp_path: Path) -> None: uri = tmp_path / "test_store_table.sqlite" storage = SQLiteFeatureStorage(uri=uri) # 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 = [ [1, 10], @@ -358,18 +397,23 @@ def test_store_table(tmp_path: Path) -> None: [5, 50], ] # 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 df = pd.DataFrame(data, columns=["f1", "f2"], index=idx) + # Store table - storage.store_table(data, meta, columns=["f1", "f2"], rows_col_name="scan") - # Store metadata - table_name = storage.store_metadata(meta) + storage.store_table( + meta_md5=meta_md5, + element=element_to_store, + data=data, + columns=["f1", "f2"], + rows_col_name="scan", + ) # Read stored table c_df = _read_sql( - table_name=table_name, + table_name=f"meta_{meta_md5}", uri=uri.as_posix(), - index_col=["element", "scan"], + index_col=["subject", "scan"], ) # Check if dataframes are equal assert_frame_equal(df, c_df) @@ -377,19 +421,24 @@ def test_store_table(tmp_path: Path) -> None: # New data to store data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]] # 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 df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new) # Check warning with pytest.warns(RuntimeWarning, match=r"Some rows"): + # 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 c_df_new = _read_sql( - table_name=table_name, + table_name=f"meta_{meta_md5}", uri=uri.as_posix(), - index_col=["element", "scan"], + index_col=["subject", "scan"], ) # Check if dataframes are equal 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" storage = SQLiteFeatureStorage(uri=uri) # 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 data = np.array( @@ -418,17 +480,17 @@ def test_store_matrix(tmp_path: Path) -> None: # Store table storage.store_matrix( + meta_md5=meta_md5, + element=element_to_store, data=data, - meta=meta, row_names=row_names, col_names=col_names, ) - stored_names = [f"{i}~{j}" for i in row_names for j in col_names] features = storage.list_features() 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) assert read_df.shape == (1, 12) @@ -437,7 +499,14 @@ def test_store_matrix(tmp_path: Path) -> None: # Store without row and column names uri = tmp_path / "test_store_table_nonames.sqlite" 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 = [ f"r{i}~c{j}" for i in range(data.shape[0]) @@ -445,20 +514,37 @@ def test_store_matrix(tmp_path: Path) -> None: ] features = storage.list_features() 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) assert list(read_df.columns) == stored_names 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"): - 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"): storage.store_matrix( + meta_md5=meta_md5, + element=element_to_store, data=data, - meta=meta, + row_names=row_names, + col_names=col_names, matrix_kind="full", diagonal=False, ) @@ -469,12 +555,17 @@ def test_store_matrix(tmp_path: Path) -> None: col_names = ["col1", "col2", "col3"] uri = tmp_path / "test_store_table_triu.sqlite" storage = SQLiteFeatureStorage(uri=uri) + # Store metadata + storage.store_metadata( + meta_md5=meta_md5, element=element_to_store, meta=meta_to_store + ) storage.store_matrix( + meta_md5=meta_md5, + element=element_to_store, data=data, - meta=meta, - matrix_kind="triu", row_names=row_names, col_names=col_names, + matrix_kind="triu", ) stored_names = [ @@ -488,7 +579,7 @@ def test_store_matrix(tmp_path: Path) -> None: features = storage.list_features() 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) assert list(read_df.columns) == stored_names assert_array_equal( @@ -498,12 +589,17 @@ def test_store_matrix(tmp_path: Path) -> None: # Store upper triangular matrix without diagonal uri = tmp_path / "test_store_table_triu_nodiagonal.sqlite" storage = SQLiteFeatureStorage(uri=uri) + # Store metadata + storage.store_metadata( + meta_md5=meta_md5, element=element_to_store, meta=meta_to_store + ) storage.store_matrix( + meta_md5=meta_md5, + element=element_to_store, data=data, - meta=meta, - matrix_kind="triu", row_names=row_names, col_names=col_names, + matrix_kind="triu", diagonal=False, ) @@ -515,7 +611,7 @@ def test_store_matrix(tmp_path: Path) -> None: features = storage.list_features() 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) assert list(read_df.columns) == stored_names assert_array_equal( @@ -528,12 +624,17 @@ def test_store_matrix(tmp_path: Path) -> None: col_names = ["col1", "col2", "col3"] uri = tmp_path / "test_store_table_tril.sqlite" storage = SQLiteFeatureStorage(uri=uri) + # Store metadata + storage.store_metadata( + meta_md5=meta_md5, element=element_to_store, meta=meta_to_store + ) storage.store_matrix( + meta_md5=meta_md5, + element=element_to_store, data=data, - meta=meta, - matrix_kind="tril", row_names=row_names, col_names=col_names, + matrix_kind="tril", ) stored_names = [ @@ -547,7 +648,7 @@ def test_store_matrix(tmp_path: Path) -> None: features = storage.list_features() 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) assert list(read_df.columns) == stored_names assert_array_equal( @@ -557,15 +658,19 @@ def test_store_matrix(tmp_path: Path) -> None: # Store lower triangular matrix without diagonal uri = tmp_path / "test_store_table_tril_nodiagonal.sqlite" storage = SQLiteFeatureStorage(uri=uri) + # Store metadata + storage.store_metadata( + meta_md5=meta_md5, element=element_to_store, meta=meta_to_store + ) storage.store_matrix( - data, - meta, - matrix_kind="tril", + meta_md5=meta_md5, + element=element_to_store, + data=data, row_names=row_names, col_names=col_names, + matrix_kind="tril", diagonal=False, ) - stored_names = [ "row2~col1", "row3~col1", @@ -574,7 +679,7 @@ def test_store_matrix(tmp_path: Path) -> None: features = storage.list_features() 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) assert list(read_df.columns) == stored_names assert_array_equal( @@ -597,18 +702,21 @@ def test_store_multiple_output(tmp_path: Path): # Metadata to store meta1 = { "element": {"subject": "test-01", "session": "ses-01"}, - "version": "0.0.1", + "dependencies": ["numpy"], "marker": {"name": "fc"}, + "type": "BOLD", } meta2 = { "element": {"subject": "test-02", "session": "ses-01"}, - "version": "0.0.1", + "dependencies": ["numpy"], "marker": {"name": "fc"}, + "type": "BOLD", } meta3 = { "element": {"subject": "test-01", "session": "ses-02"}, - "version": "0.0.1", + "dependencies": ["numpy"], "marker": {"name": "fc"}, + "type": "BOLD", } # Data to store data1 = np.array( @@ -623,35 +731,62 @@ def test_store_multiple_output(tmp_path: Path): data2 = data1 * 10 data3 = data1 * 20 # Process metadata for storage - hash1, _ = process_meta(meta1) + hash1, meta_to_store1, element_to_store1 = process_meta(meta1) # 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 df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1) # Process metadata for storage - hash2, _ = process_meta(meta2) + hash2, meta_to_store2, element_to_store2 = process_meta(meta2) # 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 df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2) # Process metadata for storage - hash3, _ = process_meta(meta3) + hash3, meta_to_store3, element_to_store3 = process_meta(meta3) # 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 df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3) # Check hash equality assert hash1 == hash2 assert hash2 == hash3 # Store tables - storage.store_table( - data1, meta1, columns=["f1", "f2"], rows_col_name="scan" + 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( - 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( - 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 assert not uri.exists() @@ -667,10 +802,9 @@ def test_store_multiple_output(tmp_path: Path): assert uri1.exists() assert uri2.exists() assert uri3.exists() - # Store metadata - table_name = storage.store_metadata(meta1) # Set index columns cols = ["subject", "session", "scan"] + table_name = f"meta_{hash1}" # Read stored tables cdf1 = _read_sql(table_name, uri1.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 meta1 = { "element": {"subject": "test-01", "session": "ses-01"}, - "version": "0.0.1", + "dependencies": ["numpy"], "marker": {"name": "fc"}, + "type": "BOLD", } meta2 = { "element": {"subject": "test-02", "session": "ses-01"}, - "version": "0.0.1", + "dependencies": ["numpy"], "marker": {"name": "fc"}, + "type": "BOLD", } meta3 = { "element": {"subject": "test-01", "session": "ses-02"}, - "version": "0.0.1", + "dependencies": ["numpy"], "marker": {"name": "fc"}, + "type": "BOLD", } # Data for storage data1 = np.array( @@ -722,14 +859,38 @@ def test_collect(tmp_path: Path) -> None: data2 = data1 * 10 data3 = data1 * 20 # Store tables - storage.store_table( - data1, meta1, columns=["f1", "f2"], rows_col_name="scan" + hash1, meta_to_store1, element_to_store1 = process_meta(meta1) + 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( - 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( - 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 prefix1 = element_to_prefix(meta1["element"]) @@ -752,7 +913,7 @@ def test_collect(tmp_path: Path) -> None: # Set index columns cols = ["subject", "session", "scan"] # Store metadata - table_name = storage.store_metadata(meta1) + table_name = f"meta_{hash1}" # Read stored tables all_df = _read_sql(table_name, uri.as_posix(), index_col=cols) cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) diff --git a/junifer/storage/tests/test_storage_base.py b/junifer/storage/tests/test_storage_base.py index cfe9a0068..d59b27086 100644 --- a/junifer/storage/tests/test_storage_base.py +++ b/junifer/storage/tests/test_storage_base.py @@ -12,7 +12,7 @@ from junifer.storage.base import BaseFeatureStorage def test_BaseFeatureStorage_abstractness() -> None: """Test BaseFeatureStorage is abstract base class.""" 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: @@ -22,13 +22,16 @@ def test_BaseFeatureStorage() -> None: """Implement concrete class.""" def __init__(self, uri, single_output=False): - storage_types = ["matrix"] + storage_types = ["matrix", "table", "timeseries"] super().__init__( uri=uri, storage_types=storage_types, single_output=single_output, ) + def get_valid_inputs(self): + return ["matrix", "table", "timeseries"] + def list_features(self): super().list_features() @@ -38,8 +41,8 @@ def test_BaseFeatureStorage() -> None: feature_md5=feature_md5, ) - def store_metadata(self, metadata): - super().store_metadata(metadata) + def store_metadata(self, meta_md5, meta, element): + super().store_metadata(meta_md5, meta, element) def collect(self): return super().collect() @@ -55,7 +58,7 @@ def test_BaseFeatureStorage() -> None: st.validate(input_=["matrix"]) # Check validate with invalid argument with pytest.raises(ValueError): - st.validate(input_=["table"]) + st.validate(input_=["duck"]) with pytest.raises(NotImplementedError): st.list_features() @@ -63,22 +66,31 @@ def test_BaseFeatureStorage() -> None: with pytest.raises(NotImplementedError): st.read_df(None) + element = {"subject": "test"} + dependencies = ["numpy"] + meta = { + "element": element, + "dependencies": dependencies, + "marker": {"name": "fc"}, + "type": "BOLD", + } + 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): st.collect() with pytest.raises(NotImplementedError): - st.store(kind="matrix") + st.store(kind="timeseries", meta=meta) with pytest.raises(NotImplementedError): - st.store(kind="timeseries") - - with pytest.raises(NotImplementedError): - st.store(kind="table") + st.store(kind="table", meta=meta) with pytest.raises(ValueError): - st.store(kind="lego") + st.store(kind="lego", meta=meta) assert st.uri == "/tmp" diff --git a/junifer/storage/tests/test_utils.py b/junifer/storage/tests/test_utils.py index 271397115..35a6f9c4b 100644 --- a/junifer/storage/tests/test_utils.py +++ b/junifer/storage/tests/test_utils.py @@ -4,17 +4,51 @@ # Synchon Mandal # License: AGPL -from typing import Dict, List, Tuple, Union +from typing import Dict, List import pytest from junifer.storage.utils import ( - element_to_index, element_to_prefix, + get_dependency_version, 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: """Test invalid metadata type check for metadata hash processing.""" meta = None @@ -25,58 +59,133 @@ def test_process_meta_invalid_metadata_type() -> None: # TODO: parameterize def test_process_meta_hash() -> None: """Test metadata hash processing.""" - meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]} - hash1, _ = process_meta(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} - hash2, _ = process_meta(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 element2 == {"foo": "baz"} - meta = {"element": "foo", "A": 1, "B": [2, 3, 1, 5, 6]} - hash3, _ = process_meta(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 element3 == element1 - meta1 = { - "element": "foo", + meta4 = { + "element": {"foo": "bar"}, "B": { "B2": [2, 3, 4, 5, 6], "B1": [9.22, 3.14, 1.41, 5.67, 6.28], "B3": (1, "car"), }, "A": 1, + "dependencies": ["numpy"], + "marker": {"name": "fc"}, + "type": "BOLD", } - meta2 = { + meta5 = { "A": 1, "B": { "B3": (1, "car"), "B1": [9.22, 3.14, 1.41, 5.67, 6.28], "B2": [2, 3, 4, 5, 6], }, - "element": "foo", + "element": {"foo": "baz"}, + "dependencies": ["numpy"], + "marker": {"name": "fc"}, + "type": "BOLD", } - hash4, _ = process_meta(meta1) - hash5, _ = process_meta(meta2) + hash4, _, _ = process_meta(meta4) + hash5, _, _ = process_meta(meta5) 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: """Test invalid metadata key check for metadata hash processing.""" 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) @pytest.mark.parametrize( "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"}, "B": [2, 3, 4, 5, 6], "A": 1, + "dependencies": ["numpy"], + "marker": {"name": "fc"}, + "type": "BOLD", }, ["subject", "session"], ), @@ -93,34 +202,30 @@ def test_process_meta_element(meta: Dict, elements: List[str]) -> None: 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 processed_meta["_element_keys"] == elements assert "A" in processed_meta assert "B" in processed_meta assert "element" not in processed_meta - hash2, processed_meta2 = process_meta(processed_meta) - - assert hash1, hash2 - assert processed_meta == processed_meta2 + assert isinstance(processed_meta["dependencies"], Dict) + assert all( + x in processed_meta["dependencies"] for x in meta["dependencies"] + ) + assert "name" in processed_meta + assert processed_meta["name"] == f"{meta['type']}_{meta['marker']['name']}" @pytest.mark.parametrize( "element,prefix", [ - ("sub-01", "element_sub-01_"), - (1, "element_1_"), ({"subject": "sub-01"}, "element_sub-01_"), ({"subject": 1}, "element_1_"), ({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"), ({"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( - element: Union[str, int, Dict, Tuple], prefix: str -) -> None: +def test_element_to_prefix(element: Dict, prefix: str) -> None: """Test converting element to prefix (for file naming). Parameters @@ -138,75 +243,5 @@ def test_element_to_prefix( def test_element_to_prefix_invalid_type() -> None: """Test element to prefix type checking.""" 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 - - -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,) diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py index 0a6fa58df..462d93bb6 100644 --- a/junifer/storage/utils.py +++ b/junifer/storage/utils.py @@ -6,14 +6,39 @@ import hashlib import json -from typing import Any, Dict, Optional, Tuple, Union - -import numpy as np -import pandas as pd +from importlib.metadata import PackageNotFoundError, version +from typing import Dict, Tuple 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: """Compute the MD5 hash of the metadata. @@ -29,6 +54,12 @@ def _meta_hash(meta: Dict) -> str: """ 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( json.dumps(meta, sort_keys=True).encode("utf-8") ).hexdigest() @@ -36,7 +67,7 @@ def _meta_hash(meta: Dict) -> str: 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. 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. dict The processed metadata for storage. + tuple + The element. Raises ------ ValueError - If `meta` is None or if it does not contain the key "element" or - "_element_keys". + If `meta` is None or if it does not contain the key "element". """ if meta is None: @@ -68,97 +100,42 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]: # Remove key "element" element = t_meta.pop("element", None) if element is None: - if "_element_keys" not in t_meta: - raise_error( - msg="`meta` must contain the key 'element' or '_element_keys'" - ) - else: - if isinstance(element, dict): - t_meta["_element_keys"] = list(element.keys()) - else: - t_meta["_element_keys"] = ["element"] + raise_error(msg="`meta` must contain the key 'element'") + if "marker" not in t_meta: + raise_error(msg="`meta` must contain the key 'marker'") + if "name" not in t_meta["marker"]: + raise_error(msg="`meta['marker']` must contain the key 'name'") + if "type" not in t_meta: + raise_error(msg="`meta` must contain the key 'type'") + + 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 = _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. Parameters ---------- - element : tuple, dict, str or int + element : dict The element to convert to prefix. Returns ------- str The element converted to prefix. - - Raises - ------ - ValueError - If invalid type is passed for `element`. - """ logger.debug(f"Converting element {element} to prefix.") prefix = "element" - if isinstance(element, tuple): - prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}" - elif isinstance(element, dict): - 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." - ) + if not isinstance(element, dict): + raise_error(msg="`element` must be a dict") + + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}" + logger.debug(f"Converted prefix: {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 diff --git a/junifer/testing/registry.py b/junifer/testing/registry.py index b721fe47b..014aa50d4 100644 --- a/junifer/testing/registry.py +++ b/junifer/testing/registry.py @@ -10,6 +10,7 @@ from .datagrabbers import ( SPMAuditoryTestingDatagrabber, ) + # Register testing datagrabber register( step="datagrabber", diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 1c3db746b..5e06dcd48 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -14,6 +14,7 @@ from warnings import warn import datalad + logger = logging.getLogger("JUNIFER") # Set up datalad logger level to warning by default @@ -262,7 +263,11 @@ def configure_logging( 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. Parameters @@ -271,10 +276,15 @@ def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: The message for the exception. klass : subclass of Exception, optional The subclass of Exception to raise using (default ValueError). + exception : Exception, optional + The original exception to follow up on (default None). """ logger.error(msg) - raise klass(msg) + if exception is not None: + raise klass(msg) from exception + else: + raise klass(msg) def warn_with_log( diff --git a/pyproject.toml b/pyproject.toml index 3ad16fd71..5e4edf8bc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ name = "junifer" description = "JUelich NeuroImaging FEature extractoR" readme = "README.md" requires-python = ">=3.8" -license = {file = "LICENSE.md"} +license = {text = "AGPL-3.0-only"} authors = [ {name = "Fede Raimondo", email = "f.raimondo@fz-juelich.de"}, {name = "Synchon Mandal", email = "s.mandal@fz-juelich.de"},