[ENH] Refactor Functional Connectivity based markers #107
13 changed files with 241 additions and 219 deletions
|
|
@ -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
|
||||||
~~~~
|
~~~~
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
8
junifer/markers/functional_connectivity/__init__.py
Normal file
8
junifer/markers/functional_connectivity/__init__.py
Normal 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
|
||||||
|
|
@ -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(
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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)
|
||||||
|
|
@ -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)
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
@ -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()
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
|
||||||
Loading…
Reference in a new issue