fix: storage #103

Merged
synchon merged 83 commits from fix/storage into main 2022-10-17 16:54:37 +00:00
20 changed files with 717 additions and 529 deletions

View file

@ -51,13 +51,17 @@ Enhancements
- Implement SphereAggregation marker (by `Fede Raimondo`_). - Implement SphereAggregation marker (by `Fede Raimondo`_).
- Implement PIOP1 and PIOP2 AOMIC datasets and refactor AOMICID1000 slightly - Implement PIOP1 and PIOP2 AOMIC datasets and refactor AOMICID1000 slightly
(:gh: `94` by `Leonard Sasse`_) (:gh: `94` by `Leonard Sasse`_)
- Implement a JuselessDataladCamCANVBM datagrabber class (:gh: `99` by `Leonard Sasse`_) - Implement a JuselessDataladCamCANVBM datagrabber class (:gh: `99` by `Leonard Sasse`_)
- Implement IXI CAT output datagrabber for juseless (:gh: `48` by `Leonard Sasse`_). - Implement IXI CAT output datagrabber for juseless (:gh: `48` by `Leonard Sasse`_).
- Upgrade storage interface for storage-like objects (:gh: `84` by `Synchon Mandal`_).
- Add missing type annotations (:gh: `74` by `Synchon Mandal`_).
Bugs Bugs
~~~~ ~~~~

View file

