[ENH] Refactor Functional Connectivity based markers #107

Merged
synchon merged 17 commits from refactor/fc-markers into main 2023-01-10 08:55:29 +00:00
13 changed files with 241 additions and 219 deletions

View file

@ -29,6 +29,9 @@ Enhancements
- Add support for ``Dosenbach`` coordinates (:gh:`168` by `Synchon Mandal`_). - Add support for ``Dosenbach`` coordinates (:gh:`168` by `Synchon Mandal`_).
- Organize functional connectivity markers in ``junifer.markers.functional_connectivity``
(:gh:`107` by `Synchon Mandal`_).
Bugs Bugs
~~~~ ~~~~

View file

@ -2,16 +2,19 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de> # Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de> # Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from .base import BaseMarker from .base import BaseMarker
from .collection import MarkerCollection from .collection import MarkerCollection
from .crossparcellation_functional_connectivity import CrossParcellationFC
from .ets_rss import RSSETSMarker from .ets_rss import RSSETSMarker
from .functional_connectivity_parcels import FunctionalConnectivityParcels
from .functional_connectivity_spheres import FunctionalConnectivitySpheres
from .parcel_aggregation import ParcelAggregation from .parcel_aggregation import ParcelAggregation
from .sphere_aggregation import SphereAggregation from .sphere_aggregation import SphereAggregation
from .functional_connectivity import (
FunctionalConnectivityParcels,
FunctionalConnectivitySpheres,
CrossParcellationFC,
)
from .reho import ReHoParcels, ReHoSpheres from .reho import ReHoParcels, ReHoSpheres
from .falff import ( from .falff import (
AmplitudeLowFrequencyFluctuationParcels, AmplitudeLowFrequencyFluctuationParcels,

View file

@ -110,6 +110,7 @@ class AmplitudeLowFrequencyFluctuationParcels(
* ``data`` : the actual computed values as a numpy.ndarray * ``data`` : the actual computed values as a numpy.ndarray
* ``columns`` : the column labels for the computed values as a list * ``columns`` : the column labels for the computed values as a list
""" """
pa = ParcelAggregation( pa = ParcelAggregation(
parcellation=self.parcellation, parcellation=self.parcellation,

View file

@ -0,0 +1,8 @@
"""Provide imports for functional connectivity sub-package."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from .functional_connectivity_parcels import FunctionalConnectivityParcels
from .functional_connectivity_spheres import FunctionalConnectivitySpheres
from .crossparcellation_functional_connectivity import CrossParcellationFC

View file

@ -8,12 +8,11 @@ from typing import Any, Dict, List, Optional
import pandas as pd import pandas as pd
from ..api.decorators import register_marker from ...api.decorators import register_marker
from ..utils import logger from ...utils import logger, raise_error
from ..utils.logging import raise_error from ..base import BaseMarker
from .base import BaseMarker from ..parcel_aggregation import ParcelAggregation
from .parcel_aggregation import ParcelAggregation from ..utils import _correlate_dataframes
from .utils import _correlate_dataframes
@register_marker @register_marker
@ -115,10 +114,10 @@ class CrossParcellationFC(BaseMarker):
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. The dictionary has the following keys: with this as a parameter. The dictionary has the following keys:
* data : the correlation values between the two parcellations as * ``data`` : the correlation values between the two parcellations
a numpy.ndarray as a numpy.ndarray
* col_names : the ROIs for first parcellation as a list * ``col_names`` : the ROIs for first parcellation as a list
* row_names : the ROIs for second parcellation as a list * ``row_names`` : the ROIs for second parcellation as a list
""" """
logger.debug( logger.debug(

View file

@ -1,28 +1,24 @@
"""Provide class for functional connectivity.""" """Provide abstract base class for functional connectivity (FC)."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Any, Dict, List, Optional, Union
from abc import abstractmethod
from typing import Any, Dict, List, Optional
from nilearn.connectome import ConnectivityMeasure from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
from ..api.decorators import register_marker from ...utils import raise_error
from .base import BaseMarker from ..base import BaseMarker
from .parcel_aggregation import ParcelAggregation
@register_marker class FunctionalConnectivityBase(BaseMarker):
class FunctionalConnectivityParcels(BaseMarker): """Abstract base class for functional connectivity markers.
"""Class for functional connectivity.
Parameters Parameters
---------- ----------
parcellation : str or list of str
The name(s) of the parcellation(s). Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`.
agg_method : str, optional agg_method : str, optional
The method to perform aggregation using. Check valid options in The method to perform aggregation using. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name` (default "mean"). :func:`junifer.stats.get_aggfunc_by_name` (default "mean").
@ -43,13 +39,13 @@ class FunctionalConnectivityParcels(BaseMarker):
name : str, optional name : str, optional
The name of the marker. If None, will use the class name (default The name of the marker. If None, will use the class name (default
None). None).
""" """
_DEPENDENCIES = {"nilearn", "scikit-learn"} _DEPENDENCIES = {"nilearn", "scikit-learn"}
def __init__( def __init__(
self, self,
parcellation: Union[str, List[str]],
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",
@ -57,7 +53,6 @@ class FunctionalConnectivityParcels(BaseMarker):
mask: Optional[str] = None, mask: Optional[str] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
self.parcellation = parcellation
self.agg_method = agg_method self.agg_method = agg_method
self.agg_method_params = agg_method_params self.agg_method_params = agg_method_params
self.cor_method = cor_method self.cor_method = cor_method
@ -68,8 +63,15 @@ class FunctionalConnectivityParcels(BaseMarker):
"empirical", False "empirical", False
) )
self.mask = mask self.mask = mask
super().__init__(on="BOLD", name=name)
super().__init__(name=name) @abstractmethod
def aggregate(self, input: Dict[str, Any]) -> Dict[str, Any]:
"""Perform aggregation."""
raise_error(
msg="Concrete classes need to implement aggregate().",
klass=NotImplementedError,
)
def get_valid_inputs(self) -> List[str]: def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input. """Get valid data types for input.
@ -118,36 +120,30 @@ class FunctionalConnectivityParcels(BaseMarker):
Returns Returns
------- -------
dict dict
The computed result as dictionary. The following data will be The computed result as dictionary. The following keys will be
included in the dictionary: included in the dictionary:
* ``data`` : functional connectivity matrix as a numpy.ndarray. * ``data`` : functional connectivity matrix as a ``numpy.ndarray``.
* ``row_names`` : row names as a list * ``row_names`` : row names as a list
* ``col_names`` : column names as a list * ``col_names`` : column names as a list
* ``matrix_kind`` : the kind of matrix (tril, triu or full) * ``matrix_kind`` : the kind of matrix (tril, triu or full)
""" """
pa = ParcelAggregation( # Perform necessary aggregation
parcellation=self.parcellation, aggregation = self.aggregate(input)
method=self.agg_method, # Compute correlation
method_params=self.agg_method_params,
mask=self.mask,
on="BOLD",
)
# get the 2D timeseries after parcel aggregation
ts = pa.compute(input)
if self.cor_method_params["empirical"]: if self.cor_method_params["empirical"]:
cm = ConnectivityMeasure( connectivity = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), # type: ignore cov_estimator=EmpiricalCovariance(),
kind=self.cor_method, kind=self.cor_method,
) )
else: else:
cm = ConnectivityMeasure(kind=self.cor_method) connectivity = ConnectivityMeasure(kind=self.cor_method)
# Create dictionary for output
out = {} out = {}
out["data"] = cm.fit_transform([ts["data"]])[0] out["data"] = connectivity.fit_transform([aggregation["data"]])[0]
# create column names # Create column names
out["row_names"] = ts["columns"] out["row_names"] = aggregation["columns"]
out["col_names"] = ts["columns"] out["col_names"] = aggregation["columns"]
out["matrix_kind"] = "tril" out["matrix_kind"] = "tril"
return out return out

View file

@ -0,0 +1,77 @@
"""Provide class for functional connectivity using parcels."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Any, Dict, List, Optional, Union
from ...api.decorators import register_marker
from ..parcel_aggregation import ParcelAggregation
from .functional_connectivity_base import FunctionalConnectivityBase
@register_marker
class FunctionalConnectivityParcels(FunctionalConnectivityBase):
"""Class for functional connectivity using parcellations.
Parameters
----------
parcellation : str or list of str
The name(s) of the parcellation(s). Check valid options by calling
:func:`junifer.data.parcellations.list_parcellations`.
agg_method : str, optional
The method to perform aggregation using. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name` (default "mean").
agg_method_params : dict, optional
Parameters to pass to the aggregation function. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name` (default None).
cor_method : str, optional
The method to perform correlation using. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure`
(default "covariance").
cor_method_params : dict, optional
Parameters to pass to the correlation function. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
mask : str, optional
The name of the mask to apply to regions before extracting signals.
Check valid options by calling :func:`junifer.data.masks.list_masks`
(default None).
name : str, optional
The name of the marker. If None, will use the class name (default
None).
"""
def __init__(
self,
parcellation: Union[str, List[str]],
agg_method: str = "mean",
agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance",
cor_method_params: Optional[Dict] = None,
mask: Optional[str] = None,
name: Optional[str] = None,
) -> None:
self.parcellation = parcellation
super().__init__(
agg_method=agg_method,
agg_method_params=agg_method_params,
cor_method=cor_method,
cor_method_params=cor_method_params,
mask=mask,
name=name,
)
def aggregate(self, input: Dict[str, Any]) -> Dict:
"""Perform parcel aggregation."""
parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation,
method=self.agg_method,
method_params=self.agg_method_params,
mask=self.mask,
on="BOLD",
)
# Return the 2D timeseries after parcel aggregation
return parcel_aggregation.compute(input)

View file

@ -0,0 +1,87 @@
"""Provide class for functional connectivity using spheres."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Any, Dict, Optional
from ...api.decorators import register_marker
from ..sphere_aggregation import SphereAggregation
from ..utils import raise_error
from .functional_connectivity_base import FunctionalConnectivityBase
@register_marker
class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
"""Class for functional connectivity using coordinates (spheres).
Parameters
----------
coords : str
The name of the coordinates list to use. See
:func:`junifer.data.coordinates.list_coordinates` for options.
radius : float, optional
The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None).
agg_method : str, optional
The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information
(default None).
agg_method_params : dict, optional
The parameters to pass to the aggregation method (default None).
cor_method : str, optional
The method to perform correlation using. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
cor_method_params : dict, optional
Parameters to pass to the correlation function. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
mask : str, optional
The name of the mask to apply to regions before extracting signals.
Check valid options by calling :func:`junifer.data.masks.list_masks`
(default None).
name : str, optional
The name of the marker. By default, it will use
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
was applied to (default None).
"""
def __init__(
self,
coords: str,
radius: Optional[float] = None,
agg_method: str = "mean",
agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance",
cor_method_params: Optional[Dict] = None,
mask: Optional[str] = None,
name: Optional[str] = None,
) -> None:
self.coords = coords
self.radius = radius
if radius is None or radius <= 0:
raise_error(f"radius should be > 0: provided {radius}")
super().__init__(
agg_method=agg_method,
agg_method_params=agg_method_params,
cor_method=cor_method,
cor_method_params=cor_method_params,
mask=mask,
name=name,
)
def aggregate(self, input: Dict[str, Any]) -> Dict:
"""Perform sphere aggregation."""
sphere_aggregation = SphereAggregation(
coords=self.coords,
radius=self.radius,
method=self.agg_method,
method_params=self.agg_method_params,
mask=self.mask,
on="BOLD",
)
# Return the 2D timeseries after sphere aggregation
return sphere_aggregation.compute(input)

View file

@ -9,9 +9,7 @@ from pathlib import Path
import pytest import pytest
from nilearn import image from nilearn import image
from junifer.markers.crossparcellation_functional_connectivity import ( from junifer.markers.functional_connectivity import CrossParcellationFC
CrossParcellationFC,
)
from junifer.storage import SQLiteFeatureStorage from junifer.storage import SQLiteFeatureStorage
from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber

View file

@ -0,0 +1,15 @@
"""Provide tests for base functional connectivity marker."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
import pytest
# done to keep line length 79
import junifer.markers.functional_connectivity as fc
def test_base_functional_connectivity_marker_abstractness() -> None:
"""Test FunctionalConnectivityBase is an abstract base class."""
with pytest.raises(TypeError, match="abstract"):
fc.functional_connectivity_base.FunctionalConnectivityBase()

View file

@ -11,7 +11,7 @@ from nilearn.connectome import ConnectivityMeasure
from nilearn.maskers import NiftiLabelsMasker from nilearn.maskers import NiftiLabelsMasker
from numpy.testing import assert_array_almost_equal, assert_array_equal from numpy.testing import assert_array_almost_equal, assert_array_equal
from junifer.markers.functional_connectivity_parcels import ( from junifer.markers.functional_connectivity import (
FunctionalConnectivityParcels, FunctionalConnectivityParcels,
) )
from junifer.markers.parcel_aggregation import ParcelAggregation from junifer.markers.parcel_aggregation import ParcelAggregation

View file

@ -13,7 +13,7 @@ from nilearn.connectome import ConnectivityMeasure
from numpy.testing import assert_array_almost_equal from numpy.testing import assert_array_almost_equal
from sklearn.covariance import EmpiricalCovariance from sklearn.covariance import EmpiricalCovariance
from junifer.markers.functional_connectivity_spheres import ( from junifer.markers.functional_connectivity import (
FunctionalConnectivitySpheres, FunctionalConnectivitySpheres,
) )
from junifer.markers.sphere_aggregation import SphereAggregation from junifer.markers.sphere_aggregation import SphereAggregation

View file

@ -1,165 +0,0 @@
"""Provide base class for functional connectivity using spheres."""
# Authors: Amir Omidvarnia <a.omidvarnia@fz-juelich.de>
# Kaustubh R. Patil <k.patil@fz-juelich.de>
# License: AGPL
from typing import Any, Dict, List, Optional
from nilearn.connectome import ConnectivityMeasure
from sklearn.covariance import EmpiricalCovariance
from ..api.decorators import register_marker
from ..utils import raise_error
from .base import BaseMarker
from .sphere_aggregation import SphereAggregation
@register_marker
class FunctionalConnectivitySpheres(BaseMarker):
"""Class for functional connectivity using coordinates (spheres).
Parameters
----------
coords : str
The name of the coordinates list to use. See
:func:`junifer.data.coordinates.list_coordinates` for options.
radius : float, optional
The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None).
agg_method : str, optional
The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information
(default None).
agg_method_params : dict, optional
The parameters to pass to the aggregation method (default None).
cor_method : str, optional
The method to perform correlation using. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
cor_method_params : dict, optional
Parameters to pass to the correlation function. Check valid options in
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
mask : str, optional
The name of the mask to apply to regions before extracting signals.
Check valid options by calling :func:`junifer.data.masks.list_masks`
(default None).
name : str, optional
The name of the marker. By default, it will use
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
was applied to (default None).
"""
_DEPENDENCIES = {"nilearn", "scikit-learn"}
def __init__(
self,
coords: str,
radius: Optional[float] = None,
agg_method: str = "mean",
agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance",
cor_method_params: Optional[Dict] = None,
mask: Optional[str] = None,
name: Optional[str] = None,
) -> None:
self.coords = coords
self.radius = radius
if radius is None or radius <= 0:
raise_error(f"radius should be > 0: provided {radius}")
self.agg_method = agg_method
self.agg_method_params = agg_method_params
self.cor_method = cor_method
self.cor_method_params = cor_method_params or {}
# default to nilearn behavior
self.cor_method_params["empirical"] = self.cor_method_params.get(
"empirical", False
)
self.mask = mask
super().__init__(name=name)
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this marker.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The data type input to the marker.
Returns
-------
str
The storage type output by the marker.
"""
return "matrix"
def compute(
self,
input: Dict[str, Any],
extra_input: Optional[Dict] = None,
) -> Dict:
"""Compute.
Parameters
----------
input : dict
A single input from the pipeline data object in which to compute
the marker.
extra_input : dict, optional
The other fields in the pipeline data object. Useful for accessing
other data kind that needs to be used in the computation. For
example, the functional connectivity markers can make use of the
confounds if available (default None).
Returns
-------
dict
The computed result as dictionary. The following keys will be
included in the dictionary:
* ``data`` : functional connectivity matrix as a numpy.ndarray.
* ``row_names`` : row names as a list
* ``col_names`` : column names as a list
* ``matrix_kind`` : the kind of matrix (tril, triu or full)
"""
sa = SphereAggregation(
coords=self.coords,
radius=self.radius,
method=self.agg_method,
method_params=self.agg_method_params,
mask=self.mask,
on="BOLD",
)
ts = sa.compute(input)
if self.cor_method_params["empirical"]:
cm = ConnectivityMeasure(
cov_estimator=EmpiricalCovariance(), # type: ignore
kind=self.cor_method,
)
else:
cm = ConnectivityMeasure(kind=self.cor_method)
out = {}
out["data"] = cm.fit_transform([ts["data"]])[0]
# create column names
out["row_names"] = ts["columns"]
out["col_names"] = ts["columns"]
out["matrix_kind"] = "tril"
return out