update: improve marker and storage interfaces #149
73 changed files with 1626 additions and 1337 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from click.testing import CliRunner
|
|||
|
||||
from junifer.api.cli import collect, run, selftest, wtf
|
||||
|
||||
|
||||
# Create click test runner
|
||||
runner = CliRunner()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@
|
|||
|
||||
from typing import List
|
||||
|
||||
import pytest
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from junifer.data.utils import closest_resolution
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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.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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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,11 +64,11 @@ 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()
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
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 {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)
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
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 {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)
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
|
|||
|
|
@ -5,3 +5,4 @@
|
|||
|
||||
from . import registry
|
||||
from .pipeline_step_mixin import PipelineStepMixin
|
||||
from .update_meta_mixin import UpdateMetaMixin
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
51
junifer/pipeline/tests/test_update_meta_mixin.py
Normal file
51
junifer/pipeline/tests/test_update_meta_mixin.py
Normal 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
|
||||
43
junifer/pipeline/update_meta_mixin.py
Normal file
43
junifer/pipeline/update_meta_mixin.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
"""
|
||||
if kind == "matrix":
|
||||
self.store_matrix(**kwargs)
|
||||
elif kind == "timeseries":
|
||||
self.store_timeseries(**kwargs)
|
||||
elif kind == "table":
|
||||
self.store_table(**kwargs)
|
||||
else:
|
||||
raise ValueError(f"I don't know how to store {kind}")
|
||||
|
||||
def store_df(self, **kwargs) -> None:
|
||||
"""Store pandas DataFrame.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
**kwargs : dict
|
||||
The keyword arguments.
|
||||
|
||||
"""
|
||||
# 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="Concrete classes need to implement store_df().",
|
||||
klass=NotImplementedError,
|
||||
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(meta_md5=meta_md5, element=t_element, **kwargs)
|
||||
elif kind == "timeseries":
|
||||
self.store_timeseries(
|
||||
meta_md5=meta_md5, element=t_element, **kwargs
|
||||
)
|
||||
elif kind == "table":
|
||||
self.store_table(meta_md5=meta_md5, element=t_element, **kwargs)
|
||||
|
||||
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().",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
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.
|
||||
|
|
|
|||
70
junifer/storage/tests/test_pandas_base.py
Normal file
70
junifer/storage/tests/test_pandas_base.py
Normal 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,)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
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())
|
||||
else:
|
||||
t_meta["_element_keys"] = ["element"]
|
||||
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):
|
||||
if not isinstance(element, dict):
|
||||
raise_error(msg="`element` must be a 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."
|
||||
)
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from .datagrabbers import (
|
|||
SPMAuditoryTestingDatagrabber,
|
||||
)
|
||||
|
||||
|
||||
# Register testing datagrabber
|
||||
register(
|
||||
step="datagrabber",
|
||||
|
|
|
|||
|
|
@ -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,9 +276,14 @@ 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)
|
||||
if exception is not None:
|
||||
raise klass(msg) from exception
|
||||
else:
|
||||
raise klass(msg)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
Loading…
Reference in a new issue