@ -2,7 +2,7 @@
Extracting root sum of squares from edge-wise timeseries. Extracting root sum of squares from edge-wise timeseries.
========================================================= =========================================================
This example uses a RSSETSMarker to compute root sum of squares This example uses a ``RSSETSMarker`` to compute root sum of squares
of the edge-wise timeseries using the Schaefer atlas of the edge-wise timeseries using the Schaefer atlas
(100 rois and 200 rois, 17 Yeo networks) for a 4D nifti BOLD file. (100 rois and 200 rois, 17 Yeo networks) for a 4D nifti BOLD file.
@ -15,6 +15,8 @@ License: BSD 3 clause
import tempfile import tempfile
import junifer.testing.registry # noqa: F401 import junifer.testing.registry # noqa: F401
from junifer.api import collect, run
from junifer.storage import SQLiteFeatureStorage
from junifer.utils import configure_logging from junifer.utils import configure_logging
@ -22,11 +24,15 @@ from junifer.utils import configure_logging
# Set the logging level to info to see extra information: # Set the logging level to info to see extra information:
configure_logging(level="INFO") configure_logging(level="INFO")
##############################################################################
# Define the datagrabber interface
datagrabber = {
"kind": "SPMAuditoryTestingDatagrabber",
}
############################################################################### ###############################################################################
# Define the markers you want: # Define the markers interface
markers = [
marker_dicts = [
{ {
"name": "Schaefer100x17_RSSETS", "name": "Schaefer100x17_RSSETS",
"kind": "RSSETSMarker", "kind": "RSSETSMarker",
@ -39,26 +45,34 @@ marker_dicts = [
}, },
] ]
############################################################################### ###############################################################################
# Create a temporary directory for junifer feature extraction: # Create a temporary directory for junifer feature extraction:
# At the end you can read the extracted data into a ``pandas.DataFrame``. # At the end you can read the extracted data into a ``pandas.DataFrame``.
with tempfile.TemporaryDirectory() as tmpdir: with tempfile.TemporaryDirectory() as tmpdir:
# Define the storage interface
storage = {
"kind": "SQLiteFeatureStorage",
"uri": f"{tmpdir}/test.db",
"single_output": False,
}
# Run the defined junifer feature extraction pipeline
run(
workdir=tmpdir,
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=["sub001"], # we calculate for one subject only
)
# Collect extracted features data
collect(storage=storage)
# Create storage object to read in extracted features
db = SQLiteFeatureStorage(
uri=storage["uri"],
single_output=True, # as we ran collect, we have single output now
)
# Read extracted features
df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS")
storage = {"kind": "SQLiteFeatureStorage", "uri": f"{tmpdir}/test.db"} ###############################################################################
# run the defined junifer feature extraction pipeline # Now we take a look at the dataframe
# TODO: needs SQLiteFeatureStorage.store_timeseries() to be df_vbm.head()
# implemented first
# run(
# workdir="/tmp",
# datagrabber={"kind": "SPMAuditoryTestingDatagrabber"},
# markers=marker_dicts,
# storage=storage,
# )
# read in extracted features and add confounds and targets
# for julearn run cross validation
# This will not run for now as store_timeseries() is not implemented yet
# collect(storage)
# db = SQLiteFeatureStorage(uri=storage["uri"], single_output=True)
# df_vbm = db.read_df(feature_name="Schaefer100x17")

View file

@ -7,5 +7,7 @@
from .base import BaseMarker from .base import BaseMarker
from .collection import MarkerCollection from .collection import MarkerCollection
from .ets_rss import RSSETSMarker from .ets_rss import RSSETSMarker
from .functional_connectivity_atlas import FunctionalConnectivityAtlas
from .functional_connectivity_spheres import FunctionalConnectivitySpheres
from .parcel import ParcelAggregation from .parcel import ParcelAggregation
from .sphere_aggregation import SphereAggregation from .sphere_aggregation import SphereAggregation

View file

@ -1,23 +1,28 @@
"""Provide base class for markers.""" """Provide abstract base class for markers."""
# 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 typing import Dict, List, Optional, Union from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from ..pipeline import PipelineStepMixin from ..pipeline import PipelineStepMixin
from ..utils import logger, raise_error from ..utils import logger, raise_error
class BaseMarker(PipelineStepMixin): if TYPE_CHECKING:
"""Base class for all markers. from junifer.storage import BaseFeatureStorage
class BaseMarker(ABC, PipelineStepMixin):
"""Abstract base class for all markers.
Parameters Parameters
---------- ----------
on : list of str on : str or list of str
The kind of data to apply the marker to. By default, will work on all The kind of data to apply the marker to. By default, will work on all
available data (default None). available data.
name : str, optional name : str, optional
The name of the marker. By default, it will use the class name as the The name of the marker. By default, it will use the class name as the
name of the marker (default None). name of the marker (default None).
@ -44,7 +49,7 @@ class BaseMarker(PipelineStepMixin):
Returns Returns
------- -------
dict dict
The metadata as a dictionary. The metadata as a dictionary with the only key 'marker'.
""" """
s_meta = super().get_meta() s_meta = super().get_meta()
@ -76,6 +81,7 @@ class BaseMarker(PipelineStepMixin):
f"\t Required (any of): {self._valid_inputs}" f"\t Required (any of): {self._valid_inputs}"
) )
@abstractmethod
def get_output_kind(self, input: List[str]) -> List[str]: def get_output_kind(self, input: List[str]) -> List[str]:
"""Get output kind. """Get output kind.
@ -92,9 +98,11 @@ class BaseMarker(PipelineStepMixin):
""" """
raise_error( raise_error(
msg="get_output_kind() not implemented", klass=NotImplementedError msg="Concrete classes need to implement get_output_kind().",
klass=NotImplementedError,
) )
@abstractmethod
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict: def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
"""Compute. """Compute.
@ -117,32 +125,54 @@ class BaseMarker(PipelineStepMixin):
with this as a parameter. with this as a parameter.
""" """
raise_error(msg="compute() not implemented", klass=NotImplementedError) raise_error(
msg="Concrete classes need to implement compute().",
klass=NotImplementedError,
)
# TODO: complete type annotations @abstractmethod
def store(self, kind: str, out: Dict, storage) -> None: def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
input kind : str
out The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class.
""" """
raise_error(msg="store() not implemented", klass=NotImplementedError) raise_error(
msg="Concrete classes need to implement store().",
klass=NotImplementedError,
)
# TODO: complete type annotations def fit_transform(
def fit_transform(self, input: Dict[str, Dict], storage=None) -> Dict: self,
input: Dict[str, Dict],
storage: "BaseFeatureStorage" = None,
) -> Dict:
"""Fit and transform. """Fit and transform.
Parameters Parameters
---------- ----------
input input : dict
storage The Junifer Data object.
storage : storage-like, optional
The storage class, for example, SQLiteFeatureStorage.
Returns Returns
------- -------
dict dict
The processed output as a dictionary. If `storage` is provided,
empty dictionary is returned.
""" """
out = {} out = {}
@ -156,11 +186,11 @@ class BaseMarker(PipelineStepMixin):
t_meta = meta.copy() t_meta = meta.copy()
t_meta.update(t_input.get("meta", {})) t_meta.update(t_input.get("meta", {}))
t_meta.update(self.get_meta(kind)) t_meta.update(self.get_meta(kind))
t_out = self.compute(t_input, extra_input) t_out = self.compute(input=t_input, extra_input=extra_input)
t_out.update(meta=t_meta) t_out.update(meta=t_meta)
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, t_out, storage) self.store(kind=kind, 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[kind] = t_out

View file

@ -5,7 +5,7 @@
# License: AGPL # License: AGPL
from collections import Counter from collections import Counter
from typing import Dict, List, Optional from typing import TYPE_CHECKING, Dict, List, Optional
from ..datareader.default import DefaultDataReader from ..datareader.default import DefaultDataReader
from ..markers.base import BaseMarker from ..markers.base import BaseMarker
@ -14,15 +14,23 @@ from ..storage.base import BaseFeatureStorage
from ..utils import logger from ..utils import logger
if TYPE_CHECKING:
from junifer.datagrabber import BaseDataGrabber
class MarkerCollection: class MarkerCollection:
"""Class for marker collection. """Class for marker collection.
Parameters Parameters
---------- ----------
markers markers : list of marker-like
datareader The markers to compute.
preprocessing datareader : datareader-like, optional
storage The datareader to use (default None).
preprocessing : preprocessing-like, optional
The preprocessing steps to apply.
storage : storage-like, optional
The storage to use (default None).
""" """
@ -60,17 +68,23 @@ class MarkerCollection:
Returns Returns
------- -------
output : dict or None dict or None
The output of the pipeline. Each key represents a marker name and The output of the pipeline. Each key represents a marker name and
the values are the computer marker values. If the pipeline has a the values are the computer marker values. If the pipeline has a
storage configured, then the output will be None. storage configured, then the output will be None.
""" """
logger.info("Fitting pipeline") logger.info("Fitting pipeline")
# Fetch actual data using datareader
data = self._datareader.fit_transform(input) data = self._datareader.fit_transform(input)
# Apply preprocessing steps
if self._preprocessing is not None: if self._preprocessing is not None:
logger.info("Preprocessing data") logger.info("Preprocessing data")
data = self._preprocessing.fit_transform(data) data = self._preprocessing.fit_transform(data)
# Compute markers
out = {} out = {}
for marker in self._markers: for marker in self._markers:
logger.info(f"Fitting marker {marker.name}") logger.info(f"Fitting marker {marker.name}")
@ -78,20 +92,21 @@ class MarkerCollection:
if self._storage is None: if self._storage is None:
out[marker.name] = m_value out[marker.name] = m_value
logger.info("Marker collection fitting done") logger.info("Marker collection fitting done")
return None if self._storage else out return None if self._storage else out
# TODO: complete type annotations def validate(self, datagrabber: "BaseDataGrabber") -> None:
def validate(self, datagrabber) -> None:
"""Validate the pipeline. """Validate the pipeline.
Without doing any computation, check if the Marker Collection can Without doing any computation, check if the marker collection can
be fit without problems. That is, the data required for each marker is be fit without problems. That is, the data required for each marker is
present and streamed down the steps. Also, if a storage is configured, present and streamed down the steps. Also, if a storage is configured,
check that the storage can handle the markers output. check that the storage can handle the markers output.
Parameters Parameters
---------- ----------
datagrabber datagrabber : datagrabber-like
The datagrabber to validate.
""" """
logger.info("Validating Marker Collection") logger.info("Validating Marker Collection")
@ -104,8 +119,11 @@ class MarkerCollection:
fraimondo commented 2022-10-14 10:48:29 +00:00 (Migrated from github.com)

This should still be validate.

There are two methods:

validate_input: checks the input
get_output_kind: gives the output kind (given the input)

validate calls both of them

This should still be validate. There are two methods: `validate_input`: checks the input `get_output_kind`: gives the output kind (given the input) `validate` calls both of them
fraimondo commented 2022-10-14 10:49:18 +00:00 (Migrated from github.com)

Same here, the method should be validate

Same here, the method should be `validate`
synchon commented 2022-10-14 12:21:20 +00:00 (Migrated from github.com)

Okay makes sense since it comes from PipelineStepMixin.

Okay makes sense since it comes from `PipelineStepMixin`.
synchon commented 2022-10-17 08:54:43 +00:00 (Migrated from github.com)

Done.

Done.
for marker in self._markers: for marker in self._markers:
logger.info(f"Validating Marker: {marker.name}") logger.info(f"Validating Marker: {marker.name}")
m_data = marker.validate(t_data) # Validate marker
m_data = marker.validate(input=t_data)
logger.info(f"Marker output type: {m_data}") logger.info(f"Marker output type: {m_data}")
# Check storage for the marker
if self._storage is not None: if self._storage is not None:
logger.info(f"Validating storage for {marker.name}") logger.info(f"Validating storage for {marker.name}")
self._storage.validate(m_data) # Validate storage
self._storage.validate(input_=m_data)

View file

@ -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 Dict, List from typing import TYPE_CHECKING, Any, Dict, List, Optional
import numpy as np import numpy as np
@ -17,6 +17,10 @@ from .parcel 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):
"""Class for root sum of squares of edgewise timeseries. """Class for root sum of squares of edgewise timeseries.
@ -24,9 +28,14 @@ class RSSETSMarker(BaseMarker):
Parameters Parameters
---------- ----------
atlas : str atlas : str
The name of the atlas. The name of the atlas. Check valid options by calling
:func:`junifer.data.list_atlases`.
aggregation_method : str, optional aggregation_method : str, optional
The aggregation method (default "mean"). The method to perform aggregation using. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name` (default "mean").
name : str, optional
The name of the marker. If None, will use the class name (default
None).
""" """
@ -58,21 +67,32 @@ class RSSETSMarker(BaseMarker):
""" """
return ["timeseries"] return ["timeseries"]
# TODO: complete type annotations def store(
def store(self, kind: str, out, storage) -> None: self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
kind kind : {"BOLD"}
out The data kind to store.
storage 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}") logger.debug(f"Storing BOLD in {storage}")
storage.store_timeseries(**out) storage.store(kind="timeseries", **out)
def compute(self, input: Dict) -> Dict: def compute(
self,
input: Dict[str, Any],
extra_input: Optional[Dict] = None,
) -> Dict:
"""Compute. """Compute.
Take a timeseries of brain areas, and calculate timeseries for each Take a timeseries of brain areas, and calculate timeseries for each
@ -81,12 +101,19 @@ class RSSETSMarker(BaseMarker):
Parameters Parameters
---------- ----------
input : dict of the BOLD data input : dict
The BOLD data as dictionary.
extra_input : dict, optional
The other fields in the pipeline data object (default None).
Returns Returns
------- -------
dict dict
The computed result as dictionary. The computed result as dictionary. 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"
References References
---------- ----------
@ -103,9 +130,10 @@ class RSSETSMarker(BaseMarker):
method=self.aggregation_method, method=self.aggregation_method,
) )
# Compute the parcel aggregation # Compute the parcel aggregation
out = parcel_aggregation.compute(input) out = parcel_aggregation.compute(input=input, extra_input=extra_input)
edge_ts = _ets(out["data"]) edge_ts = _ets(out["data"])
# Compute the RSS # Compute the RSS
out["data"] = np.sum(edge_ts**2, 1) ** 0.5 out["data"] = np.sum(edge_ts**2, 1) ** 0.5
# Set correct column label
out["columns"] = ["root_sum_of_squares_ets"]
return out return out

View file

@ -4,7 +4,7 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Dict, List, Optional from typing import TYPE_CHECKING, Any, Dict, List, Optional
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
@ -15,29 +15,45 @@ from .base import BaseMarker
from .parcel import ParcelAggregation from .parcel import ParcelAggregation
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker @register_marker
class FunctionalConnectivityAtlas(BaseMarker): class FunctionalConnectivityAtlas(BaseMarker):
"""Class for functional connectivity. """Class for functional connectivity.
Parameters Parameters
---------- ----------
atlas atlas : str
agg_method The name of the atlas. Check valid options by calling
agg_method_params :func:`junifer.data.list_atlases`.
cor_method agg_method : str, optional
cor_method_params The method to perform aggregation using. Check valid options in
name :func:`junifer.stats.get_aggfunc_by_name` (default "mean").
agg_method_params : dict, optional
Parameters to pass to the aggregation function. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name` (default None).
cor_method : str, optional
The method to perform correlation using. Check valid options in
:func:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
cor_method_params : dict, optional
Parameters to pass to the correlation function. Check valid options in
:func:`nilearn.connectome.ConnectivityMeasure` (default None).
name : str, optional
The name of the marker. If None, will use the class name (default
None).
""" """
def __init__( def __init__(
self, self,
atlas, atlas: str,
agg_method="mean", agg_method: str = "mean",
agg_method_params=None, agg_method_params: Optional[Dict] = None,
cor_method="covariance", cor_method: str = "covariance",
cor_method_params=None, cor_method_params: Optional[Dict] = None,
name=None, name: Optional[str] = None,
) -> None: ) -> None:
"""Initialize the class.""" """Initialize the class."""
self.atlas = atlas self.atlas = atlas
@ -75,15 +91,19 @@ class FunctionalConnectivityAtlas(BaseMarker):
outputs = ["matrix"] outputs = ["matrix"]
return outputs return outputs
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict: def compute(
self,
input: Dict[str, Any],
extra_input: Optional[Dict] = None,
) -> Dict:
"""Compute. """Compute.
Parameters Parameters
---------- ----------
input : Dict[str, Dict] input : dict
A single input from the pipeline data object in which to compute A single input from the pipeline data object in which to compute
the marker. the marker.
extra_input : Dict, optional extra_input : dict, optional
The other fields in the pipeline data object. Useful for accessing The other fields in the pipeline data object. Useful for accessing
other data kind that needs to be used in the computation. For other data kind that needs to be used in the computation. For
example, the functional connectivity markers can make use of the example, the functional connectivity markers can make use of the
@ -94,10 +114,11 @@ class FunctionalConnectivityAtlas(BaseMarker):
dict dict
The computed result as dictionary. The following data will be The computed result as dictionary. The following data will be
included in the dictionary: included in the dictionary:
- 'data': FC matrix as a 2D numpy array. - data: functional connectivity matrix as a numpy.ndarray.
- 'row_names': Row names as a list. - row_names: row names as a list
- 'col_names': Col names as a list. - col_names: column names as a list
- 'kind': The kind of matrix (tril, triu or full) - matrix_kind: the kind of matrix (tril, triu or full)
""" """
pa = ParcelAggregation( pa = ParcelAggregation(
atlas=self.atlas, atlas=self.atlas,
@ -120,18 +141,26 @@ class FunctionalConnectivityAtlas(BaseMarker):
# create column names # create column names
out["row_names"] = ts["columns"] out["row_names"] = ts["columns"]
out["col_names"] = ts["columns"] out["col_names"] = ts["columns"]
out["kind"] = "tril" out["matrix_kind"] = "tril"
return out return out
# TODO: complete type annotations def store(
def store(self, kind: str, out: Dict, storage) -> None: self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
input kind : {"BOLD"}
out 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}") logger.debug(f"Storing {kind} in {storage}")
storage.store_matrix2d(**out) storage.store(kind="matrix", **out)

View file

@ -4,7 +4,7 @@
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Dict, List, Optional from typing import TYPE_CHECKING, Any, Dict, List, Optional
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
@ -15,6 +15,10 @@ 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):
"""Class for functional connectivity using coordinates (spheres). """Class for functional connectivity using coordinates (spheres).
@ -23,16 +27,23 @@ class FunctionalConnectivitySpheres(BaseMarker):
---------- ----------
coords : str coords : str
The name of the coordinates list to use. See The name of the coordinates list to use. See
:mod:`junifer.data.coordinates` :mod:`junifer.data.coordinates` for options.
radius : float radius : float, optional
The radius of the sphere in mm. If None, the signal will be extracted The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information. for more information (default None).
agg_method : str agg_method : str, optional
The aggregation method to use. The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information. See :func:`junifer.stats.get_aggfunc_by_name` for more information
(default None).
agg_method_params : dict, optional agg_method_params : dict, optional
The parameters to pass to the aggregation method. The parameters to pass to the aggregation method (default None).
cor_method : str, optional
The method to perform correlation using. Check valid options in
:func:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
cor_method_params : dict, optional
Parameters to pass to the correlation function. Check valid options in
:func:`nilearn.connectome.ConnectivityMeasure` (default None).
name : str, optional name : str, optional
The name of the marker. By default, it will use The name of the marker. By default, it will use
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
@ -43,7 +54,7 @@ class FunctionalConnectivitySpheres(BaseMarker):
def __init__( def __init__(
self, self,
coords: str, coords: str,
radius: float, radius: Optional[float] = None,
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance", cor_method: str = "covariance",
@ -88,12 +99,16 @@ class FunctionalConnectivitySpheres(BaseMarker):
outputs = ["matrix"] outputs = ["matrix"]
return outputs return outputs
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict: def compute(
self,
input: Dict[str, Any],
extra_input: Optional[Dict] = None,
) -> Dict:
"""Compute. """Compute.
Parameters Parameters
---------- ----------
input : dict[str, dict] input : dict
A single input from the pipeline data object in which to compute A single input from the pipeline data object in which to compute
the marker. the marker.
extra_input : dict, optional extra_input : dict, optional
@ -107,10 +122,10 @@ class FunctionalConnectivitySpheres(BaseMarker):
dict dict
The computed result as dictionary. The following keys will be The computed result as dictionary. The following keys will be
included in the dictionary: included in the dictionary:
- 'data': FC matrix as a 2D numpy array. - data: functional connectivity matrix as a numpy.ndarray.
- 'row_names': Row names as a list. - row_names: row names as a list
- 'col_names': Col names as a list. - col_names: column names as a list
- 'kind': The kind of matrix (tril, triu or full) - matrix_kind: the kind of matrix (tril, triu or full)
""" """
sa = SphereAggregation( sa = SphereAggregation(
@ -135,18 +150,26 @@ class FunctionalConnectivitySpheres(BaseMarker):
# create column names # create column names
out["row_names"] = ts["columns"] out["row_names"] = ts["columns"]
out["col_names"] = ts["columns"] out["col_names"] = ts["columns"]
out["kind"] = "tril" out["matrix_kind"] = "tril"
return out return out
# TODO: complete type annotations def store(
def store(self, kind: str, out: Dict, storage) -> None: self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
input kind : {"BOLD"}
out 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}") logger.debug(f"Storing {kind} in {storage}")
storage.store_matrix2d(**out) storage.store(kind="matrix", **out)

View file

@ -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 Dict, List, Optional from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import numpy as np import numpy as np
from nilearn.image import math_img, resample_to_img from nilearn.image import math_img, resample_to_img
@ -17,22 +17,42 @@ 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):
"""Class for parcel aggregation. """Class for parcel aggregation.
Parameters Parameters
---------- ----------
atlas atlas : str
method The name of the atlas. Check valid options by calling
method_params :func:`junifer.data.list_atlases`.
on method : str
name The method to perform aggregation using. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name`.
method_params : dict, optional
Parameters to pass to the aggregation function. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name`.
on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} or list
of the options, optional
The kind of data to apply the marker to. If None, will work on all
available data (default None).
name : str, optional
The name of the marker. If None, will use the class name (default
None).
""" """
def __init__( def __init__(
self, atlas, method, method_params=None, on=None, name=None self,
atlas: str,
method: str,
method_params: Optional[Dict[str, Any]] = None,
on: Union[List[str], str, None] = None,
name: Optional[str] = None,
) -> None: ) -> None:
"""Initialize the class.""" """Initialize the class."""
self.atlas = atlas self.atlas = atlas
@ -52,8 +72,8 @@ class ParcelAggregation(BaseMarker):
Returns Returns
------- -------
str list of str
The kind of output. The list of storage kinds.
""" """
outputs = [] outputs = []
@ -66,32 +86,43 @@ class ParcelAggregation(BaseMarker):
raise ValueError(f"Unknown input kind for {t_input}") raise ValueError(f"Unknown input kind for {t_input}")
return outputs return outputs
# TODO: complete type annotations def store(
def store(self, kind: str, out, storage) -> None: self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
kind kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
out The data kind to store.
storage 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}") logger.debug(f"Storing {kind} in {storage}")
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
storage.store_table(**out) storage.store(kind="table", **out)
if kind in ["BOLD"]: elif kind in ["BOLD"]:
storage.store_timeseries(**out) storage.store(kind="timeseries", **out)
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict: def compute(
self,
input: Dict[str, Any],
extra_input: Optional[Dict] = None,
) -> Dict:
"""Compute. """Compute.
Parameters Parameters
---------- ----------
input : Dict[str, Dict] input : dict
A single input from the pipeline data object in which to compute A single input from the pipeline data object in which to compute
the marker. the marker.
extra_input : Dict, optional extra_input : dict, optional
The other fields in the pipeline data object. Useful for accessing The other fields in the pipeline data object. Useful for accessing
other data kind that needs to be used in the computation. For other data kind that needs to be used in the computation. For
example, the functional connectivity markers can make use of the example, the functional connectivity markers can make use of the
@ -102,17 +133,24 @@ class ParcelAggregation(BaseMarker):
dict dict
The computed result as dictionary. This will be either returned The computed result as dictionary. This will be either returned
to the user or stored in the storage by calling the store method to the user or stored in the storage by calling the store method
with this as a parameter. 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"
""" """
t_input = input["data"] t_input = input["data"]
logger.debug(f"Parcel aggregation using {self.method}") logger.debug(f"Parcel aggregation using {self.method}")
agg_func = get_aggfunc_by_name( agg_func = get_aggfunc_by_name(
self.method, func_params=self.method_params name=self.method,
func_params=self.method_params,
) )
# Get the min of the voxels sizes and use it as the resolution # Get the min of the voxels sizes and use it as the resolution
resolution = np.min(t_input.header.get_zooms()[:3]) # type: ignore resolution = np.min(t_input.header.get_zooms()[:3])
t_atlas, t_labels, _ = load_atlas(self.atlas, resolution=resolution) t_atlas, t_labels, _ = load_atlas(
name=self.atlas,
resolution=resolution,
)
atlas_img_res = resample_to_img( atlas_img_res = resample_to_img(
t_atlas, t_atlas,
t_input, t_input,

View file

@ -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 Dict, List, Optional from typing import TYPE_CHECKING, Any, Dict, List, Optional
from nilearn.maskers import NiftiSpheresMasker from nilearn.maskers import NiftiSpheresMasker
@ -14,6 +14,10 @@ from ..utils import logger, raise_error
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):
"""Class for sphere aggregation. """Class for sphere aggregation.
@ -22,16 +26,17 @@ class SphereAggregation(BaseMarker):
---------- ----------
coords: str coords: str
The name of the coordinates list to use. See The name of the coordinates list to use. See
:mod:`junifer.data.coordinates` :mod:`junifer.data.coordinates` for options.
radius: float radius: float, optional
The radius of the sphere in mm. If None, the signal will be extracted The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information. for more information (default None).
method: str method: str, optional
The aggregation method to use. The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information. See :func:`junifer.stats.get_aggfunc_by_name` for more information
method_params: Dict, optional (default "mean").
The parameters to pass to the aggregation method. method_params: dict, optional
The parameters to pass to the aggregation method (default None).
on: list of str, optional on: list of str, optional
The kind of data to apply the marker to. By default, will work on all The kind of data to apply the marker to. By default, will work on all
available data (default None). available data (default None).
@ -44,8 +49,8 @@ class SphereAggregation(BaseMarker):
def __init__( def __init__(
self, self,
coords: str, coords: str,
radius: float, radius: Optional[float] = None,
method: str, method: str = "mean",
method_params: Optional[Dict] = None, method_params: Optional[Dict] = None,
on: Optional[List[str]] = None, on: Optional[List[str]] = None,
name: Optional[str] = None, name: Optional[str] = None,
@ -92,31 +97,43 @@ class SphereAggregation(BaseMarker):
raise ValueError(f"Unknown input kind for {t_input}") raise ValueError(f"Unknown input kind for {t_input}")
return outputs return outputs
def store(self, kind: str, out, storage) -> None: def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store. """Store.
Parameters Parameters
---------- ----------
kind kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
out The data kind to store.
storage 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}") logger.debug(f"Storing {kind} in {storage}")
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
storage.store_table(**out) storage.store(kind="table", **out)
elif kind in ["BOLD"]: elif kind in ["BOLD"]:
storage.store_timeseries(**out) storage.store(kind="timeseries", **out)
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict: def compute(
self,
input: Dict[str, Any],
extra_input: Optional[Dict] = None,
) -> Dict:
"""Compute. """Compute.
Parameters Parameters
---------- ----------
input : Dict[str, Dict] input : dict
A single input from the pipeline data object in which to compute A single input from the pipeline data object in which to compute
the marker. the marker.
extra_input : Dict, optional extra_input : dict, optional
The other fields in the pipeline data object. Useful for accessing The other fields in the pipeline data object. Useful for accessing
other data kind that needs to be used in the computation. For other data kind that needs to be used in the computation. For
example, the functional connectivity markers can make use of the example, the functional connectivity markers can make use of the
@ -125,7 +142,12 @@ class SphereAggregation(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The computed result as dictionary. This will be either returned
to the user or stored in the storage by calling the store method
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"
""" """
t_input = input["data"] t_input = input["data"]
@ -135,8 +157,8 @@ class SphereAggregation(BaseMarker):
# ) # )
coords, out_labels = load_coordinates(self.coords) coords, out_labels = load_coordinates(self.coords)
masker = NiftiSpheresMasker( masker = NiftiSpheresMasker(
coords, seeds=coords,
self.radius, radius=self.radius,
mask_img=None, # TODO: support this (needs #79) mask_img=None, # TODO: support this (needs #79)
) )

View file

@ -4,6 +4,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path
import pytest import pytest
from numpy.testing import assert_array_equal from numpy.testing import assert_array_equal
@ -94,7 +96,7 @@ def test_marker_collection():
) )
def test_MarkerCollection_storage(tmp_path) -> None: def test_marker_collection_storage(tmp_path: Path) -> None:
"""Test marker collection with storage. """Test marker collection with storage.
Parameters Parameters
@ -123,7 +125,9 @@ def test_MarkerCollection_storage(tmp_path) -> None:
uri = tmp_path / "test_marker_collection_storage.db" uri = tmp_path / "test_marker_collection_storage.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True) storage = SQLiteFeatureStorage(uri=uri, single_output=True)
mc = MarkerCollection( mc = MarkerCollection(
markers=markers, storage=storage, datareader=DefaultDataReader() markers=markers,
storage=storage,
datareader=DefaultDataReader(),
) )
mc.validate(dg) mc.validate(dg)
assert mc._storage is not None assert mc._storage is not None

View file

@ -13,27 +13,36 @@ from nilearn.maskers import NiftiLabelsMasker
from junifer.data import load_atlas from junifer.data import load_atlas
from junifer.markers.ets_rss import RSSETSMarker from junifer.markers.ets_rss import RSSETSMarker
from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
def test_compute() -> None: # Set atlas
"""Test RSS ETS.""" ATLAS = "Schaefer100x17"
atlas = "Schaefer100x17"
test_atlas, _, _ = load_atlas(atlas)
def test_compute() -> None:
"""Test RSS ETS compute()."""
with SPMAuditoryTestingDatagrabber() as dg: with SPMAuditoryTestingDatagrabber() as dg:
# Fetch element
out = dg["sub001"] out = dg["sub001"]
# Load BOLD image
niimg = image.load_img(str(out["BOLD"]["path"].absolute())) niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
# Create input data
input_dict = {"data": niimg, "path": out["BOLD"]["path"]} input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
# Compute the RSSETSMarker # Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(atlas=atlas) ets_rss_marker = RSSETSMarker(atlas=ATLAS)
new_out = ets_rss_marker.compute(input_dict) new_out = ets_rss_marker.compute(input_dict)
# Load atlas
test_atlas, _, _ = load_atlas(ATLAS)
# Compute the NiftiLabelsMasker # Compute the NiftiLabelsMasker
test_masker = NiftiLabelsMasker(test_atlas) test_masker = NiftiLabelsMasker(test_atlas)
test_ts = test_masker.fit_transform(niimg) test_ts = test_masker.fit_transform(niimg)
# Assert the dimension of timeseries # Assert the dimension of timeseries
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 # Assert the meta
meta = ets_rss_marker.get_meta("BOLD")["marker"] meta = ets_rss_marker.get_meta("BOLD")["marker"]
assert meta["atlas"] == "Schaefer100x17" assert meta["atlas"] == "Schaefer100x17"
@ -42,10 +51,8 @@ def test_compute() -> None:
def test_get_output_kind() -> None: def test_get_output_kind() -> None:
"""Test get_output_kind.""" """Test RSS ETS get_output_kind()."""
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
atlas = "Schaefer100x17"
ets_rss_marker = RSSETSMarker(atlas=atlas)
input_list = ["BOLD"] input_list = ["BOLD"]
input_list = ets_rss_marker.get_output_kind(input_list) input_list = ets_rss_marker.get_output_kind(input_list)
assert len(input_list) == 1 assert len(input_list) == 1
@ -53,19 +60,26 @@ def test_get_output_kind() -> None:
def test_store(tmp_path: Path) -> None: def test_store(tmp_path: Path) -> None:
"""Test store.""" """Test RSS ETS store().
atlas = "Schaefer100x17" Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
with SPMAuditoryTestingDatagrabber() as dg: with SPMAuditoryTestingDatagrabber() as dg:
# Fetch element
out = dg["sub001"] out = dg["sub001"]
# Load BOLD image
niimg = image.load_img(str(out["BOLD"]["path"].absolute())) niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
input_dict = {"data": niimg, "path": out["BOLD"]["path"]} input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
# Compute the RSSETSMarker # Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(atlas=atlas) ets_rss_marker = RSSETSMarker(atlas=ATLAS)
_ = ets_rss_marker.compute(input_dict) # Create storage
# TODO: Needs store_timeseries implemented for SQLiteFeatureStorage storage = SQLiteFeatureStorage(
# storage = { uri=str((tmp_path / "test.db").absolute()),
# "kind": "SQLiteFeatureStorage", single_output=True,
# "uri": str((tmp_path / "test.db").absolute()), )
# } # Store
# ets_rss_marker.store("SQLiteFeatureStorage", new_out, storage) ets_rss_marker.fit_transform(input=input_dict, storage=storage)

View file

@ -1,4 +1,4 @@
"""Provide test for parcel aggregation.""" """Provide tests for functional connectivity atlas."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
@ -19,8 +19,14 @@ from junifer.storage import SQLiteFeatureStorage
def test_FunctionalConnectivityAtlas(tmp_path: Path) -> None: def test_FunctionalConnectivityAtlas(tmp_path: Path) -> None:
"""Test FunctionalConnectivityAtlas.""" """Test FunctionalConnectivityAtlas.
Parameters
----------
tmp_path : pathlib.Path
The path to the test directory.
"""
# get a dataset # get a dataset
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
fmri_img = image.concat_imgs(ni_data.func) # type: ignore fmri_img = image.concat_imgs(ni_data.func) # type: ignore

View file

@ -1,18 +1,17 @@
"""Provide test for functional connectivity spheres.""" """Provide tests for functional connectivity spheres."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de> # Kaustubh R. Patil <k.patil@fz-juelich.de>
# Federico Raimondo <f.raimondo@fz-juelich.de> # Federico Raimondo <f.raimondo@fz-juelich.de>
# License: AGPL # License: AGPL
import pytest
from pathlib import Path from pathlib import Path
from numpy.testing import assert_array_almost_equal
from sklearn.covariance import EmpiricalCovariance import pytest
from nilearn import datasets, image from nilearn import datasets, image
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from numpy.testing import assert_array_almost_equal
from sklearn.covariance import EmpiricalCovariance
from junifer.markers.functional_connectivity_spheres import ( from junifer.markers.functional_connectivity_spheres import (
FunctionalConnectivitySpheres, FunctionalConnectivitySpheres,
@ -30,7 +29,6 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
The path to the test directory. The path to the test directory.
""" """
# get a dataset # get a dataset
ni_data = datasets.fetch_spm_auditory(subject_id="sub001") ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
fmri_img = image.concat_imgs(ni_data.func) # type: ignore fmri_img = image.concat_imgs(ni_data.func) # type: ignore
@ -118,8 +116,7 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
# Check that FC are almost equal when using nileran # Check that FC are almost equal when using nileran
cm = ConnectivityMeasure( cm = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), # type: ignore cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
kind="correlation"
) )
out_ni = cm.fit_transform([ts["data"]])[0] out_ni = cm.fit_transform([ts["data"]])[0]
assert_array_almost_equal(out_ni, out["data"], decimal=3) assert_array_almost_equal(out_ni, out["data"], decimal=3)

View file

@ -4,94 +4,57 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import List, Optional
import pytest import pytest
from junifer.markers.base import BaseMarker from junifer.markers.base import BaseMarker
@pytest.mark.parametrize( def test_base_marker_abstractness() -> None:
"on, name, kind, expected_class, expected_name", """Test BaseMarker is abstract base class."""
[ with pytest.raises(TypeError, match=r"abstract"):
(["bold", "dwi"], None, "bold", "BaseMarker", "bold_BaseMarker"), BaseMarker(on=["BOLD"])
(["bold", "dwi"], "mymarker", "dwi", "BaseMarker", "dwi_mymarker"),
],
)
def test_base_marker_meta(
on: List[str],
name: Optional[str],
kind: str,
expected_class: str,
expected_name: str,
) -> None:
"""Test metadata for BaseMarker.
Parameters
----------
on : list of str
The parametrized kind of data to work on.
name : str or None
The parametrized name of the marker.
kind : str
The parametrized kind of data to get metadata for.
expected_class : str
The paramtrized expected class of the marker.
expected_name : str
The parametrized expected name of the marker.
"""
base = BaseMarker(on=on, name=name)
t_meta = base.get_meta(kind=kind)
assert t_meta["marker"]["class"] == expected_class
assert t_meta["marker"]["name"] == expected_name
def test_compute_parameters() -> None: def test_base_datagrabber_subclassing() -> None:
"""Test compute parameters.""" """Test proper subclassing of BaseMarker."""
base = BaseMarker(on=["bold", "dwi"], name="mymarker") # Create concrete class
base.compute = lambda x, y: { # type: ignore class MyBaseMarker(BaseMarker):
"data": x.keys(), def get_output_kind(self, input):
"extra": y.keys(), return ["timeseries"]
def compute(self, input, extra_input):
return {
"data": "data",
"columns": "columns",
"row_names": "row_names",
}
def store(self, kind, out, storage):
return super().store(kind=kind, out=out, storage=storage)
# Create input for marker
input_ = {
"meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"},
"BOLD": {
"path": ".",
"data": "data",
},
} }
input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}} marker = MyBaseMarker(on=["BOLD"])
out = base.fit_transform( output = marker.fit_transform(input=input_) # process
input_, # Check output
) assert "BOLD" in output
assert list(out["bold"]["data"]) == ["path"] assert "data" in output["BOLD"]
assert list(out["bold"]["extra"]) == ["t2"] assert "columns" in output["BOLD"]
assert "row_names" in output["BOLD"]
assert "meta" in output["BOLD"]
def test_BaseMarker() -> None: assert "datagrabber" in output["BOLD"]["meta"]
"""Test base class.""" assert "element" in output["BOLD"]["meta"]
base = BaseMarker(on=["bold", "dwi"], name="mymarker") assert "datareader" in output["BOLD"]["meta"]
input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}}
base.validate_input(list(input_.keys()))
wrong_input = {"t2": {"path": "test"}}
with pytest.raises(ValueError):
base.validate_input(list(wrong_input.keys()))
# Check no implementation check
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
base.get_output_kind(list(wrong_input.keys())) marker.store(kind="kind", out="out", storage="storage")
with pytest.raises(NotImplementedError): # Check attributes
base.fit_transform(input_) assert marker.name == "MyBaseMarker"
with pytest.raises(NotImplementedError):
base.store("bold", {}, None)
base.compute = lambda x, y: {"data": 1} # type: ignore
out = base.fit_transform(input_)
assert out["bold"]["data"] == 1
assert out["bold"]["meta"]["marker"]["name"] == "bold_mymarker"
assert out["bold"]["meta"]["marker"]["class"] == "BaseMarker"
assert "dwi" not in out
base2 = BaseMarker(on="bold", name="mymarker")
base2.compute = lambda x, y: {"data": 1} # type: ignore
out2 = base2.fit_transform(input_)
assert out2["bold"]["data"] == 1
assert out2["bold"]["meta"]["marker"]["name"] == "bold_mymarker"
assert out2["bold"]["meta"]["marker"]["class"] == "BaseMarker"

