[ENH]: Introduce JuniferConnectivityMeasure #348
19 changed files with 1866 additions and 221 deletions
1
docs/changes/newsfragments/348.change
Normal file
1
docs/changes/newsfragments/348.change
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
For :class:`.CrossParcellationFC`, ``aggregation_method`` and ``correlation_method`` have been renamed to ``agg_method`` and ``corr_method`` respectively and ``agg_method_params`` has been added; for ``FunctionalConnectivityBase``, :class:`.FunctionalConnectivityParcels`, :class:`.FunctionalConnectivitySpheres`, :class:`.EdgeCentricFCParcels` and :class:`.EdgeCentricFCSpheres`, ``cor_method`` and ``cor_method_params`` have been renamed to ``conn_method`` and ``conn_method_params`` by `Synchon Mandal`_
|
||||||
1
docs/changes/newsfragments/348.enh
Normal file
1
docs/changes/newsfragments/348.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
``FunctionalConnectivity``-family Markers now use :class:`sklearn.covariance.EmpiricalCovariance` as the default covariance estimator and ``correlation`` as the default connecivity matrix kind by `Synchon Mandal`_
|
||||||
1
docs/changes/newsfragments/348.feature
Normal file
1
docs/changes/newsfragments/348.feature
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Introduce :class:`.JuniferConnectivityMeasure` for customising functional connectivity matrix kinds and measurements by `Synchon Mandal`_
|
||||||
|
|
@ -170,6 +170,7 @@ numpydoc_xref_ignore = {
|
||||||
"Engine",
|
"Engine",
|
||||||
"positive",
|
"positive",
|
||||||
"negative",
|
"negative",
|
||||||
|
"estimator",
|
||||||
}
|
}
|
||||||
# numpydoc_validation_checks = {
|
# numpydoc_validation_checks = {
|
||||||
# "all",
|
# "all",
|
||||||
|
|
|
||||||
|
|
@ -5,3 +5,4 @@ sepulcre
|
||||||
arange
|
arange
|
||||||
sinc
|
sinc
|
||||||
whit
|
whit
|
||||||
|
amin
|
||||||
|
|
|
||||||
3
junifer/external/nilearn/__init__.py
vendored
3
junifer/external/nilearn/__init__.py
vendored
|
|
@ -4,6 +4,7 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from .junifer_nifti_spheres_masker import JuniferNiftiSpheresMasker
|
from .junifer_nifti_spheres_masker import JuniferNiftiSpheresMasker
|
||||||
|
from .junifer_connectivity_measure import JuniferConnectivityMeasure
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["JuniferNiftiSpheresMasker"]
|
__all__ = ["JuniferNiftiSpheresMasker", "JuniferConnectivityMeasure"]
|
||||||
|
|
|
||||||
483
junifer/external/nilearn/junifer_connectivity_measure.py
vendored
Normal file
483
junifer/external/nilearn/junifer_connectivity_measure.py
vendored
Normal file
|
|
@ -0,0 +1,483 @@
|
||||||
|
"""Provide JuniferConnectivityMeasure class."""
|
||||||
|
|
||||||
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
|
# License: AGPL
|
||||||
|
|
||||||
|
from typing import Callable, List, Optional
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from nilearn import signal
|
||||||
|
from nilearn.connectome import (
|
||||||
|
ConnectivityMeasure,
|
||||||
|
cov_to_corr,
|
||||||
|
prec_to_partial,
|
||||||
|
sym_matrix_to_vec,
|
||||||
|
)
|
||||||
|
from scipy import linalg
|
||||||
|
from sklearn.base import clone
|
||||||
|
from sklearn.covariance import EmpiricalCovariance
|
||||||
|
|
||||||
|
from ...utils import logger, raise_error, warn_with_log
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["JuniferConnectivityMeasure"]
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_COV_ESTIMATOR = EmpiricalCovariance(store_precision=False)
|
||||||
|
|
||||||
|
|
||||||
|
# New BSD License
|
||||||
|
|
||||||
|
# Copyright (c) The nilearn developers.
|
||||||
|
# All rights reserved.
|
||||||
|
|
||||||
|
|
||||||
|
# Redistribution and use in source and binary forms, with or without
|
||||||
|
# modification, are permitted provided that the following conditions are met:
|
||||||
|
|
||||||
|
# a. Redistributions of source code must retain the above copyright notice,
|
||||||
|
# this list of conditions and the following disclaimer.
|
||||||
|
# b. Redistributions in binary form must reproduce the above copyright
|
||||||
|
# notice, this list of conditions and the following disclaimer in the
|
||||||
|
# documentation and/or other materials provided with the distribution.
|
||||||
|
# c. Neither the name of the nilearn developers nor the names of
|
||||||
|
# its contributors may be used to endorse or promote products
|
||||||
|
# derived from this software without specific prior written
|
||||||
|
# permission.
|
||||||
|
|
||||||
|
|
||||||
|
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||||
|
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||||
|
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
|
||||||
|
# ARE DISCLAIMED. IN NO EVENT SHALL THE REGENTS OR CONTRIBUTORS BE LIABLE FOR
|
||||||
|
# ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||||
|
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||||
|
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||||
|
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT
|
||||||
|
# LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY
|
||||||
|
# OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH
|
||||||
|
# DAMAGE.
|
||||||
|
|
||||||
|
|
||||||
|
def _check_square(matrix: np.ndarray) -> None:
|
||||||
|
"""Raise a ValueError if the input matrix is square.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
matrix : numpy.ndarray
|
||||||
|
Input array.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError
|
||||||
|
If ``matrix`` is not a square matrix.
|
||||||
|
|
||||||
|
"""
|
||||||
|
if matrix.ndim != 2 or (matrix.shape[0] != matrix.shape[-1]):
|
||||||
|
raise_error(
|
||||||
|
f"Expected a square matrix, got array of shape {matrix.shape}."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_spd(M: np.ndarray, decimal: int = 15) -> bool: # noqa: N803
|
||||||
|
"""Check that input matrix is symmetric positive definite.
|
||||||
|
|
||||||
|
``M`` must be symmetric down to specified ``decimal`` places.
|
||||||
|
The check is performed by checking that all eigenvalues are positive.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
M : numpy.ndarray
|
||||||
|
Input matrix to check for symmetric positive definite.
|
||||||
|
decimal : int, optional
|
||||||
|
Decimal places to check (default 15).
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
bool
|
||||||
|
True if matrix is symmetric positive definite, False otherwise.
|
||||||
|
|
||||||
|
"""
|
||||||
|
if not np.allclose(M, M.T, atol=0, rtol=10**-decimal):
|
||||||
|
logger.debug(f"matrix not symmetric to {decimal:d} decimals")
|
||||||
|
return False
|
||||||
|
eigvalsh = np.linalg.eigvalsh(M)
|
||||||
|
ispd = eigvalsh.min() > 0
|
||||||
|
if not ispd:
|
||||||
|
logger.debug(f"matrix has a negative eigenvalue: {eigvalsh.min():.3f}")
|
||||||
|
return ispd
|
||||||
|
|
||||||
|
|
||||||
|
def _check_spd(matrix: np.ndarray) -> None:
|
||||||
|
"""Check ``matrix`` is symmetric positive definite.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
matrix : numpy.ndarray
|
||||||
|
Input array.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError
|
||||||
|
If the input matrix is not symmetric positive definite.
|
||||||
|
|
||||||
|
"""
|
||||||
|
if not is_spd(matrix, decimal=7):
|
||||||
|
raise_error("Expected a symmetric positive definite matrix.")
|
||||||
|
|
||||||
|
|
||||||
|
def _form_symmetric(
|
||||||
|
function: Callable[[np.ndarray], np.ndarray],
|
||||||
|
eigenvalues: np.ndarray,
|
||||||
|
eigenvectors: np.ndarray,
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""Return the symmetric matrix.
|
||||||
|
|
||||||
|
Apply ``function`` to ``eigenvalues``, construct symmetric matrix with it
|
||||||
|
and ``eigenvectors`` and return the constructed symmetric matrix.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
function : callable (function numpy.ndarray -> numpy.ndarray)
|
||||||
|
The transform to apply to the eigenvalues.
|
||||||
|
eigenvalues : numpy.ndarray of shape (n_features, )
|
||||||
|
Input argument of the function.
|
||||||
|
eigenvectors : numpy.ndarray of shape (n_features, n_features)
|
||||||
|
Unitary matrix.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
numpy.ndarray of shape (n_features, n_features)
|
||||||
|
The symmetric matrix obtained after transforming the eigenvalues, while
|
||||||
|
keeping the same eigenvectors.
|
||||||
|
|
||||||
|
"""
|
||||||
|
return np.dot(eigenvectors * function(eigenvalues), eigenvectors.T)
|
||||||
|
|
||||||
|
|
||||||
|
def _map_eigenvalues(
|
||||||
|
function: Callable[[np.ndarray], np.ndarray], symmetric: np.ndarray
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""Matrix function, for real symmetric matrices.
|
||||||
|
|
||||||
|
The function is applied to the eigenvalues of ``symmetric``.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
function : callable (function numpy.ndarray -> numpy.ndarray)
|
||||||
|
The transform to apply to the eigenvalues.
|
||||||
|
symmetric : numpy.ndarray of shape (n_features, n_features)
|
||||||
|
The input symmetric matrix.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
numpy.ndarray of shape (n_features, n_features)
|
||||||
|
The new symmetric matrix obtained after transforming the eigenvalues,
|
||||||
|
while keeping the same eigenvectors.
|
||||||
|
|
||||||
|
Notes
|
||||||
|
-----
|
||||||
|
If input matrix is not real symmetric, no error is reported but result will
|
||||||
|
be wrong.
|
||||||
|
|
||||||
|
"""
|
||||||
|
eigenvalues, eigenvectors = linalg.eigh(symmetric)
|
||||||
|
return _form_symmetric(function, eigenvalues, eigenvectors)
|
||||||
|
|
||||||
|
|
||||||
|
def _geometric_mean(
|
||||||
|
matrices: List[np.ndarray],
|
||||||
|
init: Optional[np.ndarray] = None,
|
||||||
|
max_iter: int = 10,
|
||||||
|
tol: Optional[float] = 1e-7,
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""Compute the geometric mean of symmetric positive definite matrices.
|
||||||
|
|
||||||
|
The geometric mean of ``n`` positive definite matrices
|
||||||
|
``M_1, ..., M_n`` is the minimizer of the sum of squared distances from an
|
||||||
|
arbitrary matrix to each input matrix ``M_k``
|
||||||
|
|
||||||
|
.. math:: gmean(M_1, ..., M_n) = argmin_X sum_{k=1}^N dist(X, M_k)^2
|
||||||
|
|
||||||
|
where the used distance is related to matrices logarithm
|
||||||
|
|
||||||
|
.. math:: dist(X, M_k) = ||log(X^{-1/2} M_k X^{-1/2)}||
|
||||||
|
|
||||||
|
In case of positive numbers, this mean is the usual geometric mean.
|
||||||
|
|
||||||
|
See Algorithm 3 of [1]_ .
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
matrices : list of numpy.ndarray, all of shape (n_features, n_features)
|
||||||
|
List of matrices whose geometric mean to compute. Raise an error if the
|
||||||
|
matrices are not all symmetric positive definite of the same shape.
|
||||||
|
init : numpy.ndarray of shape (n_features, n_features), optional
|
||||||
|
Initialization matrix, default to the arithmetic mean of matrices.
|
||||||
|
Raise an error if the matrix is not symmetric positive definite of the
|
||||||
|
same shape as the elements of matrices (default None).
|
||||||
|
max_iter : int, optional
|
||||||
|
Maximal number of iterations (default 10).
|
||||||
|
tol : positive float or None, optional
|
||||||
|
The tolerance to declare convergence: if the gradient norm goes below
|
||||||
|
this value, the gradient descent is stopped. If None, no check is
|
||||||
|
performed (default 1e-7).
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
gmean : numpy.ndarray of shape (n_features, n_features)
|
||||||
|
Geometric mean of the matrices.
|
||||||
|
|
||||||
|
References
|
||||||
|
----------
|
||||||
|
.. [1] Fletcher, T., P., Joshi, S.
|
||||||
|
Riemannian geometry for the statistical analysis of diffusion tensor
|
||||||
|
data.
|
||||||
|
Signal Processing, Volume 87, Issue 2, 2007, Pages 250-262
|
||||||
|
https://doi.org/10.1016/j.sigpro.2005.12.018.
|
||||||
|
|
||||||
|
"""
|
||||||
|
# Shape and symmetry positive definiteness checks
|
||||||
|
n_features = matrices[0].shape[0]
|
||||||
|
for matrix in matrices:
|
||||||
|
_check_square(matrix)
|
||||||
|
if matrix.shape[0] != n_features:
|
||||||
|
raise_error("Matrices are not of the same shape.")
|
||||||
|
_check_spd(matrix)
|
||||||
|
|
||||||
|
# Initialization
|
||||||
|
matrices = np.array(matrices)
|
||||||
|
if init is None:
|
||||||
|
gmean = np.mean(matrices, axis=0)
|
||||||
|
else:
|
||||||
|
_check_square(init)
|
||||||
|
if init.shape[0] != n_features:
|
||||||
|
raise_error("Initialization has incorrect shape.")
|
||||||
|
_check_spd(init)
|
||||||
|
gmean = init
|
||||||
|
|
||||||
|
norm_old = np.inf
|
||||||
|
step = 1.0
|
||||||
|
|
||||||
|
# Gradient descent
|
||||||
|
for _ in range(max_iter):
|
||||||
|
# Computation of the gradient
|
||||||
|
vals_gmean, vecs_gmean = linalg.eigh(gmean)
|
||||||
|
gmean_inv_sqrt = _form_symmetric(np.sqrt, 1.0 / vals_gmean, vecs_gmean)
|
||||||
|
whitened_matrices = [
|
||||||
|
gmean_inv_sqrt.dot(matrix).dot(gmean_inv_sqrt)
|
||||||
|
for matrix in matrices
|
||||||
|
]
|
||||||
|
logs = [_map_eigenvalues(np.log, w_mat) for w_mat in whitened_matrices]
|
||||||
|
# Covariant derivative is - gmean.dot(logms_mean)
|
||||||
|
logs_mean = np.mean(logs, axis=0)
|
||||||
|
if np.any(np.isnan(logs_mean)):
|
||||||
|
raise_error(
|
||||||
|
klass=FloatingPointError,
|
||||||
|
msg="Nan value after logarithm operation.",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Norm of the covariant derivative on the tangent space at point gmean
|
||||||
|
norm = np.linalg.norm(logs_mean)
|
||||||
|
|
||||||
|
# Update of the minimizer
|
||||||
|
vals_log, vecs_log = linalg.eigh(logs_mean)
|
||||||
|
gmean_sqrt = _form_symmetric(np.sqrt, vals_gmean, vecs_gmean)
|
||||||
|
# Move along the geodesic
|
||||||
|
gmean = gmean_sqrt.dot(
|
||||||
|
_form_symmetric(np.exp, vals_log * step, vecs_log)
|
||||||
|
).dot(gmean_sqrt)
|
||||||
|
|
||||||
|
# Update the norm and the step size
|
||||||
|
if norm < norm_old:
|
||||||
|
norm_old = norm
|
||||||
|
elif norm > norm_old:
|
||||||
|
step = step / 2.0
|
||||||
|
norm = norm_old
|
||||||
|
if tol is not None and norm / gmean.size < tol:
|
||||||
|
break
|
||||||
|
if tol is not None and norm / gmean.size >= tol:
|
||||||
|
warn_with_log(
|
||||||
|
f"Maximum number of iterations {max_iter} reached without "
|
||||||
|
f"getting to the requested tolerance level {tol}."
|
||||||
|
)
|
||||||
|
|
||||||
|
return gmean
|
||||||
|
|
||||||
|
|
||||||
|
class JuniferConnectivityMeasure(ConnectivityMeasure):
|
||||||
|
"""Class for custom ConnectivityMeasure.
|
||||||
|
|
||||||
|
Differs from :class:`nilearn.connectome.ConnectivityMeasure` in the
|
||||||
|
following ways:
|
||||||
|
|
||||||
|
* default ``cov_estimator`` is
|
||||||
|
:class:`sklearn.covariance.EmpiricalCovariance`
|
||||||
|
* default ``kind`` is ``"correlation"``
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
cov_estimator : estimator object, optional
|
||||||
|
The covariance estimator
|
||||||
|
(default ``EmpiricalCovariance(store_precision=False)``).
|
||||||
|
kind : {"covariance", "correlation", "partial correlation", \
|
||||||
|
"tangent", "precision"}, optional
|
||||||
|
The matrix kind. For the use of ``"tangent"`` see [1]_
|
||||||
|
(default "correlation").
|
||||||
|
vectorize : bool, optional
|
||||||
|
If True, connectivity matrices are reshaped into 1D arrays and only
|
||||||
|
their flattened lower triangular parts are returned (default False).
|
||||||
|
discard_diagonal : bool, optional
|
||||||
|
If True, vectorized connectivity coefficients do not include the
|
||||||
|
matrices diagonal elements. Used only when vectorize is set to True
|
||||||
|
(default False).
|
||||||
|
standardize : bool, optional
|
||||||
|
If standardize is True, the data are centered and normed: their mean
|
||||||
|
is put to 0 and their variance is put to 1 in the time dimension
|
||||||
|
(default True).
|
||||||
|
|
||||||
|
.. note::
|
||||||
|
|
||||||
|
Added to control passing value to ``standardize`` of
|
||||||
|
``signal.clean`` to call new behavior since passing ``"zscore"`` or
|
||||||
|
True (default) is deprecated. This parameter will be deprecated in
|
||||||
|
version 0.13 and removed in version 0.15.
|
||||||
|
|
||||||
|
Attributes
|
||||||
|
----------
|
||||||
|
cov_estimator_ : estimator object
|
||||||
|
A new covariance estimator with the same parameters as
|
||||||
|
``cov_estimator``.
|
||||||
|
mean_ : numpy.ndarray
|
||||||
|
The mean connectivity matrix across subjects. For ``"tangent"`` kind,
|
||||||
|
it is the geometric mean of covariances (a group covariance
|
||||||
|
matrix that captures information from both correlation and partial
|
||||||
|
correlation matrices). For other values for ``kind``, it is the
|
||||||
|
mean of the corresponding matrices.
|
||||||
|
whitening_ : numpy.ndarray
|
||||||
|
The inverted square-rooted geometric mean of the covariance matrices.
|
||||||
|
|
||||||
|
References
|
||||||
|
----------
|
||||||
|
.. [1] Varoquaux, G., Baronnet, F., Kleinschmidt, A. et al.
|
||||||
|
Detection of brain functional-connectivity difference in
|
||||||
|
post-stroke patients using group-level covariance modeling.
|
||||||
|
In Tianzi Jiang, Nassir Navab, Josien P. W. Pluim, and
|
||||||
|
Max A. Viergever, editors, Medical image computing and
|
||||||
|
computer-assisted intervention - MICCAI 2010, Lecture notes
|
||||||
|
in computer science, Pages 200-208. Berlin, Heidelberg, 2010.
|
||||||
|
Springer.
|
||||||
|
doi:10/cn2h9c.
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
cov_estimator=DEFAULT_COV_ESTIMATOR,
|
||||||
|
kind="correlation",
|
||||||
|
vectorize=False,
|
||||||
|
discard_diagonal=False,
|
||||||
|
standardize=True,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
cov_estimator=cov_estimator,
|
||||||
|
kind=kind,
|
||||||
|
vectorize=vectorize,
|
||||||
|
discard_diagonal=discard_diagonal,
|
||||||
|
standardize=standardize,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _fit_transform(
|
||||||
|
self,
|
||||||
|
X, # noqa: N803
|
||||||
|
do_transform=False,
|
||||||
|
do_fit=False,
|
||||||
|
confounds=None,
|
||||||
|
):
|
||||||
|
"""Avoid duplication of computation."""
|
||||||
|
self._check_input(X, confounds=confounds)
|
||||||
|
if do_fit:
|
||||||
|
self.cov_estimator_ = clone(self.cov_estimator)
|
||||||
|
|
||||||
|
# Compute all the matrices, stored in "connectivities"
|
||||||
|
if self.kind == "correlation":
|
||||||
|
covariances_std = [
|
||||||
|
self.cov_estimator_.fit(
|
||||||
|
signal.standardize_signal(
|
||||||
|
x,
|
||||||
|
detrend=False,
|
||||||
|
standardize=self.standardize,
|
||||||
|
)
|
||||||
|
).covariance_
|
||||||
|
for x in X
|
||||||
|
]
|
||||||
|
connectivities = [cov_to_corr(cov) for cov in covariances_std]
|
||||||
|
else:
|
||||||
|
covariances = [self.cov_estimator_.fit(x).covariance_ for x in X]
|
||||||
|
if self.kind in ("covariance", "tangent"):
|
||||||
|
connectivities = covariances
|
||||||
|
elif self.kind == "precision":
|
||||||
|
connectivities = [linalg.inv(cov) for cov in covariances]
|
||||||
|
elif self.kind == "partial correlation":
|
||||||
|
connectivities = [
|
||||||
|
prec_to_partial(linalg.inv(cov)) for cov in covariances
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
allowed_kinds = (
|
||||||
|
"correlation",
|
||||||
|
"partial correlation",
|
||||||
|
"tangent",
|
||||||
|
"covariance",
|
||||||
|
"precision",
|
||||||
|
)
|
||||||
|
raise_error(
|
||||||
|
f"Allowed connectivity kinds are {allowed_kinds}. "
|
||||||
|
f"Got kind {self.kind}."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Store the mean
|
||||||
|
if do_fit:
|
||||||
|
if self.kind == "tangent":
|
||||||
|
self.mean_ = _geometric_mean(
|
||||||
|
covariances, max_iter=30, tol=1e-7
|
||||||
|
)
|
||||||
|
self.whitening_ = _map_eigenvalues(
|
||||||
|
lambda x: 1.0 / np.sqrt(x), self.mean_
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.mean_ = np.mean(connectivities, axis=0)
|
||||||
|
# Fight numerical instabilities: make symmetric
|
||||||
|
self.mean_ = self.mean_ + self.mean_.T
|
||||||
|
self.mean_ *= 0.5
|
||||||
|
|
||||||
|
# Compute the vector we return on transform
|
||||||
|
if do_transform:
|
||||||
|
if self.kind == "tangent":
|
||||||
|
connectivities = [
|
||||||
|
_map_eigenvalues(
|
||||||
|
np.log, self.whitening_.dot(cov).dot(self.whitening_)
|
||||||
|
)
|
||||||
|
for cov in connectivities
|
||||||
|
]
|
||||||
|
|
||||||
|
connectivities = np.array(connectivities)
|
||||||
|
|
||||||
|
if confounds is not None and not self.vectorize:
|
||||||
|
error_message = (
|
||||||
|
"'confounds' are provided but vectorize=False. "
|
||||||
|
"Confounds are only cleaned on vectorized matrices "
|
||||||
|
"as second level connectome regression "
|
||||||
|
"but not on symmetric matrices."
|
||||||
|
)
|
||||||
|
raise_error(error_message)
|
||||||
|
|
||||||
|
if self.vectorize:
|
||||||
|
connectivities = sym_matrix_to_vec(
|
||||||
|
connectivities, discard_diagonal=self.discard_diagonal
|
||||||
|
)
|
||||||
|
if confounds is not None:
|
||||||
|
connectivities = signal.clean(
|
||||||
|
connectivities, confounds=confounds
|
||||||
|
)
|
||||||
|
|
||||||
|
return connectivities
|
||||||
1089
junifer/external/nilearn/tests/test_junifer_connectivity_measure.py
vendored
Normal file
1089
junifer/external/nilearn/tests/test_junifer_connectivity_measure.py
vendored
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -28,18 +28,24 @@ class CrossParcellationFC(BaseMarker):
|
||||||
The name of the first parcellation.
|
The name of the first parcellation.
|
||||||
parcellation_two : str
|
parcellation_two : str
|
||||||
The name of the second parcellation.
|
The name of the second parcellation.
|
||||||
aggregation_method : str, optional
|
agg_method : str, optional
|
||||||
The aggregation method (default "mean").
|
The method to perform aggregation using.
|
||||||
correlation_method : str, optional
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
|
(default "mean").
|
||||||
|
agg_method_params : dict, optional
|
||||||
|
Parameters to pass to the aggregation function.
|
||||||
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
|
(default None).
|
||||||
|
corr_method : str, optional
|
||||||
Any method that can be passed to
|
Any method that can be passed to
|
||||||
:any:`pandas.DataFrame.corr` (default "pearson").
|
:meth:`pandas.DataFrame.corr` (default "pearson").
|
||||||
masks : str, dict or list of dict or str, optional
|
masks : str, dict or list of dict or str, optional
|
||||||
The specification of the masks to apply to regions before extracting
|
The specification of the masks to apply to regions before extracting
|
||||||
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
||||||
If None, will not apply any mask (default None).
|
If None, will not apply any mask (default None).
|
||||||
name : str, optional
|
name : str, optional
|
||||||
The name of the marker. If None, will use the class name
|
The name of the marker. If None, will use
|
||||||
(default None).
|
``BOLD_CrossParcellationFC`` (default None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -55,8 +61,9 @@ class CrossParcellationFC(BaseMarker):
|
||||||
self,
|
self,
|
||||||
parcellation_one: str,
|
parcellation_one: str,
|
||||||
parcellation_two: str,
|
parcellation_two: str,
|
||||||
aggregation_method: str = "mean",
|
agg_method: str = "mean",
|
||||||
correlation_method: str = "pearson",
|
agg_method_params: Optional[Dict] = None,
|
||||||
|
corr_method: str = "pearson",
|
||||||
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -66,8 +73,9 @@ class CrossParcellationFC(BaseMarker):
|
||||||
)
|
)
|
||||||
self.parcellation_one = parcellation_one
|
self.parcellation_one = parcellation_one
|
||||||
self.parcellation_two = parcellation_two
|
self.parcellation_two = parcellation_two
|
||||||
self.aggregation_method = aggregation_method
|
self.agg_method = agg_method
|
||||||
self.correlation_method = correlation_method
|
self.agg_method_params = agg_method_params
|
||||||
|
self.corr_method = corr_method
|
||||||
self.masks = masks
|
self.masks = masks
|
||||||
super().__init__(on=["BOLD"], name=name)
|
super().__init__(on=["BOLD"], name=name)
|
||||||
|
|
||||||
|
|
@ -115,13 +123,17 @@ class CrossParcellationFC(BaseMarker):
|
||||||
# Perform aggregation using two parcellations
|
# Perform aggregation using two parcellations
|
||||||
aggregation_parcellation_one = ParcelAggregation(
|
aggregation_parcellation_one = ParcelAggregation(
|
||||||
parcellation=self.parcellation_one,
|
parcellation=self.parcellation_one,
|
||||||
method=self.aggregation_method,
|
method=self.agg_method,
|
||||||
|
method_params=self.agg_method_params,
|
||||||
masks=self.masks,
|
masks=self.masks,
|
||||||
|
on="BOLD",
|
||||||
).compute(input, extra_input=extra_input)
|
).compute(input, extra_input=extra_input)
|
||||||
aggregation_parcellation_two = ParcelAggregation(
|
aggregation_parcellation_two = ParcelAggregation(
|
||||||
parcellation=self.parcellation_two,
|
parcellation=self.parcellation_two,
|
||||||
method=self.aggregation_method,
|
method=self.agg_method,
|
||||||
|
method_params=self.agg_method_params,
|
||||||
masks=self.masks,
|
masks=self.masks,
|
||||||
|
on="BOLD",
|
||||||
).compute(input, extra_input=extra_input)
|
).compute(input, extra_input=extra_input)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|
@ -133,7 +145,7 @@ class CrossParcellationFC(BaseMarker):
|
||||||
pd.DataFrame(
|
pd.DataFrame(
|
||||||
aggregation_parcellation_two["aggregation"]["data"]
|
aggregation_parcellation_two["aggregation"]["data"]
|
||||||
),
|
),
|
||||||
method=self.correlation_method,
|
method=self.corr_method,
|
||||||
).values,
|
).values,
|
||||||
# Columns should be named after parcellation 1
|
# Columns should be named after parcellation 1
|
||||||
"col_names": aggregation_parcellation_one["aggregation"][
|
"col_names": aggregation_parcellation_one["aggregation"][
|
||||||
|
|
|
||||||
|
|
@ -22,36 +22,40 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
parcellation : str or list of str
|
parcellation : str or list of str
|
||||||
The name(s) of the parcellation(s). Check valid options by calling
|
The name(s) of the parcellation(s) to use.
|
||||||
:func:`.list_parcellations`.
|
See :func:`.list_parcellations` for options.
|
||||||
agg_method : str, optional
|
agg_method : str, optional
|
||||||
The method to perform aggregation of BOLD time series.
|
The method to perform aggregation using.
|
||||||
Check valid options in :func:`.get_aggfunc_by_name`
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
(default "mean").
|
(default "mean").
|
||||||
agg_method_params : dict, optional
|
agg_method_params : dict, optional
|
||||||
Parameters to pass to the aggregation function. Check valid options in
|
Parameters to pass to the aggregation function.
|
||||||
:func:`.get_aggfunc_by_name` (default None).
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
cor_method : str, optional
|
(default None).
|
||||||
The method to perform correlation. Check valid options in
|
conn_method : str, optional
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure`
|
The method to perform connectivity measure using.
|
||||||
(default "covariance").
|
See :class:`.JuniferConnectivityMeasure` for options
|
||||||
cor_method_params : dict, optional
|
(default "correlation").
|
||||||
Parameters to pass to the correlation function. Check valid options in
|
conn_method_params : dict, optional
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
|
Parameters to pass to :class:`.JuniferConnectivityMeasure`.
|
||||||
|
If None, ``{"empirical": True}`` will be used, which would mean
|
||||||
|
:class:`sklearn.covariance.EmpiricalCovariance` is used to compute
|
||||||
|
covariance. If usage of :class:`sklearn.covariance.LedoitWolf` is
|
||||||
|
desired, ``{"empirical": False}`` should be passed
|
||||||
|
(default None).
|
||||||
masks : str, dict or list of dict or str, optional
|
masks : str, dict or list of dict or str, optional
|
||||||
The specification of the masks to apply to regions before extracting
|
The specification of the masks to apply to regions before extracting
|
||||||
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
||||||
If None, will not apply any mask (default None).
|
If None, will not apply any mask (default None).
|
||||||
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
|
||||||
None).
|
``BOLD_EdgeCentricFCParcels`` (default None).
|
||||||
|
|
||||||
References
|
References
|
||||||
----------
|
----------
|
||||||
.. [1] Jo et al. (2021)
|
.. [1] Jo et al. (2021)
|
||||||
Subject identification using
|
Subject identification using edge-centric functional connectivity.
|
||||||
edge-centric functional connectivity
|
https://doi.org/10.1016/j.neuroimage.2021.118204
|
||||||
doi: https://doi.org/10.1016/j.neuroimage.2021.118204
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -60,8 +64,8 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
|
||||||
parcellation: Union[str, List[str]],
|
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",
|
conn_method: str = "correlation",
|
||||||
cor_method_params: Optional[Dict] = None,
|
conn_method_params: Optional[Dict] = None,
|
||||||
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -69,8 +73,8 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
agg_method=agg_method,
|
agg_method=agg_method,
|
||||||
agg_method_params=agg_method_params,
|
agg_method_params=agg_method_params,
|
||||||
cor_method=cor_method,
|
conn_method=conn_method,
|
||||||
cor_method_params=cor_method_params,
|
conn_method_params=conn_method_params,
|
||||||
masks=masks,
|
masks=masks,
|
||||||
name=name,
|
name=name,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -22,42 +22,48 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
coords : str
|
coords : str
|
||||||
The name of the coordinates list to use. See
|
The name of the coordinates list to use.
|
||||||
:func:`.list_coordinates` for options.
|
See :func:`.list_coordinates` for options.
|
||||||
radius : float, optional
|
radius : positive float, optional
|
||||||
The radius of the sphere in mm. If None, the signal will be extracted
|
The radius of the sphere around each coordinates in millimetres.
|
||||||
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
|
If None, the signal will be extracted from a single voxel.
|
||||||
for more information (default None).
|
See :class:`.JuniferNiftiSpheresMasker` for more information
|
||||||
|
(default None).
|
||||||
allow_overlap : bool, optional
|
allow_overlap : bool, optional
|
||||||
Whether to allow overlapping spheres. If False, an error is raised if
|
Whether to allow overlapping spheres. If False, an error is raised if
|
||||||
the spheres overlap (default is False).
|
the spheres overlap (default False).
|
||||||
agg_method : str, optional
|
agg_method : str, optional
|
||||||
The aggregation method to use.
|
The method to perform aggregation using.
|
||||||
See :func:`.get_aggfunc_by_name` for more information
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
(default None).
|
(default "mean").
|
||||||
agg_method_params : dict, optional
|
agg_method_params : dict, optional
|
||||||
The parameters to pass to the aggregation method (default None).
|
Parameters to pass to the aggregation function.
|
||||||
cor_method : str, optional
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
The method to perform correlation using. Check valid options in
|
(default None).
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
|
conn_method : str, optional
|
||||||
cor_method_params : dict, optional
|
The method to perform connectivity measure using.
|
||||||
Parameters to pass to the correlation function. Check valid options in
|
See :class:`.JuniferConnectivityMeasure` for options
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
|
(default "correlation").
|
||||||
|
conn_method_params : dict, optional
|
||||||
|
Parameters to pass to :class:`.JuniferConnectivityMeasure`.
|
||||||
|
If None, ``{"empirical": True}`` will be used, which would mean
|
||||||
|
:class:`sklearn.covariance.EmpiricalCovariance` is used to compute
|
||||||
|
covariance. If usage of :class:`sklearn.covariance.LedoitWolf` is
|
||||||
|
desired, ``{"empirical": False}`` should be passed
|
||||||
|
(default None).
|
||||||
masks : str, dict or list of dict or str, optional
|
masks : str, dict or list of dict or str, optional
|
||||||
The specification of the masks to apply to regions before extracting
|
The specification of the masks to apply to regions before extracting
|
||||||
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
||||||
If None, will not apply any mask (default None).
|
If None, will not apply any mask (default None).
|
||||||
name : str, optional
|
name : str, optional
|
||||||
The name of the marker. By default, it will use
|
The name of the marker. If None, will use
|
||||||
KIND_EdgeCentricFCSpheres where KIND is the kind of data it
|
``BOLD_EdgeCentricFCSpheres`` (default None).
|
||||||
was applied to (default None).
|
|
||||||
|
|
||||||
References
|
References
|
||||||
----------
|
----------
|
||||||
.. [1] Jo et al. (2021)
|
.. [1] Jo et al. (2021)
|
||||||
Subject identification using
|
Subject identification using edge-centric functional connectivity.
|
||||||
edge-centric functional connectivity
|
https://doi.org/10.1016/j.neuroimage.2021.118204
|
||||||
doi: https://doi.org/10.1016/j.neuroimage.2021.118204
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -68,8 +74,8 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
|
||||||
allow_overlap: bool = False,
|
allow_overlap: bool = False,
|
||||||
agg_method: str = "mean",
|
agg_method: str = "mean",
|
||||||
agg_method_params: Optional[Dict] = None,
|
agg_method_params: Optional[Dict] = None,
|
||||||
cor_method: str = "covariance",
|
conn_method: str = "correlation",
|
||||||
cor_method_params: Optional[Dict] = None,
|
conn_method_params: Optional[Dict] = None,
|
||||||
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -81,8 +87,8 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
agg_method=agg_method,
|
agg_method=agg_method,
|
||||||
agg_method_params=agg_method_params,
|
agg_method_params=agg_method_params,
|
||||||
cor_method=cor_method,
|
conn_method=conn_method,
|
||||||
cor_method_params=cor_method_params,
|
conn_method_params=conn_method_params,
|
||||||
masks=masks,
|
masks=masks,
|
||||||
name=name,
|
name=name,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -7,9 +7,9 @@
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from typing import Any, ClassVar, Dict, List, Optional, Set, Union
|
from typing import Any, ClassVar, Dict, List, Optional, Set, Union
|
||||||
|
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from sklearn.covariance import EmpiricalCovariance, LedoitWolf
|
||||||
from sklearn.covariance import EmpiricalCovariance
|
|
||||||
|
|
||||||
|
from ...external.nilearn import JuniferConnectivityMeasure
|
||||||
from ...utils import raise_error
|
from ...utils import raise_error
|
||||||
from ..base import BaseMarker
|
from ..base import BaseMarker
|
||||||
|
|
||||||
|
|
@ -23,25 +23,31 @@ class FunctionalConnectivityBase(BaseMarker):
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
agg_method : str, optional
|
agg_method : str, optional
|
||||||
The method to perform aggregation using. Check valid options in
|
The method to perform aggregation using.
|
||||||
:func:`.get_aggfunc_by_name` (default "mean").
|
Check valid options in :func:`.get_aggfunc_by_name`
|
||||||
|
(default "mean").
|
||||||
agg_method_params : dict, optional
|
agg_method_params : dict, optional
|
||||||
Parameters to pass to the aggregation function. Check valid options in
|
Parameters to pass to the aggregation function.
|
||||||
:func:`.get_aggfunc_by_name` (default None).
|
Check valid options in :func:`.get_aggfunc_by_name`
|
||||||
cor_method : str, optional
|
(default None).
|
||||||
The method to perform correlation using. Check valid options in
|
conn_method : str, optional
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure`
|
The method to perform connectivity measure using.
|
||||||
(default "covariance").
|
Check valid options in :class:`.JuniferConnectivityMeasure`
|
||||||
cor_method_params : dict, optional
|
(default "correlation").
|
||||||
Parameters to pass to the correlation function. Check valid options in
|
conn_method_params : dict, optional
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
|
Parameters to pass to :class:`.JuniferConnectivityMeasure`.
|
||||||
|
If None, ``{"empirical": True}`` will be used, which would mean
|
||||||
|
:class:`sklearn.covariance.EmpiricalCovariance` is used to compute
|
||||||
|
covariance. If usage of :class:`sklearn.covariance.LedoitWolf` is
|
||||||
|
desired, ``{"empirical": False}`` should be passed
|
||||||
|
(default None).
|
||||||
masks : str, dict or list of dict or str, optional
|
masks : str, dict or list of dict or str, optional
|
||||||
The specification of the masks to apply to regions before extracting
|
The specification of the masks to apply to regions before extracting
|
||||||
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
||||||
If None, will not apply any mask (default None).
|
If None, will not apply any mask (default None).
|
||||||
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 ``BOLD_<class_name>``
|
||||||
None).
|
(default None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -57,19 +63,18 @@ class FunctionalConnectivityBase(BaseMarker):
|
||||||
self,
|
self,
|
||||||
agg_method: str = "mean",
|
agg_method: str = "mean",
|
||||||
agg_method_params: Optional[Dict] = None,
|
agg_method_params: Optional[Dict] = None,
|
||||||
cor_method: str = "covariance",
|
conn_method: str = "correlation",
|
||||||
cor_method_params: Optional[Dict] = None,
|
conn_method_params: Optional[Dict] = None,
|
||||||
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
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.conn_method = conn_method
|
||||||
self.cor_method_params = cor_method_params or {}
|
self.conn_method_params = conn_method_params or {}
|
||||||
|
# Reverse of nilearn behavior
|
||||||
# default to nilearn behavior
|
self.conn_method_params["empirical"] = self.conn_method_params.get(
|
||||||
self.cor_method_params["empirical"] = self.cor_method_params.get(
|
"empirical", True
|
||||||
"empirical", False
|
|
||||||
)
|
)
|
||||||
self.masks = masks
|
self.masks = masks
|
||||||
super().__init__(on="BOLD", name=name)
|
super().__init__(on="BOLD", name=name)
|
||||||
|
|
@ -121,14 +126,21 @@ class FunctionalConnectivityBase(BaseMarker):
|
||||||
"""
|
"""
|
||||||
# Perform necessary aggregation
|
# Perform necessary aggregation
|
||||||
aggregation = self.aggregate(input, extra_input=extra_input)
|
aggregation = self.aggregate(input, extra_input=extra_input)
|
||||||
# Compute correlation
|
# Set covariance estimator
|
||||||
if self.cor_method_params["empirical"]:
|
if self.conn_method_params["empirical"]:
|
||||||
connectivity = ConnectivityMeasure(
|
cov_estimator = EmpiricalCovariance(store_precision=False)
|
||||||
cov_estimator=EmpiricalCovariance(), # type: ignore
|
|
||||||
kind=self.cor_method,
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
connectivity = ConnectivityMeasure(kind=self.cor_method)
|
cov_estimator = LedoitWolf(store_precision=False)
|
||||||
|
# Compute correlation
|
||||||
|
connectivity = JuniferConnectivityMeasure(
|
||||||
|
cov_estimator=cov_estimator,
|
||||||
|
kind=self.conn_method,
|
||||||
|
**{
|
||||||
|
k: v
|
||||||
|
for k, v in self.conn_method_params.items()
|
||||||
|
if k != "empirical"
|
||||||
|
},
|
||||||
|
)
|
||||||
# Create dictionary for output
|
# Create dictionary for output
|
||||||
return {
|
return {
|
||||||
"functional_connectivity": {
|
"functional_connectivity": {
|
||||||
|
|
|
||||||
|
|
@ -22,28 +22,34 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
parcellation : str or list of str
|
parcellation : str or list of str
|
||||||
The name(s) of the parcellation(s). Check valid options by calling
|
The name(s) of the parcellation(s) to use.
|
||||||
:func:`.list_parcellations`.
|
See :func:`.list_parcellations` for options.
|
||||||
agg_method : str, optional
|
agg_method : str, optional
|
||||||
The method to perform aggregation using. Check valid options in
|
The method to perform aggregation using.
|
||||||
:func:`.get_aggfunc_by_name` (default "mean").
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
|
(default "mean").
|
||||||
agg_method_params : dict, optional
|
agg_method_params : dict, optional
|
||||||
Parameters to pass to the aggregation function. Check valid options in
|
Parameters to pass to the aggregation function.
|
||||||
:func:`.get_aggfunc_by_name` (default None).
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
cor_method : str, optional
|
(default None).
|
||||||
The method to perform correlation using. Check valid options in
|
conn_method : str, optional
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure`
|
The method to perform connectivity measure using.
|
||||||
(default "covariance").
|
See :class:`.JuniferConnectivityMeasure` for options
|
||||||
cor_method_params : dict, optional
|
(default "correlation").
|
||||||
Parameters to pass to the correlation function. Check valid options in
|
conn_method_params : dict, optional
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
|
Parameters to pass to :class:`.JuniferConnectivityMeasure`.
|
||||||
|
If None, ``{"empirical": True}`` will be used, which would mean
|
||||||
|
:class:`sklearn.covariance.EmpiricalCovariance` is used to compute
|
||||||
|
covariance. If usage of :class:`sklearn.covariance.LedoitWolf` is
|
||||||
|
desired, ``{"empirical": False}`` should be passed
|
||||||
|
(default None).
|
||||||
masks : str, dict or list of dict or str, optional
|
masks : str, dict or list of dict or str, optional
|
||||||
The specification of the masks to apply to regions before extracting
|
The specification of the masks to apply to regions before extracting
|
||||||
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
||||||
If None, will not apply any mask (default None).
|
If None, will not apply any mask (default None).
|
||||||
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
|
||||||
None).
|
``BOLD_FunctionalConnectivityParcels`` (default None).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -52,8 +58,8 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
|
||||||
parcellation: Union[str, List[str]],
|
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",
|
conn_method: str = "correlation",
|
||||||
cor_method_params: Optional[Dict] = None,
|
conn_method_params: Optional[Dict] = None,
|
||||||
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -61,8 +67,8 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
agg_method=agg_method,
|
agg_method=agg_method,
|
||||||
agg_method_params=agg_method_params,
|
agg_method_params=agg_method_params,
|
||||||
cor_method=cor_method,
|
conn_method=conn_method,
|
||||||
cor_method_params=cor_method_params,
|
conn_method_params=conn_method_params,
|
||||||
masks=masks,
|
masks=masks,
|
||||||
name=name,
|
name=name,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -23,35 +23,42 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
coords : str
|
coords : str
|
||||||
The name of the coordinates list to use. See
|
The name of the coordinates list to use.
|
||||||
:func:`.list_coordinates` for options.
|
See :func:`.list_coordinates` for options.
|
||||||
radius : float, optional
|
radius : positive float, optional
|
||||||
The radius of the sphere in mm. If None, the signal will be extracted
|
The radius of the sphere around each coordinates in millimetres.
|
||||||
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
|
If None, the signal will be extracted from a single voxel.
|
||||||
for more information (default None).
|
See :class:`.JuniferNiftiSpheresMasker` for more information
|
||||||
|
(default None).
|
||||||
allow_overlap : bool, optional
|
allow_overlap : bool, optional
|
||||||
Whether to allow overlapping spheres. If False, an error is raised if
|
Whether to allow overlapping spheres. If False, an error is raised if
|
||||||
the spheres overlap (default is False).
|
the spheres overlap (default False).
|
||||||
agg_method : str, optional
|
agg_method : str, optional
|
||||||
The aggregation method to use.
|
The method to perform aggregation using.
|
||||||
See :func:`.get_aggfunc_by_name` for more information
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
(default None).
|
(default "mean").
|
||||||
agg_method_params : dict, optional
|
agg_method_params : dict, optional
|
||||||
The parameters to pass to the aggregation method (default None).
|
Parameters to pass to the aggregation function.
|
||||||
cor_method : str, optional
|
See :func:`.get_aggfunc_by_name` for options
|
||||||
The method to perform correlation using. Check valid options in
|
(default None).
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default "covariance").
|
conn_method : str, optional
|
||||||
cor_method_params : dict, optional
|
The method to perform connectivity measure using.
|
||||||
Parameters to pass to the correlation function. Check valid options in
|
See :class:`.JuniferConnectivityMeasure` for options
|
||||||
:class:`nilearn.connectome.ConnectivityMeasure` (default None).
|
(default "correlation").
|
||||||
|
conn_method_params : dict, optional
|
||||||
|
Parameters to pass to :class:`.JuniferConnectivityMeasure`.
|
||||||
|
If None, ``{"empirical": True}`` will be used, which would mean
|
||||||
|
:class:`sklearn.covariance.EmpiricalCovariance` is used to compute
|
||||||
|
covariance. If usage of :class:`sklearn.covariance.LedoitWolf` is
|
||||||
|
desired, ``{"empirical": False}`` should be passed
|
||||||
|
(default None).
|
||||||
masks : str, dict or list of dict or str, optional
|
masks : str, dict or list of dict or str, optional
|
||||||
The specification of the masks to apply to regions before extracting
|
The specification of the masks to apply to regions before extracting
|
||||||
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
signals. Check :ref:`Using Masks <using_masks>` for more details.
|
||||||
If None, will not apply any mask (default None).
|
If None, will not apply any mask (default None).
|
||||||
name : str, optional
|
name : str, optional
|
||||||
The name of the marker. By default, it will use
|
The name of the marker. If None, will use
|
||||||
KIND_FunctionalConnectivitySpheres where KIND is the kind of data it
|
``BOLD_FunctionalConnectivitySpheres`` (default None).
|
||||||
was applied to (default None).
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
@ -62,8 +69,8 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
|
||||||
allow_overlap: bool = False,
|
allow_overlap: bool = False,
|
||||||
agg_method: str = "mean",
|
agg_method: str = "mean",
|
||||||
agg_method_params: Optional[Dict] = None,
|
agg_method_params: Optional[Dict] = None,
|
||||||
cor_method: str = "covariance",
|
conn_method: str = "correlation",
|
||||||
cor_method_params: Optional[Dict] = None,
|
conn_method_params: Optional[Dict] = None,
|
||||||
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -75,8 +82,8 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
agg_method=agg_method,
|
agg_method=agg_method,
|
||||||
agg_method_params=agg_method_params,
|
agg_method_params=agg_method_params,
|
||||||
cor_method=cor_method,
|
conn_method=conn_method,
|
||||||
cor_method_params=cor_method_params,
|
conn_method_params=conn_method_params,
|
||||||
masks=masks,
|
masks=masks,
|
||||||
name=name,
|
name=name,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -27,7 +27,7 @@ def test_init() -> None:
|
||||||
CrossParcellationFC(
|
CrossParcellationFC(
|
||||||
parcellation_one="a",
|
parcellation_one="a",
|
||||||
parcellation_two="a",
|
parcellation_two="a",
|
||||||
correlation_method="pearson",
|
corr_method="pearson",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -58,7 +58,7 @@ def test_compute(tmp_path: Path) -> None:
|
||||||
crossparcellation = CrossParcellationFC(
|
crossparcellation = CrossParcellationFC(
|
||||||
parcellation_one=parcellation_one,
|
parcellation_one=parcellation_one,
|
||||||
parcellation_two=parcellation_two,
|
parcellation_two=parcellation_two,
|
||||||
correlation_method="spearman",
|
corr_method="spearman",
|
||||||
)
|
)
|
||||||
out = crossparcellation.compute(element_data["BOLD"])[
|
out = crossparcellation.compute(element_data["BOLD"])[
|
||||||
"functional_connectivity"
|
"functional_connectivity"
|
||||||
|
|
@ -86,7 +86,7 @@ def test_store(tmp_path: Path) -> None:
|
||||||
crossparcellation = CrossParcellationFC(
|
crossparcellation = CrossParcellationFC(
|
||||||
parcellation_one=parcellation_one,
|
parcellation_one=parcellation_one,
|
||||||
parcellation_two=parcellation_two,
|
parcellation_two=parcellation_two,
|
||||||
correlation_method="spearman",
|
corr_method="spearman",
|
||||||
)
|
)
|
||||||
storage = SQLiteFeatureStorage(
|
storage = SQLiteFeatureStorage(
|
||||||
uri=tmp_path / "test_crossparcellation.sqlite", upsert="ignore"
|
uri=tmp_path / "test_crossparcellation.sqlite", upsert="ignore"
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,9 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Dict
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
from junifer.datareader import DefaultDataReader
|
from junifer.datareader import DefaultDataReader
|
||||||
from junifer.markers.functional_connectivity import EdgeCentricFCParcels
|
from junifer.markers.functional_connectivity import EdgeCentricFCParcels
|
||||||
|
|
@ -12,20 +15,35 @@ from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
|
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
|
||||||
|
|
||||||
|
|
||||||
def test_EdgeCentricFCParcels(tmp_path: Path) -> None:
|
@pytest.mark.parametrize(
|
||||||
|
"conn_method_params",
|
||||||
|
[
|
||||||
|
{"empirical": False},
|
||||||
|
{"empirical": True},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_EdgeCentricFCParcels(
|
||||||
|
tmp_path: Path,
|
||||||
|
conn_method_params: Dict[str, bool],
|
||||||
|
) -> None:
|
||||||
"""Test EdgeCentricFCParcels.
|
"""Test EdgeCentricFCParcels.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
tmp_path : pathlib.Path
|
tmp_path : pathlib.Path
|
||||||
The path to the test directory.
|
The path to the test directory.
|
||||||
|
conn_method_params : dict
|
||||||
|
The parametrized parameters to connectivity measure method.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
with PartlyCloudyTestingDataGrabber() as dg:
|
with PartlyCloudyTestingDataGrabber() as dg:
|
||||||
|
# Get element data
|
||||||
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
||||||
|
# Setup marker
|
||||||
marker = EdgeCentricFCParcels(
|
marker = EdgeCentricFCParcels(
|
||||||
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
|
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
|
||||||
cor_method_params={"empirical": True},
|
conn_method="correlation",
|
||||||
|
conn_method_params=conn_method_params,
|
||||||
)
|
)
|
||||||
# Check correct output
|
# Check correct output
|
||||||
assert "matrix" == marker.get_output_type(
|
assert "matrix" == marker.get_output_type(
|
||||||
|
|
@ -41,8 +59,7 @@ def test_EdgeCentricFCParcels(tmp_path: Path) -> None:
|
||||||
assert "data" in edge_fc_bold
|
assert "data" in edge_fc_bold
|
||||||
assert "row_names" in edge_fc_bold
|
assert "row_names" in edge_fc_bold
|
||||||
assert "col_names" in edge_fc_bold
|
assert "col_names" in edge_fc_bold
|
||||||
assert edge_fc_bold["data"].shape[0] == n_edges
|
assert edge_fc_bold["data"].shape == (n_edges, n_edges)
|
||||||
assert edge_fc_bold["data"].shape[1] == n_edges
|
|
||||||
assert len(set(edge_fc_bold["row_names"])) == n_edges
|
assert len(set(edge_fc_bold["row_names"])) == n_edges
|
||||||
assert len(set(edge_fc_bold["col_names"])) == n_edges
|
assert len(set(edge_fc_bold["col_names"])) == n_edges
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,9 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Dict
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
from junifer.datareader import DefaultDataReader
|
from junifer.datareader import DefaultDataReader
|
||||||
from junifer.markers.functional_connectivity import EdgeCentricFCSpheres
|
from junifer.markers.functional_connectivity import EdgeCentricFCSpheres
|
||||||
|
|
@ -12,19 +15,36 @@ from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
|
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
|
||||||
|
|
||||||
|
|
||||||
def test_EdgeCentricFCSpheres(tmp_path: Path) -> None:
|
@pytest.mark.parametrize(
|
||||||
|
"conn_method_params",
|
||||||
|
[
|
||||||
|
{"empirical": False},
|
||||||
|
{"empirical": True},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_EdgeCentricFCSpheres(
|
||||||
|
tmp_path: Path,
|
||||||
|
conn_method_params: Dict[str, bool],
|
||||||
|
) -> None:
|
||||||
"""Test EdgeCentricFCSpheres.
|
"""Test EdgeCentricFCSpheres.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
tmp_path : pathlib.Path
|
tmp_path : pathlib.Path
|
||||||
The path to the test directory.
|
The path to the test directory.
|
||||||
|
conn_method_params : dict
|
||||||
|
The parametrized parameters to connectivity measure method.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
with SPMAuditoryTestingDataGrabber() as dg:
|
with SPMAuditoryTestingDataGrabber() as dg:
|
||||||
|
# Get element data
|
||||||
element_data = DefaultDataReader().fit_transform(dg["sub001"])
|
element_data = DefaultDataReader().fit_transform(dg["sub001"])
|
||||||
|
# Setup marker
|
||||||
marker = EdgeCentricFCSpheres(
|
marker = EdgeCentricFCSpheres(
|
||||||
coords="DMNBuckner", radius=5.0, cor_method="correlation"
|
coords="DMNBuckner",
|
||||||
|
radius=5.0,
|
||||||
|
conn_method="correlation",
|
||||||
|
conn_method_params=conn_method_params,
|
||||||
)
|
)
|
||||||
# Check correct output
|
# Check correct output
|
||||||
assert "matrix" == marker.get_output_type(
|
assert "matrix" == marker.get_output_type(
|
||||||
|
|
@ -45,13 +65,6 @@ def test_EdgeCentricFCSpheres(tmp_path: Path) -> None:
|
||||||
assert len(set(edge_fc_bold["row_names"])) == n_edges
|
assert len(set(edge_fc_bold["row_names"])) == n_edges
|
||||||
assert len(set(edge_fc_bold["col_names"])) == n_edges
|
assert len(set(edge_fc_bold["col_names"])) == n_edges
|
||||||
|
|
||||||
# Check empirical correlation method parameters
|
|
||||||
marker = EdgeCentricFCSpheres(
|
|
||||||
coords="DMNBuckner",
|
|
||||||
radius=5.0,
|
|
||||||
cor_method="correlation",
|
|
||||||
cor_method_params={"empirical": True},
|
|
||||||
)
|
|
||||||
# Store
|
# Store
|
||||||
storage = SQLiteFeatureStorage(
|
storage = SQLiteFeatureStorage(
|
||||||
uri=tmp_path / "test_edge_fc_spheres.sqlite", upsert="ignore"
|
uri=tmp_path / "test_edge_fc_spheres.sqlite", upsert="ignore"
|
||||||
|
|
|
||||||
|
|
@ -6,10 +6,13 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Dict, Type
|
||||||
|
|
||||||
|
import pytest
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
from nilearn.maskers import NiftiLabelsMasker
|
from nilearn.maskers import NiftiLabelsMasker
|
||||||
from numpy.testing import assert_array_almost_equal
|
from numpy.testing import assert_array_almost_equal
|
||||||
|
from sklearn.covariance import EmpiricalCovariance, LedoitWolf
|
||||||
|
|
||||||
from junifer.data import get_parcellation
|
from junifer.data import get_parcellation
|
||||||
from junifer.datareader import DefaultDataReader
|
from junifer.datareader import DefaultDataReader
|
||||||
|
|
@ -20,19 +23,42 @@ from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
|
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
|
||||||
|
|
||||||
|
|
||||||
def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
|
if TYPE_CHECKING:
|
||||||
|
from sklearn.base import BaseEstimator
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"conn_method_params, cov_estimator",
|
||||||
|
[
|
||||||
|
({"empirical": False}, LedoitWolf(store_precision=False)),
|
||||||
|
({"empirical": True}, EmpiricalCovariance(store_precision=False)),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_FunctionalConnectivityParcels(
|
||||||
|
tmp_path: Path,
|
||||||
|
conn_method_params: Dict[str, bool],
|
||||||
|
cov_estimator: Type["BaseEstimator"],
|
||||||
|
) -> None:
|
||||||
"""Test FunctionalConnectivityParcels.
|
"""Test FunctionalConnectivityParcels.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
tmp_path : pathlib.Path
|
tmp_path : pathlib.Path
|
||||||
The path to the test directory.
|
The path to the test directory.
|
||||||
|
conn_method_params : dict
|
||||||
|
The parametrized parameters to connectivity measure method.
|
||||||
|
cov_estimator : estimator object
|
||||||
|
The parametrized covariance estimator.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
with PartlyCloudyTestingDataGrabber() as dg:
|
with PartlyCloudyTestingDataGrabber() as dg:
|
||||||
|
# Get element data
|
||||||
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
||||||
|
# Setup marker
|
||||||
marker = FunctionalConnectivityParcels(
|
marker = FunctionalConnectivityParcels(
|
||||||
parcellation="TianxS1x3TxMNInonlinear2009cAsym"
|
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
|
||||||
|
conn_method="correlation",
|
||||||
|
conn_method_params=conn_method_params,
|
||||||
)
|
)
|
||||||
# Check correct output
|
# Check correct output
|
||||||
assert "matrix" == marker.get_output_type(
|
assert "matrix" == marker.get_output_type(
|
||||||
|
|
@ -65,7 +91,7 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
|
||||||
)
|
)
|
||||||
# Compute the connectivity measure
|
# Compute the connectivity measure
|
||||||
connectivity_measure = ConnectivityMeasure(
|
connectivity_measure = ConnectivityMeasure(
|
||||||
kind="covariance"
|
cov_estimator=cov_estimator, kind="correlation" # type: ignore
|
||||||
).fit_transform([extracted_timeseries])[0]
|
).fit_transform([extracted_timeseries])[0]
|
||||||
|
|
||||||
# Check that FC are almost equal
|
# Check that FC are almost equal
|
||||||
|
|
@ -73,11 +99,6 @@ def test_FunctionalConnectivityParcels(tmp_path: Path) -> None:
|
||||||
connectivity_measure, fc_bold["data"], decimal=3
|
connectivity_measure, fc_bold["data"], decimal=3
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check empirical correlation method parameters
|
|
||||||
marker = FunctionalConnectivityParcels(
|
|
||||||
parcellation="TianxS1x3TxMNInonlinear2009cAsym",
|
|
||||||
cor_method_params={"empirical": True},
|
|
||||||
)
|
|
||||||
# Store
|
# Store
|
||||||
storage = SQLiteFeatureStorage(
|
storage = SQLiteFeatureStorage(
|
||||||
uri=tmp_path / "test_fc_parcels.sqlite", upsert="ignore"
|
uri=tmp_path / "test_fc_parcels.sqlite", upsert="ignore"
|
||||||
|
|
|
||||||
|
|
@ -7,12 +7,13 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Dict, Type
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from nilearn.connectome import ConnectivityMeasure
|
from nilearn.connectome import ConnectivityMeasure
|
||||||
from nilearn.maskers import NiftiSpheresMasker
|
from nilearn.maskers import NiftiSpheresMasker
|
||||||
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, LedoitWolf
|
||||||
|
|
||||||
from junifer.data import get_coordinates
|
from junifer.data import get_coordinates
|
||||||
from junifer.datareader import DefaultDataReader
|
from junifer.datareader import DefaultDataReader
|
||||||
|
|
@ -23,19 +24,43 @@ from junifer.storage import SQLiteFeatureStorage
|
||||||
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
|
from junifer.testing.datagrabbers import SPMAuditoryTestingDataGrabber
|
||||||
|
|
||||||
|
|
||||||
def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
if TYPE_CHECKING:
|
||||||
|
from sklearn.base import BaseEstimator
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"conn_method_params, cov_estimator",
|
||||||
|
[
|
||||||
|
({"empirical": False}, LedoitWolf(store_precision=False)),
|
||||||
|
({"empirical": True}, EmpiricalCovariance(store_precision=False)),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_FunctionalConnectivitySpheres(
|
||||||
|
tmp_path: Path,
|
||||||
|
conn_method_params: Dict[str, bool],
|
||||||
|
cov_estimator: Type["BaseEstimator"],
|
||||||
|
) -> None:
|
||||||
"""Test FunctionalConnectivitySpheres.
|
"""Test FunctionalConnectivitySpheres.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
tmp_path : pathlib.Path
|
tmp_path : pathlib.Path
|
||||||
The path to the test directory.
|
The path to the test directory.
|
||||||
|
conn_method_params : dict
|
||||||
|
The parametrized parameters to connectivity measure method.
|
||||||
|
cov_estimator : estimator object
|
||||||
|
The parametrized covariance estimator.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
with SPMAuditoryTestingDataGrabber() as dg:
|
with SPMAuditoryTestingDataGrabber() as dg:
|
||||||
|
# Get element data
|
||||||
element_data = DefaultDataReader().fit_transform(dg["sub001"])
|
element_data = DefaultDataReader().fit_transform(dg["sub001"])
|
||||||
|
# Setup marker
|
||||||
marker = FunctionalConnectivitySpheres(
|
marker = FunctionalConnectivitySpheres(
|
||||||
coords="DMNBuckner", radius=5.0, cor_method="correlation"
|
coords="DMNBuckner",
|
||||||
|
radius=5.0,
|
||||||
|
conn_method="correlation",
|
||||||
|
conn_method_params=conn_method_params,
|
||||||
)
|
)
|
||||||
# Check correct output
|
# Check correct output
|
||||||
assert "matrix" == marker.get_output_type(
|
assert "matrix" == marker.get_output_type(
|
||||||
|
|
@ -67,7 +92,7 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
||||||
)
|
)
|
||||||
# Compute the connectivity measure
|
# Compute the connectivity measure
|
||||||
connectivity_measure = ConnectivityMeasure(
|
connectivity_measure = ConnectivityMeasure(
|
||||||
kind="correlation"
|
cov_estimator=cov_estimator, kind="correlation" # type: ignore
|
||||||
).fit_transform([extracted_timeseries])[0]
|
).fit_transform([extracted_timeseries])[0]
|
||||||
|
|
||||||
# Check that FC are almost equal
|
# Check that FC are almost equal
|
||||||
|
|
@ -88,65 +113,9 @@ def test_FunctionalConnectivitySpheres(tmp_path: Path) -> None:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_FunctionalConnectivitySpheres_empirical(tmp_path: Path) -> None:
|
|
||||||
"""Test FunctionalConnectivitySpheres with empirical covariance.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
tmp_path : pathlib.Path
|
|
||||||
The path to the test directory.
|
|
||||||
|
|
||||||
"""
|
|
||||||
with SPMAuditoryTestingDataGrabber() as dg:
|
|
||||||
element_data = DefaultDataReader().fit_transform(dg["sub001"])
|
|
||||||
marker = FunctionalConnectivitySpheres(
|
|
||||||
coords="DMNBuckner",
|
|
||||||
radius=5.0,
|
|
||||||
cor_method="correlation",
|
|
||||||
cor_method_params={"empirical": True},
|
|
||||||
)
|
|
||||||
# Check correct output
|
|
||||||
assert "matrix" == marker.get_output_type(
|
|
||||||
input_type="BOLD", output_feature="functional_connectivity"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Fit-transform the data
|
|
||||||
fc = marker.fit_transform(element_data)
|
|
||||||
fc_bold = fc["BOLD"]["functional_connectivity"]
|
|
||||||
|
|
||||||
assert "data" in fc_bold
|
|
||||||
assert "row_names" in fc_bold
|
|
||||||
assert "col_names" in fc_bold
|
|
||||||
assert fc_bold["data"].shape == (6, 6)
|
|
||||||
assert len(set(fc_bold["row_names"])) == 6
|
|
||||||
assert len(set(fc_bold["col_names"])) == 6
|
|
||||||
|
|
||||||
# Compare with nilearn
|
|
||||||
# Load testing coordinates for the target data
|
|
||||||
testing_coords, _ = get_coordinates(
|
|
||||||
coords="DMNBuckner", target_data=element_data["BOLD"]
|
|
||||||
)
|
|
||||||
# Extract timeseries
|
|
||||||
nifti_spheres_masker = NiftiSpheresMasker(
|
|
||||||
seeds=testing_coords, radius=5.0
|
|
||||||
)
|
|
||||||
extracted_timeseries = nifti_spheres_masker.fit_transform(
|
|
||||||
element_data["BOLD"]["data"]
|
|
||||||
)
|
|
||||||
# Compute the connectivity measure
|
|
||||||
connectivity_measure = ConnectivityMeasure(
|
|
||||||
cov_estimator=EmpiricalCovariance(), kind="correlation" # type: ignore
|
|
||||||
).fit_transform([extracted_timeseries])[0]
|
|
||||||
|
|
||||||
# Check that FC are almost equal
|
|
||||||
assert_array_almost_equal(
|
|
||||||
connectivity_measure, fc_bold["data"], decimal=3
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_FunctionalConnectivitySpheres_error() -> None:
|
def test_FunctionalConnectivitySpheres_error() -> None:
|
||||||
"""Test FunctionalConnectivitySpheres errors."""
|
"""Test FunctionalConnectivitySpheres errors."""
|
||||||
with pytest.raises(ValueError, match="radius should be > 0"):
|
with pytest.raises(ValueError, match="radius should be > 0"):
|
||||||
FunctionalConnectivitySpheres(
|
FunctionalConnectivitySpheres(
|
||||||
coords="DMNBuckner", radius=-0.1, cor_method="correlation"
|
coords="DMNBuckner", radius=-0.1, conn_method="correlation"
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue