Skip to content

Commit 22e4161

Browse files
authored
Merge pull request #3677 from chrishalcrow/add-sorting-time-slice
Add `time_slice` method to `BaseSorting`
2 parents 4e34fc7 + 1facc03 commit 22e4161

3 files changed

Lines changed: 69 additions & 2 deletions

File tree

src/spikeinterface/core/baserecording.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -772,7 +772,7 @@ def frame_slice(self, start_frame: int | None, end_frame: int | None) -> BaseRec
772772

773773
def time_slice(self, start_time: float | None, end_time: float) -> BaseRecording:
774774
"""
775-
Returns a new recording with sliced time. Note that this operation is not in place.
775+
Returns a new recording object, restricted to the time interval [start_time, end_time].
776776
777777
Parameters
778778
----------

src/spikeinterface/core/basesorting.py

Lines changed: 40 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -179,7 +179,9 @@ def get_unit_spike_train(
179179
return spike_frames
180180

181181
def register_recording(self, recording, check_spike_frames=True):
182-
"""Register a recording to the sorting.
182+
"""
183+
Register a recording to the sorting. If the sorting and recording both contain
184+
time information, the recording’s time information will be used.
183185
184186
Parameters
185187
----------
@@ -472,6 +474,43 @@ def frame_slice(self, start_frame, end_frame, check_spike_frames=True):
472474
)
473475
return sub_sorting
474476

477+
def time_slice(self, start_time: float | None, end_time: float | None) -> BaseSorting:
478+
"""
479+
Returns a new sorting object, restricted to the time interval [start_time, end_time].
480+
481+
Parameters
482+
----------
483+
start_time : float | None, default: None
484+
The start time in seconds. If not provided it is set to 0.
485+
end_time : float | None, default: None
486+
The end time in seconds. If not provided it is set to the total duration.
487+
488+
Returns
489+
-------
490+
BaseSorting
491+
A new sorting object with only samples between start_time and end_time
492+
"""
493+
494+
assert self.get_num_segments() == 1, "Time slicing is only supported for single segment sortings."
495+
496+
start_frame = self.time_to_sample_index(start_time, segment_index=0) if start_time else None
497+
end_frame = self.time_to_sample_index(end_time, segment_index=0) if end_time else None
498+
499+
return self.frame_slice(start_frame=start_frame, end_frame=end_frame)
500+
501+
def time_to_sample_index(self, time, segment_index=0):
502+
"""
503+
Transform time in seconds into sample index
504+
"""
505+
if self.has_recording():
506+
sample_index = self._recording.time_to_sample_index(time, segment_index=segment_index)
507+
else:
508+
segment = self._sorting_segments[segment_index]
509+
t_start = segment._t_start if segment._t_start is not None else 0
510+
sample_index = round((time - t_start) * self.get_sampling_frequency())
511+
512+
return sample_index
513+
475514
def get_all_spike_trains(self, outputs="unit_id"):
476515
"""
477516
Return all spike trains concatenated.

src/spikeinterface/core/tests/test_basesorting.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,8 @@
2525
from spikeinterface.core.testing import check_sorted_arrays_equal, check_sortings_equal
2626
from spikeinterface.core.generate import generate_sorting
2727

28+
from spikeinterface.core import generate_recording, generate_ground_truth_recording
29+
2830

2931
def test_BaseSorting(create_cache_folder):
3032
cache_folder = create_cache_folder
@@ -206,6 +208,32 @@ def test_empty_sorting():
206208
assert spikes.shape == (0,)
207209

208210

211+
def test_time_slice():
212+
213+
sampling_frequency = 10_000.0
214+
215+
# no recording attached to sorting
216+
sorting = generate_sorting(durations=[1.0], sampling_frequency=sampling_frequency)
217+
218+
sliced_sorting_times = sorting.time_slice(start_time=0.1, end_time=0.8)
219+
sliced_sorting_frames = sorting.frame_slice(start_frame=1000, end_frame=8000)
220+
221+
assert np.allclose(
222+
sliced_sorting_times.to_spike_vector()["sample_index"], sliced_sorting_frames.to_spike_vector()["sample_index"]
223+
)
224+
225+
# with recording
226+
recording, sorting = generate_ground_truth_recording(durations=[1.0], sampling_frequency=sampling_frequency)
227+
sorting.register_recording(recording)
228+
229+
sliced_sorting_times = sorting.time_slice(start_time=0.1, end_time=0.8)
230+
sliced_sorting_frames = sorting.frame_slice(start_frame=1000, end_frame=8000)
231+
232+
assert np.allclose(
233+
sliced_sorting_times.to_spike_vector()["sample_index"], sliced_sorting_frames.to_spike_vector()["sample_index"]
234+
)
235+
236+
209237
if __name__ == "__main__":
210238
test_BaseSorting()
211239
test_npy_sorting()

0 commit comments

Comments
 (0)