[ENH]: Rename 'table' to 'vector' in storage API #181

Merged
synchon merged 51 commits from refactor/storage/table-to-vector into main 2023-02-17 12:51:26 +00:00
30 changed files with 381 additions and 326 deletions

View file

@ -49,4 +49,8 @@ API changes
~~~~~~~~~~~
- Add ``confounds_format`` parameter to :class:`junifer.datagrabber.PatternDataGrabber` constructor for
improved handling of confounds specified via ``BOLD_confounds`` data type (:gh:`158` by `Synchon Mandal`_).
improved handling of confounds specified via ``BOLD_confounds`` data type (:gh:`158` by `Synchon Mandal`_).
- Rename ``store_table()`` to ``store_vector()`` for storage-like objects and adapt marker-like objects to
use ``"vector"`` in place of ``"table"`` for storage. Also, improve the logic of storing vectors
(:gh:`181` by `Synchon Mandal`_).

View file

@ -15,10 +15,10 @@ that object else they are kept in memory.
Storage is meant to be used inside the datagrabber context but you can operate on them outside the context as long
as the processed data is in the memory and the Python runtime has not garbage-collected it.
The :ref:`Markers <marker>` are responsible for defining what *storage kind* (``matrix``, ``table``, ``timeseries``)
The :ref:`Markers <marker>` are responsible for defining what *storage kind* (``matrix``, ``vector``, ``timeseries``)
they support for which :ref:`data type <data_types>` by overriding its ``store`` method. The storage object in turn
declares and provides implementation for specific *storage kind*. For example, :class:`junifer.storage.SQLiteFeatureStorage`
supports saving ``matrix``, ``table`` and ``timeseries`` via ``store_matrix``, ``store_table`` and ``store_timeseries``
supports saving ``matrix``, ``vector`` and ``timeseries`` via ``store_matrix``, ``store_vector`` and ``store_timeseries``
methods respectively.
For storage interfaces not supported by junifer yet, you can either make your own ``Storage`` by providing a concrete
@ -42,15 +42,15 @@ Currently supported storage types
- A 2D matrix with row and column names
- ``col_names``, ``row_names``, ``matrix_kind``, ``diagonal``
- :meth:`junifer.storage.BaseFeatureStorage.store_matrix`
* - ``table``
* - ``vector``
- A vector of values with column names
- ``columns``, ``row_names``
- :meth:`junifer.storage.BaseFeatureStorage.store_table`
- :meth:`junifer.storage.BaseFeatureStorage.store_vector`
* - ``timeseries``
- A 2D matrix of values with column names
- ``columns``, ``row_names``
- :meth:`junifer.storage.BaseFeatureStorage.store_timeseries`
.. _storage_interfaces:
Currently supported storage interfaces

View file

@ -7,24 +7,23 @@ from pathlib import Path
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
List,
Optional,
Tuple,
Union,
Callable,
)
import numpy as np
import nibabel as nib
from nilearn.masking import (
compute_brain_mask,
compute_background_mask,
compute_epi_mask,
)
import numpy as np
from nilearn.datasets import fetch_icbm152_brain_gm_mask
from nilearn.image import resample_to_img
from nilearn.masking import (
compute_background_mask,
compute_brain_mask,
compute_epi_mask,
)
from ..utils.logging import logger, raise_error
from .utils import closest_resolution

View file

