Skip to content

Commit 356b5a9

Browse files
committed
Refactor to make time_to_sample_index for BaseSorting
1 parent df9f4a7 commit 356b5a9

1 file changed

Lines changed: 13 additions & 6 deletions

File tree

src/spikeinterface/core/basesorting.py

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -491,16 +491,23 @@ def time_slice(self, start_time: float | None, end_time: float | None) -> BaseSo
491491

492492
assert self.get_num_segments() == 1, "Time slicing is only supported for single segment sortings."
493493

494+
start_frame = self.time_to_sample_index(start_time, segment_index=0) if start_time else None
495+
end_frame = self.time_to_sample_index(end_time, segment_index=0) if end_time else None
496+
497+
return self.frame_slice(start_frame=start_frame, end_frame=end_frame)
498+
499+
def time_to_sample_index(self, time, segment_index=0):
500+
"""
501+
Transform time in seconds into sample index
502+
"""
494503
if self.has_recording():
495-
start_frame = self._recording.time_to_sample_index(start_time) if start_time else None
496-
end_frame = self._recording.time_to_sample_index(end_time) if end_time else None
504+
sample_index = self._recording.time_to_sample_index(time)
497505
else:
498-
segment = self._sorting_segments[0]
506+
segment = self._sorting_segments[segment_index]
499507
t_start = segment._t_start if segment._t_start is not None else 0
500-
start_frame = round((start_time - t_start) * self.get_sampling_frequency())
501-
end_frame = round((end_time - t_start) * self.get_sampling_frequency())
508+
sample_index = round((time - t_start) * self.get_sampling_frequency())
502509

503-
return self.frame_slice(start_frame=start_frame, end_frame=end_frame)
510+
return sample_index
504511

505512
def get_all_spike_trains(self, outputs="unit_id"):
506513
"""

0 commit comments

Comments
 (0)