[ENH]: Support "xi" correlation metric for functional connectivity #333

Merged
synchon merged 17 commits from feat/xi-correlation-fc into main 2025-04-01 12:02:43 +00:00
4 changed files with 85 additions and 8 deletions

View file

@ -0,0 +1 @@
Add support for Chatterjee's xi correlation in :class:`.JuniferConnectivityMeasure` by `Synchon Mandal`_

View file

@ -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"):

View file

@ -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)

View file

@ -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"
),
}, },
} }