fix: storage #103
20 changed files with 717 additions and 529 deletions
|
|
@ -51,13 +51,17 @@ Enhancements
|
||||||
|
|
||||||
- Implement SphereAggregation marker (by `Fede Raimondo`_).
|
- Implement SphereAggregation marker (by `Fede Raimondo`_).
|
||||||
|
|
||||||
- Implement PIOP1 and PIOP2 AOMIC datasets and refactor AOMICID1000 slightly
|
- Implement PIOP1 and PIOP2 AOMIC datasets and refactor AOMICID1000 slightly
|
||||||
(:gh: `94` by `Leonard Sasse`_)
|
(:gh: `94` by `Leonard Sasse`_)
|
||||||
|
|
||||||
- Implement a JuselessDataladCamCANVBM datagrabber class (:gh: `99` by `Leonard Sasse`_)
|
- Implement a JuselessDataladCamCANVBM datagrabber class (:gh: `99` by `Leonard Sasse`_)
|
||||||
|
|
||||||
- Implement IXI CAT output datagrabber for juseless (:gh: `48` by `Leonard Sasse`_).
|
- Implement IXI CAT output datagrabber for juseless (:gh: `48` by `Leonard Sasse`_).
|
||||||
|
|
||||||
|
- Upgrade storage interface for storage-like objects (:gh: `84` by `Synchon Mandal`_).
|
||||||
|
|
||||||
|
- Add missing type annotations (:gh: `74` by `Synchon Mandal`_).
|
||||||
|
|
||||||
Bugs
|
Bugs
|
||||||
~~~~
|
~~~~
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@
|
||||||
|
|
||||||
Extracting root sum of squares from edge-wise timeseries.
|
Extracting root sum of squares from edge-wise timeseries.
|
||||||
=========================================================
|
=========================================================
|
||||||
This example uses a RSSETSMarker to compute root sum of squares
|
This example uses a ``RSSETSMarker`` to compute root sum of squares
|
||||||
of the edge-wise timeseries using the Schaefer atlas
|
of the edge-wise timeseries using the Schaefer atlas
|
||||||
(100 rois and 200 rois, 17 Yeo networks) for a 4D nifti BOLD file.
|
(100 rois and 200 rois, 17 Yeo networks) for a 4D nifti BOLD file.
|
||||||
|
|
||||||
|
|
@ -15,6 +15,8 @@ License: BSD 3 clause
|
||||||
import tempfile
|
import tempfile
|
||||||
|
|
||||||
import junifer.testing.registry # noqa: F401
|
import junifer.testing.registry # noqa: F401
|
||||||
|
from junifer.api import collect, run
|
||||||
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.utils import configure_logging
|
from junifer.utils import configure_logging
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -22,11 +24,15 @@ from junifer.utils import configure_logging
|
||||||
# Set the logging level to info to see extra information:
|
# Set the logging level to info to see extra information:
|
||||||
configure_logging(level="INFO")
|
configure_logging(level="INFO")
|
||||||
|
|
||||||
|
##############################################################################
|
||||||
|
# Define the datagrabber interface
|
||||||
|
datagrabber = {
|
||||||
|
"kind": "SPMAuditoryTestingDatagrabber",
|
||||||
|
}
|
||||||
|
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Define the markers you want:
|
# Define the markers interface
|
||||||
|
markers = [
|
||||||
marker_dicts = [
|
|
||||||
{
|
{
|
||||||
"name": "Schaefer100x17_RSSETS",
|
"name": "Schaefer100x17_RSSETS",
|
||||||
"kind": "RSSETSMarker",
|
"kind": "RSSETSMarker",
|
||||||
|
|
@ -39,26 +45,34 @@ marker_dicts = [
|
||||||
},
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Create a temporary directory for junifer feature extraction:
|
# Create a temporary directory for junifer feature extraction:
|
||||||
# At the end you can read the extracted data into a ``pandas.DataFrame``.
|
# At the end you can read the extracted data into a ``pandas.DataFrame``.
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
|
# Define the storage interface
|
||||||
|
storage = {
|
||||||
|
"kind": "SQLiteFeatureStorage",
|
||||||
|
"uri": f"{tmpdir}/test.db",
|
||||||
|
"single_output": False,
|
||||||
|
}
|
||||||
|
# Run the defined junifer feature extraction pipeline
|
||||||
|
run(
|
||||||
|
workdir=tmpdir,
|
||||||
|
datagrabber=datagrabber,
|
||||||
|
markers=markers,
|
||||||
|
storage=storage,
|
||||||
|
elements=["sub001"], # we calculate for one subject only
|
||||||
|
)
|
||||||
|
# Collect extracted features data
|
||||||
|
collect(storage=storage)
|
||||||
|
# Create storage object to read in extracted features
|
||||||
|
db = SQLiteFeatureStorage(
|
||||||
|
uri=storage["uri"],
|
||||||
|
single_output=True, # as we ran collect, we have single output now
|
||||||
|
)
|
||||||
|
# Read extracted features
|
||||||
|
df_vbm = db.read_df(feature_name="BOLD_Schaefer100x17_RSSETS")
|
||||||
|
|
||||||
storage = {"kind": "SQLiteFeatureStorage", "uri": f"{tmpdir}/test.db"}
|
###############################################################################
|
||||||
# run the defined junifer feature extraction pipeline
|
# Now we take a look at the dataframe
|
||||||
# TODO: needs SQLiteFeatureStorage.store_timeseries() to be
|
df_vbm.head()
|
||||||
# implemented first
|
|
||||||
# run(
|
|
||||||
# workdir="/tmp",
|
|
||||||
# datagrabber={"kind": "SPMAuditoryTestingDatagrabber"},
|
|
||||||
# markers=marker_dicts,
|
|
||||||
# storage=storage,
|
|
||||||
# )
|
|
||||||
|
|
||||||
# read in extracted features and add confounds and targets
|
|
||||||
# for julearn run cross validation
|
|
||||||
# This will not run for now as store_timeseries() is not implemented yet
|
|
||||||
# collect(storage)
|
|
||||||
# db = SQLiteFeatureStorage(uri=storage["uri"], single_output=True)
|
|
||||||
# df_vbm = db.read_df(feature_name="Schaefer100x17")
|
|
||||||
|
|
|
||||||
|
|
@ -7,5 +7,7 @@
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
from .collection import MarkerCollection
|
from .collection import MarkerCollection
|
||||||
from .ets_rss import RSSETSMarker
|
from .ets_rss import RSSETSMarker
|
||||||
|
from .functional_connectivity_atlas import FunctionalConnectivityAtlas
|
||||||
|
from .functional_connectivity_spheres import FunctionalConnectivitySpheres
|
||||||
from .parcel import ParcelAggregation
|
from .parcel import ParcelAggregation
|
||||||
from .sphere_aggregation import SphereAggregation
|
from .sphere_aggregation import SphereAggregation
|
||||||
|
|
|
||||||
|
|
@ -1,23 +1,28 @@
|
||||||
"""Provide base class for markers."""
|
"""Provide abstract base class for markers."""
|
||||||
|
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List, Optional, Union
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from ..pipeline import PipelineStepMixin
|
from ..pipeline import PipelineStepMixin
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
class BaseMarker(PipelineStepMixin):
|
if TYPE_CHECKING:
|
||||||
"""Base class for all markers.
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
|
class BaseMarker(ABC, PipelineStepMixin):
|
||||||
|
"""Abstract base class for all markers.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
on : list of str
|
on : str or list of str
|
||||||
The kind of data to apply the marker to. By default, will work on all
|
The kind of data to apply the marker to. By default, will work on all
|
||||||
available data (default None).
|
available data.
|
||||||
name : str, optional
|
name : str, optional
|
||||||
The name of the marker. By default, it will use the class name as the
|
The name of the marker. By default, it will use the class name as the
|
||||||
name of the marker (default None).
|
name of the marker (default None).
|
||||||
|
|
@ -44,7 +49,7 @@ class BaseMarker(PipelineStepMixin):
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict
|
dict
|
||||||
The metadata as a dictionary.
|
The metadata as a dictionary with the only key 'marker'.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
s_meta = super().get_meta()
|
s_meta = super().get_meta()
|
||||||
|
|
@ -76,6 +81,7 @@ class BaseMarker(PipelineStepMixin):
|
||||||
f"\t Required (any of): {self._valid_inputs}"
|
f"\t Required (any of): {self._valid_inputs}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
def get_output_kind(self, input: List[str]) -> List[str]:
|
def get_output_kind(self, input: List[str]) -> List[str]:
|
||||||
"""Get output kind.
|
"""Get output kind.
|
||||||
|
|
||||||
|
|
@ -92,9 +98,11 @@ class BaseMarker(PipelineStepMixin):
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
msg="get_output_kind() not implemented", klass=NotImplementedError
|
msg="Concrete classes need to implement get_output_kind().",
|
||||||
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
|
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
|
|
@ -117,32 +125,54 @@ class BaseMarker(PipelineStepMixin):
|
||||||
with this as a parameter.
|
with this as a parameter.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(msg="compute() not implemented", klass=NotImplementedError)
|
raise_error(
|
||||||
|
msg="Concrete classes need to implement compute().",
|
||||||
|
klass=NotImplementedError,
|
||||||
|
)
|
||||||
|
|
||||||
# TODO: complete type annotations
|
@abstractmethod
|
||||||
def store(self, kind: str, out: Dict, storage) -> None:
|
def store(
|
||||||
|
self,
|
||||||
|
kind: str,
|
||||||
|
out: Dict[str, Any],
|
||||||
|
storage: "BaseFeatureStorage",
|
||||||
|
) -> None:
|
||||||
"""Store.
|
"""Store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input
|
kind : str
|
||||||
out
|
The data kind to store.
|
||||||
|
out : dict
|
||||||
|
The computed result as a dictionary to store.
|
||||||
|
storage : storage-like
|
||||||
|
The storage class.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(msg="store() not implemented", klass=NotImplementedError)
|
raise_error(
|
||||||
|
msg="Concrete classes need to implement store().",
|
||||||
|
klass=NotImplementedError,
|
||||||
|
)
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def fit_transform(
|
||||||
def fit_transform(self, input: Dict[str, Dict], storage=None) -> Dict:
|
self,
|
||||||
|
input: Dict[str, Dict],
|
||||||
|
storage: "BaseFeatureStorage" = None,
|
||||||
|
) -> Dict:
|
||||||
"""Fit and transform.
|
"""Fit and transform.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input
|
input : dict
|
||||||
storage
|
The Junifer Data object.
|
||||||
|
storage : storage-like, optional
|
||||||
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict
|
dict
|
||||||
|
The processed output as a dictionary. If `storage` is provided,
|
||||||
|
empty dictionary is returned.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
out = {}
|
out = {}
|
||||||
|
|
@ -156,11 +186,11 @@ class BaseMarker(PipelineStepMixin):
|
||||||
t_meta = meta.copy()
|
t_meta = meta.copy()
|
||||||
t_meta.update(t_input.get("meta", {}))
|
t_meta.update(t_input.get("meta", {}))
|
||||||
t_meta.update(self.get_meta(kind))
|
t_meta.update(self.get_meta(kind))
|
||||||
t_out = self.compute(t_input, extra_input)
|
t_out = self.compute(input=t_input, extra_input=extra_input)
|
||||||
t_out.update(meta=t_meta)
|
t_out.update(meta=t_meta)
|
||||||
if storage is not None:
|
if storage is not None:
|
||||||
logger.info(f"Storing in {storage}")
|
logger.info(f"Storing in {storage}")
|
||||||
self.store(kind, t_out, storage)
|
self.store(kind=kind, out=t_out, storage=storage)
|
||||||
else:
|
else:
|
||||||
logger.info("No storage specified, returning dictionary")
|
logger.info("No storage specified, returning dictionary")
|
||||||
out[kind] = t_out
|
out[kind] = t_out
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,7 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from collections import Counter
|
from collections import Counter
|
||||||
from typing import Dict, List, Optional
|
from typing import TYPE_CHECKING, Dict, List, Optional
|
||||||
|
|
||||||
from ..datareader.default import DefaultDataReader
|
from ..datareader.default import DefaultDataReader
|
||||||
from ..markers.base import BaseMarker
|
from ..markers.base import BaseMarker
|
||||||
|
|
@ -14,15 +14,23 @@ from ..storage.base import BaseFeatureStorage
|
||||||
from ..utils import logger
|
from ..utils import logger
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from junifer.datagrabber import BaseDataGrabber
|
||||||
|
|
||||||
|
|
||||||
class MarkerCollection:
|
class MarkerCollection:
|
||||||
"""Class for marker collection.
|
"""Class for marker collection.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
markers
|
markers : list of marker-like
|
||||||
datareader
|
The markers to compute.
|
||||||
preprocessing
|
datareader : datareader-like, optional
|
||||||
storage
|
The datareader to use (default None).
|
||||||
|
preprocessing : preprocessing-like, optional
|
||||||
|
The preprocessing steps to apply.
|
||||||
|
storage : storage-like, optional
|
||||||
|
The storage to use (default None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -60,17 +68,23 @@ class MarkerCollection:
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
output : dict or None
|
dict or None
|
||||||
The output of the pipeline. Each key represents a marker name and
|
The output of the pipeline. Each key represents a marker name and
|
||||||
the values are the computer marker values. If the pipeline has a
|
the values are the computer marker values. If the pipeline has a
|
||||||
storage configured, then the output will be None.
|
storage configured, then the output will be None.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.info("Fitting pipeline")
|
logger.info("Fitting pipeline")
|
||||||
|
|
||||||
|
# Fetch actual data using datareader
|
||||||
data = self._datareader.fit_transform(input)
|
data = self._datareader.fit_transform(input)
|
||||||
|
|
||||||
|
# Apply preprocessing steps
|
||||||
if self._preprocessing is not None:
|
if self._preprocessing is not None:
|
||||||
logger.info("Preprocessing data")
|
logger.info("Preprocessing data")
|
||||||
data = self._preprocessing.fit_transform(data)
|
data = self._preprocessing.fit_transform(data)
|
||||||
|
|
||||||
|
# Compute markers
|
||||||
out = {}
|
out = {}
|
||||||
for marker in self._markers:
|
for marker in self._markers:
|
||||||
logger.info(f"Fitting marker {marker.name}")
|
logger.info(f"Fitting marker {marker.name}")
|
||||||
|
|
@ -78,20 +92,21 @@ class MarkerCollection:
|
||||||
if self._storage is None:
|
if self._storage is None:
|
||||||
out[marker.name] = m_value
|
out[marker.name] = m_value
|
||||||
logger.info("Marker collection fitting done")
|
logger.info("Marker collection fitting done")
|
||||||
|
|
||||||
return None if self._storage else out
|
return None if self._storage else out
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def validate(self, datagrabber: "BaseDataGrabber") -> None:
|
||||||
def validate(self, datagrabber) -> None:
|
|
||||||
"""Validate the pipeline.
|
"""Validate the pipeline.
|
||||||
|
|
||||||
Without doing any computation, check if the Marker Collection can
|
Without doing any computation, check if the marker collection can
|
||||||
be fit without problems. That is, the data required for each marker is
|
be fit without problems. That is, the data required for each marker is
|
||||||
present and streamed down the steps. Also, if a storage is configured,
|
present and streamed down the steps. Also, if a storage is configured,
|
||||||
check that the storage can handle the markers output.
|
check that the storage can handle the markers output.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
datagrabber
|
datagrabber : datagrabber-like
|
||||||
|
The datagrabber to validate.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.info("Validating Marker Collection")
|
logger.info("Validating Marker Collection")
|
||||||
|
|
@ -104,8 +119,11 @@ class MarkerCollection:
|
||||||
|
|
||||||
|
Okay makes sense since it comes from Okay makes sense since it comes from `PipelineStepMixin`.
Done. Done.
|
|||||||
for marker in self._markers:
|
for marker in self._markers:
|
||||||
logger.info(f"Validating Marker: {marker.name}")
|
logger.info(f"Validating Marker: {marker.name}")
|
||||||
m_data = marker.validate(t_data)
|
# Validate marker
|
||||||
|
m_data = marker.validate(input=t_data)
|
||||||
logger.info(f"Marker output type: {m_data}")
|
logger.info(f"Marker output type: {m_data}")
|
||||||
|
# Check storage for the marker
|
||||||
if self._storage is not None:
|
if self._storage is not None:
|
||||||
logger.info(f"Validating storage for {marker.name}")
|
logger.info(f"Validating storage for {marker.name}")
|
||||||
self._storage.validate(m_data)
|
# Validate storage
|
||||||
|
self._storage.validate(input_=m_data)
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
|
|
@ -17,6 +17,10 @@ from .parcel import ParcelAggregation
|
||||||
from .utils import _ets
|
from .utils import _ets
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class RSSETSMarker(BaseMarker):
|
class RSSETSMarker(BaseMarker):
|
||||||
"""Class for root sum of squares of edgewise timeseries.
|
"""Class for root sum of squares of edgewise timeseries.
|
||||||
|
|
@ -24,9 +28,14 @@ class RSSETSMarker(BaseMarker):
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
atlas : str
|
atlas : str
|
||||||
The name of the atlas.
|
The name of the atlas. Check valid options by calling
|
||||||
|
:func:`junifer.data.list_atlases`.
|
||||||
aggregation_method : str, optional
|
aggregation_method : str, optional
|
||||||
The aggregation method (default "mean").
|
The method to perform aggregation using. Check valid options in
|
||||||
|
:func:`junifer.stats.get_aggfunc_by_name` (default "mean").
|
||||||
|
name : str, optional
|
||||||
|
The name of the marker. If None, will use the class name (default
|
||||||
|
None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -58,21 +67,32 @@ class RSSETSMarker(BaseMarker):
|
||||||
"""
|
"""
|
||||||
return ["timeseries"]
|
return ["timeseries"]
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store(
|
||||||
def store(self, kind: str, out, storage) -> None:
|
self,
|
||||||
|
kind: str,
|
||||||
|
out: Dict[str, Any],
|
||||||
|
storage: "BaseFeatureStorage",
|
||||||
|
) -> None:
|
||||||
"""Store.
|
"""Store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
kind
|
kind : {"BOLD"}
|
||||||
out
|
The data kind to store.
|
||||||
storage
|
out : dict
|
||||||
|
The computed result as a dictionary to store.
|
||||||
|
storage : storage-like
|
||||||
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Storing BOLD in {storage}")
|
logger.debug(f"Storing BOLD in {storage}")
|
||||||
storage.store_timeseries(**out)
|
storage.store(kind="timeseries", **out)
|
||||||
|
|
||||||
def compute(self, input: Dict) -> Dict:
|
def compute(
|
||||||
|
self,
|
||||||
|
input: Dict[str, Any],
|
||||||
|
extra_input: Optional[Dict] = None,
|
||||||
|
) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
Take a timeseries of brain areas, and calculate timeseries for each
|
Take a timeseries of brain areas, and calculate timeseries for each
|
||||||
|
|
@ -81,12 +101,19 @@ class RSSETSMarker(BaseMarker):
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : dict of the BOLD data
|
input : dict
|
||||||
|
The BOLD data as dictionary.
|
||||||
|
extra_input : dict, optional
|
||||||
|
The other fields in the pipeline data object (default None).
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict
|
dict
|
||||||
The computed result as dictionary.
|
The computed result as dictionary. The dictionary has the following
|
||||||
|
keys:
|
||||||
|
- data : the actual computed values as a numpy.ndarray
|
||||||
|
- columns : the column labels for the computed values as a list
|
||||||
|
- row_names (if more than one row is present in data): "scan"
|
||||||
|
|
||||||
References
|
References
|
||||||
----------
|
----------
|
||||||
|
|
@ -103,9 +130,10 @@ class RSSETSMarker(BaseMarker):
|
||||||
method=self.aggregation_method,
|
method=self.aggregation_method,
|
||||||
)
|
)
|
||||||
# Compute the parcel aggregation
|
# Compute the parcel aggregation
|
||||||
out = parcel_aggregation.compute(input)
|
out = parcel_aggregation.compute(input=input, extra_input=extra_input)
|
||||||
edge_ts = _ets(out["data"])
|
edge_ts = _ets(out["data"])
|
||||||
# Compute the RSS
|
# Compute the RSS
|
||||||
out["data"] = np.sum(edge_ts**2, 1) ** 0.5
|
out["data"] = np.sum(edge_ts**2, 1) ** 0.5
|
||||||
|
# Set correct column label
|
||||||
|
out["columns"] = ["root_sum_of_squares_ets"]
|
||||||
return out
|
return out
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@
|
||||||
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
from sklearn.covariance import EmpiricalCovariance
|
from sklearn.covariance import EmpiricalCovariance
|
||||||
|
|
@ -15,29 +15,45 @@ from .base import BaseMarker
|
||||||
from .parcel import ParcelAggregation
|
from .parcel import ParcelAggregation
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class FunctionalConnectivityAtlas(BaseMarker):
|
class FunctionalConnectivityAtlas(BaseMarker):
|
||||||
"""Class for functional connectivity.
|
"""Class for functional connectivity.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
atlas
|
atlas : str
|
||||||
agg_method
|
The name of the atlas. Check valid options by calling
|
||||||
agg_method_params
|
:func:`junifer.data.list_atlases`.
|
||||||
cor_method
|
agg_method : str, optional
|
||||||
cor_method_params
|
The method to perform aggregation using. Check valid options in
|
||||||
name
|
:func:`junifer.stats.get_aggfunc_by_name` (default "mean").
|
||||||
|
agg_method_params : dict, optional
|
||||||
|
Parameters to pass to the aggregation function. Check valid options in
|
||||||
|
:func:`junifer.stats.get_aggfunc_by_name` (default None).
|
||||||
|
cor_method : str, optional
|
||||||
|
The method to perform correlation using. Check valid options in
|
||||||
|
:func:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
|
||||||
|
cor_method_params : dict, optional
|
||||||
|
Parameters to pass to the correlation function. Check valid options in
|
||||||
|
:func:`nilearn.connectome.ConnectivityMeasure` (default None).
|
||||||
|
name : str, optional
|
||||||
|
The name of the marker. If None, will use the class name (default
|
||||||
|
None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
atlas,
|
atlas: str,
|
||||||
agg_method="mean",
|
agg_method: str = "mean",
|
||||||
agg_method_params=None,
|
agg_method_params: Optional[Dict] = None,
|
||||||
cor_method="covariance",
|
cor_method: str = "covariance",
|
||||||
cor_method_params=None,
|
cor_method_params: Optional[Dict] = None,
|
||||||
name=None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Initialize the class."""
|
"""Initialize the class."""
|
||||||
self.atlas = atlas
|
self.atlas = atlas
|
||||||
|
|
@ -75,15 +91,19 @@ class FunctionalConnectivityAtlas(BaseMarker):
|
||||||
outputs = ["matrix"]
|
outputs = ["matrix"]
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
|
def compute(
|
||||||
|
self,
|
||||||
|
input: Dict[str, Any],
|
||||||
|
extra_input: Optional[Dict] = None,
|
||||||
|
) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : Dict[str, Dict]
|
input : dict
|
||||||
A single input from the pipeline data object in which to compute
|
A single input from the pipeline data object in which to compute
|
||||||
the marker.
|
the marker.
|
||||||
extra_input : Dict, optional
|
extra_input : dict, optional
|
||||||
The other fields in the pipeline data object. Useful for accessing
|
The other fields in the pipeline data object. Useful for accessing
|
||||||
other data kind that needs to be used in the computation. For
|
other data kind that needs to be used in the computation. For
|
||||||
example, the functional connectivity markers can make use of the
|
example, the functional connectivity markers can make use of the
|
||||||
|
|
@ -94,10 +114,11 @@ class FunctionalConnectivityAtlas(BaseMarker):
|
||||||
dict
|
dict
|
||||||
The computed result as dictionary. The following data will be
|
The computed result as dictionary. The following data will be
|
||||||
included in the dictionary:
|
included in the dictionary:
|
||||||
- 'data': FC matrix as a 2D numpy array.
|
- data: functional connectivity matrix as a numpy.ndarray.
|
||||||
- 'row_names': Row names as a list.
|
- row_names: row names as a list
|
||||||
- 'col_names': Col names as a list.
|
- col_names: column names as a list
|
||||||
- 'kind': The kind of matrix (tril, triu or full)
|
- matrix_kind: the kind of matrix (tril, triu or full)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
pa = ParcelAggregation(
|
pa = ParcelAggregation(
|
||||||
atlas=self.atlas,
|
atlas=self.atlas,
|
||||||
|
|
@ -120,18 +141,26 @@ class FunctionalConnectivityAtlas(BaseMarker):
|
||||||
# create column names
|
# create column names
|
||||||
out["row_names"] = ts["columns"]
|
out["row_names"] = ts["columns"]
|
||||||
out["col_names"] = ts["columns"]
|
out["col_names"] = ts["columns"]
|
||||||
out["kind"] = "tril"
|
out["matrix_kind"] = "tril"
|
||||||
return out
|
return out
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store(
|
||||||
def store(self, kind: str, out: Dict, storage) -> None:
|
self,
|
||||||
|
kind: str,
|
||||||
|
out: Dict[str, Any],
|
||||||
|
storage: "BaseFeatureStorage",
|
||||||
|
) -> None:
|
||||||
"""Store.
|
"""Store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input
|
kind : {"BOLD"}
|
||||||
out
|
The data kind to store.
|
||||||
|
out : dict
|
||||||
|
The computed result as a dictionary to store.
|
||||||
|
storage : storage-like
|
||||||
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
logger.debug(f"Storing {kind} in {storage}")
|
||||||
storage.store_matrix2d(**out)
|
storage.store(kind="matrix", **out)
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@
|
||||||
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
from sklearn.covariance import EmpiricalCovariance
|
from sklearn.covariance import EmpiricalCovariance
|
||||||
|
|
@ -15,6 +15,10 @@ from .base import BaseMarker
|
||||||
from .sphere_aggregation import SphereAggregation
|
from .sphere_aggregation import SphereAggregation
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class FunctionalConnectivitySpheres(BaseMarker):
|
class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
"""Class for functional connectivity using coordinates (spheres).
|
"""Class for functional connectivity using coordinates (spheres).
|
||||||
|
|
@ -23,16 +27,23 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
----------
|
----------
|
||||||
coords : str
|
coords : str
|
||||||
The name of the coordinates list to use. See
|
The name of the coordinates list to use. See
|
||||||
:mod:`junifer.data.coordinates`
|
:mod:`junifer.data.coordinates` for options.
|
||||||
radius : float
|
radius : float, optional
|
||||||
The radius of the sphere in mm. If None, the signal will be extracted
|
The radius of the sphere in mm. If None, the signal will be extracted
|
||||||
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
|
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
|
||||||
for more information.
|
for more information (default None).
|
||||||
agg_method : str
|
agg_method : str, optional
|
||||||
The aggregation method to use.
|
The aggregation method to use.
|
||||||
See :func:`junifer.stats.get_aggfunc_by_name` for more information.
|
See :func:`junifer.stats.get_aggfunc_by_name` for more information
|
||||||
|
(default None).
|
||||||
agg_method_params : dict, optional
|
agg_method_params : dict, optional
|
||||||
The parameters to pass to the aggregation method.
|
The parameters to pass to the aggregation method (default None).
|
||||||
|
cor_method : str, optional
|
||||||
|
The method to perform correlation using. Check valid options in
|
||||||
|
:func:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
|
||||||
|
cor_method_params : dict, optional
|
||||||
|
Parameters to pass to the correlation function. Check valid options in
|
||||||
|
:func:`nilearn.connectome.ConnectivityMeasure` (default None).
|
||||||
name : str, optional
|
name : str, optional
|
||||||
The name of the marker. By default, it will use
|
The name of the marker. By default, it will use
|
||||||
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
|
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
|
||||||
|
|
@ -43,7 +54,7 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
coords: str,
|
coords: str,
|
||||||
radius: float,
|
radius: Optional[float] = None,
|
||||||
agg_method: str = "mean",
|
agg_method: str = "mean",
|
||||||
agg_method_params: Optional[Dict] = None,
|
agg_method_params: Optional[Dict] = None,
|
||||||
cor_method: str = "covariance",
|
cor_method: str = "covariance",
|
||||||
|
|
@ -88,12 +99,16 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
outputs = ["matrix"]
|
outputs = ["matrix"]
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
|
def compute(
|
||||||
|
self,
|
||||||
|
input: Dict[str, Any],
|
||||||
|
extra_input: Optional[Dict] = None,
|
||||||
|
) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : dict[str, dict]
|
input : dict
|
||||||
A single input from the pipeline data object in which to compute
|
A single input from the pipeline data object in which to compute
|
||||||
the marker.
|
the marker.
|
||||||
extra_input : dict, optional
|
extra_input : dict, optional
|
||||||
|
|
@ -107,10 +122,10 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
dict
|
dict
|
||||||
The computed result as dictionary. The following keys will be
|
The computed result as dictionary. The following keys will be
|
||||||
included in the dictionary:
|
included in the dictionary:
|
||||||
- 'data': FC matrix as a 2D numpy array.
|
- data: functional connectivity matrix as a numpy.ndarray.
|
||||||
- 'row_names': Row names as a list.
|
- row_names: row names as a list
|
||||||
- 'col_names': Col names as a list.
|
- col_names: column names as a list
|
||||||
- 'kind': The kind of matrix (tril, triu or full)
|
- matrix_kind: the kind of matrix (tril, triu or full)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
sa = SphereAggregation(
|
sa = SphereAggregation(
|
||||||
|
|
@ -135,18 +150,26 @@ class FunctionalConnectivitySpheres(BaseMarker):
|
||||||
# create column names
|
# create column names
|
||||||
out["row_names"] = ts["columns"]
|
out["row_names"] = ts["columns"]
|
||||||
out["col_names"] = ts["columns"]
|
out["col_names"] = ts["columns"]
|
||||||
out["kind"] = "tril"
|
out["matrix_kind"] = "tril"
|
||||||
return out
|
return out
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store(
|
||||||
def store(self, kind: str, out: Dict, storage) -> None:
|
self,
|
||||||
|
kind: str,
|
||||||
|
out: Dict[str, Any],
|
||||||
|
storage: "BaseFeatureStorage",
|
||||||
|
) -> None:
|
||||||
"""Store.
|
"""Store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input
|
kind : {"BOLD"}
|
||||||
out
|
The data kind to store.
|
||||||
|
out : dict
|
||||||
|
The computed result as a dictionary to store.
|
||||||
|
storage : storage-like
|
||||||
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
logger.debug(f"Storing {kind} in {storage}")
|
||||||
storage.store_matrix2d(**out)
|
storage.store(kind="matrix", **out)
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from nilearn.image import math_img, resample_to_img
|
from nilearn.image import math_img, resample_to_img
|
||||||
|
|
@ -17,22 +17,42 @@ from ..utils import logger
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class ParcelAggregation(BaseMarker):
|
class ParcelAggregation(BaseMarker):
|
||||||
"""Class for parcel aggregation.
|
"""Class for parcel aggregation.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
atlas
|
atlas : str
|
||||||
method
|
The name of the atlas. Check valid options by calling
|
||||||
method_params
|
:func:`junifer.data.list_atlases`.
|
||||||
on
|
method : str
|
||||||
name
|
The method to perform aggregation using. Check valid options in
|
||||||
|
:func:`junifer.stats.get_aggfunc_by_name`.
|
||||||
|
method_params : dict, optional
|
||||||
|
Parameters to pass to the aggregation function. Check valid options in
|
||||||
|
:func:`junifer.stats.get_aggfunc_by_name`.
|
||||||
|
on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} or list
|
||||||
|
of the options, optional
|
||||||
|
The kind of data to apply the marker to. If None, will work on all
|
||||||
|
available data (default None).
|
||||||
|
name : str, optional
|
||||||
|
The name of the marker. If None, will use the class name (default
|
||||||
|
None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, atlas, method, method_params=None, on=None, name=None
|
self,
|
||||||
|
atlas: str,
|
||||||
|
method: str,
|
||||||
|
method_params: Optional[Dict[str, Any]] = None,
|
||||||
|
on: Union[List[str], str, None] = None,
|
||||||
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Initialize the class."""
|
"""Initialize the class."""
|
||||||
self.atlas = atlas
|
self.atlas = atlas
|
||||||
|
|
@ -52,8 +72,8 @@ class ParcelAggregation(BaseMarker):
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
str
|
list of str
|
||||||
The kind of output.
|
The list of storage kinds.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
outputs = []
|
outputs = []
|
||||||
|
|
@ -66,32 +86,43 @@ class ParcelAggregation(BaseMarker):
|
||||||
raise ValueError(f"Unknown input kind for {t_input}")
|
raise ValueError(f"Unknown input kind for {t_input}")
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store(
|
||||||
def store(self, kind: str, out, storage) -> None:
|
self,
|
||||||
|
kind: str,
|
||||||
|
out: Dict[str, Any],
|
||||||
|
storage: "BaseFeatureStorage",
|
||||||
|
) -> None:
|
||||||
"""Store.
|
"""Store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
kind
|
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
|
||||||
out
|
The data kind to store.
|
||||||
storage
|
out : dict
|
||||||
|
The computed result as a dictionary to store.
|
||||||
|
storage : storage-like
|
||||||
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
logger.debug(f"Storing {kind} in {storage}")
|
||||||
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
||||||
storage.store_table(**out)
|
storage.store(kind="table", **out)
|
||||||
if kind in ["BOLD"]:
|
elif kind in ["BOLD"]:
|
||||||
storage.store_timeseries(**out)
|
storage.store(kind="timeseries", **out)
|
||||||
|
|
||||||
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
|
def compute(
|
||||||
|
self,
|
||||||
|
input: Dict[str, Any],
|
||||||
|
extra_input: Optional[Dict] = None,
|
||||||
|
) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : Dict[str, Dict]
|
input : dict
|
||||||
A single input from the pipeline data object in which to compute
|
A single input from the pipeline data object in which to compute
|
||||||
the marker.
|
the marker.
|
||||||
extra_input : Dict, optional
|
extra_input : dict, optional
|
||||||
The other fields in the pipeline data object. Useful for accessing
|
The other fields in the pipeline data object. Useful for accessing
|
||||||
other data kind that needs to be used in the computation. For
|
other data kind that needs to be used in the computation. For
|
||||||
example, the functional connectivity markers can make use of the
|
example, the functional connectivity markers can make use of the
|
||||||
|
|
@ -102,17 +133,24 @@ class ParcelAggregation(BaseMarker):
|
||||||
dict
|
dict
|
||||||
The computed result as dictionary. This will be either returned
|
The computed result as dictionary. This will be either returned
|
||||||
to the user or stored in the storage by calling the store method
|
to the user or stored in the storage by calling the store method
|
||||||
with this as a parameter.
|
with this as a parameter. The dictionary has the following keys:
|
||||||
|
- data : the actual computed values as a numpy.ndarray
|
||||||
|
- columns : the column labels for the computed values as a list
|
||||||
|
- row_names (if more than one row is present in data): "scan"
|
||||||
|
|
||||||
"""
|
"""
|
||||||
t_input = input["data"]
|
t_input = input["data"]
|
||||||
logger.debug(f"Parcel aggregation using {self.method}")
|
logger.debug(f"Parcel aggregation using {self.method}")
|
||||||
agg_func = get_aggfunc_by_name(
|
agg_func = get_aggfunc_by_name(
|
||||||
self.method, func_params=self.method_params
|
name=self.method,
|
||||||
|
func_params=self.method_params,
|
||||||
)
|
)
|
||||||
# Get the min of the voxels sizes and use it as the resolution
|
# Get the min of the voxels sizes and use it as the resolution
|
||||||
resolution = np.min(t_input.header.get_zooms()[:3]) # type: ignore
|
resolution = np.min(t_input.header.get_zooms()[:3])
|
||||||
t_atlas, t_labels, _ = load_atlas(self.atlas, resolution=resolution)
|
t_atlas, t_labels, _ = load_atlas(
|
||||||
|
name=self.atlas,
|
||||||
|
resolution=resolution,
|
||||||
|
)
|
||||||
atlas_img_res = resample_to_img(
|
atlas_img_res = resample_to_img(
|
||||||
t_atlas,
|
t_atlas,
|
||||||
t_input,
|
t_input,
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
from nilearn.maskers import NiftiSpheresMasker
|
from nilearn.maskers import NiftiSpheresMasker
|
||||||
|
|
||||||
|
|
@ -14,6 +14,10 @@ from ..utils import logger, raise_error
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from junifer.storage import BaseFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
@register_marker
|
@register_marker
|
||||||
class SphereAggregation(BaseMarker):
|
class SphereAggregation(BaseMarker):
|
||||||
"""Class for sphere aggregation.
|
"""Class for sphere aggregation.
|
||||||
|
|
@ -22,16 +26,17 @@ class SphereAggregation(BaseMarker):
|
||||||
----------
|
----------
|
||||||
coords: str
|
coords: str
|
||||||
The name of the coordinates list to use. See
|
The name of the coordinates list to use. See
|
||||||
:mod:`junifer.data.coordinates`
|
:mod:`junifer.data.coordinates` for options.
|
||||||
radius: float
|
radius: float, optional
|
||||||
The radius of the sphere in mm. If None, the signal will be extracted
|
The radius of the sphere in mm. If None, the signal will be extracted
|
||||||
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
|
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
|
||||||
for more information.
|
for more information (default None).
|
||||||
method: str
|
method: str, optional
|
||||||
The aggregation method to use.
|
The aggregation method to use.
|
||||||
See :func:`junifer.stats.get_aggfunc_by_name` for more information.
|
See :func:`junifer.stats.get_aggfunc_by_name` for more information
|
||||||
method_params: Dict, optional
|
(default "mean").
|
||||||
The parameters to pass to the aggregation method.
|
method_params: dict, optional
|
||||||
|
The parameters to pass to the aggregation method (default None).
|
||||||
on: list of str, optional
|
on: list of str, optional
|
||||||
The kind of data to apply the marker to. By default, will work on all
|
The kind of data to apply the marker to. By default, will work on all
|
||||||
available data (default None).
|
available data (default None).
|
||||||
|
|
@ -44,8 +49,8 @@ class SphereAggregation(BaseMarker):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
coords: str,
|
coords: str,
|
||||||
radius: float,
|
radius: Optional[float] = None,
|
||||||
method: str,
|
method: str = "mean",
|
||||||
method_params: Optional[Dict] = None,
|
method_params: Optional[Dict] = None,
|
||||||
on: Optional[List[str]] = None,
|
on: Optional[List[str]] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
|
|
@ -92,31 +97,43 @@ class SphereAggregation(BaseMarker):
|
||||||
raise ValueError(f"Unknown input kind for {t_input}")
|
raise ValueError(f"Unknown input kind for {t_input}")
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def store(self, kind: str, out, storage) -> None:
|
def store(
|
||||||
|
self,
|
||||||
|
kind: str,
|
||||||
|
out: Dict[str, Any],
|
||||||
|
storage: "BaseFeatureStorage",
|
||||||
|
) -> None:
|
||||||
"""Store.
|
"""Store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
kind
|
kind : {"BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"}
|
||||||
out
|
The data kind to store.
|
||||||
storage
|
out : dict
|
||||||
|
The computed result as a dictionary to store.
|
||||||
|
storage : storage-like
|
||||||
|
The storage class, for example, SQLiteFeatureStorage.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Storing {kind} in {storage}")
|
logger.debug(f"Storing {kind} in {storage}")
|
||||||
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]:
|
||||||
storage.store_table(**out)
|
storage.store(kind="table", **out)
|
||||||
elif kind in ["BOLD"]:
|
elif kind in ["BOLD"]:
|
||||||
storage.store_timeseries(**out)
|
storage.store(kind="timeseries", **out)
|
||||||
|
|
||||||
def compute(self, input: Dict, extra_input: Optional[Dict] = None) -> Dict:
|
def compute(
|
||||||
|
self,
|
||||||
|
input: Dict[str, Any],
|
||||||
|
extra_input: Optional[Dict] = None,
|
||||||
|
) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : Dict[str, Dict]
|
input : dict
|
||||||
A single input from the pipeline data object in which to compute
|
A single input from the pipeline data object in which to compute
|
||||||
the marker.
|
the marker.
|
||||||
extra_input : Dict, optional
|
extra_input : dict, optional
|
||||||
The other fields in the pipeline data object. Useful for accessing
|
The other fields in the pipeline data object. Useful for accessing
|
||||||
other data kind that needs to be used in the computation. For
|
other data kind that needs to be used in the computation. For
|
||||||
example, the functional connectivity markers can make use of the
|
example, the functional connectivity markers can make use of the
|
||||||
|
|
@ -125,7 +142,12 @@ class SphereAggregation(BaseMarker):
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
dict
|
dict
|
||||||
The computed result as dictionary.
|
The computed result as dictionary. This will be either returned
|
||||||
|
to the user or stored in the storage by calling the store method
|
||||||
|
with this as a parameter. The dictionary has the following keys:
|
||||||
|
- data : the actual computed values as a numpy.ndarray
|
||||||
|
- columns : the column labels for the computed values as a list
|
||||||
|
- row_names (if more than one row is present in data): "scan"
|
||||||
|
|
||||||
"""
|
"""
|
||||||
t_input = input["data"]
|
t_input = input["data"]
|
||||||
|
|
@ -135,8 +157,8 @@ class SphereAggregation(BaseMarker):
|
||||||
# )
|
# )
|
||||||
coords, out_labels = load_coordinates(self.coords)
|
coords, out_labels = load_coordinates(self.coords)
|
||||||
masker = NiftiSpheresMasker(
|
masker = NiftiSpheresMasker(
|
||||||
coords,
|
seeds=coords,
|
||||||
self.radius,
|
radius=self.radius,
|
||||||
mask_img=None, # TODO: support this (needs #79)
|
mask_img=None, # TODO: support this (needs #79)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from numpy.testing import assert_array_equal
|
from numpy.testing import assert_array_equal
|
||||||
|
|
||||||
|
|
@ -94,7 +96,7 @@ def test_marker_collection():
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_MarkerCollection_storage(tmp_path) -> None:
|
def test_marker_collection_storage(tmp_path: Path) -> None:
|
||||||
"""Test marker collection with storage.
|
"""Test marker collection with storage.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -123,7 +125,9 @@ def test_MarkerCollection_storage(tmp_path) -> None:
|
||||||
uri = tmp_path / "test_marker_collection_storage.db"
|
uri = tmp_path / "test_marker_collection_storage.db"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
||||||
mc = MarkerCollection(
|
mc = MarkerCollection(
|
||||||
markers=markers, storage=storage, datareader=DefaultDataReader()
|
markers=markers,
|
||||||
|
storage=storage,
|
||||||
|
datareader=DefaultDataReader(),
|
||||||
)
|
)
|
||||||
mc.validate(dg)
|
mc.validate(dg)
|
||||||
assert mc._storage is not None
|
assert mc._storage is not None
|
||||||
|
|
|
||||||
|
|
@ -13,27 +13,36 @@ from nilearn.maskers import NiftiLabelsMasker
|
||||||
|
|
||||||
from junifer.data import load_atlas
|
from junifer.data import load_atlas
|
||||||
from junifer.markers.ets_rss import RSSETSMarker
|
from junifer.markers.ets_rss import RSSETSMarker
|
||||||
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
|
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber
|
||||||
|
|
||||||
|
|
||||||
def test_compute() -> None:
|
# Set atlas
|
||||||
"""Test RSS ETS."""
|
ATLAS = "Schaefer100x17"
|
||||||
atlas = "Schaefer100x17"
|
|
||||||
test_atlas, _, _ = load_atlas(atlas)
|
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute() -> None:
|
||||||
|
"""Test RSS ETS compute()."""
|
||||||
with SPMAuditoryTestingDatagrabber() as dg:
|
with SPMAuditoryTestingDatagrabber() as dg:
|
||||||
|
# Fetch element
|
||||||
out = dg["sub001"]
|
out = dg["sub001"]
|
||||||
|
# Load BOLD image
|
||||||
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
|
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
|
||||||
|
# Create input data
|
||||||
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
|
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
|
||||||
# Compute the RSSETSMarker
|
# Compute the RSSETSMarker
|
||||||
ets_rss_marker = RSSETSMarker(atlas=atlas)
|
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
|
||||||
new_out = ets_rss_marker.compute(input_dict)
|
new_out = ets_rss_marker.compute(input_dict)
|
||||||
|
|
||||||
|
# Load atlas
|
||||||
|
test_atlas, _, _ = load_atlas(ATLAS)
|
||||||
# Compute the NiftiLabelsMasker
|
# Compute the NiftiLabelsMasker
|
||||||
test_masker = NiftiLabelsMasker(test_atlas)
|
test_masker = NiftiLabelsMasker(test_atlas)
|
||||||
test_ts = test_masker.fit_transform(niimg)
|
test_ts = test_masker.fit_transform(niimg)
|
||||||
# Assert the dimension of timeseries
|
# Assert the dimension of timeseries
|
||||||
n_time, _ = test_ts.shape
|
n_time, _ = test_ts.shape
|
||||||
assert n_time == len(new_out["data"])
|
assert n_time == len(new_out["data"])
|
||||||
|
|
||||||
# Assert the meta
|
# Assert the meta
|
||||||
meta = ets_rss_marker.get_meta("BOLD")["marker"]
|
meta = ets_rss_marker.get_meta("BOLD")["marker"]
|
||||||
assert meta["atlas"] == "Schaefer100x17"
|
assert meta["atlas"] == "Schaefer100x17"
|
||||||
|
|
@ -42,10 +51,8 @@ def test_compute() -> None:
|
||||||
|
|
||||||
|
|
||||||
def test_get_output_kind() -> None:
|
def test_get_output_kind() -> None:
|
||||||
"""Test get_output_kind."""
|
"""Test RSS ETS get_output_kind()."""
|
||||||
|
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
|
||||||
atlas = "Schaefer100x17"
|
|
||||||
ets_rss_marker = RSSETSMarker(atlas=atlas)
|
|
||||||
input_list = ["BOLD"]
|
input_list = ["BOLD"]
|
||||||
input_list = ets_rss_marker.get_output_kind(input_list)
|
input_list = ets_rss_marker.get_output_kind(input_list)
|
||||||
assert len(input_list) == 1
|
assert len(input_list) == 1
|
||||||
|
|
@ -53,19 +60,26 @@ def test_get_output_kind() -> None:
|
||||||
|
|
||||||
|
|
||||||
def test_store(tmp_path: Path) -> None:
|
def test_store(tmp_path: Path) -> None:
|
||||||
"""Test store."""
|
"""Test RSS ETS store().
|
||||||
|
|
||||||
atlas = "Schaefer100x17"
|
Parameters
|
||||||
|
----------
|
||||||
|
tmp_path : pathlib.Path
|
||||||
|
The path to the test directory.
|
||||||
|
|
||||||
|
"""
|
||||||
with SPMAuditoryTestingDatagrabber() as dg:
|
with SPMAuditoryTestingDatagrabber() as dg:
|
||||||
|
# Fetch element
|
||||||
out = dg["sub001"]
|
out = dg["sub001"]
|
||||||
|
# Load BOLD image
|
||||||
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
|
niimg = image.load_img(str(out["BOLD"]["path"].absolute()))
|
||||||
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
|
input_dict = {"data": niimg, "path": out["BOLD"]["path"]}
|
||||||
# Compute the RSSETSMarker
|
# Compute the RSSETSMarker
|
||||||
ets_rss_marker = RSSETSMarker(atlas=atlas)
|
ets_rss_marker = RSSETSMarker(atlas=ATLAS)
|
||||||
_ = ets_rss_marker.compute(input_dict)
|
# Create storage
|
||||||
# TODO: Needs store_timeseries implemented for SQLiteFeatureStorage
|
storage = SQLiteFeatureStorage(
|
||||||
# storage = {
|
uri=str((tmp_path / "test.db").absolute()),
|
||||||
# "kind": "SQLiteFeatureStorage",
|
single_output=True,
|
||||||
# "uri": str((tmp_path / "test.db").absolute()),
|
)
|
||||||
# }
|
# Store
|
||||||
# ets_rss_marker.store("SQLiteFeatureStorage", new_out, storage)
|
ets_rss_marker.fit_transform(input=input_dict, storage=storage)
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
"""Provide test for parcel aggregation."""
|
"""Provide tests for functional connectivity atlas."""
|
||||||
|
|
||||||
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
|
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
|
||||||
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
||||||
|
|
@ -19,8 +19,14 @@ from junifer.storage import SQLiteFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
def test_FunctionalConnectivityAtlas(tmp_path: Path) -> None:
|
def test_FunctionalConnectivityAtlas(tmp_path: Path) -> None:
|
||||||
"""Test FunctionalConnectivityAtlas."""
|
"""Test FunctionalConnectivityAtlas.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
tmp_path : pathlib.Path
|
||||||
|
The path to the test directory.
|
||||||
|
|
||||||
|
"""
|
||||||
# get a dataset
|
# get a dataset
|
||||||
ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
|
ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
|
||||||
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
|
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
|
||||||
|
|
|
||||||
|
|
@ -1,18 +1,17 @@
|
||||||
"""Provide test for functional connectivity spheres."""
|
"""Provide tests for functional connectivity spheres."""
|
||||||
|
|
||||||
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
|
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
|
||||||
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
# Kaustubh R. Patil <k.patil@fz-juelich.de>
|
||||||
# Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from numpy.testing import assert_array_almost_equal
|
|
||||||
|
|
||||||
from sklearn.covariance import EmpiricalCovariance
|
import pytest
|
||||||
from nilearn import datasets, image
|
from nilearn import datasets, image
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
|
from numpy.testing import assert_array_almost_equal
|
||||||
|
from sklearn.covariance import EmpiricalCovariance
|
||||||
|
|
||||||
from junifer.markers.functional_connectivity_spheres import (
|
from junifer.markers.functional_connectivity_spheres import (
|
||||||
FunctionalConnectivitySpheres,
|
FunctionalConnectivitySpheres,
|
||||||
|
|
@ -30,7 +29,6 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
||||||
The path to the test directory.
|
The path to the test directory.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# get a dataset
|
# get a dataset
|
||||||
ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
|
ni_data = datasets.fetch_spm_auditory(subject_id="sub001")
|
||||||
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
|
fmri_img = image.concat_imgs(ni_data.func) # type: ignore
|
||||||
|
|
@ -118,8 +116,7 @@ def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
|
||||||
|
|
||||||
# Check that FC are almost equal when using nileran
|
# Check that FC are almost equal when using nileran
|
||||||
cm = ConnectivityMeasure(
|
cm = ConnectivityMeasure(
|
||||||
cov_estimator=EmpiricalCovariance(), # type: ignore
|
cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
|
||||||
kind="correlation"
|
|
||||||
)
|
)
|
||||||
out_ni = cm.fit_transform([ts["data"]])[0]
|
out_ni = cm.fit_transform([ts["data"]])[0]
|
||||||
assert_array_almost_equal(out_ni, out["data"], decimal=3)
|
assert_array_almost_equal(out_ni, out["data"], decimal=3)
|
||||||
|
|
|
||||||
|
|
@ -4,94 +4,57 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from junifer.markers.base import BaseMarker
|
from junifer.markers.base import BaseMarker
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
def test_base_marker_abstractness() -> None:
|
||||||
"on, name, kind, expected_class, expected_name",
|
"""Test BaseMarker is abstract base class."""
|
||||||
[
|
with pytest.raises(TypeError, match=r"abstract"):
|
||||||
(["bold", "dwi"], None, "bold", "BaseMarker", "bold_BaseMarker"),
|
BaseMarker(on=["BOLD"])
|
||||||
(["bold", "dwi"], "mymarker", "dwi", "BaseMarker", "dwi_mymarker"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_base_marker_meta(
|
|
||||||
on: List[str],
|
|
||||||
name: Optional[str],
|
|
||||||
kind: str,
|
|
||||||
expected_class: str,
|
|
||||||
expected_name: str,
|
|
||||||
) -> None:
|
|
||||||
"""Test metadata for BaseMarker.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
on : list of str
|
|
||||||
The parametrized kind of data to work on.
|
|
||||||
name : str or None
|
|
||||||
The parametrized name of the marker.
|
|
||||||
kind : str
|
|
||||||
The parametrized kind of data to get metadata for.
|
|
||||||
expected_class : str
|
|
||||||
The paramtrized expected class of the marker.
|
|
||||||
expected_name : str
|
|
||||||
The parametrized expected name of the marker.
|
|
||||||
|
|
||||||
"""
|
|
||||||
base = BaseMarker(on=on, name=name)
|
|
||||||
t_meta = base.get_meta(kind=kind)
|
|
||||||
assert t_meta["marker"]["class"] == expected_class
|
|
||||||
assert t_meta["marker"]["name"] == expected_name
|
|
||||||
|
|
||||||
|
|
||||||
def test_compute_parameters() -> None:
|
def test_base_datagrabber_subclassing() -> None:
|
||||||
"""Test compute parameters."""
|
"""Test proper subclassing of BaseMarker."""
|
||||||
base = BaseMarker(on=["bold", "dwi"], name="mymarker")
|
# Create concrete class
|
||||||
base.compute = lambda x, y: { # type: ignore
|
class MyBaseMarker(BaseMarker):
|
||||||
"data": x.keys(),
|
def get_output_kind(self, input):
|
||||||
"extra": y.keys(),
|
return ["timeseries"]
|
||||||
|
|
||||||
|
def compute(self, input, extra_input):
|
||||||
|
return {
|
||||||
|
"data": "data",
|
||||||
|
"columns": "columns",
|
||||||
|
"row_names": "row_names",
|
||||||
|
}
|
||||||
|
|
||||||
|
def store(self, kind, out, storage):
|
||||||
|
return super().store(kind=kind, out=out, storage=storage)
|
||||||
|
|
||||||
|
# Create input for marker
|
||||||
|
input_ = {
|
||||||
|
"meta": {"datagrabber": "dg", "element": "elem", "datareader": "dr"},
|
||||||
|
"BOLD": {
|
||||||
|
"path": ".",
|
||||||
|
"data": "data",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}}
|
marker = MyBaseMarker(on=["BOLD"])
|
||||||
out = base.fit_transform(
|
output = marker.fit_transform(input=input_) # process
|
||||||
input_,
|
# Check output
|
||||||
)
|
assert "BOLD" in output
|
||||||
assert list(out["bold"]["data"]) == ["path"]
|
assert "data" in output["BOLD"]
|
||||||
assert list(out["bold"]["extra"]) == ["t2"]
|
assert "columns" in output["BOLD"]
|
||||||
|
assert "row_names" in output["BOLD"]
|
||||||
|
assert "meta" in output["BOLD"]
|
||||||
def test_BaseMarker() -> None:
|
assert "datagrabber" in output["BOLD"]["meta"]
|
||||||
"""Test base class."""
|
assert "element" in output["BOLD"]["meta"]
|
||||||
base = BaseMarker(on=["bold", "dwi"], name="mymarker")
|
assert "datareader" in output["BOLD"]["meta"]
|
||||||
input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}}
|
|
||||||
base.validate_input(list(input_.keys()))
|
|
||||||
|
|
||||||
wrong_input = {"t2": {"path": "test"}}
|
|
||||||
with pytest.raises(ValueError):
|
|
||||||
base.validate_input(list(wrong_input.keys()))
|
|
||||||
|
|
||||||
|
# Check no implementation check
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
base.get_output_kind(list(wrong_input.keys()))
|
marker.store(kind="kind", out="out", storage="storage")
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
# Check attributes
|
||||||
base.fit_transform(input_)
|
assert marker.name == "MyBaseMarker"
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
base.store("bold", {}, None)
|
|
||||||
|
|
||||||
base.compute = lambda x, y: {"data": 1} # type: ignore
|
|
||||||
|
|
||||||
out = base.fit_transform(input_)
|
|
||||||
assert out["bold"]["data"] == 1
|
|
||||||
assert out["bold"]["meta"]["marker"]["name"] == "bold_mymarker"
|
|
||||||
assert out["bold"]["meta"]["marker"]["class"] == "BaseMarker"
|
|
||||||
assert "dwi" not in out
|
|
||||||
|
|
||||||
base2 = BaseMarker(on="bold", name="mymarker")
|
|
||||||
base2.compute = lambda x, y: {"data": 1} # type: ignore
|
|
||||||
out2 = base2.fit_transform(input_)
|
|
||||||
assert out2["bold"]["data"] == 1
|
|
||||||
assert out2["bold"]["meta"]["marker"]["name"] == "bold_mymarker"
|
|
||||||
assert out2["bold"]["meta"]["marker"]["class"] == "BaseMarker"
|
|
||||||
|
|
|
||||||
|
|
@ -138,19 +138,17 @@ def test_SphereAggregation_storage(tmp_path: Path) -> None:
|
||||||
|
|
||||||
marker.fit_transform(input, storage=storage)
|
marker.fit_transform(input, storage=storage)
|
||||||
|
|
||||||
# TODO: Needs store_timeseries implemented for SQLiteFeatureStorage
|
meta = {
|
||||||
|
"element": "test",
|
||||||
|
"version": "0.0.1",
|
||||||
|
"marker": {"name": "BOLD_fcname"},
|
||||||
|
}
|
||||||
|
# Get the SPM auditory data:
|
||||||
|
subject_data = datasets.fetch_spm_auditory()
|
||||||
|
fmri_img = concat_imgs(subject_data.func) # type: ignore
|
||||||
|
input = {"BOLD": {"data": fmri_img}, "meta": meta}
|
||||||
|
marker = SphereAggregation(
|
||||||
|
coords="DMNBuckner", method="mean", radius=8, on="BOLD"
|
||||||
|
)
|
||||||
|
|
||||||
# meta = {
|
marker.fit_transform(input, storage=storage)
|
||||||
# "element": "test",
|
|
||||||
# "version": "0.0.1",
|
|
||||||
# "marker": {"name": "BOLD_fcname"},
|
|
||||||
# }
|
|
||||||
# # Get the SPM auditory data:
|
|
||||||
# subject_data = datasets.fetch_spm_auditory()
|
|
||||||
# fmri_img = concat_imgs(subject_data.func) # type: ignore
|
|
||||||
# input = {"BOLD": {"data": fmri_img}, "meta": meta}
|
|
||||||
# marker = SphereAggregation(
|
|
||||||
# coords="DMNBuckner", method="mean", radius=8, on="BOLD"
|
|
||||||
# )
|
|
||||||
|
|
||||||
# marker.fit_transform(input, storage=storage)
|
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, Iterable, List, Optional, Union
|
from typing import Dict, List, Optional, Union
|
||||||
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
|
|
@ -24,16 +24,24 @@ class BaseFeatureStorage(ABC):
|
||||||
----------
|
----------
|
||||||
uri : str or pathlib.Path
|
uri : str or pathlib.Path
|
||||||
The path to the storage.
|
The path to the storage.
|
||||||
|
storage_types : str or list of str
|
||||||
|
The available storage types for the class.
|
||||||
single_output : bool, optional
|
single_output : bool, optional
|
||||||
Whether to have single output (default False).
|
Whether to have single output (default False).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, uri: Union[str, Path], single_output: bool = False
|
self,
|
||||||
|
uri: Union[str, Path],
|
||||||
|
storage_types: Union[List[str], str],
|
||||||
|
single_output: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Initialize the class."""
|
"""Initialize the class."""
|
||||||
self.uri = uri
|
self.uri = uri
|
||||||
|
if not isinstance(storage_types, list):
|
||||||
|
storage_types = [storage_types]
|
||||||
|
self._valid_inputs = storage_types
|
||||||
self.single_output = single_output
|
self.single_output = single_output
|
||||||
|
|
||||||
def get_meta(self) -> Dict:
|
def get_meta(self) -> Dict:
|
||||||
|
|
@ -51,31 +59,26 @@ class BaseFeatureStorage(ABC):
|
||||||
}
|
}
|
||||||
return meta
|
return meta
|
||||||
|
|
||||||
# TODO: is raising ValueError required?
|
def validate(self, input_: List[str]) -> None:
|
||||||
@abstractmethod
|
|
||||||
def validate(self, input_: List[str]) -> bool:
|
|
||||||
"""Validate the input to the pipeline step.
|
"""Validate the input to the pipeline step.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input_ : list
|
input_ : list of str
|
||||||
The input to the pipeline step.
|
The input to the pipeline step.
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
bool
|
|
||||||
Whether the `input` is valid or not.
|
|
||||||
|
|
||||||
Raises
|
Raises
|
||||||
------
|
------
|
||||||
ValueError
|
ValueError
|
||||||
If the input does not have the required data.
|
If the `input_` is invalid.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
if not any(x in input_ for x in self._valid_inputs):
|
||||||
msg="Concrete classes need to implement validate_input().",
|
raise_error(
|
||||||
klass=NotImplementedError,
|
"Input does not have the required data."
|
||||||
)
|
f"\t Input: {input}"
|
||||||
|
f"\t Required (any of): {self._valid_inputs}"
|
||||||
|
)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def list_features(
|
def list_features(
|
||||||
|
|
@ -108,7 +111,7 @@ class BaseFeatureStorage(ABC):
|
||||||
feature_name: Optional[str] = None,
|
feature_name: Optional[str] = None,
|
||||||
feature_md5: Optional[bool] = None,
|
feature_md5: Optional[bool] = None,
|
||||||
) -> pd.DataFrame:
|
) -> pd.DataFrame:
|
||||||
"""Read feature from the storage.
|
"""Read feature into a pandas DataFrame.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
|
@ -140,7 +143,7 @@ class BaseFeatureStorage(ABC):
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
str
|
str
|
||||||
The metadata column
|
The metadata column.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
|
|
@ -148,83 +151,38 @@ class BaseFeatureStorage(ABC):
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store(self, kind: str, **kwargs) -> None:
|
||||||
@abstractmethod
|
"""Store extracted features data.
|
||||||
def store_matrix2d(
|
|
||||||
self,
|
|
||||||
data,
|
|
||||||
meta: Dict,
|
|
||||||
col_names: Optional[Iterable[str]] = None,
|
|
||||||
row_names: Optional[Iterable[str]] = None,
|
|
||||||
kind: Optional[str] = "full",
|
|
||||||
diagonal: bool = True,
|
|
||||||
) -> None:
|
|
||||||
"""Store 2D matrix.
|
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
data
|
kind : {"matrix", "timeseries", "table"}
|
||||||
meta : dict
|
The storage kind.
|
||||||
The metadata as a dictionary.
|
**kwargs
|
||||||
col_names : list or tuple of str, optional
|
The keyword arguments.
|
||||||
The column names (default None).
|
|
||||||
row_names : list of tuple of str, optional
|
Raises
|
||||||
The row names (default None).
|
------
|
||||||
kind : str, optional
|
ValueError
|
||||||
The kind of matrix:
|
If `kind` is invalid.
|
||||||
- 'triu': store upper triangular only.
|
|
||||||
- 'tril': store lower triangular.
|
|
||||||
- 'full': full matrix (default 'full').
|
|
||||||
diagonal : bool, optional
|
|
||||||
Whether to store the diagonal (default True).
|
|
||||||
If kind == 'full', setting this to false will raise
|
|
||||||
an error
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
if kind == "matrix":
|
||||||
msg="Concrete classes need to implement store_matrix2d().",
|
self.store_matrix(**kwargs)
|
||||||
klass=NotImplementedError,
|
elif kind == "timeseries":
|
||||||
)
|
self.store_timeseries(**kwargs)
|
||||||
|
elif kind == "table":
|
||||||
|
self.store_table(**kwargs)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"I don't know how to store {kind}")
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store_df(self, **kwargs) -> None:
|
||||||
@abstractmethod
|
"""Store pandas DataFrame.
|
||||||
def store_table(
|
|
||||||
self,
|
|
||||||
data,
|
|
||||||
meta: Dict,
|
|
||||||
columns: Optional[Iterable[str]] = None,
|
|
||||||
rows_col_name: Optional[str] = None,
|
|
||||||
) -> None:
|
|
||||||
"""Store table.
|
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
data
|
**kwargs : dict
|
||||||
meta : dict
|
The keyword arguments.
|
||||||
The metadata as a dictionary.
|
|
||||||
columns : list or tuple of str, optional
|
|
||||||
The columns (default None).
|
|
||||||
rows_col_name : str, optional
|
|
||||||
The column name to use in case number of rows greater than 1.
|
|
||||||
If None and number of rows greater than 1, then the name will be
|
|
||||||
"index" (default None).
|
|
||||||
|
|
||||||
"""
|
|
||||||
raise_error(
|
|
||||||
msg="Concrete classes need to implement store_table().",
|
|
||||||
klass=NotImplementedError,
|
|
||||||
)
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def store_df(self, df: pd.DataFrame, meta: Dict) -> None:
|
|
||||||
"""Store pandas DataFerame.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
df : pandas.DataFrame
|
|
||||||
The DataFrame to store.
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
|
|
@ -232,16 +190,41 @@ class BaseFeatureStorage(ABC):
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
)
|
)
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store_matrix(self, **kwargs) -> None:
|
||||||
@abstractmethod
|
"""Store matrix.
|
||||||
def store_timeseries(self, data, meta: Dict) -> None:
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
**kwargs : dict
|
||||||
|
The keyword arguments.
|
||||||
|
|
||||||
|
"""
|
||||||
|
raise_error(
|
||||||
|
msg="Concrete classes need to implement store_matrix2d().",
|
||||||
|
klass=NotImplementedError,
|
||||||
|
)
|
||||||
|
|
||||||
|
def store_table(self, **kwargs) -> None:
|
||||||
|
"""Store table.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
**kwargs : dict
|
||||||
|
The keyword arguments.
|
||||||
|
|
||||||
|
"""
|
||||||
|
raise_error(
|
||||||
|
msg="Concrete classes need to implement store_table().",
|
||||||
|
klass=NotImplementedError,
|
||||||
|
)
|
||||||
|
|
||||||
|
def store_timeseries(self, **kwargs) -> None:
|
||||||
"""Store timeseries.
|
"""Store timeseries.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
data
|
**kwargs : dict
|
||||||
meta : dict
|
The keyword arguments.
|
||||||
The metadata as a dictionary.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
raise_error(
|
||||||
|
|
|
||||||
|
|
@ -59,6 +59,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
**kwargs: str,
|
**kwargs: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Initialize the class."""
|
"""Initialize the class."""
|
||||||
|
# Check upsert argument value
|
||||||
if upsert not in ["update", "ignore"]:
|
if upsert not in ["update", "ignore"]:
|
||||||
raise_error(
|
raise_error(
|
||||||
msg=(
|
msg=(
|
||||||
|
|
@ -76,9 +77,16 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
"does not exist, creating now."
|
"does not exist, creating now."
|
||||||
)
|
)
|
||||||
uri.parent.mkdir(parents=True, exist_ok=True)
|
uri.parent.mkdir(parents=True, exist_ok=True)
|
||||||
super().__init__(uri=uri, single_output=single_output, **kwargs)
|
# Available storage kinds
|
||||||
|
storage_types = ["table", "timeseries", "matrix"]
|
||||||
|
super().__init__(
|
||||||
|
uri=uri,
|
||||||
|
storage_types=storage_types,
|
||||||
|
single_output=single_output,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
# Set upsert
|
||||||
self._upsert = upsert
|
self._upsert = upsert
|
||||||
self._valid_inputs = ["table", "timeseries", "matrix"]
|
|
||||||
|
|
||||||
def get_engine(self, meta: Optional[Dict] = None) -> "Engine":
|
def get_engine(self, meta: Optional[Dict] = None) -> "Engine":
|
||||||
"""Get engine.
|
"""Get engine.
|
||||||
|
|
@ -109,9 +117,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
)
|
)
|
||||||
prefix = element_to_prefix(element)
|
prefix = element_to_prefix(element)
|
||||||
# Format URI for engine creation
|
# Format URI for engine creation
|
||||||
uri = (
|
uri = f"sqlite:///{self.uri.parent}/{prefix}{self.uri.name}"
|
||||||
"sqlite:///" f"{self.uri.parent}/{prefix}{self.uri.name}"
|
|
||||||
) # type: ignore
|
|
||||||
return create_engine(uri, echo=False)
|
return create_engine(uri, echo=False)
|
||||||
|
|
||||||
def _save_upsert(
|
def _save_upsert(
|
||||||
|
|
@ -203,10 +209,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
msg=f"Invalid option {if_exists} for if_exists."
|
msg=f"Invalid option {if_exists} for if_exists."
|
||||||
)
|
)
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def _store_2d(
|
||||||
def store_2d(
|
|
||||||
self,
|
self,
|
||||||
data,
|
data: Dict,
|
||||||
meta: Dict,
|
meta: Dict,
|
||||||
columns: Optional[Iterable[str]] = None,
|
columns: Optional[Iterable[str]] = None,
|
||||||
rows_col_name: Optional[str] = None,
|
rows_col_name: Optional[str] = None,
|
||||||
|
|
@ -215,7 +220,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
data
|
data : dict
|
||||||
|
The data to store.
|
||||||
meta : dict
|
meta : dict
|
||||||
The metadata as a dictionary.
|
The metadata as a dictionary.
|
||||||
columns : list or tuple of str, optional
|
columns : list or tuple of str, optional
|
||||||
|
|
@ -232,32 +238,10 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
meta=meta, n_rows=n_rows, rows_col_name=rows_col_name
|
meta=meta, n_rows=n_rows, rows_col_name=rows_col_name
|
||||||
)
|
)
|
||||||
# Prepare new dataframe
|
# Prepare new dataframe
|
||||||
data_df = pd.DataFrame(
|
data_df = pd.DataFrame(data, columns=columns, index=idx)
|
||||||
data, columns=columns, index=idx
|
|
||||||
) # type: ignore
|
|
||||||
# Store dataframe
|
# Store dataframe
|
||||||
self.store_df(df=data_df, meta=meta)
|
self.store_df(df=data_df, meta=meta)
|
||||||
|
|
||||||
def validate(self, input_: List[str]) -> bool:
|
|
||||||
"""Implement input validation.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
input_ : list of str
|
|
||||||
The input to the pipeline step.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
bool
|
|
||||||
Whether the `input` is valid or not.
|
|
||||||
|
|
||||||
"""
|
|
||||||
# Convert input to list
|
|
||||||
if not isinstance(input_, list):
|
|
||||||
input_ = [input_]
|
|
||||||
|
|
||||||
return all(x in self._valid_inputs for x in input_)
|
|
||||||
|
|
||||||
def list_features(
|
def list_features(
|
||||||
self, return_df: bool = False
|
self, return_df: bool = False
|
||||||
) -> Union[Dict[str, Dict], pd.DataFrame]:
|
) -> Union[Dict[str, Dict], pd.DataFrame]:
|
||||||
|
|
@ -293,7 +277,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
feature_name: Optional[str] = None,
|
feature_name: Optional[str] = None,
|
||||||
feature_md5: Optional[str] = None,
|
feature_md5: Optional[str] = None,
|
||||||
) -> pd.DataFrame:
|
) -> pd.DataFrame:
|
||||||
"""Implement feature reading from the storage.
|
"""Implement feature reading into a pandas DataFrame.
|
||||||
|
|
||||||
Either one of `feature_name` or `feature_md5` needs to be specified.
|
Either one of `feature_name` or `feature_md5` needs to be specified.
|
||||||
|
|
||||||
|
|
@ -399,130 +383,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
self._save_upsert(meta_df, "meta", engine)
|
self._save_upsert(meta_df, "meta", engine)
|
||||||
return f"meta_{meta_md5}"
|
return f"meta_{meta_md5}"
|
||||||
|
|
||||||
# TODO: complete type annotations
|
|
||||||
def store_matrix2d(
|
|
||||||
self,
|
|
||||||
data,
|
|
||||||
meta: Dict,
|
|
||||||
col_names: Optional[List[str]] = None,
|
|
||||||
row_names: Optional[List[str]] = None,
|
|
||||||
kind: Optional[str] = "full",
|
|
||||||
diagonal: bool = True,
|
|
||||||
) -> None:
|
|
||||||
"""Implement 2D matrix storing.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
data
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
col_names : list or tuple of str, optional
|
|
||||||
The column names (default None).
|
|
||||||
row_names : str, optional
|
|
||||||
The column name to use in case number of rows greater than 1.
|
|
||||||
If None and number of rows greater than 1, then the name will be
|
|
||||||
"index" (default None).
|
|
||||||
kind : str, optional
|
|
||||||
The kind of matrix:
|
|
||||||
- 'triu: store upper triangular only.
|
|
||||||
- 'tril': store lower triangular.
|
|
||||||
- 'full': full matrix (default 'full').
|
|
||||||
diagonal : bool, optional
|
|
||||||
Whether to store the diagonal (default True).
|
|
||||||
If kind == 'full', setting this to false will raise
|
|
||||||
an error
|
|
||||||
"""
|
|
||||||
if diagonal is False and kind not in ["triu", "tril"]:
|
|
||||||
raise_error(
|
|
||||||
msg="Diagonal cannot be False if kind is not full",
|
|
||||||
klass=ValueError,
|
|
||||||
)
|
|
||||||
|
|
||||||
if kind in ["triu", "tril"]:
|
|
||||||
if data.shape[0] != data.shape[1]:
|
|
||||||
raise_error(
|
|
||||||
"Cannot store a non-square matrix as a triangular matrix",
|
|
||||||
klass=ValueError,
|
|
||||||
)
|
|
||||||
|
|
||||||
if kind == "triu":
|
|
||||||
k = 0 if diagonal is True else 1
|
|
||||||
data_idx = np.triu_indices(data.shape[0], k=k)
|
|
||||||
elif kind == "tril":
|
|
||||||
k = 0 if diagonal is True else -1
|
|
||||||
data_idx = np.tril_indices(data.shape[0], k=k)
|
|
||||||
elif kind == "full":
|
|
||||||
data_idx = (
|
|
||||||
np.repeat(np.arange(data.shape[0]), data.shape[1]),
|
|
||||||
np.tile(np.arange(data.shape[1]), data.shape[0]),
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
raise_error(msg=f"Invalid kind {kind}", klass=ValueError)
|
|
||||||
if row_names is None:
|
|
||||||
row_names = [f"r{i}" for i in range(data.shape[0])]
|
|
||||||
elif len(row_names) != data.shape[0]:
|
|
||||||
raise_error(
|
|
||||||
msg="Number of row names does not match number of rows",
|
|
||||||
klass=ValueError,
|
|
||||||
)
|
|
||||||
|
|
||||||
if col_names is None:
|
|
||||||
col_names = [f"c{i}" for i in range(data.shape[1])]
|
|
||||||
elif len(col_names) != data.shape[1]:
|
|
||||||
raise_error(
|
|
||||||
msg="Number of column names does not match number of columns",
|
|
||||||
klass=ValueError,
|
|
||||||
)
|
|
||||||
|
|
||||||
flat_data = data[data_idx]
|
|
||||||
columns = [
|
|
||||||
f"{row_names[i]}~{col_names[j]}"
|
|
||||||
for i, j in zip(data_idx[0], data_idx[1])
|
|
||||||
]
|
|
||||||
|
|
||||||
# Convert element metadata to index
|
|
||||||
n_rows = 1
|
|
||||||
idx = element_to_index(meta=meta, n_rows=n_rows, rows_col_name=None)
|
|
||||||
# Prepare new dataframe
|
|
||||||
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
|
|
||||||
|
|
||||||
if len(columns) > 2000: # TODO: check SQLITE_MAX_COLUMN
|
|
||||||
data_df = data_df.stack()
|
|
||||||
new_names = [x for x in data_df.index.names[:-1]]
|
|
||||||
new_names.append("pair")
|
|
||||||
data_df.index.names = new_names
|
|
||||||
# Store dataframe
|
|
||||||
self.store_df(df=data_df, meta=meta) # type: ignore
|
|
||||||
|
|
||||||
# TODO: complete type annotations
|
|
||||||
def store_table(
|
|
||||||
self,
|
|
||||||
data,
|
|
||||||
meta: Dict,
|
|
||||||
columns: Optional[Iterable[str]] = None,
|
|
||||||
rows_col_name: Optional[str] = None,
|
|
||||||
) -> None:
|
|
||||||
"""Implement table storing.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
data
|
|
||||||
meta : dict
|
|
||||||
The metadata as a dictionary.
|
|
||||||
columns : list or tuple of str, optional
|
|
||||||
The columns (default None).
|
|
||||||
rows_col_name : str, optional
|
|
||||||
The column name to use in case number of rows greater than 1.
|
|
||||||
If None and number of rows greater than 1, then the name will be
|
|
||||||
"index" (default None).
|
|
||||||
|
|
||||||
"""
|
|
||||||
self.store_2d(
|
|
||||||
data=data, meta=meta, columns=columns, rows_col_name=rows_col_name
|
|
||||||
)
|
|
||||||
|
|
||||||
def store_df(self, df: pd.DataFrame, meta: Dict) -> None:
|
def store_df(self, df: pd.DataFrame, meta: Dict) -> None:
|
||||||
"""Implement dataframe storing.
|
"""Implement pandas DataFrame storing.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
|
@ -568,20 +430,158 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage):
|
||||||
# Save data
|
# Save data
|
||||||
self._save_upsert(df, table_name, engine)
|
self._save_upsert(df, table_name, engine)
|
||||||
|
|
||||||
# TODO: complete type annotations
|
def store_matrix(
|
||||||
def store_timeseries(self, data, meta: Dict) -> None:
|
self,
|
||||||
|
data: Dict,
|
||||||
|
meta: Dict,
|
||||||
|
col_names: Optional[List[str]] = None,
|
||||||
|
row_names: Optional[List[str]] = None,
|
||||||
|
matrix_kind: Optional[str] = "full",
|
||||||
|
diagonal: bool = True,
|
||||||
|
) -> None:
|
||||||
|
"""Implement matrix storing.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
data : dict
|
||||||
|
The matrix data to store.
|
||||||
|
meta : dict
|
||||||
|
The metadata as a dictionary.
|
||||||
|
col_names : list or tuple of str, optional
|
||||||
|
The column names (default None).
|
||||||
|
row_names : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1.
|
||||||
|
If None and number of rows greater than 1, then the name will be
|
||||||
|
"index" (default None).
|
||||||
|
matrix_kind : str, optional
|
||||||
|
The kind of matrix:
|
||||||
|
- "triu" : store upper triangular only
|
||||||
|
- "tril" : store lower triangular
|
||||||
|
- "full" : full matrix
|
||||||
|
(default "full").
|
||||||
|
diagonal : bool, optional
|
||||||
|
Whether to store the diagonal. If `matrix_kind` is "full", setting
|
||||||
|
this to False will raise an error (default True)..
|
||||||
|
|
||||||
|
"""
|
||||||
|
if diagonal is False and matrix_kind not in ["triu", "tril"]:
|
||||||
|
raise_error(
|
||||||
|
msg="Diagonal cannot be False if kind is not full",
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
|
||||||
|
if matrix_kind in ["triu", "tril"]:
|
||||||
|
if data.shape[0] != data.shape[1]:
|
||||||
|
raise_error(
|
||||||
|
"Cannot store a non-square matrix as a triangular matrix",
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
|
||||||
|
if matrix_kind == "triu":
|
||||||
|
k = 0 if diagonal is True else 1
|
||||||
|
data_idx = np.triu_indices(data.shape[0], k=k)
|
||||||
|
elif matrix_kind == "tril":
|
||||||
|
k = 0 if diagonal is True else -1
|
||||||
|
data_idx = np.tril_indices(data.shape[0], k=k)
|
||||||
|
elif matrix_kind == "full":
|
||||||
|
data_idx = (
|
||||||
|
np.repeat(np.arange(data.shape[0]), data.shape[1]),
|
||||||
|
np.tile(np.arange(data.shape[1]), data.shape[0]),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise_error(msg=f"Invalid kind {matrix_kind}", klass=ValueError)
|
||||||
|
|
||||||
|
if row_names is None:
|
||||||
|
row_names = [f"r{i}" for i in range(data.shape[0])]
|
||||||
|
elif len(row_names) != data.shape[0]:
|
||||||
|
raise_error(
|
||||||
|
msg="Number of row names does not match number of rows",
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
|
||||||
|
if col_names is None:
|
||||||
|
col_names = [f"c{i}" for i in range(data.shape[1])]
|
||||||
|
elif len(col_names) != data.shape[1]:
|
||||||
|
raise_error(
|
||||||
|
msg="Number of column names does not match number of columns",
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
|
||||||
|
flat_data = data[data_idx]
|
||||||
|
columns = [
|
||||||
|
f"{row_names[i]}~{col_names[j]}"
|
||||||
|
for i, j in zip(data_idx[0], data_idx[1])
|
||||||
|
]
|
||||||
|
|
||||||
|
# Convert element metadata to index
|
||||||
|
n_rows = 1
|
||||||
|
idx = element_to_index(meta=meta, n_rows=n_rows, rows_col_name=None)
|
||||||
|
# Prepare new dataframe
|
||||||
|
data_df = pd.DataFrame(flat_data[None, :], columns=columns, index=idx)
|
||||||
|
|
||||||
|
if len(columns) > 2000: # TODO: check SQLITE_MAX_COLUMN
|
||||||
|
data_df = data_df.stack()
|
||||||
|
new_names = [x for x in data_df.index.names[:-1]]
|
||||||
|
new_names.append("pair")
|
||||||
|
data_df.index.names = new_names
|
||||||
|
|
||||||
|
# Store dataframe
|
||||||
|
self.store_df(df=data_df, meta=meta)
|
||||||
|
|
||||||
|
def store_table(
|
||||||
|
self,
|
||||||
|
data: Dict,
|
||||||
|
meta: Dict,
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
rows_col_name: Optional[str] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Implement table storing.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
data : dict
|
||||||
|
The table data to store.
|
||||||
|
meta : dict
|
||||||
|
The metadata as a dictionary.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The columns (default None).
|
||||||
|
rows_col_name : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1.
|
||||||
|
If None and number of rows greater than 1, then the name will be
|
||||||
|
"index" (default None).
|
||||||
|
|
||||||
|
"""
|
||||||
|
self._store_2d(
|
||||||
|
data=data, meta=meta, columns=columns, rows_col_name=rows_col_name
|
||||||
|
)
|
||||||
|
|
||||||
|
def store_timeseries(
|
||||||
|
self,
|
||||||
|
data: Dict,
|
||||||
|
meta: Dict,
|
||||||
|
columns: Optional[Iterable[str]] = None,
|
||||||
|
row_names: str = "timepoint",
|
||||||
|
) -> None:
|
||||||
"""Implement timeseries storing.
|
"""Implement timeseries storing.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
data
|
data: dict
|
||||||
|
The timeseries data to store.
|
||||||
meta : dict
|
meta : dict
|
||||||
The metadata as a dictionary.
|
The metadata as a dictionary.
|
||||||
|
columns : list or tuple of str, optional
|
||||||
|
The column labels (default None).
|
||||||
|
row_names : str, optional
|
||||||
|
The column name to use in case number of rows greater than 1
|
||||||
|
(default "timepoint").
|
||||||
|
|
||||||
"""
|
"""
|
||||||
raise_error(
|
self._store_2d(
|
||||||
msg="store_timeseries() not implemented.",
|
data=data,
|
||||||
klass=NotImplementedError,
|
meta=meta,
|
||||||
|
columns=columns,
|
||||||
|
rows_col_name="timepoint", # explicit so as to stop overriding
|
||||||
)
|
)
|
||||||
|
|
||||||
def collect(self) -> None:
|
def collect(self) -> None:
|
||||||
|
|
|
||||||
|
|
@ -406,8 +406,8 @@ def test_store_table(tmp_path: Path) -> None:
|
||||||
assert_frame_equal(df_new, c_df_new)
|
assert_frame_equal(df_new, c_df_new)
|
||||||
|
|
||||||
|
|
||||||
def test_store_matrix2d(tmp_path: Path) -> None:
|
def test_store_matrix(tmp_path: Path) -> None:
|
||||||
"""Test 2D Matrix store.
|
"""Test matrix store.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
|
@ -428,8 +428,11 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
col_names = ["col1", "col2", "col3"]
|
col_names = ["col1", "col2", "col3"]
|
||||||
|
|
||||||
# Store table
|
# Store table
|
||||||
storage.store_matrix2d(
|
storage.store_matrix(
|
||||||
data, meta, row_names=row_names, col_names=col_names
|
data=data,
|
||||||
|
meta=meta,
|
||||||
|
row_names=row_names,
|
||||||
|
col_names=col_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [f"{i}~{j}" for i in row_names for j in col_names]
|
stored_names = [f"{i}~{j}" for i in row_names for j in col_names]
|
||||||
|
|
@ -445,7 +448,7 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
# Store without row and column names
|
# Store without row and column names
|
||||||
uri = tmp_path / "test_store_table_nonames.db"
|
uri = tmp_path / "test_store_table_nonames.db"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
||||||
storage.store_matrix2d(data, meta)
|
storage.store_matrix(data=data, meta=meta)
|
||||||
stored_names = [
|
stored_names = [
|
||||||
f"r{i}~c{j}"
|
f"r{i}~c{j}"
|
||||||
for i in range(data.shape[0])
|
for i in range(data.shape[0])
|
||||||
|
|
@ -458,13 +461,18 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
assert list(read_df.columns) == stored_names
|
assert list(read_df.columns) == stored_names
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="Invalid kind"):
|
with pytest.raises(ValueError, match="Invalid kind"):
|
||||||
storage.store_matrix2d(data, meta, kind="wrong")
|
storage.store_matrix(data=data, meta=meta, matrix_kind="wrong")
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="non-square"):
|
with pytest.raises(ValueError, match="non-square"):
|
||||||
storage.store_matrix2d(data, meta, kind="triu")
|
storage.store_matrix(data=data, meta=meta, matrix_kind="triu")
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="cannot be False"):
|
with pytest.raises(ValueError, match="cannot be False"):
|
||||||
storage.store_matrix2d(data, meta, kind="full", diagonal=False)
|
storage.store_matrix(
|
||||||
|
data=data,
|
||||||
|
meta=meta,
|
||||||
|
matrix_kind="full",
|
||||||
|
diagonal=False,
|
||||||
|
)
|
||||||
|
|
||||||
# Store upper triangular matrix
|
# Store upper triangular matrix
|
||||||
data = np.array([[1, 2, 3], [11, 22, 33], [111, 222, 333]])
|
data = np.array([[1, 2, 3], [11, 22, 33], [111, 222, 333]])
|
||||||
|
|
@ -472,8 +480,12 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
col_names = ["col1", "col2", "col3"]
|
col_names = ["col1", "col2", "col3"]
|
||||||
uri = tmp_path / "test_store_table_triu.db"
|
uri = tmp_path / "test_store_table_triu.db"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
||||||
storage.store_matrix2d(
|
storage.store_matrix(
|
||||||
data, meta, kind="triu", row_names=row_names, col_names=col_names
|
data=data,
|
||||||
|
meta=meta,
|
||||||
|
matrix_kind="triu",
|
||||||
|
row_names=row_names,
|
||||||
|
col_names=col_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [
|
stored_names = [
|
||||||
|
|
@ -497,10 +509,10 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
# Store upper triangular matrix without diagonal
|
# Store upper triangular matrix without diagonal
|
||||||
uri = tmp_path / "test_store_table_triu_nodiagonal.db"
|
uri = tmp_path / "test_store_table_triu_nodiagonal.db"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
||||||
storage.store_matrix2d(
|
storage.store_matrix(
|
||||||
data,
|
data=data,
|
||||||
meta,
|
meta=meta,
|
||||||
kind="triu",
|
matrix_kind="triu",
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
diagonal=False,
|
diagonal=False,
|
||||||
|
|
@ -527,8 +539,12 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
col_names = ["col1", "col2", "col3"]
|
col_names = ["col1", "col2", "col3"]
|
||||||
uri = tmp_path / "test_store_table_tril.db"
|
uri = tmp_path / "test_store_table_tril.db"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
||||||
storage.store_matrix2d(
|
storage.store_matrix(
|
||||||
data, meta, kind="tril", row_names=row_names, col_names=col_names
|
data=data,
|
||||||
|
meta=meta,
|
||||||
|
matrix_kind="tril",
|
||||||
|
row_names=row_names,
|
||||||
|
col_names=col_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_names = [
|
stored_names = [
|
||||||
|
|
@ -552,10 +568,10 @@ def test_store_matrix2d(tmp_path: Path) -> None:
|
||||||
# Store lower triangular matrix without diagonal
|
# Store lower triangular matrix without diagonal
|
||||||
uri = tmp_path / "test_store_table_tril_nodiagonal.db"
|
uri = tmp_path / "test_store_table_tril_nodiagonal.db"
|
||||||
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
storage = SQLiteFeatureStorage(uri=uri, single_output=True)
|
||||||
storage.store_matrix2d(
|
storage.store_matrix(
|
||||||
data,
|
data,
|
||||||
meta,
|
meta,
|
||||||
kind="tril",
|
matrix_kind="tril",
|
||||||
row_names=row_names,
|
row_names=row_names,
|
||||||
col_names=col_names,
|
col_names=col_names,
|
||||||
diagonal=False,
|
diagonal=False,
|
||||||
|
|
|
||||||
|
|
@ -12,51 +12,50 @@ from junifer.storage.base import BaseFeatureStorage
|
||||||
def test_BaseFeatureStorage_abstractness() -> None:
|
def test_BaseFeatureStorage_abstractness() -> None:
|
||||||
"""Test BaseFeatureStorage is abstract base class."""
|
"""Test BaseFeatureStorage is abstract base class."""
|
||||||
with pytest.raises(TypeError, match=r"abstract"):
|
with pytest.raises(TypeError, match=r"abstract"):
|
||||||
BaseFeatureStorage(uri="/tmp") # type: ignore
|
BaseFeatureStorage(uri="/tmp", storage_types=["matrix"])
|
||||||
|
|
||||||
|
|
||||||
def test_BaseFeatureStorage() -> None:
|
def test_BaseFeatureStorage() -> None:
|
||||||
"""Test BaseFeatureStorage."""
|
"""Test proper subclassing of BaseFeatureStorage."""
|
||||||
# Create concrete class
|
# Create concrete class
|
||||||
class MyFeatureStorage(BaseFeatureStorage):
|
class MyFeatureStorage(BaseFeatureStorage):
|
||||||
def __init__(self, uri, single_output=False):
|
"""Implement concrete class."""
|
||||||
super().__init__(uri, single_output=single_output)
|
|
||||||
|
|
||||||
def validate(self, input):
|
def __init__(self, uri, single_output=False):
|
||||||
super().validate(input)
|
storage_types = ["matrix"]
|
||||||
|
super().__init__(
|
||||||
|
uri=uri,
|
||||||
|
storage_types=storage_types,
|
||||||
|
single_output=single_output,
|
||||||
|
)
|
||||||
|
|
||||||
def list_features(self):
|
def list_features(self):
|
||||||
super().list_features()
|
super().list_features()
|
||||||
|
|
||||||
def read_df(self, feature_name=None, feature_md5=None):
|
def read_df(self, feature_name=None, feature_md5=None):
|
||||||
super().read_df(feature_name=feature_name, feature_md5=feature_md5)
|
super().read_df(
|
||||||
|
feature_name=feature_name,
|
||||||
|
feature_md5=feature_md5,
|
||||||
|
)
|
||||||
|
|
||||||
def store_metadata(self, metadata):
|
def store_metadata(self, metadata):
|
||||||
super().store_metadata(metadata)
|
super().store_metadata(metadata)
|
||||||
|
|
||||||
def store_matrix2d(self, matrix, meta):
|
|
||||||
super().store_matrix2d(matrix, meta)
|
|
||||||
|
|
||||||
def store_table(self, table, meta):
|
|
||||||
super().store_table(table, meta)
|
|
||||||
|
|
||||||
def store_df(self, df, meta):
|
|
||||||
super().store_df(df, meta)
|
|
||||||
|
|
||||||
def store_timeseries(self, timeseries, meta):
|
|
||||||
super().store_timeseries(timeseries, meta)
|
|
||||||
|
|
||||||
def collect(self):
|
def collect(self):
|
||||||
return super().collect()
|
return super().collect()
|
||||||
|
|
||||||
|
# Check single_output is False
|
||||||
st = MyFeatureStorage(uri="/tmp")
|
st = MyFeatureStorage(uri="/tmp")
|
||||||
assert st.single_output is False
|
assert st.single_output is False
|
||||||
|
# Check single_output is True
|
||||||
st = MyFeatureStorage(uri="/tmp", single_output=True)
|
st = MyFeatureStorage(uri="/tmp", single_output=True)
|
||||||
assert st.single_output is True
|
assert st.single_output is True
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
# Check validate with valid argument
|
||||||
st.validate(None)
|
st.validate(input_=["matrix"])
|
||||||
|
# Check validate with invalid argument
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
st.validate(input_=["table"])
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.list_features()
|
st.list_features()
|
||||||
|
|
@ -67,19 +66,19 @@ def test_BaseFeatureStorage() -> None:
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.store_metadata(None)
|
st.store_metadata(None)
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
st.store_matrix2d(None, None)
|
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
st.store_table(None, None)
|
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
st.store_df(None, None) # type: ignore
|
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
st.store_timeseries(None, None)
|
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
with pytest.raises(NotImplementedError):
|
||||||
st.collect()
|
st.collect()
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError):
|
||||||
|
st.store(kind="matrix")
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError):
|
||||||
|
st.store(kind="timeseries")
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError):
|
||||||
|
st.store(kind="table")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
st.store(kind="lego")
|
||||||
|
|
||||||
assert st.uri == "/tmp"
|
assert st.uri == "/tmp"
|
||||||
|
|
|
||||||
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