[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>
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
import sys
|
||||||
|
from itertools import product
|
||||||
from typing import Callable, Optional
|
from typing import Callable, Optional
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
@ -322,11 +324,13 @@ class JuniferConnectivityMeasure(ConnectivityMeasure):
|
||||||
The covariance estimator
|
The covariance estimator
|
||||||
(default ``EmpiricalCovariance(store_precision=False)``).
|
(default ``EmpiricalCovariance(store_precision=False)``).
|
||||||
kind : {"covariance", "correlation", "spearman correlation", \
|
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.
|
The matrix kind. The default value uses Pearson's correlation.
|
||||||
If ``"spearman correlation"`` is used, the data will be ranked before
|
If ``"spearman correlation"`` is used, the data will be ranked before
|
||||||
estimating the covariance. For the use of ``"tangent"`` see [1]_
|
estimating the covariance. For ``"xi correlation"``, the coefficient
|
||||||
(default "correlation").
|
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
|
vectorize : bool, optional
|
||||||
If True, connectivity matrices are reshaped into 1D arrays and only
|
If True, connectivity matrices are reshaped into 1D arrays and only
|
||||||
their flattened lower triangular parts are returned (default False).
|
their flattened lower triangular parts are returned (default False).
|
||||||
|
|
@ -372,6 +376,12 @@ class JuniferConnectivityMeasure(ConnectivityMeasure):
|
||||||
Springer.
|
Springer.
|
||||||
doi:10/cn2h9c.
|
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__(
|
def __init__(
|
||||||
|
|
@ -420,6 +430,25 @@ class JuniferConnectivityMeasure(ConnectivityMeasure):
|
||||||
covariances_std.append(self.cov_estimator_.fit(x).covariance_)
|
covariances_std.append(self.cov_estimator_.fit(x).covariance_)
|
||||||
|
|
||||||
connectivities = [cov_to_corr(cov) for cov in covariances_std]
|
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:
|
else:
|
||||||
covariances = [self.cov_estimator_.fit(x).covariance_ for x in X]
|
covariances = [self.cov_estimator_.fit(x).covariance_ for x in X]
|
||||||
if self.kind in ("covariance", "tangent"):
|
if self.kind in ("covariance", "tangent"):
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
import copy
|
import copy
|
||||||
|
import sys
|
||||||
import warnings
|
import warnings
|
||||||
from math import cosh, exp, log, sinh, sqrt
|
from math import cosh, exp, log, sinh, sqrt
|
||||||
from typing import TYPE_CHECKING, Optional, Union
|
from typing import TYPE_CHECKING, Optional, Union
|
||||||
|
|
@ -12,7 +13,11 @@ import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
from nilearn.connectome.connectivity_matrices import sym_matrix_to_vec
|
from nilearn.connectome.connectivity_matrices import sym_matrix_to_vec
|
||||||
from nilearn.tests.test_signal import generate_signals
|
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 pandas import DataFrame
|
||||||
from scipy import linalg
|
from scipy import linalg
|
||||||
from sklearn.covariance import EmpiricalCovariance, LedoitWolf
|
from sklearn.covariance import EmpiricalCovariance, LedoitWolf
|
||||||
|
|
@ -1088,3 +1093,42 @@ def test_connectivity_measure_standardize(
|
||||||
).fit_transform(signals)
|
).fit_transform(signals)
|
||||||
for m in record:
|
for m in record:
|
||||||
assert match not in m.message
|
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
|
# Create dictionary for output
|
||||||
|
labels = aggregation["aggregation"]["col_names"]
|
||||||
return {
|
return {
|
||||||
"functional_connectivity": {
|
"functional_connectivity": {
|
||||||
"data": connectivity.fit_transform(
|
"data": connectivity.fit_transform(
|
||||||
[aggregation["aggregation"]["data"]]
|
[aggregation["aggregation"]["data"]]
|
||||||
)[0],
|
)[0],
|
||||||
# Create column names
|
"row_names": labels,
|
||||||
"row_names": aggregation["aggregation"]["col_names"],
|
"col_names": labels,
|
||||||
"col_names": aggregation["aggregation"]["col_names"],
|
# xi correlation coefficient is not symmetric
|
||||||
"matrix_kind": "tril",
|
"matrix_kind": (
|
||||||
|
"full" if self.conn_method == "xi correlation" else "tril"
|
||||||
|
),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue