"""Functions for getting connectivity measures in a labeled array format."""
import warnings
from collections.abc import Hashable, Mapping, Sequence
from dataclasses import dataclass
from logging import getLogger
from typing import Any, Literal
import numpy as np
import xarray as xr
from numpy.typing import DTypeLike, NDArray
from spectral_connectivity._frequency_bands import (
_nyquist_frequency,
_reduce_frequency_bands,
_select_and_reduce_frequencies,
)
from spectral_connectivity._input_handling import (
_UNSET,
_is_real_numeric_dtype,
_SignalLabel,
_SignalMetadata,
_unwrap_fourier_input,
_unwrap_xarray_input,
_validated_signal_labels,
)
from spectral_connectivity._measure_registry import (
_MEASURE_SPECS,
_is_group_measure,
_measure_description,
_requested_methods,
_requires_two_sided,
_validate_method_names,
)
from spectral_connectivity._provenance import _canonical_json, _shared_provenance_attrs
from spectral_connectivity._result_formatting import (
UnsupportedMeasureError,
_connectivity_result_to_xarray,
)
from spectral_connectivity.connectivity import (
_MIGRATION_GUIDE_URL,
_ORIENTATION_WARNING_PREFIX,
_SILENCE_ORIENTATION_WARNING,
Connectivity,
DirectedOrientationWarning,
_frequencies_in_band,
_validated_flag,
)
from spectral_connectivity.transforms import Multitaper
from spectral_connectivity.utils import stacklevel_outside_package, to_numpy
# The public API. UnsupportedMeasureError is defined in _result_formatting
# (which raises it) and re-exported here; listing it makes the API reference
# document it with the wrapper (see autosummary_ignore_module_all in docs/conf.py).
__all__ = [
"DEFAULT_METHODS",
"MeasureInfo",
"UnsupportedMeasureError",
"connectivity_to_xarray",
"fourier_connectivity",
"frequency_band_reduce",
"list_measures",
"multitaper_connectivity",
]
logger = getLogger(__name__)
#: Measures computed when ``method`` is omitted, in result-variable order.
DEFAULT_METHODS: tuple[str, ...] = tuple(
name for name, spec in _MEASURE_SPECS.items() if spec.is_default
)
[docs]
@dataclass(frozen=True)
class MeasureInfo:
"""A single connectivity measure the high-level wrapper can compute.
Attributes
----------
name : str
Value to pass as ``method`` to :func:`multitaper_connectivity` or
:func:`fourier_connectivity`, and the name of the corresponding
:class:`~spectral_connectivity.Connectivity` method.
category : str
The output-shape contract, one of ``"pairwise"``, ``"power"``,
``"group_pairwise"``, ``"multivariate_components"``, ``"delay"``,
``"global"``, ``"group_delay"``, or ``"phase_slope"``.
description : str
One-line summary taken from the ``Connectivity`` method's docstring.
is_default : bool
Whether the measure is in the default set computed when ``method`` is
omitted (see ``DEFAULT_METHODS``).
is_directed : bool
Whether the measure is directional (``source -> target`` asymmetric).
requires_two_sided : bool
Whether the measure requires a full two-sided spectrum, including
negative-frequency bins.
long_name : str
Human-readable name, also the ``long_name`` attribute of the result.
units : str
Units of the values: ``"1"`` for a dimensionless score, ``"rad"``,
``"s"``, or ``"(input units)^2/Hz"`` for a spectral density.
value_range : tuple of float
``(lower, upper)`` bounds of the values, or of their magnitude when
``is_complex``; ``inf`` marks an unbounded side.
is_complex : bool
Whether the values are complex.
dims : tuple of str
Dimensions of the measure's variable in the wrapper's result, before
band reduction or squeezing. Rich results (``canonical_coherency``,
``global_coherence``, ...) are Datasets whose variable named ``name``
has these dimensions.
interpretation : str
How to read the values, including the sign convention where one exists.
"""
name: str
category: str
description: str
is_default: bool
is_directed: bool
requires_two_sided: bool
long_name: str
units: str
value_range: tuple[float, float]
is_complex: bool
dims: tuple[str, ...]
interpretation: str
[docs]
def list_measures(
*,
category: str | None = None,
default_only: bool = False,
directed: bool | None = None,
) -> list[MeasureInfo]:
"""List the connectivity measures the high-level wrapper can compute.
This is the discovery entry point: it enumerates every valid ``method``
name for :func:`multitaper_connectivity` and :func:`fourier_connectivity`,
together with each measure's output category, a one-line description, and
whether it is in the default set and/or directional, and what its values
mean: units, range, result dimensions, and how to read the direction.
Parameters
----------
category : str, optional
Return only measures with this output category (e.g. ``"pairwise"``,
``"power"``, ``"group_pairwise"``). Raises ``ValueError`` for an
unknown category.
default_only : bool, default False
Return only the measures computed when ``method`` is omitted.
directed : bool, optional
If ``True``, return only directional measures; if ``False``, only
non-directional ones; if ``None`` (default), return both. Non-directional
does not necessarily mean symmetric: for example, phase-valued measures
may be antisymmetric and complex coherency is Hermitian.
Returns
-------
measures : list of MeasureInfo
One record per measure, in the wrapper's canonical order.
Examples
--------
>>> from spectral_connectivity import list_measures
>>> [m.name for m in list_measures(default_only=True)][:3]
['coherence_magnitude', 'coherence_phase', 'debiased_squared_phase_lag_index']
>>> next(m for m in list_measures() if m.name == "power").description
'Return the one-sided power spectral density of the signal.'
"""
valid_categories = {spec.output_kind for spec in _MEASURE_SPECS.values()}
if category is not None and category not in valid_categories:
msg = (
f"Unknown category {category!r}. Valid categories are: "
f"{', '.join(sorted(valid_categories))}."
)
raise ValueError(msg)
measures = []
for name, spec in _MEASURE_SPECS.items():
if default_only and not spec.is_default:
continue
if category is not None and spec.output_kind != category:
continue
if directed is not None and spec.is_directed != directed:
continue
measures.append(
MeasureInfo(
name=name,
category=spec.output_kind,
description=_measure_description(name),
is_default=spec.is_default,
is_directed=spec.is_directed,
requires_two_sided=spec.requires_two_sided,
long_name=spec.long_name,
units=spec.units_for("input units"),
value_range=spec.value_range,
is_complex=spec.is_complex,
dims=spec.dims,
interpretation=spec.interpretation,
)
)
return measures
[docs]
def frequency_band_reduce(
result: xr.DataArray | xr.Dataset,
bands: Mapping[str, tuple[float, float]],
*,
reduction: Literal["mean", "integral"] = "mean",
circular: bool | None = None,
) -> xr.DataArray | xr.Dataset:
"""Reduce a frequency-resolved result into labeled frequency bands.
``reduction="mean"`` averages the already-computed connectivity score over
the bins in each inclusive band. Phase is treated specially: a
``coherence_phase`` result uses a circular mean, while complex-valued
measures use their ordinary complex (vector) mean. ``reduction="integral"``
integrates a spectral density over ``[low, high]`` and is intentionally
restricted to ``power`` and ``cross_spectral_density``, where it represents
band power/covariance rather than a frequency-averaged score. Each bin
stands for the frequency cell between the midpoints to its neighbours, and
contributes its density times the part of that cell inside the band, so
band edges need not fall on bins, a one-bin band is not zero, and adjacent
bands add up to their union. On a one-sided grid starting at 0 Hz the DC
bin and the Nyquist bin (present for an even FFT length) own only the
half-cell toward their neighbour but count with a full spacing, matching
the one-sided convention in which those two bins are not doubled;
integrating ``power`` over a band that covers every bin therefore
reproduces the total signal power (Parseval).
Parameters
----------
result : xarray.DataArray or xarray.Dataset
Result with a one-dimensional ``frequency`` coordinate.
bands : mapping of str to (float, float)
Inclusive lower and upper frequency bounds in the coordinate's units.
reduction : {"mean", "integral"}, default="mean"
Scientifically defined reduction to apply within each band.
circular : bool, optional
Use a circular mean (for phase angles in radians). By default it is
inferred per variable: ``coherence_phase`` results and variables with
``units="rad"`` are averaged circularly. Pass ``True``/``False`` to
override, e.g. for a phase array whose name and attrs were removed.
Returns
-------
xarray.DataArray or xarray.Dataset
Same type as ``result`` with the ``frequency`` dimension replaced by a
``band`` dimension holding the band names, and the band definitions
recorded in ``attrs["frequency_bands_json"]``. An integral is in the
density's ``units`` without the ``/Hz`` (e.g. ``(uV)^2/Hz`` becomes
``(uV)^2``) and is labeled ``"Band power"`` or ``"Band cross-power"``.
Notes
-----
Spatial filters, spatial patterns, and global-coherence vectors have an
arbitrary sign or complex phase independently at each frequency. A Dataset
containing those variables is therefore rejected instead of averaging them
into a scientifically undefined band projection. Select the scalar score
variable from the Dataset and reduce that DataArray when only band scores
are needed.
A band is undefined wherever any of its bins is ``NaN`` (for example an
edge-invalid ``MorletWavelet`` bin under ``edge_mode="nan"``): both
reductions return ``NaN`` there rather than silently reducing the valid
bins only. When the input carries a ``valid_time_frequency`` coordinate the
result gains a ``valid_time_band`` coordinate that is ``True`` only where
every bin of the band had full support.
The Nyquist frequency is half the sampling rate recorded in ``result``'s
provenance attrs (e.g. ``mt_sampling_frequency``). If none is recorded,
the last bin of a grid starting at 0 Hz is taken to be the Nyquist bin, so
reduce such a result before cropping it.
"""
return _reduce_frequency_bands(
result,
bands,
reduction=reduction,
circular=circular,
nyquist_frequency=_nyquist_frequency(result),
)
# The directed measures whose 2.x wrapper labels were transposed. The 2.x
# wrapper rejected every method named "directed", so the transfer-function
# measures had no 2.x labels to flip.
_LABEL_ORIENTATION_CHANGED_MEASURES = frozenset(
{"pairwise_spectral_granger_prediction", "subset_pairwise_spectral_granger_prediction"}
)
def _warn_label_orientation_changed(methods: Sequence[str]) -> None:
"""Warn that ``sel(source, target)`` of the 2.x Granger measures flipped.
Remove with :class:`DirectedOrientationWarning` in 3.2.
"""
changed = sorted(_LABEL_ORIENTATION_CHANGED_MEASURES.intersection(methods))
if changed:
warnings.warn(
f"{_ORIENTATION_WARNING_PREFIX}sel(source=a, target=b) of "
f"{', '.join(changed)} is a -> b; 2.x returned b -> a. Review code written "
"for 2.x that selects these results. See the migration guide: "
f"{_MIGRATION_GUIDE_URL}. {_SILENCE_ORIENTATION_WARNING}",
DirectedOrientationWarning,
stacklevel=stacklevel_outside_package(),
)
[docs]
def connectivity_to_xarray(
m: Any,
method: str = "coherence_magnitude",
signal_names: Sequence[_SignalLabel] | None = None,
squeeze: bool = False,
**kwargs: Any,
) -> xr.DataArray | xr.Dataset:
"""Calculate one connectivity measure and return a labeled array.
Ordinary pairwise measures return a DataArray; component-resolved or
multi-quantity measures return a Dataset with explicit semantic axes.
.. versionchanged:: 3.0
For ``pairwise_spectral_granger_prediction`` and
``subset_pairwise_spectral_granger_prediction``,
``sel(source=a, target=b)`` is ``a -> b``; 2.x returned ``b -> a``. The
2.x wrapper rejected the transfer-function measures, and
``phase_slope_index``, ``group_delay``, and ``delay`` are unchanged.
Parameters
----------
m : transform
A spectral transform (e.g. ``Multitaper``, ``MorletWavelet``, ``Welch``)
whose coefficients the measure is computed from.
method : str, default="coherence_magnitude"
Measure name from :func:`list_measures`.
signal_names : sequence, optional
Labels for the ``source``/``target`` coordinates; defaults to
``"0"``, ``"1"``, ....
squeeze : bool, default=False
With exactly 2 signals, reduce a pairwise measure to the ordered pair
(first source, last target), keeping ``source`` and ``target`` as
scalar coordinates.
**kwargs
Keyword arguments for the measure (e.g. ``group_labels``).
Returns
-------
xarray.DataArray or xarray.Dataset
The labeled result with provenance in ``attrs``; see
:func:`multitaper_connectivity` for the dimensions and orientation
(``sel(source=a, target=b)`` is the influence ``a -> b``).
Examples
--------
>>> import numpy as np
>>> from spectral_connectivity.transforms import Multitaper
>>> data = np.random.default_rng(0).standard_normal((100, 5, 3))
>>> mt = Multitaper(data, sampling_frequency=1000)
>>> connectivity_to_xarray(mt).dims
('time', 'frequency', 'source', 'target')
"""
_validate_method_names([method])
metadata = m._provenance_metadata()
connectivity = Connectivity.from_transform(m)
signal_labels = _validated_signal_labels(signal_names, connectivity.n_signals)
shared_attrs = _shared_provenance_attrs(
connectivity,
metadata,
transform_prefix=getattr(m, "_provenance_prefix", "mt_"),
)
result = _connectivity_result_to_xarray(
connectivity, method, signal_labels, squeeze, shared_attrs, **kwargs
)
valid_time_frequency = getattr(m, "valid_time_frequency", None)
if valid_time_frequency is not None:
validity = to_numpy(valid_time_frequency).astype(bool)
expected_shape = (len(connectivity.time), len(connectivity.frequencies))
if validity.shape != expected_shape:
msg = (
"transform.valid_time_frequency must have shape "
f"{expected_shape}, got {validity.shape}."
)
raise ValueError(msg)
validity_attrs = {"long_name": "Full wavelet and smoothing support is in-record"}
if "frequency" in result.dims:
full_validity = xr.DataArray(
validity,
coords={
"time": np.asarray(connectivity.time),
"frequency": np.asarray(connectivity.frequencies),
},
dims=("time", "frequency"),
attrs=validity_attrs,
)
# Delay and other nonstandard schemas may retain only a requested
# frequency band. Select the matching validity bins rather than
# attaching the transform's full frequency axis to the result.
result_frequencies = np.asarray(result.coords["frequency"])
aligned_validity = full_validity.sel(frequency=result_frequencies)
result = result.assign_coords(valid_time_frequency=aligned_validity)
elif "time" in result.dims:
# PSI and group delay aggregate a frequency band. They have no
# frequency dimension on which a 2-D coordinate can live, so expose
# whether every frequency contributing to each time point has full
# wavelet/smoothing support.
frequencies = np.asarray(connectivity.frequencies)
frequency_band = kwargs.get("frequencies_of_interest")
if frequency_band is None:
frequency_index = np.ones(frequencies.shape, dtype=bool)
else:
frequency_index = _frequencies_in_band(frequencies, frequency_band)
valid_time = validity[:, frequency_index].all(axis=1)
result = result.assign_coords(valid_time=(("time",), valid_time, validity_attrs))
_warn_label_orientation_changed([method])
return result
def _combine_formatted_results(
results: Sequence[xr.DataArray | xr.Dataset],
shared_attrs: Mapping[str, Any],
) -> xr.Dataset:
"""Merge heterogeneous formatted measures without losing sub-variables."""
datasets = [
result.to_dataset(name=result.name) if isinstance(result, xr.DataArray) else result
for result in results
]
try:
combined = xr.merge(datasets, compat="no_conflicts", join="exact")
except ValueError as error:
msg = (
"Requested measures produced conflicting xarray variables or "
"coordinates; request them separately or use compatible group labels."
)
raise ValueError(msg) from error
combined.attrs = dict(shared_attrs)
return combined
def _resolve_group_labels(
methods: Sequence[str],
connectivity_kwargs: Mapping[str, Any] | None,
group_labels: Sequence[Hashable] | NDArray[Any] | None,
n_signals: int,
) -> tuple[Sequence[Hashable] | NDArray[Any] | None, dict[str, Any], frozenset[str]]:
"""Resolve ``group_labels`` from the argument or ``connectivity_kwargs``.
Runs before any spectral work so a missing or misplaced label fails fast.
Parameters
----------
methods : sequence of str
The requested measure names.
connectivity_kwargs : mapping or None
The caller's measure keyword arguments; copied, never mutated.
group_labels : sequence, array or None
The ``group_labels=`` argument.
n_signals : int
Number of signals, quoted in the missing-labels message.
Returns
-------
group_labels : sequence, array or None
The labels from whichever form supplied them.
label_free_kwargs : dict
A copy of ``connectivity_kwargs`` without ``group_labels``.
group_methods : frozenset of str
The requested measures that take ``group_labels``.
Raises
------
ValueError
If the labels are given both ways, given with no group measure
requested, or missing while a group measure is requested.
"""
label_free_kwargs = dict(connectivity_kwargs or {})
legacy_labels = label_free_kwargs.pop("group_labels", None)
if group_labels is not None and legacy_labels is not None:
msg = (
"group_labels was given both as an argument and inside "
"connectivity_kwargs; pass it once, as group_labels=..."
)
raise ValueError(msg)
group_labels = group_labels if group_labels is not None else legacy_labels
group_methods = [method for method in methods if _is_group_measure(method)]
if group_labels is not None and not group_methods:
msg = (
"group_labels was given, but none of the requested measures compares "
f"groups of signals: {list(methods)!r}. Request a group measure (see "
"list_measures(category='group_pairwise') or "
"list_measures(category='multivariate_components')) or drop group_labels."
)
raise ValueError(msg)
if group_methods and group_labels is None:
msg = (
f"group_labels is required for {', '.join(group_methods)}: group measures "
"compare groups of signals and need one label per signal "
f"({n_signals} here) naming the group each signal belongs to, e.g. "
"group_labels=['CA1', 'CA1', 'PFC', 'PFC'].\n"
"Pass it as group_labels=... (the same labels apply to every group "
"measure in this call)."
)
raise ValueError(msg)
return group_labels, label_free_kwargs, frozenset(group_methods)
def _format_and_reduce_measures(
connectivity: Connectivity,
methods: list[str],
*,
return_dataarray: bool,
signal_labels: NDArray[Any],
squeeze: bool,
shared_attrs: Mapping[str, Any],
connectivity_kwargs: Mapping[str, Any],
group_labels: Sequence[Hashable] | NDArray[Any] | None,
group_methods: frozenset[str],
transform_settings_hint: str,
frequency_range: tuple[float, float] | None,
frequency_decimation: int,
frequency_bands: Mapping[str, tuple[float, float]] | None,
frequency_reduction: Literal["mean", "integral"],
signal_metadata: _SignalMetadata | None = None,
) -> xr.DataArray | xr.Dataset:
"""Format the requested measures to xarray and apply frequency reduction.
Shared tail of :func:`multitaper_connectivity` and :func:`fourier_connectivity`:
passes the already-resolved ``group_labels`` (see :func:`_resolve_group_labels`)
to ``group_methods`` only, honors ``squeeze`` only for a single-measure
DataArray, formats each measure (skipping structurally-unsupported ones in a
multi-measure batch), merges the survivors, and applies any frequency
crop/decimation/band reduction. ``connectivity_kwargs`` must be label-free;
``transform_settings_hint`` tells a rejected keyword argument where transform
settings go for the calling wrapper.
"""
def measure_kwargs(method: str) -> Mapping[str, Any]:
"""Keyword arguments for one measure: ``group_labels`` reaches group measures only."""
if method in group_methods:
return {**connectivity_kwargs, "group_labels": group_labels}
return connectivity_kwargs
if squeeze and not return_dataarray:
# squeeze reduces a pairwise measure to a (time, frequency) array whose
# source/target become scalar coordinates; in a Dataset those scalars are
# shared across variables and collide with a sibling's axes, so squeeze is
# honored only for a single-method DataArray.
warnings.warn(
"squeeze=True is ignored for multi-measure results (a Dataset); "
"request a single method (a string) to get a squeezed DataArray.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
squeeze = False
if return_dataarray:
result: xr.DataArray | xr.Dataset = _connectivity_result_to_xarray(
connectivity,
methods[0],
signal_labels,
squeeze,
shared_attrs,
signal_metadata=signal_metadata,
transform_settings_hint=transform_settings_hint,
**measure_kwargs(methods[0]),
)
else:
formatted_results: list[xr.DataArray | xr.Dataset] = []
for this_method in methods:
try:
formatted_results.append(
_connectivity_result_to_xarray(
connectivity,
this_method,
signal_labels,
False,
shared_attrs,
signal_metadata=signal_metadata,
transform_settings_hint=transform_settings_hint,
**measure_kwargs(this_method),
)
)
except UnsupportedMeasureError as error: # noqa: PERF203 -- per-measure skip
# A measure whose result shape does not fit the xarray layout can
# be skipped in a batch. _connectivity_result_to_xarray raises
# UnsupportedMeasureError from its shape check after the measure
# has run (an unregistered extension returning a non-pairwise
# shape); a genuine NotImplementedError is not caught, so a broken
# measure fails loudly instead of silently vanishing from the
# Dataset.
if len(methods) == 1:
raise
logger.warning("Skipping %s: %s", this_method, error)
if not formatted_results:
msg = (
"None of the requested methods produced a compatible result "
f"for the xarray interface: {methods!r}."
)
raise UnsupportedMeasureError(msg)
result = _combine_formatted_results(formatted_results, shared_attrs)
return _select_and_reduce_frequencies(
result,
frequency_range=frequency_range,
frequency_decimation=frequency_decimation,
frequency_bands=frequency_bands,
frequency_reduction=frequency_reduction,
)
[docs]
def multitaper_connectivity(
time_series: NDArray[np.floating] | xr.DataArray,
sampling_frequency: float | None = None,
time_window_duration: float | None = None,
method: str | list[str] | None = None,
signal_names: Sequence[_SignalLabel] | None = None,
squeeze: bool = False,
connectivity_kwargs: dict[str, Any] | None = None,
*,
group_labels: Sequence[Hashable] | NDArray[Any] | None = None,
frequency_range: tuple[float, float] | None = None,
frequency_decimation: int = 1,
frequency_bands: Mapping[str, tuple[float, float]] | None = None,
frequency_reduction: Literal["mean", "integral"] = "mean",
time_dim: Hashable | None = None,
trial_dim: Hashable | None = None,
signal_dim: Hashable | None = None,
**kwargs: Any,
) -> xr.DataArray | xr.Dataset:
"""
Compute connectivity measures with multitaper spectral estimation.
This is the main high-level function for connectivity analysis. It performs
multitaper spectral analysis on the input time series and computes the
requested connectivity measures, returning results as labeled xarray objects.
.. versionchanged:: 3.0
For ``pairwise_spectral_granger_prediction`` and
``subset_pairwise_spectral_granger_prediction``,
``sel(source=a, target=b)`` is ``a -> b``; 2.x returned ``b -> a``. The
2.x wrapper rejected the transfer-function measures, and
``phase_slope_index``, ``group_delay``, and ``delay`` are unchanged.
Parameters
----------
time_series : NDArray[floating] or xarray.DataArray,
shape (n_times, n_trials, n_channels) or (n_times, n_channels)
Time series data. For multiple trials, trials are averaged in spectral
domain. For a DataArray, common time/trial/signal dimension names are
inferred and transposed automatically; use ``time_dim``, ``trial_dim``,
and ``signal_dim`` for domain-specific names. Ambiguous names raise rather
than falling back to dimension position, though a single unrecognized
dimension left for the one remaining role is assigned by elimination with
a warning (a spectral name such as ``frequency`` or ``band`` is rejected
instead, since it marks an already-transformed input). A dask-backed
DataArray is rejected (materialize it first with
``DataArray.compute()``). A numeric time index is
interpreted as elapsed seconds (a ``sample`` index as sample numbers)
and used to label output window centers. When ``sampling_frequency`` is
given it is validated against the index; when it is omitted, an
elapsed-seconds ``time`` coordinate infers it (a ``sample`` index cannot,
having no time scale). Datetime, timedelta, and object-valued time
coordinates are not yet supported and must first be converted to numeric
elapsed seconds.
When ``signal_names`` is omitted,
labels from a 1-D index coordinate on the signal dimension are carried to
the output's ``source`` and ``target`` coordinates without changing their
type; if such labels are present but unusable a warning is issued and
default string labels are used.
sampling_frequency : float, optional
Sampling rate in Hz of the time series data. Required for array input;
for a DataArray it may be omitted and inferred from a sufficiently precise
numeric elapsed-seconds ``time`` coordinate. Pass it explicitly when a
low-precision or large-offset coordinate cannot resolve the rate reliably.
time_window_duration : float, optional
Duration of sliding window in seconds for time-resolved analysis.
If None, analyzes entire time series (no time resolution).
method : str or list of str, optional
Connectivity method(s) to compute. If None, computes the default set of
real-valued measures that fit the xarray/NetCDF interface (see
``DEFAULT_METHODS``) — not every measure. ``coherency`` is left out of the
default because complex arrays are not portably serializable across all
supported xarray versions and NetCDF engines, but it can be requested by
name. Every directed measure is opt-in by name, including
``pairwise_spectral_granger_prediction`` and the
directed-transfer-function family (``directed_transfer_function``,
``directed_coherence``, ``partial_directed_coherence``,
``generalized_partial_directed_coherence``,
``direct_directed_transfer_function``); the spectral Granger and
transfer-function measures factorize the spectrum and dominate the cost
(see the Notes on directed orientation). Measures with nonstandard layouts,
including ``global_coherence``, ``phase_slope_index``, ``group_delay``,
``delay``, ``canonical_coherence``, and blockwise spectral Granger, are
available by name and return labeled DataArrays or Datasets with their
component, group, candidate-delay, or frequency-reduced dimensions.
Examples:
"coherence_magnitude", "imaginary_coherence", "phase_locking_value".
signal_names : sequence of scalar, optional
Scalar, non-missing, unique xarray-compatible coordinate labels for signal
channels. Integer labels must fit the signed 32-bit range for portable
NetCDF3 serialization. Nested or structured labels are not supported. If
None, uses the DataArray signal index when available, otherwise stringified
indices.
squeeze : bool, default=False
Only honored when a single ``method`` (a string) is requested, so the
result is a DataArray. If there are exactly 2 channels, reduce a pairwise
measure to the single ordered pair (first source, last target), returning
a ``(time, frequency)`` array whose selected ``source`` and ``target`` are
retained as scalar coordinates -- so the pair (and, for directed measures,
the direction) is still recorded. With more than 2 channels a warning is
issued and the full matrix is returned; for ``power`` (no target axis)
squeeze is a no-op. For multi-measure requests (which return a Dataset,
whose variables can have incompatible axes such as ``power``'s), squeeze
is ignored with a warning.
connectivity_kwargs : dict, optional
Extra keyword arguments for the *measure* (a ``Connectivity`` method),
e.g. ``pairs`` for ``subset_pairwise_spectral_granger_prediction`` or
``n_components`` for ``canonical_coherency``; passed to every requested
measure. Transform settings do not go here (see ``**kwargs``).
group_labels : sequence, optional
One label per signal naming the group it belongs to; labels are scalars
such as integers or area names (``"CA1"``), and missing values
(``None``, NaN) are rejected. Required by
the group measures (``canonical_coherence``, ``canonical_coherency``,
``maximized_imaginary_coherency``, ``multivariate_interaction_measure``,
``blockwise_spectral_granger_prediction`` and
``maximized_imaginary_coherency_components``); an error is raised if no
requested measure takes it (labels are passed only to the measures that
do).
frequency_range : (float, float), optional
Inclusive frequency interval retained in the labeled result.
frequency_decimation : int, default=1
Keep every Nth frequency bin after applying ``frequency_range``.
frequency_bands : mapping of str to (float, float), optional
Reduce the selected bins into named, inclusive bands. With
``frequency_reduction="mean"``, scores are averaged, complex measures
use a complex vector mean, and ``coherence_phase`` uses a circular mean.
frequency_reduction : {"mean", "integral"}, default="mean"
Band reduction. Integration is restricted to ``power`` and
``cross_spectral_density``, where it yields band power/covariance.
time_dim : hashable, optional
DataArray dimension containing time samples. Common names such as
``"time"`` and ``"sample"`` are inferred automatically.
trial_dim : hashable, optional
DataArray dimension containing trials or epochs. Required for a 3-D
DataArray when its role cannot be inferred unambiguously.
signal_dim : hashable, optional
DataArray dimension containing signals or channels. Common names such as
``"signal"`` and ``"channel"`` are inferred automatically.
**kwargs
Extra keyword arguments for the *transform* (``Multitaper``), e.g.
``time_halfbandwidth_product``, ``n_tapers``, ``time_window_step``,
``taper_weighting`` (or ``fft_workers=-1`` to parallelize the CPU FFT
across all cores). Measure settings do not go here (see
``connectivity_kwargs``).
Returns
-------
result : xarray.DataArray or xarray.Dataset
A plain single-quantity method returns a DataArray. Component-resolved
and multi-quantity methods return a Dataset even when requested alone;
multiple methods are merged into one Dataset without flattening their
semantic dimensions.
Examples
--------
>>> import numpy as np
>>> rng = np.random.default_rng(0)
>>> # Generate coupled oscillator data
>>> t = np.arange(0, 1, 1/500) # 500 Hz, 1 second
>>> sig1 = np.sin(2*np.pi*10*t) + 0.1*rng.standard_normal(len(t))
>>> sig2 = np.sin(2*np.pi*10*t + np.pi/4) + 0.1*rng.standard_normal(len(t))
>>> # Shape (n_time, n_channels); a single trial of 2 signals. The 2-D form
>>> # is promoted to a single-trial 3-D array internally.
>>> data = np.stack([sig1, sig2], axis=-1) # (500, 2)
>>>
>>> # Compute coherence
>>> coherence = multitaper_connectivity(
... data, sampling_frequency=500,
... method="coherence_magnitude",
... signal_names=["Signal_1", "Signal_2"]
... )
>>> coherence.dims
('time', 'frequency', 'source', 'target')
>>> # Compute multiple measures
>>> measures = multitaper_connectivity(
... data, sampling_frequency=500,
... method=["coherence_magnitude", "imaginary_coherence"]
... )
>>> list(measures.data_vars)
['coherence_magnitude', 'imaginary_coherence']
>>> # An xarray.DataArray labels axes by dimension name and can supply the
>>> # sampling rate and channel labels itself (no sampling_frequency needed).
>>> import xarray as xr
>>> da = xr.DataArray(
... data,
... dims=("time", "channel"),
... coords={"time": t, "channel": ["Signal_1", "Signal_2"]},
... )
>>> coherence = multitaper_connectivity(da, method="coherence_magnitude")
>>> coherence.coords["source"].values.tolist()
['Signal_1', 'Signal_2']
Notes
-----
Uses multitaper spectral estimation for robust power spectral density
estimation before computing connectivity measures. This provides better
spectral estimates than single-taper methods, especially for short time series.
For directed measures (e.g. ``pairwise_spectral_granger_prediction``) the
``source`` and ``target`` axes are oriented so that
``result.sel(source=a, target=b)`` is the influence *from* ``a`` *to* ``b``,
the same order as the underlying ``Connectivity`` arrays, where
``result[..., i, j]`` is ``i -> j``. Signed undirected phase measures
(``coherence_phase``, ``imaginary_coherency``, ``phase_lag_index``,
``weighted_phase_lag_index``) are positive at ``sel(source=a, target=b)``
when ``a`` leads ``b``.
Every variable has ``long_name`` and ``units`` attrs (``"1"`` for
dimensionless scores, ``"rad"`` for phase, ``"s"`` for delay; spectral
densities are ``"(<units>)^2/Hz"`` when an input DataArray states its
``units``, and ``"(<units>)^2"`` once integrated over a band). Non-index coordinates on an input DataArray's signal dimension
(e.g. ``region``) are carried as ``source_<name>``/``target_<name>``.
Real-valued results write with any NetCDF engine (booleans are stored as
0/1). Complex results (``coherency``, ``cross_spectral_density``,
``canonical_coherency``, and the global-coherence vectors) need an engine
that stores complex data, e.g.
``result.to_netcdf("result.h5", engine="h5netcdf", invalid_netcdf=True)``,
or netCDF4 >= 1.7 with ``engine="netcdf4", auto_complex=True`` (open with the
same option).
The result records provenance as NetCDF-safe attributes so a saved file is
self-describing:
- ``mt_*`` -- the Multitaper transform parameters.
- ``measure`` and ``measure_kwargs_json`` -- the measure name and a canonical,
JSON-normalized representation of its keyword arguments.
- ``arg_<key>`` / ``arg_<key>_json`` -- each measure keyword argument
individually for quick inspection; a scalar is stored as-is under
``arg_<key>``, while a structured or non-finite value is stored as a JSON
string under ``arg_<key>_json`` (parse with ``json.loads``;
``measure_kwargs_json`` is the canonical record).
- ``package``, ``package_version``, ``backend``, ``expectation_type`` --
software provenance.
- ``input_attrs_json`` -- a canonical, JSON-normalized record of attributes
carried over from an input ``xarray.DataArray`` (e.g. subject or session
metadata). Keeping the complete mapping in one record preserves arbitrary
keys without collisions or invalid NetCDF attribute names.
JSON records are canonical for scalar, numpy, mapping, and sequence values.
A value outside those kinds is recorded best-effort via its ``repr``, which
may embed a memory address and is therefore not guaranteed reproducible
across runs.
References
----------
.. [1] Thomson, D. J. (1982). Spectrum estimation and harmonic analysis.
Proceedings of the IEEE, 70(9), 1055-1096.
.. [2] Percival, D. B., & Walden, A. T. (1993). Spectral Analysis for Physical
Applications: Multitaper and Conventional Univariate Techniques.
"""
explicit_start_time = kwargs.get("start_time", _UNSET)
(
time_series_data,
signal_names,
inferred_sampling_frequency,
inferred_start_time,
input_attrs,
signal_metadata,
) = _unwrap_xarray_input(
time_series,
signal_names,
sampling_frequency,
time_dim=time_dim,
trial_dim=trial_dim,
signal_dim=signal_dim,
explicit_start_time=explicit_start_time,
)
if inferred_sampling_frequency is not None:
sampling_frequency = inferred_sampling_frequency
if sampling_frequency is None:
msg = (
"sampling_frequency is required unless the input is an "
"xarray.DataArray with a numeric 'time' coordinate (in elapsed "
"seconds) to infer it from."
)
raise ValueError(msg)
if inferred_start_time is not None and explicit_start_time is _UNSET:
kwargs["start_time"] = inferred_start_time
# The default set (DEFAULT_METHODS) is portably serializable; complex,
# component/group, frequency-reduced, and directed-transfer-function results
# remain opt-in by name.
method, return_dataarray = _requested_methods(method, DEFAULT_METHODS)
# Accept the documented (n_times, n_channels) 2-D form by inserting a
# singleton trial axis; Multitaper requires 3-D (n_times, n_trials,
# n_signals).
if getattr(time_series_data, "ndim", None) == 2:
time_series_data = time_series_data[:, np.newaxis, :]
m = Multitaper(
time_series=time_series_data,
sampling_frequency=sampling_frequency,
time_window_duration=time_window_duration,
**kwargs,
)
# Resolve group labels before Connectivity.from_multitaper runs the FFT, so
# a missing or misplaced label fails fast. The constructor has validated the
# input shape, so m.n_signals is the true signal count.
group_labels, connectivity_kwargs, group_methods = _resolve_group_labels(
method, connectivity_kwargs, group_labels, m.n_signals
)
# Capture metadata and build the shared calculation object from the same
# immutable transform. The private formatter below never accepts a separate
# Multitaper, so data and labels cannot be paired accidentally.
metadata = m._provenance_metadata()
shared_connectivity = Connectivity.from_multitaper(m)
# Validate labels and build shared provenance once; both are invariant across
# the requested measures.
signal_labels = _validated_signal_labels(signal_names, shared_connectivity.n_signals)
shared_attrs = _shared_provenance_attrs(
shared_connectivity, metadata, input_attrs=input_attrs
)
result = _format_and_reduce_measures(
shared_connectivity,
method,
return_dataarray=return_dataarray,
signal_labels=signal_labels,
squeeze=squeeze,
shared_attrs=shared_attrs,
connectivity_kwargs=connectivity_kwargs,
group_labels=group_labels,
group_methods=group_methods,
transform_settings_hint="pass it to multitaper_connectivity directly",
frequency_range=frequency_range,
frequency_decimation=frequency_decimation,
frequency_bands=frequency_bands,
frequency_reduction=frequency_reduction,
signal_metadata=signal_metadata,
)
_warn_label_orientation_changed(method)
return result
[docs]
def fourier_connectivity(
fourier_coefficients: NDArray[np.complexfloating] | xr.DataArray,
frequencies: NDArray[np.floating] | None = None,
time: NDArray[np.floating] | None = None,
method: str | list[str] | None = None,
signal_names: Sequence[_SignalLabel] | None = None,
squeeze: bool = False,
connectivity_kwargs: dict[str, Any] | None = None,
is_one_sided: bool | None = None,
*,
group_labels: Sequence[Hashable] | NDArray[Any] | None = None,
frequency_range: tuple[float, float] | None = None,
frequency_decimation: int = 1,
frequency_bands: Mapping[str, tuple[float, float]] | None = None,
frequency_reduction: Literal["mean", "integral"] = "mean",
time_dim: Hashable | None = None,
trial_dim: Hashable | None = None,
taper_dim: Hashable | None = None,
frequency_dim: Hashable | None = None,
signal_dim: Hashable | None = None,
dtype: DTypeLike = np.complex128,
minimum_phase_tolerance: float = 1e-8,
minimum_phase_max_iterations: int = 500,
) -> xr.DataArray | xr.Dataset:
"""Compute labeled connectivity from externally estimated FFT coefficients.
The labeled output has one time axis, so the expectation is always
``"trials_tapers"``; use :class:`Connectivity` directly for expectations
that retain trial/taper axes or average over time.
Parameters
----------
fourier_coefficients : array or xarray.DataArray
Complex coefficients in ``(n_observations, n_frequencies, n_signals)``,
``(n_trials, n_tapers, n_frequencies, n_signals)``, or the core's full
``(n_time, n_trials, n_tapers, n_frequencies, n_signals)`` layout. A
DataArray is transposed by semantic dimension names (or the explicit
``*_dim`` arguments) and its frequency, time, and signal coordinates
and attributes are preserved.
frequencies : array, shape (n_frequencies,), optional
Frequency of each bin in Hz. A two-sided coordinate must be in standard
FFT order; a non-negative, strictly increasing coordinate is treated as
one-sided when ``is_one_sided`` is omitted. Taken from the DataArray
coordinate when not given.
time : array, shape (n_time,), optional
Center time of each window in seconds; defaults to window indices.
method : str or list of str, optional
Measure name(s) from :func:`list_measures`. A single name returns a
DataArray; a list (or ``None`` for :data:`DEFAULT_METHODS`) returns a
Dataset with one variable per measure.
signal_names : sequence, optional
Labels for the ``source``/``target`` coordinates; defaults to the
DataArray signal coordinate or ``"0"``, ``"1"``, ....
squeeze : bool, default=False
Only honored when a single ``method`` (a string) is requested. If there
are exactly 2 signals, reduce a pairwise measure to the single ordered
pair (first source, last target), returning a ``(time, frequency)`` array
that keeps the selected ``source`` and ``target`` as scalar coordinates.
With more than 2 signals a warning is issued and the full matrix is
returned; for ``power`` squeeze is a no-op. Length-one ``time`` is kept.
connectivity_kwargs : dict, optional
Extra keyword arguments for the *measure* (a ``Connectivity`` method),
e.g. ``pairs`` for ``subset_pairwise_spectral_granger_prediction`` or
``n_components`` for ``canonical_coherency``; passed to every requested
measure. Transform settings do not go here: they belong to the transform
that produced ``fourier_coefficients``.
is_one_sided : bool, optional
Declare whether the coefficients cover only non-negative frequencies.
When no frequency coordinate is available the sidedness cannot be
inferred: pass ``True`` for one-sided input (e.g. ``rfft`` output) or
``False`` for a full FFT-order spectrum. ``False`` is honored as a
declaration, so measures that require a two-sided spectrum run on the
coefficients as given; they warn if the coefficients are not
conjugate-symmetric, as the FFT of real-valued signals is, because
mislabeled one-sided input makes those measures wrong. Leaving it unset in that case assumes two-sided,
warns, and refuses those measures because the assumption cannot be
checked. With a frequency coordinate it is inferred from
``frequencies``.
group_labels : sequence, optional
One label per signal naming the group it belongs to; labels are scalars
such as integers or area names (``"CA1"``), and missing values
(``None``, NaN) are rejected. Required by
the group measures (``canonical_coherence``, ``canonical_coherency``,
``maximized_imaginary_coherency``, ``multivariate_interaction_measure``,
``blockwise_spectral_granger_prediction`` and
``maximized_imaginary_coherency_components``); an error is raised if no
requested measure takes it (labels are passed only to the measures that
do).
frequency_range : (float, float), optional
Inclusive ``(low, high)`` bounds in Hz to keep before any decimation
or band reduction.
frequency_decimation : int, default=1
Keep every ``frequency_decimation``-th frequency bin.
frequency_bands : mapping of str to (float, float), optional
Named inclusive bands to reduce the frequency axis into; see
:func:`frequency_band_reduce`.
frequency_reduction : {"mean", "integral"}, default="mean"
Within-band reduction used with ``frequency_bands``.
time_dim, trial_dim, taper_dim, frequency_dim, signal_dim : hashable, optional
DataArray dimension names for each axis role, when they cannot be
inferred from common names.
dtype : numpy.dtype, default=complex128
Working precision for the connectivity computations.
minimum_phase_tolerance : float, default=1e-8
Relative convergence tolerance of the Wilson factorization used by the
directed measures.
minimum_phase_max_iterations : int, default=500
Maximum Wilson iterations for the directed measures.
Returns
-------
xarray.DataArray or xarray.Dataset
Labeled result with ``time``, ``frequency`` (or ``band``),
and measure-specific dimensions such as ``source``/``target``; a
DataArray for a single ``method`` name, otherwise a Dataset. Directed
measures are oriented so ``sel(source=a, target=b)`` is the influence
from ``a`` to ``b``. One-sided coefficients support functional
measures, but measures that need a full two-sided spectrum raise.
"""
(
coefficient_data,
frequencies,
time,
signal_names,
input_attrs,
signal_metadata,
) = _unwrap_fourier_input(
fourier_coefficients,
frequencies=frequencies,
time=time,
signal_names=signal_names,
time_dim=time_dim,
trial_dim=trial_dim,
taper_dim=taper_dim,
frequency_dim=frequency_dim,
signal_dim=signal_dim,
)
if time is not None and not _is_real_numeric_dtype(np.asarray(time).dtype):
msg = (
"time must contain numeric elapsed seconds (window centers); "
f"got dtype {np.asarray(time).dtype!r}. Convert a datetime axis to "
"elapsed seconds, e.g. (t - t[0]) / np.timedelta64(1, 's')."
)
raise TypeError(msg)
inferred_one_sided = False
if is_one_sided is not None:
is_one_sided = _validated_flag("is_one_sided", is_one_sided)
if frequencies is not None:
frequency_values = np.asarray(frequencies, dtype=float)
if frequency_values.ndim != 1:
msg = "frequencies must be a one-dimensional coordinate."
raise ValueError(msg)
inferred_one_sided = bool(
frequency_values.size > 0 and not np.any(frequency_values < 0)
)
one_sided = inferred_one_sided if is_one_sided is None else bool(is_one_sided)
else:
if is_one_sided is None:
warnings.warn(
"fourier_connectivity received no frequency coordinate and no "
"is_one_sided flag; assuming a two-sided spectrum in standard "
"FFT order. For rfft or wavelet coefficients (non-negative "
"frequencies only) pass is_one_sided=True, otherwise "
"is_one_sided=False to silence this warning.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
one_sided = bool(is_one_sided) if is_one_sided is not None else False
connectivity = Connectivity(
coefficient_data,
expectation_type="trials_tapers",
frequencies=frequencies,
time=time,
dtype=dtype,
minimum_phase_tolerance=minimum_phase_tolerance,
minimum_phase_max_iterations=minimum_phase_max_iterations,
is_one_sided=one_sided,
)
# No default measure needs a two-sided spectrum, so the defaults suit any input.
methods, return_dataarray = _requested_methods(method, DEFAULT_METHODS)
group_labels, connectivity_kwargs, group_methods = _resolve_group_labels(
methods, connectivity_kwargs, group_labels, connectivity.n_signals
)
if frequencies is None:
# Without a frequency coordinate two-sidedness cannot be verified, so an
# *assumed* two-sided spectrum (is_one_sided=None) must not let one-sided
# input (e.g. rfft/wavelet coefficients) reach Wilson factorization and
# produce a silently wrong result. Reject methods that declare the
# full-spectrum requirement unless the caller declared the spectrum
# two-sided with is_one_sided=False; other directional measures such as
# dPLI and PSI remain valid on one-sided coefficients.
two_sided_methods = [name for name in methods if _requires_two_sided(name)]
if two_sided_methods and one_sided:
# The caller already declared one-sided input, so no frequency vector
# would enable Wilson factorization -- give the accurate reason.
msg = (
f"Measures {sorted(set(two_sided_methods))} require a full "
"two-sided spectrum in standard FFT order. One-sided transforms "
"(is_one_sided=True) support functional connectivity measures but "
"not Wilson-factorized measures. Request only one-sided-compatible "
"measures, or supply full two-sided coefficients."
)
raise ValueError(msg)
if two_sided_methods and is_one_sided is None:
msg = (
f"Measures {sorted(set(two_sided_methods))} require a full "
"two-sided spectrum in standard FFT order, which cannot be verified "
"without a frequency coordinate. Pass `frequencies` (the FFT "
"frequency vector, including negative bins) so two-sidedness can be "
"checked, pass is_one_sided=False to declare a full FFT-order "
"spectrum, or request only one-sided-compatible measures."
)
raise ValueError(msg)
if two_sided_methods and is_one_sided is not None:
# The declaration is trusted here, so check what it implies for real
# signals: bin k is the conjugate of bin -k (FFT order). Measured
# relative residuals: <= 4e-16 for FFTs of real noise in complex128
# and <= 2e-7 in complex64 (up to 65536 bins), versus 1.3-1.45 for
# rfft output declared two-sided and for complex-valued signals. A
# threshold of 1e-3 leaves over three orders of magnitude on each side.
positive_bins = coefficient_data[..., 1:, :]
mirrored_bins = coefficient_data[..., :0:-1, :].conj()
asymmetry = float((abs(positive_bins - mirrored_bins) ** 2).sum()) ** 0.5
scale = float((abs(positive_bins) ** 2).sum()) ** 0.5
if asymmetry > 1e-3 * scale:
warnings.warn(
"The Fourier coefficients declared two-sided (is_one_sided="
"False) are not conjugate-symmetric along the frequency axis "
f"(relative residual {asymmetry / scale:.2g}), as the FFT of "
"real-valued signals is. This happens when one-sided "
"coefficients (e.g. rfft or wavelet output) are declared "
f"two-sided, and then {sorted(set(two_sided_methods))}, which "
"require a two-sided spectrum, are wrong; or when the signals "
"are complex-valued. Pass is_one_sided=True for one-sided "
"coefficients and request only one-sided-compatible measures, "
"or pass the full two-sided FFT of real-valued signals.",
UserWarning,
stacklevel=stacklevel_outside_package(),
)
signal_labels = _validated_signal_labels(signal_names, connectivity.n_signals)
metadata: dict[str, Any] = {
"source": "external_fourier_coefficients",
"coefficient_shape_json": _canonical_json(tuple(coefficient_data.shape)),
"frequency_coordinate": "provided" if frequencies is not None else "normalized",
"time_coordinate": "provided" if time is not None else "index",
"is_one_sided": one_sided,
"one_sided_inferred": is_one_sided is None and inferred_one_sided,
}
all_frequencies = connectivity.all_frequencies
if not one_sided and all_frequencies.size > 1:
# FFT-order bins are sampling_frequency / n_fft apart (bin 1 is -spacing
# when n_fft == 2), so the rate follows and places the Nyquist bin,
# which only an even n_fft has.
metadata["sampling_frequency"] = float(all_frequencies.size * abs(all_frequencies[1]))
shared_attrs = _shared_provenance_attrs(
connectivity,
metadata,
input_attrs=input_attrs,
transform_prefix="fourier_",
)
return _format_and_reduce_measures(
connectivity,
methods,
return_dataarray=return_dataarray,
signal_labels=signal_labels,
squeeze=squeeze,
shared_attrs=shared_attrs,
connectivity_kwargs=connectivity_kwargs,
group_labels=group_labels,
group_methods=group_methods,
transform_settings_hint=("set it on the transform that produced fourier_coefficients"),
frequency_range=frequency_range,
frequency_decimation=frequency_decimation,
frequency_bands=frequency_bands,
frequency_reduction=frequency_reduction,
signal_metadata=signal_metadata,
)