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 PIOP1 and PIOP2 AOMIC datasets and refactor AOMICID1000 slightly
- Implement PIOP1 and PIOP2 AOMIC datasets and refactor AOMICID1000 slightly
(:gh: `94` 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`_).
- Upgrade storage interface for storage-like objects (:gh: `84` by `Synchon Mandal`_).
- Add missing type annotations (:gh: `74` by `Synchon Mandal`_).
Bugs
~~~~

View file

@ -2,7 +2,7 @@
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
(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 junifer.testing.registry # noqa: F401
from junifer.api import collect, run
from junifer.storage import SQLiteFeatureStorage
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:
configure_logging(level="INFO")
##############################################################################
# Define the datagrabber interface
datagrabber = {
"kind": "SPMAuditoryTestingDatagrabber",
}
###############################################################################
# Define the markers you want:
marker_dicts = [
# Define the markers interface
markers = [
{
"name": "Schaefer100x17_RSSETS",
"kind": "RSSETSMarker",
@ -39,26 +45,34 @@ marker_dicts = [
},
]
###############################################################################
# Create a temporary directory for junifer feature extraction:
# At the end you can read the extracted data into a ``pandas.DataFrame``.
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
# TODO: needs SQLiteFeatureStorage.store_timeseries() to be
# 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")
###############################################################################
# Now we take a look at the dataframe
df_vbm.head()

View file

@ -7,5 +7,7 @@
from .base import BaseMarker
from .collection import MarkerCollection
from .ets_rss import RSSETSMarker
from .functional_connectivity_atlas import FunctionalConnectivityAtlas
from .functional_connectivity_spheres import FunctionalConnectivitySpheres
from .parcel import ParcelAggregation
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>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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 ..utils import logger, raise_error
class BaseMarker(PipelineStepMixin):
"""Base class for all markers.
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
class BaseMarker(ABC, PipelineStepMixin):
"""Abstract base class for all markers.
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
available data (default None).
available data.
name : str, optional
The name of the marker. By default, it will use the class name as the
name of the marker (default None).
@ -44,7 +49,7 @@ class BaseMarker(PipelineStepMixin):
Returns
-------
dict
The metadata as a dictionary.
The metadata as a dictionary with the only key 'marker'.
"""
s_meta = super().get_meta()
@ -76,6 +81,7 @@ class BaseMarker(PipelineStepMixin):
f"\t Required (any of): {self._valid_inputs}"
)
@abstractmethod
def get_output_kind(self, input: List[str]) -> List[str]:
"""Get output kind.
@ -92,9 +98,11 @@ class BaseMarker(PipelineStepMixin):
"""
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:
"""Compute.
@ -117,32 +125,54 @@ class BaseMarker(PipelineStepMixin):
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
def store(self, kind: str, out: Dict, storage) -> None:
@abstractmethod
def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
input
out
kind : str
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(self, input: Dict[str, Dict], storage=None) -> Dict:
def fit_transform(
self,
input: Dict[str, Dict],
storage: "BaseFeatureStorage" = None,
) -> Dict:
"""Fit and transform.
Parameters
----------
input
storage
input : dict
The Junifer Data object.
storage : storage-like, optional
The storage class, for example, SQLiteFeatureStorage.
Returns
-------
dict
The processed output as a dictionary. If `storage` is provided,
empty dictionary is returned.
"""
out = {}
@ -156,11 +186,11 @@ class BaseMarker(PipelineStepMixin):
t_meta = meta.copy()
t_meta.update(t_input.get("meta", {}))
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)
if storage is not None:
logger.info(f"Storing in {storage}")
self.store(kind, t_out, storage)
self.store(kind=kind, out=t_out, storage=storage)
else:
logger.info("No storage specified, returning dictionary")
out[kind] = t_out

View file

@ -5,7 +5,7 @@
# License: AGPL
from collections import Counter
from typing import Dict, List, Optional
from typing import TYPE_CHECKING, Dict, List, Optional
from ..datareader.default import DefaultDataReader
from ..markers.base import BaseMarker
@ -14,15 +14,23 @@ from ..storage.base import BaseFeatureStorage
from ..utils import logger
if TYPE_CHECKING:
from junifer.datagrabber import BaseDataGrabber
class MarkerCollection:
"""Class for marker collection.
Parameters
----------
markers
datareader
preprocessing
storage
markers : list of marker-like
The markers to compute.
datareader : datareader-like, optional
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
-------
output : dict or None
dict or None
The output of the pipeline. Each key represents a marker name and
the values are the computer marker values. If the pipeline has a
storage configured, then the output will be None.
"""
logger.info("Fitting pipeline")
# Fetch actual data using datareader
data = self._datareader.fit_transform(input)
# Apply preprocessing steps
if self._preprocessing is not None:
logger.info("Preprocessing data")
data = self._preprocessing.fit_transform(data)
# Compute markers
out = {}
for marker in self._markers:
logger.info(f"Fitting marker {marker.name}")
@ -78,20 +92,21 @@ class MarkerCollection:
if self._storage is None:
out[marker.name] = m_value
logger.info("Marker collection fitting done")
return None if self._storage else out
# TODO: complete type annotations
def validate(self, datagrabber) -> None:
def validate(self, datagrabber: "BaseDataGrabber") -> None:
"""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
present and streamed down the steps. Also, if a storage is configured,
check that the storage can handle the markers output.
Parameters
----------
datagrabber
datagrabber : datagrabber-like
The datagrabber to validate.
"""
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:
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}")
# Check storage for the marker
if self._storage is not None:
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>
# License: AGPL
from typing import Dict, List
from typing import TYPE_CHECKING, Any, Dict, List, Optional
import numpy as np
@ -17,6 +17,10 @@ from .parcel import ParcelAggregation
from .utils import _ets
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker
class RSSETSMarker(BaseMarker):
"""Class for root sum of squares of edgewise timeseries.
@ -24,9 +28,14 @@ class RSSETSMarker(BaseMarker):
Parameters
----------
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
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"]
# TODO: complete type annotations
def store(self, kind: str, out, storage) -> None:
def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
kind
out
storage
kind : {"BOLD"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing BOLD in {storage}")
storage.store_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.
Take a timeseries of brain areas, and calculate timeseries for each
@ -81,12 +101,19 @@ class RSSETSMarker(BaseMarker):
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
-------
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
----------
@ -103,9 +130,10 @@ class RSSETSMarker(BaseMarker):
method=self.aggregation_method,
)
# Compute the parcel aggregation
out = parcel_aggregation.compute(input)
out = parcel_aggregation.compute(input=input, extra_input=extra_input)
edge_ts = _ets(out["data"])
# Compute the RSS
out["data"] = np.sum(edge_ts**2, 1) ** 0.5
# Set correct column label
out["columns"] = ["root_sum_of_squares_ets"]
return out

View file

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

View file

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

View file

@ -4,7 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Dict, List, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import numpy as np
from nilearn.image import math_img, resample_to_img
@ -17,22 +17,42 @@ from ..utils import logger
from .base import BaseMarker
if TYPE_CHECKING:
from junifer.storage import BaseFeatureStorage
@register_marker
class ParcelAggregation(BaseMarker):
"""Class for parcel aggregation.
Parameters
----------
atlas
method
method_params
on
name
atlas : str
The name of the atlas. Check valid options by calling
:func:`junifer.data.list_atlases`.
method : str
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__(
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:
"""Initialize the class."""
self.atlas = atlas
@ -52,8 +72,8 @@ class ParcelAggregation(BaseMarker):
Returns
-------
str
The kind of output.
list of str
The list of storage kinds.
"""
outputs = []
@ -66,32 +86,43 @@ class ParcelAggregation(BaseMarker):
raise ValueError(f"Unknown input kind for {t_input}")
return outputs
# TODO: complete type annotations
def store(self, kind: str, out, storage) -> None:
def store(
self,
kind: str,
out: Dict[str, Any],
storage: "BaseFeatureStorage",
) -> None:
"""Store.
Parameters
----------
kind
out
storage
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
The data kind to store.
out : dict
The computed result as a dictionary to store.
storage : storage-like
The storage class, for example, SQLiteFeatureStorage.
"""
logger.debug(f"Storing {kind} in {storage}")
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
storage.store_table(**out)
if kind in ["BOLD"]:
storage.store_timeseries(**out)
storage.store(kind="table", **out)
elif kind in ["BOLD"]:
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.
Parameters
----------
input : Dict[str, Dict]
input : dict
A single input from the pipeline data object in which to compute
the marker.
extra_input : Dict, optional
extra_input : dict, optional
The other fields in the pipeline data object. Useful for accessing
other data kind that needs to be used in the computation. For
example, the functional connectivity markers can make use of the
@ -102,17 +133,24 @@ class ParcelAggregation(BaseMarker):
dict
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.
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"]
logger.debug(f"Parcel aggregation using {self.method}")
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
resolution = np.min(t_input.header.get_zooms()[:3]) # type: ignore
t_atlas, t_labels, _ = load_atlas(self.atlas, resolution=resolution)
resolution = np.min(t_input.header.get_zooms()[:3])
t_atlas, t_labels, _ = load_atlas(
name=self.atlas,
resolution=resolution,
)
atlas_img_res = resample_to_img(
t_atlas,
t_input,

View file

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

View file

@ -4,6 +4,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from pathlib import Path
import pytest
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.
Parameters
@ -123,7 +125,9 @@ def test_MarkerCollection_storage(tmp_path) -> None:
uri = tmp_path / "test_marker_collection_storage.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
mc = MarkerCollection(
markers=markers, storage=storage, datareader=DefaultDataReader()
markers=markers,
storage=storage,
datareader=DefaultDataReader(),
)
mc.validate(dg)
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.markers.ets_rss import RSSETSMarker
from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
def test_compute() -> None:
"""Test RSS ETS."""
atlas = "Schaefer100x17"
test_atlas, _, _ = load_atlas(atlas)
# Set atlas
ATLAS = "Schaefer100x17"
def test_compute() -> None:
"""Test RSS ETS compute()."""
with SPMAuditoryTestingDatagrabber() as dg:
# Fetch element
out = dg["sub001"]
# Load BOLD image
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
# Create input data
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
# Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(atlas=atlas)
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
new_out = ets_rss_marker.compute(input_dict)
# Load atlas
test_atlas, _, _ = load_atlas(ATLAS)
# Compute the NiftiLabelsMasker
test_masker = NiftiLabelsMasker(test_atlas)
test_ts = test_masker.fit_transform(niimg)
# Assert the dimension of timeseries
n_time, _ = test_ts.shape
assert n_time == len(new_out["data"])
# Assert the meta
meta = ets_rss_marker.get_meta("BOLD")["marker"]
assert meta["atlas"] == "Schaefer100x17"
@ -42,10 +51,8 @@ def test_compute() -> None:
def test_get_output_kind() -> None:
"""Test get_output_kind."""
atlas = "Schaefer100x17"
ets_rss_marker = RSSETSMarker(atlas=atlas)
"""Test RSS ETS get_output_kind()."""
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
input_list = ["BOLD"]
input_list = ets_rss_marker.get_output_kind(input_list)
assert len(input_list) == 1
@ -53,19 +60,26 @@ def test_get_output_kind() -> 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:
# Fetch element
out = dg["sub001"]
# Load BOLD image
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
# Compute the RSSETSMarker
ets_rss_marker = RSSETSMarker(atlas=atlas)
_ = ets_rss_marker.compute(input_dict)
# TODO: Needs store_timeseries implemented for SQLiteFeatureStorage
# storage = {
# "kind": "SQLiteFeatureStorage",
# "uri": str((tmp_path / "test.db").absolute()),
# }
# ets_rss_marker.store("SQLiteFeatureStorage", new_out, storage)
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
# Create storage
storage = SQLiteFeatureStorage(
uri=str((tmp_path / "test.db").absolute()),
single_output=True,
)
# Store
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>
# Kaustubh R. Patil <k.patil@fz-juelich.de>
@ -19,8 +19,14 @@ from junifer.storage import SQLiteFeatureStorage
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
ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
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>
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# Federico Raimondo <f.raimondo@fz-juelich.de>
# License: AGPL
import pytest
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.connectome import ConnectivityMeasure
from numpy.testing import assert_array_almost_equal
from sklearn.covariance import EmpiricalCovariance
from junifer.markers.functional_connectivity_spheres import (
FunctionalConnectivitySpheres,
@ -30,7 +29,6 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
The path to the test directory.
"""
# get a dataset
ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
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
cm = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), # type: ignore
kind="correlation"
cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
)
out_ni = cm.fit_transform([ts["data"]])[0]
assert_array_almost_equal(out_ni, out["data"], decimal=3)

View file

@ -4,94 +4,57 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import List, Optional
import pytest
from junifer.markers.base import BaseMarker
@pytest.mark.parametrize(
"on, name, kind, expected_class, expected_name",
[
(["bold", "dwi"], None, "bold", "BaseMarker", "bold_BaseMarker"),
(["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_base_marker_abstractness() -> None:
"""Test BaseMarker is abstract base class."""
with pytest.raises(TypeError, match=r"abstract"):
BaseMarker(on=["BOLD"])
def test_compute_parameters() -> None:
"""Test compute parameters."""
base = BaseMarker(on=["bold", "dwi"], name="mymarker")
base.compute = lambda x, y: { # type: ignore
"data": x.keys(),
"extra": y.keys(),
def test_base_datagrabber_subclassing() -> None:
"""Test proper subclassing of BaseMarker."""
# Create concrete class
class MyBaseMarker(BaseMarker):
def get_output_kind(self, input):
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"}}
out = base.fit_transform(
input_,
)
assert list(out["bold"]["data"]) == ["path"]
assert list(out["bold"]["extra"]) == ["t2"]
def test_BaseMarker() -> None:
"""Test base class."""
base = BaseMarker(on=["bold", "dwi"], name="mymarker")
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()))
marker = MyBaseMarker(on=["BOLD"])
output = marker.fit_transform(input=input_) # process
# Check output
assert "BOLD" in output
assert "data" in output["BOLD"]
assert "columns" in output["BOLD"]
assert "row_names" in output["BOLD"]
assert "meta" in output["BOLD"]
assert "datagrabber" in output["BOLD"]["meta"]
assert "element" in output["BOLD"]["meta"]
assert "datareader" in output["BOLD"]["meta"]
# Check no implementation check
with pytest.raises(NotImplementedError):
base.get_output_kind(list(wrong_input.keys()))
marker.store(kind="kind", out="out", storage="storage")
with pytest.raises(NotImplementedError):
base.fit_transform(input_)
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"
# Check attributes
assert marker.name == "MyBaseMarker"

View file

@ -138,19 +138,17 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
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 = {
# "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)
marker.fit_transform(input, storage=storage)

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

View file

@ -59,6 +59,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
**kwargs: str,
) -> None:
"""Initialize the class."""
# Check upsert argument value
if upsert not in ["update", "ignore"]:
raise_error(
msg=(
@ -76,9 +77,16 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
"does not exist, creating now."
)
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._valid_inputs = ["table", "timeseries", "matrix"]
def get_engine(self, meta: Optional[Dict] = None) -> "Engine":
"""Get engine.
@ -109,9 +117,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
)
prefix = element_to_prefix(element)
# Format URI for engine creation
uri = (
"sqlite:///" f"{self.uri.parent}/{prefix}{self.uri.name}"
) # type: ignore
uri = f"sqlite:///{self.uri.parent}/{prefix}{self.uri.name}"
return create_engine(uri, echo=False)
def _save_upsert(
@ -203,10 +209,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
msg=f"Invalid option {if_exists} for if_exists."
)
# TODO: complete type annotations
def store_2d(
def _store_2d(
self,
data,
data: Dict,
meta: Dict,
columns: Optional[Iterable[str]] = None,
rows_col_name: Optional[str] = None,
@ -215,7 +220,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
Parameters
----------
data
data : dict
The data to store.
meta : dict
The metadata as a dictionary.
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
)
# Prepare new dataframe
data_df = pd.DataFrame(
data, columns=columns, index=idx
) # type: ignore
data_df = pd.DataFrame(data, columns=columns, index=idx)
# Store dataframe
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(
self, return_df: bool = False
) -> Union[Dict[str, Dict], pd.DataFrame]:
@ -293,7 +277,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
feature_name: Optional[str] = None,
feature_md5: Optional[str] = None,
) -> 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.
@ -399,130 +383,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
self._save_upsert(meta_df, "meta", engine)
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:
"""Implement dataframe storing.
"""Implement pandas DataFrame storing.
Parameters
----------
@ -568,20 +430,158 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
# Save data
self._save_upsert(df, table_name, engine)
# TODO: complete type annotations
def store_timeseries(self, data, meta: Dict) -> None:
def store_matrix(
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.
Parameters
----------
data
data: dict
The timeseries data to store.
meta : dict
The metadata as a dictionary.
columns : list or tuple of str, optional
The column labels (default None).
row_names : str, optional
The column name to use in case number of rows greater than 1
(default "timepoint").
"""
raise_error(
msg="store_timeseries() not implemented.",
klass=NotImplementedError,
self._store_2d(
data=data,
meta=meta,
columns=columns,
rows_col_name="timepoint", # explicit so as to stop overriding
)
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)
def test_store_matrix2d(tmp_path: Path) -> None:
"""Test 2D Matrix store.
def test_store_matrix(tmp_path: Path) -> None:
"""Test matrix store.
Parameters
----------
@ -428,8 +428,11 @@ def test_store_matrix2d(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"]
# Store table
storage.store_matrix2d(
data, meta, row_names=row_names, col_names=col_names
storage.store_matrix(
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]
@ -445,7 +448,7 @@ def test_store_matrix2d(tmp_path: Path) -> None:
# Store without row and column names
uri = tmp_path / "test_store_table_nonames.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d(data, meta)
storage.store_matrix(data=data, meta=meta)
stored_names = [
f"r{i}~c{j}"
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
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"):
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"):
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
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"]
uri = tmp_path / "test_store_table_triu.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d(
data, meta, kind="triu", row_names=row_names, col_names=col_names
storage.store_matrix(
data=data,
meta=meta,
matrix_kind="triu",
row_names=row_names,
col_names=col_names,
)
stored_names = [
@ -497,10 +509,10 @@ def test_store_matrix2d(tmp_path: Path) -> None:
# Store upper triangular matrix without diagonal
uri = tmp_path / "test_store_table_triu_nodiagonal.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d(
data,
meta,
kind="triu",
storage.store_matrix(
data=data,
meta=meta,
matrix_kind="triu",
row_names=row_names,
col_names=col_names,
diagonal=False,
@ -527,8 +539,12 @@ def test_store_matrix2d(tmp_path: Path) -> None:
col_names = ["col1", "col2", "col3"]
uri = tmp_path / "test_store_table_tril.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d(
data, meta, kind="tril", row_names=row_names, col_names=col_names
storage.store_matrix(
data=data,
meta=meta,
matrix_kind="tril",
row_names=row_names,
col_names=col_names,
)
stored_names = [
@ -552,10 +568,10 @@ def test_store_matrix2d(tmp_path: Path) -> None:
# Store lower triangular matrix without diagonal
uri = tmp_path / "test_store_table_tril_nodiagonal.db"
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
storage.store_matrix2d(
storage.store_matrix(
data,
meta,
kind="tril",
matrix_kind="tril",
row_names=row_names,
col_names=col_names,
diagonal=False,

View file

@ -12,51 +12,50 @@ from junifer.storage.base import BaseFeatureStorage
def test_BaseFeatureStorage_abstractness() -> None:
"""Test BaseFeatureStorage is abstract base class."""
with pytest.raises(TypeError, match=r"abstract"):
BaseFeatureStorage(uri="/tmp") # type: ignore
BaseFeatureStorage(uri="/tmp", storage_types=["matrix"])
def test_BaseFeatureStorage() -> None:
"""Test BaseFeatureStorage."""
"""Test proper subclassing of BaseFeatureStorage."""
# Create concrete class
class MyFeatureStorage(BaseFeatureStorage):
def __init__(self, uri, single_output=False):
super().__init__(uri, single_output=single_output)
"""Implement concrete class."""
def validate(self, input):
super().validate(input)
def __init__(self, uri, single_output=False):
storage_types = ["matrix"]
super().__init__(
uri=uri,
storage_types=storage_types,
single_output=single_output,
)
def list_features(self):
super().list_features()
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):
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):
return super().collect()
# Check single_output is False
st = MyFeatureStorage(uri="/tmp")
assert st.single_output is False
# Check single_output is True
st = MyFeatureStorage(uri="/tmp", single_output=True)
assert st.single_output is True
with pytest.raises(NotImplementedError):
st.validate(None)
# Check validate with valid argument
st.validate(input_=["matrix"])
# Check validate with invalid argument
with pytest.raises(ValueError):
st.validate(input_=["table"])
with pytest.raises(NotImplementedError):
st.list_features()
@ -67,19 +66,19 @@ def test_BaseFeatureStorage() -> None:
with pytest.raises(NotImplementedError):
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):
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"