diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index b418273f5..4d83e2d8b 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -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 ~~~~ diff --git a/examples/run_ets_rss_marker.py b/examples/run_ets_rss_marker.py index dc8901bdd..79b2b1006 100644 --- a/examples/run_ets_rss_marker.py +++ b/examples/run_ets_rss_marker.py @@ -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() diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index 0fe3ca032..6fcc780b9 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -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 diff --git a/junifer/markers/base.py b/junifer/markers/base.py index bb4095bcb..12227ce79 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -1,23 +1,28 @@ -"""Provide base class for markers.""" +"""Provide abstract base class for markers.""" # Authors: Federico Raimondo # Synchon Mandal # 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 diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index 1897783f7..4fe2a0f33 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -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: 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) diff --git a/junifer/markers/ets_rss.py b/junifer/markers/ets_rss.py index c9c3d24ae..668773201 100644 --- a/junifer/markers/ets_rss.py +++ b/junifer/markers/ets_rss.py @@ -6,7 +6,7 @@ # Synchon Mandal # 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 diff --git a/junifer/markers/functional_connectivity_atlas.py b/junifer/markers/functional_connectivity_atlas.py index a0ee6ffcf..4a84f215c 100644 --- a/junifer/markers/functional_connectivity_atlas.py +++ b/junifer/markers/functional_connectivity_atlas.py @@ -4,7 +4,7 @@ # Kaustubh R. Patil # 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) diff --git a/junifer/markers/functional_connectivity_spheres.py b/junifer/markers/functional_connectivity_spheres.py index 36c401e25..e65a6f10b 100644 --- a/junifer/markers/functional_connectivity_spheres.py +++ b/junifer/markers/functional_connectivity_spheres.py @@ -4,7 +4,7 @@ # Kaustubh R. Patil # 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) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index c0f5544c8..b4c64a150 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -4,7 +4,7 @@ # Synchon Mandal # 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, diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index 6c3797942..e25478c52 100644 --- a/junifer/markers/sphere_aggregation.py +++ b/junifer/markers/sphere_aggregation.py @@ -4,7 +4,7 @@ # Synchon Mandal # 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) ) diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 8ef21777d..ea3cd9f76 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -4,6 +4,8 @@ # Synchon Mandal # 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 diff --git a/junifer/markers/tests/test_ets_rss.py b/junifer/markers/tests/test_ets_rss.py index 8a38c56e8..2c7b793b9 100644 --- a/junifer/markers/tests/test_ets_rss.py +++ b/junifer/markers/tests/test_ets_rss.py @@ -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) diff --git a/junifer/markers/tests/test_functional_connectivity_atlas.py b/junifer/markers/tests/test_functional_connectivity_atlas.py index 7aa435183..4b14b17c1 100644 --- a/junifer/markers/tests/test_functional_connectivity_atlas.py +++ b/junifer/markers/tests/test_functional_connectivity_atlas.py @@ -1,4 +1,4 @@ -"""Provide test for parcel aggregation.""" +"""Provide tests for functional connectivity atlas.""" # Authors: Amir Omidvarnia # Kaustubh R. Patil @@ -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 diff --git a/junifer/markers/tests/test_functional_connectivity_spheres.py b/junifer/markers/tests/test_functional_connectivity_spheres.py index a58b53927..96a41dfd7 100644 --- a/junifer/markers/tests/test_functional_connectivity_spheres.py +++ b/junifer/markers/tests/test_functional_connectivity_spheres.py @@ -1,18 +1,17 @@ -"""Provide test for functional connectivity spheres.""" +"""Provide tests for functional connectivity spheres.""" # Authors: Amir Omidvarnia # Kaustubh R. Patil # Federico Raimondo # 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) diff --git a/junifer/markers/tests/test_markers_base.py b/junifer/markers/tests/test_markers_base.py index c86c074d0..12d0b583e 100644 --- a/junifer/markers/tests/test_markers_base.py +++ b/junifer/markers/tests/test_markers_base.py @@ -4,94 +4,57 @@ # Synchon Mandal # 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" diff --git a/junifer/markers/tests/test_sphere_aggregation.py b/junifer/markers/tests/test_sphere_aggregation.py index d7733592b..9b5d1a0f0 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -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) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index f9bda0886..16513ba7c 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -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( diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 4d0ee310d..80a3fb1cd 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -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: diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 0d3850f12..c9eb73c34 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -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, diff --git a/junifer/storage/tests/test_storage_base.py b/junifer/storage/tests/test_storage_base.py index c9af9db04..cfe9a0068 100644 --- a/junifer/storage/tests/test_storage_base.py +++ b/junifer/storage/tests/test_storage_base.py @@ -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"