[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>
# 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"):

View file

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

View file

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