From 75944d4bb8061f2664e2589177e73ea7eb268578 Mon Sep 17 00:00:00 2001 From: Aditya Singh Date: Wed, 5 Aug 2026 05:48:37 -0700 Subject: [PATCH 1/3] 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") From d5376c1df857f16615de20b8ea88884157c8f3d8 Mon Sep 17 00:00:00 2001 From: Benjamin Pettit Date: Sat, 3 Oct 2026 14:29:19 +1000 Subject: [PATCH 2/3] fix(analysis): support named markers and grouped difference waveforms Resolve event names and condition groups while preserving numeric marker support. Add compatibility tests and remove change-history commentary from the plotting helper and tests. --- eegnb/analysis/analysis_utils.py | 40 +++++++----- tests/test_analysis_plots.py | 105 ++++++++++++++++++++++++------- 2 files changed, 104 insertions(+), 41 deletions(-) diff --git a/eegnb/analysis/analysis_utils.py b/eegnb/analysis/analysis_utils.py index e1af456f9..ff7a74e15 100644 --- a/eegnb/analysis/analysis_utils.py +++ b/eegnb/analysis/analysis_utils.py @@ -230,7 +230,7 @@ def plot_conditions( Keyword Args: conditions (OrderedDict): dictionary that contains the names of the conditions to plot as keys, and the list of corresponding marker - numbers as value. E.g., + numbers or names from epochs.event_id as value. E.g., conditions = {'Non-target': [0, 1], 'Target': [2, 3, 4]} ci (float): confidence interval in range [0, 100] @@ -238,8 +238,8 @@ def plot_conditions( title (str): title of the figure palette (list): color palette to use for conditions ylim (tuple): (ymin, ymax) - diff_waveform (tuple or None): tuple of ints indicating which - conditions to subtract for producing the difference waveform. + diff_waveform (tuple or None): pair of marker numbers, event names or + condition keys. The second waveform minus the first is plotted. If None, do not plot a difference waveform channel_count (int): number of channels to plot. Default set to 4 for backward compatibility with Muse implementations @@ -256,6 +256,24 @@ def plot_conditions( if isinstance(conditions, dict): conditions = OrderedDict(conditions) + def resolve_marker(marker): + if isinstance(marker, str): + if marker not in epochs.event_id: + raise ValueError(f"Unknown event name: {marker!r}") + return epochs.event_id[marker] + return marker + + conditions = OrderedDict( + (name, [resolve_marker(marker) for marker in markers]) + for name, markers in conditions.items() + ) + if diff_waveform: + diff_codes = [ + conditions[marker] if isinstance(marker, str) and marker in conditions + else [resolve_marker(marker)] + for marker in diff_waveform + ] + if palette is None: palette = sns.color_palette("hls", len(conditions) + 1) @@ -269,7 +287,6 @@ def plot_conditions( midaxis = math.ceil(channel_count / 2) fig, axes = plt.subplots(2, midaxis, figsize=[12, 6], sharex=True, sharey=False) - # get individual plot axis plot_axes = [] for axis_y in range(midaxis): for axis_x in range(2): @@ -278,11 +295,6 @@ 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") ) @@ -301,8 +313,8 @@ def plot_conditions( axes[ch].set(xlabel='Time (s)', ylabel='Amplitude (uV)', title=epochs.ch_names[channel_order[ch]]) if diff_waveform: - diff = np.nanmean(X[y == diff_waveform[1], ch], axis=0) - np.nanmean( - X[y == diff_waveform[0], ch], axis=0 + diff = np.nanmean(X[y.isin(diff_codes[1]), ch], axis=0) - np.nanmean( + X[y.isin(diff_codes[0]), ch], axis=0 ) axes[ch].plot(times, diff, color="k", lw=1) @@ -312,12 +324,6 @@ 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( diff --git a/tests/test_analysis_plots.py b/tests/test_analysis_plots.py index e7e846445..e1c703698 100644 --- a/tests/test_analysis_plots.py +++ b/tests/test_analysis_plots.py @@ -22,7 +22,7 @@ CH_NAMES = ["TP9", "AF7", "AF8", "TP10"] -def _make_epochs(data, codes, sfreq=256.0): +def _make_epochs(data, codes, sfreq=256.0, event_id=None): """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 @@ -33,7 +33,7 @@ def _make_epochs(data, codes, sfreq=256.0): data, info, events=events, - event_id={"Non-Target": 1, "Target": 2}, + event_id=event_id if event_id is not None else {"Non-Target": 1, "Target": 2}, tmin=-0.1, verbose="error", ) @@ -63,15 +63,9 @@ def _legend_entries(ax): return entries -@pytest.mark.parametrize("diff_waveform", [None, (1, 2)]) +@pytest.mark.parametrize("diff_waveform", [None, (1, 2), ("Non-Target", "Target")]) 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. - """ + """Each legend label matches its line colour.""" epochs = _noise_epochs() conditions = OrderedDict(NonTarget=[1], Target=[2]) @@ -104,27 +98,18 @@ def test_plot_conditions_legend_matches_lines(diff_waveform): 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. - """ +@pytest.mark.parametrize("markers", [(1, 2), ("Non-Target", "Target"), (1, "Target")]) +def test_plot_conditions_plots_the_condition_average(markers): + """Each condition's plotted waveform matches its epoch average.""" 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]) + conditions = OrderedDict(NonTarget=[markers[0]], Target=[markers[1]]) _, axes = plot_conditions( epochs, @@ -134,7 +119,6 @@ def test_plot_conditions_plots_the_condition_average(): 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)] @@ -149,3 +133,76 @@ def test_plot_conditions_plots_the_condition_average(): ) matplotlib.pyplot.close("all") + + +def test_plot_conditions_names_match_numeric_waveforms(): + epochs = _noise_epochs() + named_conditions = OrderedDict(NonTarget=["Non-Target"], Target=["Target"]) + numeric_conditions = OrderedDict(NonTarget=[1], Target=[2]) + numeric_fig, numeric_axes = plot_conditions( + epochs, conditions=numeric_conditions, diff_waveform=(1, 2), n_boot=10 + ) + named_fig, named_axes = plot_conditions( + epochs, + conditions=named_conditions, + diff_waveform=("Non-Target", "Target"), + n_boot=10, + ) + try: + expected_difference = epochs.get_data()[epochs.events[:, -1] == 2].mean(axis=0) + expected_difference -= epochs.get_data()[epochs.events[:, -1] == 1].mean(axis=0) + for ch, (numeric_ax, named_ax) in enumerate(zip(numeric_axes, named_axes)): + numeric_lines = [line for line in numeric_ax.lines if len(line.get_xdata()) == len(epochs.times)] + named_lines = [line for line in named_ax.lines if len(line.get_xdata()) == len(epochs.times)] + assert len(numeric_lines) == len(named_lines) == 3 + for numeric_line, named_line in zip(numeric_lines, named_lines): + np.testing.assert_allclose(named_line.get_xdata(), epochs.times) + np.testing.assert_allclose(named_line.get_ydata(), numeric_line.get_ydata()) + np.testing.assert_allclose(named_lines[2].get_ydata(), expected_difference[ch] * 1e6) + assert named_conditions == OrderedDict(NonTarget=["Non-Target"], Target=["Target"]) + finally: + matplotlib.pyplot.close(numeric_fig) + matplotlib.pyplot.close(named_fig) + + +def test_plot_conditions_grouped_cueing_difference(): + rng = np.random.RandomState(1) + codes = np.array([11, 12, 21, 22, 21, 11, 22, 12]) + data = rng.randn(8, len(CH_NAMES), 16) * 1e-6 + data += codes[:, None, None] * 1e-6 + epochs = _make_epochs(data, codes, event_id={ + "InvalidTarget_Left": 11, "InvalidTarget_Right": 12, + "ValidTarget_Left": 21, "ValidTarget_Right": 22, + }) + conditions = OrderedDict( + ValidTarget=["ValidTarget_Left", "ValidTarget_Right"], + InvalidTarget=["InvalidTarget_Left", "InvalidTarget_Right"], + ) + fig, axes = plot_conditions( + epochs, conditions=conditions, + diff_waveform=("ValidTarget", "InvalidTarget"), n_boot=10, + ) + try: + valid = data[np.isin(codes, [21, 22])].mean(axis=0) * 1e6 + invalid = data[np.isin(codes, [11, 12])].mean(axis=0) * 1e6 + for ch, ax in enumerate(axes): + drawn = [line.get_ydata() for line in ax.lines if len(line.get_xdata()) == len(epochs.times)] + assert len(drawn) == 3 + np.testing.assert_allclose(drawn[0], valid[ch]) + np.testing.assert_allclose(drawn[1], invalid[ch]) + np.testing.assert_allclose(drawn[2], invalid[ch] - valid[ch]) + assert _legend_entries(axes[-1])[-1][0] == "InvalidTarget - ValidTarget" + finally: + matplotlib.pyplot.close(fig) + + +@pytest.mark.parametrize( + "conditions, diff_waveform", + [ + (OrderedDict(Unknown=["missing"]), None), + (OrderedDict(NonTarget=[1], Target=[2]), ("Non-Target", "missing")), + ], +) +def test_plot_conditions_unknown_name_raises(conditions, diff_waveform): + with pytest.raises(ValueError, match="Unknown event name: 'missing'"): + plot_conditions(_noise_epochs(), conditions=conditions, diff_waveform=diff_waveform) From b089497306cb24534bd886c8a515801ffdcaf6e2 Mon Sep 17 00:00:00 2001 From: Benjamin Pettit Date: Sat, 3 Oct 2026 14:46:10 +1000 Subject: [PATCH 3/3] fix(examples): display P300 waveforms on a microvolt scale Remove oversized axis limits and use the plotting helper's default microvolt range so the P300 averages and difference waveform are visible. --- examples/visual_p300/01r__p300_viz.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/examples/visual_p300/01r__p300_viz.py b/examples/visual_p300/01r__p300_viz.py index 138f20d64..8780cf58a 100644 --- a/examples/visual_p300/01r__p300_viz.py +++ b/examples/visual_p300/01r__p300_viz.py @@ -99,12 +99,7 @@ fig, ax = plot_conditions(epochs, conditions=conditions, ci=97.5, n_boot=1000, title='', - channel_order=[1,0,2,3],ylim=[-2E6,2.5E6], + channel_order=[1,0,2,3], diff_waveform = diffwav) -# Manually adjust the ylims -for i in [0,2]: ax[i].set_ylim([-0.5e6,0.5e6]) -for i in [1,3]: ax[i].set_ylim([-1.5e6,2.5e6]) - plt.tight_layout() -