fix: storage #103
20 changed files with 717 additions and 529 deletions
|
|
@ -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
|
||||
~~~~
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
|||
|
||||
|
Okay makes sense since it comes from Okay makes sense since it comes from `PipelineStepMixin`.
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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Reference in a new issue
This should still be validate.
There are two methods:
validate_input: checks the inputget_output_kind: gives the output kind (given the input)validatecalls both of themSame here, the method should be
validate