|
13 | 13 | from numpy.testing import assert_allclose, assert_array_equal |
14 | 14 |
|
15 | 15 | from mne_lsl.datasets import testing |
16 | | -from mne_lsl.lsl import StreamInfo, StreamOutlet |
| 16 | +from mne_lsl.lsl import StreamInfo, StreamOutlet, local_clock |
17 | 17 | from mne_lsl.stream import EpochsStream, StreamLSL |
18 | 18 | from mne_lsl.stream.epochs import ( |
19 | 19 | _check_baseline, |
|
34 | 34 | from numpy.typing import NDArray |
35 | 35 |
|
36 | 36 |
|
| 37 | +def _wait_for_epochs( |
| 38 | + epochs: EpochsStream, |
| 39 | + n: int, |
| 40 | + *, |
| 41 | + timeout: float = 20, |
| 42 | + settle: float = 1.5, |
| 43 | + interval: float = 0.1, |
| 44 | +) -> None: |
| 45 | + """Wait until 'n' epochs were acquired, then check that the count is stable. |
| 46 | +
|
| 47 | + Unlike ``while epochs.n_new_epochs != n``, which exits on the first observation |
| 48 | + matching 'n', this reports an over-count instead of hiding it. |
| 49 | + """ |
| 50 | + manual = epochs._acquisition_delay is None |
| 51 | + start = time.monotonic() |
| 52 | + while epochs.n_new_epochs < n: |
| 53 | + if timeout < time.monotonic() - start: |
| 54 | + raise AssertionError( |
| 55 | + f"Only {epochs.n_new_epochs} epoch(s) were acquired within {timeout} " |
| 56 | + f"seconds, expected {n}. Events in the buffer: {epochs.events}." |
| 57 | + ) |
| 58 | + if manual: |
| 59 | + epochs.acquire() |
| 60 | + time.sleep(interval) |
| 61 | + start = time.monotonic() |
| 62 | + while time.monotonic() - start < settle: |
| 63 | + if manual: |
| 64 | + epochs.acquire() |
| 65 | + time.sleep(interval) |
| 66 | + assert epochs.n_new_epochs == n, ( |
| 67 | + f"{epochs.n_new_epochs} epoch(s) were acquired, expected {n}. Events in the " |
| 68 | + f"buffer: {epochs.events}." |
| 69 | + ) |
| 70 | + |
| 71 | + |
37 | 72 | def test_ensure_event_id() -> None: |
38 | 73 | """Test validation of event dictionary.""" |
39 | 74 | assert _ensure_event_id(5, None) == {"5": 5} |
@@ -775,7 +810,9 @@ def test_epochs_with_irregular_numerical_event_stream( |
775 | 810 | n = epochs.n_new_epochs |
776 | 811 | data = epochs.get_data() |
777 | 812 | assert_allclose(data[:-n, :, :], np.zeros((10 - n, data.shape[1], data.shape[2]))) |
778 | | - data_channels = data[-n:, 1:-1, 2:-2] # give 2 sample of jitter |
| 813 | + # the epoch spans exactly the 100 samples set to 101, so trim 10 samples per side |
| 814 | + # for the timestamp jitter. A misaligned epoch still leaves a long run of 0. |
| 815 | + data_channels = data[-n:, 1:-1, 10:-10] |
779 | 816 | assert_allclose(data_channels, np.ones(data_channels.shape) * 101) |
780 | 817 | epochs.disconnect() |
781 | 818 | stream.disconnect() |
@@ -1076,15 +1113,9 @@ def test_epochs_single_event( |
1076 | 1113 | time.sleep(0.5) |
1077 | 1114 | epochs.acquire() |
1078 | 1115 | assert epochs.n_new_epochs == 0 |
1079 | | - # push a single event, then loop until the epoch is acquired or timeout. The epoch |
1080 | | - # needs tmax=1.5 seconds of data after the event before it can be completed, so the |
1081 | | - # loop must wait at least that long. |
| 1116 | + # push a single event; the epoch needs tmax=1.5 seconds of data after it to complete |
1082 | 1117 | outlet_marker.push_sample(np.array([1], dtype=sinfo.dtype)) |
1083 | | - start = time.monotonic() |
1084 | | - while epochs.n_new_epochs == 0 and time.monotonic() - start < 5: |
1085 | | - epochs.acquire() |
1086 | | - time.sleep(0.2) |
1087 | | - assert epochs.n_new_epochs == 1 |
| 1118 | + _wait_for_epochs(epochs, 1) |
1088 | 1119 | assert event_stream.n_new_samples == 1 |
1089 | 1120 | epochs.disconnect() |
1090 | 1121 | event_stream.disconnect() |
@@ -1122,10 +1153,17 @@ def test_epochs_with_more_events_than_buffer_size( |
1122 | 1153 | for _ in range(5): |
1123 | 1154 | outlet_marker.push_sample(np.array([1], dtype=sinfo.dtype)) |
1124 | 1155 | time.sleep(0.1) |
1125 | | - # wait for the event stream and data stream to buffer all events, so that a |
1126 | | - # single acquire() call sees all 5 events at once and triggers the warning. |
1127 | | - # The event_stream bufsize=10 ensures all 5 events are retained in its ringbuffer. |
1128 | | - time.sleep(1.0) |
| 1156 | + # to trigger the warning, a single acquire() must see all 5 events at once, i.e. the |
| 1157 | + # event stream pulled all 5 markers and the data stream buffered past the last one. |
| 1158 | + # Wait on both rather than on a fixed sleep, which a loaded runner can outlast. |
| 1159 | + ts_last_event = local_clock() |
| 1160 | + start = time.monotonic() |
| 1161 | + while ( |
| 1162 | + event_stream.n_new_samples < 5 or stream._timestamps[-1] <= ts_last_event |
| 1163 | + ) and time.monotonic() - start < 20: |
| 1164 | + time.sleep(0.1) |
| 1165 | + assert event_stream.n_new_samples == 5 |
| 1166 | + assert ts_last_event < stream._timestamps[-1] |
1129 | 1167 | with pytest.warns(RuntimeWarning, match="number of new epochs to add.*is greater"): |
1130 | 1168 | epochs.acquire() |
1131 | 1169 | epochs.disconnect() |
@@ -1170,11 +1208,8 @@ def test_epochs_with_irregular_numerical_event_stream_and_event_id( |
1170 | 1208 | time.sleep(0.1) |
1171 | 1209 | outlet_marker.push_sample(np.array([1], dtype=sinfo.dtype)) |
1172 | 1210 | time.sleep(0.1) |
1173 | | - start = time.monotonic() |
1174 | | - while epochs.n_new_epochs != 3 and time.monotonic() - start < 3: |
1175 | | - epochs.acquire() |
1176 | | - time.sleep(0.5) |
1177 | 1211 | # check that we got 3 epochs on the code '1' |
| 1212 | + _wait_for_epochs(epochs, 3) |
1178 | 1213 | events = epochs.events |
1179 | 1214 | assert_array_equal(events[events != 0], [1, 1, 1]) |
1180 | 1215 | epochs.disconnect() |
@@ -1208,7 +1243,9 @@ def test_epochs_with_irregular_numerical_event_stream_with_2_ch_and_event_id( |
1208 | 1243 | 10, name=mock_lsl_stream.name, source_id=mock_lsl_stream.source_id |
1209 | 1244 | ).connect(acquisition_delay=0.1) |
1210 | 1245 | sinfo = outlet_marker_3_channel.get_sinfo() |
1211 | | - event_stream = StreamLSL(5, name=sinfo.name, source_id=sinfo.source_id).connect( |
| 1246 | + # the buffer must hold all 6 events pushed below, else the events seen by acquire() |
| 1247 | + # depend on how fast the acquisition thread drains the inlet. |
| 1248 | + event_stream = StreamLSL(10, name=sinfo.name, source_id=sinfo.source_id).connect( |
1212 | 1249 | acquisition_delay=0.1 |
1213 | 1250 | ) |
1214 | 1251 | epochs = EpochsStream( |
@@ -1245,12 +1282,8 @@ def test_epochs_with_irregular_numerical_event_stream_with_2_ch_and_event_id( |
1245 | 1282 | time.sleep(0.1) |
1246 | 1283 | outlet_marker_3_channel.push_sample(np.array([1, 2, 0], dtype=sinfo.dtype)) |
1247 | 1284 | time.sleep(0.1) |
1248 | | - start = time.monotonic() |
1249 | | - while epochs.n_new_epochs != 2 and time.monotonic() - start < 3: |
1250 | | - epochs.acquire() |
1251 | | - time.sleep(0.5) |
1252 | 1285 | # check that we got 2 epochs on the code '1' |
1253 | | - assert epochs.n_new_epochs == 2 |
| 1286 | + _wait_for_epochs(epochs, 2) |
1254 | 1287 | events = epochs.events |
1255 | 1288 | assert_array_equal(events[events != 0], [1, 1]) |
1256 | 1289 | epochs.disconnect() |
|
0 commit comments