diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index d7d5c9a64..0fa9aa7d2 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -29,6 +29,9 @@ Enhancements - 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 ~~~~ diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index bc069e3db..de37c5a26 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -2,16 +2,19 @@ # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL from .base import BaseMarker from .collection import MarkerCollection -from .crossparcellation_functional_connectivity import CrossParcellationFC from .ets_rss import RSSETSMarker -from .functional_connectivity_parcels import FunctionalConnectivityParcels -from .functional_connectivity_spheres import FunctionalConnectivitySpheres from .parcel_aggregation import ParcelAggregation from .sphere_aggregation import SphereAggregation +from .functional_connectivity import ( + FunctionalConnectivityParcels, + FunctionalConnectivitySpheres, + CrossParcellationFC, +) from .reho import ReHoParcels, ReHoSpheres from .falff import ( AmplitudeLowFrequencyFluctuationParcels, diff --git a/junifer/markers/falff/falff_parcels.py b/junifer/markers/falff/falff_parcels.py index 0dd3080b8..4015540f8 100644 --- a/junifer/markers/falff/falff_parcels.py +++ b/junifer/markers/falff/falff_parcels.py @@ -110,6 +110,7 @@ class AmplitudeLowFrequencyFluctuationParcels( * ``data`` : the actual computed values as a numpy.ndarray * ``columns`` : the column labels for the computed values as a list + """ pa = ParcelAggregation( parcellation=self.parcellation, diff --git a/junifer/markers/functional_connectivity/__init__.py b/junifer/markers/functional_connectivity/__init__.py new file mode 100644 index 000000000..8763a661d --- /dev/null +++ b/junifer/markers/functional_connectivity/__init__.py @@ -0,0 +1,8 @@ +"""Provide imports for functional connectivity sub-package.""" + +# Authors: Synchon Mandal +# License: AGPL + +from .functional_connectivity_parcels import FunctionalConnectivityParcels +from .functional_connectivity_spheres import FunctionalConnectivitySpheres +from .crossparcellation_functional_connectivity import CrossParcellationFC diff --git a/junifer/markers/crossparcellation_functional_connectivity.py b/junifer/markers/functional_connectivity/crossparcellation_functional_connectivity.py similarity index 90% rename from junifer/markers/crossparcellation_functional_connectivity.py rename to junifer/markers/functional_connectivity/crossparcellation_functional_connectivity.py index e57e9180e..1f1ca429c 100644 --- a/junifer/markers/crossparcellation_functional_connectivity.py +++ b/junifer/markers/functional_connectivity/crossparcellation_functional_connectivity.py @@ -8,12 +8,11 @@ from typing import Any, Dict, List, Optional import pandas as pd -from ..api.decorators import register_marker -from ..utils import logger -from ..utils.logging import raise_error -from .base import BaseMarker -from .parcel_aggregation import ParcelAggregation -from .utils import _correlate_dataframes +from ...api.decorators import register_marker +from ...utils import logger, raise_error +from ..base import BaseMarker +from ..parcel_aggregation import ParcelAggregation +from ..utils import _correlate_dataframes @register_marker @@ -115,10 +114,10 @@ class CrossParcellationFC(BaseMarker): 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 correlation values between the two parcellations as - a numpy.ndarray - * col_names : the ROIs for first parcellation as a list - * row_names : the ROIs for second parcellation as a list + * ``data`` : the correlation values between the two parcellations + as a numpy.ndarray + * ``col_names`` : the ROIs for first parcellation as a list + * ``row_names`` : the ROIs for second parcellation as a list """ logger.debug( diff --git a/junifer/markers/functional_connectivity_parcels.py b/junifer/markers/functional_connectivity/functional_connectivity_base.py similarity index 70% rename from junifer/markers/functional_connectivity_parcels.py rename to junifer/markers/functional_connectivity/functional_connectivity_base.py index 649038e6d..94ffa2538 100644 --- a/junifer/markers/functional_connectivity_parcels.py +++ b/junifer/markers/functional_connectivity/functional_connectivity_base.py @@ -1,28 +1,24 @@ -"""Provide class for functional connectivity.""" +"""Provide abstract base class for functional connectivity (FC).""" -# Authors: Amir Omidvarnia -# Kaustubh R. Patil +# Authors: Synchon Mandal # 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 sklearn.covariance import EmpiricalCovariance -from ..api.decorators import register_marker -from .base import BaseMarker -from .parcel_aggregation import ParcelAggregation +from ...utils import raise_error +from ..base import BaseMarker -@register_marker -class FunctionalConnectivityParcels(BaseMarker): - """Class for functional connectivity. +class FunctionalConnectivityBase(BaseMarker): + """Abstract base class for functional connectivity markers. 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"). @@ -43,13 +39,13 @@ class FunctionalConnectivityParcels(BaseMarker): name : str, optional The name of the marker. If None, will use the class name (default None). + """ _DEPENDENCIES = {"nilearn", "scikit-learn"} def __init__( self, - parcellation: Union[str, List[str]], agg_method: str = "mean", agg_method_params: Optional[Dict] = None, cor_method: str = "covariance", @@ -57,7 +53,6 @@ class FunctionalConnectivityParcels(BaseMarker): mask: Optional[str] = None, name: Optional[str] = None, ) -> None: - self.parcellation = parcellation self.agg_method = agg_method self.agg_method_params = agg_method_params self.cor_method = cor_method @@ -68,8 +63,15 @@ class FunctionalConnectivityParcels(BaseMarker): "empirical", False ) 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]: """Get valid data types for input. @@ -118,36 +120,30 @@ class FunctionalConnectivityParcels(BaseMarker): Returns ------- 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: - * ``data`` : functional connectivity matrix as a numpy.ndarray. + * ``data`` : functional connectivity matrix as a ``numpy.ndarray``. * ``row_names`` : row names as a list * ``col_names`` : column names as a list * ``matrix_kind`` : the kind of matrix (tril, triu or full) """ - pa = ParcelAggregation( - parcellation=self.parcellation, - method=self.agg_method, - method_params=self.agg_method_params, - mask=self.mask, - on="BOLD", - ) - # get the 2D timeseries after parcel aggregation - ts = pa.compute(input) - + # Perform necessary aggregation + aggregation = self.aggregate(input) + # Compute correlation if self.cor_method_params["empirical"]: - cm = ConnectivityMeasure( - cov_estimator=EmpiricalCovariance(), # type: ignore + connectivity = ConnectivityMeasure( + cov_estimator=EmpiricalCovariance(), kind=self.cor_method, ) else: - cm = ConnectivityMeasure(kind=self.cor_method) + connectivity = ConnectivityMeasure(kind=self.cor_method) + # Create dictionary for output out = {} - out["data"] = cm.fit_transform([ts["data"]])[0] - # create column names - out["row_names"] = ts["columns"] - out["col_names"] = ts["columns"] + out["data"] = connectivity.fit_transform([aggregation["data"]])[0] + # Create column names + out["row_names"] = aggregation["columns"] + out["col_names"] = aggregation["columns"] out["matrix_kind"] = "tril" return out diff --git a/junifer/markers/functional_connectivity/functional_connectivity_parcels.py b/junifer/markers/functional_connectivity/functional_connectivity_parcels.py new file mode 100644 index 000000000..ef0147eaa --- /dev/null +++ b/junifer/markers/functional_connectivity/functional_connectivity_parcels.py @@ -0,0 +1,77 @@ +"""Provide class for functional connectivity using parcels.""" + +# Authors: Amir Omidvarnia +# Kaustubh R. Patil +# Synchon Mandal +# 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) diff --git a/junifer/markers/functional_connectivity/functional_connectivity_spheres.py b/junifer/markers/functional_connectivity/functional_connectivity_spheres.py new file mode 100644 index 000000000..9155df655 --- /dev/null +++ b/junifer/markers/functional_connectivity/functional_connectivity_spheres.py @@ -0,0 +1,87 @@ +"""Provide class for functional connectivity using spheres.""" + +# Authors: Amir Omidvarnia +# Kaustubh R. Patil +# Synchon Mandal +# 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) diff --git a/junifer/markers/tests/test_crossparcellation_functional_connectivity.py b/junifer/markers/functional_connectivity/tests/test_crossparcellation_functional_connectivity.py similarity index 96% rename from junifer/markers/tests/test_crossparcellation_functional_connectivity.py rename to junifer/markers/functional_connectivity/tests/test_crossparcellation_functional_connectivity.py index 485a4900a..608897482 100644 --- a/junifer/markers/tests/test_crossparcellation_functional_connectivity.py +++ b/junifer/markers/functional_connectivity/tests/test_crossparcellation_functional_connectivity.py @@ -9,9 +9,7 @@ from pathlib import Path import pytest from nilearn import image -from junifer.markers.crossparcellation_functional_connectivity import ( - CrossParcellationFC, -) +from junifer.markers.functional_connectivity import CrossParcellationFC from junifer.storage import SQLiteFeatureStorage from junifer.testing.datagrabbers import SPMAuditoryTestingDatagrabber diff --git a/junifer/markers/functional_connectivity/tests/test_functional_connectivity_base.py b/junifer/markers/functional_connectivity/tests/test_functional_connectivity_base.py new file mode 100644 index 000000000..634267fad --- /dev/null +++ b/junifer/markers/functional_connectivity/tests/test_functional_connectivity_base.py @@ -0,0 +1,15 @@ +"""Provide tests for base functional connectivity marker.""" + +# Authors: Synchon Mandal +# 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() diff --git a/junifer/markers/tests/test_functional_connectivity_parcels.py b/junifer/markers/functional_connectivity/tests/test_functional_connectivity_parcels.py similarity index 98% rename from junifer/markers/tests/test_functional_connectivity_parcels.py rename to junifer/markers/functional_connectivity/tests/test_functional_connectivity_parcels.py index 5b5758970..aa438d8a5 100644 --- a/junifer/markers/tests/test_functional_connectivity_parcels.py +++ b/junifer/markers/functional_connectivity/tests/test_functional_connectivity_parcels.py @@ -11,7 +11,7 @@ from nilearn.connectome import ConnectivityMeasure from nilearn.maskers import NiftiLabelsMasker 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, ) from junifer.markers.parcel_aggregation import ParcelAggregation diff --git a/junifer/markers/tests/test_functional_connectivity_spheres.py b/junifer/markers/functional_connectivity/tests/test_functional_connectivity_spheres.py similarity index 98% rename from junifer/markers/tests/test_functional_connectivity_spheres.py rename to junifer/markers/functional_connectivity/tests/test_functional_connectivity_spheres.py index 819d4a6ad..ddc4c842b 100644 --- a/junifer/markers/tests/test_functional_connectivity_spheres.py +++ b/junifer/markers/functional_connectivity/tests/test_functional_connectivity_spheres.py @@ -13,7 +13,7 @@ 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 import ( FunctionalConnectivitySpheres, ) from junifer.markers.sphere_aggregation import SphereAggregation diff --git a/junifer/markers/functional_connectivity_spheres.py b/junifer/markers/functional_connectivity_spheres.py deleted file mode 100644 index 97b485e7b..000000000 --- a/junifer/markers/functional_connectivity_spheres.py +++ /dev/null @@ -1,165 +0,0 @@ -"""Provide base class for functional connectivity using spheres.""" - -# Authors: Amir Omidvarnia -# Kaustubh R. Patil -# 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