@ -5,33 +5,27 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from pathlib import Path
from typing import Callable, Dict, Union
from pathlib import Path
import pytest
from numpy.testing import (
assert_array_almost_equal,
assert_array_equal,
)
from nilearn.datasets import fetch_icbm152_brain_gm_mask
from nilearn.image import resample_to_img
from nilearn.masking import (
compute_brain_mask,
compute_background_mask,
compute_brain_mask,
compute_epi_mask,
)
from nilearn.datasets import fetch_icbm152_brain_gm_mask
from numpy.testing import assert_array_almost_equal, assert_array_equal
from junifer.data.masks import (
_available_masks,
_load_vickery_patil_mask,
get_mask,
list_masks,
load_mask,
register_mask,
get_mask,
_available_masks,
)
from junifer.datareader import DefaultDataReader
from junifer.testing.datagrabbers import (
OasisVBMTestingDatagrabber,

View file

@ -3,11 +3,11 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Optional, Iterable
from typing import Iterable, Optional
import pytest
from junifer.datagrabber.hcp import DataladHCP1200, HCP1200
from junifer.datagrabber.hcp import HCP1200, DataladHCP1200
from junifer.utils import configure_logging
@ -332,9 +332,7 @@ def test_hcp1200_datagrabber_elements(
)
with dg:
# Get all elements
expected_subjects = [
f"sub-{x:02d}" for x in range(1, 10)
]
expected_subjects = [f"sub-{x:02d}" for x in range(1, 10)]
found_subjects = []
all_elements = dg.get_elements()
# Check only specified task and phase encoding are found

View file

@ -110,8 +110,7 @@ class RSSETSMarker(BaseMarker):
keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list
* ``row_names`` (if more than one row is present in data): "scan"
* ``col_names`` : the column labels for the computed values as list
References
----------
@ -135,5 +134,5 @@ class RSSETSMarker(BaseMarker):
# Compute the RSS
out["data"] = np.sum(edge_ts**2, 1) ** 0.5
# Set correct column label
out["columns"] = ["root_sum_of_squares_ets"]
out["col_names"] = ["root_sum_of_squares_ets"]
return out

View file

@ -104,7 +104,7 @@ class AmplitudeLowFrequencyFluctuationBase(BaseMarker):
The storage type output by the marker.
"""
return "table"
return "vector"
def compute(
self,
@ -132,8 +132,7 @@ class AmplitudeLowFrequencyFluctuationBase(BaseMarker):
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list
* ``row_names`` (if more than one row is present in data): "scan"
* ``col_names`` : the column labels for the computed values as list
"""
if self.use_afni is None:

View file

@ -109,7 +109,7 @@ class AmplitudeLowFrequencyFluctuationParcels(
keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list
* ``col_names`` : the column labels for the computed values as list
"""
pa = ParcelAggregation(

View file

@ -115,7 +115,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list
* ``col_names`` : the column labels for the computed values as list
"""
pa = SphereAggregation(

View file

@ -150,6 +150,6 @@ class CrossParcellationFC(BaseMarker):
return {
"data": result,
"col_names": parcellation_one_dict["columns"],
"row_names": parcellation_two_dict["columns"],
"col_names": parcellation_one_dict["col_names"],
"row_names": parcellation_two_dict["col_names"],
}

View file

@ -84,7 +84,7 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
bold_aggregated = parcel_aggregation.compute(input)
ets, edge_names = _ets(
bold_aggregated["data"], bold_aggregated["columns"]
bold_aggregated["data"], bold_aggregated["col_names"]
)
return dict(data=ets, columns=edge_names)
return {"data": ets, "col_names": edge_names}

View file

@ -91,7 +91,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
)
bold_aggregated = sphere_aggregation.compute(input)
ets, edge_names = _ets(
bold_aggregated["data"], bold_aggregated["columns"]
bold_aggregated["data"], bold_aggregated["col_names"]
)
return dict(data=ets, columns=edge_names)
return {"data": ets, "col_names": edge_names}

View file

@ -143,7 +143,7 @@ class FunctionalConnectivityBase(BaseMarker):
out = {}
out["data"] = connectivity.fit_transform([aggregation["data"]])[0]
# Create column names
out["row_names"] = aggregation["columns"]
out["col_names"] = aggregation["columns"]
out["row_names"] = aggregation["col_names"]
out["col_names"] = aggregation["col_names"]
out["matrix_kind"] = "tril"
return out

View file

@ -91,7 +91,7 @@ class ParcelAggregation(BaseMarker):
"""
if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
return "table"
return "vector"
elif input_type == "BOLD":
return "timeseries"
else:
@ -123,7 +123,7 @@ class ParcelAggregation(BaseMarker):
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list
* ``col_names`` : the column labels for the computed values as list
"""
t_input_img = input["data"]
@ -218,5 +218,5 @@ class ParcelAggregation(BaseMarker):
out_labels.append(labels[t_v - 1])
out_values = np.array(out_values).T
out = {"data": out_values, "columns": out_labels}
out = {"data": out_values, "col_names": out_labels}
return out

View file

@ -70,7 +70,7 @@ class ReHoBase(BaseMarker):
The storage type output by the marker.
"""
return "table"
return "vector"
def compute_reho_map(
self,

View file

@ -122,8 +122,7 @@ class ReHoParcels(ReHoBase):
keys:
* ``data`` : the actual computed values as a 1D numpy.ndarray
* ``columns`` : the column labels for the parcels as a list
* ``row_names`` : ``None``
* ``col_names`` : the column labels for the parcels as a list
"""
logger.info("Calculating ReHo for parcels.")

View file

@ -129,8 +129,7 @@ class ReHoSpheres(ReHoBase):
keys:
* ``data`` : the actual computed values as a 1D numpy.ndarray
* ``columns`` : the column labels for the spheres as a list
* ``rows_col_name`` : ``None``
* ``col_names`` : the column labels for the spheres as a list
"""
logger.info("Calculating ReHo for spheres.")

View file

@ -35,7 +35,7 @@ def test_reho_parcels_computation() -> None:
reho_parcels_output_bold = reho_parcels_output["BOLD"]
# Assert BOLD output keys
assert "data" in reho_parcels_output_bold
assert "columns" in reho_parcels_output_bold
assert "col_names" in reho_parcels_output_bold
reho_parcels_output_bold_data = reho_parcels_output_bold["data"]
# Assert BOLD output data dimension

View file

@ -35,7 +35,7 @@ def test_reho_spheres_computation() -> None:
reho_spheres_output_bold = reho_spheres_output["BOLD"]
# Assert BOLD output keys
assert "data" in reho_spheres_output_bold
assert "columns" in reho_spheres_output_bold
assert "col_names" in reho_spheres_output_bold
reho_spheres_output_bold_data = reho_spheres_output_bold["data"]
# Assert BOLD output data dimension

View file

@ -7,7 +7,7 @@
from typing import Any, Dict, List, Optional, Union
from ..api.decorators import register_marker
from ..data import load_coordinates, get_mask
from ..data import get_mask, load_coordinates
from ..external.nilearn import JuniferNiftiSpheresMasker
from ..stats import get_aggfunc_by_name
from ..utils import logger
@ -94,7 +94,7 @@ class SphereAggregation(BaseMarker):
"""
if input_type in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
return "table"
return "vector"
elif input_type == "BOLD":
return "timeseries"
else:
@ -126,7 +126,7 @@ class SphereAggregation(BaseMarker):
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list
* ``col_names`` : the column labels for the computed values as list
"""
t_input_img = input["data"]
@ -151,5 +151,5 @@ class SphereAggregation(BaseMarker):
# Fit and transform the marker on the data
out_values = masker.fit_transform(t_input_img)
# Format the output
out = {"data": out_values, "columns": out_labels}
out = {"data": out_values, "col_names": out_labels}
return out

View file

@ -87,7 +87,7 @@ def test_marker_collection() -> None:
assert "VBM_GM" in out[t_name]
t_vbm = out[t_name]["VBM_GM"]
assert "data" in t_vbm
assert "columns" in t_vbm
assert "col_names" in t_vbm
assert "meta" in t_vbm
# Test preprocessing
@ -199,19 +199,19 @@ def test_marker_collection_storage(tmp_path: Path) -> None:
t_feature = storage.read_df(feature_md5=feature_md5)
fname = "gmd_schaefer100x7_mean"
t_data = out[fname]["VBM_GM"]["data"] # type: ignore
cols = out[fname]["VBM_GM"]["columns"] # type: ignore
cols = out[fname]["VBM_GM"]["col_names"] # type: ignore
assert_array_equal(t_feature[cols].values, t_data) # type: ignore
feature_md5 = list(features.keys())[1]
t_feature = storage.read_df(feature_md5=feature_md5)
fname = "gmd_schaefer100x7_std"
t_data = out[fname]["VBM_GM"]["data"] # type: ignore
cols = out[fname]["VBM_GM"]["columns"] # type: ignore
cols = out[fname]["VBM_GM"]["col_names"] # type: ignore
assert_array_equal(t_feature[cols].values, t_data) # type: ignore
feature_md5 = list(features.keys())[2]
t_feature = storage.read_df(feature_md5=feature_md5)
fname = "gmd_schaefer100x7_trim_mean90"
t_data = out[fname]["VBM_GM"]["data"] # type: ignore
cols = out[fname]["VBM_GM"]["columns"] # type: ignore
cols = out[fname]["VBM_GM"]["col_names"] # type: ignore
assert_array_equal(t_feature[cols].values, t_data) # type: ignore

View file

@ -26,7 +26,7 @@ def test_ParcelAggregation_input_output() -> None:
marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", on="VBM_GM"
)
for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]:
for in_, out_ in [("VBM_GM", "vector"), ("BOLD", "timeseries")]:
assert marker.get_output_type(in_) == out_
with pytest.raises(ValueError, match="Unknown input"):
@ -364,7 +364,7 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
# Data and labels should be the same
assert_array_equal(orig_mean_data, split_mean_data)
assert orig_mean["columns"] == split_mean["columns"]
assert orig_mean["col_names"] == split_mean["col_names"]
def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
@ -446,5 +446,5 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
assert_array_equal(orig_mean_data, split_mean_data)
# Labels should be "low" for the first 50 and "high" for the second 50
assert all(x.startswith("low") for x in split_mean["columns"][:50])
assert all(x.startswith("high") for x in split_mean["columns"][50:])
assert all(x.startswith("low") for x in split_mean["col_names"][:50])
assert all(x.startswith("high") for x in split_mean["col_names"][50:])

View file

@ -27,7 +27,7 @@ RADIUS = 8
def test_SphereAggregation_input_output() -> None:
"""Test SphereAggregation input and output types."""
marker = SphereAggregation(coords="DMNBuckner", method="mean", on="VBM_GM")
for in_, out_ in [("VBM_GM", "table"), ("BOLD", "timeseries")]:
for in_, out_ in [("VBM_GM", "vector"), ("BOLD", "timeseries")]:
assert marker.get_output_type(in_) == out_
with pytest.raises(ValueError, match="Unknown input"):

View file

@ -6,7 +6,7 @@
from abc import ABC, abstractmethod
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Union
from typing import Any, Dict, Iterable, List, Optional, Union
import numpy as np
import pandas as pd
@ -57,8 +57,9 @@ class BaseFeatureStorage(ABC):
Returns
-------
list of str
The list of storage types that can be used as input for this "
"storage.
The list of storage types that can be used as input for this
storage interface.
"""
raise_error(
msg="Concrete classes need to implement get_valid_inputs().",
@ -87,15 +88,15 @@ class BaseFeatureStorage(ABC):
)
@abstractmethod
def list_features(self) -> Dict:
def list_features(self) -> Dict[str, Dict[str, Any]]:
"""List the features in the storage.
Returns
-------
dict
List of features in the storage. The keys are the feature names to
be used in read_features() and the values are the metadata of each
feature.
List of features in the storage. The keys are the feature MD5 to
be used in :meth:`junifer.storage.BaseFeatureStorage.read_df`
and the values are the metadata of each feature.
"""
raise_error(
@ -107,7 +108,7 @@ class BaseFeatureStorage(ABC):
def read_df(
self,
feature_name: Optional[str] = None,
feature_md5: Optional[bool] = None,
feature_md5: Optional[str] = None,
) -> pd.DataFrame:
"""Read feature into a pandas DataFrame.
@ -152,7 +153,7 @@ class BaseFeatureStorage(ABC):
Parameters
----------
kind : {"matrix", "timeseries", "table"}
kind : {"matrix", "timeseries", "vector"}
The storage kind.
**kwargs
The keyword arguments.
@ -179,8 +180,8 @@ class BaseFeatureStorage(ABC):
self.store_timeseries(
meta_md5=meta_md5, element=t_element, **kwargs
)
elif kind == "table":
self.store_table(meta_md5=meta_md5, element=t_element, **kwargs)
elif kind == "vector":
self.store_vector(meta_md5=meta_md5, element=t_element, **kwargs)
def store_matrix(
self,
@ -203,11 +204,9 @@ class BaseFeatureStorage(ABC):
data : numpy.ndarray
The matrix data to store.
col_names : list or tuple of str, optional
The column names (default None).
The column labels (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).
The row labels (default None).
matrix_kind : str, optional
The kind of matrix:
@ -217,23 +216,23 @@ class BaseFeatureStorage(ABC):
(default "full").
diagonal : bool, optional
Whether to store the diagonal. If `matrix_kind` is "full", setting
Whether to store the diagonal. If ``matrix_kind = full``, setting
this to False will raise an error (default True).
"""
raise_error(
msg="Concrete classes need to implement store_matrix2d().",
msg="Concrete classes need to implement store_matrix().",
klass=NotImplementedError,
)
def store_table(
def store_vector(
self,
meta_md5: str,
element: Dict,
data: Union[np.ndarray, List],
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
col_names: Optional[Iterable[str]] = None,
) -> None:
"""Store table.
"""Store vector.
Parameters
----------
@ -242,16 +241,13 @@ class BaseFeatureStorage(ABC):
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).
The vector data to store.
col_names : list or tuple of str, optional
The column labels (default None).
"""
raise_error(
msg="Concrete classes need to implement store_table().",
msg="Concrete classes need to implement store_vector().",
klass=NotImplementedError,
)
@ -260,9 +256,9 @@ class BaseFeatureStorage(ABC):
meta_md5: str,
element: Dict,
data: np.ndarray,
columns: Optional[Iterable[str]] = None,
col_names: Optional[Iterable[str]] = None,
) -> None:
"""Implement timeseries storing.
"""Store timeseries.
Parameters
----------
@ -272,8 +268,9 @@ class BaseFeatureStorage(ABC):
The element as a dictionary.
data : numpy.ndarray
The timeseries data to store.
columns : list or tuple of str, optional
col_names : list or tuple of str, optional
The column labels (default None).
"""
raise_error(
msg="Concrete classes need to implement store_timeseries().",

View file

@ -6,11 +6,12 @@
import json
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Union
from typing import Dict, Iterable, List, Optional, Union
import numpy as np
import pandas as pd
from ..utils import raise_error
from .base import BaseFeatureStorage
@ -46,10 +47,11 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
Returns
-------
list of str
The list of storage types that can be used as input for this "
"storage.
The list of storage types that can be used as input for this
storage interface.
"""
return ["matrix", "table", "timeseries"]
return ["matrix", "vector", "timeseries"]
def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame:
"""Convert the metadata to a pandas DataFrame.
@ -76,7 +78,7 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
@staticmethod
def element_to_index(
element: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None
) -> pd.MultiIndex:
) -> Union[pd.Index, pd.MultiIndex]:
"""Convert the element metadata to index.
Parameters
@ -86,31 +88,40 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
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).
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
pandas.Index or 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] = {
# Make mapping between element access keys and values
elem_idx: Dict[str, Iterable[str]] = {
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))
)
# Set rows_col_name if n_rows > 1 (timeseries)
if n_rows > 1:
# Set rows_col_name if None
if rows_col_name is None:
rows_col_name = "idx"
# Set extra column for variable number of rows per element
elem_idx[rows_col_name] = np.arange(n_rows)
# Create correct index for elements with single access variable
if len(elem_idx) == 1:
# Create normal index for vector
index = pd.Index(
data=list(elem_idx.values())[0], name=list(elem_idx.keys())[0]
)
else:
# Create multiindex for timeseries
index = pd.MultiIndex.from_frame(
pd.DataFrame(elem_idx, index=range(n_rows))
)
return index
def store_df(
@ -132,17 +143,20 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
generated from the metadata.
"""
raise NotImplementedError("Implement in subclass.")
raise_error(
msg="Concrete classes need to implement store_df().",
klass=NotImplementedError,
)
def _store_2d(
self,
meta_md5: str,
element: Dict,
data: Union[np.ndarray, List],
columns: Optional[Iterable[str]] = None,
col_names: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Store 2D dataframe.
"""Store 2D data.
Parameters
----------
@ -150,37 +164,35 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
The metadata MD5 hash.
element : dict
The element as a dictionary.
data : numpy.ndarray or List
data : numpy.ndarray or list
The data to store.
columns : list or tuple of str, optional
The columns (default None).
col_names : list or tuple of str, optional
The column labels (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).
"idx" (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
element=element, n_rows=len(data), rows_col_name=rows_col_name
)
# Prepare new dataframe
data_df = pd.DataFrame( # type: ignore
data, columns=columns, index=idx # type: ignore
df = pd.DataFrame(
data=data, columns=col_names, index=idx # type: ignore
)
# Store dataframe
self.store_df(meta_md5=meta_md5, element=element, df=data_df)
self.store_df(meta_md5=meta_md5, element=element, df=df)
def store_table(
def store_vector(
self,
meta_md5: str,
element: Dict,
data: Union[np.ndarray, List],
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
col_names: Optional[Iterable[str]] = None,
) -> None:
"""Implement table storing.
"""Store vector.
Parameters
----------
@ -188,21 +200,27 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
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).
data : numpy.ndarray or list
The vector data to store.
col_names : list or tuple of str, optional
The column labels (default None).
"""
if isinstance(data, list):
# Flatten out list and convert to np.ndarray
processed_data = np.array(np.ravel(data))
elif isinstance(data, np.ndarray):
# Flatten out array
processed_data = data.ravel()
# Make it 2D
processed_data = processed_data[np.newaxis, :]
self._store_2d(
meta_md5=meta_md5,
element=element,
data=data,
columns=columns,
rows_col_name=rows_col_name,
col_names=col_names,
)
def store_timeseries(
@ -210,9 +228,9 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
meta_md5: str,
element: Dict,
data: np.ndarray,
columns: Optional[Iterable[str]] = None,
col_names: Optional[Iterable[str]] = None,
) -> None:
"""Implement timeseries storing.
"""Store timeseries.
Parameters
----------
@ -222,13 +240,14 @@ class PandasBaseFeatureStorage(BaseFeatureStorage):
The element as a dictionary.
data : numpy.ndarray
The timeseries data to store.
columns : list or tuple of str, optional
col_names : list or tuple of str, optional
The column labels (default None).
"""
self._store_2d(
meta_md5=meta_md5,
element=element,
data=data,
columns=columns,
col_names=col_names,
rows_col_name="timepoint",
)

View file

@ -6,7 +6,7 @@
import json
from pathlib import Path
from typing import TYPE_CHECKING, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import numpy as np
import pandas as pd
@ -34,12 +34,12 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
uri : str or pathlib.Path
The path to the file to be used.
single_output : bool, optional
If False, will create one file per element. The name
If False, will create one SQLite file per element. The name
of the file will be prefixed with the respective element.
If True, will create only one file as specified in the `uri` and
store all the elements in the same file. This behaviour is only
suitable for non-parallel executions. SQLite does not support
concurrency (default True).
If True, will create only one SQLite file as specified in the
``uri`` and store all the elements in the same file. This behaviour
is only suitable for non-parallel executions. SQLite does not
support concurrency (default True).
upsert : {"ignore", "update"}, optional
Upsert mode. If "ignore" is used, the existing elements are ignored.
If "update", the existing elements are updated (default "update").
@ -78,7 +78,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
)
uri.parent.mkdir(parents=True, exist_ok=True)
# Available storage kinds
storage_types = ["table", "timeseries", "matrix"]
storage_types = ["vector", "timeseries", "matrix"]
super().__init__(
uri=uri,
storage_types=storage_types,
@ -93,8 +93,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
Parameters
----------
meta : dict, optional
The metadata as dictionary (default None).
element : dict, optional
The element as dictionary (default None).
Returns
-------
@ -205,24 +205,30 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
msg=f"Invalid option {if_exists} for if_exists."
)
def list_features(self) -> Dict:
def list_features(self) -> Dict[str, Dict[str, Any]]:
"""List the features in the storage.
Returns
-------
dict
List of features in the storage. The keys are the feature names to
be used in read_features() and the values are the metadata of each
feature.
List of features in the storage. The keys are the feature MD5 to
be used in :meth:`junifer.storage.SQLiteFeatureStorage.read_df`
and the values are the metadata of each feature.
"""
# Retrieve meta table from storage
meta_df = pd.read_sql(
sql="meta",
con=self.get_engine(),
index_col="meta_md5",
)
# Format index names for retrieved data
meta_df.index = meta_df.index.str.replace(r"meta_", "")
out = meta_df.to_dict(orient="index") # type: ignore
# Convert dataframe to dictionary
out: Dict[str, Dict[str, str]] = meta_df.to_dict(
orient="index"
) # type: ignore
# Format output
for md5, t_meta in out.items():
for k, v in t_meta.items():
out[md5][k] = json.loads(v)
@ -396,7 +402,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
matrix_kind: Optional[str] = "full",
diagonal: bool = True,
) -> None:
"""Implement matrix storing.
"""Store matrix.
Parameters
----------
@ -409,11 +415,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
meta : dict
The metadata as a dictionary.
col_names : list or tuple of str, optional
The column names (default None).
The column labels (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).
The row labels (optional None).
matrix_kind : str, optional
The kind of matrix:
@ -423,7 +427,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
(default "full").
diagonal : bool, optional
Whether to store the diagonal. If `matrix_kind` is "full", setting
Whether to store the diagonal. If ``matrix_kind = full``, setting
this to False will raise an error (default True).
"""
@ -470,7 +474,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
klass=ValueError,
)
# Subset data
flat_data = data[data_idx]
# Generate flat 1D row X column names
columns = [
f"{row_names[i]}~{col_names[j]}"
for i, j in zip(data_idx[0], data_idx[1])
@ -484,6 +490,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
# Prepare new dataframe
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
# SQLite's SQLITE_MAX_COLUMN is 2000, so if more than that,
# convert it to long format
if len(columns) > 2000:
warn_with_log(
msg="The number of columns is greater than 2000. "
@ -510,7 +518,10 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
"""
if self.single_output is True:
raise_error(msg="collect() is not implemented for single output.")
raise_error(
msg="collect() is not implemented for single output.",
klass=IOError,
)
logger.info(
"Collecting data from "
f"{self.uri.parent}/*{self.uri.name}" # type: ignore

View file

@ -9,62 +9,95 @@ from junifer.storage.pandas_base import PandasBaseFeatureStorage
def test_element_to_index() -> None:
"""Test element to index."""
# First element
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,)
# First test; no extra column
index = PandasBaseFeatureStorage.element_to_index(element=element)
# Check index name
assert index.name == "foo"
# # Check index values
assert all(x == "bar" for x in index.values) # type: ignore
# Check index values shape
assert index.shape == (1,) # type: ignore
# Second test; add extra column
index = PandasBaseFeatureStorage.element_to_index(
element, n_rows=1, rows_col_name="scan"
element=element, n_rows=10
)
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,)
# Check index names
assert index.names == ["foo", "idx"]
# Check first index level name
assert index.levels[0].name == "foo" # type: ignore
# Check first index level values
assert all(x == "bar" for x in index.levels[0].values) # type: ignore
# Check first index level values shape
assert index.levels[0].values.shape == (1,) # type: ignore
# Check second index level name
assert index.levels[1].name == "idx" # type: ignore
# Check second index level values
assert all(
x == i for i, x in enumerate(index.levels[1].values) # type: ignore
)
# Check second index level values shape
assert index.levels[1].values.shape == (10,) # type: ignore
# Third test; custom extra column name has no effect
index = PandasBaseFeatureStorage.element_to_index(
element, n_rows=7, rows_col_name="scan"
element=element, n_rows=1, rows_col_name="scan"
)
# Check index name
assert index.name == "foo"
# Check index values
assert all(x == "bar" for x in index.values) # type: ignore
# Check index values shape
assert index.shape == (1,) # type: ignore
# Fourth test; custom extra column name has effect
index = PandasBaseFeatureStorage.element_to_index(
element=element, n_rows=7, rows_col_name="scan"
)
# Check index names
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,)
# Check first index level name
assert index.levels[0].name == "foo" # type: ignore
# Check first index level values
assert all(x == "bar" for x in index.levels[0].values) # type: ignore
# Check first index level values shape
assert index.levels[0].values.shape == (1,) # type: ignore
# Check second index level name
assert index.levels[1].name == "scan" # type: ignore
# Check second index level values
assert all(
x == i for i, x in enumerate(index.levels[1].values) # type: ignore
)
# Check second index level values shape
assert index.levels[1].values.shape == (7,) # type: ignore
# Second element
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,)
# Fifth test; default name for extra column and multi-level element access
index = PandasBaseFeatureStorage.element_to_index(
element=element, n_rows=10
)
# Check first index level name
assert index.levels[0].name == "subject" # type: ignore
# Check first index level values
assert all(x == "sub-01" for x in index.levels[0].values) # type: ignore
# Check first index level values shape
assert index.levels[0].values.shape == (1,) # type: ignore
# Check second index level name
assert index.levels[1].name == "session" # type: ignore
# Check second index level values
assert all(x == "ses-01" for x in index.levels[1].values) # type: ignore
# Check second index level values shape
assert index.levels[1].values.shape == (1,) # type: ignore
# Check third index level name
assert index.levels[2].name == "idx" # type: ignore
# Check third index level values
assert all(
x == i for i, x in enumerate(index.levels[2].values) # type: ignore
)
# Check third index level values shape
assert index.levels[2].values.shape == (10,) # type: ignore

View file

@ -1,4 +1,4 @@
"""Provide tests for sqlite."""
"""Provide tests for SQLite storage interface."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
@ -361,8 +361,8 @@ def test_store_metadata(tmp_path: Path) -> None:
assert meta_md5 == feature_md5
def test_store_table(tmp_path: Path) -> None:
"""Test table store.
def test_store_vector(tmp_path: Path) -> None:
"""Test vector store.
Parameters
----------
@ -370,7 +370,7 @@ def test_store_table(tmp_path: Path) -> None:
The path to the test directory.
"""
uri = tmp_path / "test_store_table.sqlite"
uri = tmp_path / "test_store_vector.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Metadata to store
element = {"subject": "test"}
@ -389,60 +389,29 @@ def test_store_table(tmp_path: Path) -> None:
)
# Data to store
data = [
[1, 10],
[2, 20],
[3, 30],
[4, 40],
[5, 50],
]
data = [[10, 20, 30, 40, 50]]
col_names = ["f1", "f2", "f3", "f4", "f5"]
# Convert element to index
idx = storage.element_to_index(element, n_rows=5, rows_col_name="scan")
idx = storage.element_to_index(element=element)
# Create dataframe
df = pd.DataFrame(data, columns=["f1", "f2"], index=idx)
df = pd.DataFrame(data=data, columns=col_names, index=idx)
# Store table
storage.store_table(
storage.store_vector(
meta_md5=meta_md5,
element=element_to_store,
data=data,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=col_names,
)
# Read stored table
c_df = _read_sql(
table_name=f"meta_{meta_md5}",
uri=uri.as_posix(),
index_col=["subject", "scan"],
index_col=["subject"],
)
# Check if dataframes are equal
assert_frame_equal(df, c_df)
# New data to store
data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]]
# Convert element to index
idx_new = storage.element_to_index(element, n_rows=6, rows_col_name="scan")
# Create dataframe
df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new)
# Check warning
with pytest.warns(RuntimeWarning, match=r"Some rows"):
# Store table
storage.store_table(
meta_md5=meta_md5,
element=element_to_store,
data=data_new,
columns=["f1", "f2"],
rows_col_name="scan",
)
# Read stored table
c_df_new = _read_sql(
table_name=f"meta_{meta_md5}",
uri=uri.as_posix(),
index_col=["subject", "scan"],
)
# Check if dataframes are equal
assert_frame_equal(df_new, c_df_new)
def test_store_matrix(tmp_path: Path) -> None:
"""Test matrix store.
@ -453,7 +422,7 @@ def test_store_matrix(tmp_path: Path) -> None:
The path to the test directory.
"""
uri = tmp_path / "test_store_table.sqlite"
uri = tmp_path / "test_store_matrix.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Metadata to store
element = {"subject": "test"}
@ -478,7 +447,7 @@ def test_store_matrix(tmp_path: Path) -> None:
row_names = ["row1", "row2", "row3", "row4"]
col_names = ["col1", "col2", "col3"]
# Store table
# Store matrix
storage.store_matrix(
meta_md5=meta_md5,
element=element_to_store,
@ -497,13 +466,13 @@ def test_store_matrix(tmp_path: Path) -> None:
assert_array_equal(read_df.values[0], data.flatten())
assert list(read_df.columns) == stored_names
# Store without row and column names
uri = tmp_path / "test_store_table_nonames.sqlite"
uri = tmp_path / "test_store_matrix_nonames.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
meta_md5=meta_md5, element=element_to_store, meta=meta_to_store
)
# Store table
# Store matrix
storage.store_matrix(
meta_md5=meta_md5, element=element_to_store, data=data
)
@ -553,7 +522,7 @@ def test_store_matrix(tmp_path: Path) -> None:
data = np.array([[1, 2, 3], [11, 22, 33], [111, 222, 333]])
row_names = ["row1", "row2", "row3"]
col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_triu.sqlite"
uri = tmp_path / "test_store_matrix_triu.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
@ -587,7 +556,7 @@ def test_store_matrix(tmp_path: Path) -> None:
)
# Store upper triangular matrix without diagonal
uri = tmp_path / "test_store_table_triu_nodiagonal.sqlite"
uri = tmp_path / "test_store_matrix_triu_nodiagonal.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
@ -622,7 +591,7 @@ def test_store_matrix(tmp_path: Path) -> None:
data = np.array([[1, 2, 3], [11, 22, 33], [111, 222, 333]])
row_names = ["row1", "row2", "row3"]
col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_tril.sqlite"
uri = tmp_path / "test_store_matrix_tril.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
@ -656,7 +625,7 @@ def test_store_matrix(tmp_path: Path) -> None:
)
# Store lower triangular matrix without diagonal
uri = tmp_path / "test_store_table_tril_nodiagonal.sqlite"
uri = tmp_path / "test_store_matrix_tril_nodiagonal.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Store metadata
storage.store_metadata(
@ -687,6 +656,60 @@ def test_store_matrix(tmp_path: Path) -> None:
)
def test_store_timeseries(tmp_path: Path) -> None:
"""Test timeseries store.
Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
uri = tmp_path / "test_store_timeseries.sqlite"
storage = SQLiteFeatureStorage(uri=uri)
# Metadata to store
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 = np.array([[10], [20], [30], [40], [50]])
col_names = ["signal"]
# Convert element to index
idx = storage.element_to_index(
element=element, n_rows=5, rows_col_name="timepoint"
)
# Create dataframe
df = pd.DataFrame(data=data, columns=col_names, index=idx)
# Store table
storage.store_timeseries(
meta_md5=meta_md5,
element=element_to_store,
data=data,
col_names=col_names,
)
# Read stored table
c_df = _read_sql(
table_name=f"meta_{meta_md5}",
uri=uri.as_posix(),
index_col=["subject", "timepoint"],
)
# Check if dataframes are equal
assert_frame_equal(df, c_df)
# TODO: can the test be parametrized?
def test_store_multiple_output(tmp_path: Path):
"""Test storing using single_output=False.
@ -719,41 +742,28 @@ def test_store_multiple_output(tmp_path: Path):
"type": "BOLD",
}
# Data to store
data1 = np.array(
[
[1, 10],
[2, 20],
[3, 30],
[4, 40],
[5, 50],
]
)
data1 = np.array([[10, 20, 30, 40, 50]])
data2 = data1 * 10
data3 = data1 * 20
col_names = ["f1", "f2", "f3", "f4", "f5"]
# Process metadata for storage
hash1, meta_to_store1, element_to_store1 = process_meta(meta1)
# Convert element to index
idx1 = storage.element_to_index(
element_to_store1, n_rows=5, rows_col_name="scan"
)
idx1 = storage.element_to_index(element=element_to_store1)
# Create dataframe
df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1)
df1 = pd.DataFrame(data1, columns=col_names, index=idx1)
# Process metadata for storage
hash2, meta_to_store2, element_to_store2 = process_meta(meta2)
# Convert element to index
idx2 = storage.element_to_index(
element_to_store2, n_rows=5, rows_col_name="scan"
)
idx2 = storage.element_to_index(element=element_to_store2)
# Create dataframe
df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2)
df2 = pd.DataFrame(data2, columns=col_names, index=idx2)
# Process metadata for storage
hash3, meta_to_store3, element_to_store3 = process_meta(meta3)
# Convert element to index
idx3 = storage.element_to_index(
element_to_store3, n_rows=5, rows_col_name="scan"
)
idx3 = storage.element_to_index(element=element_to_store3)
# Create dataframe
df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3)
df3 = pd.DataFrame(data3, columns=col_names, index=idx3)
# Check hash equality
assert hash1 == hash2
assert hash2 == hash3
@ -767,26 +777,23 @@ def test_store_multiple_output(tmp_path: Path):
storage.store_metadata(
meta_md5=hash3, element=element_to_store3, meta=meta_to_store3
)
storage.store_table(
storage.store_vector(
meta_md5=hash1,
element=element_to_store1,
data=data1,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=col_names,
)
storage.store_table(
storage.store_vector(
meta_md5=hash2,
element=element_to_store2,
data=data2,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=col_names,
)
storage.store_table(
storage.store_vector(
meta_md5=hash3,
element=element_to_store3,
data=data3,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=col_names,
)
# Check that URI does not exist yet
assert not uri.exists()
@ -803,12 +810,12 @@ def test_store_multiple_output(tmp_path: Path):
assert uri2.exists()
assert uri3.exists()
# Set index columns
cols = ["subject", "session", "scan"]
idx_cols = ["subject", "session"]
table_name = f"meta_{hash1}"
# Read stored tables
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols)
cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols)
cdf3 = _read_sql(table_name, uri3.as_posix(), index_col=cols)
cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=idx_cols)
cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=idx_cols)
cdf3 = _read_sql(table_name, uri3.as_posix(), index_col=idx_cols)
# Check if dataframes are equal
assert_frame_equal(df1, cdf1)
assert_frame_equal(df2, cdf2)
@ -871,26 +878,23 @@ def test_collect(tmp_path: Path) -> None:
storage.store_metadata(
meta_md5=hash3, element=element_to_store3, meta=meta_to_store3
)
storage.store_table(
storage.store_vector(
meta_md5=hash1,
element=element_to_store1,
data=data1,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=["f1", "f2"],
)
storage.store_table(
storage.store_vector(
meta_md5=hash2,
element=element_to_store2,
data=data2,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=["f1", "f2"],
)
storage.store_table(
storage.store_vector(
meta_md5=hash3,
element=element_to_store3,
data=data3,
columns=["f1", "f2"],
rows_col_name="scan",
col_names=["f1", "f2"],
)
# Convert element to prefix
prefix1 = element_to_prefix(meta1["element"])
@ -911,7 +915,7 @@ def test_collect(tmp_path: Path) -> None:
# Check that URI exists now
assert uri.exists()
# Set index columns
cols = ["subject", "session", "scan"]
cols = ["subject", "session"]
# Store metadata
table_name = f"meta_{hash1}"
# Read stored tables

View file

@ -12,7 +12,9 @@ from junifer.storage.base import BaseFeatureStorage
def test_BaseFeatureStorage_abstractness() -> None:
"""Test BaseFeatureStorage is abstract base class."""
with pytest.raises(TypeError, match=r"abstract"):
BaseFeatureStorage(uri="/tm", storage_types=["matrix"]) # type: ignore
BaseFeatureStorage(
uri="/tmp", storage_types=["matrix"] # type: ignore
)
def test_BaseFeatureStorage() -> None:
@ -21,8 +23,8 @@ def test_BaseFeatureStorage() -> None:
class MyFeatureStorage(BaseFeatureStorage):
"""Implement concrete class."""
def __init__(self, uri, single_output=False):
storage_types = ["matrix", "table", "timeseries"]
def __init__(self, uri, single_output=True):
storage_types = ["matrix", "vector", "timeseries"]
super().__init__(
uri=uri,
storage_types=storage_types,
@ -30,7 +32,7 @@ def test_BaseFeatureStorage() -> None:
)
def get_valid_inputs(self):
return ["matrix", "table", "timeseries"]
return ["matrix", "vector", "timeseries"]
def list_features(self):
super().list_features()
@ -48,7 +50,7 @@ def test_BaseFeatureStorage() -> None:
return super().collect()
# Check single_output is False
st = MyFeatureStorage(uri="/tmp")
st = MyFeatureStorage(uri="/tmp", single_output=False)
assert st.single_output is False
# Check single_output is True
st = MyFeatureStorage(uri="/tmp", single_output=True)
@ -66,11 +68,9 @@ def test_BaseFeatureStorage() -> None:
with pytest.raises(NotImplementedError):
st.read_df(None)
element = {"subject": "test"}
dependencies = ["numpy"]
meta = {
"element": element,
"dependencies": dependencies,
"element": {"subject": "test"},
"dependencies": ["numpy"],
"marker": {"name": "fc"},
"type": "BOLD",
}
@ -88,7 +88,7 @@ def test_BaseFeatureStorage() -> None:
st.store(kind="timeseries", meta=meta)
with pytest.raises(NotImplementedError):
st.store(kind="table", meta=meta)
st.store(kind="vector", meta=meta)
with pytest.raises(ValueError):
st.store(kind="lego", meta=meta)

View file

@ -84,13 +84,13 @@ def process_meta(meta: Dict) -> Tuple[str, Dict, Dict]:
The MD5 hash of the metadata.
dict
The processed metadata for storage.
tuple
dict
The element.
Raises
------
ValueError
If `meta` is None or if it does not contain the key "element".
If ``meta`` is None or if it does not contain the key "element".
"""
if meta is None:
@ -98,7 +98,7 @@ def process_meta(meta: Dict) -> Tuple[str, Dict, Dict]:
# Copy the metadata
t_meta = meta.copy()
# Remove key "element"
element = t_meta.pop("element", None)
element: Dict = t_meta.pop("element", None)
if element is None:
raise_error(msg="`meta` must contain the key 'element'")
if "marker" not in t_meta: