From 75944d4bb8061f2664e2589177e73ea7eb268578 Mon Sep 17 00:00:00 2001 From: Aditya Singh Date: Wed, 5 Aug 2026 05:48:37 -0700 Subject: [PATCH] Fix plot_conditions plotting one epoch instead of the condition average plot_conditions built a wide epochs-by-time DataFrame and passed the channel number as seaborn's `y`. Seaborn read that as a column key, so it selected column number `ch` of the frame. Those columns are epochs, not channels, so each subplot showed a single arbitrary epoch with no averaging and no confidence interval. The legend had a related problem. It was built from a bare list of labels, which matplotlib pairs with whatever artists it finds in draw order. The labels were written difference-first while the artists are drawn conditions-first, and seaborn's confidence bands are picked up as handles too, so labels ended up on the wrong lines. Reshape the data to long form so seaborn averages over epochs, and build the legend from explicit Line2D handles so every label carries its own colour. --- eegnb/analysis/analysis_utils.py | 44 +++++++-- tests/test_analysis_plots.py | 151 +++++++++++++++++++++++++++++++ 2 files changed, 187 insertions(+), 8 deletions(-) create mode 100644 tests/test_analysis_plots.py diff --git a/eegnb/analysis/analysis_utils.py b/eegnb/analysis/analysis_utils.py index c48e021a2..e1af456f9 100644 --- a/eegnb/analysis/analysis_utils.py +++ b/eegnb/analysis/analysis_utils.py @@ -18,6 +18,7 @@ from mne.channels import make_standard_montage from mne.filter import create_filter from matplotlib import pyplot as plt +from matplotlib import lines as mlines from scipy import stats from scipy.signal import lfilter, lfilter_zi @@ -277,10 +278,21 @@ def plot_conditions( for ch in range(channel_count): for cond, color in zip(conditions.values(), palette): + # Hand seaborn long-form data: one row per (epoch, time) sample, so + # that it averages over the epochs of this condition and bootstraps + # a confidence interval around that average. A wide frame with the + # channel number as `y` would instead select a single column, i.e. + # plot one arbitrary epoch and no interval at all. + epoch_by_time = pd.DataFrame( + X[y.isin(cond), ch].T, index=pd.Index(times, name="time") + ) + samples = epoch_by_time.melt( + ignore_index=False, var_name="epoch", value_name="amplitude" + ).reset_index() sns.lineplot( - data=pd.DataFrame(X[y.isin(cond), ch].T, index=times), - x=times, - y=ch, + data=samples, + x="time", + y="amplitude", color=color, n_boot=n_boot, ax=axes[ch], @@ -300,14 +312,30 @@ def plot_conditions( x=0, ymin=ylim[0], ymax=ylim[1], color="k", lw=1, label="_nolegend_" ) + # Build the legend from explicit handles rather than from a bare list of + # labels. Passing labels alone makes matplotlib pair them with whatever + # artists it finds on the axis, in draw order, which does not match the + # order the labels are written in and also picks up the confidence-interval + # bands seaborn draws. Pairing each label with its own handle keeps the + # legend correct no matter how many artists the plotting calls add. + legs = [] + for cond_name, color in zip(conditions.keys(), palette): + legs.append( + mlines.Line2D([], [], color=color, marker="", ls="-", label=cond_name) + ) if diff_waveform: - legend = ["{} - {}".format(diff_waveform[1], diff_waveform[0])] + list( - conditions.keys() + legs.append( + mlines.Line2D( + [], + [], + color="k", + marker="", + ls="-", + label="{} - {}".format(diff_waveform[1], diff_waveform[0]), + ) ) - else: - legend = conditions.keys() axes[-1].legend( - legend, bbox_to_anchor=(1.05, 1), loc="upper left", borderaxespad=0.0 + handles=legs, bbox_to_anchor=(1.05, 1), loc="upper left", borderaxespad=0.0 ) sns.despine() plt.tight_layout() diff --git a/tests/test_analysis_plots.py b/tests/test_analysis_plots.py new file mode 100644 index 000000000..e7e846445 --- /dev/null +++ b/tests/test_analysis_plots.py @@ -0,0 +1,151 @@ +""" +Tests for the plotting helpers in eegnb.analysis. + +These run on synthetic MNE epochs, so no EEG hardware and no downloaded +dataset is needed. +""" + +from collections import OrderedDict + +import matplotlib + +matplotlib.use("Agg") + +import matplotlib.lines as mlines +import mne +import numpy as np +import pytest + +from eegnb.analysis.analysis_utils import plot_conditions + + +CH_NAMES = ["TP9", "AF7", "AF8", "TP10"] + + +def _make_epochs(data, codes, sfreq=256.0): + """Wrap raw arrays into MNE epochs with event codes 1 and 2.""" + info = mne.create_info(CH_NAMES, sfreq, ch_types="eeg") + n_epochs, _, n_times = data.shape + events = np.column_stack( + [np.arange(n_epochs) * n_times, np.zeros(n_epochs, int), codes] + ) + return mne.EpochsArray( + data, + info, + events=events, + event_id={"Non-Target": 1, "Target": 2}, + tmin=-0.1, + verbose="error", + ) + + +def _noise_epochs(n_epochs=20, n_times=64): + rng = np.random.RandomState(0) + data = rng.randn(n_epochs, len(CH_NAMES), n_times) * 1e-6 + codes = np.array([1, 2] * (n_epochs // 2)) + data[codes == 2] += 5e-6 + return _make_epochs(data, codes) + + +def _legend_entries(ax): + """Return [(label, rgba of the handle)] for the legend on `ax`.""" + legend = ax.get_legend() + assert legend is not None, "expected a legend on the last axis" + + entries = [] + for text, handle in zip(legend.get_texts(), legend.legend_handles): + assert isinstance( + handle, mlines.Line2D + ), "legend handles should be lines, not confidence-interval bands" + entries.append( + (text.get_text(), tuple(matplotlib.colors.to_rgba(handle.get_color()))) + ) + return entries + + +@pytest.mark.parametrize("diff_waveform", [None, (1, 2)]) +def test_plot_conditions_legend_matches_lines(diff_waveform): + """Each legend label must sit next to the colour it actually describes. + + Regression test for #226. The legend used to be built from a bare list of + labels, which matplotlib paired with the artists in draw order. The + resulting legend named the difference waveform with a condition colour and + left the black difference line unlabelled. + """ + epochs = _noise_epochs() + conditions = OrderedDict(NonTarget=[1], Target=[2]) + + import seaborn as sns + + palette = sns.color_palette("hls", len(conditions) + 1) + + _, axes = plot_conditions( + epochs, + conditions=conditions, + diff_waveform=diff_waveform, + channel_count=4, + n_boot=10, + ) + + expected = [ + (name, tuple(matplotlib.colors.to_rgba(color))) + for name, color in zip(conditions.keys(), palette) + ] + if diff_waveform: + expected.append( + ( + "{} - {}".format(diff_waveform[1], diff_waveform[0]), + tuple(matplotlib.colors.to_rgba("k")), + ) + ) + + assert _legend_entries(axes[-1]) == expected + + matplotlib.pyplot.close("all") + + +def test_plot_conditions_plots_the_condition_average(): + """The line drawn per condition must be the average over that condition's + epochs, not one arbitrary epoch. + + Regression test for #226. The epochs-by-time frame was handed to seaborn in + wide form with the channel number as `y`, so seaborn selected column number + `ch` of that frame. Column `ch` is epoch number `ch`, which meant channel 0 + showed epoch 0, channel 1 showed epoch 1, and so on, with no averaging and + no confidence interval. + """ + n_epochs, n_times = 8, 16 + + # Every epoch is a distinct constant, so the average of a condition and any + # single epoch of it are all different numbers and cannot be confused. + data = np.zeros((n_epochs, len(CH_NAMES), n_times)) + for i in range(n_epochs): + data[i, :, :] = (i + 1) * 1e-6 + codes = np.array([1, 2] * (n_epochs // 2)) + + epochs = _make_epochs(data, codes) + conditions = OrderedDict(NonTarget=[1], Target=[2]) + + _, axes = plot_conditions( + epochs, + conditions=conditions, + diff_waveform=None, + channel_count=len(CH_NAMES), + n_boot=10, + ) + + # Values are scaled to microvolts inside plot_conditions. + scaled = data[:, 0, 0] * 1e6 + expected_means = [scaled[codes == code].mean() for code in (1, 2)] + + for ch, ax in enumerate(axes[: len(CH_NAMES)]): + drawn = [line.get_ydata() for line in ax.get_lines() if len(line.get_ydata()) == n_times] + assert len(drawn) == len(conditions), ( + f"channel {ch}: expected one line per condition, got {len(drawn)}" + ) + for ydata, expected in zip(drawn, expected_means): + assert np.allclose(ydata, expected), ( + f"channel {ch}: plotted {ydata[0]} but the condition average is {expected}" + ) + + matplotlib.pyplot.close("all")