Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions src/spikeinterface/core/core_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -779,3 +779,8 @@ def is_path_remote(path: str | Path) -> bool:
def ms_to_samples(ms: float, sampling_frequency: float) -> int:
"""Convert a duration in milliseconds to the nearest number of samples."""
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
68 changes: 53 additions & 15 deletions src/spikeinterface/core/node_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,14 @@
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:

# If False (general case) then compute(traces_chunk, *node_input_args)
# If True then compute(traces_chunk, start_frame, end_frame, segment_index, max_margin, *node_input_args)
name = None
_compute_has_extended_signature = False

def __init__(
Expand Down Expand Up @@ -297,8 +298,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 +322,42 @@ 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
self.sparse_waveforms = False


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 +370,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 +391,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 +408,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 +432,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 +456,8 @@ def __init__(
parents=parents,
ms_before=ms_before,
ms_after=ms_after,
nbefore=nbefore,
nafter=nafter,
return_output=return_output,
)

Expand All @@ -434,6 +471,7 @@ def __init__(
self.radius_um = radius_um
self.neighbours_mask = self.channel_distance <= radius_um
self.max_num_chans = np.max(np.sum(self.neighbours_mask, axis=1))
self.sparse_waveforms = True

def get_margin(self):
return max(self.nbefore, self.nafter)
Expand Down
Loading
Loading