Skip to content
5 changes: 5 additions & 0 deletions src/spikeinterface/core/core_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -782,6 +782,11 @@ def ms_to_samples(ms: float, sampling_frequency: float) -> int:
return round(ms * sampling_frequency / 1000.0)


def samples_to_ms(samples: int, sampling_frequency: float) -> float:
"""Convert a duration in samples to milliseconds."""
return samples / sampling_frequency * 1000.0


def slice_rows(array: np.ndarray | zarr.Array, row_indices: np.ndarray | list) -> np.ndarray:
"""
Slice a 2D array to select specific rows based on provided indices.
Expand Down
65 changes: 50 additions & 15 deletions src/spikeinterface/core/node_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from spikeinterface.core import BaseRecording, get_chunk_with_margin
from spikeinterface.core.job_tools import TimeSeriesChunkExecutor, fix_job_kwargs, _shared_job_kwargs_doc
from spikeinterface.core import get_channel_distances
from spikeinterface.core.core_tools import ms_to_samples
from spikeinterface.core.core_tools import ms_to_samples, samples_to_ms


class PipelineNode:
Expand Down Expand Up @@ -297,8 +297,10 @@ class WaveformsNode(PipelineNode):
def __init__(
self,
recording: BaseRecording,
ms_before: float,
ms_after: float,
ms_before: float | None = None,
ms_after: float | None = None,
nbefore: int | None = None,
nafter: int | None = None,
parents: list[PipelineNode] | None = None,
return_output: bool = False,
):
Expand All @@ -319,22 +321,41 @@ def __init__(
return_output : bool, default: False
Whether or not the output of the node is returned by the pipeline
"""
if ms_before is None and nbefore is None:
raise ValueError("Either ms_before or nbefore must be provided.")
if ms_after is None and nafter is None:
raise ValueError("Either ms_after or nafter must be provided.")
if ms_before is not None and nbefore is not None:
raise ValueError("Only one of ms_before or nbefore should be provided.")
if ms_after is not None and nafter is not None:
raise ValueError("Only one of ms_after or nafter should be provided.")

PipelineNode.__init__(self, recording, parents=parents, return_output=return_output)
self.recording = recording
self.ms_before = ms_before
self.ms_after = ms_after
self.nbefore = ms_to_samples(ms_before, recording.get_sampling_frequency())
self.nafter = ms_to_samples(ms_after, recording.get_sampling_frequency())
sampling_frequency = recording.sampling_frequency
if nbefore is not None:
self.nbefore = nbefore
self.ms_before = samples_to_ms(nbefore, sampling_frequency)
else:
self.ms_before = ms_before
self.nbefore = ms_to_samples(ms_before, sampling_frequency)
if nafter is not None:
self.nafter = nafter
self.ms_after = samples_to_ms(nafter, sampling_frequency)
else:
self.ms_after = ms_after
self.nafter = ms_to_samples(ms_after, sampling_frequency)
self.neighbours_mask = None


class ExtractDenseWaveforms(WaveformsNode):
def __init__(
self,
recording: BaseRecording,
ms_before: float,
ms_after: float,
ms_before: float | None = None,
ms_after: float | None = None,
nbefore: int | None = None,
nafter: int | None = None,
parents: list[PipelineNode] | None = None,
return_output: bool = False,
):
Expand All @@ -347,10 +368,14 @@ def __init__(
----------
recording : BaseRecording
The recording object.
ms_before : float
ms_before : float | None
The number of milliseconds to include before the peak of the spike
ms_after : float
ms_after : float | None
The number of milliseconds to include after the peak of the spike
nbefore : int | None, default: None
The number of samples to include before the peak of the spike
nafter : int | None, default: None
The number of samples to include after the peak of the spike
parents : list[PipelineNode] | None, default: None
Pass parents nodes to perform a previous computation
return_output : bool, default: False
Expand All @@ -364,6 +389,8 @@ def __init__(
parents=parents,
ms_before=ms_before,
ms_after=ms_after,
nbefore=nbefore,
nafter=nafter,
return_output=return_output,
)

Expand All @@ -379,8 +406,10 @@ class ExtractSparseWaveforms(WaveformsNode):
def __init__(
self,
recording: BaseRecording,
ms_before: float,
ms_after: float,
ms_before: float | None = None,
ms_after: float | None = None,
nbefore: int | None = None,
nafter: int | None = None,
parents: list[PipelineNode] | None = None,
return_output: bool = False,
radius_um: float = 100.0,
Expand All @@ -401,10 +430,14 @@ def __init__(
----------
recording : BaseRecording
The recording object
ms_before : float
ms_before : float | None
The number of milliseconds to include before the peak of the spike
ms_after : float
ms_after : float | None
The number of milliseconds to include after the peak of the spike
nbefore : int | None, default: None
The number of samples to include before the peak of the spike
nafter : int | None, default: None
The number of samples to include after the peak of the spike
parents : list[PipelineNode] | None, default: None
Pass parents nodes to perform a previous computation
return_output : bool, default: False
Expand All @@ -421,6 +454,8 @@ def __init__(
parents=parents,
ms_before=ms_before,
ms_after=ms_after,
nbefore=nbefore,
nafter=nafter,
return_output=return_output,
)

Expand Down
Loading
Loading