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:
|
junifer is released under the AGPL v3 license:
|
||||||
|
|
||||||
julearn, FZJuelich AML neuroimaging feature extraction library.
|
junifer, FZJuelich AML neuroimaging feature extraction library.
|
||||||
Copyright (C) 2022, authors of junifer.
|
Copyright (C) 2022, authors of junifer.
|
||||||
|
|
||||||
This program is free software: you can redistribute it and/or modify
|
This program is free software: you can redistribute it and/or modify
|
||||||
|
|
|
||||||
|
|
@ -15,12 +15,11 @@ Thus, only a few methods are required:
|
||||||
1. ``get_valid_inputs``: a method to obtain the list of valid inputs for the marker. This is used to check that the
|
1. ``get_valid_inputs``: a method to obtain the list of valid inputs for the marker. This is used to check that the
|
||||||
inputs provided by the user are valid. This method should return a list of strings, representing
|
inputs provided by the user are valid. This method should return a list of strings, representing
|
||||||
:ref:`data types <data_types>`
|
:ref:`data types <data_types>`
|
||||||
2. ``get_output_kind``: a method to obtain the kind of output of the marker. This is used to check that the output
|
2. ``get_output_type``: a method to obtain the kind of output of the marker. This is used to check that the output
|
||||||
of the marker is compatible with the storage. This method should return a string, representing
|
of the marker is compatible with the storage. This method should return a string, representing
|
||||||
:ref:`storage types <storage_types>`
|
:ref:`storage types <storage_types>`
|
||||||
3. ``compute``: the method that given the data, computes the marker.
|
3. ``compute``: the method that given the data, computes the marker.
|
||||||
4. ``store``: the method that stores the computed marker.
|
4. ``__init__``: the initialization method, where the marker is configured.
|
||||||
5. ``__init__``: the initialization method, where the marker is configured.
|
|
||||||
|
|
||||||
As an example, we will develop a Parcel Mean marker, that is, a marker that first applies a parcellation and
|
As an example, we will develop a Parcel Mean marker, that is, a marker that first applies a parcellation and
|
||||||
then computes the mean of the data in each parcel. This is a very simple example, but it will show you how to create
|
then computes the mean of the data in each parcel. This is a very simple example, but it will show you how to create
|
||||||
|
|
@ -44,7 +43,7 @@ it will be ``table``. Thus, we can define the output as:
|
||||||
|
|
||||||
.. code-block:: python
|
.. code-block:: python
|
||||||
|
|
||||||
def get_output_kind(self, input_kind):
|
def get_output_type(self, input_kind):
|
||||||
if input_kind == 'BOLD':
|
if input_kind == 'BOLD':
|
||||||
return 'timeseries'
|
return 'timeseries'
|
||||||
else:
|
else:
|
||||||
|
|
@ -143,38 +142,24 @@ This dictionary will later be passed onto the ``store`` method.
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
.. _extending_markers_store:
|
.. _extending_markers_finalize:
|
||||||
|
|
||||||
Step 4: Store the marker
|
Step 4: Finalize the marker
|
||||||
------------------------
|
---------------------------
|
||||||
|
|
||||||
In this step, we will define the method that stores the marker. This method will be called by junifer when needed,
|
Once all of the above steps are done, we just need to give our marker a name, state its *dependencies* and register it
|
||||||
using the data provided by the ``compute`` method. The method ``store`` has three arguments:
|
using the ``@register_marker`` decorator.
|
||||||
|
|
||||||
* ``kind``: A string indicating the :ref:`data type <data_types>` that was used to compute the marker.
|
The *dependencies* are the core packages that are required to compute the marker. This will be later used to keep track
|
||||||
* ``out``: The output of the ``compute`` method.
|
of the versions of the packages used to compute the marker. To inform junifer about the dependencies of a marker,
|
||||||
* ``storage``: The storage object, that will be used to store the marker.
|
we need to define a ``_DEPENDENCIES`` attribute in the class. This attribute must be a set, with the names of the
|
||||||
|
packages as strings. For example, the ``ParcelMean`` marker has the following dependencies:
|
||||||
|
|
||||||
.. code-block:: python
|
.. code-block:: python
|
||||||
|
|
||||||
def store(self, kind, out, storage):
|
_DEPENDENCIES = {"nilearn"}
|
||||||
if kind in ["VBM_GM", "VBM_WM"]:
|
|
||||||
storage.store(kind="table", **out)
|
|
||||||
elif kind in ["BOLD"]:
|
|
||||||
storage.store(kind="timeseries", **out)
|
|
||||||
|
|
||||||
.. hint:: Check the hint on :ref:`extending_markers_compute`. If the output of the ``compute`` method is a dictionary
|
Finally, we need to register the marker using the ``@register_marker`` decorator. This decorator takes the name of the
|
||||||
with keys based on the :ref:`storage types <storage_types>`, the ``store`` method can simply call the right
|
|
||||||
storage function, based on the ``kind`` parameter, with ``**out``.
|
|
||||||
|
|
||||||
|
|
||||||
.. _extending_markers_finalize:
|
|
||||||
|
|
||||||
Step 5: Finalize the marker
|
|
||||||
---------------------------
|
|
||||||
|
|
||||||
Once all of the above steps are done, we just need to give our marker a name an register it using the
|
|
||||||
``@register_marker`` decorator:
|
|
||||||
|
|
||||||
.. code-block:: python
|
.. code-block:: python
|
||||||
|
|
||||||
|
|
@ -186,6 +171,8 @@ Once all of the above steps are done, we just need to give our marker a name an
|
||||||
@register_marker
|
@register_marker
|
||||||
class ParcelMean(BaseMarker):
|
class ParcelMean(BaseMarker):
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn", "numpy"}
|
||||||
|
|
||||||
def __init__(self, parcellation_name, on=None, name=None):
|
def __init__(self, parcellation_name, on=None, name=None):
|
||||||
self.parcellation_name = parcellation_name
|
self.parcellation_name = parcellation_name
|
||||||
super().__init__(on=on, name=name)
|
super().__init__(on=on, name=name)
|
||||||
|
|
@ -193,7 +180,7 @@ Once all of the above steps are done, we just need to give our marker a name an
|
||||||
def get_valid_inputs(self):
|
def get_valid_inputs(self):
|
||||||
return ['BOLD', 'VBM_WM', 'VBM_GM']
|
return ['BOLD', 'VBM_WM', 'VBM_GM']
|
||||||
|
|
||||||
def get_output_kind(self, input_kind):
|
def get_output_type(self, input_kind):
|
||||||
if input_kind == 'BOLD':
|
if input_kind == 'BOLD':
|
||||||
return 'timeseries'
|
return 'timeseries'
|
||||||
else:
|
else:
|
||||||
|
|
@ -231,12 +218,6 @@ Once all of the above steps are done, we just need to give our marker a name an
|
||||||
out["row_names"] = "scan"
|
out["row_names"] = "scan"
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def store(self, kind, out, storage):
|
|
||||||
if kind in ["VBM_GM", "VBM_WM"]:
|
|
||||||
storage.store(kind="table", **out)
|
|
||||||
elif kind in ["BOLD"]:
|
|
||||||
storage.store(kind="timeseries", **out)
|
|
||||||
|
|
||||||
|
|
||||||
.. _extending_markers_template:
|
.. _extending_markers_template:
|
||||||
|
|
||||||
|
|
@ -260,7 +241,7 @@ Template for a custom Marker
|
||||||
valid = []
|
valid = []
|
||||||
return valid
|
return valid
|
||||||
|
|
||||||
def get_output_kind(self, input_kind):
|
def get_output_type(self, input_kind):
|
||||||
# TODO: Return the valid output kind for each input kind
|
# TODO: Return the valid output kind for each input kind
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
@ -270,7 +251,3 @@ Template for a custom Marker
|
||||||
# Create the output dictionary
|
# Create the output dictionary
|
||||||
out = {"data": None, "columns": None}
|
out = {"data": None, "columns": None}
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def store(self, kind, out, storage):
|
|
||||||
# TODO: store out using the storage object, based on the kind of data
|
|
||||||
pass
|
|
||||||
|
|
@ -20,4 +20,4 @@ as the actual data is in the memory and the Python runtime has not garbage-colle
|
||||||
|
|
||||||
If you are interested in using already provided markers, please go to :doc:`../builtin`. And, if you want to implement
|
If you are interested in using already provided markers, please go to :doc:`../builtin`. And, if you want to implement
|
||||||
your own marker, you need to provide concrete implementation of :class:`junifer.markers.BaseMarker`. Specifically, you
|
your own marker, you need to provide concrete implementation of :class:`junifer.markers.BaseMarker`. Specifically, you
|
||||||
need to override ``get_output_kind``, ``store`` and ``compute`` methods.
|
need to override ``get_output_type``, ``store`` and ``compute`` methods.
|
||||||
|
|
|
||||||
|
|
@ -42,7 +42,10 @@ marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
|
||||||
|
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Prepare the input
|
# Prepare the input
|
||||||
input = {"BOLD": {"data": fmri_img}, "VBM_GM": {"data": vbm_img}}
|
input = {
|
||||||
|
"BOLD": {"data": fmri_img, "meta": {"element": "subject1"}},
|
||||||
|
"VBM_GM": {"data": vbm_img, "meta": {"element": "subject1"}},
|
||||||
|
}
|
||||||
|
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Fit transform the data
|
# Fit transform the data
|
||||||
|
|
|
||||||
|
|
@ -66,6 +66,9 @@ with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
collect(storage=storage)
|
collect(storage=storage)
|
||||||
# Create storage object to read in extracted features
|
# Create storage object to read in extracted features
|
||||||
db = SQLiteFeatureStorage(uri=storage["uri"])
|
db = SQLiteFeatureStorage(uri=storage["uri"])
|
||||||
|
|
||||||
|
# List all the features
|
||||||
|
print(db.list_features())
|
||||||
# Read extracted features
|
# Read extracted features
|
||||||
df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS")
|
df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -16,8 +16,8 @@ import yaml
|
||||||
from ..utils.logging import (
|
from ..utils.logging import (
|
||||||
configure_logging,
|
configure_logging,
|
||||||
logger,
|
logger,
|
||||||
warn_with_log,
|
|
||||||
raise_error,
|
raise_error,
|
||||||
|
warn_with_log,
|
||||||
)
|
)
|
||||||
from .functions import collect as api_collect
|
from .functions import collect as api_collect
|
||||||
from .functions import queue as api_queue
|
from .functions import queue as api_queue
|
||||||
|
|
@ -69,7 +69,8 @@ def _parse_elements(element: str, config: Dict) -> Union[List, None]:
|
||||||
raise_error(
|
raise_error(
|
||||||
"The 'elements' key is set in the configuration, but its value"
|
"The 'elements' key is set in the configuration, but its value"
|
||||||
" is 'None'. It is likely that there is an empty 'elements' "
|
" is 'None'. It is likely that there is an empty 'elements' "
|
||||||
"section in the yaml configuration file.")
|
"section in the yaml configuration file."
|
||||||
|
)
|
||||||
return elements
|
return elements
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -7,11 +7,11 @@
|
||||||
|
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
|
import textwrap
|
||||||
import typing
|
import typing
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Optional, Tuple, Union
|
from typing import Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import textwrap
|
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
from ..datagrabber.base import BaseDataGrabber
|
from ..datagrabber.base import BaseDataGrabber
|
||||||
|
|
|
||||||
|
|
@ -5,11 +5,11 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
import importlib
|
import importlib
|
||||||
from pathlib import Path
|
|
||||||
from typing import Dict, Union
|
|
||||||
import importlib.util
|
import importlib.util
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, Union
|
||||||
|
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
|
|
@ -45,7 +45,8 @@ def parse_yaml(filepath: Union[str, Path]) -> Dict:
|
||||||
if contents["elements"] is None:
|
if contents["elements"] is None:
|
||||||
raise_error(
|
raise_error(
|
||||||
"The elements key was defined but its content is empty. "
|
"The elements key was defined but its content is empty. "
|
||||||
"Please define the elements to operate on or remove the key.")
|
"Please define the elements to operate on or remove the key."
|
||||||
|
)
|
||||||
# load modules
|
# load modules
|
||||||
if "with" in contents:
|
if "with" in contents:
|
||||||
to_load = contents["with"]
|
to_load = contents["with"]
|
||||||
|
|
@ -58,9 +59,11 @@ def parse_yaml(filepath: Union[str, Path]) -> Dict:
|
||||||
file_path = Path(os.getcwd()) / t_module
|
file_path = Path(os.getcwd()) / t_module
|
||||||
if not file_path.exists():
|
if not file_path.exists():
|
||||||
raise_error(
|
raise_error(
|
||||||
f"File in 'with' section does not exist: {file_path}")
|
f"File in 'with' section does not exist: {file_path}"
|
||||||
|
)
|
||||||
spec = importlib.util.spec_from_file_location(
|
spec = importlib.util.spec_from_file_location(
|
||||||
t_module, file_path)
|
t_module, file_path
|
||||||
|
)
|
||||||
module = importlib.util.module_from_spec(spec) # type: ignore
|
module = importlib.util.module_from_spec(spec) # type: ignore
|
||||||
sys.modules[t_module] = module
|
sys.modules[t_module] = module
|
||||||
spec.loader.exec_module(module) # type: ignore
|
spec.loader.exec_module(module) # type: ignore
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ from click.testing import CliRunner
|
||||||
|
|
||||||
from junifer.api.cli import collect, run, selftest, wtf
|
from junifer.api.cli import collect, run, selftest, wtf
|
||||||
|
|
||||||
|
|
||||||
# Create click test runner
|
# Create click test runner
|
||||||
runner = CliRunner()
|
runner = CliRunner()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import pytest
|
||||||
from junifer.configs.juseless.datagrabbers import JuselessDataladAOMICID1000VBM
|
from junifer.configs.juseless.datagrabbers import JuselessDataladAOMICID1000VBM
|
||||||
from junifer.utils.logging import configure_logging
|
from junifer.utils.logging import configure_logging
|
||||||
|
|
||||||
|
|
||||||
# Check if the test is running on juseless
|
# Check if the test is running on juseless
|
||||||
if socket.gethostname() != "juseless":
|
if socket.gethostname() != "juseless":
|
||||||
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import pytest
|
||||||
from junifer.configs.juseless.datagrabbers import JuselessDataladCamCANVBM
|
from junifer.configs.juseless.datagrabbers import JuselessDataladCamCANVBM
|
||||||
from junifer.utils.logging import configure_logging
|
from junifer.utils.logging import configure_logging
|
||||||
|
|
||||||
|
|
||||||
# Check if the test is running on juseless
|
# Check if the test is running on juseless
|
||||||
if socket.gethostname() != "juseless":
|
if socket.gethostname() != "juseless":
|
||||||
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import pytest
|
||||||
from junifer.configs.juseless.datagrabbers import JuselessDataladIXIVBM
|
from junifer.configs.juseless.datagrabbers import JuselessDataladIXIVBM
|
||||||
from junifer.utils.logging import configure_logging
|
from junifer.utils.logging import configure_logging
|
||||||
|
|
||||||
|
|
||||||
# Check if the test is running on juseless
|
# Check if the test is running on juseless
|
||||||
if socket.gethostname() != "juseless":
|
if socket.gethostname() != "juseless":
|
||||||
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ import pytest
|
||||||
from junifer.configs.juseless.datagrabbers import JuselessUCLA
|
from junifer.configs.juseless.datagrabbers import JuselessUCLA
|
||||||
from junifer.utils.logging import configure_logging
|
from junifer.utils.logging import configure_logging
|
||||||
|
|
||||||
|
|
||||||
# Check if the test is running on juseless
|
# Check if the test is running on juseless
|
||||||
if socket.gethostname() != "juseless":
|
if socket.gethostname() != "juseless":
|
||||||
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import pytest
|
||||||
from junifer.configs.juseless.datagrabbers import JuselessDataladUKBVBM
|
from junifer.configs.juseless.datagrabbers import JuselessDataladUKBVBM
|
||||||
from junifer.utils.logging import configure_logging
|
from junifer.utils.logging import configure_logging
|
||||||
|
|
||||||
|
|
||||||
# Check if the test is running on juseless
|
# Check if the test is running on juseless
|
||||||
if socket.gethostname() != "juseless":
|
if socket.gethostname() != "juseless":
|
||||||
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import pytest
|
||||||
from junifer.datagrabber.hcp import DataladHCP1200
|
from junifer.datagrabber.hcp import DataladHCP1200
|
||||||
from junifer.utils.logging import configure_logging
|
from junifer.utils.logging import configure_logging
|
||||||
|
|
||||||
|
|
||||||
# Check if the test is running on juseless
|
# Check if the test is running on juseless
|
||||||
if socket.gethostname() != "juseless":
|
if socket.gethostname() != "juseless":
|
||||||
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ from numpy.typing import ArrayLike
|
||||||
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
# Path to the VOIs
|
# Path to the VOIs
|
||||||
_vois_path = Path(__file__).parent / "VOIs"
|
_vois_path = Path(__file__).parent / "VOIs"
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
|
||||||
from .utils import closest_resolution
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error
|
||||||
|
from .utils import closest_resolution
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|
@ -58,10 +58,9 @@ def register_mask(
|
||||||
if name in _available_masks:
|
if name in _available_masks:
|
||||||
if overwrite is True:
|
if overwrite is True:
|
||||||
logger.info(f"Overwriting {name} mask")
|
logger.info(f"Overwriting {name} mask")
|
||||||
if (_available_masks[name]["family"] != "CustomUserMask"):
|
if _available_masks[name]["family"] != "CustomUserMask":
|
||||||
raise_error(
|
raise_error(
|
||||||
f"Cannot overwrite {name} mask. "
|
f"Cannot overwrite {name} mask. " "It is a built-in mask."
|
||||||
"It is a built-in mask."
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise_error(
|
raise_error(
|
||||||
|
|
@ -117,8 +116,7 @@ def load_mask(
|
||||||
"""
|
"""
|
||||||
if name not in _available_masks:
|
if name not in _available_masks:
|
||||||
raise_error(
|
raise_error(
|
||||||
f"Mask {name} not found. "
|
f"Mask {name} not found. " f"Valid options are: {list_masks()}"
|
||||||
f"Valid options are: {list_masks()}"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
mask_definition = _available_masks[name].copy()
|
mask_definition = _available_masks[name].copy()
|
||||||
|
|
@ -126,12 +124,10 @@ def load_mask(
|
||||||
|
|
||||||
if t_family == "CustomUserMask":
|
if t_family == "CustomUserMask":
|
||||||
mask_fname = Path(mask_definition["path"])
|
mask_fname = Path(mask_definition["path"])
|
||||||
elif t_family == 'Vickery-Patil':
|
elif t_family == "Vickery-Patil":
|
||||||
mask_fname = _load_vickery_patil_mask(name, resolution)
|
mask_fname = _load_vickery_patil_mask(name, resolution)
|
||||||
else:
|
else:
|
||||||
raise_error(
|
raise_error(f"I don't know about the {t_family} mask family.")
|
||||||
f"I don't know about the {t_family} mask family."
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.info(f"Loading mask {mask_fname.absolute()}")
|
logger.info(f"Loading mask {mask_fname.absolute()}")
|
||||||
|
|
||||||
|
|
@ -167,8 +163,9 @@ def _load_vickery_patil_mask(
|
||||||
available_resolutions = [1.5, 3.0]
|
available_resolutions = [1.5, 3.0]
|
||||||
to_load = closest_resolution(resolution, available_resolutions)
|
to_load = closest_resolution(resolution, available_resolutions)
|
||||||
if to_load == 3.0:
|
if to_load == 3.0:
|
||||||
mask_fname = \
|
mask_fname = (
|
||||||
"CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz"
|
"CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean_3mm.nii.gz"
|
||||||
|
)
|
||||||
elif to_load == 1.5:
|
elif to_load == 1.5:
|
||||||
mask_fname = "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz"
|
mask_fname = "CAT12_IXI555_MNI152_TMP_GS_GMprob0.2_clean.nii.gz"
|
||||||
else:
|
else:
|
||||||
|
|
@ -178,9 +175,7 @@ def _load_vickery_patil_mask(
|
||||||
elif name == "GM_prob0.2_cortex":
|
elif name == "GM_prob0.2_cortex":
|
||||||
mask_fname = "GMprob0.2_cortex_3mm_NA_rm.nii.gz"
|
mask_fname = "GMprob0.2_cortex_3mm_NA_rm.nii.gz"
|
||||||
else:
|
else:
|
||||||
raise_error(
|
raise_error(f"Cannot find a Vickery-Patil mask called {name}")
|
||||||
f"Cannot find a Vickery-Patil mask called {name}"
|
|
||||||
)
|
|
||||||
mask_fname = _masks_path / "vickery-patil" / mask_fname
|
mask_fname = _masks_path / "vickery-patil" / mask_fname
|
||||||
|
|
||||||
return mask_fname
|
return mask_fname
|
||||||
|
|
|
||||||
|
|
@ -18,8 +18,9 @@ import pandas as pd
|
||||||
import requests
|
import requests
|
||||||
from nilearn import datasets
|
from nilearn import datasets
|
||||||
|
|
||||||
from .utils import closest_resolution
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error
|
||||||
|
from .utils import closest_resolution
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nibabel import Nifti1Image
|
from nibabel import Nifti1Image
|
||||||
|
|
|
||||||
|
|
@ -5,8 +5,8 @@
|
||||||
|
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import pytest
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
|
||||||
from junifer.data.utils import closest_resolution
|
from junifer.data.utils import closest_resolution
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -11,10 +11,10 @@ import pytest
|
||||||
from numpy.testing import assert_array_almost_equal
|
from numpy.testing import assert_array_almost_equal
|
||||||
|
|
||||||
from junifer.data.masks import (
|
from junifer.data.masks import (
|
||||||
|
_load_vickery_patil_mask,
|
||||||
|
list_masks,
|
||||||
load_mask,
|
load_mask,
|
||||||
register_mask,
|
register_mask,
|
||||||
list_masks,
|
|
||||||
_load_vickery_patil_mask,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -8,11 +8,10 @@
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import pytest
|
|
||||||
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
import pytest
|
||||||
from nilearn.image import new_img_like
|
from nilearn.image import new_img_like
|
||||||
|
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
||||||
|
|
||||||
from junifer.data.parcellations import (
|
from junifer.data.parcellations import (
|
||||||
_retrieve_parcellation,
|
_retrieve_parcellation,
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,5 @@
|
||||||
"""Provide utilities for data module."""
|
"""Provide utilities for data module."""
|
||||||
from typing import Optional, Union, List
|
from typing import List, Optional, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -111,8 +111,8 @@ def test_aomic1000_datagrabber() -> None:
|
||||||
assert out["DWI"]["path"].is_file()
|
assert out["DWI"]["path"].is_file()
|
||||||
|
|
||||||
# asserts meta
|
# asserts meta
|
||||||
assert "meta" in out
|
assert "meta" in out["BOLD"]
|
||||||
meta = out["meta"]
|
meta = out["BOLD"]["meta"]
|
||||||
assert "element" in meta
|
assert "element" in meta
|
||||||
assert "subject" in meta["element"]
|
assert "subject" in meta["element"]
|
||||||
assert test_element == meta["element"]["subject"]
|
assert test_element == meta["element"]["subject"]
|
||||||
|
|
|
||||||
|
|
@ -125,8 +125,8 @@ def test_aomic_piop1_datagrabber() -> None:
|
||||||
assert out["DWI"]["path"].is_file()
|
assert out["DWI"]["path"].is_file()
|
||||||
|
|
||||||
# asserts meta
|
# asserts meta
|
||||||
assert "meta" in out
|
assert "meta" in out["BOLD"]
|
||||||
meta = out["meta"]
|
meta = out["BOLD"]["meta"]
|
||||||
assert "element" in meta
|
assert "element" in meta
|
||||||
assert "subject" in meta["element"]
|
assert "subject" in meta["element"]
|
||||||
assert sub == meta["element"]["subject"]
|
assert sub == meta["element"]["subject"]
|
||||||
|
|
|
||||||
|
|
@ -121,8 +121,8 @@ def test_aomic_piop2_datagrabber() -> None:
|
||||||
assert out["DWI"]["path"].is_file()
|
assert out["DWI"]["path"].is_file()
|
||||||
|
|
||||||
# asserts meta
|
# asserts meta
|
||||||
assert "meta" in out
|
assert "meta" in out["BOLD"]
|
||||||
meta = out["meta"]
|
meta = out["BOLD"]["meta"]
|
||||||
assert "element" in meta
|
assert "element" in meta
|
||||||
assert "subject" in meta["element"]
|
assert "subject" in meta["element"]
|
||||||
assert sub == meta["element"]["subject"]
|
assert sub == meta["element"]["subject"]
|
||||||
|
|
|
||||||
|
|
@ -9,11 +9,12 @@ from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, Iterator, List, Tuple, Union
|
from typing import Dict, Iterator, List, Tuple, Union
|
||||||
|
|
||||||
|
from ..pipeline import UpdateMetaMixin
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
from .utils import validate_types
|
from .utils import validate_types
|
||||||
|
|
||||||
|
|
||||||
class BaseDataGrabber(ABC):
|
class BaseDataGrabber(ABC, UpdateMetaMixin):
|
||||||
"""Abstract base class for datagrabber.
|
"""Abstract base class for datagrabber.
|
||||||
|
|
||||||
For every interface that is required, one needs to provide a concrete
|
For every interface that is required, one needs to provide a concrete
|
||||||
|
|
@ -78,10 +79,11 @@ class BaseDataGrabber(ABC):
|
||||||
named_element = dict(zip(self.get_element_keys(), element))
|
named_element = dict(zip(self.get_element_keys(), element))
|
||||||
logger.debug(f"Named element: {named_element}")
|
logger.debug(f"Named element: {named_element}")
|
||||||
out = self.get_item(**named_element)
|
out = self.get_item(**named_element)
|
||||||
out["meta"] = {
|
|
||||||
"datagrabber": self.get_meta(),
|
for _, t_val in out.items():
|
||||||
"element": named_element,
|
self.update_meta(t_val, "datagrabber")
|
||||||
}
|
t_val["meta"]["element"] = named_element
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def __enter__(self) -> "BaseDataGrabber":
|
def __enter__(self) -> "BaseDataGrabber":
|
||||||
|
|
@ -103,22 +105,6 @@ class BaseDataGrabber(ABC):
|
||||||
"""
|
"""
|
||||||
return self.types.copy()
|
return self.types.copy()
|
||||||
|
|
||||||
def get_meta(self) -> Dict:
|
|
||||||
"""Get metadata.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
dict
|
|
||||||
The metadata as dictionary.
|
|
||||||
|
|
||||||
"""
|
|
||||||
t_meta = {}
|
|
||||||
t_meta["class"] = self.__class__.__name__
|
|
||||||
for k, v in vars(self).items():
|
|
||||||
if not k.startswith("_"):
|
|
||||||
t_meta[k] = v
|
|
||||||
return t_meta
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def datadir(self) -> Path:
|
def datadir(self) -> Path:
|
||||||
"""Get data directory path.
|
"""Get data directory path.
|
||||||
|
|
|
||||||
|
|
@ -83,7 +83,7 @@ class DataladDataGrabber(BaseDataGrabber):
|
||||||
self._rootdir = rootdir
|
self._rootdir = rootdir
|
||||||
# Flag to indicate if the dataset was cloned before and it might be
|
# Flag to indicate if the dataset was cloned before and it might be
|
||||||
# dirty
|
# dirty
|
||||||
self._dataset_dirty = False
|
self.datalad_dirty = False
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def datadir(self) -> Path:
|
def datadir(self) -> Path:
|
||||||
|
|
@ -175,7 +175,7 @@ class DataladDataGrabber(BaseDataGrabber):
|
||||||
# Check for dirty datasets:
|
# Check for dirty datasets:
|
||||||
status = self._dataset.status()
|
status = self._dataset.status()
|
||||||
if any([x["state"] != "clean" for x in status]):
|
if any([x["state"] != "clean" for x in status]):
|
||||||
self._dataset_dirty = True
|
self.datalad_dirty = True
|
||||||
warn_with_log(
|
warn_with_log(
|
||||||
"At least one file is not clean, Junifer will "
|
"At least one file is not clean, Junifer will "
|
||||||
"consider this dataset as dirty."
|
"consider this dataset as dirty."
|
||||||
|
|
@ -191,11 +191,10 @@ class DataladDataGrabber(BaseDataGrabber):
|
||||||
logger.debug("Dataset installed")
|
logger.debug("Dataset installed")
|
||||||
self._was_cloned = not isinstalled
|
self._was_cloned = not isinstalled
|
||||||
|
|
||||||
self._datalad_commit_id = (
|
self.datalad_commit_id = self._dataset.repo.get_hexsha( # type: ignore
|
||||||
self._dataset.repo.get_hexsha( # type: ignore
|
self._dataset.repo.get_corresponding_branch() # type: ignore
|
||||||
self._dataset.repo.get_corresponding_branch() # type: ignore
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
self.datalad_id = self._dataset.id
|
||||||
|
|
||||||
def cleanup(self) -> None:
|
def cleanup(self) -> None:
|
||||||
"""Cleanup the datalad dataset."""
|
"""Cleanup the datalad dataset."""
|
||||||
|
|
@ -245,21 +244,3 @@ class DataladDataGrabber(BaseDataGrabber):
|
||||||
logger.debug("Cleaning up dataset")
|
logger.debug("Cleaning up dataset")
|
||||||
self.cleanup()
|
self.cleanup()
|
||||||
logger.debug("Dataset state restored")
|
logger.debug("Dataset state restored")
|
||||||
|
|
||||||
def get_meta(self) -> Dict:
|
|
||||||
"""Get metadata.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
dict
|
|
||||||
The metadata as dictionary.
|
|
||||||
|
|
||||||
"""
|
|
||||||
t_meta = super().get_meta()
|
|
||||||
t_meta["datalad_commit_id"] = self._datalad_commit_id
|
|
||||||
|
|
||||||
t_meta["datalad_id"] = self._dataset.id
|
|
||||||
|
|
||||||
# Set a flag to indicate that the dataset was dirty
|
|
||||||
t_meta["datalad_dirty"] = self._dataset_dirty
|
|
||||||
return t_meta
|
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,6 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Dict, List, Tuple, Union
|
from typing import Dict, List, Tuple, Union
|
||||||
|
|
||||||
from .base import BaseDataGrabber
|
from .base import BaseDataGrabber
|
||||||
|
|
@ -38,7 +37,7 @@ class MultipleDataGrabber(BaseDataGrabber):
|
||||||
raise ValueError("Datagrabbers have overlapping types.")
|
raise ValueError("Datagrabbers have overlapping types.")
|
||||||
self._datagrabbers = datagrabbers
|
self._datagrabbers = datagrabbers
|
||||||
|
|
||||||
def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]:
|
def __getitem__(self, element: Union[str, Tuple]) -> Dict:
|
||||||
"""Implement indexing.
|
"""Implement indexing.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -58,9 +57,20 @@ class MultipleDataGrabber(BaseDataGrabber):
|
||||||
"""
|
"""
|
||||||
|
|
||||||
out = {}
|
out = {}
|
||||||
|
metas = []
|
||||||
for dg in self._datagrabbers:
|
for dg in self._datagrabbers:
|
||||||
t_out = dg[element]
|
t_out = dg[element]
|
||||||
out.update(t_out)
|
out.update(t_out)
|
||||||
|
# Now get the meta for this datagrabber
|
||||||
|
t_meta = {}
|
||||||
|
dg.update_meta(t_meta, "datagrabber")
|
||||||
|
# Store all the sub-datagrabbers meta
|
||||||
|
metas.append(t_meta["meta"]["datagrabber"])
|
||||||
|
|
||||||
|
# Update all the metas again
|
||||||
|
for kind in out:
|
||||||
|
self.update_meta(out[kind], "datagrabber")
|
||||||
|
out[kind]["meta"]["datagrabber"]["datagrabbers"] = metas
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def get_item(self, **element: Dict) -> Dict[str, Dict]:
|
def get_item(self, **element: Dict) -> Dict[str, Dict]:
|
||||||
|
|
@ -136,17 +146,3 @@ class MultipleDataGrabber(BaseDataGrabber):
|
||||||
"""
|
"""
|
||||||
types = [x for dg in self._datagrabbers for x in dg.get_types()]
|
types = [x for dg in self._datagrabbers for x in dg.get_types()]
|
||||||
return types
|
return types
|
||||||
|
|
||||||
def get_meta(self) -> Dict:
|
|
||||||
"""Get metadata.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
dict
|
|
||||||
The metadata as dictionary.
|
|
||||||
|
|
||||||
"""
|
|
||||||
t_meta = {}
|
|
||||||
t_meta["class"] = self.__class__.__name__
|
|
||||||
t_meta["datagrabbers"] = [dg.get_meta() for dg in self._datagrabbers]
|
|
||||||
return t_meta
|
|
||||||
|
|
|
||||||
|
|
@ -23,7 +23,7 @@ def test_BaseDataGrabber() -> None:
|
||||||
# Create concrete class.
|
# Create concrete class.
|
||||||
class MyDataGrabber(BaseDataGrabber):
|
class MyDataGrabber(BaseDataGrabber):
|
||||||
def get_item(self, subject):
|
def get_item(self, subject):
|
||||||
return {}
|
return {"BOLD": {}}
|
||||||
|
|
||||||
def get_elements(self):
|
def get_elements(self):
|
||||||
return super().get_elements()
|
return super().get_elements()
|
||||||
|
|
@ -31,22 +31,24 @@ def test_BaseDataGrabber() -> None:
|
||||||
def get_element_keys(self):
|
def get_element_keys(self):
|
||||||
return ["subject"]
|
return ["subject"]
|
||||||
|
|
||||||
dg = MyDataGrabber(datadir="/tmp", types=["func"])
|
dg = MyDataGrabber(datadir="/tmp", types=["BOLD"])
|
||||||
elem = dg["elem"]
|
elem = dg["sub01"]
|
||||||
assert "meta" in elem
|
assert "BOLD" in elem
|
||||||
assert "datagrabber" in elem["meta"]
|
assert "meta" in elem["BOLD"]
|
||||||
assert "class" in elem["meta"]["datagrabber"]
|
meta = elem["BOLD"]["meta"]
|
||||||
assert MyDataGrabber.__name__ in elem["meta"]["datagrabber"]["class"]
|
assert "datagrabber" in meta
|
||||||
assert "element" in elem["meta"]
|
assert "class" in meta["datagrabber"]
|
||||||
assert "subject" in elem["meta"]["element"]
|
assert MyDataGrabber.__name__ in meta["datagrabber"]["class"]
|
||||||
assert "elem" in elem["meta"]["element"]["subject"]
|
assert "element" in meta
|
||||||
|
assert "subject" in meta["element"]
|
||||||
|
assert "sub01" in meta["element"]["subject"]
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
dg.get_elements()
|
dg.get_elements()
|
||||||
|
|
||||||
with dg:
|
with dg:
|
||||||
assert dg.datadir == Path("/tmp")
|
assert dg.datadir == Path("/tmp")
|
||||||
assert dg.types == ["func"]
|
assert dg.types == ["BOLD"]
|
||||||
|
|
||||||
class MyDataGrabber2(BaseDataGrabber):
|
class MyDataGrabber2(BaseDataGrabber):
|
||||||
def get_item(self, subject):
|
def get_item(self, subject):
|
||||||
|
|
@ -58,7 +60,7 @@ def test_BaseDataGrabber() -> None:
|
||||||
def get_element_keys(self):
|
def get_element_keys(self):
|
||||||
return super().get_element_keys()
|
return super().get_element_keys()
|
||||||
|
|
||||||
dg = MyDataGrabber2(datadir="/tmp", types=["func"])
|
dg = MyDataGrabber2(datadir="/tmp", types=["BOLD"])
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
dg.get_element_keys()
|
dg.get_element_keys()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import pytest
|
||||||
|
|
||||||
from junifer.datagrabber.datalad_base import DataladDataGrabber
|
from junifer.datagrabber.datalad_base import DataladDataGrabber
|
||||||
|
|
||||||
|
|
||||||
_testing_dataset = {
|
_testing_dataset = {
|
||||||
"example_bids": {
|
"example_bids": {
|
||||||
"uri": "https://gin.g-node.org/juaml/datalad-example-bids",
|
"uri": "https://gin.g-node.org/juaml/datalad-example-bids",
|
||||||
|
|
@ -143,10 +144,11 @@ def test_datalad_clone_cleanup(
|
||||||
assert elem1_t1w.is_file() is False
|
assert elem1_t1w.is_file() is False
|
||||||
assert elem1_t1w.is_symlink() is True
|
assert elem1_t1w.is_symlink() is True
|
||||||
elem1 = dg["sub-01"]
|
elem1 = dg["sub-01"]
|
||||||
assert "meta" in elem1
|
assert "meta" in elem1["BOLD"]
|
||||||
assert "datagrabber" in elem1["meta"]
|
meta = elem1["BOLD"]["meta"]
|
||||||
assert "datalad_dirty" in elem1["meta"]["datagrabber"]
|
assert "datagrabber" in meta
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False
|
assert "datalad_dirty" in meta["datagrabber"]
|
||||||
|
assert meta["datagrabber"]["datalad_dirty"] is False
|
||||||
assert hasattr(dg, "_got_files") is False
|
assert hasattr(dg, "_got_files") is False
|
||||||
assert datadir.exists() is True
|
assert datadir.exists() is True
|
||||||
assert elem1_bold.is_file() is True
|
assert elem1_bold.is_file() is True
|
||||||
|
|
@ -198,14 +200,15 @@ def test_datalad_previously_cloned(
|
||||||
assert datadir.exists() is True
|
assert datadir.exists() is True
|
||||||
assert dg._was_cloned is False
|
assert dg._was_cloned is False
|
||||||
elem1 = dg["sub-01"]
|
elem1 = dg["sub-01"]
|
||||||
assert "meta" in elem1
|
assert "meta" in elem1["BOLD"]
|
||||||
assert "datagrabber" in elem1["meta"]
|
meta = elem1["BOLD"]["meta"]
|
||||||
assert "datalad_dirty" in elem1["meta"]["datagrabber"]
|
assert "datagrabber" in meta
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False
|
assert "datalad_dirty" in meta["datagrabber"]
|
||||||
assert "datalad_commit_id" in elem1["meta"]["datagrabber"]
|
assert meta["datagrabber"]["datalad_dirty"] is False
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit
|
assert "datalad_commit_id" in meta["datagrabber"]
|
||||||
assert "datalad_id" in elem1["meta"]["datagrabber"]
|
assert meta["datagrabber"]["datalad_commit_id"] == commit
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id
|
assert "datalad_id" in meta["datagrabber"]
|
||||||
|
assert meta["datagrabber"]["datalad_id"] == remote_id
|
||||||
|
|
||||||
assert hasattr(dg, "_got_files") is True
|
assert hasattr(dg, "_got_files") is True
|
||||||
# Files are there and symlinks are fixed
|
# Files are there and symlinks are fixed
|
||||||
|
|
@ -264,7 +267,8 @@ def test_datalad_previously_cloned_and_get(
|
||||||
assert elem1_t1w.is_file() is False
|
assert elem1_t1w.is_file() is False
|
||||||
|
|
||||||
dl.get( # type: ignore
|
dl.get( # type: ignore
|
||||||
elem1_t1w, dataset=datadir, result_renderer="disabled")
|
elem1_t1w, dataset=datadir, result_renderer="disabled"
|
||||||
|
)
|
||||||
|
|
||||||
assert elem1_bold.is_symlink() is True
|
assert elem1_bold.is_symlink() is True
|
||||||
assert elem1_bold.is_file() is False
|
assert elem1_bold.is_file() is False
|
||||||
|
|
@ -275,14 +279,15 @@ def test_datalad_previously_cloned_and_get(
|
||||||
assert datadir.exists() is True
|
assert datadir.exists() is True
|
||||||
assert dg._was_cloned is False
|
assert dg._was_cloned is False
|
||||||
elem1 = dg["sub-01"]
|
elem1 = dg["sub-01"]
|
||||||
assert "meta" in elem1
|
assert "meta" in elem1["BOLD"]
|
||||||
assert "datagrabber" in elem1["meta"]
|
meta = elem1["BOLD"]["meta"]
|
||||||
assert "datalad_dirty" in elem1["meta"]["datagrabber"]
|
assert "datagrabber" in meta
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is False
|
assert "datalad_dirty" in meta["datagrabber"]
|
||||||
assert "datalad_commit_id" in elem1["meta"]["datagrabber"]
|
assert meta["datagrabber"]["datalad_dirty"] is False
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit
|
assert "datalad_commit_id" in meta["datagrabber"]
|
||||||
assert "datalad_id" in elem1["meta"]["datagrabber"]
|
assert meta["datagrabber"]["datalad_commit_id"] == commit
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id
|
assert "datalad_id" in meta["datagrabber"]
|
||||||
|
assert meta["datagrabber"]["datalad_id"] == remote_id
|
||||||
|
|
||||||
assert hasattr(dg, "_got_files") is True
|
assert hasattr(dg, "_got_files") is True
|
||||||
# Files are there and symlinks are fixed
|
# Files are there and symlinks are fixed
|
||||||
|
|
@ -344,7 +349,8 @@ def test_datalad_previously_cloned_and_get_dirty(
|
||||||
assert elem1_t1w.is_file() is False
|
assert elem1_t1w.is_file() is False
|
||||||
|
|
||||||
dl.get( # type: ignore
|
dl.get( # type: ignore
|
||||||
elem1_t1w, dataset=datadir, result_renderer="disabled")
|
elem1_t1w, dataset=datadir, result_renderer="disabled"
|
||||||
|
)
|
||||||
|
|
||||||
assert elem1_bold.is_symlink() is True
|
assert elem1_bold.is_symlink() is True
|
||||||
assert elem1_bold.is_file() is False
|
assert elem1_bold.is_file() is False
|
||||||
|
|
@ -359,14 +365,15 @@ def test_datalad_previously_cloned_and_get_dirty(
|
||||||
assert datadir.exists() is True
|
assert datadir.exists() is True
|
||||||
assert dg._was_cloned is False
|
assert dg._was_cloned is False
|
||||||
elem1 = dg["sub-01"]
|
elem1 = dg["sub-01"]
|
||||||
assert "meta" in elem1
|
assert "meta" in elem1["BOLD"]
|
||||||
assert "datagrabber" in elem1["meta"]
|
meta = elem1["BOLD"]["meta"]
|
||||||
assert "datalad_dirty" in elem1["meta"]["datagrabber"]
|
assert "datagrabber" in meta
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_dirty"] is True
|
assert "datalad_dirty" in meta["datagrabber"]
|
||||||
assert "datalad_commit_id" in elem1["meta"]["datagrabber"]
|
assert meta["datagrabber"]["datalad_dirty"] is True
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_commit_id"] == commit
|
assert "datalad_commit_id" in meta["datagrabber"]
|
||||||
assert "datalad_id" in elem1["meta"]["datagrabber"]
|
assert meta["datagrabber"]["datalad_commit_id"] == commit
|
||||||
assert elem1["meta"]["datagrabber"]["datalad_id"] == remote_id
|
assert "datalad_id" in meta["datagrabber"]
|
||||||
|
assert meta["datagrabber"]["datalad_id"] == remote_id
|
||||||
|
|
||||||
assert hasattr(dg, "_got_files") is True
|
assert hasattr(dg, "_got_files") is True
|
||||||
# Files are there and symlinks are fixed
|
# Files are there and symlinks are fixed
|
||||||
|
|
@ -384,17 +391,18 @@ def test_datalad_previously_cloned_and_get_dirty(
|
||||||
assert datadir.exists() is True
|
assert datadir.exists() is True
|
||||||
assert dg._was_cloned is False
|
assert dg._was_cloned is False
|
||||||
elem2 = dg["sub-02"]
|
elem2 = dg["sub-02"]
|
||||||
assert "meta" in elem2
|
assert "meta" in elem1["BOLD"]
|
||||||
assert "datagrabber" in elem2["meta"]
|
meta = elem2["BOLD"]["meta"]
|
||||||
assert "datalad_dirty" in elem2["meta"]["datagrabber"]
|
assert "datagrabber" in meta
|
||||||
|
assert "datalad_dirty" in meta["datagrabber"]
|
||||||
|
|
||||||
# Dataset is still dirty due to subject sub-01
|
# Dataset is still dirty due to subject sub-01
|
||||||
assert elem2["meta"]["datagrabber"]["datalad_dirty"] is True
|
assert meta["datagrabber"]["datalad_dirty"] is True
|
||||||
|
|
||||||
assert "datalad_commit_id" in elem2["meta"]["datagrabber"]
|
assert "datalad_commit_id" in meta["datagrabber"]
|
||||||
assert elem2["meta"]["datagrabber"]["datalad_commit_id"] == commit
|
assert meta["datagrabber"]["datalad_commit_id"] == commit
|
||||||
assert "datalad_id" in elem2["meta"]["datagrabber"]
|
assert "datalad_id" in meta["datagrabber"]
|
||||||
assert elem2["meta"]["datagrabber"]["datalad_id"] == remote_id
|
assert meta["datagrabber"]["datalad_id"] == remote_id
|
||||||
|
|
||||||
assert hasattr(dg, "_got_files") is True
|
assert hasattr(dg, "_got_files") is True
|
||||||
# Files are there and symlinks are fixed
|
# Files are there and symlinks are fixed
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ import pytest
|
||||||
from junifer.datagrabber.hcp import DataladHCP1200
|
from junifer.datagrabber.hcp import DataladHCP1200
|
||||||
from junifer.utils import configure_logging
|
from junifer.utils import configure_logging
|
||||||
|
|
||||||
|
|
||||||
URI = "https://gin.g-node.org/juaml/datalad-example-hcp1200"
|
URI = "https://gin.g-node.org/juaml/datalad-example-hcp1200"
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -79,8 +80,8 @@ def test_dataladhcp1200_datagrabber(
|
||||||
# Assert data file path is a file
|
# Assert data file path is a file
|
||||||
assert out["BOLD"]["path"].is_file()
|
assert out["BOLD"]["path"].is_file()
|
||||||
# Assert metadata
|
# Assert metadata
|
||||||
assert "meta" in out
|
assert "meta" in out["BOLD"]
|
||||||
meta = out["meta"]
|
meta = out["BOLD"]["meta"]
|
||||||
assert "element" in meta
|
assert "element" in meta
|
||||||
assert "subject" in meta["element"]
|
assert "subject" in meta["element"]
|
||||||
assert test_element[0] == meta["element"]["subject"]
|
assert test_element[0] == meta["element"]["subject"]
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ import pytest
|
||||||
|
|
||||||
from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber
|
from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber
|
||||||
|
|
||||||
|
|
||||||
_testing_dataset = {
|
_testing_dataset = {
|
||||||
"example_bids": {
|
"example_bids": {
|
||||||
"uri": "https://gin.g-node.org/juaml/datalad-example-bids",
|
"uri": "https://gin.g-node.org/juaml/datalad-example-bids",
|
||||||
|
|
@ -63,17 +64,17 @@ def test_multiple() -> None:
|
||||||
subs = [x for x in dg]
|
subs = [x for x in dg]
|
||||||
assert set(subs) == set(expected_subs)
|
assert set(subs) == set(expected_subs)
|
||||||
|
|
||||||
data = dg[("sub-01", "ses-01")]
|
elem = dg[("sub-01", "ses-01")]
|
||||||
assert "T1w" in data
|
assert "T1w" in elem
|
||||||
assert "BOLD" in data
|
assert "BOLD" in elem
|
||||||
|
assert "meta" in elem["BOLD"]
|
||||||
meta = dg.get_meta()
|
meta = elem["BOLD"]["meta"]["datagrabber"]
|
||||||
assert "class" in meta
|
assert "class" in meta
|
||||||
assert meta["class"] == "MultipleDataGrabber"
|
assert meta["class"] == "MultipleDataGrabber"
|
||||||
assert "datagrabbers" in meta
|
assert "datagrabbers" in meta
|
||||||
assert len(meta["datagrabbers"]) == 2
|
assert len(meta["datagrabbers"]) == 2
|
||||||
assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber"
|
assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber"
|
||||||
assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber"
|
assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber"
|
||||||
|
|
||||||
|
|
||||||
def test_multiple_no_intersection() -> None:
|
def test_multiple_no_intersection() -> None:
|
||||||
|
|
|
||||||
|
|
@ -256,6 +256,6 @@ def test_PatternDataGrabber(tmp_path: Path) -> None:
|
||||||
out1 = datagrabber[("sub000", "ses000", "task002")]
|
out1 = datagrabber[("sub000", "ses000", "task002")]
|
||||||
out2 = datagrabber[("sub000", "ses000", "task003")]
|
out2 = datagrabber[("sub000", "ses000", "task003")]
|
||||||
|
|
||||||
assert out1["func"] == out2["func"]
|
assert out1["func"]["path"] == out2["func"]["path"]
|
||||||
assert out1["anat"] == out2["anat"]
|
assert out1["anat"]["path"] == out2["anat"]["path"]
|
||||||
assert out1["vbm"] != out2["vbm"]
|
assert out1["vbm"]["path"] != out2["vbm"]["path"]
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import pytest
|
||||||
|
|
||||||
from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber
|
from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber
|
||||||
|
|
||||||
|
|
||||||
_testing_dataset = {
|
_testing_dataset = {
|
||||||
"example_bids": {
|
"example_bids": {
|
||||||
"uri": "https://gin.g-node.org/juaml/datalad-example-bids",
|
"uri": "https://gin.g-node.org/juaml/datalad-example-bids",
|
||||||
|
|
@ -81,9 +82,10 @@ def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None:
|
||||||
dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz"
|
dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz"
|
||||||
)
|
)
|
||||||
|
|
||||||
assert "meta" in t_sub
|
assert "meta" in t_sub["BOLD"]
|
||||||
assert "datagrabber" in t_sub["meta"]
|
meta = t_sub["BOLD"]["meta"]
|
||||||
dg_meta = t_sub["meta"]["datagrabber"]
|
assert "datagrabber" in meta
|
||||||
|
dg_meta = meta["datagrabber"]
|
||||||
assert "class" in dg_meta
|
assert "class" in dg_meta
|
||||||
assert dg_meta["class"] == "PatternDataladDataGrabber"
|
assert dg_meta["class"] == "PatternDataladDataGrabber"
|
||||||
assert "uri" in dg_meta
|
assert "uri" in dg_meta
|
||||||
|
|
|
||||||
|
|
@ -11,10 +11,11 @@ import nibabel as nib
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from ..api.decorators import register_datareader
|
from ..api.decorators import register_datareader
|
||||||
from ..pipeline import PipelineStepMixin
|
from ..pipeline import PipelineStepMixin, UpdateMetaMixin
|
||||||
from ..utils.logging import logger, warn_with_log
|
from ..utils.logging import logger, warn_with_log
|
||||||
|
|
||||||
# Map each file extension to a kind
|
|
||||||
|
# Map each file extension to a type
|
||||||
_extensions = {
|
_extensions = {
|
||||||
".nii": "NIFTI",
|
".nii": "NIFTI",
|
||||||
".nii.gz": "NIFTI",
|
".nii.gz": "NIFTI",
|
||||||
|
|
@ -22,7 +23,7 @@ _extensions = {
|
||||||
".tsv": "TSV",
|
".tsv": "TSV",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Map each kind to a function and arguments
|
# Map each type to a function and arguments
|
||||||
_readers = {}
|
_readers = {}
|
||||||
_readers["NIFTI"] = {"func": nib.load, "params": None}
|
_readers["NIFTI"] = {"func": nib.load, "params": None}
|
||||||
_readers["CSV"] = {"func": pd.read_csv, "params": None}
|
_readers["CSV"] = {"func": pd.read_csv, "params": None}
|
||||||
|
|
@ -30,7 +31,7 @@ _readers["TSV"] = {"func": pd.read_csv, "params": {"sep": "\t"}}
|
||||||
|
|
||||||
|
|
||||||
@register_datareader
|
@register_datareader
|
||||||
class DefaultDataReader(PipelineStepMixin):
|
class DefaultDataReader(PipelineStepMixin, UpdateMetaMixin):
|
||||||
"""Mixin class for default data reader."""
|
"""Mixin class for default data reader."""
|
||||||
|
|
||||||
def validate_input(self, input: List[str]) -> None:
|
def validate_input(self, input: List[str]) -> None:
|
||||||
|
|
@ -46,8 +47,8 @@ class DefaultDataReader(PipelineStepMixin):
|
||||||
# Nothing to validate, any input is fine
|
# Nothing to validate, any input is fine
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input: List[str]) -> List[str]:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
|
@ -58,10 +59,10 @@ class DefaultDataReader(PipelineStepMixin):
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
list of str
|
||||||
The updated list of output kinds, as reading possibilities.
|
The updated list of output types, as reading possibilities.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# It will output the same kind of data as the input
|
# It will output the same type of data as the input
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def fit_transform(
|
def fit_transform(
|
||||||
|
|
@ -82,36 +83,33 @@ class DefaultDataReader(PipelineStepMixin):
|
||||||
-------
|
-------
|
||||||
dict
|
dict
|
||||||
The processed output as dictionary. The "data" key is added to
|
The processed output as dictionary. The "data" key is added to
|
||||||
each data type dictionary except "meta".
|
each data type dictionary.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# For each kind of data, try to read it
|
# For each type of data, try to read it
|
||||||
out = input.copy()
|
out = input.copy()
|
||||||
if params is None:
|
if params is None:
|
||||||
params = {}
|
params = {}
|
||||||
for kind in input.keys():
|
for type_ in input.keys():
|
||||||
if kind == "meta":
|
if "path" not in input[type_]:
|
||||||
out["meta"] = input["meta"]
|
|
||||||
continue
|
|
||||||
if "path" not in input[kind]:
|
|
||||||
warn_with_log(
|
warn_with_log(
|
||||||
f"Input kind {kind} does not provide a path. Skipping."
|
f"Input type {type_} does not provide a path. Skipping."
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
t_path = input[kind]["path"]
|
t_path = input[type_]["path"]
|
||||||
t_params = params.get(kind, {})
|
t_params = params.get(type_, {})
|
||||||
|
|
||||||
# Convert to Path if datareader is not well done
|
# Convert to Path if datareader is not well done
|
||||||
if not isinstance(t_path, Path):
|
if not isinstance(t_path, Path):
|
||||||
t_path = Path(t_path)
|
t_path = Path(t_path)
|
||||||
out[kind]["path"] = t_path
|
out[type_]["path"] = t_path
|
||||||
logger.info(f"Reading {kind} from {t_path.as_posix()}")
|
logger.info(f"Reading {type_} from {t_path.as_posix()}")
|
||||||
fread = None
|
fread = None
|
||||||
|
|
||||||
fname = t_path.name.lower()
|
fname = t_path.name.lower()
|
||||||
for ext, ftype in _extensions.items():
|
for ext, ftype in _extensions.items():
|
||||||
if fname.endswith(ext):
|
if fname.endswith(ext):
|
||||||
logger.info(f"{kind} is type {ftype}")
|
logger.info(f"{type_} is type {ftype}")
|
||||||
reader_func = _readers[ftype]["func"]
|
reader_func = _readers[ftype]["func"]
|
||||||
reader_params = _readers[ftype]["params"]
|
reader_params = _readers[ftype]["params"]
|
||||||
if reader_params is not None:
|
if reader_params is not None:
|
||||||
|
|
@ -123,8 +121,6 @@ class DefaultDataReader(PipelineStepMixin):
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Unknown file type {t_path.as_posix()}, skipping reading"
|
f"Unknown file type {t_path.as_posix()}, skipping reading"
|
||||||
)
|
)
|
||||||
out[kind]["data"] = fread
|
out[type_]["data"] = fread
|
||||||
if "meta" not in out:
|
self.update_meta(out[type_], "datareader")
|
||||||
out["meta"] = {}
|
|
||||||
out["meta"]["datareader"] = self.get_meta()
|
|
||||||
return out
|
return out
|
||||||
|
|
|
||||||
|
|
@ -17,37 +17,36 @@ from junifer.datareader import DefaultDataReader
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"kind", [["T1w", "BOLD", "T2", "dwi"], [], None, ["whatever"]]
|
"type_", [["T1w", "BOLD", "T2", "dwi"], [], ["whatever"]]
|
||||||
)
|
)
|
||||||
def test_validation(kind) -> None:
|
def test_validation(type_) -> None:
|
||||||
"""Test validating input/output.
|
"""Test validating input/output.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
kind : list of str or str or None
|
type_ : list of str or str or None
|
||||||
The parametrized kind of data.
|
The parametrized type_ of data.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
reader = DefaultDataReader()
|
reader = DefaultDataReader()
|
||||||
assert reader.validate_input(kind) is None
|
assert reader.validate_input(type_) is None
|
||||||
assert reader.get_output_kind(kind) == kind
|
assert reader.get_output_type(type_) == type_
|
||||||
assert reader.validate(kind) == kind
|
assert reader.validate(type_) == type_
|
||||||
|
|
||||||
|
|
||||||
def test_meta() -> None:
|
def test_meta() -> None:
|
||||||
"""Test reader metadata."""
|
"""Test reader metadata."""
|
||||||
reader = DefaultDataReader()
|
reader = DefaultDataReader()
|
||||||
t_meta = reader.get_meta()
|
|
||||||
assert t_meta["class"] == "DefaultDataReader"
|
|
||||||
|
|
||||||
nib_data_path = Path(nib_testing.data_path)
|
nib_data_path = Path(nib_testing.data_path)
|
||||||
t_path = nib_data_path / "example4d.nii.gz"
|
t_path = nib_data_path / "example4d.nii.gz"
|
||||||
input = {"BOLD": {"path": t_path}}
|
input = {"BOLD": {"path": t_path}}
|
||||||
output = reader.fit_transform(input)
|
output = reader.fit_transform(input)
|
||||||
assert "meta" in output
|
assert "meta" in output["BOLD"]
|
||||||
assert "datareader" in output["meta"]
|
meta = output["BOLD"]["meta"]
|
||||||
assert "class" in output["meta"]["datareader"]
|
assert "datareader" in meta
|
||||||
assert output["meta"]["datareader"]["class"] == "DefaultDataReader"
|
assert "class" in meta["datareader"]
|
||||||
|
assert meta["datareader"]["class"] == "DefaultDataReader"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
|
|
|
||||||
|
|
@ -7,14 +7,15 @@
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from ..pipeline import PipelineStepMixin
|
from ..pipeline import PipelineStepMixin, UpdateMetaMixin
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from junifer.storage import BaseFeatureStorage
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
class BaseMarker(ABC, PipelineStepMixin):
|
class BaseMarker(ABC, PipelineStepMixin, UpdateMetaMixin):
|
||||||
"""Abstract base class for all markers.
|
"""Abstract base class for all markers.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -57,27 +58,6 @@ class BaseMarker(ABC, PipelineStepMixin):
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_meta(self, kind: str) -> Dict:
|
|
||||||
"""Get metadata.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : str
|
|
||||||
The kind of pipeline step.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
dict
|
|
||||||
The metadata as a dictionary with the only key 'marker'.
|
|
||||||
|
|
||||||
"""
|
|
||||||
s_meta = super().get_meta()
|
|
||||||
# same marker can be "fit"ted into different kinds, so the name
|
|
||||||
# is created from the kind and the name of the marker
|
|
||||||
s_meta["name"] = f"{kind}_{self.name}"
|
|
||||||
s_meta["kind"] = kind
|
|
||||||
return {"marker": s_meta}
|
|
||||||
|
|
||||||
def validate_input(self, input: List[str]) -> None:
|
def validate_input(self, input: List[str]) -> None:
|
||||||
"""Validate input.
|
"""Validate input.
|
||||||
|
|
||||||
|
|
@ -101,23 +81,22 @@ class BaseMarker(ABC, PipelineStepMixin):
|
||||||
)
|
)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The input to the marker. The list must contain the
|
The data type input to the marker.
|
||||||
available Junifer Data dictionary keys.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The updated list of output kinds, as storage possibilities.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement get_output_kind().",
|
msg="Concrete classes need to implement get_output_type().",
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -149,10 +128,9 @@ class BaseMarker(ABC, PipelineStepMixin):
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def store(
|
def store(
|
||||||
self,
|
self,
|
||||||
kind: str,
|
type_: str,
|
||||||
out: Dict[str, Any],
|
out: Dict[str, Any],
|
||||||
storage: "BaseFeatureStorage",
|
storage: "BaseFeatureStorage",
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -160,18 +138,17 @@ class BaseMarker(ABC, PipelineStepMixin):
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
kind : str
|
type_ : str
|
||||||
The data kind to store.
|
The data type to store.
|
||||||
out : dict
|
out : dict
|
||||||
The computed result as a dictionary to store.
|
The computed result as a dictionary to store.
|
||||||
storage : storage-like
|
storage : storage-like
|
||||||
The storage class.
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
output_type_ = self.get_output_type(type_)
|
||||||
msg="Concrete classes need to implement store().",
|
logger.debug(f"Storing {output_type_} in {storage}")
|
||||||
klass=NotImplementedError,
|
storage.store(kind=output_type_, **out)
|
||||||
)
|
|
||||||
|
|
||||||
def fit_transform(
|
def fit_transform(
|
||||||
self,
|
self,
|
||||||
|
|
@ -195,23 +172,25 @@ class BaseMarker(ABC, PipelineStepMixin):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
out = {}
|
out = {}
|
||||||
meta = input.get("meta", {})
|
for type_ in self._on:
|
||||||
for kind in self._on:
|
if type_ in input.keys():
|
||||||
if kind in input.keys():
|
logger.info(f"Computing {type_}")
|
||||||
logger.info(f"Computing {kind}")
|
t_input = input[type_]
|
||||||
t_input = input[kind]
|
|
||||||
extra_input = input.copy()
|
extra_input = input.copy()
|
||||||
extra_input.pop(kind)
|
extra_input.pop(type_)
|
||||||
t_meta = meta.copy()
|
t_meta = t_input["meta"].copy()
|
||||||
t_meta.update(t_input.get("meta", {}))
|
t_meta["type"] = type_
|
||||||
t_meta.update(self.get_meta(kind))
|
|
||||||
t_out = self.compute(input=t_input, extra_input=extra_input)
|
t_out = self.compute(input=t_input, extra_input=extra_input)
|
||||||
t_out.update(meta=t_meta)
|
t_out["meta"] = t_meta
|
||||||
|
|
||||||
|
self.update_meta(t_out, "marker")
|
||||||
|
|
||||||
if storage is not None:
|
if storage is not None:
|
||||||
logger.info(f"Storing in {storage}")
|
logger.info(f"Storing in {storage}")
|
||||||
self.store(kind=kind, out=t_out, storage=storage)
|
self.store(type_=type_, out=t_out, storage=storage)
|
||||||
else:
|
else:
|
||||||
logger.info("No storage specified, returning dictionary")
|
logger.info("No storage specified, returning dictionary")
|
||||||
out[kind] = t_out
|
out[type_] = t_out
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ from ..pipeline import PipelineStepMixin
|
||||||
from ..storage.base import BaseFeatureStorage
|
from ..storage.base import BaseFeatureStorage
|
||||||
from ..utils import logger
|
from ..utils import logger
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from junifer.datagrabber import BaseDataGrabber
|
from junifer.datagrabber import BaseDataGrabber
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -9,7 +9,6 @@ from typing import Any, Dict, List, Optional
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from ..api.decorators import register_marker
|
from ..api.decorators import register_marker
|
||||||
from ..storage import BaseFeatureStorage
|
|
||||||
from ..utils import logger
|
from ..utils import logger
|
||||||
from ..utils.logging import raise_error
|
from ..utils.logging import raise_error
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
|
|
@ -41,6 +40,8 @@ class CrossParcellationFC(BaseMarker):
|
||||||
(default None).
|
(default None).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
parcellation_one: str,
|
parcellation_one: str,
|
||||||
|
|
@ -72,43 +73,21 @@ class CrossParcellationFC(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["BOLD"]
|
return ["BOLD"]
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The input to the marker. The list must contain the
|
The data type input to the marker.
|
||||||
available Junifer Data dictionary keys.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The updated list of output kinds, as storage possibilities.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
return ["matrix"]
|
return "matrix"
|
||||||
|
|
||||||
def store(
|
|
||||||
self,
|
|
||||||
kind,
|
|
||||||
out: Dict[str, Any],
|
|
||||||
storage: "BaseFeatureStorage",
|
|
||||||
) -> None:
|
|
||||||
"""Store.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : {"BOLD"}
|
|
||||||
The data kind to store.
|
|
||||||
out : dict
|
|
||||||
The computed result as a dictionary to store.
|
|
||||||
storage : storage-like
|
|
||||||
The storage class, for example, SQLiteFeatureStorage.
|
|
||||||
|
|
||||||
"""
|
|
||||||
logger.debug(f"Storing BOLD-based marker in {storage}")
|
|
||||||
storage.store(kind="matrix", **out)
|
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
|
|
@ -16,9 +16,6 @@ from .base import BaseMarker
|
||||||
from .parcel_aggregation import ParcelAggregation
|
from .parcel_aggregation import ParcelAggregation
|
||||||
from .utils import _ets
|
from .utils import _ets
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from junifer.storage import BaseFeatureStorage
|
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class RSSETSMarker(BaseMarker):
|
class RSSETSMarker(BaseMarker):
|
||||||
|
|
@ -45,6 +42,8 @@ class RSSETSMarker(BaseMarker):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
parcellation: Union[str, List[str]],
|
parcellation: Union[str, List[str]],
|
||||||
|
|
@ -70,43 +69,21 @@ class RSSETSMarker(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["BOLD"]
|
return ["BOLD"]
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The input to the marker. The list must contain the
|
The data type input to the marker.
|
||||||
available Junifer Data dictionary keys.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The updated list of output kinds, as storage possibilities.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
return ["timeseries"]
|
return "timeseries"
|
||||||
|
|
||||||
def store(
|
|
||||||
self,
|
|
||||||
kind: str,
|
|
||||||
out: Dict[str, Any],
|
|
||||||
storage: "BaseFeatureStorage",
|
|
||||||
) -> None:
|
|
||||||
"""Store.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : {"BOLD"}
|
|
||||||
The data kind to store.
|
|
||||||
out : dict
|
|
||||||
The computed result as a dictionary to store.
|
|
||||||
storage : storage-like
|
|
||||||
The storage class, for example, SQLiteFeatureStorage.
|
|
||||||
|
|
||||||
"""
|
|
||||||
logger.debug(f"Storing BOLD in {storage}")
|
|
||||||
storage.store(kind="timeseries", **out)
|
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self,
|
||||||
|
|
@ -150,7 +127,7 @@ class RSSETSMarker(BaseMarker):
|
||||||
parcellation=self.parcellation,
|
parcellation=self.parcellation,
|
||||||
method=self.agg_method,
|
method=self.agg_method,
|
||||||
method_params=self.agg_method_params,
|
method_params=self.agg_method_params,
|
||||||
mask=self.mask
|
mask=self.mask,
|
||||||
)
|
)
|
||||||
# Compute the parcel aggregation
|
# Compute the parcel aggregation
|
||||||
out = parcel_aggregation.compute(input=input, extra_input=extra_input)
|
out = parcel_aggregation.compute(input=input, extra_input=extra_input)
|
||||||
|
|
|
||||||
|
|
@ -4,19 +4,15 @@
|
||||||
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
from sklearn.covariance import EmpiricalCovariance
|
from sklearn.covariance import EmpiricalCovariance
|
||||||
|
|
||||||
from ..api.decorators import register_marker
|
from ..api.decorators import register_marker
|
||||||
from ..utils import logger
|
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
from .parcel_aggregation import ParcelAggregation
|
from .parcel_aggregation import ParcelAggregation
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from junifer.storage import BaseFeatureStorage
|
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class FunctionalConnectivityParcels(BaseMarker):
|
class FunctionalConnectivityParcels(BaseMarker):
|
||||||
|
|
@ -49,6 +45,8 @@ class FunctionalConnectivityParcels(BaseMarker):
|
||||||
None).
|
None).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn", "scikit-learn"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
parcellation: Union[str, List[str]],
|
parcellation: Union[str, List[str]],
|
||||||
|
|
@ -83,22 +81,21 @@ class FunctionalConnectivityParcels(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["BOLD"]
|
return ["BOLD"]
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The input to the marker. The list must contain the
|
The data type input to the marker.
|
||||||
available Junifer Data dictionary keys.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The updated list of output kinds, as storage possibilities.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
outputs = ["matrix"]
|
return "matrix"
|
||||||
return outputs
|
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self,
|
||||||
|
|
@ -154,24 +151,3 @@ class FunctionalConnectivityParcels(BaseMarker):
|
||||||
out["col_names"] = ts["columns"]
|
out["col_names"] = ts["columns"]
|
||||||
out["matrix_kind"] = "tril"
|
out["matrix_kind"] = "tril"
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def store(
|
|
||||||
self,
|
|
||||||
kind: str,
|
|
||||||
out: Dict[str, Any],
|
|
||||||
storage: "BaseFeatureStorage",
|
|
||||||
) -> None:
|
|
||||||
"""Store.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : {"BOLD"}
|
|
||||||
The data kind to store.
|
|
||||||
out : dict
|
|
||||||
The computed result as a dictionary to store.
|
|
||||||
storage : storage-like
|
|
||||||
The storage class, for example, SQLiteFeatureStorage.
|
|
||||||
|
|
||||||
"""
|
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
|
||||||
storage.store(kind="matrix", **out)
|
|
||||||
|
|
|
||||||
|
|
@ -4,19 +4,16 @@
|
||||||
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
from sklearn.covariance import EmpiricalCovariance
|
from sklearn.covariance import EmpiricalCovariance
|
||||||
|
|
||||||
from ..api.decorators import register_marker
|
from ..api.decorators import register_marker
|
||||||
from ..utils import logger, raise_error
|
from ..utils import raise_error
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
from .sphere_aggregation import SphereAggregation
|
from .sphere_aggregation import SphereAggregation
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from junifer.storage import BaseFeatureStorage
|
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class FunctionalConnectivitySpheres(BaseMarker):
|
class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
|
|
@ -54,6 +51,8 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn", "scikit-learn"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
coords: str,
|
coords: str,
|
||||||
|
|
@ -93,23 +92,21 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["BOLD"]
|
return ["BOLD"]
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The input to the marker. The list must contain the
|
The data type input to the marker.
|
||||||
available Junifer Data dictionary keys.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The updated list of output kinds, as storage possibilities.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
outputs = ["matrix"]
|
return "matrix"
|
||||||
return outputs
|
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self,
|
||||||
|
|
@ -166,24 +163,3 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
out["col_names"] = ts["columns"]
|
out["col_names"] = ts["columns"]
|
||||||
out["matrix_kind"] = "tril"
|
out["matrix_kind"] = "tril"
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def store(
|
|
||||||
self,
|
|
||||||
kind: str,
|
|
||||||
out: Dict[str, Any],
|
|
||||||
storage: "BaseFeatureStorage",
|
|
||||||
) -> None:
|
|
||||||
"""Store.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : {"BOLD"}
|
|
||||||
The data kind to store.
|
|
||||||
out : dict
|
|
||||||
The computed result as a dictionary to store.
|
|
||||||
storage : storage-like
|
|
||||||
The storage class, for example, SQLiteFeatureStorage.
|
|
||||||
|
|
||||||
"""
|
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
|
||||||
storage.store(kind="matrix", **out)
|
|
||||||
|
|
|
||||||
|
|
@ -4,21 +4,18 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from nilearn.image import math_img, resample_to_img, new_img_like
|
from nilearn.image import math_img, new_img_like, resample_to_img
|
||||||
from nilearn.maskers import NiftiMasker
|
from nilearn.maskers import NiftiMasker
|
||||||
|
|
||||||
from ..api.decorators import register_marker
|
from ..api.decorators import register_marker
|
||||||
from ..data import load_parcellation, load_mask
|
from ..data import load_mask, load_parcellation
|
||||||
from ..stats import get_aggfunc_by_name
|
from ..stats import get_aggfunc_by_name
|
||||||
from ..utils import logger
|
from ..utils import logger
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from junifer.storage import BaseFeatureStorage
|
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class ParcelAggregation(BaseMarker):
|
class ParcelAggregation(BaseMarker):
|
||||||
|
|
@ -48,6 +45,8 @@ class ParcelAggregation(BaseMarker):
|
||||||
None).
|
None).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn", "numpy"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
parcellation: Union[str, List[str]],
|
parcellation: Union[str, List[str]],
|
||||||
|
|
@ -76,53 +75,27 @@ class ParcelAggregation(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]
|
return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The kind of data to work on.
|
The data type input to the marker.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The list of storage kinds.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
outputs = []
|
|
||||||
for t_input in input:
|
|
||||||
if t_input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
|
||||||
outputs.append("table")
|
|
||||||
elif t_input in ["BOLD"]:
|
|
||||||
outputs.append("timeseries")
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown input kind for {t_input}")
|
|
||||||
return outputs
|
|
||||||
|
|
||||||
def store(
|
if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
||||||
self,
|
return "table"
|
||||||
kind: str,
|
elif input_type == "BOLD":
|
||||||
out: Dict[str, Any],
|
return "timeseries"
|
||||||
storage: "BaseFeatureStorage",
|
else:
|
||||||
) -> None:
|
raise ValueError(f"Unknown input kind for {input_type}")
|
||||||
"""Store.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
|
|
||||||
The data kind to store.
|
|
||||||
out : dict
|
|
||||||
The computed result as a dictionary to store.
|
|
||||||
storage : storage-like
|
|
||||||
The storage class, for example, SQLiteFeatureStorage.
|
|
||||||
|
|
||||||
"""
|
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
|
||||||
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
|
||||||
storage.store(kind="table", **out)
|
|
||||||
elif kind in ["BOLD"]:
|
|
||||||
storage.store(kind="timeseries", **out)
|
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self,
|
||||||
|
|
@ -252,6 +225,4 @@ class ParcelAggregation(BaseMarker):
|
||||||
|
|
||||||
out_values = np.array(out_values).T
|
out_values = np.array(out_values).T
|
||||||
out = {"data": out_values, "columns": out_labels}
|
out = {"data": out_values, "columns": out_labels}
|
||||||
if out_values.shape[0] > 1:
|
|
||||||
out["row_names"] = "scan"
|
|
||||||
return out
|
return out
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from ..api.decorators import register_marker
|
from ..api.decorators import register_marker
|
||||||
from ..data import load_coordinates, load_mask
|
from ..data import load_coordinates, load_mask
|
||||||
|
|
@ -13,9 +13,6 @@ from ..stats import get_aggfunc_by_name
|
||||||
from ..utils import logger
|
from ..utils import logger
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from junifer.storage import BaseFeatureStorage
|
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class SphereAggregation(BaseMarker):
|
class SphereAggregation(BaseMarker):
|
||||||
|
|
@ -51,6 +48,8 @@ class SphereAggregation(BaseMarker):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"nilearn", "numpy"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
coords: str,
|
coords: str,
|
||||||
|
|
@ -79,53 +78,27 @@ class SphereAggregation(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]
|
return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The kind of data to work on.
|
The data type input to the marker.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The list of storage kinds.
|
The storage type output by the marker.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
outputs = []
|
|
||||||
for t_input in input:
|
|
||||||
if t_input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
|
||||||
outputs.append("table")
|
|
||||||
elif t_input in ["BOLD"]:
|
|
||||||
outputs.append("timeseries")
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown input kind for {t_input}")
|
|
||||||
return outputs
|
|
||||||
|
|
||||||
def store(
|
if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
||||||
self,
|
return "table"
|
||||||
kind: str,
|
elif input_type == "BOLD":
|
||||||
out: Dict[str, Any],
|
return "timeseries"
|
||||||
storage: "BaseFeatureStorage",
|
else:
|
||||||
) -> None:
|
raise ValueError(f"Unknown input kind for {input_type}")
|
||||||
"""Store.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
|
|
||||||
The data kind to store.
|
|
||||||
out : dict
|
|
||||||
The computed result as a dictionary to store.
|
|
||||||
storage : storage-like
|
|
||||||
The storage class, for example, SQLiteFeatureStorage.
|
|
||||||
|
|
||||||
"""
|
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
|
||||||
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
|
||||||
storage.store(kind="table", **out)
|
|
||||||
elif kind in ["BOLD"]:
|
|
||||||
storage.store(kind="timeseries", **out)
|
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self,
|
||||||
|
|
@ -180,6 +153,4 @@ class SphereAggregation(BaseMarker):
|
||||||
out_values = masker.fit_transform(t_input)
|
out_values = masker.fit_transform(t_input)
|
||||||
# Format the output
|
# Format the output
|
||||||
out = {"data": out_values, "columns": out_labels}
|
out = {"data": out_values, "columns": out_labels}
|
||||||
if out_values.shape[0] > 1:
|
|
||||||
out["row_names"] = "scan"
|
|
||||||
return out
|
return out
|
||||||
|
|
|
||||||
|
|
@ -15,6 +15,7 @@ from junifer.markers.crossparcellation_functional_connectivity import (
|
||||||
from junifer.storage import SQLiteFeatureStorage
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
|
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
|
||||||
|
|
||||||
|
|
||||||
parcellation_ONE = "Schaefer100x17"
|
parcellation_ONE = "Schaefer100x17"
|
||||||
parcellation_TWO = "Schaefer200x17"
|
parcellation_TWO = "Schaefer200x17"
|
||||||
|
|
||||||
|
|
@ -30,8 +31,7 @@ def test_compute() -> None:
|
||||||
"data": niimg,
|
"data": niimg,
|
||||||
"path": out["BOLD"]["path"],
|
"path": out["BOLD"]["path"],
|
||||||
"meta": {"element": "sub001"},
|
"meta": {"element": "sub001"},
|
||||||
},
|
}
|
||||||
"meta": {"element": "sub001"},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
crossparcellation = CrossParcellationFC(
|
crossparcellation = CrossParcellationFC(
|
||||||
|
|
@ -43,12 +43,6 @@ def test_compute() -> None:
|
||||||
assert out["data"].shape == (200, 100)
|
assert out["data"].shape == (200, 100)
|
||||||
assert len(out["col_names"]) == 100
|
assert len(out["col_names"]) == 100
|
||||||
assert len(out["row_names"]) == 200
|
assert len(out["row_names"]) == 200
|
||||||
meta = crossparcellation.get_meta("BOLD")["marker"]
|
|
||||||
assert meta["aggregation_method"] == "mean"
|
|
||||||
assert meta["class"] == "CrossParcellationFC"
|
|
||||||
assert meta["parcellation_one"] == "Schaefer100x17"
|
|
||||||
assert meta["parcellation_two"] == "Schaefer200x17"
|
|
||||||
assert meta["correlation_method"] == "spearman"
|
|
||||||
|
|
||||||
|
|
||||||
def test_store(tmp_path: Path) -> None:
|
def test_store(tmp_path: Path) -> None:
|
||||||
|
|
@ -62,16 +56,10 @@ def test_store(tmp_path: Path) -> None:
|
||||||
"""
|
"""
|
||||||
|
|
||||||
with SPMAuditoryTestingDatagrabber() as dg:
|
with SPMAuditoryTestingDatagrabber() as dg:
|
||||||
out = dg["sub001"]
|
input_dict = dg["sub001"]
|
||||||
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
|
niimg = image.load_img(str(input_dict["BOLD"]["path"].absolute()))
|
||||||
input_dict = {
|
|
||||||
"BOLD": {
|
input_dict["BOLD"]["data"] = niimg
|
||||||
"data": niimg,
|
|
||||||
"path": out["BOLD"]["path"],
|
|
||||||
"meta": {"element": "sub001"},
|
|
||||||
},
|
|
||||||
"meta": {"element": "sub001"},
|
|
||||||
}
|
|
||||||
|
|
||||||
crossparcellation = CrossParcellationFC(
|
crossparcellation = CrossParcellationFC(
|
||||||
parcellation_one=parcellation_ONE,
|
parcellation_one=parcellation_ONE,
|
||||||
|
|
@ -80,19 +68,23 @@ def test_store(tmp_path: Path) -> None:
|
||||||
)
|
)
|
||||||
uri = tmp_path / "test_crossparcellation.sqlite"
|
uri = tmp_path / "test_crossparcellation.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
out = crossparcellation.fit_transform(input_dict, storage=storage)
|
crossparcellation.fit_transform(input_dict, storage=storage)
|
||||||
|
features = storage.list_features()
|
||||||
|
assert any(
|
||||||
|
x["name"] == "BOLD_CrossParcellationFC"
|
||||||
|
for x in features.values()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_get_output_kind() -> None:
|
def test_get_output_type() -> None:
|
||||||
"""Test CrossParcellationFC get_output_kind()."""
|
"""Test CrossParcellationFC get_output_type()."""
|
||||||
|
|
||||||
crossparcellation = CrossParcellationFC(
|
crossparcellation = CrossParcellationFC(
|
||||||
parcellation_one=parcellation_ONE, parcellation_two=parcellation_TWO
|
parcellation_one=parcellation_ONE, parcellation_two=parcellation_TWO
|
||||||
)
|
)
|
||||||
input_list = ["BOLD"]
|
input_ = "BOLD"
|
||||||
input_list = crossparcellation.get_output_kind(input_list)
|
output = crossparcellation.get_output_type(input_)
|
||||||
assert len(input_list) == 1
|
assert output == "matrix"
|
||||||
assert input_list[0] in ["matrix"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_init_() -> None:
|
def test_init_() -> None:
|
||||||
|
|
|
||||||
|
|
@ -16,6 +16,7 @@ from junifer.markers.ets_rss import RSSETSMarker
|
||||||
from junifer.storage import SQLiteFeatureStorage
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
|
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
|
||||||
|
|
||||||
|
|
||||||
# Set parcellation
|
# Set parcellation
|
||||||
PARCELLATION = "Schaefer100x17"
|
PARCELLATION = "Schaefer100x17"
|
||||||
|
|
||||||
|
|
@ -42,21 +43,13 @@ def test_compute() -> None:
|
||||||
n_time, _ = test_ts.shape
|
n_time, _ = test_ts.shape
|
||||||
assert n_time == len(new_out["data"])
|
assert n_time == len(new_out["data"])
|
||||||
|
|
||||||
# Assert the meta
|
|
||||||
meta = ets_rss_marker.get_meta("BOLD")["marker"]
|
|
||||||
assert meta["parcellation"] == "Schaefer100x17"
|
|
||||||
assert meta["agg_method"] == "mean"
|
|
||||||
assert meta["agg_method_params"] is None
|
|
||||||
assert meta["class"] == "RSSETSMarker"
|
|
||||||
|
|
||||||
|
def test_get_output_type() -> None:
|
||||||
def test_get_output_kind() -> None:
|
"""Test RSS ETS get_output_type()."""
|
||||||
"""Test RSS ETS get_output_kind()."""
|
|
||||||
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION)
|
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION)
|
||||||
input_list = ["BOLD"]
|
input_ = "BOLD"
|
||||||
input_list = ets_rss_marker.get_output_kind(input_list)
|
output = ets_rss_marker.get_output_type(input_)
|
||||||
assert len(input_list) == 1
|
assert output == "timeseries"
|
||||||
assert input_list[0] in ["timeseries"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_store(tmp_path: Path) -> None:
|
def test_store(tmp_path: Path) -> None:
|
||||||
|
|
@ -70,14 +63,18 @@ def test_store(tmp_path: Path) -> None:
|
||||||
"""
|
"""
|
||||||
with SPMAuditoryTestingDatagrabber() as dg:
|
with SPMAuditoryTestingDatagrabber() as dg:
|
||||||
# Fetch element
|
# Fetch element
|
||||||
out = dg["sub001"]
|
elem = dg["sub001"]
|
||||||
# Load BOLD image
|
# Load BOLD image
|
||||||
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
|
niimg = image.load_img(str(elem["BOLD"]["path"].absolute()))
|
||||||
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
|
elem["BOLD"]["data"] = niimg
|
||||||
# Compute the RSSETSMarker
|
# Compute the RSSETSMarker
|
||||||
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION)
|
ets_rss_marker = RSSETSMarker(parcellation=PARCELLATION)
|
||||||
# Create storage
|
# Create storage
|
||||||
storage = SQLiteFeatureStorage(
|
storage = SQLiteFeatureStorage(
|
||||||
uri=str((tmp_path / "test.sqlite").absolute()))
|
uri=str((tmp_path / "test.sqlite").absolute())
|
||||||
|
)
|
||||||
# Store
|
# Store
|
||||||
ets_rss_marker.fit_transform(input=input_dict, storage=storage)
|
ets_rss_marker.fit_transform(input=elem, storage=storage)
|
||||||
|
|
||||||
|
features = storage.list_features()
|
||||||
|
assert any(x["name"] == "BOLD_RSSETSMarker" for x in features.values())
|
||||||
|
|
|
||||||
|
|
@ -32,7 +32,7 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
|
||||||
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
|
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
|
||||||
|
|
||||||
fc = FunctionalConnectivityParcels(parcellation="Schaefer100x7")
|
fc = FunctionalConnectivityParcels(parcellation="Schaefer100x7")
|
||||||
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}})
|
all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": {}}})
|
||||||
|
|
||||||
out = all_out["BOLD"]
|
out = all_out["BOLD"]
|
||||||
|
|
||||||
|
|
@ -48,7 +48,11 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
|
||||||
pa = ParcelAggregation(
|
pa = ParcelAggregation(
|
||||||
parcellation="Schaefer100x7", method="mean", on="BOLD"
|
parcellation="Schaefer100x7", method="mean", on="BOLD"
|
||||||
)
|
)
|
||||||
ts = pa.compute({"data": fmri_img})
|
meta = {
|
||||||
|
"element": {"subject": "sub001"},
|
||||||
|
"dependencies": {"nilearn"},
|
||||||
|
}
|
||||||
|
ts = pa.compute({"data": fmri_img, "meta": meta})
|
||||||
|
|
||||||
# compare with nilearn
|
# compare with nilearn
|
||||||
# Get the testing parcellation (for nilearn)
|
# Get the testing parcellation (for nilearn)
|
||||||
|
|
@ -69,22 +73,24 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
|
||||||
assert_array_almost_equal(out_ni, out["data"], decimal=3)
|
assert_array_almost_equal(out_ni, out["data"], decimal=3)
|
||||||
|
|
||||||
# check correct output
|
# check correct output
|
||||||
assert fc.get_output_kind(["BOLD"]) == ["matrix"]
|
assert fc.get_output_type("BOLD") == "matrix"
|
||||||
|
|
||||||
# Check empirical correlation method parameters
|
# Check empirical correlation method parameters
|
||||||
fc = FunctionalConnectivityParcels(
|
fc = FunctionalConnectivityParcels(
|
||||||
parcellation="Schaefer100x7", cor_method_params={"empirical": True}
|
parcellation="Schaefer100x7", cor_method_params={"empirical": True}
|
||||||
)
|
)
|
||||||
|
|
||||||
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}})
|
all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": meta}})
|
||||||
|
|
||||||
uri = tmp_path / "test_fc_parcellation.sqlite"
|
uri = tmp_path / "test_fc_parcellation.sqlite"
|
||||||
# Single storage, must be the uri
|
# Single storage, must be the uri
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
meta = {
|
meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}}
|
||||||
"element": "test",
|
input = {"BOLD": {"data": fmri_img, "meta": meta}}
|
||||||
"version": "0.0.1",
|
|
||||||
"marker": {"name": "fcname"},
|
|
||||||
}
|
|
||||||
input = {"BOLD": {"data": fmri_img}, "meta": meta}
|
|
||||||
all_out = fc.fit_transform(input, storage=storage)
|
all_out = fc.fit_transform(input, storage=storage)
|
||||||
|
|
||||||
|
features = storage.list_features()
|
||||||
|
assert any(
|
||||||
|
x["name"] == "BOLD_FunctionalConnectivityParcels"
|
||||||
|
for x in features.values()
|
||||||
|
)
|
||||||
|
|
|
||||||
|
|
@ -36,7 +36,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
||||||
fc = FunctionalConnectivitySpheres(
|
fc = FunctionalConnectivitySpheres(
|
||||||
coords="DMNBuckner", radius=5.0, cor_method="correlation"
|
coords="DMNBuckner", radius=5.0, cor_method="correlation"
|
||||||
)
|
)
|
||||||
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}})
|
all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": {}}})
|
||||||
|
|
||||||
out = all_out["BOLD"]
|
out = all_out["BOLD"]
|
||||||
|
|
||||||
|
|
@ -52,7 +52,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
||||||
sa = SphereAggregation(
|
sa = SphereAggregation(
|
||||||
coords="DMNBuckner", radius=5.0, method="mean", on="BOLD"
|
coords="DMNBuckner", radius=5.0, method="mean", on="BOLD"
|
||||||
)
|
)
|
||||||
ts = sa.compute({"data": fmri_img})
|
ts = sa.compute({"data": fmri_img, "meta": {}})
|
||||||
|
|
||||||
# Check that FC are almost equal when using nileran
|
# Check that FC are almost equal when using nileran
|
||||||
cm = ConnectivityMeasure(kind="correlation")
|
cm = ConnectivityMeasure(kind="correlation")
|
||||||
|
|
@ -60,19 +60,24 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
||||||
assert_array_almost_equal(out_ni, out["data"], decimal=3)
|
assert_array_almost_equal(out_ni, out["data"], decimal=3)
|
||||||
|
|
||||||
# check correct output
|
# check correct output
|
||||||
assert fc.get_output_kind(["BOLD"]) == ["matrix"]
|
assert fc.get_output_type("BOLD") == "matrix"
|
||||||
|
|
||||||
uri = tmp_path / "test_fc_parcel.sqlite"
|
uri = tmp_path / "test_fc_parcel.sqlite"
|
||||||
# Single storage, must be the uri
|
# Single storage, must be the uri
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
meta = {
|
meta = {
|
||||||
"element": "test",
|
"element": {"subject": "test"},
|
||||||
"version": "0.0.1",
|
"dependencies": {"numpy", "nilearn"},
|
||||||
"marker": {"name": "fcname"},
|
|
||||||
}
|
}
|
||||||
input = {"BOLD": {"data": fmri_img}, "meta": meta}
|
input = {"BOLD": {"data": fmri_img, "meta": meta}}
|
||||||
all_out = fc.fit_transform(input, storage=storage)
|
all_out = fc.fit_transform(input, storage=storage)
|
||||||
|
|
||||||
|
features = storage.list_features()
|
||||||
|
assert any(
|
||||||
|
x["name"] == "BOLD_FunctionalConnectivitySpheres"
|
||||||
|
for x in features.values()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
|
def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
|
||||||
"""Test FunctionalConnectivitySpheres with empirical covariance.
|
"""Test FunctionalConnectivitySpheres with empirical covariance.
|
||||||
|
|
@ -94,7 +99,7 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
|
||||||
cor_method="correlation",
|
cor_method="correlation",
|
||||||
cor_method_params={"empirical": True},
|
cor_method_params={"empirical": True},
|
||||||
)
|
)
|
||||||
all_out = fc.fit_transform({"BOLD": {"data": fmri_img}})
|
all_out = fc.fit_transform({"BOLD": {"data": fmri_img, "meta": {}}})
|
||||||
|
|
||||||
out = all_out["BOLD"]
|
out = all_out["BOLD"]
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -19,10 +19,14 @@ def test_base_marker_subclassing() -> None:
|
||||||
"""Test proper subclassing of BaseMarker."""
|
"""Test proper subclassing of BaseMarker."""
|
||||||
# Create concrete class
|
# Create concrete class
|
||||||
class MyBaseMarker(BaseMarker):
|
class MyBaseMarker(BaseMarker):
|
||||||
|
def __init__(self, on, name=None) -> None:
|
||||||
|
self.parameter = 1
|
||||||
|
super().__init__(on, name)
|
||||||
|
|
||||||
def get_valid_inputs(self):
|
def get_valid_inputs(self):
|
||||||
return ["BOLD", "T1w"]
|
return ["BOLD", "T1w"]
|
||||||
|
|
||||||
def get_output_kind(self, input):
|
def get_output_type(self, input):
|
||||||
return ["timeseries"]
|
return ["timeseries"]
|
||||||
|
|
||||||
def compute(self, input, extra_input):
|
def compute(self, input, extra_input):
|
||||||
|
|
@ -32,18 +36,19 @@ def test_base_marker_subclassing() -> None:
|
||||||
"row_names": "row_names",
|
"row_names": "row_names",
|
||||||
}
|
}
|
||||||
|
|
||||||
def store(self, kind, out, storage):
|
|
||||||
return super().store(kind=kind, out=out, storage=storage)
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"):
|
with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"):
|
||||||
MyBaseMarker(on=["BOLD", "T2w"])
|
MyBaseMarker(on=["BOLD", "T2w"])
|
||||||
|
|
||||||
# Create input for marker
|
# Create input for marker
|
||||||
input_ = {
|
input_ = {
|
||||||
"meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"},
|
|
||||||
"BOLD": {
|
"BOLD": {
|
||||||
"path": ".",
|
"path": ".",
|
||||||
"data": "data",
|
"data": "data",
|
||||||
|
"meta": {
|
||||||
|
"datagrabber": "dg",
|
||||||
|
"element": "elem",
|
||||||
|
"datareader": "dr",
|
||||||
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
marker = MyBaseMarker(on=["BOLD"])
|
marker = MyBaseMarker(on=["BOLD"])
|
||||||
|
|
@ -57,14 +62,16 @@ def test_base_marker_subclassing() -> None:
|
||||||
assert "data" in output["BOLD"]
|
assert "data" in output["BOLD"]
|
||||||
assert "columns" in output["BOLD"]
|
assert "columns" in output["BOLD"]
|
||||||
assert "row_names" in output["BOLD"]
|
assert "row_names" in output["BOLD"]
|
||||||
assert "meta" in output["BOLD"]
|
|
||||||
assert "datagrabber" in output["BOLD"]["meta"]
|
|
||||||
assert "element" in output["BOLD"]["meta"]
|
|
||||||
assert "datareader" in output["BOLD"]["meta"]
|
|
||||||
|
|
||||||
# Check no implementation check
|
assert "meta" in output["BOLD"]
|
||||||
with pytest.raises(NotImplementedError):
|
meta = output["BOLD"]["meta"]
|
||||||
marker.store(kind="kind", out="out", storage="storage") # type: ignore
|
assert "datagrabber" in meta
|
||||||
|
assert "element" in meta
|
||||||
|
assert "datareader" in meta
|
||||||
|
assert "marker" in meta
|
||||||
|
assert "name" in meta["marker"]
|
||||||
|
assert "parameter" in meta["marker"]
|
||||||
|
assert meta["marker"]["parameter"] == 1
|
||||||
|
|
||||||
# Check attributes
|
# Check attributes
|
||||||
assert marker.name == "MyBaseMarker"
|
assert marker.name == "MyBaseMarker"
|
||||||
|
|
|
||||||
|
|
@ -3,18 +3,21 @@
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
from pathlib import Path
|
|
||||||
from nilearn import datasets
|
from nilearn import datasets
|
||||||
from nilearn.image import concat_imgs, math_img, resample_to_img, new_img_like
|
from nilearn.image import concat_imgs, math_img, new_img_like, resample_to_img
|
||||||
from nilearn.maskers import NiftiLabelsMasker, NiftiMasker
|
from nilearn.maskers import NiftiLabelsMasker, NiftiMasker
|
||||||
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
||||||
from scipy.stats import trim_mean
|
from scipy.stats import trim_mean
|
||||||
|
|
||||||
from junifer.data import load_mask, load_parcellation, register_parcellation
|
from junifer.data import load_mask, load_parcellation, register_parcellation
|
||||||
from junifer.markers.parcel_aggregation import ParcelAggregation
|
from junifer.markers.parcel_aggregation import ParcelAggregation
|
||||||
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
def test_ParcelAggregation_input_output() -> None:
|
def test_ParcelAggregation_input_output() -> None:
|
||||||
|
|
@ -22,12 +25,11 @@ def test_ParcelAggregation_input_output() -> None:
|
||||||
marker = ParcelAggregation(
|
marker = ParcelAggregation(
|
||||||
parcellation="Schaefer100x7", method="mean", on="VBM_GM"
|
parcellation="Schaefer100x7", method="mean", on="VBM_GM"
|
||||||
)
|
)
|
||||||
|
for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]:
|
||||||
output = marker.get_output_kind(["VBM_GM", "BOLD"])
|
assert marker.get_output_type(in_) == out_
|
||||||
assert output == ["table", "timeseries"]
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="Unknown input"):
|
with pytest.raises(ValueError, match="Unknown input"):
|
||||||
marker.get_output_kind(["VBM_GM", "BOLD", "unknown"])
|
marker.get_output_type("unknown")
|
||||||
|
|
||||||
|
|
||||||
def test_ParcelAggregation_3D() -> None:
|
def test_ParcelAggregation_3D() -> None:
|
||||||
|
|
@ -78,22 +80,13 @@ def test_ParcelAggregation_3D() -> None:
|
||||||
name="gmd_schaefer100x7_mean",
|
name="gmd_schaefer100x7_mean",
|
||||||
on="VBM_GM",
|
on="VBM_GM",
|
||||||
) # Test passing "on" as a keyword argument
|
) # Test passing "on" as a keyword argument
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
|
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
|
||||||
|
|
||||||
assert jun_values3d_mean.ndim == 2
|
assert jun_values3d_mean.ndim == 2
|
||||||
assert jun_values3d_mean.shape[0] == 1
|
assert jun_values3d_mean.shape[0] == 1
|
||||||
assert_array_equal(manual, jun_values3d_mean)
|
assert_array_equal(manual, jun_values3d_mean)
|
||||||
|
|
||||||
meta = marker.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
# Test using another function (std)
|
# Test using another function (std)
|
||||||
manual = []
|
manual = []
|
||||||
for t_v in sorted(np.unique(parcellation_values)):
|
for t_v in sorted(np.unique(parcellation_values)):
|
||||||
|
|
@ -103,22 +96,13 @@ def test_ParcelAggregation_3D() -> None:
|
||||||
|
|
||||||
# Use the ParcelAggregation object
|
# Use the ParcelAggregation object
|
||||||
marker = ParcelAggregation(parcellation="Schaefer100x7", method="std")
|
marker = ParcelAggregation(parcellation="Schaefer100x7", method="std")
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"]
|
jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"]
|
||||||
|
|
||||||
assert jun_values3d_std.ndim == 2
|
assert jun_values3d_std.ndim == 2
|
||||||
assert jun_values3d_std.shape[0] == 1
|
assert jun_values3d_std.shape[0] == 1
|
||||||
assert_array_equal(manual, jun_values3d_std)
|
assert_array_equal(manual, jun_values3d_std)
|
||||||
|
|
||||||
meta = marker.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "std"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_ParcelAggregation"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
# Test using another function with parameters
|
# Test using another function with parameters
|
||||||
manual = []
|
manual = []
|
||||||
for t_v in sorted(np.unique(parcellation_values)):
|
for t_v in sorted(np.unique(parcellation_values)):
|
||||||
|
|
@ -136,22 +120,13 @@ def test_ParcelAggregation_3D() -> None:
|
||||||
method="trim_mean",
|
method="trim_mean",
|
||||||
method_params={"proportiontocut": 0.1},
|
method_params={"proportiontocut": 0.1},
|
||||||
)
|
)
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"]
|
jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"]
|
||||||
|
|
||||||
assert jun_values3d_tm.ndim == 2
|
assert jun_values3d_tm.ndim == 2
|
||||||
assert jun_values3d_tm.shape[0] == 1
|
assert jun_values3d_tm.shape[0] == 1
|
||||||
assert_array_equal(manual, jun_values3d_tm)
|
assert_array_equal(manual, jun_values3d_tm)
|
||||||
|
|
||||||
meta = marker.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "trim_mean"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_ParcelAggregation"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {"proportiontocut": 0.1}
|
|
||||||
|
|
||||||
|
|
||||||
def test_ParcelAggregation_4D():
|
def test_ParcelAggregation_4D():
|
||||||
"""Test ParcelAggregation object on 4D images."""
|
"""Test ParcelAggregation object on 4D images."""
|
||||||
|
|
@ -170,21 +145,63 @@ def test_ParcelAggregation_4D():
|
||||||
|
|
||||||
# Create ParcelAggregation object
|
# Create ParcelAggregation object
|
||||||
marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
|
marker = ParcelAggregation(parcellation="Schaefer100x7", method="mean")
|
||||||
input = dict(BOLD=dict(data=fmri_img))
|
input = {"BOLD": {"data": fmri_img, "meta": {}}}
|
||||||
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
|
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
|
||||||
|
|
||||||
assert jun_values4d.ndim == 2
|
assert jun_values4d.ndim == 2
|
||||||
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
||||||
assert_array_equal(auto4d, jun_values4d)
|
assert_array_equal(auto4d, jun_values4d)
|
||||||
|
|
||||||
meta = marker.get_meta("BOLD")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
def test_ParcelAggregation_storage(tmp_path: Path) -> None:
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
"""Test ParcelAggregation storage.
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "BOLD_ParcelAggregation"
|
Parameters
|
||||||
assert meta["class"] == "ParcelAggregation"
|
----------
|
||||||
assert meta["kind"] == "BOLD"
|
tmp_path : pathlib.Path
|
||||||
assert meta["method_params"] == {}
|
The path to the test directory.
|
||||||
|
|
||||||
|
"""
|
||||||
|
# Get the oasis VBM data
|
||||||
|
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
|
||||||
|
vbm = oasis_dataset.gray_matter_maps[0]
|
||||||
|
img = nib.load(vbm)
|
||||||
|
uri = tmp_path / "test_sphere_storage_3D.sqlite"
|
||||||
|
|
||||||
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
|
meta = {
|
||||||
|
"element": {"subject": "sub-01", "session": "ses-01"},
|
||||||
|
"dependencies": {"nilearn", "nibabel"},
|
||||||
|
}
|
||||||
|
input = {"VBM_GM": {"data": img, "meta": meta}}
|
||||||
|
marker = ParcelAggregation(
|
||||||
|
parcellation="Schaefer100x7", method="mean", on="VBM_GM"
|
||||||
|
)
|
||||||
|
|
||||||
|
marker.fit_transform(input, storage=storage)
|
||||||
|
|
||||||
|
features = storage.list_features()
|
||||||
|
assert any(
|
||||||
|
x["name"] == "VBM_GM_ParcelAggregation" for x in features.values()
|
||||||
|
)
|
||||||
|
|
||||||
|
meta = {
|
||||||
|
"element": {"subject": "sub-01", "session": "ses-01"},
|
||||||
|
"dependencies": {"nilearn", "nibabel"},
|
||||||
|
}
|
||||||
|
# Get the SPM auditory data
|
||||||
|
subject_data = datasets.fetch_spm_auditory()
|
||||||
|
fmri_img = concat_imgs(subject_data.func) # type: ignore
|
||||||
|
input = {"BOLD": {"data": fmri_img, "meta": meta}}
|
||||||
|
marker = ParcelAggregation(
|
||||||
|
parcellation="Schaefer100x7", method="mean", on="BOLD"
|
||||||
|
)
|
||||||
|
|
||||||
|
marker.fit_transform(input, storage=storage)
|
||||||
|
features = storage.list_features()
|
||||||
|
assert any(
|
||||||
|
x["name"] == "BOLD_ParcelAggregation" for x in features.values()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_ParcelAggregation_3D_mask() -> None:
|
def test_ParcelAggregation_3D_mask() -> None:
|
||||||
|
|
@ -215,22 +232,13 @@ def test_ParcelAggregation_3D_mask() -> None:
|
||||||
name="gmd_schaefer100x7_mean",
|
name="gmd_schaefer100x7_mean",
|
||||||
on="VBM_GM",
|
on="VBM_GM",
|
||||||
) # Test passing "on" as a keyword argument
|
) # Test passing "on" as a keyword argument
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
|
jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"]
|
||||||
|
|
||||||
assert jun_values3d_mean.ndim == 2
|
assert jun_values3d_mean.ndim == 2
|
||||||
assert jun_values3d_mean.shape[0] == 1
|
assert jun_values3d_mean.shape[0] == 1
|
||||||
assert_array_almost_equal(auto, jun_values3d_mean)
|
assert_array_almost_equal(auto, jun_values3d_mean)
|
||||||
|
|
||||||
meta = marker.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
|
||||||
assert meta["mask"] == "GM_prob0.2"
|
|
||||||
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
|
|
||||||
def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
|
def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
|
||||||
"""Test ParcelAggregation with multiple non-overlapping parcellations.
|
"""Test ParcelAggregation with multiple non-overlapping parcellations.
|
||||||
|
|
@ -281,7 +289,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
|
||||||
name="gmd_schaefer100x7_mean",
|
name="gmd_schaefer100x7_mean",
|
||||||
on="VBM_GM",
|
on="VBM_GM",
|
||||||
) # Test passing "on" as a keyword argument
|
) # Test passing "on" as a keyword argument
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
|
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
|
||||||
|
|
||||||
orig_mean_data = orig_mean["data"]
|
orig_mean_data = orig_mean["data"]
|
||||||
|
|
@ -290,15 +298,6 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
|
||||||
assert orig_mean_data.shape[1] == 100
|
assert orig_mean_data.shape[1] == 100
|
||||||
# assert_array_almost_equal(auto, jun_values3d_mean)
|
# assert_array_almost_equal(auto, jun_values3d_mean)
|
||||||
|
|
||||||
meta = marker_original.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
# Use the ParcelAggregation object on the two parcellations
|
# Use the ParcelAggregation object on the two parcellations
|
||||||
marker_split = ParcelAggregation(
|
marker_split = ParcelAggregation(
|
||||||
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"],
|
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"],
|
||||||
|
|
@ -306,7 +305,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
|
||||||
name="gmd_schaefer100x7_mean",
|
name="gmd_schaefer100x7_mean",
|
||||||
on="VBM_GM",
|
on="VBM_GM",
|
||||||
) # Test passing "on" as a keyword argument
|
) # Test passing "on" as a keyword argument
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
split_mean = marker_split.fit_transform(input)["VBM_GM"]
|
split_mean = marker_split.fit_transform(input)["VBM_GM"]
|
||||||
split_mean_data = split_mean["data"]
|
split_mean_data = split_mean["data"]
|
||||||
|
|
||||||
|
|
@ -314,15 +313,6 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
|
||||||
assert split_mean_data.shape[0] == 1
|
assert split_mean_data.shape[0] == 1
|
||||||
assert split_mean_data.shape[1] == 100
|
assert split_mean_data.shape[1] == 100
|
||||||
|
|
||||||
meta = marker_split.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7_low", "Schaefer100x7_high"]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
# Data and labels should be the same
|
# Data and labels should be the same
|
||||||
assert_array_equal(orig_mean_data, split_mean_data)
|
assert_array_equal(orig_mean_data, split_mean_data)
|
||||||
assert orig_mean["columns"] == split_mean["columns"]
|
assert orig_mean["columns"] == split_mean["columns"]
|
||||||
|
|
@ -379,7 +369,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
|
||||||
name="gmd_schaefer100x7_mean",
|
name="gmd_schaefer100x7_mean",
|
||||||
on="VBM_GM",
|
on="VBM_GM",
|
||||||
) # Test passing "on" as a keyword argument
|
) # Test passing "on" as a keyword argument
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
|
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
|
||||||
|
|
||||||
orig_mean_data = orig_mean["data"]
|
orig_mean_data = orig_mean["data"]
|
||||||
|
|
@ -388,15 +378,6 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
|
||||||
assert orig_mean_data.shape[1] == 100
|
assert orig_mean_data.shape[1] == 100
|
||||||
# assert_array_almost_equal(auto, jun_values3d_mean)
|
# assert_array_almost_equal(auto, jun_values3d_mean)
|
||||||
|
|
||||||
meta = marker_original.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["parcellation"] == ["Schaefer100x7"]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
# Use the ParcelAggregation object on the two parcellations
|
# Use the ParcelAggregation object on the two parcellations
|
||||||
marker_split = ParcelAggregation(
|
marker_split = ParcelAggregation(
|
||||||
parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"],
|
parcellation=["Schaefer100x7_low2", "Schaefer100x7_high2"],
|
||||||
|
|
@ -404,7 +385,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
|
||||||
name="gmd_schaefer100x7_mean",
|
name="gmd_schaefer100x7_mean",
|
||||||
on="VBM_GM",
|
on="VBM_GM",
|
||||||
) # Test passing "on" as a keyword argument
|
) # Test passing "on" as a keyword argument
|
||||||
input = dict(VBM_GM=dict(data=img))
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
split_mean = marker_split.fit_transform(input)["VBM_GM"]
|
split_mean = marker_split.fit_transform(input)["VBM_GM"]
|
||||||
split_mean_data = split_mean["data"]
|
split_mean_data = split_mean["data"]
|
||||||
|
|
||||||
|
|
@ -412,18 +393,6 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
|
||||||
assert split_mean_data.shape[0] == 1
|
assert split_mean_data.shape[0] == 1
|
||||||
assert split_mean_data.shape[1] == 100
|
assert split_mean_data.shape[1] == 100
|
||||||
|
|
||||||
meta = marker_split.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["parcellation"] == [
|
|
||||||
"Schaefer100x7_low2",
|
|
||||||
"Schaefer100x7_high2",
|
|
||||||
]
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean"
|
|
||||||
assert meta["class"] == "ParcelAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
# Data should be the same
|
# Data should be the same
|
||||||
assert_array_equal(orig_mean_data, split_mean_data)
|
assert_array_equal(orig_mean_data, split_mean_data)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,8 @@
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
import typing
|
||||||
|
from typing import Dict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
|
@ -16,6 +18,7 @@ from junifer.data import load_coordinates, load_mask
|
||||||
from junifer.markers.sphere_aggregation import SphereAggregation
|
from junifer.markers.sphere_aggregation import SphereAggregation
|
||||||
from junifer.storage import SQLiteFeatureStorage
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
# Define common variables
|
# Define common variables
|
||||||
COORDS = "DMNBuckner"
|
COORDS = "DMNBuckner"
|
||||||
RADIUS = 8
|
RADIUS = 8
|
||||||
|
|
@ -23,15 +26,12 @@ RADIUS = 8
|
||||||
|
|
||||||
def test_SphereAggregation_input_output() -> None:
|
def test_SphereAggregation_input_output() -> None:
|
||||||
"""Test SphereAggregation input and output types."""
|
"""Test SphereAggregation input and output types."""
|
||||||
marker = SphereAggregation(
|
marker = SphereAggregation(coords="DMNBuckner", method="mean", on="VBM_GM")
|
||||||
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
|
for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]:
|
||||||
)
|
assert marker.get_output_type(in_) == out_
|
||||||
|
|
||||||
output = marker.get_output_kind(["VBM_GM", "BOLD"])
|
|
||||||
assert output == ["table", "timeseries"]
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="Unknown input"):
|
with pytest.raises(ValueError, match="Unknown input"):
|
||||||
marker.get_output_kind(["VBM_GM", "BOLD", "unknown"])
|
marker.get_output_type("unknown")
|
||||||
|
|
||||||
|
|
||||||
def test_SphereAggregation_3D() -> None:
|
def test_SphereAggregation_3D() -> None:
|
||||||
|
|
@ -52,23 +52,13 @@ def test_SphereAggregation_3D() -> None:
|
||||||
marker = SphereAggregation(
|
marker = SphereAggregation(
|
||||||
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
|
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
|
||||||
)
|
)
|
||||||
input = {"VBM_GM": {"data": img}}
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
|
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
|
||||||
|
|
||||||
assert jun_values4d.ndim == 2
|
assert jun_values4d.ndim == 2
|
||||||
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
||||||
assert_array_equal(auto4d, jun_values4d)
|
assert_array_equal(auto4d, jun_values4d)
|
||||||
|
|
||||||
meta = marker.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["coords"] == COORDS
|
|
||||||
assert meta["radius"] == RADIUS
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "VBM_GM_SphereAggregation"
|
|
||||||
assert meta["class"] == "SphereAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
|
|
||||||
def test_SphereAggregation_4D() -> None:
|
def test_SphereAggregation_4D() -> None:
|
||||||
"""Test SphereAggregation object on 4D images."""
|
"""Test SphereAggregation object on 4D images."""
|
||||||
|
|
@ -84,26 +74,14 @@ def test_SphereAggregation_4D() -> None:
|
||||||
auto4d = nifti_masker.fit_transform(fmri_img)
|
auto4d = nifti_masker.fit_transform(fmri_img)
|
||||||
|
|
||||||
# Create SphereAggregation object
|
# Create SphereAggregation object
|
||||||
marker = SphereAggregation(
|
marker = SphereAggregation(coords=COORDS, method="mean", radius=RADIUS)
|
||||||
coords=COORDS, method="mean", radius=RADIUS
|
input = {"BOLD": {"data": fmri_img, "meta": {}}}
|
||||||
)
|
|
||||||
input = {"BOLD": {"data": fmri_img}}
|
|
||||||
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
|
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
|
||||||
|
|
||||||
assert jun_values4d.ndim == 2
|
assert jun_values4d.ndim == 2
|
||||||
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
||||||
assert_array_equal(auto4d, jun_values4d)
|
assert_array_equal(auto4d, jun_values4d)
|
||||||
|
|
||||||
meta = marker.get_meta("BOLD")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["coords"] == COORDS
|
|
||||||
assert meta["radius"] == RADIUS
|
|
||||||
assert meta["mask"] is None
|
|
||||||
assert meta["name"] == "BOLD_SphereAggregation"
|
|
||||||
assert meta["class"] == "SphereAggregation"
|
|
||||||
assert meta["kind"] == "BOLD"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
||||||
|
|
||||||
def test_SphereAggregation_storage(tmp_path: Path) -> None:
|
def test_SphereAggregation_storage(tmp_path: Path) -> None:
|
||||||
"""Test SphereAggregation storage.
|
"""Test SphereAggregation storage.
|
||||||
|
|
@ -122,31 +100,38 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
|
||||||
|
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
meta = {
|
meta = {
|
||||||
"element": "test",
|
"element": {"subject": "sub-01", "session": "ses-01"},
|
||||||
"version": "0.0.1",
|
"dependencies": {"nilearn", "nibabel"},
|
||||||
"marker": {"name": "fcname"},
|
|
||||||
}
|
}
|
||||||
input = {"VBM_GM": {"data": img}, "meta": meta}
|
input = {"VBM_GM": {"data": img, "meta": meta}}
|
||||||
marker = SphereAggregation(
|
marker = SphereAggregation(
|
||||||
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
|
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM"
|
||||||
)
|
)
|
||||||
|
|
||||||
marker.fit_transform(input, storage=storage)
|
marker.fit_transform(input, storage=storage)
|
||||||
|
|
||||||
|
features: Dict = typing.cast(Dict, storage.list_features())
|
||||||
|
assert any(
|
||||||
|
x["name"] == "VBM_GM_SphereAggregation" for x in features.values()
|
||||||
|
)
|
||||||
|
|
||||||
meta = {
|
meta = {
|
||||||
"element": "test",
|
"element": {"subject": "sub-01", "session": "ses-01"},
|
||||||
"version": "0.0.1",
|
"dependencies": {"nilearn", "nibabel"},
|
||||||
"marker": {"name": "BOLD_fcname"},
|
|
||||||
}
|
}
|
||||||
# Get the SPM auditory data
|
# Get the SPM auditory data
|
||||||
subject_data = datasets.fetch_spm_auditory()
|
subject_data = datasets.fetch_spm_auditory()
|
||||||
fmri_img = concat_imgs(subject_data.func) # type: ignore
|
fmri_img = concat_imgs(subject_data.func) # type: ignore
|
||||||
input = {"BOLD": {"data": fmri_img}, "meta": meta}
|
input = {"BOLD": {"data": fmri_img, "meta": meta}}
|
||||||
marker = SphereAggregation(
|
marker = SphereAggregation(
|
||||||
coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
|
coords=COORDS, method="mean", radius=RADIUS, on="BOLD"
|
||||||
)
|
)
|
||||||
|
|
||||||
marker.fit_transform(input, storage=storage)
|
marker.fit_transform(input, storage=storage)
|
||||||
|
features: Dict = typing.cast(Dict, storage.list_features())
|
||||||
|
assert any(
|
||||||
|
x["name"] == "BOLD_SphereAggregation" for x in features.values()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_SphereAggregation_3D_mask() -> None:
|
def test_SphereAggregation_3D_mask() -> None:
|
||||||
|
|
@ -164,26 +149,21 @@ def test_SphereAggregation_3D_mask() -> None:
|
||||||
|
|
||||||
# Create NiftSpheresMasker
|
# Create NiftSpheresMasker
|
||||||
nifti_masker = NiftiSpheresMasker(
|
nifti_masker = NiftiSpheresMasker(
|
||||||
seeds=coordinates, radius=RADIUS, mask_img=mask_img)
|
seeds=coordinates, radius=RADIUS, mask_img=mask_img
|
||||||
|
)
|
||||||
auto4d = nifti_masker.fit_transform(img)
|
auto4d = nifti_masker.fit_transform(img)
|
||||||
|
|
||||||
# Create SphereAggregation object
|
# Create SphereAggregation object
|
||||||
marker = SphereAggregation(
|
marker = SphereAggregation(
|
||||||
coords=COORDS, method="mean", radius=RADIUS, on="VBM_GM",
|
coords=COORDS,
|
||||||
mask="GM_prob0.2"
|
method="mean",
|
||||||
|
radius=RADIUS,
|
||||||
|
on="VBM_GM",
|
||||||
|
mask="GM_prob0.2",
|
||||||
)
|
)
|
||||||
input = {"VBM_GM": {"data": img}}
|
input = {"VBM_GM": {"data": img, "meta": {}}}
|
||||||
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
|
jun_values4d = marker.fit_transform(input)["VBM_GM"]["data"]
|
||||||
|
|
||||||
assert jun_values4d.ndim == 2
|
assert jun_values4d.ndim == 2
|
||||||
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
assert_array_equal(auto4d.shape, jun_values4d.shape)
|
||||||
assert_array_equal(auto4d, jun_values4d)
|
assert_array_equal(auto4d, jun_values4d)
|
||||||
|
|
||||||
meta = marker.get_meta("VBM_GM")["marker"]
|
|
||||||
assert meta["method"] == "mean"
|
|
||||||
assert meta["coords"] == COORDS
|
|
||||||
assert meta["radius"] == RADIUS
|
|
||||||
assert meta["name"] == "VBM_GM_SphereAggregation"
|
|
||||||
assert meta["class"] == "SphereAggregation"
|
|
||||||
assert meta["kind"] == "VBM_GM"
|
|
||||||
assert meta["method_params"] == {}
|
|
||||||
|
|
|
||||||
|
|
@ -5,3 +5,4 @@
|
||||||
|
|
||||||
from . import registry
|
from . import registry
|
||||||
from .pipeline_step_mixin import PipelineStepMixin
|
from .pipeline_step_mixin import PipelineStepMixin
|
||||||
|
from .update_meta_mixin import UpdateMetaMixin
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
from importlib.util import find_spec
|
||||||
from typing import Dict, List
|
from typing import Dict, List
|
||||||
|
|
||||||
from ..utils import raise_error
|
from ..utils import raise_error
|
||||||
|
|
@ -12,22 +13,6 @@ from ..utils import raise_error
|
||||||
class PipelineStepMixin:
|
class PipelineStepMixin:
|
||||||
"""Mixin class for pipeline."""
|
"""Mixin class for pipeline."""
|
||||||
|
|
||||||
def get_meta(self) -> Dict:
|
|
||||||
"""Get metadata.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
|
|
||||||
"""
|
|
||||||
t_meta = {}
|
|
||||||
t_meta["class"] = self.__class__.__name__
|
|
||||||
for k, v in vars(self).items():
|
|
||||||
if not k.startswith("_"):
|
|
||||||
t_meta[k] = v
|
|
||||||
return t_meta
|
|
||||||
|
|
||||||
def validate_input(self, input: List[str]) -> None:
|
def validate_input(self, input: List[str]) -> None:
|
||||||
"""Validate the input to the pipeline step.
|
"""Validate the input to the pipeline step.
|
||||||
|
|
||||||
|
|
@ -48,24 +33,22 @@ class PipelineStepMixin:
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input_type: str) -> str:
|
||||||
"""Get the kind of the pipeline step.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : list of str
|
input_type : str
|
||||||
The input to the pipeline step. The list must contain the
|
The data type input to the marker.
|
||||||
available Junifer Data dictionary keys.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
list of str
|
str
|
||||||
The updated list of available Junifer Data dictionary keys after
|
The storage type output by the marker.
|
||||||
the pipeline step.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement get_output_kind().",
|
msg="Concrete classes need to implement get_output_type().",
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -85,11 +68,30 @@ class PipelineStepMixin:
|
||||||
Raises
|
Raises
|
||||||
------
|
------
|
||||||
ValueError
|
ValueError
|
||||||
If the input does not have the required data.
|
If the pipeline step object is missing dependencies required for
|
||||||
|
its working or if the input does not have the required data.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
# Check if _DEPENDENCIES attribute is found;
|
||||||
|
# (markers and preprocessors will have them but not datareaders
|
||||||
|
# as of now)
|
||||||
|
dependencies_not_found = []
|
||||||
|
if hasattr(self, "_DEPENDENCIES"):
|
||||||
|
# Check if dependencies are importable
|
||||||
|
for dependency in self._DEPENDENCIES: # type: ignore
|
||||||
|
if find_spec(dependency) is None:
|
||||||
|
dependencies_not_found.append(dependency)
|
||||||
|
# Raise error if any dependency is not found
|
||||||
|
if dependencies_not_found:
|
||||||
|
raise_error(
|
||||||
|
msg=f"{dependencies_not_found} are not installed but are "
|
||||||
|
"required for using {self.name}.",
|
||||||
|
klass=ImportError,
|
||||||
|
)
|
||||||
|
|
||||||
self.validate_input(input=input)
|
self.validate_input(input=input)
|
||||||
return self.get_output_kind(input=input)
|
outputs = [self.get_output_type(t_input) for t_input in input]
|
||||||
|
return outputs
|
||||||
|
|
||||||
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
|
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
|
||||||
"""Fit and transform.
|
"""Fit and transform.
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Union
|
||||||
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from ..datagrabber import BaseDataGrabber
|
from ..datagrabber import BaseDataGrabber
|
||||||
from ..storage import BaseFeatureStorage
|
from ..storage import BaseFeatureStorage
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from junifer.pipeline.pipeline_step_mixin import PipelineStepMixin
|
from junifer.pipeline.pipeline_step_mixin import PipelineStepMixin
|
||||||
|
|
@ -15,13 +17,49 @@ def test_PipelineStepMixin() -> None:
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
mixin.validate_input([])
|
mixin.validate_input([])
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
mixin.get_output_kind([])
|
mixin.get_output_type("")
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
mixin.fit_transform({})
|
mixin.fit_transform({})
|
||||||
|
|
||||||
|
|
||||||
def test_pipeline_step_mixin_meta():
|
def test_pipeline_step_mixin_validate_correct_dependencies() -> None:
|
||||||
"""Test metadata for PipelineStepMixin."""
|
"""Test validate with correct dependencies."""
|
||||||
pipemixin = PipelineStepMixin()
|
|
||||||
t_meta = pipemixin.get_meta()
|
class CorrectMixer(PipelineStepMixin):
|
||||||
assert t_meta["class"] == "PipelineStepMixin"
|
"""Test class for validation."""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"setuptools"}
|
||||||
|
|
||||||
|
def validate_input(self, input: List[str]) -> None:
|
||||||
|
print(input)
|
||||||
|
|
||||||
|
def get_output_type(self, input_type: str) -> str:
|
||||||
|
return input_type
|
||||||
|
|
||||||
|
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
|
||||||
|
return {"input": input}
|
||||||
|
|
||||||
|
mixer = CorrectMixer()
|
||||||
|
mixer.validate([])
|
||||||
|
|
||||||
|
|
||||||
|
def test_pipeline_step_mixin_validate_incorrect_dependencies() -> None:
|
||||||
|
"""Test validate with incorrect dependencies."""
|
||||||
|
|
||||||
|
class IncorrectMixer(PipelineStepMixin):
|
||||||
|
"""Test class for validation."""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"foobar"}
|
||||||
|
|
||||||
|
def validate_input(self, input: List[str]) -> None:
|
||||||
|
print(input)
|
||||||
|
|
||||||
|
def get_output_type(self, input_type: str) -> str:
|
||||||
|
return input_type
|
||||||
|
|
||||||
|
def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]:
|
||||||
|
return {"input": input}
|
||||||
|
|
||||||
|
mixer = IncorrectMixer()
|
||||||
|
with pytest.raises(ImportError, match="not installed"):
|
||||||
|
mixer.validate([])
|
||||||
|
|
|
||||||
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 abc import ABC, abstractmethod
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
from ..pipeline import PipelineStepMixin
|
from ..pipeline import PipelineStepMixin, UpdateMetaMixin
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
class BasePreprocessor(ABC, PipelineStepMixin):
|
class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
|
||||||
"""Provide abstract base class for all preprocessors.
|
"""Provide abstract base class for all preprocessors.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -58,8 +58,8 @@ class BasePreprocessor(ABC, PipelineStepMixin):
|
||||||
)
|
)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input: List[str]) -> List[str]:
|
||||||
"""Get output kind.
|
"""Get output type.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
|
@ -75,7 +75,7 @@ class BasePreprocessor(ABC, PipelineStepMixin):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement get_output_kind().",
|
msg="Concrete classes need to implement get_output_type().",
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -93,22 +93,6 @@ class BasePreprocessor(ABC, PipelineStepMixin):
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_meta(self, kind: str) -> Dict:
|
|
||||||
"""Get metadata.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
kind : str
|
|
||||||
The kind of pipeline step.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
dict
|
|
||||||
The metadata as a dictionary with the only key 'preprocess'.
|
|
||||||
"""
|
|
||||||
s_meta = super().get_meta()
|
|
||||||
return {"preprocess": s_meta}
|
|
||||||
|
|
||||||
def fit_transform(
|
def fit_transform(
|
||||||
self,
|
self,
|
||||||
input: Dict[str, Dict],
|
input: Dict[str, Dict],
|
||||||
|
|
@ -127,19 +111,17 @@ class BasePreprocessor(ABC, PipelineStepMixin):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
out = input
|
out = input
|
||||||
for kind in self._on:
|
for type_ in self._on:
|
||||||
if kind in input.keys():
|
if type_ in input.keys():
|
||||||
logger.info(f"Computing {kind}")
|
logger.info(f"Computing {type_}")
|
||||||
t_input = input[kind]
|
t_input = input[type_]
|
||||||
extra_input = input.copy()
|
extra_input = input.copy()
|
||||||
extra_input.pop(kind)
|
extra_input.pop(type_)
|
||||||
t_meta = t_input.get("meta", {}) # input kind meta
|
|
||||||
t_meta.update(self.get_meta(kind))
|
|
||||||
key, t_out = self.preprocess(
|
key, t_out = self.preprocess(
|
||||||
input=t_input, extra_input=extra_input
|
input=t_input, extra_input=extra_input
|
||||||
)
|
)
|
||||||
t_out.update(meta=t_meta)
|
|
||||||
out[key] = t_out
|
out[key] = t_out
|
||||||
|
self.update_meta(out[key], "preprocess")
|
||||||
return out
|
return out
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
|
|
|
||||||
|
|
@ -17,6 +17,7 @@ from ...api.decorators import register_preprocessor
|
||||||
from ...utils import logger, raise_error
|
from ...utils import logger, raise_error
|
||||||
from ..base import BasePreprocessor
|
from ..base import BasePreprocessor
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nibabel import MGHImage, Nifti1Image, Nifti2Image
|
from nibabel import MGHImage, Nifti1Image, Nifti2Image
|
||||||
|
|
||||||
|
|
@ -141,6 +142,8 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_DEPENDENCIES = {"numpy", "nilearn"}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
strategy: Optional[Dict[str, str]] = None,
|
strategy: Optional[Dict[str, str]] = None,
|
||||||
|
|
@ -222,7 +225,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
||||||
klass=ValueError,
|
klass=ValueError,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_type(self, input: List[str]) -> List[str]:
|
||||||
"""Get the kind of the pipeline step.
|
"""Get the kind of the pipeline step.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
|
||||||
|
|
@ -62,7 +62,7 @@ def test_fMRIPrepConfoundRemover_validate_input() -> None:
|
||||||
confound_remover.validate_input(input)
|
confound_remover.validate_input(input)
|
||||||
|
|
||||||
|
|
||||||
def test_fMRIPrepConfoundRemover_get_output_kind() -> None:
|
def test_fMRIPrepConfoundRemover_get_output_type() -> None:
|
||||||
"""Test fMRIPrepConfoundRemover validate_input."""
|
"""Test fMRIPrepConfoundRemover validate_input."""
|
||||||
confound_remover = fMRIPrepConfoundRemover()
|
confound_remover = fMRIPrepConfoundRemover()
|
||||||
inputs = [
|
inputs = [
|
||||||
|
|
@ -72,7 +72,7 @@ def test_fMRIPrepConfoundRemover_get_output_kind() -> None:
|
||||||
]
|
]
|
||||||
# Confound remover works in place
|
# Confound remover works in place
|
||||||
for input in inputs:
|
for input in inputs:
|
||||||
assert confound_remover.get_output_kind(input) == input
|
assert confound_remover.get_output_type(input) == input
|
||||||
|
|
||||||
|
|
||||||
def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None:
|
def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None:
|
||||||
|
|
@ -460,8 +460,8 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None:
|
||||||
clean_bold = typing.cast(nib.Nifti1Image, clean_bold)
|
clean_bold = typing.cast(nib.Nifti1Image, clean_bold)
|
||||||
# TODO: Find a better way to test functionality here
|
# TODO: Find a better way to test functionality here
|
||||||
assert (
|
assert (
|
||||||
clean_bold.header.get_zooms() == # type: ignore
|
clean_bold.header.get_zooms() # type: ignore
|
||||||
raw_bold.header.get_zooms()
|
== raw_bold.header.get_zooms() # type: ignore
|
||||||
)
|
)
|
||||||
assert clean_bold.get_fdata().shape == raw_bold.get_fdata().shape
|
assert clean_bold.get_fdata().shape == raw_bold.get_fdata().shape
|
||||||
# TODO: Test confound remover with mask, needs #79 to be implemented
|
# TODO: Test confound remover with mask, needs #79 to be implemented
|
||||||
|
|
@ -521,7 +521,6 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None:
|
||||||
AssertionError, assert_array_equal, orig_bold, trans_bold
|
AssertionError, assert_array_equal, orig_bold, trans_bold
|
||||||
)
|
)
|
||||||
|
|
||||||
assert output["meta"] == input["meta"] # general meta does not change
|
|
||||||
assert "meta" in output["BOLD"]
|
assert "meta" in output["BOLD"]
|
||||||
assert "preprocess" in output["BOLD"]["meta"]
|
assert "preprocess" in output["BOLD"]["meta"]
|
||||||
t_meta = output["BOLD"]["meta"]["preprocess"]
|
t_meta = output["BOLD"]["meta"]["preprocess"]
|
||||||
|
|
@ -535,3 +534,7 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None:
|
||||||
assert t_meta["high_pass"] is None
|
assert t_meta["high_pass"] is None
|
||||||
assert t_meta["t_r"] is None
|
assert t_meta["t_r"] is None
|
||||||
assert t_meta["mask_img"] is None
|
assert t_meta["mask_img"] is None
|
||||||
|
|
||||||
|
assert "dependencies" in output["BOLD"]["meta"]
|
||||||
|
dependencies = output["BOLD"]["meta"]["dependencies"]
|
||||||
|
assert dependencies == {"numpy", "nilearn"}
|
||||||
|
|
|
||||||
|
|
@ -19,10 +19,14 @@ def test_base_preprocessor_subclassing() -> None:
|
||||||
"""Test proper subclassing of BasePreprocessor."""
|
"""Test proper subclassing of BasePreprocessor."""
|
||||||
# Create concrete class
|
# Create concrete class
|
||||||
class MyBasePreprocessor(BasePreprocessor):
|
class MyBasePreprocessor(BasePreprocessor):
|
||||||
|
def __init__(self, on):
|
||||||
|
self.parameter = 1
|
||||||
|
super().__init__(on=on)
|
||||||
|
|
||||||
def get_valid_inputs(self):
|
def get_valid_inputs(self):
|
||||||
return ["BOLD", "T1w"]
|
return ["BOLD", "T1w"]
|
||||||
|
|
||||||
def get_output_kind(self, input):
|
def get_output_type(self, input):
|
||||||
return ["timeseries"]
|
return ["timeseries"]
|
||||||
|
|
||||||
def preprocess(self, input, extra_input):
|
def preprocess(self, input, extra_input):
|
||||||
|
|
@ -37,14 +41,23 @@ def test_base_preprocessor_subclassing() -> None:
|
||||||
|
|
||||||
# Create input for marker
|
# Create input for marker
|
||||||
input_ = {
|
input_ = {
|
||||||
"meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"},
|
|
||||||
"BOLD": {
|
"BOLD": {
|
||||||
"path": ".",
|
"path": ".",
|
||||||
"data": "data",
|
"data": "data",
|
||||||
|
"meta": {
|
||||||
|
"datagrabber": "dg",
|
||||||
|
"element": "elem",
|
||||||
|
"datareader": "dr",
|
||||||
|
},
|
||||||
},
|
},
|
||||||
"T1w": {
|
"T1w": {
|
||||||
"path": ".",
|
"path": ".",
|
||||||
"data": "data",
|
"data": "data",
|
||||||
|
"meta": {
|
||||||
|
"datagrabber": "dg",
|
||||||
|
"element": "elem",
|
||||||
|
"datareader": "dr",
|
||||||
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
prep = MyBasePreprocessor(on=["BOLD"])
|
prep = MyBasePreprocessor(on=["BOLD"])
|
||||||
|
|
@ -60,6 +73,13 @@ def test_base_preprocessor_subclassing() -> None:
|
||||||
assert "path" in output["BOLD"]
|
assert "path" in output["BOLD"]
|
||||||
assert "meta" in output["BOLD"]
|
assert "meta" in output["BOLD"]
|
||||||
|
|
||||||
|
meta = output["BOLD"]["meta"]
|
||||||
|
assert "preprocess" in meta
|
||||||
|
assert "class" in meta["preprocess"]
|
||||||
|
assert "MyBasePreprocessor" == meta["preprocess"]["class"]
|
||||||
|
assert "parameter" in meta["preprocess"]
|
||||||
|
assert 1 == meta["preprocess"]["parameter"]
|
||||||
|
|
||||||
assert "T1w" in output
|
assert "T1w" in output
|
||||||
assert "data" in output["T1w"]
|
assert "data" in output["T1w"]
|
||||||
assert output["T1w"]["data"] == "data"
|
assert output["T1w"]["data"] == "data"
|
||||||
|
|
|
||||||
|
|
@ -6,12 +6,13 @@
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Optional, Union
|
from typing import Dict, Iterable, List, Optional, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from .._version import __version__
|
|
||||||
from ..utils import raise_error
|
from ..utils import raise_error
|
||||||
|
from .utils import process_meta
|
||||||
|
|
||||||
|
|
||||||
class BaseFeatureStorage(ABC):
|
class BaseFeatureStorage(ABC):
|
||||||
|
|
@ -40,23 +41,29 @@ class BaseFeatureStorage(ABC):
|
||||||
self.uri = uri
|
self.uri = uri
|
||||||
if not isinstance(storage_types, list):
|
if not isinstance(storage_types, list):
|
||||||
storage_types = [storage_types]
|
storage_types = [storage_types]
|
||||||
|
if any(x not in self.get_valid_inputs() for x in storage_types):
|
||||||
|
wrong_storage_types = [
|
||||||
|
x for x in storage_types if x not in self.get_valid_inputs()
|
||||||
|
]
|
||||||
|
raise ValueError(
|
||||||
|
f"{self.__class__.__name__} cannot store {wrong_storage_types}"
|
||||||
|
)
|
||||||
self._valid_inputs = storage_types
|
self._valid_inputs = storage_types
|
||||||
self.single_output = single_output
|
self.single_output = single_output
|
||||||
|
|
||||||
def get_meta(self) -> Dict:
|
def get_valid_inputs(self) -> List[str]:
|
||||||
"""Get metadata.
|
"""Get valid storage types for input.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict
|
list of str
|
||||||
The metadata as a dictionary.
|
The list of storage types that can be used as input for this "
|
||||||
|
"storage.
|
||||||
"""
|
"""
|
||||||
meta = {}
|
raise_error(
|
||||||
meta["versions"] = {
|
msg="Concrete classes need to implement get_valid_inputs().",
|
||||||
"junifer": __version__,
|
klass=NotImplementedError,
|
||||||
}
|
)
|
||||||
return meta
|
|
||||||
|
|
||||||
def validate(self, input_: List[str]) -> None:
|
def validate(self, input_: List[str]) -> None:
|
||||||
"""Validate the input to the pipeline step.
|
"""Validate the input to the pipeline step.
|
||||||
|
|
@ -80,23 +87,15 @@ class BaseFeatureStorage(ABC):
|
||||||
)
|
)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def list_features(
|
def list_features(self) -> Dict:
|
||||||
self, return_df: bool = False
|
|
||||||
) -> Union[Dict[str, Dict], pd.DataFrame]:
|
|
||||||
"""List the features in the storage.
|
"""List the features in the storage.
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
return_df : bool, optional
|
|
||||||
If True, returns a pandas DataFrame. If False, returns a
|
|
||||||
dictionary (default False).
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict or pandas.DataFrame
|
dict
|
||||||
List of features in the storage. If dictionary is returned, the
|
List of features in the storage. The keys are the feature names to
|
||||||
keys are the feature names to be used in read_features() and the
|
be used in read_features() and the values are the metadata of each
|
||||||
values are the metadata of each feature.
|
feature.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
|
|
@ -131,19 +130,17 @@ class BaseFeatureStorage(ABC):
|
||||||
)
|
)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def store_metadata(self, meta: Dict) -> str:
|
def store_metadata(self, meta_md5: str, element: Dict, meta: Dict) -> None:
|
||||||
"""Store metadata.
|
"""Store metadata.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
meta_md5 : str
|
||||||
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
meta : dict
|
meta : dict
|
||||||
The metadata as a dictionary.
|
The metadata as a dictionary.
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
str
|
|
||||||
The metadata column.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement store_metadata().",
|
msg="Concrete classes need to implement store_metadata().",
|
||||||
|
|
@ -166,65 +163,117 @@ class BaseFeatureStorage(ABC):
|
||||||
If ``kind`` is invalid.
|
If ``kind`` is invalid.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
# Do the check before calling the abstract methods, otherwise the
|
||||||
|
# meta might be stored even if the data is not stored.
|
||||||
|
if kind not in self._valid_inputs:
|
||||||
|
raise_error(
|
||||||
|
msg=f"I don't know how to store {kind}.",
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
t_meta = kwargs.pop("meta")
|
||||||
|
meta_md5, t_meta, t_element = process_meta(t_meta)
|
||||||
|
self.store_metadata(meta_md5=meta_md5, element=t_element, meta=t_meta)
|
||||||
if kind == "matrix":
|
if kind == "matrix":
|
||||||
self.store_matrix(**kwargs)
|
self.store_matrix(meta_md5=meta_md5, element=t_element, **kwargs)
|
||||||
elif kind == "timeseries":
|
elif kind == "timeseries":
|
||||||
self.store_timeseries(**kwargs)
|
self.store_timeseries(
|
||||||
|
meta_md5=meta_md5, element=t_element, **kwargs
|
||||||
|
)
|
||||||
elif kind == "table":
|
elif kind == "table":
|
||||||
self.store_table(**kwargs)
|
self.store_table(meta_md5=meta_md5, element=t_element, **kwargs)
|
||||||
else:
|
|
||||||
raise ValueError(f"I don't know how to store {kind}")
|
|
||||||
|
|
||||||
def store_df(self, **kwargs) -> None:
|
def store_matrix(
|
||||||
"""Store pandas DataFrame.
|
self,
|
||||||
|
meta_md5: str,
|
||||||
Parameters
|
element: Dict,
|
||||||
----------
|
data: np.ndarray,
|
||||||
**kwargs : dict
|
col_names: Optional[Iterable[str]] = None,
|
||||||
The keyword arguments.
|
row_names: Optional[Iterable[str]] = None,
|
||||||
|
matrix_kind: Optional[str] = "full",
|
||||||
"""
|
diagonal: bool = True,
|
||||||
raise_error(
|
) -> None:
|
||||||
msg="Concrete classes need to implement store_df().",
|
|
||||||
klass=NotImplementedError,
|
|
||||||
)
|
|
||||||
|
|
||||||
def store_matrix(self, **kwargs) -> None:
|
|
||||||
"""Store matrix.
|
"""Store matrix.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
**kwargs : dict
|
meta_md5 : str
|
||||||
The keyword arguments.
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
data : numpy.ndarray
|
||||||
|
The matrix data to store.
|
||||||
|
col_names : list or tuple of str, optional
|
||||||
|
The column names (default None).
|
||||||
|
row_names : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1.
|
||||||
|
If None and number of rows greater than 1, then the name will be
|
||||||
|
"index" (default None).
|
||||||
|
matrix_kind : str, optional
|
||||||
|
The kind of matrix:
|
||||||
|
|
||||||
|
* ``triu`` : store upper triangular only
|
||||||
|
* ``tril`` : store lower triangular
|
||||||
|
* ``full`` : full matrix
|
||||||
|
|
||||||
|
(default "full").
|
||||||
|
diagonal : bool, optional
|
||||||
|
Whether to store the diagonal. If `matrix_kind` is "full", setting
|
||||||
|
this to False will raise an error (default True).
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement store_matrix2d().",
|
msg="Concrete classes need to implement store_matrix2d().",
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
def store_table(self, **kwargs) -> None:
|
def store_table(
|
||||||
|
self,
|
||||||
|
meta_md5: str,
|
||||||
|
element: Dict,
|
||||||
|
data: Union[np.ndarray, List],
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
rows_col_name: Optional[str] = None,
|
||||||
|
) -> None:
|
||||||
"""Store table.
|
"""Store table.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
**kwargs : dict
|
meta_md5 : str
|
||||||
The keyword arguments.
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
data : numpy.ndarray or list
|
||||||
|
The table data to store.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The columns (default None).
|
||||||
|
rows_col_name : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1.
|
||||||
|
If None and number of rows greater than 1, then the name will be
|
||||||
|
"index" (default None).
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement store_table().",
|
msg="Concrete classes need to implement store_table().",
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
def store_timeseries(self, **kwargs) -> None:
|
def store_timeseries(
|
||||||
"""Store timeseries.
|
self,
|
||||||
|
meta_md5: str,
|
||||||
|
element: Dict,
|
||||||
|
data: np.ndarray,
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Implement timeseries storing.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
**kwargs : dict
|
meta_md5 : str
|
||||||
The keyword arguments.
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
data : numpy.ndarray
|
||||||
|
The timeseries data to store.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The column labels (default None).
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="Concrete classes need to implement store_timeseries().",
|
msg="Concrete classes need to implement store_timeseries().",
|
||||||
|
|
|
||||||
|
|
@ -6,8 +6,9 @@
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, Union
|
from typing import Any, Dict, Iterable, List, Optional, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from .base import BaseFeatureStorage
|
from .base import BaseFeatureStorage
|
||||||
|
|
@ -39,6 +40,17 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(uri=uri, single_output=single_output, **kwargs)
|
super().__init__(uri=uri, single_output=single_output, **kwargs)
|
||||||
|
|
||||||
|
def get_valid_inputs(self) -> List[str]:
|
||||||
|
"""Get valid storage types for input.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
list of str
|
||||||
|
The list of storage types that can be used as input for this "
|
||||||
|
"storage.
|
||||||
|
"""
|
||||||
|
return ["matrix", "table", "timeseries"]
|
||||||
|
|
||||||
def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame:
|
def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame:
|
||||||
"""Convert the metadata to a pandas DataFrame.
|
"""Convert the metadata to a pandas DataFrame.
|
||||||
|
|
||||||
|
|
@ -57,8 +69,166 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
|
||||||
data_df = {}
|
data_df = {}
|
||||||
for k, v in meta.items():
|
for k, v in meta.items():
|
||||||
data_df[k] = json.dumps(v, sort_keys=True)
|
data_df[k] = json.dumps(v, sort_keys=True)
|
||||||
if "marker" in meta:
|
|
||||||
data_df["name"] = meta["marker"]["name"]
|
|
||||||
df = pd.DataFrame(data_df, index=[meta_md5])
|
df = pd.DataFrame(data_df, index=[meta_md5])
|
||||||
df.index.name = "meta_md5"
|
df.index.name = "meta_md5"
|
||||||
return df
|
return df
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def element_to_index(
|
||||||
|
element: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None
|
||||||
|
) -> pd.MultiIndex:
|
||||||
|
"""Convert the element metadata to index.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
n_rows : int, optional
|
||||||
|
Number of rows to create (default 1).
|
||||||
|
rows_col_name: str, optional
|
||||||
|
The column name to use in case `n_rows` > 1. If None and
|
||||||
|
n_rows > 1, the name will be "idx" (default None).
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
pandas.MultiIndex
|
||||||
|
The index of the dataframe to store.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError
|
||||||
|
If `meta` does not contain the key "element".
|
||||||
|
|
||||||
|
"""
|
||||||
|
# Check rows_col_name
|
||||||
|
if rows_col_name is None:
|
||||||
|
rows_col_name = "idx"
|
||||||
|
elem_idx: Dict[Any, Any] = {
|
||||||
|
k: [v] * n_rows for k, v in element.items()
|
||||||
|
}
|
||||||
|
elem_idx[rows_col_name] = np.arange(n_rows)
|
||||||
|
# Create index
|
||||||
|
index = pd.MultiIndex.from_frame(
|
||||||
|
pd.DataFrame(elem_idx, index=range(n_rows))
|
||||||
|
)
|
||||||
|
return index
|
||||||
|
|
||||||
|
def store_df(
|
||||||
|
self, meta_md5: str, element: Dict, df: Union[pd.DataFrame, pd.Series]
|
||||||
|
) -> None:
|
||||||
|
"""Implement pandas DataFrame storing.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
df : pandas.DataFrame or pandas.Series
|
||||||
|
The pandas DataFrame or Series to store.
|
||||||
|
meta : dict
|
||||||
|
The metadata as a dictionary.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError
|
||||||
|
If the dataframe index has items that are not in the index
|
||||||
|
generated from the metadata.
|
||||||
|
|
||||||
|
"""
|
||||||
|
raise NotImplementedError("Implement in subclass.")
|
||||||
|
|
||||||
|
def _store_2d(
|
||||||
|
self,
|
||||||
|
meta_md5: str,
|
||||||
|
element: Dict,
|
||||||
|
data: Union[np.ndarray, List],
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
rows_col_name: Optional[str] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Store 2D dataframe.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
meta_md5 : str
|
||||||
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
data : numpy.ndarray or List
|
||||||
|
The data to store.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The columns (default None).
|
||||||
|
rows_col_name : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1.
|
||||||
|
If None and number of rows greater than 1, then the name will be
|
||||||
|
"index" (default None).
|
||||||
|
|
||||||
|
"""
|
||||||
|
n_rows = len(data)
|
||||||
|
# Convert element metadata to index
|
||||||
|
idx = self.element_to_index(
|
||||||
|
element=element, n_rows=n_rows, rows_col_name=rows_col_name
|
||||||
|
)
|
||||||
|
# Prepare new dataframe
|
||||||
|
data_df = pd.DataFrame( # type: ignore
|
||||||
|
data, columns=columns, index=idx # type: ignore
|
||||||
|
)
|
||||||
|
# Store dataframe
|
||||||
|
self.store_df(meta_md5=meta_md5, element=element, df=data_df)
|
||||||
|
|
||||||
|
def store_table(
|
||||||
|
self,
|
||||||
|
meta_md5: str,
|
||||||
|
element: Dict,
|
||||||
|
data: Union[np.ndarray, List],
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
rows_col_name: Optional[str] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Implement table storing.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
meta_md5 : str
|
||||||
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
data : numpy.ndarray or List
|
||||||
|
The table data to store.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The columns (default None).
|
||||||
|
rows_col_name : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1.
|
||||||
|
If None and number of rows greater than 1, then the name will be
|
||||||
|
"index" (default None).
|
||||||
|
"""
|
||||||
|
self._store_2d(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element,
|
||||||
|
data=data,
|
||||||
|
columns=columns,
|
||||||
|
rows_col_name=rows_col_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
def store_timeseries(
|
||||||
|
self,
|
||||||
|
meta_md5: str,
|
||||||
|
element: Dict,
|
||||||
|
data: np.ndarray,
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Implement timeseries storing.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
meta_md5 : str
|
||||||
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
|
data : numpy.ndarray
|
||||||
|
The timeseries data to store.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The column labels (default None).
|
||||||
|
"""
|
||||||
|
self._store_2d(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element,
|
||||||
|
data=data,
|
||||||
|
columns=columns,
|
||||||
|
rows_col_name="timepoint",
|
||||||
|
)
|
||||||
|
|
|
||||||
|
|
@ -5,8 +5,9 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Dict, Iterable, List, Optional, Union
|
from typing import TYPE_CHECKING, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
import json
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from pandas.core.base import NoNewAttributesMixin
|
from pandas.core.base import NoNewAttributesMixin
|
||||||
|
|
@ -17,7 +18,8 @@ from tqdm import tqdm
|
||||||
from ..api.decorators import register_storage
|
from ..api.decorators import register_storage
|
||||||
from ..utils import logger, raise_error, warn_with_log
|
from ..utils import logger, raise_error, warn_with_log
|
||||||
from .pandas_base import PandasBaseFeatureStorage
|
from .pandas_base import PandasBaseFeatureStorage
|
||||||
from .utils import element_to_index, element_to_prefix, process_meta
|
from .utils import element_to_prefix
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sqlalchemy.engine import Engine
|
from sqlalchemy.engine import Engine
|
||||||
|
|
@ -86,7 +88,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
# Set upsert
|
# Set upsert
|
||||||
self._upsert = upsert
|
self._upsert = upsert
|
||||||
|
|
||||||
def get_engine(self, meta: Optional[Dict] = None) -> "Engine":
|
def get_engine(self, element: Optional[Dict] = None) -> "Engine":
|
||||||
"""Get engine.
|
"""Get engine.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -100,11 +102,6 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
The sqlalchemy engine.
|
The sqlalchemy engine.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# Set metadata as empty dictionary if None
|
|
||||||
if meta is None:
|
|
||||||
meta = {}
|
|
||||||
# Retrieve element key from metadata
|
|
||||||
element = meta.get("element", None)
|
|
||||||
# Prefixed elements
|
# Prefixed elements
|
||||||
prefix = ""
|
prefix = ""
|
||||||
if self.single_output is False:
|
if self.single_output is False:
|
||||||
|
|
@ -208,58 +205,15 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
msg=f"Invalid option {if_exists} for if_exists."
|
msg=f"Invalid option {if_exists} for if_exists."
|
||||||
)
|
)
|
||||||
|
|
||||||
def _store_2d(
|
def list_features(self) -> Dict:
|
||||||
self,
|
"""List the features in the storage.
|
||||||
data: Dict,
|
|
||||||
meta: Dict,
|
|
||||||
columns: Optional[Iterable[str]] = None,
|
|
||||||
rows_col_name: Optional[str] = None,
|
|
||||||
) -> None:
|
|
||||||
"""Store 2D dataframe.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
data : dict
|
|
||||||
The data to store.
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
columns : list or tuple of str, optional
|
|
||||||
The columns (default None).
|
|
||||||
rows_col_name : str, optional
|
|
||||||
The column name to use in case number of rows greater than 1.
|
|
||||||
If None and number of rows greater than 1, then the name will be
|
|
||||||
"index" (default None).
|
|
||||||
|
|
||||||
"""
|
|
||||||
n_rows = len(data)
|
|
||||||
# Convert element metadata to index
|
|
||||||
idx = element_to_index(
|
|
||||||
meta=meta, n_rows=n_rows, rows_col_name=rows_col_name
|
|
||||||
)
|
|
||||||
# Prepare new dataframe
|
|
||||||
data_df = pd.DataFrame( # type: ignore
|
|
||||||
data, columns=columns, index=idx # type: ignore
|
|
||||||
)
|
|
||||||
# Store dataframe
|
|
||||||
self.store_df(df=data_df, meta=meta)
|
|
||||||
|
|
||||||
def list_features(
|
|
||||||
self, return_df: bool = False
|
|
||||||
) -> Union[Dict, pd.DataFrame]:
|
|
||||||
"""Implement features listing from the storage.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
return_df : bool, optional
|
|
||||||
If True, returns a pandas DataFrame. If False, returns a
|
|
||||||
dictionary (default False).
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict or pandas.DataFrame
|
dict
|
||||||
List of features in the storage. If dictionary is returned, the
|
List of features in the storage. The keys are the feature names to
|
||||||
keys are the feature names to be used in read_features() and the
|
be used in read_features() and the values are the metadata of each
|
||||||
values are the metadata of each feature.
|
feature.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
meta_df = pd.read_sql(
|
meta_df = pd.read_sql(
|
||||||
|
|
@ -267,10 +221,11 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
con=self.get_engine(),
|
con=self.get_engine(),
|
||||||
index_col="meta_md5",
|
index_col="meta_md5",
|
||||||
)
|
)
|
||||||
out = meta_df
|
meta_df.index = meta_df.index.str.replace(r"meta_", "")
|
||||||
# Return dictionary
|
out = meta_df.to_dict(orient="index") # type: ignore
|
||||||
if return_df is False:
|
for md5, t_meta in out.items():
|
||||||
out = meta_df.to_dict(orient="index") # type: ignore
|
for k, v in t_meta.items():
|
||||||
|
out[md5][k] = json.loads(v)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def read_df(
|
def read_df(
|
||||||
|
|
@ -327,7 +282,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
con=engine,
|
con=engine,
|
||||||
index_col="meta_md5",
|
index_col="meta_md5",
|
||||||
)
|
)
|
||||||
t_df = meta_df.query(f"name == '{feature_name}'")
|
|
||||||
|
# Wrap in double quotes as the fields are in JSON format
|
||||||
|
t_df = meta_df.query(f"name == '\"{feature_name}\"'")
|
||||||
if len(t_df) == 0:
|
if len(t_df) == 0:
|
||||||
raise_error(msg=f"Feature {feature_name} not found")
|
raise_error(msg=f"Feature {feature_name} not found")
|
||||||
elif len(t_df) > 1:
|
elif len(t_df) > 1:
|
||||||
|
|
@ -339,8 +296,11 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
table_name = f"meta_{t_df.index[0]}"
|
table_name = f"meta_{t_df.index[0]}"
|
||||||
|
if table_name not in inspect(engine).get_table_names():
|
||||||
|
raise_error(msg=f"Feature MD5 {feature_md5} not found")
|
||||||
# Read metadata from table
|
# Read metadata from table
|
||||||
df = pd.read_sql(sql=table_name, con=engine)
|
df = pd.read_sql(sql=table_name, con=engine)
|
||||||
|
|
||||||
# Read the index
|
# Read the index
|
||||||
query = (
|
query = (
|
||||||
"SELECT ii.name FROM sqlite_master AS m, "
|
"SELECT ii.name FROM sqlite_master AS m, "
|
||||||
|
|
@ -356,36 +316,30 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
df = df.set_index(index_names)
|
df = df.set_index(index_names)
|
||||||
return df
|
return df
|
||||||
|
|
||||||
def store_metadata(self, meta: Dict) -> str:
|
def store_metadata(self, meta_md5: str, element: Dict, meta: Dict) -> None:
|
||||||
r"""Implement metadata storing in the storage.
|
"""Implement metadata storing in the storage.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
meta_md5 : str
|
||||||
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
meta : dict
|
meta : dict
|
||||||
The metadata as a dictionary.
|
The metadata as a dictionary.
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
str
|
|
||||||
The MD5 hash of the metadata prefixed with "meta\_".
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# Copy metadata
|
|
||||||
t_meta = meta.copy()
|
|
||||||
# Update metadata
|
|
||||||
t_meta.update(self.get_meta())
|
|
||||||
# Process metadata
|
|
||||||
meta_md5, t_meta_row = process_meta(t_meta)
|
|
||||||
# Get sqlalchemy engine
|
# Get sqlalchemy engine
|
||||||
engine = self.get_engine(meta=t_meta)
|
engine = self.get_engine(element=element)
|
||||||
if meta_md5 not in inspect(engine).get_table_names():
|
table_name = f"meta_{meta_md5}"
|
||||||
|
if table_name not in inspect(engine).get_table_names():
|
||||||
# Convert metadata to dataframe
|
# Convert metadata to dataframe
|
||||||
meta_df = self._meta_row(meta=t_meta_row, meta_md5=meta_md5)
|
meta_df = self._meta_row(meta=meta, meta_md5=meta_md5)
|
||||||
# Save dataframe
|
# Save dataframe
|
||||||
self._save_upsert(meta_df, "meta", engine)
|
self._save_upsert(meta_df, "meta", engine)
|
||||||
return f"meta_{meta_md5}"
|
|
||||||
|
|
||||||
def store_df(self, df: Union[pd.DataFrame, pd.Series], meta: Dict) -> None:
|
def store_df(
|
||||||
|
self, meta_md5: str, element: Dict, df: Union[pd.DataFrame, pd.Series]
|
||||||
|
) -> None:
|
||||||
"""Implement pandas DataFrame storing.
|
"""Implement pandas DataFrame storing.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -405,7 +359,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
# TODO: Test this function
|
# TODO: Test this function
|
||||||
# Check that the index generated by meta matches the one in
|
# Check that the index generated by meta matches the one in
|
||||||
# the dataframe.
|
# the dataframe.
|
||||||
idx = element_to_index(meta)
|
idx = self.element_to_index(element)
|
||||||
# Given the meta, we might not know if there is an extra column added
|
# Given the meta, we might not know if there is an extra column added
|
||||||
# when storing a timeseries or 2d elements. We need to check if the
|
# when storing a timeseries or 2d elements. We need to check if the
|
||||||
# extra element is only one.
|
# extra element is only one.
|
||||||
|
|
@ -418,24 +372,25 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
elif len(extra) == 1:
|
elif len(extra) == 1:
|
||||||
# The df has one extra index item, this should be the new name
|
# The df has one extra index item, this should be the new name
|
||||||
# of the missing element in the index
|
# of the missing element in the index
|
||||||
idx = element_to_index(meta, rows_col_name=extra[0])
|
idx = self.element_to_index(element, rows_col_name=extra[0])
|
||||||
|
|
||||||
if any(x not in df.index.names for x in idx.names):
|
if any(x not in df.index.names for x in idx.names):
|
||||||
raise_error(
|
raise_error(
|
||||||
"The index of the dataframe is missing index items that are "
|
"The index of the dataframe is missing index items that are "
|
||||||
"generated from the metadata."
|
"generated from the metadata."
|
||||||
)
|
)
|
||||||
# Get table name
|
|
||||||
table_name = self.store_metadata(meta)
|
table_name = f"meta_{meta_md5}"
|
||||||
# Get sqlalchemy engine
|
# Get sqlalchemy engine
|
||||||
engine = self.get_engine(meta)
|
engine = self.get_engine(element)
|
||||||
# Save data
|
# Save data
|
||||||
self._save_upsert(df, table_name, engine)
|
self._save_upsert(df, table_name, engine)
|
||||||
|
|
||||||
def store_matrix(
|
def store_matrix(
|
||||||
self,
|
self,
|
||||||
|
meta_md5: str,
|
||||||
|
element: Dict,
|
||||||
data: np.ndarray,
|
data: np.ndarray,
|
||||||
meta: Dict,
|
|
||||||
col_names: Optional[List[str]] = None,
|
col_names: Optional[List[str]] = None,
|
||||||
row_names: Optional[List[str]] = None,
|
row_names: Optional[List[str]] = None,
|
||||||
matrix_kind: Optional[str] = "full",
|
matrix_kind: Optional[str] = "full",
|
||||||
|
|
@ -445,6 +400,10 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
meta_md5 : str
|
||||||
|
The metadata MD5 hash.
|
||||||
|
element : dict
|
||||||
|
The element as a dictionary.
|
||||||
data : numpy.ndarray
|
data : numpy.ndarray
|
||||||
The matrix data to store.
|
The matrix data to store.
|
||||||
meta : dict
|
meta : dict
|
||||||
|
|
@ -465,7 +424,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
(default "full").
|
(default "full").
|
||||||
diagonal : bool, optional
|
diagonal : bool, optional
|
||||||
Whether to store the diagonal. If `matrix_kind` is "full", setting
|
Whether to store the diagonal. If `matrix_kind` is "full", setting
|
||||||
this to False will raise an error (default True)..
|
this to False will raise an error (default True).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
if diagonal is False and matrix_kind not in ["triu", "tril"]:
|
if diagonal is False and matrix_kind not in ["triu", "tril"]:
|
||||||
|
|
@ -519,74 +478,27 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
|
|
||||||
# Convert element metadata to index
|
# Convert element metadata to index
|
||||||
n_rows = 1
|
n_rows = 1
|
||||||
idx = element_to_index(meta=meta, n_rows=n_rows, rows_col_name=None)
|
idx = self.element_to_index(
|
||||||
|
element=element, n_rows=n_rows, rows_col_name=None
|
||||||
|
)
|
||||||
# Prepare new dataframe
|
# Prepare new dataframe
|
||||||
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
|
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
|
||||||
|
|
||||||
if len(columns) > 2000: # TODO: check SQLITE_MAX_COLUMN
|
if len(columns) > 2000:
|
||||||
|
warn_with_log(
|
||||||
|
msg="The number of columns is greater than 2000. "
|
||||||
|
"The data will be stored in long format. "
|
||||||
|
"This will make it slower to collect the data. "
|
||||||
|
"Future versions of junifer will provide additional storage "
|
||||||
|
"options that will not raise this warning.",
|
||||||
|
)
|
||||||
data_df = data_df.stack()
|
data_df = data_df.stack()
|
||||||
new_names = [x for x in data_df.index.names[:-1]]
|
new_names = [x for x in data_df.index.names[:-1]]
|
||||||
new_names.append("pair")
|
new_names.append("pair")
|
||||||
data_df.index.names = new_names
|
data_df.index.names = new_names
|
||||||
|
|
||||||
# Store dataframe
|
# Store dataframe
|
||||||
self.store_df(df=data_df, meta=meta)
|
self.store_df(meta_md5=meta_md5, element=element, df=data_df)
|
||||||
|
|
||||||
def store_table(
|
|
||||||
self,
|
|
||||||
data: Dict,
|
|
||||||
meta: Dict,
|
|
||||||
columns: Optional[Iterable[str]] = None,
|
|
||||||
rows_col_name: Optional[str] = None,
|
|
||||||
) -> None:
|
|
||||||
"""Implement table storing.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
data : dict
|
|
||||||
The table data to store.
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
columns : list or tuple of str, optional
|
|
||||||
The columns (default None).
|
|
||||||
rows_col_name : str, optional
|
|
||||||
The column name to use in case number of rows greater than 1.
|
|
||||||
If None and number of rows greater than 1, then the name will be
|
|
||||||
"index" (default None).
|
|
||||||
|
|
||||||
"""
|
|
||||||
self._store_2d(
|
|
||||||
data=data, meta=meta, columns=columns, rows_col_name=rows_col_name
|
|
||||||
)
|
|
||||||
|
|
||||||
def store_timeseries(
|
|
||||||
self,
|
|
||||||
data: Dict,
|
|
||||||
meta: Dict,
|
|
||||||
columns: Optional[Iterable[str]] = None,
|
|
||||||
row_names: str = "timepoint",
|
|
||||||
) -> None:
|
|
||||||
"""Implement timeseries storing.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
data : dict
|
|
||||||
The timeseries data to store.
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
columns : list or tuple of str, optional
|
|
||||||
The column labels (default None).
|
|
||||||
row_names : str, optional
|
|
||||||
The column name to use in case number of rows greater than 1
|
|
||||||
(default "timepoint").
|
|
||||||
|
|
||||||
"""
|
|
||||||
self._store_2d(
|
|
||||||
data=data,
|
|
||||||
meta=meta,
|
|
||||||
columns=columns,
|
|
||||||
rows_col_name="timepoint", # explicit so as to stop overriding
|
|
||||||
)
|
|
||||||
|
|
||||||
def collect(self) -> None:
|
def collect(self) -> None:
|
||||||
"""Implement data collection.
|
"""Implement data collection.
|
||||||
|
|
|
||||||
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 sqlalchemy import create_engine
|
||||||
|
|
||||||
from junifer.storage.sqlite import SQLiteFeatureStorage
|
from junifer.storage.sqlite import SQLiteFeatureStorage
|
||||||
from junifer.storage.utils import (
|
from junifer.storage.utils import element_to_prefix, process_meta
|
||||||
element_to_index,
|
|
||||||
element_to_prefix,
|
|
||||||
process_meta,
|
|
||||||
)
|
|
||||||
|
|
||||||
df1 = pd.DataFrame(
|
df1 = pd.DataFrame(
|
||||||
{
|
{
|
||||||
"element": [1, 2, 3, 4, 5],
|
"subject": [1, 2, 3, 4, 5],
|
||||||
"pk2": ["a", "b", "c", "d", "e"],
|
"pk2": ["a", "b", "c", "d", "e"],
|
||||||
"col1": [11, 22, 33, 44, 55],
|
"col1": [11, 22, 33, 44, 55],
|
||||||
"col2": [111, 222, 333, 444, 555],
|
"col2": [111, 222, 333, 444, 555],
|
||||||
}
|
}
|
||||||
).set_index(["element", "pk2"])
|
).set_index(["subject", "pk2"])
|
||||||
|
|
||||||
df2 = pd.DataFrame(
|
df2 = pd.DataFrame(
|
||||||
{
|
{
|
||||||
"element": [2, 5, 6],
|
"subject": [2, 5, 6],
|
||||||
"pk2": ["b", "e", "f"],
|
"pk2": ["b", "e", "f"],
|
||||||
"col1": [2222, 5555, 66],
|
"col1": [2222, 5555, 66],
|
||||||
"col2": [22222, 55555, 666],
|
"col2": [22222, 55555, 666],
|
||||||
}
|
}
|
||||||
).set_index(["element", "pk2"])
|
).set_index(["subject", "pk2"])
|
||||||
|
|
||||||
df_update = pd.DataFrame(
|
df_update = pd.DataFrame(
|
||||||
{
|
{
|
||||||
"element": [1, 2, 3, 4, 5, 6],
|
"subject": [1, 2, 3, 4, 5, 6],
|
||||||
"pk2": ["a", "b", "c", "d", "e", "f"],
|
"pk2": ["a", "b", "c", "d", "e", "f"],
|
||||||
"col1": [11, 2222, 33, 44, 5555, 66],
|
"col1": [11, 2222, 33, 44, 5555, 66],
|
||||||
"col2": [111, 22222, 333, 444, 55555, 666],
|
"col2": [111, 22222, 333, 444, 55555, 666],
|
||||||
}
|
}
|
||||||
).set_index(["element", "pk2"])
|
).set_index(["subject", "pk2"])
|
||||||
|
|
||||||
df_ignore = pd.DataFrame(
|
df_ignore = pd.DataFrame(
|
||||||
{
|
{
|
||||||
"element": [1, 2, 3, 4, 5, 6],
|
"subject": [1, 2, 3, 4, 5, 6],
|
||||||
"pk2": ["a", "b", "c", "d", "e", "f"],
|
"pk2": ["a", "b", "c", "d", "e", "f"],
|
||||||
"col1": [11, 22, 33, 44, 55, 66],
|
"col1": [11, 22, 33, 44, 55, 66],
|
||||||
"col2": [111, 222, 333, 444, 555, 666],
|
"col2": [111, 222, 333, 444, 555, 666],
|
||||||
}
|
}
|
||||||
).set_index(["element", "pk2"])
|
).set_index(["subject", "pk2"])
|
||||||
|
|
||||||
|
|
||||||
def _read_sql(
|
def _read_sql(
|
||||||
|
|
@ -148,15 +145,14 @@ def test_upsert_replace(tmp_path: Path) -> None:
|
||||||
uri = tmp_path / "test_upsert_replace.sqlite"
|
uri = tmp_path / "test_upsert_replace.sqlite"
|
||||||
# Single storage, must be the uri
|
# Single storage, must be the uri
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
# Metadata to store
|
|
||||||
meta = {"element": "test", "version": "0.0.1"}
|
|
||||||
# Save to database
|
# Save to database
|
||||||
storage.store_df(df=df1, meta=meta)
|
storage.store_df(
|
||||||
# Store metadata
|
meta_md5="table_name", element={"subject": "test"}, df=df1
|
||||||
table_name = storage.store_metadata(meta)
|
)
|
||||||
|
table_name = "meta_table_name"
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_df1 = _read_sql(
|
c_df1 = _read_sql(
|
||||||
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"]
|
table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(df1, c_df1)
|
assert_frame_equal(df1, c_df1)
|
||||||
|
|
@ -164,7 +160,7 @@ def test_upsert_replace(tmp_path: Path) -> None:
|
||||||
storage._save_upsert(df=df2, name=table_name, if_exists="replace")
|
storage._save_upsert(df=df2, name=table_name, if_exists="replace")
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_df2 = _read_sql(
|
c_df2 = _read_sql(
|
||||||
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"]
|
table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(df2, c_df2)
|
assert_frame_equal(df2, c_df2)
|
||||||
|
|
@ -182,24 +178,25 @@ def test_upsert_ignore(tmp_path: Path) -> None:
|
||||||
uri = tmp_path / "test_upsert_ignore.sqlite"
|
uri = tmp_path / "test_upsert_ignore.sqlite"
|
||||||
# Single storage, must be the uri
|
# Single storage, must be the uri
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
# Metadata to store
|
|
||||||
meta = {"element": "test", "version": "0.0.1"}
|
|
||||||
# Save to database
|
# Save to database
|
||||||
storage.store_df(df=df1, meta=meta)
|
storage.store_df(
|
||||||
# Store metadata
|
meta_md5="table_name", element={"subject": "test"}, df=df1
|
||||||
table_name = storage.store_metadata(meta=meta)
|
)
|
||||||
|
table_name = "meta_table_name"
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_df1 = _read_sql(
|
c_df1 = _read_sql(
|
||||||
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"]
|
table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(df1, c_df1)
|
assert_frame_equal(df1, c_df1)
|
||||||
# Check for warning
|
# Check for warning
|
||||||
with pytest.warns(RuntimeWarning, match="are already present"):
|
with pytest.warns(RuntimeWarning, match="are already present"):
|
||||||
storage.store_df(df2, meta)
|
storage.store_df(
|
||||||
|
meta_md5="table_name", element={"subject": "test"}, df=df2
|
||||||
|
)
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_dfignore = _read_sql(
|
c_dfignore = _read_sql(
|
||||||
table_name, uri=uri.as_posix(), index_col=["element", "pk2"]
|
table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(c_dfignore, df_ignore)
|
assert_frame_equal(c_dfignore, df_ignore)
|
||||||
|
|
@ -219,23 +216,24 @@ def test_upsert_update(tmp_path: Path) -> None:
|
||||||
"""
|
"""
|
||||||
uri = tmp_path / "test_upsert_delete.sqlite"
|
uri = tmp_path / "test_upsert_delete.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
# Metadata to store
|
|
||||||
meta = {"element": "test", "version": "0.0.1"}
|
|
||||||
# Save to database
|
# Save to database
|
||||||
storage.store_df(df1, meta)
|
storage.store_df(
|
||||||
# Store metadata
|
meta_md5="table_name", element={"subject": "test"}, df=df1
|
||||||
table_name = storage.store_metadata(meta)
|
)
|
||||||
|
table_name = "meta_table_name"
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_df1 = _read_sql(
|
c_df1 = _read_sql(
|
||||||
table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"]
|
table_name=table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(df1, c_df1)
|
assert_frame_equal(df1, c_df1)
|
||||||
# Save to database
|
# Save to database
|
||||||
storage.store_df(df2, meta)
|
storage.store_df(
|
||||||
|
meta_md5="table_name", element={"subject": "test"}, df=df2
|
||||||
|
)
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_dfupdate = _read_sql(
|
c_dfupdate = _read_sql(
|
||||||
table_name, uri=uri.as_posix(), index_col=["element", "pk2"]
|
table_name, uri=uri.as_posix(), index_col=["subject", "pk2"]
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(c_dfupdate, df_update)
|
assert_frame_equal(c_dfupdate, df_update)
|
||||||
|
|
@ -268,50 +266,66 @@ def test_store_df_and_read_df(tmp_path: Path) -> None:
|
||||||
uri = tmp_path / "test_store_df_and_read_df.sqlite"
|
uri = tmp_path / "test_store_df_and_read_df.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
# Metadata to store
|
# Metadata to store
|
||||||
|
meta_md5 = "feature_md5"
|
||||||
|
element = {"subject": "test"}
|
||||||
meta = {
|
meta = {
|
||||||
"element": "test",
|
"element": element,
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fcname"},
|
"marker": {
|
||||||
|
"name": "markername",
|
||||||
|
},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
_, meta_to_store, element_to_store = process_meta(meta)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5="feature_md5", element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
|
|
||||||
# Columns to store
|
# Columns to store
|
||||||
to_store = df1[["col1", "col2"]]
|
to_store = df1[["col1", "col2"]]
|
||||||
# Check for error while storing
|
# Check for error while storing
|
||||||
with pytest.raises(ValueError, match=r"missing index items"):
|
with pytest.raises(ValueError, match=r"missing index items"):
|
||||||
storage.store_df(to_store.set_index("col1"), meta)
|
storage.store_df(
|
||||||
|
meta_md5=meta_md5, element=element, df=to_store.set_index("col1")
|
||||||
|
)
|
||||||
# Set index
|
# Set index
|
||||||
to_store = df1.reset_index().set_index(["element", "pk2", "col1"])
|
to_store = df1.reset_index().set_index(["subject", "pk2", "col1"])
|
||||||
# Check for error while storing
|
# Check for error while storing
|
||||||
with pytest.raises(ValueError, match=r"extra items"):
|
with pytest.raises(ValueError, match=r"extra items"):
|
||||||
storage.store_df(to_store, meta)
|
storage.store_df(meta_md5=meta_md5, element=element, df=to_store)
|
||||||
# Convert element to index
|
# Convert element to index
|
||||||
idx = element_to_index(meta, n_rows=len(to_store))
|
idx = storage.element_to_index(element=element, n_rows=len(to_store))
|
||||||
# Set index
|
# Set index
|
||||||
to_store = to_store.set_index(idx)
|
to_store = to_store.set_index(idx)
|
||||||
# Store dataframe
|
# Store dataframe
|
||||||
storage.store_df(to_store, meta)
|
storage.store_df(meta_md5=meta_md5, element=element, df=to_store)
|
||||||
# Store metadata
|
|
||||||
table_name = storage.store_metadata(meta)
|
|
||||||
# List stored features
|
# List stored features
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
# Check correct usage
|
# Check correct usage
|
||||||
assert len(features) == 1
|
assert len(features) == 1
|
||||||
assert table_name.replace("meta_", "") in features
|
assert "feature_md5" in features
|
||||||
# Check for missing feature
|
# Check for missing feature
|
||||||
with pytest.raises(ValueError, match="not found"):
|
with pytest.raises(ValueError, match="not found"):
|
||||||
storage.read_df("wrong_md5")
|
storage.read_df(feature_md5="wrong_md5")
|
||||||
|
with pytest.raises(ValueError, match="not found"):
|
||||||
|
storage.read_df(feature_name="wrong_name")
|
||||||
# Check for missing feature to fetch
|
# Check for missing feature to fetch
|
||||||
with pytest.raises(ValueError, match="least one"):
|
with pytest.raises(ValueError, match="least one"):
|
||||||
storage.read_df()
|
storage.read_df()
|
||||||
# Check for multiple features to fetch
|
# Check for multiple features to fetch
|
||||||
with pytest.raises(ValueError, match="Only one"):
|
with pytest.raises(ValueError, match="Only one"):
|
||||||
storage.read_df("wrong_md5", "wrong_name")
|
storage.read_df(feature_name="wrong_name", feature_md5="wrong_md5")
|
||||||
# Get MD5 hash of features
|
# Get MD5 hash of features
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
|
assert "feature_md5" == feature_md5
|
||||||
# Check for key
|
# Check for key
|
||||||
assert "fcname" == features[feature_md5]["name"]
|
assert "BOLD_markername" == features[feature_md5]["name"]
|
||||||
# Read into dataframes
|
# Read into dataframes
|
||||||
read_df1 = storage.read_df(feature_md5=feature_md5)
|
read_df1 = storage.read_df(feature_md5=feature_md5)
|
||||||
read_df2 = storage.read_df(feature_name="fcname")
|
read_df2 = storage.read_df(feature_name="BOLD_markername")
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(read_df1, read_df2)
|
assert_frame_equal(read_df1, read_df2)
|
||||||
assert_frame_equal(read_df1, to_store)
|
assert_frame_equal(read_df1, to_store)
|
||||||
|
|
@ -330,10 +344,21 @@ def test_store_metadata(tmp_path: Path) -> None:
|
||||||
# Single storage, must be the uri
|
# Single storage, must be the uri
|
||||||
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
storage = SQLiteFeatureStorage(uri=uri, upsert="ignore")
|
||||||
# Metadata to store
|
# Metadata to store
|
||||||
meta = {"element": "test", "version": "0.0.1"}
|
meta = {
|
||||||
|
"element": {"subject": "test"},
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "test"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
|
||||||
|
meta_md5, meta_to_store, element_to_store = process_meta(meta)
|
||||||
# Store metadata
|
# Store metadata
|
||||||
table_name = storage.store_metadata(meta)
|
storage.store_metadata(
|
||||||
assert table_name.startswith("meta_")
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
|
features = storage.list_features()
|
||||||
|
feature_md5 = list(features.keys())[0]
|
||||||
|
assert meta_md5 == feature_md5
|
||||||
|
|
||||||
|
|
||||||
def test_store_table(tmp_path: Path) -> None:
|
def test_store_table(tmp_path: Path) -> None:
|
||||||
|
|
@ -348,7 +373,21 @@ def test_store_table(tmp_path: Path) -> None:
|
||||||
uri = tmp_path / "test_store_table.sqlite"
|
uri = tmp_path / "test_store_table.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
# Metadata to store
|
# Metadata to store
|
||||||
meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}}
|
element = {"subject": "test"}
|
||||||
|
dependencies = ["numpy"]
|
||||||
|
meta = {
|
||||||
|
"element": element,
|
||||||
|
"dependencies": dependencies,
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
|
||||||
|
meta_md5, meta_to_store, element_to_store = process_meta(meta)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
|
|
||||||
# Data to store
|
# Data to store
|
||||||
data = [
|
data = [
|
||||||
[1, 10],
|
[1, 10],
|
||||||
|
|
@ -358,18 +397,23 @@ def test_store_table(tmp_path: Path) -> None:
|
||||||
[5, 50],
|
[5, 50],
|
||||||
]
|
]
|
||||||
# Convert element to index
|
# Convert element to index
|
||||||
idx = element_to_index(meta, n_rows=5, rows_col_name="scan")
|
idx = storage.element_to_index(element, n_rows=5, rows_col_name="scan")
|
||||||
# Create dataframe
|
# Create dataframe
|
||||||
df = pd.DataFrame(data, columns=["f1", "f2"], index=idx)
|
df = pd.DataFrame(data, columns=["f1", "f2"], index=idx)
|
||||||
|
|
||||||
# Store table
|
# Store table
|
||||||
storage.store_table(data, meta, columns=["f1", "f2"], rows_col_name="scan")
|
storage.store_table(
|
||||||
# Store metadata
|
meta_md5=meta_md5,
|
||||||
table_name = storage.store_metadata(meta)
|
element=element_to_store,
|
||||||
|
data=data,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
|
)
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_df = _read_sql(
|
c_df = _read_sql(
|
||||||
table_name=table_name,
|
table_name=f"meta_{meta_md5}",
|
||||||
uri=uri.as_posix(),
|
uri=uri.as_posix(),
|
||||||
index_col=["element", "scan"],
|
index_col=["subject", "scan"],
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(df, c_df)
|
assert_frame_equal(df, c_df)
|
||||||
|
|
@ -377,19 +421,24 @@ def test_store_table(tmp_path: Path) -> None:
|
||||||
# New data to store
|
# New data to store
|
||||||
data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]]
|
data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]]
|
||||||
# Convert element to index
|
# Convert element to index
|
||||||
idx_new = element_to_index(meta, n_rows=6, rows_col_name="scan")
|
idx_new = storage.element_to_index(element, n_rows=6, rows_col_name="scan")
|
||||||
# Create dataframe
|
# Create dataframe
|
||||||
df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new)
|
df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new)
|
||||||
# Check warning
|
# Check warning
|
||||||
with pytest.warns(RuntimeWarning, match=r"Some rows"):
|
with pytest.warns(RuntimeWarning, match=r"Some rows"):
|
||||||
|
# Store table
|
||||||
storage.store_table(
|
storage.store_table(
|
||||||
data_new, meta, columns=["f1", "f2"], rows_col_name="scan"
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
|
data=data_new,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
)
|
)
|
||||||
# Read stored table
|
# Read stored table
|
||||||
c_df_new = _read_sql(
|
c_df_new = _read_sql(
|
||||||
table_name=table_name,
|
table_name=f"meta_{meta_md5}",
|
||||||
uri=uri.as_posix(),
|
uri=uri.as_posix(),
|
||||||
index_col=["element", "scan"],
|
index_col=["subject", "scan"],
|
||||||
)
|
)
|
||||||
# Check if dataframes are equal
|
# Check if dataframes are equal
|
||||||
assert_frame_equal(df_new, c_df_new)
|
assert_frame_equal(df_new, c_df_new)
|
||||||
|
|
@ -407,7 +456,20 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
uri = tmp_path / "test_store_table.sqlite"
|
uri = tmp_path / "test_store_table.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
# Metadata to store
|
# Metadata to store
|
||||||
meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}}
|
element = {"subject": "test"}
|
||||||
|
dependencies = ["numpy"]
|
||||||
|
meta = {
|
||||||
|
"element": element,
|
||||||
|
"dependencies": dependencies,
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
|
||||||
|
meta_md5, meta_to_store, element_to_store = process_meta(meta)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
|
|
||||||
# Store 4 x 3 full matrix
|
# Store 4 x 3 full matrix
|
||||||
data = np.array(
|
data = np.array(
|
||||||
|
|
@ -418,17 +480,17 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
|
|
||||||
# Store table
|
# Store table
|
||||||
storage.store_matrix(
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
data=data,
|
data=data,
|
||||||
meta=meta,
|
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [f"{i}~{j}" for i in row_names for j in col_names]
|
stored_names = [f"{i}~{j}" for i in row_names for j in col_names]
|
||||||
|
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
assert "fc" == features[feature_md5]["name"]
|
assert "BOLD_fc" == features[feature_md5]["name"]
|
||||||
|
|
||||||
read_df = storage.read_df(feature_md5=feature_md5)
|
read_df = storage.read_df(feature_md5=feature_md5)
|
||||||
assert read_df.shape == (1, 12)
|
assert read_df.shape == (1, 12)
|
||||||
|
|
@ -437,7 +499,14 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
# Store without row and column names
|
# Store without row and column names
|
||||||
uri = tmp_path / "test_store_table_nonames.sqlite"
|
uri = tmp_path / "test_store_table_nonames.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
storage.store_matrix(data=data, meta=meta)
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
|
# Store table
|
||||||
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, data=data
|
||||||
|
)
|
||||||
stored_names = [
|
stored_names = [
|
||||||
f"r{i}~c{j}"
|
f"r{i}~c{j}"
|
||||||
for i in range(data.shape[0])
|
for i in range(data.shape[0])
|
||||||
|
|
@ -445,20 +514,37 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
]
|
]
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
assert "fc" == features[feature_md5]["name"]
|
assert "BOLD_fc" == features[feature_md5]["name"]
|
||||||
read_df = storage.read_df(feature_md5=feature_md5)
|
read_df = storage.read_df(feature_md5=feature_md5)
|
||||||
assert list(read_df.columns) == stored_names
|
assert list(read_df.columns) == stored_names
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="Invalid kind"):
|
with pytest.raises(ValueError, match="Invalid kind"):
|
||||||
storage.store_matrix(data=data, meta=meta, matrix_kind="wrong")
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
|
data=data,
|
||||||
|
row_names=row_names,
|
||||||
|
col_names=col_names,
|
||||||
|
matrix_kind="wrong",
|
||||||
|
)
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="non-square"):
|
with pytest.raises(ValueError, match="non-square"):
|
||||||
storage.store_matrix(data=data, meta=meta, matrix_kind="triu")
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
|
data=data,
|
||||||
|
row_names=row_names,
|
||||||
|
col_names=col_names,
|
||||||
|
matrix_kind="triu",
|
||||||
|
)
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="cannot be False"):
|
with pytest.raises(ValueError, match="cannot be False"):
|
||||||
storage.store_matrix(
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
data=data,
|
data=data,
|
||||||
meta=meta,
|
row_names=row_names,
|
||||||
|
col_names=col_names,
|
||||||
matrix_kind="full",
|
matrix_kind="full",
|
||||||
diagonal=False,
|
diagonal=False,
|
||||||
)
|
)
|
||||||
|
|
@ -469,12 +555,17 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
col_names = ["col1", "col2", "col3"]
|
col_names = ["col1", "col2", "col3"]
|
||||||
uri = tmp_path / "test_store_table_triu.sqlite"
|
uri = tmp_path / "test_store_table_triu.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
storage.store_matrix(
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
data=data,
|
data=data,
|
||||||
meta=meta,
|
|
||||||
matrix_kind="triu",
|
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
|
matrix_kind="triu",
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [
|
stored_names = [
|
||||||
|
|
@ -488,7 +579,7 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
|
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
assert "fc" == features[feature_md5]["name"]
|
assert "BOLD_fc" == features[feature_md5]["name"]
|
||||||
read_df = storage.read_df(feature_md5=feature_md5)
|
read_df = storage.read_df(feature_md5=feature_md5)
|
||||||
assert list(read_df.columns) == stored_names
|
assert list(read_df.columns) == stored_names
|
||||||
assert_array_equal(
|
assert_array_equal(
|
||||||
|
|
@ -498,12 +589,17 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
# Store upper triangular matrix without diagonal
|
# Store upper triangular matrix without diagonal
|
||||||
uri = tmp_path / "test_store_table_triu_nodiagonal.sqlite"
|
uri = tmp_path / "test_store_table_triu_nodiagonal.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
storage.store_matrix(
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
data=data,
|
data=data,
|
||||||
meta=meta,
|
|
||||||
matrix_kind="triu",
|
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
|
matrix_kind="triu",
|
||||||
diagonal=False,
|
diagonal=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -515,7 +611,7 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
|
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
assert "fc" == features[feature_md5]["name"]
|
assert "BOLD_fc" == features[feature_md5]["name"]
|
||||||
read_df = storage.read_df(feature_md5=feature_md5)
|
read_df = storage.read_df(feature_md5=feature_md5)
|
||||||
assert list(read_df.columns) == stored_names
|
assert list(read_df.columns) == stored_names
|
||||||
assert_array_equal(
|
assert_array_equal(
|
||||||
|
|
@ -528,12 +624,17 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
col_names = ["col1", "col2", "col3"]
|
col_names = ["col1", "col2", "col3"]
|
||||||
uri = tmp_path / "test_store_table_tril.sqlite"
|
uri = tmp_path / "test_store_table_tril.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
storage.store_matrix(
|
storage.store_matrix(
|
||||||
|
meta_md5=meta_md5,
|
||||||
|
element=element_to_store,
|
||||||
data=data,
|
data=data,
|
||||||
meta=meta,
|
|
||||||
matrix_kind="tril",
|
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
|
matrix_kind="tril",
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [
|
stored_names = [
|
||||||
|
|
@ -547,7 +648,7 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
|
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
assert "fc" == features[feature_md5]["name"]
|
assert "BOLD_fc" == features[feature_md5]["name"]
|
||||||
read_df = storage.read_df(feature_md5=feature_md5)
|
read_df = storage.read_df(feature_md5=feature_md5)
|
||||||
assert list(read_df.columns) == stored_names
|
assert list(read_df.columns) == stored_names
|
||||||
assert_array_equal(
|
assert_array_equal(
|
||||||
|
|
@ -557,15 +658,19 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
# Store lower triangular matrix without diagonal
|
# Store lower triangular matrix without diagonal
|
||||||
uri = tmp_path / "test_store_table_tril_nodiagonal.sqlite"
|
uri = tmp_path / "test_store_table_tril_nodiagonal.sqlite"
|
||||||
storage = SQLiteFeatureStorage(uri=uri)
|
storage = SQLiteFeatureStorage(uri=uri)
|
||||||
|
# Store metadata
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
|
||||||
|
)
|
||||||
storage.store_matrix(
|
storage.store_matrix(
|
||||||
data,
|
meta_md5=meta_md5,
|
||||||
meta,
|
element=element_to_store,
|
||||||
matrix_kind="tril",
|
data=data,
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
|
matrix_kind="tril",
|
||||||
diagonal=False,
|
diagonal=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [
|
stored_names = [
|
||||||
"row2~col1",
|
"row2~col1",
|
||||||
"row3~col1",
|
"row3~col1",
|
||||||
|
|
@ -574,7 +679,7 @@ def test_store_matrix(tmp_path: Path) -> None:
|
||||||
|
|
||||||
features = storage.list_features()
|
features = storage.list_features()
|
||||||
feature_md5 = list(features.keys())[0]
|
feature_md5 = list(features.keys())[0]
|
||||||
assert "fc" == features[feature_md5]["name"]
|
assert "BOLD_fc" == features[feature_md5]["name"]
|
||||||
read_df = storage.read_df(feature_md5=feature_md5)
|
read_df = storage.read_df(feature_md5=feature_md5)
|
||||||
assert list(read_df.columns) == stored_names
|
assert list(read_df.columns) == stored_names
|
||||||
assert_array_equal(
|
assert_array_equal(
|
||||||
|
|
@ -597,18 +702,21 @@ def test_store_multiple_output(tmp_path: Path):
|
||||||
# Metadata to store
|
# Metadata to store
|
||||||
meta1 = {
|
meta1 = {
|
||||||
"element": {"subject": "test-01", "session": "ses-01"},
|
"element": {"subject": "test-01", "session": "ses-01"},
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fc"},
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
meta2 = {
|
meta2 = {
|
||||||
"element": {"subject": "test-02", "session": "ses-01"},
|
"element": {"subject": "test-02", "session": "ses-01"},
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fc"},
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
meta3 = {
|
meta3 = {
|
||||||
"element": {"subject": "test-01", "session": "ses-02"},
|
"element": {"subject": "test-01", "session": "ses-02"},
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fc"},
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
# Data to store
|
# Data to store
|
||||||
data1 = np.array(
|
data1 = np.array(
|
||||||
|
|
@ -623,35 +731,62 @@ def test_store_multiple_output(tmp_path: Path):
|
||||||
data2 = data1 * 10
|
data2 = data1 * 10
|
||||||
data3 = data1 * 20
|
data3 = data1 * 20
|
||||||
# Process metadata for storage
|
# Process metadata for storage
|
||||||
hash1, _ = process_meta(meta1)
|
hash1, meta_to_store1, element_to_store1 = process_meta(meta1)
|
||||||
# Convert element to index
|
# Convert element to index
|
||||||
idx1 = element_to_index(meta1, n_rows=5, rows_col_name="scan")
|
idx1 = storage.element_to_index(
|
||||||
|
element_to_store1, n_rows=5, rows_col_name="scan"
|
||||||
|
)
|
||||||
# Create dataframe
|
# Create dataframe
|
||||||
df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1)
|
df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1)
|
||||||
# Process metadata for storage
|
# Process metadata for storage
|
||||||
hash2, _ = process_meta(meta2)
|
hash2, meta_to_store2, element_to_store2 = process_meta(meta2)
|
||||||
# Convert element to index
|
# Convert element to index
|
||||||
idx2 = element_to_index(meta2, n_rows=5, rows_col_name="scan")
|
idx2 = storage.element_to_index(
|
||||||
|
element_to_store2, n_rows=5, rows_col_name="scan"
|
||||||
|
)
|
||||||
# Create dataframe
|
# Create dataframe
|
||||||
df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2)
|
df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2)
|
||||||
# Process metadata for storage
|
# Process metadata for storage
|
||||||
hash3, _ = process_meta(meta3)
|
hash3, meta_to_store3, element_to_store3 = process_meta(meta3)
|
||||||
# Convert element to index
|
# Convert element to index
|
||||||
idx3 = element_to_index(meta3, n_rows=5, rows_col_name="scan")
|
idx3 = storage.element_to_index(
|
||||||
|
element_to_store3, n_rows=5, rows_col_name="scan"
|
||||||
|
)
|
||||||
# Create dataframe
|
# Create dataframe
|
||||||
df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3)
|
df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3)
|
||||||
# Check hash equality
|
# Check hash equality
|
||||||
assert hash1 == hash2
|
assert hash1 == hash2
|
||||||
assert hash2 == hash3
|
assert hash2 == hash3
|
||||||
# Store tables
|
# Store tables
|
||||||
storage.store_table(
|
storage.store_metadata(
|
||||||
data1, meta1, columns=["f1", "f2"], rows_col_name="scan"
|
meta_md5=hash1, element=element_to_store1, meta=meta_to_store1
|
||||||
|
)
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=hash2, element=element_to_store2, meta=meta_to_store2
|
||||||
|
)
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=hash3, element=element_to_store3, meta=meta_to_store3
|
||||||
)
|
)
|
||||||
storage.store_table(
|
storage.store_table(
|
||||||
data2, meta2, columns=["f1", "f2"], rows_col_name="scan"
|
meta_md5=hash1,
|
||||||
|
element=element_to_store1,
|
||||||
|
data=data1,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
)
|
)
|
||||||
storage.store_table(
|
storage.store_table(
|
||||||
data3, meta3, columns=["f1", "f2"], rows_col_name="scan"
|
meta_md5=hash2,
|
||||||
|
element=element_to_store2,
|
||||||
|
data=data2,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
|
)
|
||||||
|
storage.store_table(
|
||||||
|
meta_md5=hash3,
|
||||||
|
element=element_to_store3,
|
||||||
|
data=data3,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
)
|
)
|
||||||
# Check that URI does not exist yet
|
# Check that URI does not exist yet
|
||||||
assert not uri.exists()
|
assert not uri.exists()
|
||||||
|
|
@ -667,10 +802,9 @@ def test_store_multiple_output(tmp_path: Path):
|
||||||
assert uri1.exists()
|
assert uri1.exists()
|
||||||
assert uri2.exists()
|
assert uri2.exists()
|
||||||
assert uri3.exists()
|
assert uri3.exists()
|
||||||
# Store metadata
|
|
||||||
table_name = storage.store_metadata(meta1)
|
|
||||||
# Set index columns
|
# Set index columns
|
||||||
cols = ["subject", "session", "scan"]
|
cols = ["subject", "session", "scan"]
|
||||||
|
table_name = f"meta_{hash1}"
|
||||||
# Read stored tables
|
# Read stored tables
|
||||||
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)
|
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)
|
||||||
cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols)
|
cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols)
|
||||||
|
|
@ -696,18 +830,21 @@ def test_collect(tmp_path: Path) -> None:
|
||||||
# Metadata for storage
|
# Metadata for storage
|
||||||
meta1 = {
|
meta1 = {
|
||||||
"element": {"subject": "test-01", "session": "ses-01"},
|
"element": {"subject": "test-01", "session": "ses-01"},
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fc"},
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
meta2 = {
|
meta2 = {
|
||||||
"element": {"subject": "test-02", "session": "ses-01"},
|
"element": {"subject": "test-02", "session": "ses-01"},
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fc"},
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
meta3 = {
|
meta3 = {
|
||||||
"element": {"subject": "test-01", "session": "ses-02"},
|
"element": {"subject": "test-01", "session": "ses-02"},
|
||||||
"version": "0.0.1",
|
"dependencies": ["numpy"],
|
||||||
"marker": {"name": "fc"},
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
# Data for storage
|
# Data for storage
|
||||||
data1 = np.array(
|
data1 = np.array(
|
||||||
|
|
@ -722,14 +859,38 @@ def test_collect(tmp_path: Path) -> None:
|
||||||
data2 = data1 * 10
|
data2 = data1 * 10
|
||||||
data3 = data1 * 20
|
data3 = data1 * 20
|
||||||
# Store tables
|
# Store tables
|
||||||
storage.store_table(
|
hash1, meta_to_store1, element_to_store1 = process_meta(meta1)
|
||||||
data1, meta1, columns=["f1", "f2"], rows_col_name="scan"
|
hash2, meta_to_store2, element_to_store2 = process_meta(meta2)
|
||||||
|
hash3, meta_to_store3, element_to_store3 = process_meta(meta3)
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=hash1, element=element_to_store1, meta=meta_to_store1
|
||||||
|
)
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=hash2, element=element_to_store2, meta=meta_to_store2
|
||||||
|
)
|
||||||
|
storage.store_metadata(
|
||||||
|
meta_md5=hash3, element=element_to_store3, meta=meta_to_store3
|
||||||
)
|
)
|
||||||
storage.store_table(
|
storage.store_table(
|
||||||
data2, meta2, columns=["f1", "f2"], rows_col_name="scan"
|
meta_md5=hash1,
|
||||||
|
element=element_to_store1,
|
||||||
|
data=data1,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
)
|
)
|
||||||
storage.store_table(
|
storage.store_table(
|
||||||
data3, meta3, columns=["f1", "f2"], rows_col_name="scan"
|
meta_md5=hash2,
|
||||||
|
element=element_to_store2,
|
||||||
|
data=data2,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
|
)
|
||||||
|
storage.store_table(
|
||||||
|
meta_md5=hash3,
|
||||||
|
element=element_to_store3,
|
||||||
|
data=data3,
|
||||||
|
columns=["f1", "f2"],
|
||||||
|
rows_col_name="scan",
|
||||||
)
|
)
|
||||||
# Convert element to prefix
|
# Convert element to prefix
|
||||||
prefix1 = element_to_prefix(meta1["element"])
|
prefix1 = element_to_prefix(meta1["element"])
|
||||||
|
|
@ -752,7 +913,7 @@ def test_collect(tmp_path: Path) -> None:
|
||||||
# Set index columns
|
# Set index columns
|
||||||
cols = ["subject", "session", "scan"]
|
cols = ["subject", "session", "scan"]
|
||||||
# Store metadata
|
# Store metadata
|
||||||
table_name = storage.store_metadata(meta1)
|
table_name = f"meta_{hash1}"
|
||||||
# Read stored tables
|
# Read stored tables
|
||||||
all_df = _read_sql(table_name, uri.as_posix(), index_col=cols)
|
all_df = _read_sql(table_name, uri.as_posix(), index_col=cols)
|
||||||
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)
|
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,7 @@ from junifer.storage.base import BaseFeatureStorage
|
||||||
def test_BaseFeatureStorage_abstractness() -> None:
|
def test_BaseFeatureStorage_abstractness() -> None:
|
||||||
"""Test BaseFeatureStorage is abstract base class."""
|
"""Test BaseFeatureStorage is abstract base class."""
|
||||||
with pytest.raises(TypeError, match=r"abstract"):
|
with pytest.raises(TypeError, match=r"abstract"):
|
||||||
BaseFeatureStorage(uri="/tmp", storage_types=["matrix"])
|
BaseFeatureStorage(uri="/tm", storage_types=["matrix"]) # type: ignore
|
||||||
|
|
||||||
|
|
||||||
def test_BaseFeatureStorage() -> None:
|
def test_BaseFeatureStorage() -> None:
|
||||||
|
|
@ -22,13 +22,16 @@ def test_BaseFeatureStorage() -> None:
|
||||||
"""Implement concrete class."""
|
"""Implement concrete class."""
|
||||||
|
|
||||||
def __init__(self, uri, single_output=False):
|
def __init__(self, uri, single_output=False):
|
||||||
storage_types = ["matrix"]
|
storage_types = ["matrix", "table", "timeseries"]
|
||||||
super().__init__(
|
super().__init__(
|
||||||
uri=uri,
|
uri=uri,
|
||||||
storage_types=storage_types,
|
storage_types=storage_types,
|
||||||
single_output=single_output,
|
single_output=single_output,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def get_valid_inputs(self):
|
||||||
|
return ["matrix", "table", "timeseries"]
|
||||||
|
|
||||||
def list_features(self):
|
def list_features(self):
|
||||||
super().list_features()
|
super().list_features()
|
||||||
|
|
||||||
|
|
@ -38,8 +41,8 @@ def test_BaseFeatureStorage() -> None:
|
||||||
feature_md5=feature_md5,
|
feature_md5=feature_md5,
|
||||||
)
|
)
|
||||||
|
|
||||||
def store_metadata(self, metadata):
|
def store_metadata(self, meta_md5, meta, element):
|
||||||
super().store_metadata(metadata)
|
super().store_metadata(meta_md5, meta, element)
|
||||||
|
|
||||||
def collect(self):
|
def collect(self):
|
||||||
return super().collect()
|
return super().collect()
|
||||||
|
|
@ -55,7 +58,7 @@ def test_BaseFeatureStorage() -> None:
|
||||||
st.validate(input_=["matrix"])
|
st.validate(input_=["matrix"])
|
||||||
# Check validate with invalid argument
|
# Check validate with invalid argument
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
st.validate(input_=["table"])
|
st.validate(input_=["duck"])
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.list_features()
|
st.list_features()
|
||||||
|
|
@ -63,22 +66,31 @@ def test_BaseFeatureStorage() -> None:
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.read_df(None)
|
st.read_df(None)
|
||||||
|
|
||||||
|
element = {"subject": "test"}
|
||||||
|
dependencies = ["numpy"]
|
||||||
|
meta = {
|
||||||
|
"element": element,
|
||||||
|
"dependencies": dependencies,
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.store_metadata(None)
|
st.store(kind="matrix", meta=meta)
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError):
|
||||||
|
st.store_metadata("md5", meta=meta, element={})
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.collect()
|
st.collect()
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.store(kind="matrix")
|
st.store(kind="timeseries", meta=meta)
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.store(kind="timeseries")
|
st.store(kind="table", meta=meta)
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
st.store(kind="table")
|
|
||||||
|
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
st.store(kind="lego")
|
st.store(kind="lego", meta=meta)
|
||||||
|
|
||||||
assert st.uri == "/tmp"
|
assert st.uri == "/tmp"
|
||||||
|
|
|
||||||
|
|
@ -4,17 +4,51 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List, Tuple, Union
|
from typing import Dict, List
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from junifer.storage.utils import (
|
from junifer.storage.utils import (
|
||||||
element_to_index,
|
|
||||||
element_to_prefix,
|
element_to_prefix,
|
||||||
|
get_dependency_version,
|
||||||
process_meta,
|
process_meta,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"dependency, max_version",
|
||||||
|
[
|
||||||
|
("click", "8.2"),
|
||||||
|
("numpy", "1.24"),
|
||||||
|
("datalad", "0.18"),
|
||||||
|
("pandas", "1.6"),
|
||||||
|
("nibabel", "4.1"),
|
||||||
|
("nilearn", "1.0"),
|
||||||
|
("sqlalchemy", "1.5.0"),
|
||||||
|
("pyyaml", "7.0"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_get_dependency_version(dependency: str, max_version: str) -> None:
|
||||||
|
"""Test dependency resolution for installed dependencies.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
dependency : str
|
||||||
|
The parametrized dependency name.
|
||||||
|
max_version : str
|
||||||
|
The parametrized maximum version of the dependency.
|
||||||
|
|
||||||
|
"""
|
||||||
|
version = get_dependency_version(dependency)
|
||||||
|
assert version < max_version
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_dependency_version_invalid() -> None:
|
||||||
|
"""Test invalid package name handling for dependency resolution."""
|
||||||
|
with pytest.raises(ValueError, match="Could not obtain"):
|
||||||
|
get_dependency_version("foobar")
|
||||||
|
|
||||||
|
|
||||||
def test_process_meta_invalid_metadata_type() -> None:
|
def test_process_meta_invalid_metadata_type() -> None:
|
||||||
"""Test invalid metadata type check for metadata hash processing."""
|
"""Test invalid metadata type check for metadata hash processing."""
|
||||||
meta = None
|
meta = None
|
||||||
|
|
@ -25,58 +59,133 @@ def test_process_meta_invalid_metadata_type() -> None:
|
||||||
# TODO: parameterize
|
# TODO: parameterize
|
||||||
def test_process_meta_hash() -> None:
|
def test_process_meta_hash() -> None:
|
||||||
"""Test metadata hash processing."""
|
"""Test metadata hash processing."""
|
||||||
meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}
|
meta = {
|
||||||
hash1, _ = process_meta(meta)
|
"element": {"foo": "bar"},
|
||||||
|
"A": 1,
|
||||||
|
"B": [2, 3, 4, 5, 6],
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
hash1, _, element1 = process_meta(meta)
|
||||||
|
assert element1 == {"foo": "bar"}
|
||||||
|
|
||||||
meta = {"element": "foo", "B": [2, 3, 4, 5, 6], "A": 1}
|
meta = {
|
||||||
hash2, _ = process_meta(meta)
|
"element": {"foo": "baz"},
|
||||||
|
"B": [2, 3, 4, 5, 6],
|
||||||
|
"A": 1,
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
hash2, _, element2 = process_meta(meta)
|
||||||
assert hash1 == hash2
|
assert hash1 == hash2
|
||||||
|
assert element2 == {"foo": "baz"}
|
||||||
|
|
||||||
meta = {"element": "foo", "A": 1, "B": [2, 3, 1, 5, 6]}
|
meta = {
|
||||||
hash3, _ = process_meta(meta)
|
"element": {"foo": "bar"},
|
||||||
|
"A": 1,
|
||||||
|
"B": [2, 3, 1, 5, 6],
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
hash3, _, element3 = process_meta(meta)
|
||||||
assert hash1 != hash3
|
assert hash1 != hash3
|
||||||
|
assert element3 == element1
|
||||||
|
|
||||||
meta1 = {
|
meta4 = {
|
||||||
"element": "foo",
|
"element": {"foo": "bar"},
|
||||||
"B": {
|
"B": {
|
||||||
"B2": [2, 3, 4, 5, 6],
|
"B2": [2, 3, 4, 5, 6],
|
||||||
"B1": [9.22, 3.14, 1.41, 5.67, 6.28],
|
"B1": [9.22, 3.14, 1.41, 5.67, 6.28],
|
||||||
"B3": (1, "car"),
|
"B3": (1, "car"),
|
||||||
},
|
},
|
||||||
"A": 1,
|
"A": 1,
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
|
|
||||||
meta2 = {
|
meta5 = {
|
||||||
"A": 1,
|
"A": 1,
|
||||||
"B": {
|
"B": {
|
||||||
"B3": (1, "car"),
|
"B3": (1, "car"),
|
||||||
"B1": [9.22, 3.14, 1.41, 5.67, 6.28],
|
"B1": [9.22, 3.14, 1.41, 5.67, 6.28],
|
||||||
"B2": [2, 3, 4, 5, 6],
|
"B2": [2, 3, 4, 5, 6],
|
||||||
},
|
},
|
||||||
"element": "foo",
|
"element": {"foo": "baz"},
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
}
|
}
|
||||||
|
|
||||||
hash4, _ = process_meta(meta1)
|
hash4, _, _ = process_meta(meta4)
|
||||||
hash5, _ = process_meta(meta2)
|
hash5, _, _ = process_meta(meta5)
|
||||||
assert hash4 == hash5
|
assert hash4 == hash5
|
||||||
|
|
||||||
|
# Different element keys should give a different hash
|
||||||
|
meta6 = {
|
||||||
|
"A": 1,
|
||||||
|
"B": {
|
||||||
|
"B3": (1, "car"),
|
||||||
|
"B1": [9.22, 3.14, 1.41, 5.67, 6.28],
|
||||||
|
"B2": [2, 3, 4, 5, 6],
|
||||||
|
},
|
||||||
|
"element": {"bar": "baz"},
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
}
|
||||||
|
hash6, _, _ = process_meta(meta6)
|
||||||
|
assert hash4 != hash6
|
||||||
|
|
||||||
|
|
||||||
def test_process_meta_invalid_metadata_key() -> None:
|
def test_process_meta_invalid_metadata_key() -> None:
|
||||||
"""Test invalid metadata key check for metadata hash processing."""
|
"""Test invalid metadata key check for metadata hash processing."""
|
||||||
meta = {}
|
meta = {}
|
||||||
with pytest.raises(ValueError, match=r"_element_keys"):
|
with pytest.raises(ValueError, match=r"element"):
|
||||||
|
process_meta(meta)
|
||||||
|
|
||||||
|
meta = {"element": {}}
|
||||||
|
with pytest.raises(ValueError, match=r"marker"):
|
||||||
|
process_meta(meta)
|
||||||
|
|
||||||
|
meta = {"element": {}, "marker": {}}
|
||||||
|
with pytest.raises(ValueError, match=r"key 'name'"):
|
||||||
|
process_meta(meta)
|
||||||
|
|
||||||
|
meta = {"element": {}, "marker": {"name": "test"}}
|
||||||
|
with pytest.raises(ValueError, match=r"key 'type'"):
|
||||||
|
process_meta(meta)
|
||||||
|
|
||||||
|
meta = {"element": {}, "marker": {"name": "test"}, "type": "BOLD"}
|
||||||
|
with pytest.raises(ValueError, match=r"dependencies"):
|
||||||
process_meta(meta)
|
process_meta(meta)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"meta,elements",
|
"meta,elements",
|
||||||
[
|
[
|
||||||
({"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}, ["element"]),
|
(
|
||||||
|
{
|
||||||
|
"element": {"foo": "bar"},
|
||||||
|
"A": 1,
|
||||||
|
"B": [2, 3, 4, 5, 6],
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
|
},
|
||||||
|
["foo"],
|
||||||
|
),
|
||||||
(
|
(
|
||||||
{
|
{
|
||||||
"element": {"subject": "foo", "session": "bar"},
|
"element": {"subject": "foo", "session": "bar"},
|
||||||
"B": [2, 3, 4, 5, 6],
|
"B": [2, 3, 4, 5, 6],
|
||||||
"A": 1,
|
"A": 1,
|
||||||
|
"dependencies": ["numpy"],
|
||||||
|
"marker": {"name": "fc"},
|
||||||
|
"type": "BOLD",
|
||||||
},
|
},
|
||||||
["subject", "session"],
|
["subject", "session"],
|
||||||
),
|
),
|
||||||
|
|
@ -93,34 +202,30 @@ def test_process_meta_element(meta: Dict, elements: List[str]) -> None:
|
||||||
The parametrized elements to assert against.
|
The parametrized elements to assert against.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
hash1, processed_meta = process_meta(meta)
|
hash1, processed_meta, _ = process_meta(meta)
|
||||||
assert "_element_keys" in processed_meta
|
assert "_element_keys" in processed_meta
|
||||||
assert processed_meta["_element_keys"] == elements
|
assert processed_meta["_element_keys"] == elements
|
||||||
assert "A" in processed_meta
|
assert "A" in processed_meta
|
||||||
assert "B" in processed_meta
|
assert "B" in processed_meta
|
||||||
assert "element" not in processed_meta
|
assert "element" not in processed_meta
|
||||||
hash2, processed_meta2 = process_meta(processed_meta)
|
assert isinstance(processed_meta["dependencies"], Dict)
|
||||||
|
assert all(
|
||||||
assert hash1, hash2
|
x in processed_meta["dependencies"] for x in meta["dependencies"]
|
||||||
assert processed_meta == processed_meta2
|
)
|
||||||
|
assert "name" in processed_meta
|
||||||
|
assert processed_meta["name"] == f"{meta['type']}_{meta['marker']['name']}"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"element,prefix",
|
"element,prefix",
|
||||||
[
|
[
|
||||||
("sub-01", "element_sub-01_"),
|
|
||||||
(1, "element_1_"),
|
|
||||||
({"subject": "sub-01"}, "element_sub-01_"),
|
({"subject": "sub-01"}, "element_sub-01_"),
|
||||||
({"subject": 1}, "element_1_"),
|
({"subject": 1}, "element_1_"),
|
||||||
({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"),
|
({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"),
|
||||||
({"subject": 1, "session": 2}, "element_1_2_"),
|
({"subject": 1, "session": 2}, "element_1_2_"),
|
||||||
(("sub-01", "ses-02"), "element_sub-01_ses-02_"),
|
|
||||||
((1, 2), "element_1_2_"),
|
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_element_to_prefix(
|
def test_element_to_prefix(element: Dict, prefix: str) -> None:
|
||||||
element: Union[str, int, Dict, Tuple], prefix: str
|
|
||||||
) -> None:
|
|
||||||
"""Test converting element to prefix (for file naming).
|
"""Test converting element to prefix (for file naming).
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -138,75 +243,5 @@ def test_element_to_prefix(
|
||||||
def test_element_to_prefix_invalid_type() -> None:
|
def test_element_to_prefix_invalid_type() -> None:
|
||||||
"""Test element to prefix type checking."""
|
"""Test element to prefix type checking."""
|
||||||
element = 2.3
|
element = 2.3
|
||||||
with pytest.raises(ValueError, match=r"convert element of type"):
|
with pytest.raises(ValueError, match=r"must be a dict"):
|
||||||
element_to_prefix(element) # type: ignore
|
element_to_prefix(element) # type: ignore
|
||||||
|
|
||||||
|
|
||||||
def test_element_to_index_check_meta_invalid_key() -> None:
|
|
||||||
"""Test element to index metadata key checking."""
|
|
||||||
meta = {"noelement": "foo"}
|
|
||||||
with pytest.raises(ValueError, match=r"metadata must contain the key"):
|
|
||||||
element_to_index(meta)
|
|
||||||
|
|
||||||
|
|
||||||
def test_element_to_index() -> None:
|
|
||||||
"""Test element to index."""
|
|
||||||
meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}
|
|
||||||
index = element_to_index(meta)
|
|
||||||
assert index.names == ["element", "idx"]
|
|
||||||
assert index.levels[0].name == "element"
|
|
||||||
assert index.levels[0].values[0] == "foo"
|
|
||||||
assert all(x == "foo" for x in index.levels[0].values)
|
|
||||||
assert index.levels[0].values.shape == (1,)
|
|
||||||
assert index.levels[1].name == "idx"
|
|
||||||
assert all(x == i for i, x in enumerate(index.levels[1].values))
|
|
||||||
assert index.levels[1].values.shape == (1,)
|
|
||||||
|
|
||||||
index = element_to_index(meta, n_rows=10)
|
|
||||||
assert index.names == ["element", "idx"]
|
|
||||||
assert index.levels[0].name == "element"
|
|
||||||
assert all(x == "foo" for x in index.levels[0].values)
|
|
||||||
assert index.levels[0].values.shape == (1,)
|
|
||||||
|
|
||||||
assert index.levels[1].name == "idx"
|
|
||||||
assert all(x == i for i, x in enumerate(index.levels[1].values))
|
|
||||||
assert index.levels[1].values.shape == (10,)
|
|
||||||
|
|
||||||
index = element_to_index(meta, n_rows=1, rows_col_name="scan")
|
|
||||||
assert index.names == ["element", "scan"]
|
|
||||||
assert index.levels[0].name == "element"
|
|
||||||
assert index.levels[0].values[0] == "foo"
|
|
||||||
assert all(x == "foo" for x in index.levels[0].values)
|
|
||||||
assert index.levels[0].values.shape == (1,)
|
|
||||||
assert index.levels[1].name == "scan"
|
|
||||||
assert all(x == i for i, x in enumerate(index.levels[1].values))
|
|
||||||
assert index.levels[1].values.shape == (1,)
|
|
||||||
|
|
||||||
index = element_to_index(meta, n_rows=7, rows_col_name="scan")
|
|
||||||
assert index.names == ["element", "scan"]
|
|
||||||
assert index.levels[0].name == "element"
|
|
||||||
assert all(x == "foo" for x in index.levels[0].values)
|
|
||||||
assert index.levels[0].values.shape == (1,)
|
|
||||||
|
|
||||||
assert index.levels[1].name == "scan"
|
|
||||||
assert all(x == i for i, x in enumerate(index.levels[1].values))
|
|
||||||
assert index.levels[1].values.shape == (7,)
|
|
||||||
|
|
||||||
meta = {
|
|
||||||
"element": {"subject": "sub-01", "session": "ses-01"},
|
|
||||||
"A": 1,
|
|
||||||
"B": [2, 3, 4, 5, 6],
|
|
||||||
}
|
|
||||||
index = element_to_index(meta, n_rows=10)
|
|
||||||
|
|
||||||
assert index.levels[0].name == "subject"
|
|
||||||
assert all(x == "sub-01" for x in index.levels[0].values)
|
|
||||||
assert index.levels[0].values.shape == (1,)
|
|
||||||
|
|
||||||
assert index.levels[1].name == "session"
|
|
||||||
assert all(x == "ses-01" for x in index.levels[1].values)
|
|
||||||
assert index.levels[1].values.shape == (1,)
|
|
||||||
|
|
||||||
assert index.levels[2].name == "idx"
|
|
||||||
assert all(x == i for i, x in enumerate(index.levels[2].values))
|
|
||||||
assert index.levels[2].values.shape == (10,)
|
|
||||||
|
|
|
||||||
|
|
@ -6,14 +6,39 @@
|
||||||
|
|
||||||
import hashlib
|
import hashlib
|
||||||
import json
|
import json
|
||||||
from typing import Any, Dict, Optional, Tuple, Union
|
from importlib.metadata import PackageNotFoundError, version
|
||||||
|
from typing import Dict, Tuple
|
||||||
import numpy as np
|
|
||||||
import pandas as pd
|
|
||||||
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
|
def get_dependency_version(dependency: str) -> str:
|
||||||
|
"""Get dependency version.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
dependency : str
|
||||||
|
The dependency to fetch version for.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
str
|
||||||
|
The version of the dependency.
|
||||||
|
|
||||||
|
"""
|
||||||
|
dep_version = ""
|
||||||
|
try:
|
||||||
|
dep_version = version(dependency)
|
||||||
|
except PackageNotFoundError as e:
|
||||||
|
raise_error(
|
||||||
|
f"Could not obtain the version of {dependency}. "
|
||||||
|
"Have you specified the DEPENDENCIES variable correctly?",
|
||||||
|
exception=e,
|
||||||
|
)
|
||||||
|
|
||||||
|
return dep_version
|
||||||
|
|
||||||
|
|
||||||
def _meta_hash(meta: Dict) -> str:
|
def _meta_hash(meta: Dict) -> str:
|
||||||
"""Compute the MD5 hash of the metadata.
|
"""Compute the MD5 hash of the metadata.
|
||||||
|
|
||||||
|
|
@ -29,6 +54,12 @@ def _meta_hash(meta: Dict) -> str:
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Hashing metadata: {meta}")
|
logger.debug(f"Hashing metadata: {meta}")
|
||||||
|
if "dependencies" not in meta:
|
||||||
|
raise_error("The metadata must contain the key 'dependencies'")
|
||||||
|
# Convert dependencies set into {dependency: version} dictionary
|
||||||
|
meta["dependencies"] = {
|
||||||
|
dep: get_dependency_version(dep) for dep in meta["dependencies"]
|
||||||
|
}
|
||||||
meta_md5 = hashlib.md5(
|
meta_md5 = hashlib.md5(
|
||||||
json.dumps(meta, sort_keys=True).encode("utf-8")
|
json.dumps(meta, sort_keys=True).encode("utf-8")
|
||||||
).hexdigest()
|
).hexdigest()
|
||||||
|
|
@ -36,7 +67,7 @@ def _meta_hash(meta: Dict) -> str:
|
||||||
return meta_md5
|
return meta_md5
|
||||||
|
|
||||||
|
|
||||||
def process_meta(meta: Dict) -> Tuple[str, Dict]:
|
def process_meta(meta: Dict) -> Tuple[str, Dict, Dict]:
|
||||||
"""Process the metadata for storage.
|
"""Process the metadata for storage.
|
||||||
|
|
||||||
It removes the key "element" and adds the "_element_keys" with the keys
|
It removes the key "element" and adds the "_element_keys" with the keys
|
||||||
|
|
@ -53,12 +84,13 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]:
|
||||||
The MD5 hash of the metadata.
|
The MD5 hash of the metadata.
|
||||||
dict
|
dict
|
||||||
The processed metadata for storage.
|
The processed metadata for storage.
|
||||||
|
tuple
|
||||||
|
The element.
|
||||||
|
|
||||||
Raises
|
Raises
|
||||||
------
|
------
|
||||||
ValueError
|
ValueError
|
||||||
If `meta` is None or if it does not contain the key "element" or
|
If `meta` is None or if it does not contain the key "element".
|
||||||
"_element_keys".
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
if meta is None:
|
if meta is None:
|
||||||
|
|
@ -68,97 +100,42 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]:
|
||||||
# Remove key "element"
|
# Remove key "element"
|
||||||
element = t_meta.pop("element", None)
|
element = t_meta.pop("element", None)
|
||||||
if element is None:
|
if element is None:
|
||||||
if "_element_keys" not in t_meta:
|
raise_error(msg="`meta` must contain the key 'element'")
|
||||||
raise_error(
|
if "marker" not in t_meta:
|
||||||
msg="`meta` must contain the key 'element' or '_element_keys'"
|
raise_error(msg="`meta` must contain the key 'marker'")
|
||||||
)
|
if "name" not in t_meta["marker"]:
|
||||||
else:
|
raise_error(msg="`meta['marker']` must contain the key 'name'")
|
||||||
if isinstance(element, dict):
|
if "type" not in t_meta:
|
||||||
t_meta["_element_keys"] = list(element.keys())
|
raise_error(msg="`meta` must contain the key 'type'")
|
||||||
else:
|
|
||||||
t_meta["_element_keys"] = ["element"]
|
t_meta["_element_keys"] = list(element.keys())
|
||||||
|
type_ = t_meta["type"]
|
||||||
|
name = t_meta["marker"]["name"]
|
||||||
|
t_meta["name"] = f"{type_}_{name}"
|
||||||
# MD5 hash of the metadata
|
# MD5 hash of the metadata
|
||||||
md5_hash = _meta_hash(t_meta)
|
md5_hash = _meta_hash(t_meta)
|
||||||
return md5_hash, t_meta
|
return md5_hash, t_meta, element
|
||||||
|
|
||||||
|
|
||||||
def element_to_prefix(element: Union[Tuple, Dict, str, int]) -> str:
|
def element_to_prefix(element: Dict) -> str:
|
||||||
"""Convert the element metadata to prefix.
|
"""Convert the element metadata to prefix.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
element : tuple, dict, str or int
|
element : dict
|
||||||
The element to convert to prefix.
|
The element to convert to prefix.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
str
|
str
|
||||||
The element converted to prefix.
|
The element converted to prefix.
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
ValueError
|
|
||||||
If invalid type is passed for `element`.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Converting element {element} to prefix.")
|
logger.debug(f"Converting element {element} to prefix.")
|
||||||
prefix = "element"
|
prefix = "element"
|
||||||
if isinstance(element, tuple):
|
if not isinstance(element, dict):
|
||||||
prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}"
|
raise_error(msg="`element` must be a dict")
|
||||||
elif isinstance(element, dict):
|
|
||||||
prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}"
|
prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}"
|
||||||
elif isinstance(element, (str, int)):
|
|
||||||
prefix = f"{prefix}_{element}"
|
|
||||||
else:
|
|
||||||
raise_error(
|
|
||||||
f"Cannot convert element of type {type(element)} to prefix. "
|
|
||||||
"Must be a str, int, tuple or dict."
|
|
||||||
)
|
|
||||||
logger.debug(f"Converted prefix: {prefix}")
|
logger.debug(f"Converted prefix: {prefix}")
|
||||||
return f"{prefix}_"
|
return f"{prefix}_"
|
||||||
|
|
||||||
|
|
||||||
def element_to_index(
|
|
||||||
meta: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None
|
|
||||||
) -> pd.MultiIndex:
|
|
||||||
"""Convert the element metadata to index.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary. Must contain the key "element"."
|
|
||||||
n_rows : int, optional
|
|
||||||
Number of rows to create (default 1).
|
|
||||||
rows_col_name: str, optional
|
|
||||||
The column name to use in case `n_rows` > 1. If None and
|
|
||||||
n_rows > 1, the name will be "idx" (default None).
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
pandas.MultiIndex
|
|
||||||
The index of the dataframe to store.
|
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
ValueError
|
|
||||||
If `meta` does not contain the key "element".
|
|
||||||
|
|
||||||
"""
|
|
||||||
if "element" not in meta:
|
|
||||||
raise_error(
|
|
||||||
msg="To create and index, metadata must contain the key 'element'."
|
|
||||||
)
|
|
||||||
# Get element
|
|
||||||
element = meta["element"]
|
|
||||||
if not isinstance(element, dict):
|
|
||||||
element = {"element": element}
|
|
||||||
# Check rows_col_name
|
|
||||||
if rows_col_name is None:
|
|
||||||
rows_col_name = "idx"
|
|
||||||
elem_idx: Dict[Any, Any] = {k: [v] * n_rows for k, v in element.items()}
|
|
||||||
elem_idx[rows_col_name] = np.arange(n_rows)
|
|
||||||
# Create index
|
|
||||||
index = pd.MultiIndex.from_frame(
|
|
||||||
pd.DataFrame(elem_idx, index=range(n_rows))
|
|
||||||
)
|
|
||||||
return index
|
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ from .datagrabbers import (
|
||||||
SPMAuditoryTestingDatagrabber,
|
SPMAuditoryTestingDatagrabber,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Register testing datagrabber
|
# Register testing datagrabber
|
||||||
register(
|
register(
|
||||||
step="datagrabber",
|
step="datagrabber",
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,7 @@ from warnings import warn
|
||||||
|
|
||||||
import datalad
|
import datalad
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger("JUNIFER")
|
logger = logging.getLogger("JUNIFER")
|
||||||
|
|
||||||
# Set up datalad logger level to warning by default
|
# Set up datalad logger level to warning by default
|
||||||
|
|
@ -262,7 +263,11 @@ def configure_logging(
|
||||||
log_versions() # log versions of installed packages
|
log_versions() # log versions of installed packages
|
||||||
|
|
||||||
|
|
||||||
def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn:
|
def raise_error(
|
||||||
|
msg: str,
|
||||||
|
klass: Type[Exception] = ValueError,
|
||||||
|
exception: Optional[Exception] = None,
|
||||||
|
) -> NoReturn:
|
||||||
"""Raise error, but first log it.
|
"""Raise error, but first log it.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -271,10 +276,15 @@ def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn:
|
||||||
The message for the exception.
|
The message for the exception.
|
||||||
klass : subclass of Exception, optional
|
klass : subclass of Exception, optional
|
||||||
The subclass of Exception to raise using (default ValueError).
|
The subclass of Exception to raise using (default ValueError).
|
||||||
|
exception : Exception, optional
|
||||||
|
The original exception to follow up on (default None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.error(msg)
|
logger.error(msg)
|
||||||
raise klass(msg)
|
if exception is not None:
|
||||||
|
raise klass(msg) from exception
|
||||||
|
else:
|
||||||
|
raise klass(msg)
|
||||||
|
|
||||||
|
|
||||||
def warn_with_log(
|
def warn_with_log(
|
||||||
|
|
|
||||||
|
|
@ -11,7 +11,7 @@ name = "junifer"
|
||||||
description = "JUelich NeuroImaging FEature extractoR"
|
description = "JUelich NeuroImaging FEature extractoR"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.8"
|
requires-python = ">=3.8"
|
||||||
license = {file = "LICENSE.md"}
|
license = {text = "AGPL-3.0-only"}
|
||||||
authors = [
|
authors = [
|
||||||
{name = "Fede Raimondo", email = "f.raimondo@fz-juelich.de"},
|
{name = "Fede Raimondo", email = "f.raimondo@fz-juelich.de"},
|
||||||
{name = "Synchon Mandal", email = "s.mandal@fz-juelich.de"},
|
{name = "Synchon Mandal", email = "s.mandal@fz-juelich.de"},
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue