Skip to content

Commit ce65bba

Browse files
ENH: vectorize annotation span conversion
Implements the get_annotation_spans() design suggested by @larsoner during review of #14240.
1 parent 6de8e97 commit ce65bba

4 files changed

Lines changed: 25 additions & 28 deletions

File tree

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
Add :meth:`mne.io.BaseRaw.get_annotation_span` to convert annotation spans to the time reference used by Raw data and plotting methods, by :newcontrib:`Tim Anderson`.
1+
Add :meth:`mne.io.BaseRaw.get_annotation_spans` to convert annotation spans to the time reference used by Raw data and plotting methods, by :newcontrib:`Tim Anderson`.

‎mne/io/base.py‎

Lines changed: 8 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -740,20 +740,15 @@ def annotations(self): # noqa: D401
740740
""":class:`~mne.Annotations` for marking segments of data."""
741741
return self._annotations
742742

743-
def get_annotation_span(self, index: int) -> tuple[float, float]:
744-
"""Get an annotation span relative to the first data sample.
745-
746-
Parameters
747-
----------
748-
index : int
749-
Index of the annotation in :attr:`annotations`.
743+
def get_annotation_spans(self) -> tuple[np.ndarray, np.ndarray]:
744+
"""Get annotation spans relative to the first data sample.
750745
751746
Returns
752747
-------
753-
tmin : float
754-
Annotation onset in seconds relative to the first data sample.
755-
tmax : float
756-
Annotation end in seconds relative to the first data sample.
748+
tmin : ndarray, shape (n_annotations,)
749+
Annotation onsets in seconds relative to the first data sample.
750+
tmax : ndarray, shape (n_annotations,)
751+
Annotation ends in seconds relative to the first data sample.
757752
758753
Notes
759754
-----
@@ -763,9 +758,8 @@ def get_annotation_span(self, index: int) -> tuple[float, float]:
763758
use zero at the first available data sample. This method converts
764759
between those time references.
765760
"""
766-
_validate_type(index, "int-like", "index")
767-
tmin = float(_sync_onset(self, self.annotations.onset[index]))
768-
tmax = tmin + float(self.annotations.duration[index])
761+
tmin = _sync_onset(self, self.annotations.onset)
762+
tmax = tmin + self.annotations.duration
769763
return tmin, tmax
770764

771765
@property

‎mne/io/tests/test_raw.py‎

Lines changed: 14 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -691,7 +691,7 @@ def test_crop_by_annotations(meas_date, first_samp):
691691

692692
@pytest.mark.parametrize("meas_date", [None, 0])
693693
@pytest.mark.parametrize("first_samp", [0, 50])
694-
def test_get_annotation_span(meas_date, first_samp):
694+
def test_get_annotation_spans(meas_date, first_samp):
695695
"""Test converting annotation spans to the Raw time reference."""
696696
sfreq = 10.0
697697
data = np.arange(120.0)[np.newaxis]
@@ -702,21 +702,24 @@ def test_get_annotation_span(meas_date, first_samp):
702702
verbose="error",
703703
)
704704
raw.set_meas_date(meas_date)
705-
raw.set_annotations(mne.Annotations(3.0, 1.0, "test"))
705+
raw.set_annotations(mne.Annotations([3.0, 3.0], [1.0, 0.5], ["test", "test 2"]))
706706

707-
tmin, tmax = raw.get_annotation_span(0)
708-
assert_allclose([tmin, tmax], [3.0, 4.0])
709-
got, times = raw.get_data(tmin=tmin, tmax=tmax, return_times=True)
707+
tmin, tmax = raw.get_annotation_spans()
708+
assert_allclose(tmin, [3.0, 3.0])
709+
assert_allclose(tmax, [3.5, 4.0])
710+
got, times = raw.get_data(tmin=tmin[1], tmax=tmax[1], return_times=True)
710711
assert_array_equal(got, data[:, 30:40])
711712
assert_allclose(times, np.arange(30, 40) / sfreq)
712713

713714
cropped = raw.copy().crop(2.0, 5.0)
714-
tmin, tmax = cropped.get_annotation_span(0)
715-
assert_allclose([tmin, tmax], [1.0, 2.0])
716-
assert_array_equal(cropped.get_data(tmin=tmin, tmax=tmax), data[:, 30:40])
717-
718-
with pytest.raises(TypeError, match="index must be .* int"):
719-
raw.get_annotation_span(0.0)
715+
tmin, tmax = cropped.get_annotation_spans()
716+
assert_allclose(tmin, [1.0, 1.0])
717+
assert_allclose(tmax, [1.5, 2.0])
718+
assert_array_equal(cropped.get_data(tmin=tmin[1], tmax=tmax[1]), data[:, 30:40])
719+
720+
raw.set_annotations(None)
721+
tmin, tmax = raw.get_annotation_spans()
722+
assert tmin.shape == tmax.shape == (0,)
720723

721724

722725
@pytest.mark.parametrize(

‎mne/viz/tests/test_raw.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -795,8 +795,8 @@ def test_plot_annotation_span(browser_backend):
795795
)
796796
raw.set_annotations(Annotations(3.0, 1.0, "test"))
797797

798-
tmin, tmax = raw.get_annotation_span(0)
799-
fig = raw.plot(start=tmin, duration=tmax - tmin, show=False)
798+
tmin, tmax = raw.get_annotation_spans()
799+
fig = raw.plot(start=tmin[0], duration=tmax[0] - tmin[0], show=False)
800800
assert fig._get_start_stop() == (30, 40)
801801
browser_backend._close_all()
802802

0 commit comments

Comments
 (0)