update: improve marker and storage interfaces #149

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

View file

@ -54,7 +54,7 @@ Contributions are welcome and greatly appreciated. Please read the [guidelines](
junifer is released under the AGPL v3 license:
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

View file

@ -15,12 +15,11 @@ Thus, only a few methods are required:
1. ``get_valid_inputs``: a method to obtain the list of valid inputs for the marker. This is used to check that the
inputs provided by the user are valid. This method should return a list of strings, representing
:ref:`data types <data_types>`
2. ``get_output_kind``: a method to obtain the kind of output of the marker. This is used to check that the output
2. ``get_output_type``: a method to obtain the kind of output of the marker. This is used to check that the output
of the marker is compatible with the storage. This method should return a string, representing
:ref:`storage types <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 <data_types>` 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 <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

View file

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

View file

@ -42,7 +42,10 @@ marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
###############################################################################
# Prepare the input
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

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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

View file

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

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

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

View file

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

View file

@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import nibabel as nib
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

View file

@ -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

View file

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

View file

@ -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,
)

View file

@ -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,

View file

@ -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

View file

@ -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"]

View file

@ -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"]

View file

@ -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"]

View file

@ -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.

View file

@ -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

View file

@ -5,7 +5,6 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -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()

View file

@ -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

View file

@ -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"]

View file

@ -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:

View file

@ -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"]

View file

@ -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

View file

@ -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

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -6,7 +6,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -4,19 +4,15 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# 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)

View file

@ -4,19 +4,16 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# 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)

View file

@ -4,21 +4,18 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -4,7 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -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:

View file

@ -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())

View file

@ -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()
)

View file

@ -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"]

View file

@ -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"

View file

@ -3,18 +3,21 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -3,6 +3,8 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# 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"] == {}

View file

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

View file

@ -4,6 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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.

View file

@ -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

View file

@ -4,6 +4,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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([])

View file

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

View file

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

View file

@ -7,11 +7,11 @@
from abc import ABC, abstractmethod
from 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

View file

@ -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

View file

@ -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"}

View file

@ -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"

View file

@ -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().",

View file

@ -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",
)

View file

@ -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.

View file

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

View file

@ -15,47 +15,44 @@ from pandas.testing import assert_frame_equal
from sqlalchemy import create_engine
from 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)

View file

@ -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"

View file

@ -4,17 +4,51 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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,)

View file

@ -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

View file

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

View file

@ -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(

View file

@ -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"},