"""Compute metrics for relating signals in the frequency domain."""
import inspect
import warnings
from collections.abc import Callable, Sequence
from contextvars import ContextVar
from dataclasses import dataclass, replace
from functools import cached_property, wraps
from itertools import combinations
from logging import getLogger
from typing import TYPE_CHECKING, Any, Literal, ParamSpec, TypeVar, cast
import numpy as np
from numpy.typing import DTypeLike, NDArray
from spectral_connectivity._array_utils import (
TIKHONOV_REGULARIZATION_FACTOR,
_batched_inverse_square_root,
_complex_inner_product,
_conjugate_transpose,
_divide_where,
_NumberT,
_regularized_inverse,
_squared_magnitude,
)
from spectral_connectivity._backend import xp
from spectral_connectivity._granger import (
_estimate_all_conditional_spectral_granger,
_estimate_blockwise_spectral_granger,
_estimate_spectral_granger_prediction,
_estimate_subset_spectral_granger_prediction,
_factorize_spectrum,
_var_model_from_factor,
_warn_nan_granger_pairs,
)
from spectral_connectivity._multivariate import (
GLOBAL_COHERENCE_BATCH_CHUNK_ELEMENTS,
_canonical_coherency_components,
_estimate_canonical_coherence,
_global_coherence,
_mic_components,
_normalize_fourier_coefficients,
)
from spectral_connectivity.minimum_phase_decomposition import (
minimum_phase_reconstruction_error as _minimum_phase_reconstruction_error,
)
from spectral_connectivity.statistics import (
JackknifeResult,
adjust_for_multiple_comparisons,
coherence_significance_pvalue,
jackknife_confidence_interval,
)
from spectral_connectivity.utils import (
BackendArray,
is_positive_integer,
mark_readonly_chain_if_supported,
mark_readonly_if_supported,
stacklevel_outside_package,
to_numpy,
)
if TYPE_CHECKING:
from spectral_connectivity.transforms import SpectralTransform
logger = getLogger(__name__)
# Public helpers on Connectivity that are not connectivity measures. Keep this
# definition shared with the high-level wrapper so method discovery and
# jackknife validation cannot drift apart.
_NON_MEASURE_METHODS = frozenset(
{"clear_cache", "jackknife", "minimum_phase_reconstruction_error"}
)
# Measures whose values are magnitudes in [0, 1], so their Fisher (atanh)
# jackknife interval is clamped at 0 on the way back.
_NONNEGATIVE_MAGNITUDE_MEASURES = frozenset({"phase_locking_value", "imaginary_coherence"})
EXPECTATION_AXES = {
"time": (0,),
"trials": (1,),
"tapers": (2,),
"time_trials": (0, 1),
"time_tapers": (0, 2),
"trials_tapers": (1, 2),
"time_trials_tapers": (0, 1, 2),
}
# Per-observation functions of Im(S_ij) averaged by the phase-lag-index family;
# see Connectivity._imaginary_cross_spectrum_moments. They look ``xp`` up when
# called, not at import, so a swapped backend (the device-emulation tests) applies.
_IMAGINARY_MOMENTS: dict[str, Callable[[BackendArray], BackendArray]] = {
"sign": lambda imaginary: xp.sign(imaginary),
"imaginary": lambda imaginary: imaginary,
"absolute": lambda imaginary: xp.abs(imaginary),
"squared": lambda imaginary: imaginary**2,
}
# Im(X_j conj(X_i)) = -Im(X_i conj(X_j)), exactly in floating point too, so each
# moment is exactly antisymmetric (-1) or symmetric (+1) across the signal pair.
_IMAGINARY_MOMENT_PAIR_SYMMETRY = {"sign": -1, "imaginary": -1, "absolute": 1, "squared": 1}
[docs]
@dataclass(frozen=True)
class MultivariateConnectivityResult:
"""Component-resolved multivariate connectivity and spatial projections.
``scores`` has shape ``(..., frequency, connection, component)``. Filters
and patterns, when present, append ``(side, signal)`` where side 0 is the
first group and side 1 is the second. Entries for signals outside a side's
group are NaN. ``connections`` contains the corresponding group-label pair
for each connection and ``group_membership`` has shape ``(group, signal)``.
Attributes
----------
method : str
Name of the measure that produced the result.
scores : NDArray[number], shape (..., frequency, connection, component)
Per-component connectivity. Complex for ``canonical_coherency``
(magnitude times ``exp(-1j * phi)``), real for MIC. A component a
connection cannot supply (its smaller group has fewer channels than the
requested ``n_components``) is NaN.
connections : NDArray, shape (connection, 2)
The ``(first_group_label, second_group_label)`` pair for each connection.
group_labels : NDArray, shape (group,)
Sorted unique group labels.
group_membership : NDArray[bool], shape (group, signal)
``True`` where a signal belongs to a group.
filters : NDArray[floating] or None, shape (..., frequency, connection, component, side, signal)
Spatial filters mapping channel data to each component; NaN outside a
side's group.
patterns : NDArray[floating] or None, same shape as ``filters``
Haufe-style patterns (``within-group real CSD @ filter``) mapping each
component back to channel space.
"""
method: str
scores: NDArray[np.number]
connections: NDArray[Any]
group_labels: NDArray[Any]
group_membership: NDArray[np.bool_]
filters: NDArray[np.floating] | None = None
patterns: NDArray[np.floating] | None = None
def _validated_regularization(value: Any) -> float:
"""Return a finite non-negative scalar regularization factor."""
message = f"regularization must be a finite non-negative scalar, got {value!r}."
if isinstance(value, bool) or not isinstance(value, (int, float, np.integer, np.floating)):
raise ValueError(message)
regularization = float(value)
if not np.isfinite(regularization) or regularization < 0:
raise ValueError(message)
return regularization
def _validated_rank(rank: int | None) -> int | None:
"""Return a positive-integer rank or None, rejecting other values."""
if rank is not None and not is_positive_integer(rank):
msg = f"rank must be a positive integer or None, got {rank!r}."
raise ValueError(msg)
return rank
# Peak workspace cap for the phase-lag-index family's observation-level signal
# tiles. The tile loop holds two float buffers of this many elements, plus, one
# at a time, the three boolean comparison temporaries of the same count that
# each tile's sign and NaN reductions make. The final reduced signal-by-signal
# result is unavoidable, but the large trial/taper/time-resolved outer product
# is never materialized in full.
PHASE_LAG_INDEX_MAX_WORKSPACE_ELEMENTS = 16_000_000
# Fewest averaged observations (``Connectivity.n_observations``) at which, without
# observation weights, the first phase-lag measure reduces all four moments from
# one tile pass instead of only the requested ones. A tile holds n_observations
# times as many values as each reduced moment, so with many observations the
# tile formation dominates and two extra reductions are nearly free, while with
# few (e.g. time-resolved single-trial spectra) writing and retaining two extra
# window-resolved moments costs more than re-forming the tile for a later
# measure. Measured on 32 signals, 100 FFT bins, CPU, with windows x
# observations held near 3600: a lone phase_lag_index computing all four was
# 1.49x (3 observations), 1.33x (8), 1.23x (16), 1.15x (32), 1.08x (64) and
# 1.03-1.07x (128-500) the time of computing only its two, while the four
# phase-lag measures together took 0.80x (3) to 0.42x (500) the time. 64 is
# the smallest count at which a lone measure costs at most 10% more and is
# also faster than before one-pass reduction (0.75-0.79 s vs 0.89 s).
PHASE_LAG_ALL_MOMENTS_MIN_OBSERVATIONS = 64
# Element cap for the complex coefficients gathered per chunk of channel pairs
# by ``Connectivity._subset_cross_spectral_matrix``. The gather and its
# observation-major copy are each at most this size (32 MB at complex128), so
# peak workspace stays bounded however many pairs are requested.
SUBSET_CROSS_SPECTRUM_MAX_WORKSPACE_ELEMENTS = 2_000_000
# Machine epsilons of E[|Im S_ij|] relative to sqrt(P_i P_j) below which a pair
# has no phase lag. In-phase signals leave only rounding noise, measured at up
# to about 1 eps; a genuine lag of this size (~4e-15 rad) is not resolvable.
_ZERO_PHASE_LAG_EPSILONS = 16
# Largest relative gap between the diagonal-noise-power denominator that
# ``directed_coherence`` uses and the true power spectral density, tolerated
# before it warns that its diagonal-noise-covariance assumption is violated (see
# the note in that method). A dimension-aware criterion: unlike a pairwise
# correlation threshold, it catches many weakly-but-jointly correlated sources
# whose cross-power still omits a large fraction of the true power. Non-parametric
# estimation of a truly diagonal covariance leaves a small discrepancy from
# finite-sample noise, so the threshold sits above that floor: only a material
# omission (>= 10% of the true power) triggers the warning.
DIRECTED_COHERENCE_DISCREPANCY_TOLERANCE = 0.1
# Signature and return type of a decorated measure, preserved by its decorators.
_P = ParamSpec("_P")
_R = TypeVar("_R")
def _asnumpy(connectivity_measure: Callable[_P, _R]) -> Callable[_P, _R]:
"""Transform cupy array to numpy array.
If cupy is not installed, then return original.
Parameters
----------
connectivity_measure : callable
Connectivity measure function to wrap.
Returns
-------
callable
Wrapped function that converts output to numpy.
"""
@wraps(connectivity_measure)
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
measure = connectivity_measure(*args, **kwargs)
if measure is None:
return measure
# The type checker analyzes the NumPy path, where this is the identity.
return cast(_R, to_numpy(measure))
return wrapper
def _source_first(connectivity_measure: Callable[_P, _R]) -> Callable[_P, _R]:
"""Reorder a measure's ``[..., target, source]`` output to ``[..., source, target]``.
The Wilson-factorized kernels work in the transfer function's native
layout, where row ``i`` collects the inflow to signal ``i``. The public
measures built on them (the spectral Granger family and the transfer
function / MVAR measures) return the transpose so that ``[..., i, j]``
reads ``i -> j`` like the labeled wrapper's ``sel(source=i, target=j)``.
The lead/lag measures are computed source first and do not use it.
Apply it above :func:`_asnumpy`: swapping the host array is a free view,
whereas a swapped device array would force a contiguous device copy on
transfer.
"""
@wraps(connectivity_measure)
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
native = connectivity_measure(*args, **kwargs)
return cast(_R, np.swapaxes(cast(NDArray[Any], native), -1, -2))
return wrapper
[docs]
class DirectedOrientationWarning(UserWarning):
"""A directed measure's orientation differs from spectral_connectivity 2.x.
Emitted by the ``Connectivity`` methods that returned
``[..., target, source]`` in 2.x and now return ``[..., source, target]``,
and by :func:`~spectral_connectivity.multitaper_connectivity` and
:func:`~spectral_connectivity.connectivity_to_xarray` when they compute
pairwise or subset spectral Granger prediction, whose
``sel(source=a, target=b)`` was ``b -> a`` in 2.x (the 2.x wrapper rejected
the transfer-function measures).
Code written for 2.x keeps running but reads the opposite direction. This
warning is temporary; silence it, without hiding other warnings, with
``warnings.filterwarnings("ignore", category=DirectedOrientationWarning)``.
"""
_MIGRATION_GUIDE_URL = (
"https://github.com/Eden-Kramer-Lab/spectral_connectivity/blob/master/"
"CHANGELOG.md#migration-guide"
)
# Every DirectedOrientationWarning message starts with this, so pyproject.toml
# can filter it by message: a category filter would import the package while
# pytest reads its configuration, before coverage starts.
_ORIENTATION_WARNING_PREFIX = "Since spectral_connectivity 3.0, "
_SILENCE_ORIENTATION_WARNING = (
'Silence with warnings.filterwarnings("ignore", '
"category=spectral_connectivity.DirectedOrientationWarning)."
)
# The methods that returned [..., target, source] arrays in 2.x. The Granger
# measures added in 3.0 and the lead/lag measures were never target first.
_ORIENTATION_CHANGED_MEASURES = frozenset(
{
"pairwise_spectral_granger_prediction",
"subset_pairwise_spectral_granger_prediction",
"directed_transfer_function",
"directed_coherence",
"partial_directed_coherence",
"generalized_partial_directed_coherence",
"direct_directed_transfer_function",
}
)
# False while the labeled wrapper or a jackknife replicate runs a measure: the
# wrapper warns about its source/target labels instead, and a jackknife warns
# once, for its full estimate.
_warn_orientation_change: ContextVar[bool] = ContextVar(
"_warn_orientation_change", default=True
)
def _orientation_changed(connectivity_measure: Callable[_P, _R]) -> Callable[_P, _R]:
"""Warn that a measure's array is source first, unlike in 2.x.
Applied to the methods in ``_ORIENTATION_CHANGED_MEASURES``; remove it with
the warning in 3.2.
"""
@wraps(connectivity_measure)
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
result = connectivity_measure(*args, **kwargs)
if _warn_orientation_change.get():
warnings.warn(
f"{_ORIENTATION_WARNING_PREFIX}{connectivity_measure.__name__} "
"returns [..., source, target]: [..., i, j] is i -> j. 2.x "
"returned [..., target, source]; review the indexing and "
"normalization axes of code written for 2.x. See the migration "
f"guide: {_MIGRATION_GUIDE_URL}. {_SILENCE_ORIENTATION_WARNING}",
DirectedOrientationWarning,
stacklevel=stacklevel_outside_package(),
)
return result
return wrapper
def _ignore_nan_propagation_warnings(
connectivity_measure: Callable[_P, _R],
) -> Callable[_P, _R]:
"""Suppress NumPy invalid/divide warnings from expected NaN propagation.
The directed measures (DTF, PDC and relatives) normalize the transfer
function or MVAR coefficients by an inflow/outflow sum. When the Wilson
minimum-phase decomposition fails to converge those inputs are NaN (already
warned about at decomposition time), and the normalization then emits
``invalid value encountered in divide``. Scope the suppression to these
measures rather than silencing NumPy globally; well-conditioned inputs
produce no NaN and are unaffected.
Parameters
----------
connectivity_measure : callable
Connectivity measure method to wrap.
Returns
-------
callable
Wrapped method whose numpy arithmetic runs under a scoped errstate.
"""
@wraps(connectivity_measure)
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
with np.errstate(invalid="ignore", divide="ignore"):
return connectivity_measure(*args, **kwargs)
return wrapper
_ABSENT = object()
def _optional_transform_attribute(transform: Any, name: str, default: Any) -> Any:
"""Read an optional ``SpectralTransform`` attribute, or ``default`` if absent.
Unlike ``getattr(transform, name, default)``, an ``AttributeError`` raised
inside a property the transform does define propagates instead of being
taken for the attribute's absence, which would silently apply the default.
"""
try:
return getattr(transform, name)
except AttributeError:
if inspect.getattr_static(transform, name, _ABSENT) is not _ABSENT:
raise
return default
def _validated_flag(name: str, value: Any) -> bool:
"""Return ``value`` as a bool, or raise unless it is a boolean.
Accepts Python and NumPy bools and 0-d boolean arrays (e.g. a flag computed
as ``xp.all(...)``). ``bool()`` would read a method or the string
``"False"`` as True and ``None`` or 0 as False, silently changing how the
spectrum is treated.
"""
is_boolean_scalar = getattr(value, "ndim", None) == 0 and (
getattr(getattr(value, "dtype", None), "kind", None) == "b"
)
if not (isinstance(value, bool) or is_boolean_scalar):
msg = (
f"{name} must be a bool, got {type(value).__name__} ({value!r}). It "
f"sets how the spectrum is treated, so it is not guessed from a truth "
f"value: pass True or False (for a transform, a bool attribute or "
f"property, not a method)."
)
raise TypeError(msg)
return bool(value)
def _transform_flag(transform: Any, name: str, default: bool) -> bool:
"""Read an optional boolean ``SpectralTransform`` attribute strictly."""
return _validated_flag(
f"transform.{name}", _optional_transform_attribute(transform, name, default)
)
def _require_fft_order(frequencies: Any) -> None:
"""Raise unless two-sided ``frequencies`` are uniformly spaced in FFT order.
Two-sided coefficients are folded onto their first ``n // 2 + 1`` bins as
the non-negative frequencies, so any other layout (e.g. ``rfft`` output
labelled two-sided, or an ``fftshift``-ed spectrum) would silently drop or
mislabel frequencies.
"""
values = to_numpy(frequencies)
# Allow for the coordinate's own rounding (e.g. float32 from netCDF), which
# the grid step inherits, but never less strictly than 1e-9.
precision = values.dtype if np.issubdtype(values.dtype, np.floating) else np.float64
rtol = max(1e-9, 8 * float(np.finfo(precision).eps))
values = values.astype(float)
if values.size == 0:
return
if values.size == 1:
in_order = bool(values[0] == 0.0)
else:
step = values[1] - values[0] if values.size > 2 else abs(values[1])
expected = np.fft.fftfreq(values.size, d=1.0 / (step * values.size))
in_order = bool(step > 0) and np.allclose(
values, expected, rtol=rtol, atol=max(abs(step) * rtol, np.finfo(float).eps)
)
if not in_order:
msg = (
f"frequencies must be uniformly spaced in standard FFT order (zero and "
f"positive bins followed by negative bins, as numpy.fft.fftfreq); got "
f"{values[:3]} ... {values[-2:]}. The first n // 2 + 1 bins of a "
f"two-sided spectrum are taken as its non-negative half, so any other "
f"layout would drop or mislabel frequencies. If the coefficients hold "
f"only non-negative frequencies (e.g. numpy.fft.rfft or a wavelet "
f"transform), mark them one-sided (is_one_sided=True); otherwise order "
f"them as numpy.fft.fft does."
)
raise ValueError(msg)
[docs]
class Connectivity:
"""
Compute functional and directed connectivity measures from spectral data.
This class provides a comprehensive suite of connectivity analysis methods
based on cross-spectral matrices derived from Fourier-transformed time series.
Methods range from basic coherence to advanced Granger causality measures.
Parameters
----------
fourier_coefficients : NDArray[complexfloating], shape (n_time_windows, n_trials, n_tapers, n_frequencies, n_signals)
Complex-valued Fourier coefficients from spectral analysis. Must be
two-sided (positive and negative frequencies) for Granger methods.
Usually obtained from multitaper or other spectral estimation methods.
**Validation**: Must be 5-dimensional with at least 2 signals and
contain only finite values (no NaN/Inf).
expectation_type : {"trials_tapers", "trials", "tapers", "time",
"time_trials", "time_tapers", "time_trials_tapers"},
default="trials_tapers"
Specifies how to average the cross-spectral matrix:
- "trials_tapers": average over trials and tapers (most common)
- "trials": average over trials only (keep taper dimension)
- "tapers": average over tapers only (keep trial dimension)
- "time": average over time windows
- combinations: average over multiple dimensions
frequencies : NDArray[floating], shape (n_frequencies,), optional
Frequency values in Hz corresponding to FFT bins. If None, uses
normalized frequencies.
time : NDArray[floating], shape (n_time_windows,), optional
Time values in seconds for each time window. If None, uses indices.
dtype : np.dtype, default=complex128
Data type for internal computations. Should match input precision.
minimum_phase_tolerance : float, default=1e-8
Relative convergence tolerance for the Wilson minimum-phase
factorization used by the directed measures (spectral Granger, DTF,
PDC, and relatives).
minimum_phase_max_iterations : int, default=500
Maximum Wilson iterations. Near-singular cross-spectral matrices (highly
correlated channels) can need several hundred iterations; if the
directed measures return NaN with a non-convergence warning, increase
this value. The factorization returns early once every sub-spectrum has
converged, so a large ceiling is cheap for well-conditioned data.
is_one_sided : bool, default=False
Whether coefficients contain only non-negative frequencies. One-sided
transforms are returned without FFT half-spectrum slicing or power
doubling and cannot be used by Wilson-factorized directed measures.
observation_weights : ndarray, optional
Finite non-negative weights with shape ``(time, trial, taper,
frequency, 1)``. They are applied to every expectation and shared
across signals. Transform constructors supply these automatically when
smoothing uses a non-uniform kernel or masks invalid edge estimates.
observations_are_independent : bool, default=True
Whether the trial/taper observations are statistically independent.
Transform constructors set this ``False`` when the observation axis
holds correlated estimates (a ``MorletWavelet`` smoothing neighborhood,
or ``Welch`` segments overlapping by more than half). Measures whose
finite-sample corrections or null distributions count observations
warn in that case, and ``jackknife`` refuses to leave out tapers.
time_bins_are_independent : bool, default=True
Whether the time bins may be counted as independent observations when
the expectation averages over time. Transform constructors set this
``False`` when successive time bins are correlated (``Multitaper`` or
``ShortTimeFourierTransform`` windows overlapping by more than half, or
``MorletWavelet`` samples closer than four wavelet standard
deviations). The measures that count observations then warn for
expectations that include time; it has no effect otherwise.
Attributes
----------
n_observations : int
Number of trial/taper observations reduced by the expectation. This is
the raw count, not a weighted effective sample size, and it is not
reduced for correlated observations (see
``observations_are_independent``).
See Also
--------
spectral_connectivity.transforms.Multitaper : Produce the Fourier
coefficients this class consumes.
spectral_connectivity.wrapper.multitaper_connectivity : High-level interface
returning labeled xarray results with explicit ``source``/``target``
axes.
Notes
-----
**Array orientation**: pairwise measures end in an ``(n_signals, n_signals)``
pair (``(n_groups, n_groups)`` for group measures, ordered as the returned
labels), indexed source first. For the directed measures (the spectral
Granger family, directed transfer function, directed coherence,
(generalized) partial directed coherence, and direct directed transfer
function) ``result[..., i, j]`` is the influence of signal ``i`` on signal
``j`` (``i -> j``). For the lead/lag measures (``directed_phase_lag_index``,
``phase_slope_index``, ``group_delay``, ``delay``) and the antisymmetric
phase measures (``coherence_phase``, ``imaginary_coherency``,
``phase_lag_index``, ``weighted_phase_lag_index``) positive
``result[..., i, j]`` (above 0.5 for ``directed_phase_lag_index``) means
signal ``i`` leads signal ``j``.
The labeled wrappers :func:`~spectral_connectivity.multitaper_connectivity`
and :func:`~spectral_connectivity.fourier_connectivity` use the same order:
``result.sel(source=a, target=b)`` is ``a -> b`` (or "``a`` leads ``b``").
Prefer them unless you need this lower-level API.
Intermediates shared across measures (the expected cross-spectral matrix,
power, phase-lag moments, and the minimum-phase factor with the transfer
function, noise covariance, and MVAR coefficients derived from it) are
cached on first access. Reassigning ``fourier_coefficients``,
``expectation_type``, or ``observation_weights`` automatically invalidates
these caches, so reusing an instance for new data is safe (constructing a
new instance is still the clearer choice). Call :meth:`clear_cache` to
release them while keeping the instance.
The class supports both CPU (NumPy) and GPU (CuPy) computation depending
on the SPECTRAL_CONNECTIVITY_ENABLE_GPU environment variable. For Granger
causality measures, minimum phase decomposition [1]_ is used to estimate
transfer functions and noise covariances non-parametrically.
References
----------
.. [1] Dhamala, M., Rangarajan, G., and Ding, M. (2008). Analyzing
information flow in brain networks with nonparametric Granger
causality. NeuroImage 41, 354-362.
.. [2] Bastos, A. M., & Schoffelen, J. M. (2016). A tutorial review of
functional connectivity analysis methods and their interpretational
pitfalls. Frontiers in systems neuroscience, 9, 175.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity
>>> rng = np.random.default_rng(0)
>>> n_times, n_trials, n_tapers, n_freqs, n_signals = 50, 10, 5, 100, 2
>>> # Create complex coefficients with coherence injected at frequency bin 10
>>> phase_diff = np.pi / 4 # 45 degree phase difference
>>> coeffs = (
... rng.standard_normal((n_times, n_trials, n_tapers, n_freqs, n_signals))
... + 1j
... * rng.standard_normal((n_times, n_trials, n_tapers, n_freqs, n_signals))
... )
>>> coeffs[:, :, :, 10, 1] = coeffs[:, :, :, 10, 0] * np.exp(1j * phase_diff)
>>> conn = Connectivity(coeffs, expectation_type="trials_tapers")
>>> coherence = conn.coherence_magnitude()
>>> coherence.shape # (n_times, non-negative freqs, n_signals, n_signals)
(50, 51, 2, 2)
>>> print(f"Peak coherence: {np.max(coherence[:, 10, 0, 1]):.3f}")
Peak coherence: 1.000
"""
_observation_weights: BackendArray | None
def __init__(
self,
fourier_coefficients: NDArray[np.complexfloating],
expectation_type: str = "trials_tapers",
frequencies: NDArray[np.floating] | None = None,
time: NDArray[np.floating] | None = None,
dtype: DTypeLike = xp.complex128,
minimum_phase_tolerance: float = 1e-8,
minimum_phase_max_iterations: int = 500,
is_one_sided: bool = False,
observation_weights: NDArray[np.floating] | None = None,
observations_are_independent: bool = True,
time_bins_are_independent: bool = True,
*,
_adopt_fourier_coefficients: bool = False,
) -> None:
# fourier_coefficients and expectation_type are validated in their
# property setters (below), which also clear the cached intermediates so
# reassigning either on an existing instance cannot serve stale results.
# _adopt_fourier_coefficients is a private fast path for from_multitaper:
# the fft() output is unshared, so it is frozen in place instead of
# copied, avoiding a transient 2x peak of the largest array. It must not
# be set for a caller-owned array (see _adopt_fourier_coefficients).
if _adopt_fourier_coefficients:
self._adopt_fourier_coefficients(fourier_coefficients)
else:
self.fourier_coefficients = fourier_coefficients
self.expectation_type = expectation_type
self.observation_weights = observation_weights
# Wilson minimum-phase factorization controls, used by the directed
# measures. Near-singular cross-spectral matrices can need more than the
# default iterations to converge; exposing these lets callers recover
# (the non-convergence warning advises increasing max_iterations).
self._minimum_phase_tolerance = minimum_phase_tolerance
self._minimum_phase_max_iterations = minimum_phase_max_iterations
self._is_one_sided = _validated_flag("is_one_sided", is_one_sided)
self._observations_are_independent = _validated_flag(
"observations_are_independent", observations_are_independent
)
self._time_bins_are_independent = _validated_flag(
"time_bins_are_independent", time_bins_are_independent
)
# Fill documented defaults when coordinates are omitted: normalized
# (sampling-frequency-1) FFT frequencies and integer time-window indices.
# Otherwise coordinate-dependent methods (delay, group_delay,
# canonical_coherence) would dereference None.
n_fft_samples = self._fourier_coefficients.shape[-2]
n_time_windows = self._fourier_coefficients.shape[0]
# Supplied coordinates must be 1-D, finite, and match the data geometry.
# A wrong shape (e.g. (n, 1)) or length would silently misalign/drop
# frequency bins (phase_slope_index) or crash with a broadcasting error
# (delay, group_delay, canonical_coherence); non-finite coordinates
# would propagate NaN into the frequency step and delays.
def _validate_coordinate(name: str, coord: Any, expected_length: int) -> None:
arr = xp.asarray(coord)
if arr.ndim != 1:
msg = f"{name} must be a 1-D array, got shape {tuple(arr.shape)}."
raise ValueError(msg)
if arr.shape[0] != expected_length:
msg = f"{name} must have length {expected_length}, got {arr.shape[0]}."
raise ValueError(msg)
if not bool(xp.all(xp.isfinite(arr))):
msg = f"{name} must contain only finite values."
raise ValueError(msg)
if frequencies is not None:
_validate_coordinate("frequencies", frequencies, n_fft_samples)
frequency_values = xp.asarray(frequencies)
if self._is_one_sided and (
bool(xp.any(frequency_values < 0))
or (frequency_values.size > 1 and bool(xp.any(xp.diff(frequency_values) <= 0)))
):
msg = "One-sided frequencies must be non-negative and strictly increasing."
raise ValueError(msg)
if not self._is_one_sided:
_require_fft_order(frequency_values)
if time is not None:
_validate_coordinate("time", time, n_time_windows)
if frequencies is None:
frequencies = (
xp.linspace(0.0, 0.5, n_fft_samples)
if self._is_one_sided
else xp.fft.fftfreq(n_fft_samples)
)
time_values = xp.arange(n_time_windows) if time is None else time
# Frequencies live on the backend (``frequencies`` indexes them with a
# device index array); time is a host coordinate (see ``time`` users).
self._frequencies = xp.asarray(frequencies)
self._dtype = dtype
self.time = to_numpy(time_values)
@property
def observation_weights(self) -> BackendArray | None:
"""Non-negative weights used when averaging spectral observations.
Weights have shape ``(time, trial, taper, frequency, 1)`` and are
shared by every signal. A detached, read-only copy is returned so cached
expectations cannot be invalidated by mutation.
"""
if self._observation_weights is None:
return None
return mark_readonly_if_supported(self._observation_weights.copy())
@observation_weights.setter
def observation_weights(self, value: NDArray[np.floating] | None) -> None:
if value is None:
self._observation_weights = None
self.clear_cache()
return
weights = xp.asarray(value)
expected_shape = (*self._fourier_coefficients.shape[:-1], 1)
if tuple(weights.shape) != expected_shape:
msg = (
"observation_weights must have shape "
f"{expected_shape}, got {tuple(weights.shape)}. Weights must be "
"shared across signals."
)
raise ValueError(msg)
if not bool(xp.all(xp.isfinite(weights))) or bool(xp.any(weights < 0)):
msg = "observation_weights must contain only finite, non-negative values."
raise ValueError(msg)
real_dtype = self._fourier_coefficients.real.dtype
self._observation_weights = mark_readonly_if_supported(
weights.astype(real_dtype, copy=True)
)
self.clear_cache()
[docs]
def clear_cache(self) -> None:
"""Free the intermediates cached for reuse across measures.
Measures computed on one instance share intermediates such as the
expected cross-spectral matrix and the minimum-phase factorization,
which can each take ``n_frequencies * n_signals**2`` values for every
observation that ``expectation_type`` leaves unaveraged (each time
window by default). Call this after the last measure that needs them to
release the memory while keeping the instance; later measures recompute
them and return the same results. Replacing ``fourier_coefficients``,
``expectation_type``, or ``observation_weights`` clears the cache
automatically.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity
>>> rng = np.random.default_rng(0)
>>> fourier_coefficients = rng.standard_normal((1, 5, 3, 16, 4)) + 0j
>>> connectivity = Connectivity(fourier_coefficients)
>>> coherence = connectivity.coherence_magnitude()
>>> connectivity.clear_cache()
>>> bool(np.array_equal(connectivity.coherence_magnitude(), coherence, equal_nan=True))
True
"""
# Discovering descriptors avoids a second, hand-maintained registry that
# could omit a newly added cache. Subclass caches are cleared as well.
for klass in type(self).__mro__:
for name, descriptor in vars(klass).items():
if isinstance(descriptor, cached_property):
self.__dict__.pop(name, None)
@property
def fourier_coefficients(self) -> BackendArray:
"""Multitaper Fourier coefficients.
Shape (n_time_windows, n_trials, n_tapers, n_fft_samples, n_signals).
The instance owns an immutable snapshot so cached calculations cannot
become stale through in-place mutation. This accessor returns a detached
copy, marked read-only when the backend supports it. Assign a new array
through the setter to replace the coefficients and clear the caches.
"""
return mark_readonly_if_supported(self._fourier_coefficients.copy())
@fourier_coefficients.setter
def fourier_coefficients(self, value: NDArray[np.complexfloating]) -> None:
# Public assignment always copies defensively: the caller may keep its
# array and later mutate it in place, which would silently invalidate the
# cached intermediates.
self._set_fourier_coefficients(value, adopt=False)
def _adopt_fourier_coefficients(self, value: NDArray[np.complexfloating]) -> None:
"""Take ownership of a freshly produced, unshared array without copying.
Used only by ``from_multitaper``, where ``value`` is ``transform.fft()``,
which the ``SpectralTransform`` contract requires to be fresh and
referenced nowhere else. This avoids a full copy of the largest array in
the pipeline -- a transient ~2x peak on every construction, which can
push a memory-constrained GPU into OOM.
This is deliberately private and has no public ``copy=False`` surface: it
is safe only when the array *and its writable NumPy base buffer* are
unshared, which ``from_multitaper`` relies on the transform for and
cannot check.
"""
self._set_fourier_coefficients(value, adopt=True)
def _set_fourier_coefficients(
self, value: NDArray[np.complexfloating], *, adopt: bool
) -> None:
# Move a host array (or array-like) onto the active backend before the
# first device operation below (CuPy rejects NumPy operands). This is a
# no-op for an array already on the backend, so the adopt fast path
# stays copy-free.
value = xp.asarray(value)
if value.ndim != 5:
msg = (
f"fourier_coefficients must be 5-dimensional, got {value.ndim}D array.\n"
f"Expected shape: (n_time_windows, n_trials, n_tapers, n_fft_samples, n_signals)\n"
f"Got shape: {value.shape}\n\n"
f"If you have time series data, use the Multitaper class to transform it:\n"
f" from spectral_connectivity import Multitaper\n"
f" m = Multitaper(time_series, sampling_frequency=your_fs, ...)\n"
f" fourier_coefficients = m.fft()"
)
raise ValueError(msg)
if not xp.iscomplexobj(value):
msg = (
f"fourier_coefficients must be complex, got dtype {value.dtype}. "
f"Real-valued coefficients carry no phase, so the imaginary "
f"coherence, the phase-lag indices and the coherence phase would "
f"all be exactly 0. Pass the complex FFT output (e.g. "
f"numpy.fft.fft), not its real part or magnitude."
)
raise TypeError(msg)
# Power spectral density can be computed on single signals, but
# connectivity metrics require >= 2 signals; that is validated per-method
# in _validate_multiple_signals.
if not xp.all(xp.isfinite(value)):
warnings.warn(
"fourier_coefficients contains NaN or Inf values. This may indicate:\n"
" - NaN/Inf in your input time series data\n"
" - Issues with windowing parameters (e.g., window too short)\n"
" - Numerical instability in preprocessing\n\n"
"Suggestions:\n"
" - Check your input data for NaN/Inf values\n"
" - Consider interpolating missing data points\n"
" - Review artifact removal procedures\n"
" - Verify time_window_duration and time_halfbandwidth_product parameters",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
# Own the coefficients as an immutable snapshot: the cached intermediates
# (_power, the reduced cross-spectrum, and the directed-measure factors)
# assume the coefficients change only through the setter, which clears
# them. Marking the snapshot read-only turns an in-place edit via the
# getter into a clear error rather than silently stale results.
if adopt:
# `value` is unshared by the SpectralTransform contract but may be a
# view (Multitaper's is a swapaxes view) whose base buffer is
# writable; freeze the whole base chain, not just the outer view, or
# the data stays reachable and mutable through `.base`. No copy --
# this is the memory-saving path.
mark_readonly_chain_if_supported(value)
owned = value
else:
# copy(order="K") keeps the caller's array's memory layout, so
# downstream matmuls see the same strides and results are unchanged to
# the bit (a plain C-order copy would perturb the BLAS summation order
# by ~1e-16). CuPy may not support the writeable flag; the copy alone
# still decouples the instance from later mutation of the caller's
# array there.
owned = mark_readonly_if_supported(value.copy(order="K"))
value = owned
self._fourier_coefficients = value
self.clear_cache()
# On reassignment (not initial construction), a change in the number of
# FFT bins or time windows invalidates the stored frequency/time
# coordinates. Reset them to geometry-matching defaults so
# coordinate-dependent methods stay consistent -- stale coordinates would
# otherwise silently drop or misalign bins (e.g. phase_slope_index) or
# raise (group_delay). Warn because any user-supplied coordinates (e.g.
# Hz frequencies from from_multitaper) are discarded.
n_fft_samples = value.shape[-2]
n_time_windows = value.shape[0]
frequencies_stale = (
getattr(self, "_frequencies", None) is not None
and len(self._frequencies) != n_fft_samples
)
time_stale = (
getattr(self, "time", None) is not None and len(self.time) != n_time_windows
)
observation_weights = getattr(self, "_observation_weights", None)
expected_weight_shape = (*value.shape[:-1], 1)
weights_stale = observation_weights is not None and (
tuple(observation_weights.shape) != expected_weight_shape
)
if frequencies_stale or time_stale or weights_stale:
warnings.warn(
"Reassigning fourier_coefficients changed the FFT/time geometry; "
"incompatible frequency/time coordinates and observation weights "
"were reset to defaults or cleared. Construct a new Connectivity "
"if you need specific coordinates or weights for the new data.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
if frequencies_stale:
self._frequencies = (
xp.linspace(0.0, 0.5, n_fft_samples)
if getattr(self, "_is_one_sided", False)
else xp.fft.fftfreq(n_fft_samples)
)
if time_stale:
self.time = np.arange(n_time_windows) # host coordinate, as in __init__
if weights_stale:
self._observation_weights = None
@property
def expectation_type(self) -> str:
"""Which dimensions the cross-spectral matrix is averaged over.
Reassigning clears all cached intermediates (see :meth:`clear_cache`).
"""
return self._expectation_type
@expectation_type.setter
def expectation_type(self, value: str) -> None:
if value not in EXPECTATION_AXES:
# Detect the common mistake of the right words in the wrong order.
words = set(value.split("_"))
valid_words = {"time", "trials", "tapers"}
suggestion = None
if words.issubset(valid_words):
for valid_key in EXPECTATION_AXES:
if set(valid_key.split("_")) == words:
suggestion = valid_key
break
error_msg = (
f"Invalid expectation_type '{value}' is not supported.\n"
f"This parameter controls which dimensions to average over when computing "
f"the cross-spectral matrix.\n"
)
if suggestion:
error_msg += (
f"\nDid you mean '{suggestion}'? (The words must be in a specific order)\n"
)
error_msg += "\nValid options are:\n"
for key in sorted(EXPECTATION_AXES):
error_msg += f" - '{key}'\n"
error_msg += "\nMost common: 'trials_tapers' (average over both trials and tapers)"
raise ValueError(error_msg)
self._expectation_type = value
self.clear_cache()
[docs]
@classmethod
def from_multitaper(
cls,
multitaper_instance: "SpectralTransform",
expectation_type: str = "trials_tapers",
dtype: Any = xp.complex128,
minimum_phase_tolerance: float = 1e-8,
minimum_phase_max_iterations: int = 500,
) -> "Connectivity":
"""Construct from a spectral transform; the original name of from_transform.
Accepts any :class:`~spectral_connectivity.transforms.SpectralTransform`.
Parameters
----------
multitaper_instance : SpectralTransform
The transform; ``transform`` in :meth:`from_transform`.
expectation_type, dtype, minimum_phase_tolerance, minimum_phase_max_iterations
As in :meth:`from_transform`.
Returns
-------
Connectivity
New Connectivity instance.
"""
init_kwargs: dict[str, Any] = {
"expectation_type": expectation_type,
"dtype": dtype,
"minimum_phase_tolerance": minimum_phase_tolerance,
"minimum_phase_max_iterations": minimum_phase_max_iterations,
}
# The optional SpectralTransform attributes must always reach the
# instance: a subclass that cannot accept them fails loudly here rather
# than silently treating a one-sided, weighted, or correlated spectrum as
# two-sided, unweighted, and independent. They are passed only when
# non-default so a subclass mirroring the older signature keeps working
# with a plain two-sided transform. The flags are checked before the
# (possibly expensive) fft() so a bad one fails fast; the coordinates
# and weights are read after it, which may set them.
if _transform_flag(multitaper_instance, "is_one_sided", False):
init_kwargs["is_one_sided"] = True
if not _transform_flag(multitaper_instance, "observations_are_independent", True):
init_kwargs["observations_are_independent"] = False
if not _transform_flag(multitaper_instance, "time_bins_are_independent", True):
init_kwargs["time_bins_are_independent"] = False
init_kwargs["fourier_coefficients"] = multitaper_instance.fft()
init_kwargs["time"] = multitaper_instance.time
init_kwargs["frequencies"] = multitaper_instance.frequencies
weights = _optional_transform_attribute(
multitaper_instance, "observation_weights", None
)
if weights is not None:
init_kwargs["observation_weights"] = weights
# The SpectralTransform contract requires fft() to return a freshly
# built, unshared array, so adopt it in place instead of copying (see
# Connectivity._adopt_fourier_coefficients). Only pass the private
# keyword when the subclass has not overridden __init__:
# an overriding subclass need not accept it, and passing it would raise
# TypeError. Such a subclass falls back to the defensive-copy path.
if cls.__init__ is Connectivity.__init__:
init_kwargs["_adopt_fourier_coefficients"] = True
return cls(**init_kwargs)
def _validate_multiple_signals(self) -> None:
"""Raise if fewer than two signals are present.
Connectivity measures quantify relationships between pairs of signals
and are undefined for a single signal (they would otherwise return an
all-NaN array with no error). ``power()`` is exempt because power
spectral density is well-defined for one signal.
"""
n_signals = self._fourier_coefficients.shape[-1]
if n_signals < 2:
msg = (
f"Connectivity measures require at least 2 signals, but "
f"fourier_coefficients has {n_signals} signal "
f"(shape[-1] == {n_signals}).\n"
f"Connectivity quantifies relationships between pairs of "
f"signals; for a single signal use power() instead.\n"
f"If you sliced to one channel, keep the signal axis, e.g. "
f"fourier_coefficients[..., [channel_index]] rather than "
f"fourier_coefficients[..., channel_index]."
)
raise ValueError(msg)
def _require_uniform_frequency_grid(self, measure: str) -> None:
"""Raise if the frequency coordinate is not equally spaced.
Measures that combine adjacent bins (phase slope, delay estimates)
interpret one bin step as one frequency step; a wavelet transform can
expose an arbitrary grid, on which those combinations are meaningless.
"""
steps = np.diff(to_numpy(self.frequencies))
if steps.size and not np.allclose(steps, steps[0], rtol=1e-6, atol=0.0):
msg = (
f"{measure} requires uniformly spaced frequencies because it "
f"combines adjacent frequency bins, but the grid steps range "
f"from {steps.min():g} to {steps.max():g} Hz. Use a linearly "
"spaced frequency grid (e.g. MorletWavelet(frequencies="
"np.arange(low, high, step)))."
)
raise ValueError(msg)
def _require_uniform_observation_weights(self, measure: str, reason: str) -> None:
"""Raise if observation weights are non-uniform for a measure that
assumes equally weighted independent observations."""
if self._observation_weights_are_uniform:
return
msg = (
f"{measure} does not support non-uniform observation_weights. "
f"{reason} Non-uniform weights come from a smoothing_kernel other "
"than 'boxcar', or from edge_mode='nan' when the time axis is part "
"of the expectation (the edge mask zeroes some observations). Use "
"smoothing_kernel='boxcar' with edge_mode='trim' or 'keep', or an "
"expectation_type that keeps the time axis."
)
raise ValueError(msg)
def _validate_debiasing_observations(self, measure: str) -> None:
"""Raise if a bias-corrected measure has too few observations.
The debiased phase-lag index and pairwise phase consistency divide by
``n_observations - 1`` (respectively ``n_observations ** 2 -
n_observations``), which is zero for a single observation and would
otherwise return silent inf/NaN.
"""
n_observations = self.n_observations
if n_observations < 2:
msg = (
f"{measure} requires at least 2 observations "
f"(n_observations == n_trials * n_tapers), but got "
f"{n_observations}. This bias correction divides by a factor "
f"that is zero when n_observations < 2, so it is undefined for a "
f"single observation. Use more trials/tapers, or the "
f"non-debiased measure (phase_lag_index / phase_locking_value)."
)
raise ValueError(msg)
self._require_uniform_observation_weights(
measure,
"Its finite-sample correction assumes equally weighted independent "
"observations; use a non-debiased measure instead.",
)
self._warn_correlated_observations(
measure,
"Its finite-sample bias correction treats every observation as an "
"independent sample",
)
def _warn_correlated_observations(self, measure: str, assumption: str) -> None:
"""Warn once that ``measure`` counts correlated observations as independent.
``n_observations`` is a raw count of the averaged observations. When
the transform reports ``observations_are_independent=False`` (a
MorletWavelet smoothing neighborhood or Welch segments overlapping by
more than half on the observation axis), or the expectation averages
over time bins it reports as correlated
(``time_bins_are_independent=False``: overlapping Multitaper/STFT
windows or closely spaced Morlet samples), that count overstates the
effective sample size, biasing every consumer that uses it as a
degrees-of-freedom or bias-correction factor. ``assumption`` states
what ``measure`` uses the count for so the message is actionable.
"""
averages_correlated_time_bins = (
not self._time_bins_are_independent and 0 in self._expectation_axes
)
if self._observations_are_independent and not averages_correlated_time_bins:
return
reasons = []
if not self._observations_are_independent:
reasons.append(
"this transform's trial/taper observations are correlated "
"(observations_are_independent is False: a MorletWavelet smoothing "
"neighborhood, or Welch segments overlapping by more than half; use "
"a Multitaper transform, Welch with segment_overlap <= 0.5, or "
"MorletWavelet without smoothing)"
)
if averages_correlated_time_bins:
reasons.append(
f"expectation_type={self.expectation_type!r} averages over time "
"bins that are correlated (time_bins_are_independent is False: "
"Multitaper or ShortTimeFourierTransform windows overlapping by "
"more than half, or MorletWavelet samples closer than 4 wavelet "
"standard deviations; use time_window_step >= time_window_duration "
"/ 2, more MorletWavelet decimation, or an expectation_type that "
"does not average over time)"
)
warnings.warn(
f"{measure} assumes independent observations, but "
f"{' and '.join(reasons)}. {assumption}, so n_observations == "
f"{self.n_observations} overstates the effective sample size and the "
f"result is biased.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
def _warn_single_observation_degenerate(
self,
measure: str,
consequence: str = (
"every magnitude-normalized value is mathematically forced to 1 "
"(apparent perfect connectivity)"
),
) -> None:
"""Warn that a normalized measure is degenerate for one observation.
Coherency, the phase-locking value, and every measure derived from
them normalize each cross-spectral entry by its magnitude; the
phase-lag family averages the sign of one imaginary cross-spectrum per
observation; partial coherence inverts a rank-one cross-spectral
matrix. With a single observation (one trial times one taper/window --
e.g. a single-trial
:class:`~spectral_connectivity.transforms.MorletWavelet` transform
without ``smoothing_time``) the result is fixed by that normalization
rather than by the data (``consequence`` states how for ``measure``),
so the measure reports perfect connectivity or a perfectly consistent
lag between unrelated signals and carries no information. This is a
silent-failure trap rather than an error, so it warns instead of
raising.
"""
if self.n_observations < 2:
warnings.warn(
f"{measure} is computed from a single observation "
f"(n_observations == n_trials * n_tapers == 1), so {consequence} "
"regardless of the data, and the result carries no information. "
"Provide multiple trials/tapers, or set smoothing_time on "
"MorletWavelet to collect neighboring coefficients on the "
"observation axis.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
def _require_multiple_frequencies(self, measure: str) -> None:
"""Raise if fewer than two frequency bins are available.
``delay``, ``group_delay`` and ``phase_slope_index`` read
``frequencies[1] - frequencies[0]`` to get the frequency step; with a
single frequency bin that indexing would raise a raw ``IndexError``
instead of a clear message.
"""
n_frequencies = len(self.frequencies)
if n_frequencies < 2:
msg = (
f"{measure} requires at least 2 frequency bins, but the data has "
f"{n_frequencies}. Use a longer FFT (larger n_fft_samples / "
f"n_time_samples_per_window)."
)
raise ValueError(msg)
def _require_two_sided_spectrum(self, measure: str) -> None:
"""Reject directed factorization for positive-frequency-only inputs."""
if self._is_one_sided:
msg = (
f"{measure} requires a full two-sided spectrum in standard FFT "
"order. One-sided transforms such as Morlet wavelets support "
"functional connectivity measures, but not Wilson-factorized "
"directed measures."
)
raise ValueError(msg)
def _nonnegative_frequency_count(self, n_frequencies: int) -> int:
"""Number of non-negative-frequency bins among ``n_frequencies``.
Every bin for one-sided input; ``n_frequencies // 2 + 1`` (DC through
Nyquist) for a two-sided spectrum in standard FFT order. The single
source of truth for trimming results to non-negative frequencies.
"""
return n_frequencies if self._is_one_sided else n_frequencies // 2 + 1
def _one_sided_density(self, spectrum: BackendArray, frequency_axis: int) -> BackendArray:
"""Fold a cached (cross-)spectral density onto non-negative frequencies.
Two-sided input is trimmed to its non-negative bins and the interior
positive-frequency bins are doubled, so the one-sided density integrates
to the same total power as the two-sided spectrum. DC (bin 0) is unique;
the Nyquist bin (present only for an even FFT length) is also unique, so
neither is doubled. The scale matches the spectrum's real dtype so a
float32 (complex64) request is not silently upcast to float64.
One-sided input is returned as a copy, detached from the cache so a
caller mutating the result cannot corrupt measures that reuse it.
Parameters
----------
spectrum : array, shape (..., n_fft_samples, ...)
Cached spectrum with frequency at ``frequency_axis``.
frequency_axis : int
Negative index of the frequency axis.
Returns
-------
array, shape (..., n_frequencies, ...)
"""
if self._is_one_sided:
return spectrum.copy()
n_fft_samples = spectrum.shape[frequency_axis]
trailing = (slice(None),) * (-frequency_axis - 1)
index: tuple[Any, ...] = (
...,
slice(self._nonnegative_frequency_count(n_fft_samples)),
*trailing,
)
one_sided = spectrum[index]
scale = xp.full((one_sided.shape[frequency_axis],), 2.0, dtype=one_sided.real.dtype)
scale[0] = 1.0
if n_fft_samples % 2 == 0:
scale[-1] = 1.0
density: BackendArray = one_sided * scale.reshape((-1,) + (1,) * (-frequency_axis - 1))
return density
@property
@_asnumpy
def frequencies(self) -> NDArray[np.floating]:
"""Return non-negative frequencies of the transform.
Returns
-------
NDArray[floating], shape (n_frequencies,)
Non-negative frequency values.
"""
n_nonnegative = self._nonnegative_frequency_count(len(self._frequencies))
freqs = xp.take(self._frequencies, indices=xp.arange(n_nonnegative), axis=0)
# fftfreq returns negative Nyquist for even N, fix the sign
if len(freqs) > 0 and freqs[-1] < 0:
freqs = freqs.copy() # Don't modify the original
freqs[-1] = abs(freqs[-1])
return freqs
@property
@_asnumpy
def all_frequencies(self) -> NDArray[np.floating]:
"""Return positive and negative frequencies of the transform.
Returns
-------
NDArray[floating], shape (n_frequencies,)
All frequency values including negative frequencies.
"""
return self._frequencies
@cached_property
def _power(self) -> NDArray[np.floating]:
# Reused by coherency and directed measures; the input setters discover
# and invalidate cached_property values automatically.
return self._expectation(
self._fourier_coefficients * self._fourier_coefficients.conjugate()
).real
def _nonnegative_pairwise_power_scale(self) -> NDArray[np.floating]:
"""``sqrt(P_i P_j)`` at the non-negative frequencies.
Shape (..., n_nonnegative_frequencies, n_signals, n_signals). Built only
for the bins the normalized measures report, and recomputed per call: it
costs one outer product of ``sqrt(_power)``, which also cannot underflow
or overflow.
"""
n_nonnegative = self._nonnegative_frequency_count(self._power.shape[-2])
root_power = xp.sqrt(self._power[..., :n_nonnegative, :])
return root_power[..., :, xp.newaxis] * root_power[..., xp.newaxis, :]
def _nonnegative_fourier_coefficients(self) -> NDArray[np.complexfloating]:
"""Fourier coefficients at the non-negative frequencies, at ``self._dtype``.
A view when the dtype already matches.
"""
n_nonnegative = self._nonnegative_frequency_count(self._fourier_coefficients.shape[-2])
coefficients: NDArray[np.complexfloating] = self._fourier_coefficients[
..., :n_nonnegative, :
].astype(self._dtype, copy=False)
return coefficients
def _nonnegative_cross_spectral_matrix(self) -> NDArray[np.complexfloating]:
"""Expected cross-spectral matrix at the non-negative frequencies (a view)."""
cross_spectral_matrix = self._expectation_cross_spectral_matrix()
n_nonnegative = self._nonnegative_frequency_count(cross_spectral_matrix.shape[-3])
return cross_spectral_matrix[..., :n_nonnegative, :, :]
@property
def _cross_spectral_matrix(self) -> NDArray[np.complexfloating]:
"""Return the complex-valued linear association between fourier coefficients.
Returns
-------
cross_spectral_matrix : array
Shape (n_time_windows, n_trials, n_tapers, n_fft_samples,
n_signals, n_signals). Complex cross-spectral matrix.
"""
fourier_coefficients = self._fourier_coefficients[..., xp.newaxis]
return _complex_inner_product(
fourier_coefficients, fourier_coefficients, dtype=self._dtype
)
def _expectation_cross_spectral_matrix(self) -> NDArray[np.complexfloating]:
"""Expected cross-spectral matrix, reduced over the averaged observations.
Validates that at least two signals are present, then returns the cached
reduced cross-spectral matrix -- a single batched matmul over the
averaged time/trials/tapers axes (see ``_reduced_cross_spectral_matrix``)
rather than a per-observation outer product.
Returns
-------
array, shape (..., n_frequencies, n_signals, n_signals)
Expected cross-spectral matrix.
"""
self._validate_multiple_signals()
return self._cached_reduced_cross_spectral_matrix
def _reduced_cross_spectral_matrix(
self, fourier_coefficients: NDArray[np.complexfloating] | None = None
) -> NDArray[np.complexfloating]:
"""Expected cross-spectral matrix via a single batched matmul.
Numerically equivalent (to floating-point tolerance) to
``self._expectation(self._cross_spectral_matrix)``, but contracts the
averaged observation axes (any subset of time/trials/tapers, taken from
the active ``expectation_type``) directly instead of materializing the
full ``(..., n_signals, n_signals)`` outer product for every
observation. For the default ``trials_tapers`` expectation this replaces
a large intermediate with a small result and is markedly faster.
Parameters
----------
fourier_coefficients : array, optional
Coefficients of shape
``(n_time_windows, n_trials, n_tapers, n_fft_samples, *batch,
n_signals)`` to reduce; defaults to this instance's coefficients.
``phase_locking_value`` passes unit-normalized coefficients so the
same batched matmul yields its normalized cross-spectrum, and
``_subset_cross_spectral_matrix`` passes a pair axis as ``batch``.
The frequency axis may hold only the leading bins (e.g. the
non-negative frequencies); observation weights are sliced to match.
Returns
-------
array, shape (..., n_frequencies, *batch, n_signals, n_signals)
Expected cross-spectral matrix. The leading axes are whichever of
time/trials/tapers are *not* averaged, matching the shape produced
by the equivalent expectation over the full outer product.
"""
if fourier_coefficients is None:
fourier_coefficients = self._fourier_coefficients
average_axes = self._expectation_axes
signal_axis = fourier_coefficients.ndim - 1
frequency_axis = 3
# Frequency and any extra batch axes (axes 3 to signal_axis - 1) stay.
batch_axes = list(range(frequency_axis, signal_axis))
kept_axes = [axis for axis in range(frequency_axis) if axis not in average_axes]
# Reorder to (kept leading axes..., frequency, batch..., averaged axes...,
# signals) so the averaged axes collapse into a single observation axis
# adjacent to signals, ready for a batched matmul.
order = [*kept_axes, *batch_axes, *average_axes, signal_axis]
observations = xp.transpose(fourier_coefficients, order)
n_observations = int(
np.prod([fourier_coefficients.shape[axis] for axis in average_axes])
)
n_signals = fourier_coefficients.shape[signal_axis]
observations = observations.reshape(
(
*observations.shape[: len(kept_axes) + len(batch_axes)],
n_observations,
n_signals,
)
)
weights = None
if self._observation_weights is not None:
# Weights vary over observations and frequency, not the extra batch.
weights = xp.transpose(
self._observation_weights[
..., : fourier_coefficients.shape[frequency_axis], 0
],
[*kept_axes, frequency_axis, *average_axes],
).reshape(
(
*observations.shape[: len(kept_axes) + 1],
*(1,) * (len(batch_axes) - 1),
n_observations,
)
)
observations = observations * xp.sqrt(weights)[..., xp.newaxis]
# cross_spectral_matrix[..., i, j] = mean_obs f_i * conj(f_j), matching
# _complex_inner_product's convention, then averaged over observations.
cross_spectral_matrix: NDArray[np.complexfloating] = xp.matmul(
xp.swapaxes(observations, -1, -2),
xp.conj(observations),
dtype=self._dtype,
)
if weights is None:
return cross_spectral_matrix / n_observations
denominator = xp.sum(weights, axis=-1)[..., xp.newaxis, xp.newaxis]
return _divide_where(cross_spectral_matrix, denominator, denominator > 0, xp.nan)
@cached_property
def _cached_reduced_cross_spectral_matrix(self) -> NDArray[np.complexfloating]:
"""Cache the reduced expected cross-spectral matrix.
This is the reduced result of ``_expectation_cross_spectral_matrix()``.
It is reused within a measure (``coherency`` divides it by power) and
across measures that share one ``Connectivity`` instance
(``coherence_magnitude``, ``coherence_phase``, ``imaginary_coherence``,
pairwise spectral Granger). Only the reduced ``(..., n_signals,
n_signals)`` form is cached, never the observation-resolved
``_cross_spectral_matrix``. Invalidated with the other cached
intermediates when the inputs change; consumers treat it as read-only.
"""
return self._reduced_cross_spectral_matrix()
def _subset_cross_spectral_matrix(
self, pairs: Sequence[Sequence[int]] | NDArray[np.integer]
) -> NDArray[np.complexfloating]:
"""Compute compact expected cross-spectra for channel pairs.
Each pair's two signals are reduced over observations by the batched
matmul of :meth:`_reduced_cross_spectral_matrix`, with the pairs as a
batch axis processed in chunks of at most
``SUBSET_CROSS_SPECTRUM_MAX_WORKSPACE_ELEMENTS`` gathered coefficients.
Neither the full ``n_signals x n_signals`` matrix nor an
observation-resolved 2-by-2 matrix per pair is formed.
Parameters
----------
pairs : array_like
Pairs of channel indices.
Returns
-------
array, shape (..., n_pairs, n_frequencies, 2, 2)
One compact 2-by-2 expected cross-spectral matrix per requested
pair. The leading axes are the observation axes the configured
expectation keeps.
"""
pair_indices = xp.asarray(pairs, dtype=int)
if pair_indices.ndim != 2 or pair_indices.shape[1] != 2:
msg = "pairs must have shape (n_pairs, 2)."
raise ValueError(msg)
if pair_indices.size == 0:
msg = "pairs must contain at least one signal pair."
raise ValueError(msg)
n_signals = self._fourier_coefficients.shape[-1]
if bool(xp.any(pair_indices < 0)) or bool(xp.any(pair_indices >= n_signals)):
msg = f"pair indices must be between 0 and {n_signals - 1}."
raise IndexError(msg)
coefficients = self._fourier_coefficients
n_pairs = pair_indices.shape[0]
elements_per_pair = 2 * int(np.prod(coefficients.shape[:-1]))
pairs_per_chunk = max(
1, SUBSET_CROSS_SPECTRUM_MAX_WORKSPACE_ELEMENTS // elements_per_pair
)
# Each chunk: (time, trials, tapers, frequency, pair, 2) gathered, reduced
# to (..., frequency, pair, 2, 2).
chunks = [
self._reduced_cross_spectral_matrix(
coefficients[..., pair_indices[start : start + pairs_per_chunk]]
)
for start in range(0, n_pairs, pairs_per_chunk)
]
# Put pair before frequency so frequency remains axis -3 as required by
# the Wilson factorization.
return xp.moveaxis(xp.concatenate(chunks, axis=-3), -3, -4)
# These quantities feed every directed-connectivity measure and are
# expensive to compute (the minimum-phase decomposition in particular), so
# they are cached per instance. The validated input setters clear these
# cached properties before accepting a replacement.
@cached_property
def _minimum_phase_factor(self) -> NDArray[np.complexfloating]:
self._require_two_sided_spectrum("Directed connectivity")
return _factorize_spectrum(
self._expectation_cross_spectral_matrix(),
minimum_phase_tolerance=self._minimum_phase_tolerance,
minimum_phase_max_iterations=self._minimum_phase_max_iterations,
)
@cached_property
def _var_model(self) -> tuple[NDArray[np.complexfloating], NDArray[np.floating]]:
return _var_model_from_factor(self._minimum_phase_factor)
@property
def _transfer_function(self) -> NDArray[np.complexfloating]:
return self._var_model[0]
@property
def _noise_covariance(self) -> NDArray[np.floating]:
return self._var_model[1]
@cached_property
def _MVAR_Fourier_coefficients(self) -> NDArray[np.complexfloating]:
return _regularized_inverse(self._transfer_function)
def _expectation(self, values: BackendArray, *, frequency_axis: int = 3) -> BackendArray:
"""Average observation axes, applying optional spectral weights.
``values`` may hold only the leading (non-negative) frequency bins of
the spectrum; the weights are restricted to the same bins.
"""
if self._observation_weights is None:
expected: BackendArray = xp.mean(values, axis=self._expectation_axes)
return expected
if frequency_axis < 0:
frequency_axis += values.ndim
if frequency_axis < 3 or frequency_axis >= values.ndim:
msg = "frequency_axis must follow the three observation axes."
raise ValueError(msg)
n_frequencies = values.shape[frequency_axis]
weight_shape = [1] * values.ndim
weight_shape[0:3] = self._observation_weights.shape[0:3]
weight_shape[frequency_axis] = n_frequencies
weights = self._observation_weights[..., :n_frequencies, 0].reshape(weight_shape)
numerator = xp.sum(values * weights, axis=self._expectation_axes)
denominator = xp.sum(weights, axis=self._expectation_axes)
return _divide_where(numerator, denominator, denominator > 0, xp.nan)
@cached_property
def _observation_weights_are_uniform(self) -> bool:
"""Whether every reduced bin gives its observations equal weight."""
if self._observation_weights is None:
return True
weights = self._observation_weights[..., 0]
kept_axes = [axis for axis in range(3) if axis not in self._expectation_axes]
order = [*kept_axes, 3, *self._expectation_axes]
reordered = xp.transpose(weights, order)
flattened = reordered.reshape((*reordered.shape[: len(kept_axes) + 1], -1))
if flattened.shape[-1] < 2:
return True
return bool(xp.all(flattened == flattened[..., :1]))
@property
def _expectation_axes(self) -> tuple[int, ...]:
"""Observation axes reduced by the configured expectation."""
return EXPECTATION_AXES[self.expectation_type]
@property
def n_observations(self) -> int:
"""Return the raw number of observations averaged by the expectation.
Returns
-------
int
Product of the lengths of the averaged observation axes (for the
default ``"trials_tapers"`` expectation, ``n_trials * n_tapers``).
This is a raw count, not an effective number of independent
observations: it ignores ``observation_weights`` and is not reduced
when the transform's observations are correlated (see
``observations_are_independent``).
"""
return int(
np.prod(
[self._fourier_coefficients.shape[axis] for axis in self._expectation_axes]
)
)
@property
def n_signals(self) -> int:
"""Number of signals represented by the Fourier coefficients."""
return int(self._fourier_coefficients.shape[-1])
@property
def is_one_sided(self) -> bool:
"""Whether the input contains only non-negative frequencies."""
return self._is_one_sided
@property
def observations_are_independent(self) -> bool:
"""Whether the trial/taper observations are statistically independent.
``False`` when the transform collected correlated estimates on the
observation axis (a smoothing neighborhood of a
:class:`~spectral_connectivity.transforms.MorletWavelet`, or
:class:`~spectral_connectivity.transforms.Welch` segments overlapping by
more than half). ``n_observations`` then overstates the effective sample
size, so the measures that rely on it warn and ``jackknife`` refuses to
treat tapers as leave-one-out units.
"""
return self._observations_are_independent
@property
def time_bins_are_independent(self) -> bool:
"""Whether the time bins may be counted as independent observations.
``False`` when the transform's successive time bins are correlated
(:class:`~spectral_connectivity.transforms.Multitaper` or
:class:`~spectral_connectivity.transforms.ShortTimeFourierTransform`
windows overlapping by more than half, or
:class:`~spectral_connectivity.transforms.MorletWavelet` samples closer
than four wavelet standard deviations). It matters only when the
expectation averages over time: ``n_observations`` then counts the
correlated bins and overstates the effective sample size, so the
measures that rely on it warn.
"""
return self._time_bins_are_independent
[docs]
@_asnumpy
def minimum_phase_reconstruction_error(self) -> NDArray[np.floating]:
"""Return the relative reconstruction error of the Wilson factorization.
This diagnostic checks how faithfully the cached minimum-phase factor
reconstructs the expected cross-spectral matrix. One value is returned
per retained batch (normally time-window) dimension. Values near machine
precision indicate a faithful factorization; large values suggest that
the spectrum is too coarsely resolved for directed-connectivity measures.
Non-converged factorizations return ``NaN``.
Returns
-------
array, shape (...,)
Maximum relative reconstruction error for each sub-spectrum.
Notes
-----
A full two-sided spectrum is required, as for the directed measures that
use the Wilson factorization.
"""
self._require_two_sided_spectrum("minimum_phase_reconstruction_error")
return _minimum_phase_reconstruction_error(
self._expectation_cross_spectral_matrix(),
self._minimum_phase_factor,
)
[docs]
def jackknife(
self,
method: str,
*,
confidence_level: float = 0.95,
transformation: Literal[
"auto", "identity", "log", "fisher", "fisher_squared", "circular"
] = "auto",
**method_kwargs: Any,
) -> JackknifeResult:
"""Estimate uncertainty by leaving out one trial/taper observation.
The configured expectation must average trials, tapers, or their
combination. For ``trials_tapers`` each trial-taper eigencoefficient is
treated as one observation. The method is recomputed for every
leave-one-out sample, so it supports nonlinear measures without an
analytic variance formula (at a computational cost proportional to
``n_observations``).
Parameters
----------
method : str
Name of a public, real-valued connectivity measure to recompute,
e.g. ``"coherence_magnitude"``. Complex-valued and tuple-valued
measures are not supported.
confidence_level : float, default=0.95
Two-sided coverage of the interval, in (0, 1). The critical value
is Student t with ``n_observations - 1`` degrees of freedom.
transformation : {"auto", "identity", "log", "fisher",
"fisher_squared", "circular"}
Scale on which the interval is formed. ``"auto"`` resolves to:
- ``"log"`` for ``power``;
- ``"fisher_squared"`` (``atanh(sqrt(.))``) for the
magnitude-squared measures in ``[0, 1]``: ``coherence_magnitude``
and ``partial_coherence``;
- ``"fisher"`` (``atanh(.)``) for the magnitudes in ``[0, 1]``:
``phase_locking_value`` and ``imaginary_coherence``. For these
two measures the Fisher lower bound and bias-corrected estimate
are clamped at 0, since a magnitude cannot be negative;
- ``"circular"`` for ``coherence_phase``;
- ``"identity"`` for every other measure.
**method_kwargs
Keyword arguments forwarded to ``method`` on every replicate.
Returns
-------
JackknifeResult
Estimate, bias-corrected estimate, standard error, and confidence
bounds, each with the measure's own shape (for a pairwise measure
``(n_time, n_nonnegative_frequencies, n_signals, n_signals)``).
For directed measures ``[..., i, j]`` is the influence ``i -> j``,
the same order as the xarray wrapper's ``sel(source=i, target=j)``.
Raises
------
ValueError
If the expectation averages tapers (``"tapers"`` or
``"trials_tapers"``) but ``observations_are_independent`` is
``False``: correlated observations (a MorletWavelet smoothing
neighborhood, overlapping Welch segments) are not valid
leave-one-out units. Recompute the transform with independent
observations instead: ``Welch`` with ``segment_overlap <= 0.5``, a
``Multitaper`` transform, or ``MorletWavelet`` without
``smoothing_time`` on at least 3 trials.
(``expectation_type="trials"`` is accepted but keeps the correlated
axis as an output dimension instead of averaging it, a different
measure.)
Notes
-----
Circular confidence bounds are wrapped to ``(-pi, pi]``, so when the
interval crosses ``+/-pi`` the lower bound exceeds the upper bound;
the interval is then the arc from ``lower`` up through ``pi`` and on
to ``upper``.
If the input Fourier coefficients were produced with
``Multitaper(taper_weighting="adaptive")``, the leave-one-out replicates
reuse the full-sample Thomson weights (the adaptive weights are not
recomputed for each reduced taper set), so the interval is an
approximation in that case.
The interval describes the size of a measure; it is not a test that
the measure differs from 0. Monte Carlo coverage of the default
(``"auto"``) 95% intervals, with independent complex-Gaussian
observations and 5 to 100 observations per dataset:
- ``coherence_magnitude`` (``fisher_squared``): 94-95% when the true
``|coherency|`` is 0.3-0.8, but at zero true coherence the interval
excludes 0 in 11-15% of datasets, and more observations do not help
(the estimated magnitude's spread shrinks with ``n`` at the same
rate as its mean). "The interval excludes 0" is therefore not
evidence of nonzero coherence. Test that with the exact
zero-coherence null instead:
:func:`spectral_connectivity.statistics.coherence_significance_pvalue`
applied to ``coherency()`` and ``n_observations`` (valid for
independent, equally weighted observations), corrected across
frequencies and pairs with
:func:`spectral_connectivity.statistics.adjust_for_multiple_comparisons`.
For measures without an analytic null, use a permutation or
surrogate test.
- ``phase_locking_value`` (``fisher``): at zero true PLV the interval
excludes 0 in 5% of datasets at 5 observations but 14% at 100; at a
high true PLV it under-covers (85-95% at a true PLV of 0.82, 76-90%
at 0.93).
- ``coherence_phase`` (``circular``): under-covers when coherence is
weak (77-88% at a true ``|coherency|`` of 0.1, 84-93% at 0.2,
87-94% at 0.3, against 92-95% at 0.6).
"""
if self.expectation_type not in {"trials", "tapers", "trials_tapers"}:
msg = (
"jackknife supports expectation_type 'trials', 'tapers', or "
"'trials_tapers'; expectations involving time or retaining both "
"trial and taper axes have no single leave-one-out layout."
)
raise ValueError(msg)
if not self._observations_are_independent and self.expectation_type in {
"tapers",
"trials_tapers",
}:
msg = (
f"jackknife with expectation_type={self.expectation_type!r} leaves "
"out one taper observation at a time, but this transform's "
"observations are correlated (observations_are_independent is "
"False: a MorletWavelet smoothing neighborhood, or Welch segments "
"overlapping by more than half), so leave-one-out intervals over "
"them are invalid. For an interval on the same averaged measure, "
"recompute the transform with independent observations: Welch "
"with segment_overlap <= 0.5, a Multitaper transform, or "
"MorletWavelet without smoothing_time on at least 3 trials (the "
"trials are then the observations, and there is no time "
"smoothing). expectation_type='trials' keeps the correlated axis "
"as an output dimension instead of averaging it, so it measures "
"something else."
)
raise ValueError(msg)
method_attribute = inspect.getattr_static(type(self), method, None)
if (
method.startswith("_")
or method in _NON_MEASURE_METHODS
or not inspect.isfunction(method_attribute)
):
msg = "method must name a public connectivity measure."
raise ValueError(msg)
# The static check above guarantees ``method`` names a public function
# on the class, so the bound attribute is always callable here.
measure = getattr(self, method)
full_estimate = measure(**method_kwargs)
if isinstance(full_estimate, tuple):
msg = f"jackknife does not support tuple-valued measure {method!r}."
raise TypeError(msg)
full_estimate = np.asarray(full_estimate)
if full_estimate.dtype == object:
msg = (
f"jackknife requires a real array result from {method!r}; it "
f"returned a structured result. Jackknife the scalar score "
"measure instead (e.g. canonical_coherence rather than "
"canonical_coherency)."
)
raise TypeError(msg)
if np.iscomplexobj(full_estimate):
msg = f"jackknife requires a real-valued measure; {method!r} is complex."
raise TypeError(msg)
coefficients = self._fourier_coefficients
observation_weights = self._observation_weights
if self.expectation_type == "trials_tapers":
n_observations = coefficients.shape[1] * coefficients.shape[2]
observation_coefficients = coefficients.reshape(
coefficients.shape[0],
n_observations,
1,
coefficients.shape[3],
coefficients.shape[4],
)
if observation_weights is not None:
observation_weights = observation_weights.reshape(
observation_weights.shape[0],
n_observations,
1,
observation_weights.shape[3],
1,
)
observation_axis = 1
else:
observation_coefficients = coefficients
observation_axis = 2 if self.expectation_type == "tapers" else 1
n_observations = coefficients.shape[observation_axis]
if n_observations < 3:
msg = (
f"jackknife requires at least 3 observations, got {n_observations}. "
"With two, each leave-one-out replicate has a single observation, "
"which forces magnitude-normalized measures to 1 and makes the "
"interval degenerate (a zero standard error, or NaN under a "
"variance-stabilizing transform such as the coherence default)."
)
raise ValueError(msg)
replicates: list[NDArray[np.floating]] = []
for omitted in range(n_observations):
keep = xp.arange(n_observations) != omitted
subset = xp.compress(keep, observation_coefficients, axis=observation_axis)
subset_weights = (
None
if observation_weights is None
else xp.compress(keep, observation_weights, axis=observation_axis)
)
# ``subset`` is a fresh, unshared compress() output, so adopt it in
# place instead of copying and re-scanning it for every replicate.
replicate_connectivity = Connectivity(
subset,
expectation_type=self.expectation_type,
frequencies=self._frequencies,
time=self.time,
dtype=self._dtype,
minimum_phase_tolerance=self._minimum_phase_tolerance,
minimum_phase_max_iterations=self._minimum_phase_max_iterations,
is_one_sided=self._is_one_sided,
observation_weights=subset_weights,
observations_are_independent=self._observations_are_independent,
_adopt_fourier_coefficients=True,
)
token = _warn_orientation_change.set(False)
try:
replicate = getattr(replicate_connectivity, method)(**method_kwargs)
finally:
_warn_orientation_change.reset(token)
if isinstance(replicate, tuple) or np.iscomplexobj(replicate):
msg = f"jackknife requires a real array result from {method!r}."
raise TypeError(msg)
replicates.append(np.asarray(replicate))
resolved_transformation: Literal[
"identity", "log", "fisher", "fisher_squared", "circular"
]
if transformation == "auto":
if method == "power":
resolved_transformation = "log"
elif method in {"coherence_magnitude", "partial_coherence"}:
# These return magnitude-*squared* coherence, whose
# variance-stabilizing transform is atanh(sqrt(.)), not atanh(.).
resolved_transformation = "fisher_squared"
elif method in _NONNEGATIVE_MAGNITUDE_MEASURES:
# Unsquared magnitudes in [0, 1]: Fisher's atanh applies directly.
resolved_transformation = "fisher"
elif method == "coherence_phase":
resolved_transformation = "circular"
else:
resolved_transformation = "identity"
else:
resolved_transformation = transformation
result = jackknife_confidence_interval(
full_estimate,
np.stack(replicates, axis=0),
confidence_level=confidence_level,
transformation=resolved_transformation,
# PLV's diagonal is 1 by definition, not a saturated estimate.
_saturated_by_construction=(
np.eye(self.n_signals, dtype=bool)
if method == "phase_locking_value"
else False
),
)
if resolved_transformation == "fisher" and method in _NONNEGATIVE_MAGNITUDE_MEASURES:
# tanh maps the atanh-scale interval onto [-1, 1], but a magnitude
# cannot be negative: clamp at 0, as fisher_squared's back-transform
# does for magnitude-squared coherence.
lower, upper = result.confidence_interval
result = replace(
result,
bias_corrected=np.maximum(result.bias_corrected, 0.0),
confidence_interval=(np.maximum(lower, 0.0), upper),
)
return result
[docs]
@_asnumpy
def power(self) -> NDArray[np.floating]:
"""Return the one-sided power spectral density of the signal.
Only the non-negative frequencies are returned, with the interior
positive-frequency bins doubled so that integrating the returned
spectrum over frequency recovers the full signal power (the negative
frequencies of a real signal carry equal power). The DC bin, and the
Nyquist bin for an even FFT length, are not doubled.
For one-sided input (``is_one_sided=True``) the coefficients are
returned as provided: a one-sided transform such as
:class:`~spectral_connectivity.transforms.MorletWavelet` already
scales its coefficients to the one-sided density, and externally
supplied one-sided coefficients keep whatever scale they were given.
Returns
-------
NDArray[floating]
One-sided power spectral density for non-negative frequencies, shape
``(..., n_nonnegative_frequencies, n_signals)``.
Notes
-----
**Range**: [0, ∞). Power spectral density is always non-negative
with no finite upper bound.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> power = connectivity.power()
>>> power.shape # (n_time_windows, n_frequencies, n_signals)
(1, 501, 2)
>>> connectivity.frequencies[[0, -1]] # 0 Hz up to Nyquist in 0.5 Hz steps
array([ 0., 250.])
"""
return self._one_sided_density(self._power, frequency_axis=-2)
[docs]
@_asnumpy
def cross_spectral_density(self) -> NDArray[np.complexfloating]:
"""Return the one-sided cross-spectral density matrix.
The diagonal contains the one-sided power spectral densities returned
by :meth:`power`; off-diagonal entries retain both the amplitude and
relative-phase information between signal pairs. Interior positive
frequency bins are doubled so that the one-sided result has the same
total power as the two-sided spectrum. DC and, for an even FFT length,
Nyquist are not doubled.
Returns
-------
cross_spectral_density : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Notes
-----
The matrix is Hermitian at every time-frequency bin and has physical
units of signal squared per Hz when the input signal has physical
units. Unlike connectivity measures normalized to ``[0, 1]``, its
magnitude has no finite upper bound.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> csd = connectivity.cross_spectral_density()
>>> csd.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # The diagonal is the power spectrum; the matrix is Hermitian.
>>> bool(np.allclose(csd[..., 0, 0].real, connectivity.power()[..., 0]))
True
>>> bool(np.allclose(csd[..., 0, 1], np.conj(csd[..., 1, 0])))
True
"""
return self._one_sided_density(
self._cached_reduced_cross_spectral_matrix, frequency_axis=-3
)
[docs]
@_asnumpy
def coherency(self) -> NDArray[np.complexfloating]:
"""Return the complex-valued linear association between time series.
Computed in the frequency domain.
Returns
-------
complex_coherency : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Complex coherency between all signal pairs; Hermitian in the signal
pair, with its angle given by :meth:`coherence_phase`.
Notes
-----
**Range**: Magnitude :math:`|C_{xy}(f)|` is in [0, 1]; phase is in
[-π, π].
Values lie in the unit disk of the complex plane.
**Phase convention**: ``[..., i, j]`` is ``S_ij / sqrt(S_ii S_jj)``
with ``S_ij = E[X_i conj(X_j)]``, so a positive angle means signal
``i`` leads signal ``j``. :meth:`canonical_coherency` uses the
conjugate convention (``magnitude * exp(-1j * phi)``): with
single-channel groups its score is ``conj`` of this coherency.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> coherency = connectivity.coherency()
>>> coherency.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # At 10 Hz (bin 20) the magnitude is the coupling strength and a positive
>>> # angle at [..., 0, 1] means signal 0 leads signal 1.
>>> round(float(abs(coherency[0, 20, 0, 1])), 1)
0.8
>>> bool(np.angle(coherency[0, 20, 0, 1]) > 0)
True
"""
return self._coherency()
def _coherency(self) -> NDArray[np.complexfloating]:
"""Device-native complex coherency (see the public ``coherency``).
Kept on the active array namespace (``xp``) so internal consumers --
``coherence_magnitude``/``coherence_phase``, ``group_delay``, ``delay``,
``phase_slope_index`` -- operate on device arrays without a premature
host transfer; the public ``coherency`` converts the result to NumPy.
"""
self._warn_single_observation_degenerate("coherency")
complex_coherency = _divide_masking_zero_denominator(
self._nonnegative_cross_spectral_matrix(),
self._nonnegative_pairwise_power_scale(),
"Some signals have (near-)zero power, so coherency is undefined "
"for those pairs and is returned as NaN. This usually indicates "
"a flat/dead channel or all-zero input.",
)
n_signals = self._fourier_coefficients.shape[-1]
diagonal_ind = xp.arange(0, n_signals)
complex_coherency[..., diagonal_ind, diagonal_ind] = xp.nan
return complex_coherency
[docs]
@_asnumpy
def coherence_phase(self) -> NDArray[np.floating]:
"""Return the phase angle of the complex coherency.
Returns
-------
phase : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Phase angles in radians. Positive ``[..., i, j]`` means signal ``i``
leads signal ``j``; the result is antisymmetric in the signal pair.
Notes
-----
**Range**: [-π, π]. Phase angles in radians for complex coherency.
A pure delay ``tau`` (seconds) gives a phase of ``2 * pi * f * tau``
(wrapped into ``[-π, π]``). :meth:`canonical_coherency` reports its
phase in the conjugate convention (``magnitude * exp(-1j * phi)``), so
for single-channel groups its angle is the negative of this one.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> phase = connectivity.coherence_phase()
>>> phase.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # Signal 0 leads, so [..., 0, 1] is positive: about 2 * pi * 10 Hz * 6 ms at 10 Hz.
>>> round(float(phase[0, 20, 0, 1]), 1) # bin 20 is 10 Hz
0.4
>>> round(float(phase[0, 20, 1, 0]), 1)
-0.4
"""
phase: NDArray[np.floating] = xp.angle(self._coherency())
return phase
[docs]
@_asnumpy
def coherence_magnitude(self) -> NDArray[np.floating]:
"""Return the magnitude squared of the complex coherency.
Note that the squared modulus of coherency (originally a complex quantity)
is the magnitude-squared coherence (i.e., the normalized, real component
of coherency). This value should be bounded by 0 and 1.
Returns
-------
magnitude : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Magnitude-squared coherence values (symmetric in the signal pair).
Notes
-----
**Range**: [0, 1]. Implementation may produce tiny numerical excursions
beyond bounds due to floating-point precision.
References
----------
.. [1] Hansson-Sandsten M (2011) Cross-spectrum and coherence function
estimation using time-delayed Thomson multitapers. In: 2011 IEEE
International Conference on Acoustics, Speech and Signal
Processing (ICASSP), pp 4240-4243.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> coherence = connectivity.coherence_magnitude()
>>> coherence.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> round(float(coherence[0, 20, 0, 1]), 1) # strong coupling at 10 Hz (bin 20)
0.6
"""
magnitude = _squared_magnitude(self._coherency())
clipped: NDArray[np.floating] = xp.clip(magnitude, 0, 1)
return clipped
[docs]
@_asnumpy
def imaginary_coherence(self) -> NDArray[np.floating]:
"""Return the normalized imaginary component of the cross-spectrum.
Projects the cross-spectrum onto the imaginary axis to mitigate the
effect of volume-conducted dependencies. Assumes volume-conducted
sources arrive at sensors at the same time, resulting in
a cross-spectrum with phase angle of 0 (perfectly in-phase) or π
(anti-phase) if the sensors are on opposite sides of a dipole
source. With the imaginary coherence, in-phase and anti-phase
associations are set to zero.
Returns
-------
imaginary_coherence_magnitude : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Imaginary coherence magnitudes (symmetric in the signal pair; see
:meth:`imaginary_coherency` for the signed, lead/lag-aware version).
Notes
-----
**Range**: [0, 1]. Magnitude version of imaginary part of coherency.
Raw imaginary component ranges in [-1, 1].
References
----------
.. [1] Nolte, G., Bai, O., Wheaton, L., Mari, Z., Vorbach, S., and
Hallett, M. (2004). Identifying true brain interaction from
EEG data using the imaginary part of coherency. Clinical
Neurophysiology 115, 2292-2307.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> imaginary_coherence = connectivity.imaginary_coherence()
>>> imaginary_coherence.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # The 6 ms lag is not zero-phase, so the imaginary part survives (10 Hz = bin 20).
>>> bool(imaginary_coherence[0, 20, 0, 1] > 0.2)
True
"""
imaginary_coh = xp.abs(
_divide_masking_zero_denominator(
self._nonnegative_cross_spectral_matrix().imag,
self._nonnegative_pairwise_power_scale(),
"Some signals have (near-)zero power, so imaginary coherence is "
"undefined for those pairs and is returned as NaN. This usually "
"indicates a flat/dead channel or all-zero input.",
)
)
# abs()/clip() leave the NaN-masked zero-power entries as NaN.
return xp.clip(imaginary_coh, 0, 1)
[docs]
@_asnumpy
def imaginary_coherency(self) -> NDArray[np.floating]:
"""Return the signed imaginary component of coherency.
This is the signed counterpart of :meth:`imaginary_coherence`, which
returns its magnitude. The sign is antisymmetric across a signal pair
and preserves the pair's phase-lead/phase-lag orientation.
Returns
-------
imaginary_coherency : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Positive ``[..., i, j]`` means signal ``i`` leads signal ``j``.
Notes
-----
**Range**: ``[-1, 1]``. The diagonal and pairs involving zero-power
signals are undefined and returned as NaN, matching :meth:`coherency`.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> imaginary_coherency = connectivity.imaginary_coherency()
>>> imaginary_coherency.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # Signal 0 leads, so [..., 0, 1] is positive and [..., 1, 0] negative (10 Hz).
>>> bool(imaginary_coherency[0, 20, 0, 1] > 0 > imaginary_coherency[0, 20, 1, 0])
True
"""
imaginary = self._coherency().imag
diagonal = xp.arange(self.n_signals)
imaginary[..., diagonal, diagonal] = xp.nan
return imaginary
[docs]
@_asnumpy
def partial_coherence(
self,
regularization: float = TIKHONOV_REGULARIZATION_FACTOR,
) -> NDArray[np.floating]:
"""Return magnitude-squared coherence conditional on all other signals.
Partial coherence is computed by normalizing the off-diagonal elements
of the inverse cross-spectral density (the spectral precision matrix).
It measures the remaining linear association between each pair after
conditioning on every other observed signal.
Parameters
----------
regularization : float, default=1e-12
Non-negative relative diagonal loading applied independently to
each time-frequency cross-spectral matrix before inversion. The
absolute loading is ``regularization * rms(abs(S))``. Increase this
value for statistically rank-deficient or ill-conditioned spectra.
Returns
-------
partial_coherence : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Notes
-----
**Range**: ``[0, 1]``. The diagonal is undefined and returned as NaN.
This undirected measure is distinct from partial directed coherence.
Regularization stabilizes inversion but also changes the estimand, so
analyses should report a non-default value.
The averaged trials x tapers are the observations. With fewer
observations than signals the cross-spectral matrix is rank-deficient
and its (regularized) inverse is dominated by the null space; with
exactly one null direction every partial coherence is forced to 1 for
any data. This case warns.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1000, 20, 1))
>>> # Signals 1 and 2 are noisy copies of signal 0 and share nothing else.
>>> copies = leader + 0.5 * rng.standard_normal((1000, 20, 2))
>>> signals = np.concatenate([leader, copies], axis=-1) # (time, trials, signals)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> partial = connectivity.partial_coherence()
>>> partial.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 3, 3)
>>> # Signals 1 and 2 are coherent, but not once signal 0 is accounted for (10 Hz).
>>> bool(connectivity.coherence_magnitude()[0, 20, 1, 2] > 0.5)
True
>>> bool(partial[0, 20, 1, 2] < 0.05)
True
"""
self._validate_multiple_signals()
regularization = _validated_regularization(regularization)
self._warn_single_observation_degenerate(
"partial_coherence",
"the cross-spectral matrix has rank one and its inverse is fixed by "
"the null space: the estimate is 1 for every pair of two signals and "
"otherwise unrelated to the data",
)
n_observations = self.n_observations
# The inverse of a rank-deficient cross-spectral matrix is dominated by
# its null space: with one null direction the normalized off-diagonal
# precision has unit magnitude for every pair, and with more it is still
# unrelated to the data. (A single observation already warned above.)
if 2 <= n_observations < self.n_signals:
warnings.warn(
f"partial_coherence uses {n_observations} observations (the "
f"averaged trials x tapers), fewer than the {self.n_signals} "
"signals, so the cross-spectral matrix is rank-deficient and its "
"inverse is dominated by the null space: the estimate is forced "
"to 1 for every pair (one missing observation) or is otherwise "
"unrelated to the data. Provide more trials or tapers than "
"signals, or analyze fewer signals.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
# Drop negative frequencies before the per-bin inversion, not after.
cross_spectral_density = self._nonnegative_cross_spectral_matrix()
matrix_rms = xp.sqrt(
xp.mean(
xp.real(xp.conj(cross_spectral_density) * cross_spectral_density),
axis=(-2, -1),
keepdims=True,
)
)
zero_power = matrix_rms <= xp.finfo(matrix_rms.dtype).tiny
if bool(xp.any(zero_power)):
warnings.warn(
"Some time-frequency cross-spectral matrices have zero power, "
"so partial coherence is undefined there and is returned as NaN.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
identity = xp.eye(self.n_signals, dtype=cross_spectral_density.dtype)
safe_spectrum = xp.where(zero_power, identity, cross_spectral_density)
precision = _regularized_inverse(safe_spectrum, regularization=regularization)
precision_diagonal = xp.maximum(
xp.real(xp.diagonal(precision, axis1=-2, axis2=-1)), 0.0
)
denominator = xp.sqrt(
precision_diagonal[..., :, xp.newaxis] * precision_diagonal[..., xp.newaxis, :]
)
partial_coherency = _divide_masking_zero_denominator(
-precision,
denominator,
"Some spectral-precision diagonal entries are (near-)zero, so "
"partial coherence is undefined for those pairs and is returned as NaN.",
)
result = xp.clip(_squared_magnitude(partial_coherency), 0.0, 1.0)
result = xp.where(zero_power, xp.nan, result)
diagonal = xp.arange(self.n_signals)
result[..., diagonal, diagonal] = xp.nan
return result
def _validated_group_indices(
self, group_labels: NDArray[Any]
) -> tuple[NDArray[Any], list[NDArray[np.intp]], NDArray[np.bool_]]:
"""Validate a one-label-per-signal grouping and return its geometry."""
self._validate_multiple_signals()
labels_array = np.asarray(group_labels)
if labels_array.ndim != 1 or len(labels_array) != self.n_signals:
msg = (
f"group_labels must be one-dimensional with length "
f"n_signals ({self.n_signals}), got shape {labels_array.shape}."
)
raise ValueError(msg)
has_missing_label = False
# Check the labels as given: np.asarray turns a NaN among strings into
# the string "nan", which would otherwise become a group of its own.
for group_label in np.asarray(group_labels, dtype=object):
if group_label is None:
has_missing_label = True
break
try:
has_missing_label = not bool(group_label == group_label)
except (TypeError, ValueError):
# An indeterminate equality result (for example a nullable
# scalar) cannot define stable group membership either.
has_missing_label = True
if has_missing_label:
break
if has_missing_label:
msg = "group_labels must not contain missing values such as NaN or None."
raise ValueError(msg)
labels = np.unique(labels_array)
if len(labels) < 2:
msg = "group_labels must define at least two groups."
raise ValueError(msg)
indices = [np.flatnonzero(labels_array == label) for label in labels]
membership = np.asarray(
np.stack([labels_array == label for label in labels]), dtype=bool
)
return labels, indices, membership
[docs]
def canonical_coherence(
self, group_labels: NDArray[np.integer]
) -> tuple[NDArray[np.floating], NDArray[np.integer]]:
"""Return the historical magnitude-squared canonical correlation.
The canonical coherence finds two sets of weights such that the
coherence between the linear combination of group1 and the linear
combination of group2 is maximized.
Parameters
----------
group_labels : array-like, shape (n_signals,)
Links each signal to a group.
Returns
-------
canonical_coherence : array
Shape ``(n_time_windows, n_nonnegative_frequencies, n_groups, n_groups)``.
The maximal coherence for each group pair (symmetric; NaN diagonal).
labels : array, shape (n_groups,)
The sorted unique group labels that correspond to `n_groups`.
Notes
-----
**Range**: [0, 1]. Maximal coherence values are bounded like
coherence magnitude.
Trials x tapers are the observations. A group pair with more signals
than observations has intersecting observation subspaces, which forces
its value to 1 for any data; this case warns.
References
----------
.. [1] Stephen, E.P. (2015). Characterizing dynamically evolving
functional networks in humans with application to speech.
Boston University.
See Also
--------
canonical_coherency
Exact complex, phase-optimised Vidaurre CaCoh with component
filters and patterns.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Group "a": two noisy copies of signal 0; "b": two of its 6 ms-delayed copy.
>>> pair = np.stack([leader[3:], leader[:-3]], axis=-1)
>>> signals = np.repeat(pair, 2, axis=-1) # (time, trials, 4 signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> coherence, labels = connectivity.canonical_coherence(["a", "a", "b", "b"])
>>> coherence.shape # (n_time_windows, n_frequencies, n_groups, n_groups)
(1, 501, 2, 2)
>>> labels
array(['a', 'b'], dtype='<U1')
>>> bool(coherence[0, 20, 0, 1] > 0.5) # strong a-b coupling at 10 Hz (bin 20)
True
"""
labels, group_indices, _ = self._validated_group_indices(group_labels)
# The estimate treats trials x tapers as observations. When a group pair
# has more signals than observations, the two groups' observation
# subspaces must intersect and the value is forced to 1 for any data.
n_observations = (
self._fourier_coefficients.shape[1] * self._fourier_coefficients.shape[2]
)
degenerate_pairs = [
(labels[first], labels[second])
for first, second in combinations(range(len(labels)), 2)
if len(group_indices[first]) + len(group_indices[second]) > n_observations
]
if degenerate_pairs:
warnings.warn(
f"canonical_coherence uses {n_observations} observations (trials x "
f"tapers), fewer than the signals in group pair(s) {degenerate_pairs}; "
"their canonical coherence is forced to 1 regardless of the data. "
"Provide more trials or tapers, or use smaller groups.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
n_frequencies = self._fourier_coefficients.shape[-2]
non_negative_frequencies = xp.arange(
0, self._nonnegative_frequency_count(n_frequencies)
)
fourier_coefficients = self._fourier_coefficients[..., non_negative_frequencies, :]
observation_weights = None
if self._observation_weights is not None:
observation_weights = self._observation_weights[..., non_negative_frequencies, :]
# Canonical correlation is computed from observation covariance
# matrices. Multiplying every observation by sqrt(weight) gives the
# weighted covariance while retaining the existing SVD whitening
# implementation. A shared normalization by sum(weight) cancels
# from the canonical correlation and is therefore unnecessary.
fourier_coefficients = fourier_coefficients * xp.sqrt(observation_weights)
# Group membership was resolved on the host by _validated_group_indices
# (one ascending index array per sorted label); index the device
# coefficients with device index arrays, as ``xp.isin`` over host labels
# fails on CuPy.
normalized_fourier_coefficients = [
_normalize_fourier_coefficients(fourier_coefficients[..., xp.asarray(indices)])
for indices in group_indices
]
n_groups = len(labels)
new_shape = (self.time.size, self.frequencies.size, n_groups, n_groups)
magnitude = _squared_magnitude(
xp.stack(
[
_estimate_canonical_coherence(fourier_coefficients1, fourier_coefficients2)
for fourier_coefficients1, fourier_coefficients2 in combinations(
normalized_fourier_coefficients, 2
)
],
axis=-1,
)
)
if observation_weights is not None:
no_valid_observations = xp.sum(observation_weights[..., 0], axis=(1, 2)) <= 0
magnitude = xp.where(no_valid_observations[..., xp.newaxis], xp.nan, magnitude)
canonical_coherence_magnitude = xp.full(new_shape, xp.nan)
group_combination_ind = xp.array(list(combinations(xp.arange(n_groups), 2)))
canonical_coherence_magnitude[
..., group_combination_ind[:, 0], group_combination_ind[:, 1]
] = magnitude
canonical_coherence_magnitude[
..., group_combination_ind[:, 1], group_combination_ind[:, 0]
] = magnitude
return to_numpy(canonical_coherence_magnitude), to_numpy(labels)
[docs]
def canonical_coherency(
self,
group_labels: NDArray[Any],
*,
rank: int | None = None,
n_components: int = 1,
regularization: float = TIKHONOV_REGULARIZATION_FACTOR,
) -> MultivariateConnectivityResult:
"""Return exact complex canonical coherency (CaCoh) components.
This implements Vidaurre et al.'s phase-optimised CaCoh definition: for
each candidate phase, the real projection of the between-group CSD is
whitened by the real within-group CSDs, and the phase giving the largest
singular value is selected. Scores are complex, with magnitude equal to
the maximised coherence and phase encoded using MNE's
``magnitude * exp(-1j * phi)`` convention.
Unlike the historical :meth:`canonical_coherence`, this method returns
component-resolved scores, spatial filters, and Haufe-style patterns.
Additional components are extracted by CCA-style deflation in the
whitened space: each is sought in the orthogonal complement of the
previous components' whitened directions, so successive component
signals are uncorrelated within each group (``a_j^T Re(Caa) a_k = 0``
for ``j != k``) and every component is invariant to invertible real
mixing of a group's channels.
Parameters
----------
group_labels : array-like, shape (n_signals,)
Label assigning each signal to a group; every unordered pair of
groups is one connection.
rank : int, optional
Retain at most this many within-group whitening directions per group.
``None`` keeps every numerically non-zero direction. Applied
identically to both groups of every connection.
n_components : int, default=1
Number of coherency components to return. A connection whose smaller
group has fewer channels returns NaN for the unavailable components.
regularization : float, default=1e-12
Relative diagonal loading used by the whitening decomposition.
Returns
-------
MultivariateConnectivityResult
Complex ``scores`` of shape ``(..., frequency, connection,
component)`` plus real filters and patterns; see the class docstring.
Notes
-----
**Range**: score magnitudes lie in ``[0, 1]``. The whitening,
singular-value decomposition, and phase optimization (every local
maximum of a 74-point coarse phase grid is refined by a batched Newton
iteration and the best is kept, so near-equal lobes of the phase
objective are resolved unless two lie within about ``2 * pi / 74`` of
each other) are vectorized over the time/frequency axes on the active
``xp`` backend, so this runs on the GPU when GPU support is enabled.
**Phase convention**: a spatial filter and its negative span the same
direction, so the canonical phase is intrinsically defined only modulo
pi. Each filter's sign is fixed so that the largest-magnitude
coefficient of its spatial pattern is positive; with that convention
the score is the conjugate of the canonical coherency between the
positively oriented components, and for single-channel groups it
reduces exactly to the conjugate pairwise coherency.
**Relation to mne-connectivity**: the first component maximizes the
same objective as mne-connectivity's ``cacoh`` and matches it
(magnitude, and phase modulo pi) wherever both optimizers reach the
global maximum. Components ``>= 2`` differ from mne-connectivity 0.9,
which deflates the cross-spectrum with the previous components'
channel-space filters; the whitened-space deflation used here keeps
them uncorrelated within each group and invariant to invertible real
within-group mixing.
References
----------
.. [1] Vidaurre C, et al. (2019) Canonical maximization of coherence: A
novel tool for investigation of neuronal interactions between two
datasets. NeuroImage 201:116009.
.. [2] Haufe S, et al. (2014) On the interpretation of weight vectors of
linear models in multivariate neuroimaging. NeuroImage 87:96-110.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Group "a": two noisy copies of signal 0; "b": two of its 6 ms-delayed copy.
>>> pair = np.stack([leader[3:], leader[:-3]], axis=-1)
>>> signals = np.repeat(pair, 2, axis=-1) # (time, trials, 4 signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> result = connectivity.canonical_coherency(["a", "a", "b", "b"])
>>> result.scores.shape # (n_time_windows, n_frequencies, n_connections, n_components)
(1, 501, 1, 1)
>>> result.connections # one row per group pair
array([['a', 'b']], dtype='<U1')
>>> result.filters.shape # (..., n_connections, n_components, side, n_signals)
(1, 501, 1, 1, 2, 4)
>>> bool(abs(result.scores[0, 20, 0, 0]) > 0.5) # |CaCoh| at 10 Hz (bin 20)
True
"""
return self._multivariate_component_result(
"canonical_coherency",
group_labels,
rank=rank,
n_components=n_components,
regularization=regularization,
)
[docs]
def maximized_imaginary_coherency_components(
self,
group_labels: NDArray[Any],
*,
rank: int | None = None,
n_components: int = 1,
regularization: float = TIKHONOV_REGULARIZATION_FACTOR,
) -> MultivariateConnectivityResult:
"""Return component-resolved MIC scores, filters, and patterns.
The singular vectors of the whitened imaginary between-group CSD are
returned in descending singular-value order. Filters map channel data
to the components; patterns map the components back to channel space.
This is the component-resolved counterpart of the scalar
:meth:`maximized_imaginary_coherency`.
Parameters
----------
group_labels : array-like, shape (n_signals,)
Label assigning each signal to a group; every unordered pair of
groups is one connection.
rank : int, optional
Retain at most this many within-group whitening directions per group.
``None`` keeps every numerically non-zero direction.
n_components : int, default=1
Number of singular components to return. A connection whose smaller
group has fewer channels returns NaN for the unavailable components.
regularization : float, default=1e-12
Relative diagonal loading used by the whitening decomposition.
Returns
-------
MultivariateConnectivityResult
Real ``scores`` of shape ``(..., frequency, connection, component)``
plus filters and patterns; see the class docstring.
Notes
-----
**Range**: scores lie in ``[0, 1]``. The whitening and singular-value
decomposition are vectorized over the time/frequency axes on the active
``xp`` backend, so this runs on the GPU when GPU support is enabled.
References
----------
.. [1] Ewald A, et al. (2012) Estimating true brain connectivity from EEG/
MEG data invariant to linear and static transformations in sensor
space. NeuroImage 60(1):476-488.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Group "a": two noisy copies of signal 0; "b": two of its 6 ms-delayed copy.
>>> pair = np.stack([leader[3:], leader[:-3]], axis=-1)
>>> signals = np.repeat(pair, 2, axis=-1) # (time, trials, 4 signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> labels = ["a", "a", "b", "b"]
>>> result = connectivity.maximized_imaginary_coherency_components(labels)
>>> result.scores.shape # (n_time_windows, n_frequencies, n_connections, n_components)
(1, 501, 1, 1)
>>> result.connections # one row per group pair
array([['a', 'b']], dtype='<U1')
>>> result.patterns.shape # (..., n_connections, n_components, side, n_signals)
(1, 501, 1, 1, 2, 4)
>>> bool(result.scores[0, 20, 0, 0] > 0.2) # lagged a-b coupling at 10 Hz (bin 20)
True
"""
return self._multivariate_component_result(
"maximized_imaginary_coherency_components",
group_labels,
rank=rank,
n_components=n_components,
regularization=regularization,
)
def _multivariate_component_result(
self,
method: Literal["canonical_coherency", "maximized_imaginary_coherency_components"],
group_labels: NDArray[Any],
*,
rank: int | None,
n_components: int,
regularization: float,
) -> MultivariateConnectivityResult:
"""Compute rich CaCoh/MIC results from the expected CSD."""
labels, group_indices, membership = self._validated_group_indices(group_labels)
rank = _validated_rank(rank)
if not is_positive_integer(n_components):
msg = f"n_components must be a positive integer, got {n_components!r}."
raise ValueError(msg)
regularization = _validated_regularization(regularization)
rank_cap = rank if rank is not None else self.n_signals
# ``pairs`` is the single connection ordering shared by the capacities,
# the ``connections`` labels, and the main loop below.
pairs = list(combinations(range(len(labels)), 2))
# The number of components a connection can support is bounded by the two
# groups *in that connection* (and the requested rank), not by the
# smallest group overall. Compute the per-connection capacity and reject
# only when no connection can supply n_components; connections whose
# groups are smaller return NaN for the unavailable components.
connection_capacities = [
min(len(group_indices[first]), len(group_indices[second]), rank_cap)
for first, second in pairs
]
max_components = max(connection_capacities)
if n_components > max_components:
msg = (
f"n_components ({n_components}) must not exceed the largest "
f"per-connection group rank/size ({max_components}); no group pair "
f"is large enough to supply that many components."
)
raise ValueError(msg)
spectrum = self._nonnegative_cross_spectral_matrix()
leading_shape = spectrum.shape[:-3]
n_frequencies = spectrum.shape[-3]
connections = np.asarray([(labels[first], labels[second]) for first, second in pairs])
n_connections = len(connections)
score_dtype = xp.complex128 if method == "canonical_coherency" else float
scores = xp.full(
(*leading_shape, n_frequencies, n_connections, n_components),
xp.nan,
dtype=score_dtype,
)
projection_shape = (
*leading_shape,
n_frequencies,
n_connections,
n_components,
2,
self.n_signals,
)
filters = xp.full(projection_shape, xp.nan, dtype=float)
patterns = xp.full(projection_shape, xp.nan, dtype=float)
component_fn = (
_canonical_coherency_components
if method == "canonical_coherency"
else _mic_components
)
any_phantom = False
for connection_index, (first, second) in enumerate(pairs):
first_indices = group_indices[first]
second_indices = group_indices[second]
n_first = len(first_indices)
# This connection can only supply as many components as its smaller
# group (and the rank cap); the rest stay NaN in the pre-filled array.
component_count = min(n_components, connection_capacities[connection_index])
combined = xp.asarray(np.concatenate((first_indices, second_indices)))
# Sub-CSD over every leading/frequency bin at once: (..., freq, m, m).
subsystem = spectrum[..., combined[:, xp.newaxis], combined[xp.newaxis, :]]
# A non-finite bin (e.g. a dead channel) would make the batched
# eigendecomposition fail; compute a placeholder there and mask it
# back to NaN afterward, matching the old per-bin skip.
finite_bin = xp.all(xp.isfinite(subsystem), axis=(-2, -1))
identity = xp.eye(combined.shape[0], dtype=subsystem.dtype)
safe = xp.where(finite_bin[..., xp.newaxis, xp.newaxis], subsystem, identity)
local_scores, (filter_a, filter_b), (pattern_a, pattern_b), rank_here = (
component_fn(
safe[..., :n_first, :n_first],
safe[..., :n_first, n_first:],
safe[..., n_first:, n_first:],
rank=rank,
n_components=component_count,
regularization=regularization,
)
)
valid = finite_bin[..., xp.newaxis]
scores[..., connection_index, :component_count] = xp.where(
valid, local_scores, xp.nan
)
for side, (indices, side_filter, side_pattern) in enumerate(
(
(first_indices, filter_a, pattern_a),
(second_indices, filter_b, pattern_b),
)
):
masked_filter = xp.where(valid[..., xp.newaxis], side_filter, xp.nan)
masked_pattern = xp.where(valid[..., xp.newaxis], side_pattern, xp.nan)
signal_index = xp.asarray(indices)
# Assign one component at a time: a single trailing fancy index
# (the group's channels) with only integer indices before it keeps
# the scattered axis at the end, avoiding NumPy's mixed
# slice/advanced-index dimension reordering.
for component in range(component_count):
filters[..., connection_index, component, side, signal_index] = (
masked_filter[..., component]
)
patterns[..., connection_index, component, side, signal_index] = (
masked_pattern[..., component]
)
# A rank-deficient within-group block (collinear/duplicated channels)
# supplies fewer directions than requested; the extra "phantom"
# components come back with a ~0 score and an all-zero filter. Flag it
# via the eigenvalue-based rank (scale-invariant), not the filter norm.
if bool(xp.any(finite_bin & (rank_here < component_count))):
any_phantom = True
if any_phantom:
warnings.warn(
f"{method}: some requested components fall in the null space of a "
"rank-deficient within-group cross-spectrum (collinear or "
"duplicated channels), so they are returned with a zero score and "
"an all-zero spatial filter. Reduce n_components or pass an "
"explicit rank to avoid these phantom components.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
return MultivariateConnectivityResult(
method=method,
scores=to_numpy(scores),
connections=connections,
group_labels=np.asarray(labels),
group_membership=np.asarray(membership, dtype=bool),
filters=to_numpy(filters),
patterns=to_numpy(patterns),
)
[docs]
def maximized_imaginary_coherency(
self,
group_labels: NDArray[np.integer],
rank: int | None = None,
regularization: float = TIKHONOV_REGULARIZATION_FACTOR,
) -> tuple[NDArray[np.floating], NDArray[np.integer]]:
"""Return maximized imaginary coherency (MIC) between signal groups.
Each group's real within-group cross-spectrum is whitened before the
largest singular value of the between-group imaginary cross-spectrum is
taken. This makes the result invariant to invertible, static real-valued
mixing within either group.
Parameters
----------
group_labels : array-like, shape (n_signals,)
Label assigning each signal to a group.
rank : int, optional
Retain at most this many within-group whitening components. ``None``
retains every numerically non-zero component independently per bin.
regularization : float, default=1e-12
Relative diagonal loading used by the whitening decomposition.
Returns
-------
mic : array
Shape ``(..., n_nonnegative_frequencies, n_groups, n_groups)``.
labels : array, shape (n_groups,)
Sorted unique group labels.
Notes
-----
**Range**: ``[0, 1]``. The diagonal is returned as NaN.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Group "a": two noisy copies of signal 0; "b": two of its 6 ms-delayed copy.
>>> pair = np.stack([leader[3:], leader[:-3]], axis=-1)
>>> signals = np.repeat(pair, 2, axis=-1) # (time, trials, 4 signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> mic, labels = connectivity.maximized_imaginary_coherency(["a", "a", "b", "b"])
>>> mic.shape # (n_time_windows, n_frequencies, n_groups, n_groups)
(1, 501, 2, 2)
>>> labels
array(['a', 'b'], dtype='<U1')
>>> bool(mic[0, 20, 0, 1] > 0.2) # lagged a-b coupling at 10 Hz (bin 20)
True
"""
return self._group_imaginary_coherency(
group_labels,
lambda whitened: xp.clip(
xp.linalg.svd(whitened, full_matrices=False, compute_uv=False)[..., 0],
0.0,
1.0,
),
rank=rank,
regularization=regularization,
)
[docs]
def multivariate_interaction_measure(
self,
group_labels: NDArray[np.integer],
rank: int | None = None,
regularization: float = TIKHONOV_REGULARIZATION_FACTOR,
) -> tuple[NDArray[np.floating], NDArray[np.integer]]:
"""Return the multivariate interaction measure (MIM) between groups.
MIM sums the squared singular values of the whitened imaginary
cross-spectrum, incorporating every phase-lagged interaction component
rather than only the strongest component returned by MIC.
Parameters
----------
group_labels : array-like, shape (n_signals,)
Label assigning each signal to a group.
rank : int, optional
Retain at most this many within-group whitening components. ``None``
retains every numerically non-zero component independently per bin.
regularization : float, default=1e-12
Relative diagonal loading used by the whitening decomposition.
Returns
-------
mim : array
Shape ``(..., n_nonnegative_frequencies, n_groups, n_groups)``.
labels : array, shape (n_groups,)
Sorted unique group labels.
Notes
-----
**Range**: ``[0, min(rank_group_1, rank_group_2)]``; unlike MIC, MIM can
exceed one. The diagonal is returned as NaN.
References
----------
.. [1] Ewald, A., Marzetti, L., Zappasodi, F., Meinecke, F.C., and
Nolte, G. (2012). Estimating true brain connectivity from
EEG/MEG data invariant to linear and static transformations in
sensor space. NeuroImage 60, 476-488.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Group "a": two noisy copies of signal 0; "b": two of its 6 ms-delayed copy.
>>> pair = np.stack([leader[3:], leader[:-3]], axis=-1)
>>> signals = np.repeat(pair, 2, axis=-1) # (time, trials, 4 signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> mim, labels = connectivity.multivariate_interaction_measure(["a", "a", "b", "b"])
>>> mim.shape # (n_time_windows, n_frequencies, n_groups, n_groups)
(1, 501, 2, 2)
>>> labels
array(['a', 'b'], dtype='<U1')
>>> bool(mim[0, 20, 0, 1] > 0.05) # lagged a-b interaction at 10 Hz (bin 20)
True
"""
return self._group_imaginary_coherency(
group_labels,
lambda whitened: xp.sum(whitened**2, axis=(-2, -1)),
rank=rank,
regularization=regularization,
)
def _group_imaginary_coherency(
self,
group_labels: NDArray[np.integer],
reduce: Callable[[BackendArray], BackendArray],
*,
rank: int | None,
regularization: float,
) -> tuple[NDArray[np.floating], NDArray[np.integer]]:
"""Reduce whitened imaginary CSD blocks to a group-by-group measure.
Shared by MIC and MIM. ``reduce`` maps each group pair's whitened
imaginary cross-spectrum, shape ``(..., n_first, n_second)``, to its
value, shape ``(...)``, which fills both ``[first, second]`` and
``[second, first]`` of the returned ``(..., n_groups, n_groups)`` array
(diagonal NaN); the sorted labels are returned alongside. Invalid bins (for example the NaN edges of
an ``edge_mode="nan"`` Morlet transform) are replaced with a safe value
before the batched eigendecomposition/SVD so they cannot fail. Validity
is tracked per group and combined per connection, so a NaN confined to
one group only invalidates connections that include it -- an unrelated
connection between two healthy groups is preserved.
"""
labels, numpy_group_indices, _ = self._validated_group_indices(group_labels)
rank = _validated_rank(rank)
regularization = _validated_regularization(regularization)
spectrum = self._nonnegative_cross_spectral_matrix()
group_indices = [xp.asarray(indices) for indices in numpy_group_indices]
# Whiten each group's within-block, substituting the identity at that
# group's non-finite bins so the eigendecomposition converges there.
inverse_square_roots = []
group_finite = []
for indices in group_indices:
within = spectrum[..., indices[:, xp.newaxis], indices[xp.newaxis, :]].real
finite = xp.all(xp.isfinite(within), axis=(-2, -1))
identity = xp.eye(indices.shape[0], dtype=within.dtype)
safe_within = xp.where(finite[..., xp.newaxis, xp.newaxis], within, identity)
inverse_square_roots.append(
_batched_inverse_square_root(
safe_within,
rank=rank,
regularization=regularization,
)[0]
)
group_finite.append(finite)
result = xp.full(
(*spectrum.shape[:-2], len(labels), len(labels)), xp.nan, dtype=spectrum.real.dtype
)
for first, second in combinations(range(len(labels)), 2):
first_indices = group_indices[first]
second_indices = group_indices[second]
# A connection is valid only where both of its groups are finite;
# zero the between-block elsewhere so the SVD stays finite.
connection_finite = group_finite[first] & group_finite[second]
between = spectrum[
...,
first_indices[:, xp.newaxis],
second_indices[xp.newaxis, :],
].imag
between = xp.where(connection_finite[..., xp.newaxis, xp.newaxis], between, 0.0)
whitened = xp.matmul(
xp.matmul(inverse_square_roots[first], between),
inverse_square_roots[second],
)
value = xp.where(connection_finite, reduce(whitened), xp.nan)
result[..., first, second] = value
result[..., second, first] = value
return to_numpy(result), to_numpy(labels)
[docs]
def global_coherence(
self,
max_rank: int = 1,
max_workspace_elements: int = GLOBAL_COHERENCE_BATCH_CHUNK_ELEMENTS,
) -> tuple[NDArray[np.floating], NDArray[np.complexfloating]]:
"""Find linear combinations that capture the most coherent power.
The linear combinations of signals that capture the most coherent
power at each frequency and time window.
This is a frequency domain analog of PCA over signals at a given
frequency/time window.
Parameters
----------
max_rank : int, default=1
The number of components to keep (like the number of PC dimensions).
max_workspace_elements : int, default=16_000_000
Approximate working-set target, in array elements, for the batched
decomposition: frequency bins are processed in chunks sized so the
main intermediates stay near this many complex elements (the default
~16M ≈ 256 MB of complex128). It is a soft target, not a hard memory
cap — it counts the dominant per-bin intermediates, not the outputs or
LAPACK's internal workspace, and it never goes below one bin per chunk,
so actual peak memory is somewhat higher. Lower it to reduce peak
memory on a constrained CPU or GPU (at the cost of more, smaller
chunks); the default favors speed and does not change the result.
Ignored on the per-bin fallback path used for a large decomposition
dimension.
Returns
-------
global_coherence : ndarray
Shape (n_time_windows, n_fft_samples, n_components).
The fraction of total coherent power captured by each component
(eigenvalue of the cross-spectral matrix divided by the sum of all
eigenvalues), ordered strongest component first.
unnormalized_global_coherence : ndarray
Shape (n_time_windows, n_fft_samples, n_signals, n_components).
The global coherence vectors (left singular vectors).
Notes
-----
**Frequency axis**: unlike every other public measure, which returns
only the non-negative frequencies, this method returns all
``n_fft_samples`` bins of the (two-sided) transform in FFT order;
index it with ``all_frequencies`` rather than ``frequencies``.
**Range**: [0, 1]. Each value is the fraction of total coherent power
in that component, so the measure is scale-invariant and the components
sum to at most 1.
**Algorithm**: when the number of estimates (``n_trials * n_tapers``)
is at least ``n_signals`` and ``n_signals`` is small
(``<= 64``), the components are obtained from an eigendecomposition of
the ``(n_signals, n_signals)`` cross-spectral matrix ``A @ Aᴴ`` rather
than a singular value decomposition of ``A``. This is substantially
faster but squares the condition number, so for a nearly rank-deficient
cross-spectral matrix (near-duplicate channels) the *weakest* returned
components (large ``max_rank``) may lose relative precision. The dominant
component(s) — the usual use of this measure — are unaffected. A thin
matrix (fewer estimates than signals) uses the economy SVD directly.
References
----------
.. [1] Cimenser, A., Purdon, P.L., Pierce, E.T., Walsh, J.L.,
Salazar-Gomez, A.F., Harrell, P.G., Tavares-Stoeckel, C.,
Habeeb, K., and Brown, E.N. (2011). Tracking brain states under
general anesthesia by using global coherence analysis.
Proceedings of the National Academy of Sciences 108, 8832-8837.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> global_coherence, vectors = connectivity.global_coherence(max_rank=1)
>>> # Unlike the pairwise measures, the frequency axis spans all FFT bins.
>>> global_coherence.shape # (n_time_windows, n_fft_samples, n_components)
(1, 1000, 1)
>>> vectors.shape # (n_time_windows, n_fft_samples, n_signals, n_components)
(1, 1000, 2, 1)
>>> # One component captures most of the power of the coupled pair (10 Hz = bin 20).
>>> bool(global_coherence[0, 20, 0] > 0.8)
True
"""
self._validate_multiple_signals()
_, n_trials, n_tapers, _, n_signals = self._fourier_coefficients.shape
# A rank-r decomposition of the (n_signals, n_trials * n_tapers)
# coefficient matrix has at most min(n_signals, n_trials * n_tapers)
# non-trivial components. Requesting more than that would crash svds or
# (in the dense branch) broadcast a single component into duplicates, so
# clamp the requested rank to what is realizable.
max_available_rank = min(n_signals, n_trials * n_tapers)
if max_rank > max_available_rank:
warnings.warn(
f"max_rank={max_rank} exceeds the number of available "
f"global-coherence components "
f"(min(n_signals, n_trials * n_tapers) = {max_available_rank}); "
f"clamping to {max_available_rank}.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
max_rank = max_available_rank
# Must be a genuine positive integer: it is a floor-divided into a chunk
# size, so a float (NaN/inf included) or bool would either pick a
# nonsensical chunk or blow up later inside range(). bool is an int
# subclass, so reject it explicitly.
if not is_positive_integer(max_workspace_elements):
msg = (
f"max_workspace_elements must be a positive integer (e.g. "
f"1_000_000), got {max_workspace_elements!r}. It bounds the memory "
f"budget, in array elements, for global_coherence's batched "
f"decomposition; lower it to reduce peak memory."
)
raise ValueError(msg)
global_coherence, unnormalized_global_coherence = _global_coherence(
self._fourier_coefficients,
self._observation_weights,
max_rank,
max_workspace_elements,
)
if xp.any(xp.isnan(global_coherence)):
warnings.warn(
"Some time-frequency bins have (near-)zero total power, so "
"global coherence is undefined there and is returned as NaN. "
"This usually indicates a flat/dead channel or all-zero input.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
return to_numpy(global_coherence), to_numpy(unnormalized_global_coherence)
def _phase_locking_value(self) -> NDArray[np.complexfloating]:
# Normalize each Fourier coefficient to unit magnitude, then reuse the
# batched reduced cross-spectral matmul: because
# (z_i conj(z_j)) / |z_i conj(z_j)| = (z_i / |z_i|) conj(z_j / |z_j|),
# the mean over observations of the normalized per-observation
# cross-spectrum equals the reduced cross-spectral matrix of the
# unit-normalized coefficients -- with no per-observation outer product,
# so peak memory is O(observations * signals) rather than
# O(observations * signals**2). Kept on the active namespace (xp); the
# public ``phase_locking_value`` wrapper converts to NumPy.
self._validate_multiple_signals()
self._warn_single_observation_degenerate("phase_locking_value")
# Normalize at the computation dtype (``self._dtype``, complex128 by
# default): the previous materialized path formed the outer product at
# that dtype, so normalizing complex64 inputs at their own precision here
# would let float32 rounding push the unit magnitudes -- and thus the
# averaged PLV/PPC -- slightly past 1. copy=False avoids a copy when the
# dtype already matches (the division below allocates a fresh array).
# Only the non-negative frequencies are reported, so only those are
# normalized and reduced.
coefficients = self._nonnegative_fourier_coefficients()
magnitude = xp.abs(coefficients)
zero_magnitude = magnitude == 0
if bool(xp.any(zero_magnitude)):
warnings.warn(
"Some cross-spectrum entries have zero magnitude (e.g. a "
"flat/dead channel or all-zero input at a taper/trial), so "
"the phase-locking normalization z / |z| is undefined there "
"and is returned as NaN.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
# z / |z| is undefined where |z| == 0; those coefficients are NaN
# (rather than leaking a RuntimeWarning). A NaN coefficient at any
# observation makes every pair involving it reduce to NaN, matching the
# previous per-observation path where a zero-magnitude cross-spectrum
# entry became NaN before averaging.
normalized: NDArray[np.complexfloating] = _divide_where(
coefficients, magnitude, ~zero_magnitude, xp.nan
)
return self._reduced_cross_spectral_matrix(normalized)
[docs]
@_asnumpy
def phase_locking_value(self) -> NDArray[np.floating]:
"""Return the cross-spectrum with power scaled to magnitude 1.
The phase locking value attempts to mitigate power differences
between realizations (tapers or trials) by treating all values of
the cross-spectrum as the same power. This has the effect of
downweighting high power realizations and upweighting low power
realizations.
Returns
-------
phase_locking_value : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Phase locking values between all signal pairs.
Notes
-----
**Range**: [0, 1]. 0 indicates random phases; 1 indicates
constant phase difference.
References
----------
.. [1] Lachaux, J.-P., Rodriguez, E., Martinerie, J., Varela, F.J.,
and others (1999). Measuring phase synchrony in brain
signals. Human Brain Mapping 8, 194-208.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> plv = connectivity.phase_locking_value()
>>> plv.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> bool(plv[0, 20, 0, 1] > 0.5) # consistent phase difference at 10 Hz (bin 20)
True
"""
# Clip to the documented [0, 1] range: |mean of unit-magnitude entries|
# is <= 1 mathematically, but floating-point rounding can leave it a few
# ulp above 1 (matching the bounds clipping coherence_magnitude applies).
return xp.clip(xp.abs(self._phase_locking_value()), 0.0, 1.0)
[docs]
@_asnumpy
def corrected_imaginary_phase_locking_value(self) -> NDArray[np.floating]:
"""Return corrected imaginary phase-locking value (ciPLV).
ciPLV removes the contribution of zero- and pi-lag phase locking while
correcting the imaginary PLV for the reduction in its attainable range.
Returns
-------
corrected_imaginary_phase_locking_value : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Notes
-----
**Range**: ``[0, 1]``. Exact zero- or pi-lag locking has a zero
numerator and denominator and is defined as zero.
References
----------
.. [1] Bruña, R., Maestú, F., and Pereda, E. (2018). Phase locking
value revisited: teaching new tricks to an old dog. Journal of
Neural Engineering 15, 056011.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> ciplv = connectivity.corrected_imaginary_phase_locking_value()
>>> ciplv.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # The 6 ms lag is not zero-phase, so lagged locking remains (10 Hz = bin 20).
>>> bool(ciplv[0, 20, 0, 1] > 0.2)
True
"""
complex_plv = self._phase_locking_value()
numerator = xp.abs(complex_plv.imag)
denominator_squared = xp.maximum(0.0, 1.0 - complex_plv.real**2)
denominator = xp.sqrt(denominator_squared)
# The only mathematically valid zero-denominator case also has a zero
# numerator (perfect zero- or pi-lag locking). Define that limit as 0.
# A NaN PLV (zero-power channel) must stay NaN rather than fall into
# that zero limit, matching every other phase-locking measure.
nonzero = denominator > xp.finfo(denominator.dtype).tiny
result = _divide_where(numerator, denominator, nonzero, 0.0)
result[xp.isnan(complex_plv)] = xp.nan
return xp.clip(result, 0.0, 1.0)
@cached_property
def _imaginary_moment_cache(self) -> dict[str, BackendArray]:
"""Lazily populated phase-lag moments tied to the current inputs."""
return {}
def _has_no_phase_lag(self, mean_absolute: NDArray[np.floating]) -> NDArray[np.bool_]:
"""Pairs whose imaginary cross-spectrum is zero up to rounding.
For in-phase signals (a channel and a scaled copy of it) each
observation's ``Im(X_i conj(X_j))`` is rounding noise of order
``eps * |X_i| |X_j|``, not exactly 0, so ``E[|Im S_ij|]`` is compared
with the pair's power scale ``sqrt(P_i P_j)`` rather than with 0. The
test is relative, so it does not depend on the signal's amplitude.
Parameters
----------
mean_absolute : array, shape (..., n_nonnegative_frequencies, n_signals, n_signals)
``E[|Im S_ij|]`` from :meth:`_imaginary_cross_spectrum_moments`, at
the non-negative frequencies.
Returns
-------
no_lag : array of bool, shape (..., n_nonnegative_frequencies, n_signals, n_signals)
"""
tolerance = _ZERO_PHASE_LAG_EPSILONS * xp.finfo(mean_absolute.dtype).eps
power_scale = self._nonnegative_pairwise_power_scale()
no_lag: NDArray[np.bool_] = mean_absolute <= tolerance * power_scale
return no_lag
def _imaginary_cross_spectrum_moments(
self, *keys: str
) -> tuple[NDArray[np.floating], ...]:
"""Reduced moments of the per-observation imaginary cross-spectrum.
The phase-lag-index family (``phase_lag_index``,
``weighted_phase_lag_index``, ``debiased_squared_phase_lag_index``, and
``debiased_squared_weighted_phase_lag_index``) each average a function --
``sign``, identity, ``abs`` or square -- of the imaginary part of the
per-observation cross-spectral matrix, with the diagonal zeroed. This
returns the requested reduced moments from a cache, computing any that
are missing from signal-row tiles of the observation-level
cross-spectrum. Only the non-negative bins are reduced, and each tile
covers targets from its first row on; the strict lower triangle is then
filled by pair symmetry (``_IMAGINARY_MOMENT_PAIR_SYMMETRY``).
Which moments a call computes depends on the regime. Without
observation weights and with at least
``PHASE_LAG_ALL_MOMENTS_MIN_OBSERVATIONS`` averaged observations,
forming the tiles (memory traffic over every observation) dominates the
cost and each reduction is a single further pass over a tile, so the
first request reduces all four moments from one formation; that costs
less than forming the tiles again for a later measure that needs
different moments. The four reduced moments (each ``(n_time_windows,
n_frequencies, n_signals, n_signals)`` for the default expectation) are
then retained until :meth:`clear_cache`. In this regime a single
measure is also faster and has a lower peak than before (741 -> 536 MB
of traced allocation for one measure on a 100-trial, 5-taper,
1001-bin, 32-signal spectrum), because the tiles reuse two buffers
instead of fresh temporaries.
Otherwise only the missing requested moments are computed. With few
observations (e.g. a time-resolved single-trial spectrum with 3
tapers) each window-resolved moment is nearly as large as the tile, so
writing and retaining two extra ones cost more than re-forming the
tiles later (a lone ``phase_lag_index`` on 1199 windows x 3 tapers x
100 bins x 32 signals took 1.49x as long and 1.5x the peak memory).
With observation weights each moment is a full weighted
:meth:`_expectation` with tile-sized temporaries; computing all four
made a lone weighted measure 7-14% (``phase_lag_index``) and 33-54%
(``weighted_phase_lag_index``) slower on a Hann-weighted Morlet
transform (4000 samples, 40 trials, 16 signals, 30 frequencies).
Each tile is reduced immediately (see :meth:`_reduce_phase_lag_tile`)
into two workspace buffers allocated once per call, avoiding an
observation-resolved ``n_signals**2`` intermediate. The reduced
``n_signals**2`` outputs are unavoidable. The cached moments are
invalidated with the other cached intermediates and are treated as
read-only (callers copy before any in-place edit).
Parameters
----------
*keys : str
Any of ``"sign"``, ``"imaginary"``, ``"absolute"``, ``"squared"`` for
``E[sign(Im)]``, ``E[Im]``, ``E[|Im|]`` and ``E[Im**2]``.
Returns
-------
tuple of arrays, each shape (..., n_nonnegative_frequencies, n_signals, n_signals)
The requested moments, in the order of ``keys``, at the non-negative
frequencies the phase-lag measures report.
"""
self._validate_multiple_signals()
cache = self._imaginary_moment_cache
if any(key not in cache for key in keys):
coefficients = self._nonnegative_fourier_coefficients()
n_signals = coefficients.shape[-1]
kept_observation_axes = [
axis for axis in range(3) if axis not in self._expectation_axes
]
result_shape = (
*[coefficients.shape[axis] for axis in kept_observation_axes],
coefficients.shape[-2],
n_signals,
n_signals,
)
real_dtype = coefficients.real.dtype
# Without weights and with many observations all four moments are
# cheap single-pass reductions of the tile, so they are computed
# together (see PHASE_LAG_ALL_MOMENTS_MIN_OBSERVATIONS). Otherwise
# only the missing requested ones are: with weights each is a full
# ``_expectation`` with tile-sized temporaries, and with few
# observations the reduced moments are nearly as large as the tile.
# Filled here and cached only once complete, so an error partway
# through leaves no uninitialized moments behind.
if (
self._observation_weights is None
and self.n_observations >= PHASE_LAG_ALL_MOMENTS_MIN_OBSERVATIONS
):
computed_keys = list(_IMAGINARY_MOMENTS)
else:
computed_keys = list(dict.fromkeys(key for key in keys if key not in cache))
moments = {key: xp.empty(result_shape, dtype=real_dtype) for key in computed_keys}
observation_frequency_elements = int(np.prod(coefficients.shape[:-1]))
elements_per_source = max(1, observation_frequency_elements * n_signals)
signals_per_block = max(
1,
min(
n_signals,
PHASE_LAG_INDEX_MAX_WORKSPACE_ELEMENTS // elements_per_source,
),
)
# Nonlinear sign/abs/square transforms prevent contracting the
# observation axes before the outer product. Form only a source-row
# tile at a time, reduce it immediately, and write the small result.
# Im(X_i conj(X_j)) = Im(X_i) Re(X_j) - Re(X_i) Im(X_j), formed in
# real arithmetic rather than as a complex product. Each tile pairs
# its source rows only with targets from ``start`` on; the strict
# lower triangle is filled by pair symmetry below.
# Contiguous copies: transforms leave the signal axis slowest, which
# makes the tiles' broadcast products stride badly through memory.
real = xp.ascontiguousarray(coefficients.real)
imag = xp.ascontiguousarray(coefficients.imag)
# Each tile is a contiguous view of the leading part of these flat
# buffers, so later (narrower) tiles reuse the first tile's memory.
workspace_size = observation_frequency_elements * signals_per_block * n_signals
tile_buffer = xp.empty(workspace_size, dtype=real_dtype)
scratch_buffer = xp.empty(workspace_size, dtype=real_dtype)
observation_frequency_shape = coefficients.shape[:-1]
block_diagonal = xp.arange(signals_per_block)
# Contracts a tile's observation axes in the squared moment, e.g.
# "abcdef,abcdef->adef" (time, trial, taper, frequency, row, column).
letters = "abcdef"
kept = "".join(
letter
for axis, letter in enumerate(letters)
if axis not in self._expectation_axes
)
squared_subscripts = f"{letters},{letters}->{kept}"
for start in range(0, n_signals, signals_per_block):
stop = min(n_signals, start + signals_per_block)
tile_shape = (*observation_frequency_shape, stop - start, n_signals - start)
tile_size = (
observation_frequency_elements * (stop - start) * (n_signals - start)
)
imaginary = tile_buffer[:tile_size].reshape(tile_shape)
scratch = scratch_buffer[:tile_size].reshape(tile_shape)
xp.multiply(
imag[..., start:stop, xp.newaxis],
real[..., xp.newaxis, start:],
out=imaginary,
)
xp.multiply(
real[..., start:stop, xp.newaxis],
imag[..., xp.newaxis, start:],
out=scratch,
)
xp.subtract(imaginary, scratch, out=imaginary)
local_diagonal = block_diagonal[: stop - start]
imaginary[..., local_diagonal, local_diagonal] = 0
self._reduce_phase_lag_tile(
imaginary, scratch, moments, start, stop, squared_subscripts
)
# Overwrite the strict lower triangle from the upper one: the pairs
# each tile skipped and the in-tile lower entries alike.
rows, columns = xp.tril_indices(n_signals, k=-1)
for key, reduced in moments.items():
reduced[..., rows, columns] = (
_IMAGINARY_MOMENT_PAIR_SYMMETRY[key] * reduced[..., columns, rows]
)
cache.update(moments)
return tuple(cache[key] for key in keys)
def _reduce_phase_lag_tile(
self,
imaginary: BackendArray,
scratch: BackendArray,
moments: dict[str, BackendArray],
start: int,
stop: int,
squared_subscripts: str,
) -> None:
"""Average the four phase-lag moments of one tile into ``moments``.
Parameters
----------
imaginary : array, shape (n_time_windows, n_trials, n_tapers, n_nonnegative_frequencies, stop - start, n_signals - start)
``Im(X_i conj(X_j))`` per observation for source rows
``start:stop`` and targets ``start:``, with the diagonal zeroed.
scratch : array, same shape as ``imaginary``
Workspace; overwritten.
moments : dict of str to array, each shape (..., n_nonnegative_frequencies, n_signals, n_signals)
Reduced moments to fill, keyed as ``_IMAGINARY_MOMENTS``: all four
or the requested subset (see :meth:`_imaginary_cross_spectrum_moments`).
The block
``[..., start:stop, start:]`` of each is written.
start, stop : int
The tile's source rows.
squared_subscripts : str
``einsum`` subscripts contracting ``imaginary * imaginary`` over the
averaged observation axes (used without observation weights).
Notes
-----
Without observation weights every moment is a plain mean, reduced in a
single pass over the tile. The sign moment is counted from boolean
comparisons, which are False for NaN, so bins with any NaN observation
are set to NaN explicitly, as the mean of ``sign`` would give. The
``imaginary`` moment is the mean of the same per-observation values as
``absolute`` rather than the imaginary part of the (differently
rounded) cross-spectral matrix, so their ratio in
``weighted_phase_lag_index`` is exactly 1 in magnitude at a constant
lag above the no-lag guard (:meth:`_has_no_phase_lag`, below which it
is defined as 0). With observation weights each moment goes through
:meth:`_expectation`, which needs the per-observation values.
"""
block = (..., slice(start, stop), slice(start, None))
if self._observation_weights is not None:
for key, reduced in moments.items():
reduced[block] = self._expectation(_IMAGINARY_MOMENTS[key](imaginary))
return
observation_axes = self._expectation_axes
n_observations = self.n_observations
if "sign" in moments:
positive = xp.count_nonzero(imaginary > 0, axis=observation_axes)
negative = xp.count_nonzero(imaginary < 0, axis=observation_axes)
invalid = xp.any(xp.isnan(imaginary), axis=observation_axes)
moments["sign"][block] = xp.where(
invalid, xp.nan, (positive - negative) / n_observations
)
if "imaginary" in moments:
moments["imaginary"][block] = xp.mean(imaginary, axis=observation_axes)
if "squared" in moments:
# Sum of squares without materializing imaginary * imaginary.
squared_sum = xp.einsum(squared_subscripts, imaginary, imaginary)
moments["squared"][block] = squared_sum / n_observations
if "absolute" in moments:
xp.abs(imaginary, out=scratch)
moments["absolute"][block] = xp.mean(scratch, axis=observation_axes)
[docs]
@_asnumpy
def phase_lag_index(self) -> NDArray[np.floating]:
"""Return non-parametric synchrony measure mitigating power differences.
A non-parametric synchrony measure designed to mitigate power
differences between realizations (tapers, trials) and
volume-conduction.
The phase lag index is the average sign of the imaginary
component of the cross-spectrum. The imaginary component sets
in-phase or anti-phase signals to zero and the sign scales it to
have the same magnitude regardless of phase.
Note that this is the signed version of the phase lag index. In order
to obtain the unsigned version, as in [1], take the absolute value
of this quantity.
Returns
-------
phase_lag_index : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Signed phase lag index values for all signal pairs. Positive
``[..., i, j]`` means signal ``i`` leads signal ``j``.
Notes
-----
**Range**: [-1, 1] (signed version). For unsigned version (as in [1]),
take absolute value to get range [0, 1]. In-phase or anti-phase pairs,
whose imaginary cross-spectrum is zero up to rounding, are 0.
References
----------
.. [1] Stam, C.J., Nolte, G., and Daffertshofer, A. (2007). Phase
lag index: Assessment of functional connectivity from multi
channel EEG and MEG with diminished bias from common
sources. Human Brain Mapping 28, 1178-1193.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> pli = connectivity.phase_lag_index()
>>> pli.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # Signal 0 leads, so [..., 0, 1] is positive and [..., 1, 0] negative (10 Hz).
>>> bool(pli[0, 20, 0, 1] > 0 > pli[0, 20, 1, 0])
True
"""
self._warn_single_observation_degenerate(
"phase_lag_index",
"every value is the sign of one imaginary cross-spectrum, forced to "
"+/-1 (a perfectly consistent lag)",
)
# E[sign(Im)] of the cross-spectrum (real-valued). Pairs with no phase
# lag (in-phase signals) have Im at rounding level, whose sign is noise,
# so they are 0. xp.where returns a fresh array, not the cached moment.
mean_sign, mean_absolute = self._imaginary_cross_spectrum_moments("sign", "absolute")
pli: NDArray[np.floating] = xp.where(
self._has_no_phase_lag(mean_absolute), 0.0, mean_sign.real
)
return pli
[docs]
@_asnumpy
def directed_phase_lag_index(self) -> NDArray[np.floating]:
"""Return the directed phase-lag index (dPLI).
Values above 0.5 indicate that one signal consistently phase-leads the
other; values below 0.5 indicate that it phase-lags. A value of 0.5
represents no preferred phase-lag direction, including in-phase or
anti-phase pairs, whose imaginary cross-spectrum is zero up to rounding.
Returns
-------
directed_phase_lag_index : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
``[..., i, j]`` above 0.5 means signal ``i`` leads signal ``j``;
below 0.5 means signal ``i`` lags signal ``j``.
Notes
-----
**Range**: ``[0, 1]``. With the convention ``H(0) = 0.5``, dPLI is
``(1 + signed_PLI) / 2``. Consequently ``dPLI[i, j] = 1 - dPLI[j, i]``
and the diagonal is 0.5.
References
----------
.. [1] Stam, C.J., and van Straaten, E.C.W. (2012). Go with the flow:
use of a directed phase lag index (dPLI) to characterize
patterns of phase relations in a large-scale model of brain
dynamics. NeuroImage 62, 1415-1428.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> dpli = connectivity.directed_phase_lag_index()
>>> dpli.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # Signal 0 leads signal 1, so [..., 0, 1] > 0.5 > [..., 1, 0] (10 Hz = bin 20).
>>> bool(dpli[0, 20, 0, 1] > 0.5 > dpli[0, 20, 1, 0])
True
"""
self._warn_single_observation_degenerate(
"directed_phase_lag_index",
"every value is set by the sign of one imaginary cross-spectrum, "
"forced to 0 or 1 (a perfectly consistent lag)",
)
# Pairs with no phase lag (in-phase signals) have Im at rounding level,
# whose sign is noise, so they are neutral (0.5).
mean_sign, mean_absolute = self._imaginary_cross_spectrum_moments("sign", "absolute")
directed_pli: NDArray[np.floating] = xp.where(
self._has_no_phase_lag(mean_absolute),
0.5,
xp.clip((1.0 + mean_sign.real) / 2.0, 0.0, 1.0),
)
return directed_pli
[docs]
@_asnumpy
def weighted_phase_lag_index(self) -> NDArray[np.floating]:
"""Return weighted average of phase lag index using imaginary coherency magnitudes.
Weighted average of the phase lag index using the imaginary
coherency magnitudes as weights.
Note that this is the signed version of the weighted phase lag index
(mirroring :meth:`phase_lag_index`). In order to obtain the unsigned
version, as in [1], take the absolute value of this quantity.
Returns
-------
weighted_phase_lag_index : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Signed weighted phase lag index values. Positive ``[..., i, j]``
means signal ``i`` leads signal ``j``.
Notes
-----
**Range**: [-1, 1] (signed version). For the unsigned version (as in
[1]), take the absolute value to get range [0, 1]. The result is
antisymmetric in the signal pair (``wpli[..., i, j] = -wpli[..., j, i]``).
References
----------
.. [1] Vinck, M., Oostenveld, R., van Wingerden, M., Battaglia, F.,
and Pennartz, C.M.A. (2011). An improved index of
phase-synchronization for electrophysiological data in the
presence of volume-conduction, noise and sample-size bias.
NeuroImage 55, 1548-1565.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> wpli = connectivity.weighted_phase_lag_index()
>>> wpli.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # Signal 0 leads, so [..., 0, 1] is positive and [..., 1, 0] negative (10 Hz).
>>> bool(wpli[0, 20, 0, 1] > 0 > wpli[0, 20, 1, 0])
True
"""
self._warn_single_observation_degenerate(
"weighted_phase_lag_index",
"every value is one imaginary cross-spectrum divided by its own "
"magnitude, forced to +/-1 (a perfectly consistent lag)",
)
mean_imaginary, mean_absolute = self._imaginary_cross_spectrum_moments(
"imaginary", "absolute"
)
# Pairs with no phase lag (in-phase signals, the zeroed diagonal) have
# E[Im] and E[|Im|] both at rounding level, so their ratio is noise; the
# 0/0 is defined as 0, matching phase_lag_index's sign(0) == 0.
no_lag = self._has_no_phase_lag(mean_absolute)
return _divide_where(mean_imaginary, mean_absolute, ~no_lag, 0.0)
[docs]
@_asnumpy
def debiased_squared_phase_lag_index(self) -> NDArray[np.floating]:
"""Return square of phase lag index corrected for positive bias.
The square of the phase lag index corrected for the positive
bias induced by using the magnitude of the complex cross-spectrum.
Returns
-------
phase_lag_index : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Debiased squared phase lag index values (symmetric in the signal
pair, so they carry no lead/lag direction).
Notes
-----
**Range**: [-1 / (n_observations - 1), 1]. The unbiased finite-sample
estimate can be negative when the observed phase consistency is below
its null bias; negative values do not represent negative coupling.
Pairs whose imaginary cross-spectrum is exactly zero for every
observation (in-phase signals, and the diagonal) have no phase lag to
estimate and are returned as 0 rather than the lower bound.
References
----------
.. [1] Vinck, M., Oostenveld, R., van Wingerden, M., Battaglia, F.,
and Pennartz, C.M.A. (2011). An improved index of
phase-synchronization for electrophysiological data in the
presence of volume-conduction, noise and sample-size bias.
NeuroImage 55, 1548-1565.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> debiased_pli = connectivity.debiased_squared_phase_lag_index()
>>> debiased_pli.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> bool(debiased_pli[0, 20, 0, 1] > 0.05) # lagged coupling at 10 Hz (bin 20)
True
"""
self._validate_debiasing_observations("debiased_squared_phase_lag_index")
n_observations = self.n_observations
mean_sign, mean_absolute = self._imaginary_cross_spectrum_moments("sign", "absolute")
# Vinck's closed form assumes sign(Im) is +/-1 for every observation.
# Where there is no phase lag (in-phase signals, the zeroed diagonal)
# the signs are 0 or rounding noise and the form would report the
# spurious lower bound -1 / (n - 1) or a noise value; there is no lag
# to estimate, so 0.
debiased = (n_observations * mean_sign.real**2 - 1.0) / (n_observations - 1.0)
return xp.where(self._has_no_phase_lag(mean_absolute), 0.0, debiased)
[docs]
@_asnumpy
def debiased_squared_weighted_phase_lag_index(self) -> NDArray[np.floating]:
"""Return square of weighted phase lag index corrected for bias.
The square of the weighted phase lag index corrected for the
positive bias induced by using the magnitude of the complex
cross-spectrum.
Returns
-------
weighted_phase_lag_index : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Debiased squared weighted phase lag index values (symmetric in the
signal pair, so they carry no lead/lag direction).
Notes
-----
**Range**: [-1, 1]. The debiased finite-sample estimate can be negative
when the signed cross-products are dominated by inconsistent phase lags;
negative values do not represent negative coupling. Pairs whose
imaginary cross-spectrum is zero up to rounding (in-phase signals, the
diagonal, and the purely real DC and Nyquist bins) have no phase lag to
estimate and are returned as 0.
References
----------
.. [1] Vinck, M., Oostenveld, R., van Wingerden, M., Battaglia, F.,
and Pennartz, C.M.A. (2011). An improved index of
phase-synchronization for electrophysiological data in the
presence of volume-conduction, noise and sample-size bias.
NeuroImage 55, 1548-1565.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> debiased_wpli = connectivity.debiased_squared_weighted_phase_lag_index()
>>> debiased_wpli.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> bool(debiased_wpli[0, 20, 0, 1] > 0.3) # lagged coupling at 10 Hz (bin 20)
True
"""
self._validate_debiasing_observations("debiased_squared_weighted_phase_lag_index")
n_observations = self.n_observations
mean_imaginary, mean_squared, mean_absolute = self._imaginary_cross_spectrum_moments(
"imaginary", "squared", "absolute"
)
# Each product is a fresh array, so the cached moments are not mutated.
imaginary_csm_sum = mean_imaginary * n_observations
squared_imaginary_csm_sum = mean_squared * n_observations
imaginary_csm_magnitude_sum = mean_absolute * n_observations
weights = imaginary_csm_magnitude_sum**2 - squared_imaginary_csm_sum
# Pairs with no phase lag (in-phase signals, the zeroed diagonal) make
# the ratio one of rounding errors or 0/0; there is no lag to estimate,
# so 0, matching debiased_squared_phase_lag_index.
debiased: NDArray[np.floating] = xp.where(
self._has_no_phase_lag(mean_absolute),
0.0,
_divide_where(
imaginary_csm_sum**2 - squared_imaginary_csm_sum, weights, weights != 0, xp.nan
),
)
return debiased
[docs]
@_asnumpy
def pairwise_phase_consistency(self) -> NDArray[np.floating]:
"""Return square of phase locking value corrected for bias.
The square of the phase locking value corrected for the
positive bias induced by using the magnitude of the complex
cross-spectrum.
Returns
-------
phase_locking_value : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Pairwise phase consistency values.
Notes
-----
**Range**: [-1 / (n_observations - 1), 1]. The unbiased finite-sample
estimate can be negative when phase consistency is below its null bias;
negative values do not represent negative coupling.
References
----------
.. [1] Vinck, M., van Wingerden, M., Womelsdorf, T., Fries, P., and
Pennartz, C.M.A. (2010). The pairwise phase consistency: A
bias-free measure of rhythmic neuronal synchronization.
NeuroImage 51, 112-122.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> ppc = connectivity.pairwise_phase_consistency()
>>> ppc.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> bool(ppc[0, 20, 0, 1] > 0.3) # consistent phase difference at 10 Hz (bin 20)
True
"""
self._validate_debiasing_observations("pairwise_phase_consistency")
n_observations = self.n_observations
plv_sum = self._phase_locking_value() * n_observations
ppc = (plv_sum * plv_sum.conjugate() - n_observations) / (
n_observations**2 - n_observations
)
return ppc.real
[docs]
@_orientation_changed
@_source_first
@_asnumpy
def pairwise_spectral_granger_prediction(self) -> NDArray[np.floating]:
"""Return amount of power at a node explained by other nodes.
The amount of power at a node in a frequency explained by (is
predictive of) the power at other nodes.
Also known as spectral granger causality.
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Returns
-------
pairwise_granger : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Spectral Granger prediction values. Output ``[..., i, j]`` is the
influence of signal ``i`` on signal ``j`` (``i -> j``).
Notes
-----
**Non-negativity**: spectral Granger is ``>= 0`` by definition.
Negative estimates within roundoff of zero -- above ``-100 * eps`` of
the result dtype on the log-ratio scale, about ``-2e-14`` for float64
-- are clipped to ``0``; materially negative bins below that threshold
(a degenerate factorization) are returned as ``NaN``; use
:meth:`minimum_phase_reconstruction_error` to diagnose them. Other
packages (FieldTrip, MVGC, mne-connectivity) return such values as-is.
**Range**: [0, ∞). Non-negative values with no finite upper bound.
References
----------
.. [1] Geweke, J. (1982). Measurement of Linear Dependence and
Feedback Between Multiple Time Series. Journal of the
American Statistical Association 77, 304.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> granger = connectivity.pairwise_spectral_granger_prediction()
>>> granger.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # [..., i, j] is i -> j, so 0 -> 1 is [..., 0, 1]. It dominates at 10 Hz (bin 20).
>>> bool(granger[0, 20, 0, 1] > 10 * granger[0, 20, 1, 0])
True
"""
return self._pairwise_spectral_granger(
"pairwise_spectral_granger_prediction", time_reversed=False
)
def _pairwise_spectral_granger(
self, measure: str, *, time_reversed: bool
) -> NDArray[np.floating]:
"""Pairwise spectral Granger, optionally of the time-reversed process.
Time reversal of a real stationary process transposes its
cross-spectral matrix.
"""
self._require_two_sided_spectrum(measure)
csm = self._expectation_cross_spectral_matrix()
if time_reversed:
csm = xp.swapaxes(csm, -1, -2)
result = _estimate_spectral_granger_prediction(
self._power,
csm,
combinations(range(self.n_signals), 2),
minimum_phase_tolerance=self._minimum_phase_tolerance,
minimum_phase_max_iterations=self._minimum_phase_max_iterations,
)
_warn_nan_granger_pairs(result, measure)
return result
[docs]
@_orientation_changed
@_source_first
@_asnumpy
def subset_pairwise_spectral_granger_prediction(
self, pairs: Sequence[Sequence[int]] | NDArray[np.integer]
) -> NDArray[np.floating]:
"""Return predictive power for a subset of signal pairs.
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Parameters
----------
pairs : array_like, shape (n_pairs, 2)
Pairs of signal indices. Each pair is estimated in both directions.
Returns
-------
pairwise_granger : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Spectral Granger prediction for the specified pairs; entries for
pairs not requested (and the diagonal) are NaN. Output ``[..., i, j]``
is the influence of signal ``i`` on signal ``j`` (``i -> j``).
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz); signal 2 is noise.
>>> signals = np.stack([leader[3:], leader[:-3], np.zeros((1000, 20))], axis=-1)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> granger = connectivity.subset_pairwise_spectral_granger_prediction(pairs=[(0, 1)])
>>> granger.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 3, 3)
>>> # [..., i, j] is i -> j, so 0 -> 1 is [..., 0, 1]. It dominates at 10 Hz (bin 20).
>>> bool(granger[0, 20, 0, 1] > 10 * granger[0, 20, 1, 0])
True
>>> bool(np.isnan(granger[0, 20, 0, 2])) # pair (0, 2) was not requested
True
"""
self._require_two_sided_spectrum("subset_pairwise_spectral_granger_prediction")
pairs = np.array(pairs)
result = _estimate_subset_spectral_granger_prediction(
self._power,
self._subset_cross_spectral_matrix(pairs),
pairs,
n_signals=self._fourier_coefficients.shape[-1],
minimum_phase_tolerance=self._minimum_phase_tolerance,
minimum_phase_max_iterations=self._minimum_phase_max_iterations,
)
requested = np.zeros((self.n_signals, self.n_signals), dtype=bool)
requested[pairs[:, 0], pairs[:, 1]] = requested[pairs[:, 1], pairs[:, 0]] = True
_warn_nan_granger_pairs(
result, "subset_pairwise_spectral_granger_prediction", requested=requested
)
return result
[docs]
@_source_first
@_asnumpy
def time_reversed_spectral_granger_prediction(self) -> NDArray[np.floating]:
"""Return pairwise spectral Granger prediction after time reversal.
For a real stationary process, time reversal transposes the
cross-spectral matrix at every frequency. Contrasting this result with
:meth:`pairwise_spectral_granger_prediction` helps identify apparent
directionality caused by instantaneous mixing or data asymmetries.
Returns
-------
time_reversed_granger : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Output ``[..., i, j]`` is the influence of signal ``i`` on signal
``j`` (``i -> j``) in the time-reversed data.
Notes
-----
**Non-negativity**: spectral Granger is ``>= 0`` by definition.
Negative estimates within roundoff of zero -- above ``-100 * eps`` of
the result dtype on the log-ratio scale, about ``-2e-14`` for float64
-- are clipped to ``0``; materially negative bins below that threshold
(a degenerate factorization) are returned as ``NaN``; use
:meth:`minimum_phase_reconstruction_error` to diagnose them. Other
packages (FieldTrip, MVGC, mne-connectivity) return such values as-is.
References
----------
.. [1] Winkler, I., Panknin, D., Bartz, D., Müller, K.-R., and Haufe,
S. (2016). Validity of time reversal for testing Granger
causality. IEEE Transactions on Signal Processing 64, 2746-2760.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> reversed_granger = connectivity.time_reversed_spectral_granger_prediction()
>>> reversed_granger.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # A genuine 0 -> 1 lead flips under time reversal: 1 -> 0 ([..., 1, 0]) dominates.
>>> bool(reversed_granger[0, 20, 1, 0] > 10 * reversed_granger[0, 20, 0, 1])
True
"""
return self._pairwise_spectral_granger(
"time_reversed_spectral_granger_prediction", time_reversed=True
)
[docs]
@_source_first
@_asnumpy
def conditional_spectral_granger_prediction(self) -> NDArray[np.floating]:
"""Return pairwise spectral Granger prediction conditioned on all others.
For each ordered source-target pair, the influence of the source on the
target is measured after accounting for every remaining signal, using
the frequency-domain conditional measure of Chen, Bressler and Ding
(2006): the full model containing every signal and the reduced model
omitting the source are each spectrally factorized, and the reduced
model's innovation spectrum for the target is split into the part
explained by the target's own full-model innovations and the remainder
attributable to the source. With two signals this reduces to ordinary
pairwise spectral Granger.
Returns
-------
conditional_granger : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Output ``[..., i, j]`` is the influence of signal ``i`` on signal
``j`` (``i -> j``), conditional on every signal other than ``i``
and ``j``. The diagonal is NaN.
Notes
-----
**Non-negativity**: spectral Granger is ``>= 0`` by definition.
Negative estimates within roundoff of zero -- above ``-100 * eps`` of
the result dtype on the log-ratio scale, about ``-2e-14`` for float64
-- are clipped to ``0``; materially negative bins below that threshold
(a degenerate factorization) are returned as ``NaN``; use
:meth:`minimum_phase_reconstruction_error` to diagnose them. Other
packages (FieldTrip, MVGC, mne-connectivity) return such values as-is.
**Range**: ``[0, ∞)``. The measure is a log-ratio of a total to an
intrinsic innovation spectrum, so it is non-negative up to roundoff;
bins where either spectrum is not positive (a degenerate factorization)
are returned as NaN with a warning.
**Cost**: ``n_signals + 1`` minimum-phase factorizations (the full
system once, plus one ``(n_signals - 1)``-channel system per source),
each shared by every target. The full-system factorization is cached
and shared with the other directed measures.
References
----------
.. [1] Chen, Y., Bressler, S.L., and Ding, M. (2006). Frequency
decomposition of conditional Granger causality and application
to multivariate neural field potential data. Journal of
Neuroscience Methods 150, 228-237.
.. [2] Geweke, J.F. (1984). Measures of conditional linear dependence
and feedback between time series. Journal of the American
Statistical Association 79, 907-915.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> # Chain 0 -> 1 -> 2: each signal is the previous one delayed 3 samples, plus noise.
>>> source = rng.standard_normal((1006, 20))
>>> middle = source[:-3] + 0.5 * rng.standard_normal((1003, 20))
>>> target = middle[:-3] + 0.5 * rng.standard_normal((1000, 20))
>>> signals = np.stack([source[6:], middle[3:], target], axis=-1) # (time, trials, 3)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> conditional = connectivity.conditional_spectral_granger_prediction()
>>> conditional.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 3, 3)
>>> # [..., i, j] is i -> j. Pairwise Granger sees the indirect 0 -> 2 ([..., 0, 2]),
>>> # but conditioning on signal 1 removes it while keeping 1 -> 2 (10 Hz = bin 20).
>>> pairwise = connectivity.pairwise_spectral_granger_prediction()
>>> bool(pairwise[0, 20, 0, 2] > 0.5), bool(conditional[0, 20, 0, 2] < 0.05)
(True, True)
>>> bool(conditional[0, 20, 1, 2] > 0.5)
True
"""
self._require_two_sided_spectrum("conditional_spectral_granger_prediction")
spectrum = self._expectation_cross_spectral_matrix()
tolerance = self._minimum_phase_tolerance
max_iterations = self._minimum_phase_max_iterations
if self.n_signals == 2:
# No conditioning set: the measure is pairwise Geweke Granger.
result = _estimate_blockwise_spectral_granger(
spectrum,
[np.array([0]), np.array([1])],
minimum_phase_tolerance=tolerance,
minimum_phase_max_iterations=max_iterations,
)
else:
# The full model is the instance's cached factorization (the same
# CSM, tolerance and iteration cap), shared with the other directed
# measures.
result = _estimate_all_conditional_spectral_granger(
spectrum,
self._transfer_function,
self._noise_covariance,
minimum_phase_tolerance=tolerance,
minimum_phase_max_iterations=max_iterations,
)
_warn_nan_granger_pairs(result, "conditional_spectral_granger_prediction")
return result
[docs]
def blockwise_spectral_granger_prediction(
self, group_labels: NDArray[np.integer]
) -> tuple[NDArray[np.floating], NDArray[np.integer]]:
"""Return spectral Granger prediction between multichannel groups.
Parameters
----------
group_labels : array-like, shape (n_signals,)
Label assigning each signal to one non-overlapping group.
Returns
-------
blockwise_granger : array
Shape ``(..., n_nonnegative_frequencies, n_groups, n_groups)``.
Output ``[..., i, j]`` is the influence of group ``i`` on group ``j``
(``i -> j``), with groups ordered as in ``labels``. The diagonal is
NaN.
labels : array, shape (n_groups,)
Sorted unique group labels.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Group "a": two noisy copies of signal 0; "b": two of its 6 ms-delayed copy.
>>> pair = np.stack([leader[3:], leader[:-3]], axis=-1)
>>> signals = np.repeat(pair, 2, axis=-1) # (time, trials, 4 signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> granger, labels = connectivity.blockwise_spectral_granger_prediction(
... ["a", "a", "b", "b"]
... )
>>> granger.shape # (n_time_windows, n_frequencies, n_groups, n_groups)
(1, 501, 2, 2)
>>> labels
array(['a', 'b'], dtype='<U1')
>>> # [..., i, j] is group i -> group j, so a -> b is [..., 0, 1] (10 Hz = bin 20).
>>> bool(granger[0, 20, 0, 1] > 10 * granger[0, 20, 1, 0])
True
"""
self._require_two_sided_spectrum("blockwise_spectral_granger_prediction")
labels, indices, _ = self._validated_group_indices(group_labels)
result = _estimate_blockwise_spectral_granger(
self._expectation_cross_spectral_matrix(),
indices,
minimum_phase_tolerance=self._minimum_phase_tolerance,
minimum_phase_max_iterations=self._minimum_phase_max_iterations,
)
_warn_nan_granger_pairs(
result,
"blockwise_spectral_granger_prediction",
names=to_numpy(labels),
)
# Swap on the host, like _source_first: [..., target, source] -> source first.
return np.swapaxes(to_numpy(result), -1, -2), to_numpy(labels)
[docs]
@_orientation_changed
@_ignore_nan_propagation_warnings
@_source_first
@_asnumpy
def directed_transfer_function(self) -> NDArray[np.floating]:
"""Return transfer function coupling strength normalized by inflow.
The transfer function coupling strength normalized by the total
influence of other signals on that signal (inflow).
Characterizes the direct and indirect coupling to a node.
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Returns
-------
directed_transfer_function : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Directed transfer function values. Output ``[..., i, j]`` is the
influence of signal ``i`` on signal ``j`` (``i -> j``).
Notes
-----
**Range**: [0, 1] (normalized). Represents proportion of inflow
via transfer function; each target's values sum to 1 over sources
(``result.sum(axis=-2)`` is 1).
References
----------
.. [1] Kaminski, M., and Blinowska, K.J. (1991). A new method of
the description of the information flow in the brain
structures. Biological Cybernetics 65, 203-210.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> dtf = connectivity.directed_transfer_function()
>>> dtf.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # [..., i, j] is i -> j, so 0 -> 1 is [..., 0, 1]. It dominates at 10 Hz (bin 20).
>>> bool(dtf[0, 20, 0, 1] > 10 * dtf[0, 20, 1, 0])
True
>>> bool(np.allclose(dtf.sum(axis=-2), 1)) # each target's inflow sums to 1
True
"""
return _squared_magnitude(
self._transfer_function / _total_inflow(self._transfer_function)
)
[docs]
@_orientation_changed
@_ignore_nan_propagation_warnings
@_source_first
@_asnumpy
def directed_coherence(self) -> NDArray[np.floating]:
"""Return the squared directed coherence (noise-weighted DTF).
Like the directed transfer function, but the noise variance weights
**both** the numerator and the inflow normalization. The returned value
is the squared directed coherence ``nv_j |H_ij|^2 / sum_k nv_k |H_ik|^2``,
written in the transfer function's native ``[target, source]`` indexing
(``H_ij`` is ``j -> i``), where ``nv`` is the per-signal innovation
(noise) variance and ``H`` is the transfer function. Each target's values
sum to 1 over sources (``result.sum(axis=-2)`` is 1).
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Returns
-------
directed_coherence : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Squared directed coherence values. Output ``[..., i, j]`` is the
influence of signal ``i`` on signal ``j`` (``i -> j``).
Notes
-----
**Range**: [0, 1]. Normalized directional connectivity measure.
**Assumption**: This measure follows Baccala et al. (1998), which
assumes the MVAR innovations are uncorrelated (a diagonal noise
covariance). The denominator then equals the signal's power spectral
density ``S_ii = sum_k nv_k |H_ik|^2``. When the estimated innovation
covariance has non-negligible off-diagonal terms (common for
non-parametrically estimated MVARs), the true PSD
``S_ii = (H Cov H^H)_ii`` also contains cross-power between correlated
sources that this diagonal formula omits, so the values are approximate.
A ``UserWarning`` is emitted in that case.
References
----------
.. [1] Baccala, L., Sameshima, K., Ballester, G., Do Valle, A., and
Timo-Iaria, C. (1998). Studying the interaction between
brain structures via directed coherence and Granger
causality. Applied Signal Processing 5, 40.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> directed_coherence = connectivity.directed_coherence()
>>> directed_coherence.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # [..., i, j] is i -> j, so 0 -> 1 is [..., 0, 1]. It dominates at 10 Hz (bin 20).
>>> bool(directed_coherence[0, 20, 0, 1] > 10 * directed_coherence[0, 20, 1, 0])
True
>>> bool(np.allclose(directed_coherence.sum(axis=-2), 1)) # each target's inflow
True
"""
# Directed coherence normalizes the noise-weighted inflow over sources
# (native axis -1), so the per-source noise variance must vary along that axis.
# The squared measure is nv_j |H_ij|^2 / sum_k nv_k |H_ik|^2, which sums
# to 1 over sources like the directed transfer function. This uses only
# the diagonal of the noise covariance, which equals the PSD denominator
# only when the innovations are uncorrelated; warn when the omitted
# cross-power is a material fraction of the true PSD.
if (
_max_psd_discrepancy(self._transfer_function, self._noise_covariance)
> DIRECTED_COHERENCE_DISCREPANCY_TOLERANCE
):
warnings.warn(
"directed_coherence assumes uncorrelated MVAR innovations (a "
"diagonal noise covariance), but the estimated innovation "
"covariance has non-negligible off-diagonal terms: the diagonal "
"normalization omits a material fraction of the true power "
"spectral density (cross-power between correlated sources). "
"Interpret the values as approximate, or use "
"partial_directed_coherence instead.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
noise_variance = _get_noise_variance(self._noise_covariance, axis=-1)
directed_coherence: NDArray[np.floating] = (
noise_variance
* _squared_magnitude(self._transfer_function)
/ _total_inflow(self._transfer_function, noise_variance) ** 2
)
return directed_coherence
def _partial_directed_coherence(self) -> NDArray[np.floating]:
"""Return device-native PDC for reuse by other device-native measures."""
return _squared_magnitude(
self._MVAR_Fourier_coefficients / _total_outflow(self._MVAR_Fourier_coefficients)
)
[docs]
@_orientation_changed
@_ignore_nan_propagation_warnings
@_source_first
@_asnumpy
def partial_directed_coherence(self) -> NDArray[np.floating]:
"""Return transfer function coupling strength normalized by outflow.
The transfer function coupling strength normalized by its
strength of coupling to other signals (outflow).
The partial directed coherence tries to regress out the influence
of other observed signals, leaving only the direct coupling between
two signals.
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Returns
-------
partial_directed_coherence : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Partial directed coherence values. Output ``[..., i, j]`` is the
influence of signal ``i`` on signal ``j`` (``i -> j``).
Notes
-----
**Range**: [0, 1]. Normalized direct coupling measure; each source's
values sum to 1 over targets (``result.sum(axis=-1)`` is 1).
References
----------
.. [1] Baccala, L.A., and Sameshima, K. (2001). Partial directed
coherence: a new concept in neural structure determination.
Biological Cybernetics 84, 463-474.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> pdc = connectivity.partial_directed_coherence()
>>> pdc.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # [..., i, j] is i -> j, so 0 -> 1 is [..., 0, 1]. It dominates at 10 Hz (bin 20).
>>> bool(pdc[0, 20, 0, 1] > 10 * pdc[0, 20, 1, 0])
True
>>> bool(np.allclose(pdc.sum(axis=-1), 1)) # each source's outflow sums to 1
True
"""
return self._partial_directed_coherence()
[docs]
@_orientation_changed
@_ignore_nan_propagation_warnings
@_source_first
@_asnumpy
def generalized_partial_directed_coherence(self) -> NDArray[np.floating]:
"""Return generalized partial directed coherence.
The transfer function coupling strength normalized by its
strength of coupling to other signals (outflow).
The partial directed coherence tries to regress out the influence
of other observed signals, leaving only the direct coupling between
two signals.
The generalized partial directed coherence scales the relative
strength of coupling by the noise variance.
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Returns
-------
generalized_partial_directed_coherence : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Generalized partial directed coherence values. Output ``[..., i, j]``
is the influence of signal ``i`` on signal ``j`` (``i -> j``).
Notes
-----
**Range**: [0, 1]. Normalized, scaled by noise variance; each source's
values sum to 1 over targets (``result.sum(axis=-1)`` is 1).
References
----------
.. [1] Baccala, L.A., Sameshima, K., and Takahashi, D.Y. (2007).
Generalized partial directed coherence. In Digital Signal
Processing, 2007 15th International Conference on, (IEEE),
pp. 163-166.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> gpdc = connectivity.generalized_partial_directed_coherence()
>>> gpdc.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 2, 2)
>>> # [..., i, j] is i -> j, so 0 -> 1 is [..., 0, 1]. It dominates at 10 Hz (bin 20).
>>> bool(gpdc[0, 20, 0, 1] > 10 * gpdc[0, 20, 1, 0])
True
>>> bool(np.allclose(gpdc.sum(axis=-1), 1)) # each source's outflow sums to 1
True
"""
noise_variance = _get_noise_variance(self._noise_covariance)
return _squared_magnitude(
self._MVAR_Fourier_coefficients
/ xp.sqrt(noise_variance)
/ _total_outflow(self._MVAR_Fourier_coefficients, noise_variance)
)
[docs]
@_orientation_changed
@_ignore_nan_propagation_warnings
@_source_first
@_asnumpy
def direct_directed_transfer_function(self) -> NDArray[np.floating]:
"""Return the direct directed transfer function (dDTF).
The squared dDTF of Korzeniewska et al. (2003) multiplies the
full-frequency DTF, which keeps direct and indirect (cascade) influence,
by the squared partial coherence, which is zero between two signals
whose relation is fully explained by the others. The product therefore
keeps only direct influence:
``chi^2_ij(f) = ffDTF^2_ij(f) * kappa^2_ij(f)``, with
``ffDTF^2_ij(f) = |H_ij(f)|^2 / sum_f' sum_k |H_ik(f')|^2`` (summed over
the non-negative frequencies) and
``kappa^2_ij(f) = |G_ij(f)|^2 / (G_ii(f) G_jj(f))``, where
``G = A^H Sigma^-1 A`` is the inverse spectral matrix of the MVAR model
(``A = H^-1``, ``Sigma`` the innovation covariance). These formulas are
written in the transfer function's native ``[target, source]`` indexing
(``chi^2_ij`` is ``j -> i``).
.. versionchanged:: 3.0
The output is source first (``[..., i, j]`` is ``i -> j``); 2.x
returned the transpose (``j -> i``).
Returns
-------
direct_directed_transfer_function : array
Shape ``(..., n_nonnegative_frequencies, n_signals, n_signals)``.
Output ``[..., i, j]`` is the direct influence of signal ``i`` on
signal ``j`` (``i -> j``).
Notes
-----
**Range**: [0, 1]. Like :meth:`directed_transfer_function`, this
returns the squared quantity; SCoT and ConnectiviPy report its square
root, ``|ffDTF| * |kappa|``. The diagonal is the full-frequency DTF of
each signal with itself (``kappa_ii = 1``).
References
----------
.. [1] Korzeniewska, A., Manczak, M., Kaminski,
M., Blinowska, K.J., and Kasicki, S. (2003). Determination
of information flow direction among brain structures by a
modified directed transfer function (dDTF) method.
Journal of Neuroscience Methods 125, 195-207.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> # Chain 0 -> 1 -> 2: each signal is the previous one delayed 3 samples, plus noise.
>>> source = rng.standard_normal((1006, 20))
>>> middle = source[:-3] + 0.5 * rng.standard_normal((1003, 20))
>>> target = middle[:-3] + 0.5 * rng.standard_normal((1000, 20))
>>> signals = np.stack([source[6:], middle[3:], target], axis=-1) # (time, trials, 3)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> ddtf = connectivity.direct_directed_transfer_function()
>>> ddtf.shape # (n_time_windows, n_frequencies, n_signals, n_signals)
(1, 501, 3, 3)
>>> # [..., i, j] is i -> j. The direct 1 -> 2 ([..., 1, 2]) far exceeds both the
>>> # reverse 2 -> 1 ([..., 2, 1]) and the indirect 0 -> 2 ([..., 0, 2]) that is
>>> # relayed through signal 1 (10 Hz = bin 20).
>>> direct = ddtf[0, 20, 1, 2]
>>> bool(direct > 10 * ddtf[0, 20, 2, 1]), bool(direct > 10 * ddtf[0, 20, 0, 2])
(True, True)
"""
full_frequency_dtf = _squared_magnitude(
self._transfer_function / _total_inflow(self._transfer_function, axis=(-1, -3))
)
mvar_coefficients = self._MVAR_Fourier_coefficients
inverse_noise_covariance = _regularized_inverse(self._noise_covariance)
inverse_spectrum = xp.matmul(
xp.matmul(
_conjugate_transpose(mvar_coefficients),
inverse_noise_covariance[..., xp.newaxis, :, :],
),
mvar_coefficients,
)
inverse_spectrum_diagonal = xp.real(xp.diagonal(inverse_spectrum, axis1=-2, axis2=-1))
squared_partial_coherence: NDArray[np.floating] = _squared_magnitude(
inverse_spectrum
) / (
inverse_spectrum_diagonal[..., :, xp.newaxis]
* inverse_spectrum_diagonal[..., xp.newaxis, :]
)
return full_frequency_dtf * squared_partial_coherence
def _significant_pair_phase(
self,
measure: str,
frequencies_of_interest: NDArray[np.floating] | None,
frequency_resolution: float | None,
significance_threshold: float,
) -> tuple[np.ma.MaskedArray, NDArray[np.floating], NDArray[np.intp], int]:
"""Unwrapped coherency phase per signal pair, masked where not significant.
Shared setup of :meth:`group_delay` and :meth:`delay`: validates the
frequency grid and observation weights, bandpasses the coherency,
gathers the upper-triangle signal pairs, and masks frequencies whose
coherence is not significant.
Returns
-------
coherence_phase : masked array, shape (..., n_band_frequencies, n_pairs)
Phase unwrapped along frequency; masked where not significant.
bandpassed_frequencies : array, shape (n_band_frequencies,)
signal_combination_ind : array of int, shape (n_pairs, 2)
``(i, j)`` with ``i < j`` for each pair column.
n_signals : int
"""
frequencies = self.frequencies
self._require_multiple_frequencies(measure)
self._require_uniform_frequency_grid(measure)
self._require_uniform_observation_weights(
measure,
"Its coherence significance test uses the observation count as the "
"degrees of freedom of the zero-coherence null, which assumes equally "
"weighted observations.",
)
self._warn_correlated_observations(
measure,
"Its coherence significance test uses the observation count as the "
"degrees of freedom of the zero-coherence null",
)
frequency_difference = frequencies[1] - frequencies[0]
independent_frequency_step = _get_independent_frequency_step(
frequency_difference, frequency_resolution
)
bandpassed_coherency, bandpassed_frequencies = _bandpass(
self._coherency(), frequencies, frequencies_of_interest
)
# Significance testing and the masked phase below are NumPy operations.
# Make the GPU-to-host boundary explicit before passing data into them;
# NumPy deliberately refuses implicit conversion of CuPy arrays.
bandpassed_coherency = to_numpy(bandpassed_coherency)
bandpassed_frequencies = to_numpy(bandpassed_frequencies)
n_signals = bandpassed_coherency.shape[-1]
signal_combination_ind = np.asarray(list(combinations(np.arange(n_signals), 2)))
bandpassed_coherency = bandpassed_coherency[
..., signal_combination_ind[:, 0], signal_combination_ind[:, 1]
]
is_significant = _find_significant_frequencies(
bandpassed_coherency,
self.n_observations,
independent_frequency_step,
significance_threshold=significance_threshold,
)
coherence_phase = np.ma.masked_array(
_unwrap_across_undefined(np.angle(bandpassed_coherency), axis=-2),
mask=~is_significant,
)
return coherence_phase, bandpassed_frequencies, signal_combination_ind, n_signals
[docs]
def group_delay(
self,
frequencies_of_interest: NDArray[np.floating] | None = None,
frequency_resolution: float | None = None,
significance_threshold: float = 0.05,
) -> tuple[NDArray[np.floating], NDArray[np.floating], NDArray[np.floating]]:
"""Return the average time-delay of a broadband signal.
Parameters
----------
frequencies_of_interest : array-like, shape (2,), optional
Frequency band ``(low, high)`` to fit over, in the units of
``frequencies``. Both edges are exclusive: only bins strictly
inside the band are used, and at least one must be. ``None`` uses
every frequency.
frequency_resolution : float, optional
Frequency resolution for independent samples.
significance_threshold : float, default=0.05
P-value threshold for significance.
Returns
-------
delay : array, shape (..., n_signals, n_signals)
Time delays between signal pairs, in the reciprocal units of
``frequencies``: seconds for Hz, samples for cycles/sample.
Positive ``[..., i, j]`` means signal ``i`` leads signal ``j``. The
diagonal is NaN.
slope : array, shape (..., n_signals, n_signals)
Slope of the coherence phase vs frequency, in radians per unit of
``frequencies`` (``delay = slope / (2 * pi)``); same sign
convention as ``delay``.
r_value : array, shape (..., n_signals, n_signals)
Correlation coefficient of the linear phase-frequency fit, with the
sign of ``slope`` (so ``[..., j, i]`` is ``-[..., i, j]``).
Notes
-----
**Range**: (-∞, ∞). Time delays can be positive or negative.
References
----------
.. [1] Gotman, J. (1983). Measurement of small time differences
between EEG channels: method and application to epileptic
seizure propagation. Electroencephalography and Clinical
Neurophysiology 56, 501-514.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> delay, slope, r_value = connectivity.group_delay(frequencies_of_interest=[5, 50])
>>> delay.shape # (n_time_windows, n_signals, n_signals)
(1, 2, 2)
>>> # Signal 0 leads signal 1 by 6 ms, so [..., 0, 1] is +0.006 s.
>>> round(float(delay[0, 0, 1]), 3), round(float(delay[0, 1, 0]), 3)
(0.006, -0.006)
>>> bool(r_value[0, 0, 1] > 0.9) # phase is linear in frequency
True
"""
coherence_phase, bandpassed_frequencies, signal_combination_ind, n_signals = (
self._significant_pair_phase(
"group_delay",
frequencies_of_interest,
frequency_resolution,
significance_threshold,
)
)
# Vectorized masked linear regression of the unwrapped phase on
# frequency, per (batch, signal pair), replacing a per-slice scipy
# ``linregress`` call (``apply_along_axis`` invokes it once per slice).
# Uses the closed-form ordinary-least-squares slope and Pearson r from
# centered masked sums over the frequency axis (-2); where at least two
# distinct significant frequencies remain these match ``linregress`` to
# floating-point tolerance, and degenerate slices (fewer than two, or a
# single distinct frequency) yield NaN as ``linregress`` does.
is_valid = ~np.ma.getmaskarray(coherence_phase)
# Replace masked entries with a finite value before the arithmetic: the
# masked phase can be NaN (e.g. a zero-power bin's coherency angle), and
# ``0 * NaN`` is NaN, which would poison the whole slice's sums even
# though the valid mask already excludes those entries.
phase = np.where(is_valid, np.ma.getdata(coherence_phase), 0.0)
frequency = np.asarray(bandpassed_frequencies, dtype=float).reshape(-1, 1)
axis = -2
count = is_valid.sum(axis, keepdims=True)
safe_count = np.where(count == 0, 1, count)
# Center x and y on their per-slice means before summing squares, so the
# variance does not catastrophically cancel when the absolute
# frequencies are large relative to their spacing (raw-moment
# ``count * sum_xx - sum_x**2`` loses all precision there).
mean_x = (is_valid * frequency).sum(axis, keepdims=True) / safe_count
mean_y = (is_valid * phase).sum(axis, keepdims=True) / safe_count
centered_x = frequency - mean_x
centered_y = phase - mean_y
sum_xx = (is_valid * centered_x * centered_x).sum(axis, keepdims=True)
sum_yy = (is_valid * centered_y * centered_y).sum(axis, keepdims=True)
sum_xy = (is_valid * centered_x * centered_y).sum(axis, keepdims=True)
with np.errstate(invalid="ignore", divide="ignore"):
pair_slope = (sum_xy / sum_xx)[..., 0, :]
pair_r_value = (sum_xy / np.sqrt(sum_xx * sum_yy))[..., 0, :]
# Guard against |r| drifting just past 1 from rounding.
pair_r_value = np.clip(pair_r_value, -1.0, 1.0)
new_shape = (*coherence_phase.shape[:-2], n_signals, n_signals)
slope = np.full(new_shape, np.nan)
slope[..., signal_combination_ind[:, 0], signal_combination_ind[:, 1]] = pair_slope
slope[..., signal_combination_ind[:, 1], signal_combination_ind[:, 0]] = -pair_slope
delay = slope / (2 * np.pi)
r_value = np.ones(new_shape)
r_value[..., signal_combination_ind[:, 0], signal_combination_ind[:, 1]] = pair_r_value
# The reverse pair's phase is negated, so its correlation is too.
r_value[
..., signal_combination_ind[:, 1], signal_combination_ind[:, 0]
] = -pair_r_value
return delay, slope, r_value
[docs]
@_asnumpy
def delay(
self,
frequencies_of_interest: NDArray[np.floating] | None = None,
frequency_resolution: float | None = None,
significance_threshold: float = 0.05,
n_range: int = 3,
) -> NDArray[np.floating]:
"""Find a range of possible delays from the coherence phase.
The delay (and phase) at each frequency is indistinguishable from
2π phase jumps, but we can look at a range of possible delays
and see which one is most likely.
Parameters
----------
frequencies_of_interest : array-like, shape (2,), optional
Frequency band ``(low, high)`` to evaluate, in the units of
``frequencies``. Both edges are exclusive: only bins strictly
inside the band are returned, and at least one must be. ``None``
uses every frequency.
frequency_resolution : float, optional
Frequency resolution for independent samples.
significance_threshold : float, default=0.05
P-value threshold for significance.
n_range : int, default=3
Number of phases to consider.
Returns
-------
possible_delays : array
Shape (..., n_frequencies, (n_range * 2) + 1, n_signals, n_signals),
where ``n_frequencies`` counts only the frequencies inside
``frequencies_of_interest``. Candidate ``k`` (index ``k + n_range``)
adds ``k`` cycles of phase. Array of possible time delays in the
reciprocal units of ``frequencies`` (seconds for Hz, samples for
cycles/sample); positive ``[..., i, j]`` means signal ``i`` leads
signal ``j``. The true delay is the candidate that is consistent
(frequency-independent) across the band. Frequencies without
significant coherence, and the 0 Hz (DC) bin, are undefined and
returned as NaN.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> delays = connectivity.delay(frequencies_of_interest=[5, 50], n_range=3)
>>> delays.shape # (n_time_windows, n_band_freqs, n_candidates, n_signals, n_signals)
(1, 89, 7, 2, 2)
>>> # The zero-wrap candidate (index n_range) recovers the 6 ms lead of signal 0.
>>> round(float(np.nanmedian(delays[0, :, 3, 0, 1])), 3)
0.006
"""
coherence_phase, bandpassed_frequencies, signal_combination_ind, n_signals = (
self._significant_pair_phase(
"delay", frequencies_of_interest, frequency_resolution, significance_threshold
)
)
possible_range = 2 * np.pi * np.arange(-n_range, n_range + 1)
# Convert phase to a time delay: tau = (phase + 2*pi*k) / (2*pi*f). The
# 2*pi*k terms resolve the phase-wrapping ambiguity. Dividing only by
# 2*pi (omitting f) would return cycles, not seconds, making a constant
# physical delay appear frequency-dependent.
cycles = np.rollaxis(
(possible_range + coherence_phase[..., np.newaxis]) / (2 * np.pi), -1, -2
)
# cycles has shape (..., n_frequencies, n_candidates, n_pairs); divide by
# the frequency along the n_frequencies axis (-3). DC (f == 0) has no
# defined delay and becomes NaN.
frequency = bandpassed_frequencies[:, np.newaxis, np.newaxis]
with np.errstate(divide="ignore", invalid="ignore"):
delays = cycles / frequency
delays[..., bandpassed_frequencies == 0, :, :] = np.nan
# Fill non-significant frequencies (masked) with NaN rather than the
# masked array's underlying 0.0, so a non-significant bin is not read as
# a genuine zero-lag delay. This matches the DC handling above.
delays = np.ma.filled(delays, np.nan)
new_shape = (
*coherence_phase.shape[:-1],
len(possible_range),
n_signals,
n_signals,
)
possible_delays = np.full(new_shape, np.nan)
possible_delays[..., signal_combination_ind[:, 0], signal_combination_ind[:, 1]] = (
delays
)
# The reverse pair's phase is -phase, so its candidate k,
# (-phase + 2*pi*k) / (2*pi*f), is minus the forward candidate -k.
possible_delays[
..., signal_combination_ind[:, 1], signal_combination_ind[:, 0]
] = -delays[..., ::-1, :]
return possible_delays
[docs]
@_asnumpy
def phase_slope_index(
self,
frequencies_of_interest: NDArray[np.floating] | None = None,
frequency_resolution: float | None = None,
) -> NDArray[np.floating]:
"""Return weighted average of slopes projected onto imaginary axis.
The phase slope index sums the product of the coherency at adjacent
frequencies, ``conj(C(f)) * C(f + df)``, over the band and projects the
result onto the imaginary axis to avoid volume-conduction effects
(Nolte et al. 2008). The magnitude of the coherency at each frequency
therefore weights the contribution of that frequency step.
Parameters
----------
frequencies_of_interest : array-like, shape (2,), optional
Frequency band ``(low, high)`` to sum over, in the units of
``frequencies``. Both edges are exclusive: only bins strictly
inside the band contribute, and at least one must. ``None`` uses
every frequency.
frequency_resolution : float, optional
Frequency resolution for independent samples.
Returns
-------
phase_slope_index : array, shape (..., n_signals, n_signals)
Phase slope index values. Positive ``[..., i, j]`` means signal
``i`` leads signal ``j``; the result is antisymmetric.
Notes
-----
**Range**: (-∞, ∞). Signed directional measure with no bounds.
References
----------
.. [1] Nolte, G., Ziehe, A., Nikulin, V.V., Schlogl, A., Kramer,
N., Brismar, T., and Muller, K.-R. (2008). Robustly
Estimating the Flow Direction of Information in Complex
Physical Systems. Physical Review Letters 100.
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity import Connectivity, Multitaper
>>> rng = np.random.default_rng(0)
>>> leader = rng.standard_normal((1003, 20))
>>> # Signal 1 is signal 0 delayed by 3 samples (6 ms at 500 Hz), plus noise.
>>> signals = np.stack([leader[3:], leader[:-3]], axis=-1) # (time, trials, signals)
>>> signals += 0.5 * rng.standard_normal(signals.shape)
>>> multitaper = Multitaper(signals, sampling_frequency=500)
>>> connectivity = Connectivity.from_transform(multitaper)
>>> psi = connectivity.phase_slope_index(frequencies_of_interest=[5, 50])
>>> psi.shape # (n_time_windows, n_signals, n_signals)
(1, 2, 2)
>>> # Signal 0 leads signal 1, so [..., 0, 1] is positive and [..., 1, 0] negative.
>>> bool(psi[0, 0, 1] > 0 > psi[0, 1, 0])
True
"""
frequencies = self.frequencies
bandpassed_coherency, bandpassed_frequencies = _bandpass(
self._coherency(), frequencies, frequencies_of_interest
)
self._require_multiple_frequencies("phase_slope_index")
self._require_uniform_frequency_grid("phase_slope_index")
frequency_difference = frequencies[1] - frequencies[0]
independent_frequency_step = _get_independent_frequency_step(
frequency_difference, frequency_resolution
)
frequency_index = xp.arange(
0, bandpassed_frequencies.shape[0], independent_frequency_step
)
bandpassed_coherency = bandpassed_coherency[..., frequency_index, :, :]
# The phase slope index needs at least two frequency bins to form an
# adjacent-frequency product. With fewer, the sum below would be an empty
# sum that NumPy reports as 0, which is indistinguishable from a genuine
# "no directionality" result.
n_band_frequencies = bandpassed_coherency.shape[-3]
if n_band_frequencies < 2:
msg = (
f"phase_slope_index needs at least 2 frequency bins in the band "
f"after subsampling, but got {n_band_frequencies}. Widen "
f"frequencies_of_interest, or decrease frequency_resolution so "
f"more than one independent frequency remains."
)
raise ValueError(msg)
# Nolte et al. (2008): sum conj(C(f)) * C(f + df) over adjacent
# (independent) frequency bins, then take the imaginary part. The
# frequency axis is -3 (the two trailing axes are the signal pair).
adjacent_product: NDArray[np.complexfloating] = (
xp.conj(bandpassed_coherency[..., :-1, :, :]) * bandpassed_coherency[..., 1:, :, :]
).sum(axis=-3)
return adjacent_product.imag
def _divide_masking_zero_denominator(
numerator: NDArray[_NumberT],
denominator: NDArray[np.floating],
message: str,
) -> NDArray[_NumberT]:
"""Divide, returning NaN where the denominator is (near-)zero.
The coherence measures are undefined where a signal has (near-)zero power.
Rather than dividing by a floored epsilon (which yields spuriously large
values), this emits a single ``UserWarning`` (``message``) when any
denominator entry is at or below the dtype's smallest positive value,
divides against a safe (1.0-substituted) denominator, and sets those entries
to NaN.
Parameters
----------
numerator : NDArray
Numerator array (real or complex).
denominator : NDArray[floating]
Non-negative denominator, broadcastable against ``numerator``.
message : str
Warning text emitted once if any denominator entry is (near-)zero.
Returns
-------
NDArray
``numerator / denominator`` with (near-)zero-denominator entries NaN.
"""
zero = denominator <= xp.finfo(denominator.dtype).tiny
invalid = zero | ~xp.isfinite(denominator)
if xp.any(zero):
warnings.warn(message, UserWarning, stacklevel=stacklevel_outside_package())
# Dividing by a real array keeps the numerator's real/complex kind.
return cast(NDArray[_NumberT], _divide_where(numerator, denominator, ~invalid, xp.nan))
def _total_inflow(
transfer_function: NDArray[np.complexfloating],
noise_variance: float | NDArray[np.floating] = 1.0,
axis: int | tuple[int, ...] = -1,
) -> NDArray[np.floating]:
"""Measure effect of incoming signals onto a node via sum of squares.
Parameters
----------
transfer_function : array_like
Transfer function matrix.
noise_variance : float or array_like, default=1.0
Noise variance values.
axis : int, default=-1
Axis for summation.
Returns
-------
array_like
Total inflow values.
"""
return xp.sqrt(
xp.sum(
noise_variance * _squared_magnitude(transfer_function),
keepdims=True,
axis=axis,
)
)
def _get_noise_variance(
noise_covariance: NDArray[np.floating],
axis: int = -2,
) -> NDArray[np.floating]:
"""Extract noise variance and broadcast it along the requested signal axis.
The transfer function / MVAR coefficients use the convention
``[..., n_fft, target, source]`` (axis -2 is the target/row, axis -1 is the
source/column). Different directed measures weight by the noise variance
along different axes, so the diagonal of the covariance must be broadcast to
match: partial-directed-coherence-family measures sum over the target axis
(-2), while directed coherence sums over the source axis (-1).
Parameters
----------
noise_covariance : array, shape (..., n_signals, n_signals)
Noise covariance matrix.
axis : int, default=-2
Signal axis the noise variance should vary along, either -2 (target) or
-1 (source). A leading ``newaxis`` is always inserted for the frequency
axis.
Returns
-------
noise_variance : array
Diagonal elements (noise variances) reshaped to broadcast against the
transfer function along `axis`.
"""
noise_variance = xp.diagonal(noise_covariance, axis1=-1, axis2=-2)
if axis == -2:
return noise_variance[..., xp.newaxis, :, xp.newaxis]
if axis == -1:
return noise_variance[..., xp.newaxis, xp.newaxis, :]
msg = f"axis must be -2 (target) or -1 (source), got {axis}"
raise ValueError(msg)
def _max_psd_discrepancy(
transfer_function: NDArray[np.complexfloating],
noise_covariance: NDArray[np.floating],
) -> float:
"""Return the largest relative gap between the diagonal-noise PSD and the true PSD.
``directed_coherence`` normalizes each target ``i`` by the diagonal-only power
``D_i = sum_k Cov_kk |H_ik|^2``, which equals the true power spectral density
``S_ii = (H Cov H^H)_ii`` only when the innovations are uncorrelated. This
returns the largest relative gap ``|S_ii - D_i| / S_ii`` over all targets,
frequencies, and batch elements -- a dimension-aware measure of how much
cross-power the diagonal formula omits. Unlike a pairwise-correlation
criterion, it flags many weakly-but-jointly correlated sources whose cross
terms still omit a large fraction of the true power.
Parameters
----------
transfer_function : array, shape (..., n_fft_samples, n_signals, n_signals)
MVAR transfer function ``H`` (axis -2 target, axis -1 source).
noise_covariance : array, shape (..., n_signals, n_signals)
Estimated MVAR innovation covariance ``Cov``.
Returns
-------
float
Maximum relative PSD discrepancy, or 0.0 when it cannot be assessed
(single signal, or all entries non-finite).
"""
if transfer_function.shape[-1] < 2:
return 0.0
H = transfer_function
# True PSD S_ii = (H Cov H^H)_ii = sum_l (H Cov)_il conj(H_il). Broadcast the
# frequency-independent covariance over the frequency axis.
covariance = noise_covariance[..., xp.newaxis, :, :]
true_psd = xp.real(xp.sum(xp.matmul(H, covariance) * xp.conj(H), axis=-1))
# Diagonal-only power D_i = sum_k Cov_kk |H_ik|^2.
noise_variance = xp.real(xp.diagonal(noise_covariance, axis1=-1, axis2=-2))
diagonal_psd = xp.sum(
noise_variance[..., xp.newaxis, xp.newaxis, :] * _squared_magnitude(H),
axis=-1,
)
# The true PSD is non-negative in exact arithmetic (diagonal of a PSD
# matrix); clamp roundoff-negative values to 0 so a near-zero true power with
# nonzero diagonal power reads as an infinite (not a large-negative) gap.
true_psd = xp.maximum(true_psd, 0.0)
with np.errstate(invalid="ignore", divide="ignore"):
relative = xp.abs(true_psd - diagonal_psd) / true_psd
# Keep +inf (true power 0 but diagonal power > 0 -- the diagonal formula is
# infinitely wrong there, the strongest reason to warn); drop only 0/0 NaN
# (no power at all, so nothing to normalize against).
relative = relative[~xp.isnan(relative)]
if relative.size == 0:
return 0.0
return float(xp.max(relative))
def _total_outflow(
MVAR_Fourier_coefficients: NDArray[np.complexfloating],
noise_variance: float | NDArray[np.floating] = 1.0,
) -> NDArray[np.floating]:
"""Measure effect of outgoing signals on the node via sum of squares.
Parameters
----------
MVAR_Fourier_coefficients : array_like
MVAR Fourier coefficients.
noise_variance : float or array_like, default=1.0
Noise variance values.
Returns
-------
array_like
Total outflow values.
"""
return xp.sqrt(
xp.sum(
_squared_magnitude(MVAR_Fourier_coefficients) / noise_variance,
keepdims=True,
axis=-2,
)
)
def _unwrap_across_undefined(
phase: NDArray[np.floating], axis: int = -1
) -> NDArray[np.floating]:
"""``np.unwrap`` that skips NaN (undefined) phases instead of spreading them.
``np.unwrap`` accumulates corrections along ``axis``, so a single NaN makes
every later value NaN. Each NaN is instead held at the last defined phase
(or the first, before any is defined) while unwrapping, which unwraps the
defined values as if the undefined ones were absent, and is restored to NaN
afterwards. Without NaN the result equals ``np.unwrap``.
"""
phase = np.moveaxis(phase, axis, -1)
is_defined = ~np.isnan(phase)
positions = np.arange(phase.shape[-1])
last_defined = np.maximum.accumulate(np.where(is_defined, positions, 0), axis=-1)
first_defined = np.argmax(is_defined, axis=-1)[..., np.newaxis]
source = np.where(np.cumsum(is_defined, axis=-1) == 0, first_defined, last_defined)
filled = np.take_along_axis(phase, source, axis=-1)
unwrapped = np.where(is_defined, np.unwrap(filled, axis=-1), np.nan)
return np.moveaxis(unwrapped, -1, axis)
def _bandpass(
data: NDArray[np.complexfloating],
frequencies: NDArray[np.floating],
frequencies_of_interest: NDArray[np.floating] | None,
axis: int = -3,
) -> tuple[NDArray[np.complexfloating], NDArray[np.floating]]:
"""Filter data matrix along axis for frequencies of interest.
Filters the data matrix along an axis given a maximum and minimum
frequency of interest. Both band edges are exclusive: a frequency bin
lying exactly on an edge is dropped.
Parameters
----------
data : array, shape (..., n_fft_samples, ...)
Input data array.
frequencies : array, shape (n_fft_samples,)
Frequency values.
frequencies_of_interest : array-like, shape (2,)
Min and max frequencies of interest; two finite values with
``low < high``.
axis : int, default=-3
Axis along which to filter.
Returns
-------
filtered_data : array
Filtered data.
filtered_frequencies : array
Corresponding filtered frequencies.
Raises
------
ValueError
If ``frequencies_of_interest`` is not two finite values with
``low < high``, or no frequency lies strictly inside the band.
"""
if frequencies_of_interest is None:
return data, frequencies
band = np.asarray(frequencies_of_interest, dtype=float)
if band.shape != (2,) or not np.all(np.isfinite(band)) or not band[0] < band[1]:
msg = (
"frequencies_of_interest must be two finite values (low, high) with "
f"low < high, got {frequencies_of_interest!r}. Band edges are "
"exclusive: only frequencies strictly inside (low, high) are kept."
)
raise ValueError(msg)
frequency_index = _frequencies_in_band(frequencies, band)
if not bool(frequency_index.any()):
msg = (
f"frequencies_of_interest {frequencies_of_interest!r} contains no "
"frequency bin: band edges are exclusive, and the frequencies span "
f"{float(frequencies.min())} to {float(frequencies.max())} "
f"({frequencies.shape[0]} bins). Widen the band."
)
raise ValueError(msg)
return (
xp.take(data, frequency_index.nonzero()[0], axis=axis),
frequencies[frequency_index],
)
def _frequencies_in_band(
frequencies: NDArray[np.floating], band: NDArray[np.floating] | Sequence[float]
) -> NDArray[np.bool_]:
"""Boolean mask of ``frequencies`` strictly inside ``(low, high)``.
The single definition of the exclusive band edges used by :func:`_bandpass`
and by the wrapper's coordinates for band-restricted results.
"""
in_band: NDArray[np.bool_] = (band[0] < frequencies) & (frequencies < band[1])
return in_band
def _get_independent_frequency_step(
frequency_difference: float, frequency_resolution: float | None
) -> int:
"""Find number of points for statistically independent frequencies.
Find the number of points of a frequency axis such that they
are statistically independent given a frequency resolution.
Parameters
----------
frequency_difference : float
The distance between two frequency points.
frequency_resolution : float | None
The ability to resolve frequency points. If None, returns 1.
Returns
-------
frequency_step : int
The number of points required so that two
frequency points are statistically independent.
"""
if frequency_resolution is None:
return 1
if not np.isfinite(frequency_resolution) or frequency_resolution <= 0:
msg = (
f"frequency_resolution must be a finite positive number when "
f"provided, got {frequency_resolution}."
)
raise ValueError(msg)
return int(xp.ceil(frequency_resolution / frequency_difference))
# Element cap (rows * n_frequencies) for one chunk of the significant-frequency
# selector. The per-slice int32 run-length temporaries dominate its memory, so
# processing the flattened signal-pair slices in chunks keeps peak usage bounded
# regardless of the number of slices.
_SIGNIFICANCE_SELECTION_CHUNK_ELEMENTS = 2_000_000
def _select_largest_independent_cluster(
block: NDArray[np.bool_], frequency_step: int, min_group_size: int
) -> NDArray[np.bool_]:
"""Largest independent significant cluster per row (frequency on last axis).
``block`` has shape ``(n_rows, n_frequencies)``. See
``_largest_independent_group_along_frequency`` for the selection rule.
"""
n_frequencies = block.shape[-1]
# run_length[r, f] = length of the contiguous True-run ending at f (0 where
# False): a cumulative count that resets at each False (running count minus
# its value at the most recent False). int32 suffices (runs <= n_frequencies)
# and halves the temporaries relative to the default int64.
cumulative = np.cumsum(block, axis=-1, dtype=np.int32)
run_length = cumulative - np.maximum.accumulate(
np.where(block, np.int32(0), cumulative), axis=-1
)
# Largest run per row; run_length only reaches this value at a run's end, so
# the first index attaining it is the end of the first largest cluster
# (matching the "first cluster on ties" rule).
max_size = run_length.max(axis=-1, keepdims=True)
end_index = np.expand_dims(np.argmax(run_length == max_size, axis=-1), -1)
start_index = end_index - max_size + 1
frequency_index = np.arange(n_frequencies)
# The largest cluster is contiguous [start_index, end_index]; max_size == 0
# means no significant frequency, giving an all-False row.
in_largest_cluster = (
(frequency_index >= start_index) & (frequency_index <= end_index) & (max_size > 0)
)
# Independent points are start_index, start_index + frequency_step, ...
independent: NDArray[np.bool_] = in_largest_cluster & (
(frequency_index - start_index) % frequency_step == 0
)
count: NDArray[np.integer] = independent.sum(axis=-1, keepdims=True)
return independent & (count >= min_group_size)
def _largest_independent_group_along_frequency(
is_significant: NDArray[np.bool_], frequency_step: int, min_group_size: int
) -> NDArray[np.bool_]:
"""Largest independent significant-frequency group along axis -2.
For every slice along axis -2, keep the largest contiguous cluster of
significant frequencies (the first cluster on ties), subsample it every
``frequency_step`` points, and drop the slice to all-False if fewer than
``min_group_size`` independent points remain. Computed for all slices at
once instead of one Python call per slice; the slices are processed in
bounded chunks so peak memory stays independent of their number.
Parameters
----------
is_significant : bool array, shape (..., n_frequencies, n_signal_pairs)
frequency_step : int
Spacing (in points) between retained independent frequencies.
min_group_size : int
Minimum number of independent points for a cluster to be kept.
Returns
-------
bool array, same shape as ``is_significant``.
"""
axis = -2
n_frequencies = is_significant.shape[axis]
if is_significant.size == 0:
# An empty frequency band (or any zero-length axis) has no cluster to
# select; return an all-False array of the same (empty) shape.
return np.zeros(is_significant.shape, dtype=bool)
# Move frequency to the last axis and flatten the rest so the signal-pair
# slices can be processed in chunks. asarray(copy=False) avoids a needless
# copy of an already-boolean input; the reshape of the moved (non-contiguous)
# view then makes a single bool copy.
moved = np.moveaxis(np.asarray(is_significant, dtype=bool), axis, -1)
flattened = moved.reshape(-1, n_frequencies)
result = np.empty_like(flattened)
chunk = max(1, _SIGNIFICANCE_SELECTION_CHUNK_ELEMENTS // max(1, n_frequencies))
for start in range(0, flattened.shape[0], chunk):
block = flattened[start : start + chunk]
result[start : start + chunk] = _select_largest_independent_cluster(
block, frequency_step, min_group_size
)
return np.moveaxis(result.reshape(moved.shape), -1, axis)
def _find_significant_frequencies(
coherency: NDArray[np.complexfloating],
n_obs: int,
frequency_step: int = 1,
significance_threshold: float = 0.05,
min_group_size: int = 3,
multiple_comparisons_method: Literal[
"Benjamini_Hochberg_procedure", "Bonferroni_correction"
] = "Benjamini_Hochberg_procedure",
) -> NDArray[np.bool_]:
"""Determine the largest significant cluster along the frequency axis.
This function uses the exact zero-coherence null distribution to determine
the p-values and adjusts for multiple comparisons using the
`multiple_comparisons_method`. Only independent frequencies are
returned and there must be at least `min_group_size` frequency
points for the cluster to be returned. If there are several significant
groups, then only the largest group is returned.
Parameters
----------
coherency : array, shape (..., n_frequencies, n_signals, n_signals)
The complex coherency between signals.
n_obs : int
The number of observations used to estimate the coherency.
frequency_step : int
The number of points between each independent frequency step
significance_threshold : float
The threshold for a p-value to be considered significant.
min_group_size : int
The minimum number of independent frequency points for
multiple_comparisons_method : 'Benjamini_Hochberg_procedure' |
'Bonferroni_correction'
Procedure used to correct for multiple comparisons.
Returns
-------
is_significant : bool array, shape (..., n_frequencies,
n_signal_combintaions)
"""
# Test each frequency's coherence against zero using the exact
# magnitude-squared-coherence null distribution. The Fisher z-transform is
# miscalibrated at the zero boundary and over-rejects the null ~3-4x.
p_values = coherence_significance_pvalue(coherency, n_obs)
is_significant = adjust_for_multiple_comparisons(
p_values, alpha=significance_threshold, method=multiple_comparisons_method
)
return _largest_independent_group_along_frequency(
is_significant, frequency_step, min_group_size
)