Source code for spectral_connectivity.wrapper

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