View file

@ -138,19 +138,17 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
marker.fit_transform(input, storage=storage) marker.fit_transform(input, storage=storage)
# TODO: Needs store_timeseries implemented for SQLiteFeatureStorage meta = {
"element": "test",
"version": "0.0.1",
"marker": {"name": "BOLD_fcname"},
}
# 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 = SphereAggregation(
coords="DMNBuckner", method="mean", radius=8, on="BOLD"
)
# meta = { marker.fit_transform(input, storage=storage)
# "element": "test",
# "version": "0.0.1",
# "marker": {"name": "BOLD_fcname"},
# }
# # 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 = SphereAggregation(
# coords="DMNBuckner", method="mean", radius=8, on="BOLD"
# )
# marker.fit_transform(input, storage=storage)

View file

@ -6,7 +6,7 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from pathlib import Path from pathlib import Path
from typing import Dict, Iterable, List, Optional, Union from typing import Dict, List, Optional, Union
import pandas as pd import pandas as pd
@ -24,16 +24,24 @@ class BaseFeatureStorage(ABC):
---------- ----------
uri : str or pathlib.Path uri : str or pathlib.Path
The path to the storage. The path to the storage.
storage_types : str or list of str
The available storage types for the class.
single_output : bool, optional single_output : bool, optional
Whether to have single output (default False). Whether to have single output (default False).
""" """
def __init__( def __init__(
self, uri: Union[str, Path], single_output: bool = False self,
uri: Union[str, Path],
storage_types: Union[List[str], str],
single_output: bool = False,
) -> None: ) -> None:
"""Initialize the class.""" """Initialize the class."""
self.uri = uri self.uri = uri
if not isinstance(storage_types, list):
storage_types = [storage_types]
self._valid_inputs = storage_types
self.single_output = single_output self.single_output = single_output
def get_meta(self) -> Dict: def get_meta(self) -> Dict:
@ -51,31 +59,26 @@ class BaseFeatureStorage(ABC):
} }
return meta return meta
# TODO: is raising ValueError required? def validate(self, input_: List[str]) -> None:
@abstractmethod
def validate(self, input_: List[str]) -> bool:
"""Validate the input to the pipeline step. """Validate the input to the pipeline step.
Parameters Parameters
---------- ----------
input_ : list input_ : list of str
The input to the pipeline step. The input to the pipeline step.
Returns
-------
bool
Whether the `input` is valid or not.
Raises Raises
------ ------
ValueError ValueError
If the input does not have the required data. If the `input_` is invalid.
""" """
raise_error( if not any(x in input_ for x in self._valid_inputs):
msg="Concrete classes need to implement validate_input().", raise_error(
klass=NotImplementedError, "Input does not have the required data."
) f"\t Input: {input}"
f"\t Required (any of): {self._valid_inputs}"
)
@abstractmethod @abstractmethod
def list_features( def list_features(
@ -108,7 +111,7 @@ class BaseFeatureStorage(ABC):
feature_name: Optional[str] = None, feature_name: Optional[str] = None,
feature_md5: Optional[bool] = None, feature_md5: Optional[bool] = None,
) -> pd.DataFrame: ) -> pd.DataFrame:
"""Read feature from the storage. """Read feature into a pandas DataFrame.
Parameters Parameters
---------- ----------
@ -140,7 +143,7 @@ class BaseFeatureStorage(ABC):
Returns Returns
------- -------
str str
The metadata column The metadata column.
""" """
raise_error( raise_error(
@ -148,83 +151,38 @@ class BaseFeatureStorage(ABC):
klass=NotImplementedError, klass=NotImplementedError,
) )
# TODO: complete type annotations def store(self, kind: str, **kwargs) -> None:
@abstractmethod """Store extracted features data.
def store_matrix2d(
self,
data,
meta: Dict,
col_names: Optional[Iterable[str]] = None,
row_names: Optional[Iterable[str]] = None,
kind: Optional[str] = "full",
diagonal: bool = True,
) -> None:
"""Store 2D matrix.
Parameters Parameters
---------- ----------
data kind : {"matrix", "timeseries", "table"}
meta : dict The storage kind.
The metadata as a dictionary. **kwargs
col_names : list or tuple of str, optional The keyword arguments.
The column names (default None).
row_names : list of tuple of str, optional Raises
The row names (default None). ------
kind : str, optional ValueError
The kind of matrix: If `kind` is invalid.
- 'triu': store upper triangular only.
- 'tril': store lower triangular.
- 'full': full matrix (default 'full').
diagonal : bool, optional
Whether to store the diagonal (default True).
If kind == 'full', setting this to false will raise
an error
""" """
raise_error( if kind == "matrix":
msg="Concrete classes need to implement store_matrix2d().", self.store_matrix(**kwargs)
klass=NotImplementedError, elif kind == "timeseries":
) self.store_timeseries(**kwargs)
elif kind == "table":
self.store_table(**kwargs)
else:
raise ValueError(f"I don't know how to store {kind}")
# TODO: complete type annotations def store_df(self, **kwargs) -> None:
@abstractmethod """Store pandas DataFrame.
def store_table(
self,
data,
meta: Dict,
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Store table.
Parameters Parameters
---------- ----------
data **kwargs : dict
meta : dict The keyword arguments.
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).
"""
raise_error(
msg="Concrete classes need to implement store_table().",
klass=NotImplementedError,
)
@abstractmethod
def store_df(self, df: pd.DataFrame, meta: Dict) -> None:
"""Store pandas DataFerame.
Parameters
----------
df : pandas.DataFrame
The DataFrame to store.
meta : dict
The metadata as a dictionary.
""" """
raise_error( raise_error(
@ -232,16 +190,41 @@ class BaseFeatureStorage(ABC):
klass=NotImplementedError, klass=NotImplementedError,
) )
# TODO: complete type annotations def store_matrix(self, **kwargs) -> None:
@abstractmethod """Store matrix.
def store_timeseries(self, data, meta: Dict) -> None:
Parameters
----------
**kwargs : dict
The keyword arguments.
"""
raise_error(
msg="Concrete classes need to implement store_matrix2d().",
klass=NotImplementedError,
)
def store_table(self, **kwargs) -> None:
"""Store table.
Parameters
----------
**kwargs : dict
The keyword arguments.
"""
raise_error(
msg="Concrete classes need to implement store_table().",
klass=NotImplementedError,
)
def store_timeseries(self, **kwargs) -> None:
"""Store timeseries. """Store timeseries.
Parameters Parameters
---------- ----------
data **kwargs : dict
meta : dict The keyword arguments.
The metadata as a dictionary.
""" """
raise_error( raise_error(

View file

@ -59,6 +59,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
**kwargs: str, **kwargs: str,
) -> None: ) -> None:
"""Initialize the class.""" """Initialize the class."""
# Check upsert argument value
if upsert not in ["update", "ignore"]: if upsert not in ["update", "ignore"]:
raise_error( raise_error(
msg=( msg=(
@ -76,9 +77,16 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
"does not exist, creating now." "does not exist, creating now."
) )
uri.parent.mkdir(parents=True, exist_ok=True) uri.parent.mkdir(parents=True, exist_ok=True)
super().__init__(uri=uri, single_output=single_output, **kwargs) # Available storage kinds
storage_types = ["table", "timeseries", "matrix"]
super().__init__(
uri=uri,
storage_types=storage_types,
single_output=single_output,
**kwargs,
)
# Set upsert
self._upsert = upsert self._upsert = upsert
self._valid_inputs = ["table", "timeseries", "matrix"]
def get_engine(self, meta: Optional[Dict] = None) -> "Engine": def get_engine(self, meta: Optional[Dict] = None) -> "Engine":
"""Get engine. """Get engine.
@ -109,9 +117,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
) )
prefix = element_to_prefix(element) prefix = element_to_prefix(element)
# Format URI for engine creation # Format URI for engine creation
uri = ( uri = f"sqlite:///{self.uri.parent}/{prefix}{self.uri.name}"
"sqlite:///" f"{self.uri.parent}/{prefix}{self.uri.name}"
) # type: ignore
return create_engine(uri, echo=False) return create_engine(uri, echo=False)
def _save_upsert( def _save_upsert(
@ -203,10 +209,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
msg=f"Invalid option {if_exists} for if_exists." msg=f"Invalid option {if_exists} for if_exists."
) )
# TODO: complete type annotations def _store_2d(
def store_2d(
self, self,
data, data: Dict,
meta: Dict, meta: Dict,
columns: Optional[Iterable[str]] = None, columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None, rows_col_name: Optional[str] = None,
@ -215,7 +220,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
Parameters Parameters
---------- ----------
data data : dict
The data to store.
meta : dict meta : dict
The metadata as a dictionary. The metadata as a dictionary.
columns : list or tuple of str, optional columns : list or tuple of str, optional
@ -232,32 +238,10 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
meta=meta, n_rows=n_rows, rows_col_name=rows_col_name meta=meta, n_rows=n_rows, rows_col_name=rows_col_name
) )
# Prepare new dataframe # Prepare new dataframe
data_df = pd.DataFrame( data_df = pd.DataFrame(data, columns=columns, index=idx)
data, columns=columns, index=idx
) # type: ignore
# Store dataframe # Store dataframe
self.store_df(df=data_df, meta=meta) self.store_df(df=data_df, meta=meta)
def validate(self, input_: List[str]) -> bool:
"""Implement input validation.
Parameters
----------
input_ : list of str
The input to the pipeline step.
Returns
-------
bool
Whether the `input` is valid or not.
"""
# Convert input to list
if not isinstance(input_, list):
input_ = [input_]
return all(x in self._valid_inputs for x in input_)
def list_features( def list_features(
self, return_df: bool = False self, return_df: bool = False
) -> Union[Dict[str, Dict], pd.DataFrame]: ) -> Union[Dict[str, Dict], pd.DataFrame]:
@ -293,7 +277,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
feature_name: Optional[str] = None, feature_name: Optional[str] = None,
feature_md5: Optional[str] = None, feature_md5: Optional[str] = None,
) -> pd.DataFrame: ) -> pd.DataFrame:
"""Implement feature reading from the storage. """Implement feature reading into a pandas DataFrame.
Either one of `feature_name` or `feature_md5` needs to be specified. Either one of `feature_name` or `feature_md5` needs to be specified.
@ -399,130 +383,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
self._save_upsert(meta_df, "meta", engine) self._save_upsert(meta_df, "meta", engine)
return f"meta_{meta_md5}" return f"meta_{meta_md5}"
# TODO: complete type annotations
def store_matrix2d(
self,
data,
meta: Dict,
col_names: Optional[List[str]] = None,
row_names: Optional[List[str]] = None,
kind: Optional[str] = "full",
diagonal: bool = True,
) -> None:
"""Implement 2D matrix storing.
Parameters
----------
data
meta : dict
The metadata as a dictionary.
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).
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 (default True).
If kind == 'full', setting this to false will raise
an error
"""
if diagonal is False and kind not in ["triu", "tril"]:
raise_error(
msg="Diagonal cannot be False if kind is not full",
klass=ValueError,
)
if kind in ["triu", "tril"]:
if data.shape[0] != data.shape[1]:
raise_error(
"Cannot store a non-square matrix as a triangular matrix",
klass=ValueError,
)
if kind == "triu":
k = 0 if diagonal is True else 1
data_idx = np.triu_indices(data.shape[0], k=k)
elif kind == "tril":
k = 0 if diagonal is True else -1
data_idx = np.tril_indices(data.shape[0], k=k)
elif kind == "full":
data_idx = (
np.repeat(np.arange(data.shape[0]), data.shape[1]),
np.tile(np.arange(data.shape[1]), data.shape[0]),
)
else:
raise_error(msg=f"Invalid kind {kind}", klass=ValueError)
if row_names is None:
row_names = [f"r{i}" for i in range(data.shape[0])]
elif len(row_names) != data.shape[0]:
raise_error(
msg="Number of row names does not match number of rows",
klass=ValueError,
)
if col_names is None:
col_names = [f"c{i}" for i in range(data.shape[1])]
elif len(col_names) != data.shape[1]:
raise_error(
msg="Number of column names does not match number of columns",
klass=ValueError,
)
flat_data = data[data_idx]
columns = [
f"{row_names[i]}~{col_names[j]}"
for i, j in zip(data_idx[0], data_idx[1])
]
# Convert element metadata to index
n_rows = 1
idx = element_to_index(meta=meta, n_rows=n_rows, rows_col_name=None)
# Prepare new dataframe
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
if len(columns) > 2000: # TODO: check SQLITE_MAX_COLUMN
data_df = data_df.stack()
new_names = [x for x in data_df.index.names[:-1]]
new_names.append("pair")
data_df.index.names = new_names
# Store dataframe
self.store_df(df=data_df, meta=meta) # type: ignore
# TODO: complete type annotations
def store_table(
self,
data,
meta: Dict,
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Implement table storing.
Parameters
----------
data
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_df(self, df: pd.DataFrame, meta: Dict) -> None: def store_df(self, df: pd.DataFrame, meta: Dict) -> None:
"""Implement dataframe storing. """Implement pandas DataFrame storing.
Parameters Parameters
---------- ----------
@ -568,20 +430,158 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
# Save data # Save data
self._save_upsert(df, table_name, engine) self._save_upsert(df, table_name, engine)
# TODO: complete type annotations def store_matrix(
def store_timeseries(self, data, meta: Dict) -> None: self,
data: Dict,
meta: Dict,
col_names: Optional[List[str]] = None,
row_names: Optional[List[str]] = None,
matrix_kind: Optional[str] = "full",
diagonal: bool = True,
) -> None:
"""Implement matrix storing.
Parameters
----------
data : dict
The matrix data to store.
meta : dict
The metadata as a dictionary.
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)..
"""
if diagonal is False and matrix_kind not in ["triu", "tril"]:
raise_error(
msg="Diagonal cannot be False if kind is not full",
klass=ValueError,
)
if matrix_kind in ["triu", "tril"]:
if data.shape[0] != data.shape[1]:
raise_error(
"Cannot store a non-square matrix as a triangular matrix",
klass=ValueError,
)
if matrix_kind == "triu":
k = 0 if diagonal is True else 1
data_idx = np.triu_indices(data.shape[0], k=k)
elif matrix_kind == "tril":
k = 0 if diagonal is True else -1
data_idx = np.tril_indices(data.shape[0], k=k)
elif matrix_kind == "full":
data_idx = (
np.repeat(np.arange(data.shape[0]), data.shape[1]),
np.tile(np.arange(data.shape[1]), data.shape[0]),
)
else:
raise_error(msg=f"Invalid kind {matrix_kind}", klass=ValueError)
if row_names is None:
row_names = [f"r{i}" for i in range(data.shape[0])]
elif len(row_names) != data.shape[0]:
raise_error(
msg="Number of row names does not match number of rows",
klass=ValueError,
)
if col_names is None:
col_names = [f"c{i}" for i in range(data.shape[1])]
elif len(col_names) != data.shape[1]:
raise_error(
msg="Number of column names does not match number of columns",
klass=ValueError,
)
flat_data = data[data_idx]
columns = [
f"{row_names[i]}~{col_names[j]}"
for i, j in zip(data_idx[0], data_idx[1])
]
# Convert element metadata to index
n_rows = 1
idx = element_to_index(meta=meta, n_rows=n_rows, rows_col_name=None)
# Prepare new dataframe
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
if len(columns) > 2000: # TODO: check SQLITE_MAX_COLUMN
data_df = data_df.stack()
new_names = [x for x in data_df.index.names[:-1]]
new_names.append("pair")
data_df.index.names = new_names
# Store dataframe
self.store_df(df=data_df, meta=meta)
def store_table(
self,
data: Dict,
meta: Dict,
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
) -> None:
"""Implement table storing.
Parameters
----------
data : dict
The table data to store.
meta : dict
The metadata as a dictionary.
columns : list or tuple of str, optional
The columns (default None).
rows_col_name : str, optional
The column name to use in case number of rows greater than 1.
If None and number of rows greater than 1, then the name will be
"index" (default None).
"""
self._store_2d(
data=data, meta=meta, columns=columns, rows_col_name=rows_col_name
)
def store_timeseries(
self,
data: Dict,
meta: Dict,
columns: Optional[Iterable[str]] = None,
row_names: str = "timepoint",
) -> None:
"""Implement timeseries storing. """Implement timeseries storing.
Parameters Parameters
---------- ----------
data data: dict
The timeseries data to store.
meta : dict meta : dict
The metadata as a dictionary. 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").
""" """
raise_error( self._store_2d(
msg="store_timeseries() not implemented.", data=data,
klass=NotImplementedError, meta=meta,
columns=columns,
rows_col_name="timepoint", # explicit so as to stop overriding
) )
def collect(self) -> None: def collect(self) -> None:

View file

@ -406,8 +406,8 @@ def test_store_table(tmp_path: Path) -> None:
assert_frame_equal(df_new, c_df_new) assert_frame_equal(df_new, c_df_new)
def test_store_matrix2d(tmp_path: Path) -> None: def test_store_matrix(tmp_path: Path) -> None:
"""Test 2D Matrix store. """Test matrix store.
Parameters Parameters
---------- ----------
@ -428,8 +428,11 @@ def test_store_matrix2d(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"] col_names = ["col1", "col2", "col3"]
# Store table # Store table
storage.store_matrix2d( storage.store_matrix(
data, meta, row_names=row_names, col_names=col_names data=data,
meta=meta,
row_names=row_names,
col_names=col_names,
) )
stored_names = [f"{i}~{j}" for i in row_names for j in col_names] stored_names = [f"{i}~{j}" for i in row_names for j in col_names]
@ -445,7 +448,7 @@ def test_store_matrix2d(tmp_path: Path) -> None:
# Store without row and column names # Store without row and column names
uri = tmp_path / "test_store_table_nonames.db" uri = tmp_path / "test_store_table_nonames.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True) storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d(data, meta) storage.store_matrix(data=data, meta=meta)
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])
@ -458,13 +461,18 @@ def test_store_matrix2d(tmp_path: Path) -> None:
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_matrix2d(data, meta, kind="wrong") storage.store_matrix(data=data, meta=meta, matrix_kind="wrong")
with pytest.raises(ValueError, match="non-square"): with pytest.raises(ValueError, match="non-square"):
storage.store_matrix2d(data, meta, kind="triu") storage.store_matrix(data=data, meta=meta, matrix_kind="triu")
with pytest.raises(ValueError, match="cannot be False"): with pytest.raises(ValueError, match="cannot be False"):
storage.store_matrix2d(data, meta, kind="full", diagonal=False) storage.store_matrix(
data=data,
meta=meta,
matrix_kind="full",
diagonal=False,
)
# Store upper triangular matrix # Store upper triangular matrix
data = np.array([[1, 2, 3], [11, 22, 33], [111, 222, 333]]) data = np.array([[1, 2, 3], [11, 22, 33], [111, 222, 333]])
@ -472,8 +480,12 @@ def test_store_matrix2d(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"] col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_triu.db" uri = tmp_path / "test_store_table_triu.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True) storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d( storage.store_matrix(
data, meta, kind="triu", row_names=row_names, col_names=col_names data=data,
meta=meta,
matrix_kind="triu",
row_names=row_names,
col_names=col_names,
) )
stored_names = [ stored_names = [
@ -497,10 +509,10 @@ def test_store_matrix2d(tmp_path: Path) -> None:
# Store upper triangular matrix without diagonal # Store upper triangular matrix without diagonal
uri = tmp_path / "test_store_table_triu_nodiagonal.db" uri = tmp_path / "test_store_table_triu_nodiagonal.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True) storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d( storage.store_matrix(
data, data=data,
meta, meta=meta,
kind="triu", matrix_kind="triu",
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
diagonal=False, diagonal=False,
@ -527,8 +539,12 @@ def test_store_matrix2d(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"] col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_tril.db" uri = tmp_path / "test_store_table_tril.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True) storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d( storage.store_matrix(
data, meta, kind="tril", row_names=row_names, col_names=col_names data=data,
meta=meta,
matrix_kind="tril",
row_names=row_names,
col_names=col_names,
) )
stored_names = [ stored_names = [
@ -552,10 +568,10 @@ def test_store_matrix2d(tmp_path: Path) -> None:
# Store lower triangular matrix without diagonal # Store lower triangular matrix without diagonal
uri = tmp_path / "test_store_table_tril_nodiagonal.db" uri = tmp_path / "test_store_table_tril_nodiagonal.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True) storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d( storage.store_matrix(
data, data,
meta, meta,
kind="tril", matrix_kind="tril",
row_names=row_names, row_names=row_names,
col_names=col_names, col_names=col_names,
diagonal=False, diagonal=False,

View file

@ -12,51 +12,50 @@ 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") # type: ignore BaseFeatureStorage(uri="/tmp", storage_types=["matrix"])
def test_BaseFeatureStorage() -> None: def test_BaseFeatureStorage() -> None:
"""Test BaseFeatureStorage.""" """Test proper subclassing of BaseFeatureStorage."""
# Create concrete class # Create concrete class
class MyFeatureStorage(BaseFeatureStorage): class MyFeatureStorage(BaseFeatureStorage):
def __init__(self, uri, single_output=False): """Implement concrete class."""
super().__init__(uri, single_output=single_output)
def validate(self, input): def __init__(self, uri, single_output=False):
super().validate(input) storage_types = ["matrix"]
super().__init__(
uri=uri,
storage_types=storage_types,
single_output=single_output,
)
def list_features(self): def list_features(self):
super().list_features() super().list_features()
def read_df(self, feature_name=None, feature_md5=None): def read_df(self, feature_name=None, feature_md5=None):
super().read_df(feature_name=feature_name, feature_md5=feature_md5) super().read_df(
feature_name=feature_name,
feature_md5=feature_md5,
)
def store_metadata(self, metadata): def store_metadata(self, metadata):
super().store_metadata(metadata) super().store_metadata(metadata)
def store_matrix2d(self, matrix, meta):
super().store_matrix2d(matrix, meta)
def store_table(self, table, meta):
super().store_table(table, meta)
def store_df(self, df, meta):
super().store_df(df, meta)
def store_timeseries(self, timeseries, meta):
super().store_timeseries(timeseries, meta)
def collect(self): def collect(self):
return super().collect() return super().collect()
# Check single_output is False
st = MyFeatureStorage(uri="/tmp") st = MyFeatureStorage(uri="/tmp")
assert st.single_output is False assert st.single_output is False
# Check single_output is True
st = MyFeatureStorage(uri="/tmp", single_output=True) st = MyFeatureStorage(uri="/tmp", single_output=True)
assert st.single_output is True assert st.single_output is True
with pytest.raises(NotImplementedError): # Check validate with valid argument
st.validate(None) st.validate(input_=["matrix"])
# Check validate with invalid argument
with pytest.raises(ValueError):
st.validate(input_=["table"])
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.list_features() st.list_features()
@ -67,19 +66,19 @@ def test_BaseFeatureStorage() -> None:
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.store_metadata(None) st.store_metadata(None)
with pytest.raises(NotImplementedError):
st.store_matrix2d(None, None)
with pytest.raises(NotImplementedError):
st.store_table(None, None)
with pytest.raises(NotImplementedError):
st.store_df(None, None) # type: ignore
with pytest.raises(NotImplementedError):
st.store_timeseries(None, None)
with pytest.raises(NotImplementedError): with pytest.raises(NotImplementedError):
st.collect() st.collect()
with pytest.raises(NotImplementedError):
st.store(kind="matrix")
with pytest.raises(NotImplementedError):
st.store(kind="timeseries")
with pytest.raises(NotImplementedError):
st.store(kind="table")
with pytest.raises(ValueError):
st.store(kind="lego")
assert st.uri == "/tmp" assert st.uri == "/tmp"