Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
605ddc1
Fix psd_array_welch for good spans shorter than n_overlap (#13039)
CedricConday Jun 30, 2026
ce1f435
rename changelog fragment to PR number
CedricConday Jun 30, 2026
7a7a4e8
Drop short good-data spans in Welch PSD with a warning (#13039)
CedricConday Jun 30, 2026
150055a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jun 30, 2026
c82de77
Keep Welch PSD spans shorter than n_per_seg via a per-span window (#1…
CedricConday Jul 10, 2026
88dedc1
Fix codespell: Pre-empt -> Preempt in psd.py comment
CedricConday Jul 10, 2026
f6a5a7e
Raise a clear error for a fixed-length window array on a short span
CedricConday Jul 11, 2026
ffa1717
Zero-pad short good-data spans instead of shrinking the window
CedricConday Jul 16, 2026
15f26d1
Revert "Zero-pad short good-data spans instead of shrinking the window"
CedricConday Aug 18, 2026
3d3f154
test(psd): assert recovered band power for short good spans (#13039)
CedricConday Aug 18, 2026
45ec6b6
Merge remote-tracking branch 'upstream/main' into fix/welch-short-spa…
CedricConday Aug 26, 2026
34e8bf0
Merge upstream/main into fix/welch-short-span-overlap
CedricConday Aug 26, 2026
5f5599b
Merge branch 'main' into fix/welch-short-span-overlap
CarinaFo Sep 20, 2026
63d3549
[autofix.ci] apply automated fixes
autofix-ci[bot] Sep 20, 2026
25f69fb
Address review: name SciPy's padding, add Raw-level short-span test
CedricConday Sep 20, 2026
985fe8b
FIX: import verbose_static in time_frequency/psd.py
CedricConday Sep 20, 2026
c8b34fb
Slim the short-span handling and merge its tests
CedricConday Sep 26, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions doc/changes/dev/14003.bugfix.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Fix :func:`mne.time_frequency.psd_array_welch` (and Welch-method ``compute_psd``) so that good data spans shorter than ``n_per_seg`` no longer raise ``noverlap must be less than nperseg``; such spans are now analyzed with a window shrunk to the span length (with a warning that their spectral resolution is reduced) instead of being lost, by `Cedric Conday`_.
54 changes: 29 additions & 25 deletions mne/time_frequency/psd.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,12 @@
# License: BSD-3-Clause
# Copyright the MNE-Python contributors.

import warnings
from functools import partial

import numpy as np

from ..parallel import parallel_func
from ..utils import _check_option, _ensure_int, logger, verbose_static, warn
from ..utils import _check_option, _ensure_int, _pl, logger, verbose_static, warn
from ..utils.numerics import _mask_to_onsets_offsets


Expand Down Expand Up @@ -276,37 +275,42 @@ def psd_array_welch(
good_mask = ~nan_mask_full
t_onsets, t_offsets = _mask_to_onsets_offsets(good_mask[0])
x_splits = [x[..., t_ons:t_off] for t_ons, t_off in zip(t_onsets, t_offsets)]
# weights reflect the number of samples used from each span. For spans longer
# than `n_per_seg`, trailing samples may be discarded. For spans shorter than
# `n_per_seg`, the wrapped function (`scipy.signal.spectrogram`) automatically
# reduces `n_per_seg` to match the span length (with a warning).
# Weights reflect the number of samples used from each span (trailing
# samples that do not fill a whole window are discarded).
step = n_per_seg - n_overlap
span_lengths = [span.shape[-1] for span in x_splits]
weights = [
w if w < n_per_seg else w - ((w - n_overlap) % step) for w in span_lengths
]
# A span shorter than n_per_seg is analyzed with a window shrunk to its
# length and n_overlap clamped below it (SciPy shrinks nperseg but not
# noverlap, then raises; gh-13039). n_fft is unchanged, and SciPy zero-pads
# each segment to n_fft, so every span lands on the same frequency grid;
# short spans just get coarser resolution. A fixed-length window array
# cannot be shrunk, so that combination raises a clear error instead.
n_short = sum(w < n_per_seg for w in span_lengths)
if n_short and not isinstance(window, (str, tuple)):
raise ValueError(
f"{n_short} good data span{_pl(n_short)} shorter than n_per_seg "
f"({n_per_seg}) cannot be analyzed with a fixed-length window array; "
"pass a window name or tuple, or reduce n_per_seg."
)
if n_short:
warn(
f"{n_short} good data span{_pl(n_short)} shorter than n_per_seg "
f"({n_per_seg}) analyzed with a reduced window (lower spectral "
"resolution); reduce n_per_seg to silence this warning."
)
funcs = [
partial(_func, nperseg=min(w, n_per_seg), noverlap=min(n_overlap, w - 1))
for w in span_lengths
]
agg_func = partial(np.average, weights=weights)
if n_jobs > 1:
logger.info(
f"Data split into {len(x_splits)} (probably unequal) chunks due to "
'"bad_*" annotations. Parallelization may be sub-optimal.'
)
if (np.array(span_lengths) < n_per_seg).any():
logger.info(
"At least one good data span is shorter than n_per_seg, and will be "
"analyzed with a shorter window than the rest of the file."
)

def func(*args, **kwargs):
# swallow SciPy warnings caused by short good data spans
with warnings.catch_warnings():
warnings.filterwarnings(
action="ignore",
module="scipy",
category=UserWarning,
message=r"nperseg = \d+ is greater than input length",
)
return _func(*args, **kwargs)

else:
# Either no NaNs, or NaNs are not aligned across channels.
Expand All @@ -317,10 +321,10 @@ def func(*args, **kwargs):
)
x_splits = [arr for arr in np.array_split(x, n_jobs) if arr.size != 0]
agg_func = np.concatenate
func = _func
funcs = [_func] * len(x_splits)
f_spect = parallel(
my_spect_func(d, func=func, freq_sl=freq_sl, average=average, output=output)
for d in x_splits
my_spect_func(d, func=fn, freq_sl=freq_sl, average=average, output=output)
for d, fn in zip(x_splits, funcs)
)
psds = agg_func(f_spect, axis=0)
shape = dshape + (len(freqs),)
Expand Down
49 changes: 49 additions & 0 deletions mne/time_frequency/tests/test_psd.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
from numpy.testing import assert_allclose, assert_array_almost_equal, assert_array_equal
from scipy.signal import welch

from mne import Annotations, create_info
from mne.io import RawArray
from mne.time_frequency import psd_array_multitaper, psd_array_welch
from mne.time_frequency.multitaper import _psd_from_mt
from mne.time_frequency.psd import _median_biases
Expand Down Expand Up @@ -58,6 +60,53 @@ def test_bad_annot_handling():
np.testing.assert_allclose(got[0], want[0], rtol=1e-15, atol=0)


def test_psd_welch_short_spans():
"""Good spans shorter than n_per_seg are kept with a shrunk window (gh-13039)."""
sfreq, n_fft, n_overlap = 100.0, 256, 128
rng = np.random.default_rng(0)
times = np.arange(60 * int(sfreq)) / sfreq
x = np.sin(2 * np.pi * 10 * times) + 0.1 * rng.standard_normal((2, times.size))
kwargs = dict(sfreq=sfreq, n_fft=n_fft, n_overlap=n_overlap)
psds, freqs = psd_array_welch(x, **kwargs)
band = (freqs >= 8) & (freqs <= 12)
# mark every 100th sample bad, so every good span (99 samples) is < n_per_seg
x_frag = x.copy()
x_frag[:, 99::100] = np.nan
with pytest.warns(RuntimeWarning, match="shorter than n_per_seg"):
psds_frag, freqs_frag = psd_array_welch(x_frag, **kwargs)
assert_array_equal(freqs_frag, freqs)
# alpha-band power must match the uninterrupted signal (zero-padding would not)
assert_allclose(psds_frag[:, band].sum(-1), psds[:, band].sum(-1), rtol=0.1)
# a fixed-length window array cannot be shrunk to fit a short span
with pytest.raises(ValueError, match="fixed-length window"):
psd_array_welch(x_frag, window=np.hamming(n_fft), **kwargs)


def test_compute_psd_welch_short_span_annotations():
"""Test n_per_seg shorter than n_overlap warning and band power change."""
sfreq = 100.0
n_times = int(60 * sfreq)
rng = np.random.default_rng(42)
times = np.arange(n_times) / sfreq
data = np.sin(2 * np.pi * 10 * times) + 0.1 * rng.standard_normal((2, n_times))
raw = RawArray(data, create_info(2, sfreq, "eeg"))
kwargs = dict(method="welch", n_fft=256, n_overlap=128)
ref = raw.compute_psd(**kwargs)

# leave a 1 s good span (100 samples < n_overlap) between two bad segments
raw.set_annotations(Annotations([10.0, 21.0], [10.0, 10.0], "bad_segment"))
# test warning
with pytest.warns(RuntimeWarning, match="shorter than n_per_seg"):
spec = raw.compute_psd(**kwargs)
# alpha-band power should match the uninterrupted recording closely
band = (spec.freqs >= 8) & (spec.freqs <= 12)
assert_allclose(
spec.get_data()[:, band].sum(axis=-1),
ref.get_data()[:, band].sum(axis=-1),
rtol=0.1,
)


Comment thread
CarinaFo marked this conversation as resolved.
def _make_psd_data():
"""Make noise data with sinusoids in 2 out of 7 channels."""
rng = np.random.default_rng(0)
Expand Down
Loading