[ENH]: Support "xi" correlation metric for functional connectivity #333
4 changed files with 85 additions and 8 deletions
1
docs/changes/newsfragments/333.enh
Normal file
1
docs/changes/newsfragments/333.enh
Normal file
|
|
@ -0,0 +1 @@
|
|||
Add support for Chatterjee's xi correlation in :class:`.JuniferConnectivityMeasure` by `Synchon Mandal`_
|
||||
|
|
@ -3,6 +3,8 @@
|
|||
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||
# License: AGPL
|
||||
|
||||
import sys
|
||||
from itertools import product
|
||||
from typing import Callable, Optional
|
||||
|
||||
import numpy as np
|
||||
|
|
@ -322,11 +324,13 @@ class JuniferConnectivityMeasure(ConnectivityMeasure):
|
|||
The covariance estimator
|
||||
(default ``EmpiricalCovariance(store_precision=False)``).
|
||||
kind : {"covariance", "correlation", "spearman correlation", \
|
||||
"partial correlation", "tangent", "precision"}, optional
|
||||
"partial correlation", "xi correlation", "tangent", \
|
||||
"precision"}, optional
|
||||
The matrix kind. The default value uses Pearson's correlation.
|
||||
If ``"spearman correlation"`` is used, the data will be ranked before
|
||||
estimating the covariance. For the use of ``"tangent"`` see [1]_
|
||||
(default "correlation").
|
||||
estimating the covariance. For ``"xi correlation"``, the coefficient
|
||||
is not symmetric and should be interpreted as a measure of dependence
|
||||
[2]_ . 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).
|
||||
|
|
@ -372,6 +376,12 @@ class JuniferConnectivityMeasure(ConnectivityMeasure):
|
|||
Springer.
|
||||
doi:10/cn2h9c.
|
||||
|
||||
.. [2] Chatterjee, S.
|
||||
A new coefficient of correlation.
|
||||
Journal of the American Statistical Association 116.536 (2021):
|
||||
2009-2022.
|
||||
doi:10.1080/01621459.2020.1758115.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -420,6 +430,25 @@ class JuniferConnectivityMeasure(ConnectivityMeasure):
|
|||
covariances_std.append(self.cov_estimator_.fit(x).covariance_)
|
||||
|
||||
connectivities = [cov_to_corr(cov) for cov in covariances_std]
|
||||
elif self.kind == "xi correlation":
|
||||
if sys.version_info < (3, 10): # pragma: no cover
|
||||
raise_error(
|
||||
klass=RuntimeError,
|
||||
msg=(
|
||||
"scipy.stats.chatterjeexi is available from "
|
||||
"scipy 1.15.0 and that requires Python 3.10 and above."
|
||||
),
|
||||
)
|
||||
connectivities = []
|
||||
for x in X:
|
||||
n_rois = x.shape[1]
|
||||
connectivity = np.ones((n_rois, n_rois))
|
||||
for i, j in product(range(n_rois), range(n_rois)):
|
||||
if i != j:
|
||||
connectivity[i, j] = stats.chatterjeexi(
|
||||
x[:, i], x[:, j], y_continuous=True
|
||||
).statistic
|
||||
connectivities.append(connectivity)
|
||||
else:
|
||||
covariances = [self.cov_estimator_.fit(x).covariance_ for x in X]
|
||||
if self.kind in ("covariance", "tangent"):
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
# License: AGPL
|
||||
|
||||
import copy
|
||||
import sys
|
||||
import warnings
|
||||
from math import cosh, exp, log, sinh, sqrt
|
||||
from typing import TYPE_CHECKING, Optional, Union
|
||||
|
|
@ -12,7 +13,11 @@ import numpy as np
|
|||
import pytest
|
||||
from nilearn.connectome.connectivity_matrices import sym_matrix_to_vec
|
||||
from nilearn.tests.test_signal import generate_signals
|
||||
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
||||
from numpy.testing import (
|
||||
assert_allclose,
|
||||
assert_array_almost_equal,
|
||||
assert_array_equal,
|
||||
)
|
||||
from pandas import DataFrame
|
||||
from scipy import linalg
|
||||
from sklearn.covariance import EmpiricalCovariance, LedoitWolf
|
||||
|
|
@ -1088,3 +1093,42 @@ def test_connectivity_measure_standardize(
|
|||
).fit_transform(signals)
|
||||
for m in record:
|
||||
assert match not in m.message
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info > (3, 9),
|
||||
reason="will have correct scipy version so no error",
|
||||
)
|
||||
def test_xi_correlation_error() -> None:
|
||||
"""Check xi correlation according to paper."""
|
||||
with pytest.raises(RuntimeError, match="scipy.stats.chatterjeexi"):
|
||||
JuniferConnectivityMeasure(kind="xi correlation").fit_transform(
|
||||
np.zeros((2, 2))
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 10),
|
||||
reason=(
|
||||
"needs scipy 1.15.0 and above which in turn requires "
|
||||
"python 3.10 and above"
|
||||
),
|
||||
)
|
||||
def test_xi_correlation() -> None:
|
||||
"""Check xi correlation according to paper."""
|
||||
rng = np.random.default_rng(25982435982346983)
|
||||
x = rng.random(size=10)
|
||||
y = rng.random(size=10)
|
||||
arr = np.column_stack((x, y))
|
||||
expected = np.array(
|
||||
[
|
||||
[
|
||||
[1.0, -0.3030303],
|
||||
[-0.18181818, 1.0],
|
||||
]
|
||||
]
|
||||
)
|
||||
got = JuniferConnectivityMeasure(kind="xi correlation").fit_transform(
|
||||
[arr]
|
||||
)
|
||||
assert_allclose(expected, got)
|
||||
|
|
|
|||
|
|
@ -143,14 +143,17 @@ class FunctionalConnectivityBase(BaseMarker):
|
|||
},
|
||||
)
|
||||
# Create dictionary for output
|
||||
labels = aggregation["aggregation"]["col_names"]
|
||||
return {
|
||||
"functional_connectivity": {
|
||||
"data": connectivity.fit_transform(
|
||||
[aggregation["aggregation"]["data"]]
|
||||
)[0],
|
||||
# Create column names
|
||||
"row_names": aggregation["aggregation"]["col_names"],
|
||||
"col_names": aggregation["aggregation"]["col_names"],
|
||||
"matrix_kind": "tril",
|
||||
"row_names": labels,
|
||||
"col_names": labels,
|
||||
# xi correlation coefficient is not symmetric
|
||||
"matrix_kind": (
|
||||
"full" if self.conn_method == "xi correlation" else "tril"
|
||||
),
|
||||
},
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue