Skip to content

Commit 6590e0f

Browse files
committed
nodepipeline add skip_after_n_peaks option
1 parent b1d726a commit 6590e0f

2 files changed

Lines changed: 53 additions & 6 deletions

File tree

src/spikeinterface/core/node_pipeline.py

Lines changed: 19 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -497,6 +497,7 @@ def run_node_pipeline(
497497
folder=None,
498498
names=None,
499499
verbose=False,
500+
skip_after_n_peaks=None,
500501
):
501502
"""
502503
Common function to run pipeline with peak detector or already detected peak.
@@ -507,14 +508,19 @@ def run_node_pipeline(
507508
job_kwargs = fix_job_kwargs(job_kwargs)
508509
assert all(isinstance(node, PipelineNode) for node in nodes)
509510

511+
if skip_after_n_peaks is not None:
512+
skip_after_n_peaks_per_worker = skip_after_n_peaks / job_kwargs["n_jobs"]
513+
else:
514+
skip_after_n_peaks_per_worker = None
515+
510516
if gather_mode == "memory":
511517
gather_func = GatherToMemory()
512518
elif gather_mode == "npy":
513519
gather_func = GatherToNpy(folder, names, **gather_kwargs)
514520
else:
515521
raise ValueError(f"wrong gather_mode : {gather_mode}")
516522

517-
init_args = (recording, nodes)
523+
init_args = (recording, nodes, skip_after_n_peaks_per_worker)
518524

519525
processor = ChunkRecordingExecutor(
520526
recording,
@@ -533,19 +539,22 @@ def run_node_pipeline(
533539
return outs
534540

535541

536-
def _init_peak_pipeline(recording, nodes):
542+
def _init_peak_pipeline(recording, nodes, skip_after_n_peaks_per_worker):
537543
# create a local dict per worker
538544
worker_ctx = {}
539545
worker_ctx["recording"] = recording
540546
worker_ctx["nodes"] = nodes
541547
worker_ctx["max_margin"] = max(node.get_trace_margin() for node in nodes)
548+
worker_ctx["skip_after_n_peaks_per_worker"] = skip_after_n_peaks_per_worker
549+
worker_ctx["num_peaks"] = 0
542550
return worker_ctx
543551

544552

545553
def _compute_peak_pipeline_chunk(segment_index, start_frame, end_frame, worker_ctx):
546554
recording = worker_ctx["recording"]
547555
max_margin = worker_ctx["max_margin"]
548556
nodes = worker_ctx["nodes"]
557+
skip_after_n_peaks_per_worker = worker_ctx["skip_after_n_peaks_per_worker"]
549558

550559
recording_segment = recording._recording_segments[segment_index]
551560
node0 = nodes[0]
@@ -557,7 +566,11 @@ def _compute_peak_pipeline_chunk(segment_index, start_frame, end_frame, worker_c
557566
else:
558567
# PeakDetector always need traces
559568
load_trace_and_compute = True
560-
569+
570+
if skip_after_n_peaks_per_worker is not None:
571+
if worker_ctx["num_peaks"] > skip_after_n_peaks_per_worker:
572+
load_trace_and_compute = False
573+
561574
if load_trace_and_compute:
562575
traces_chunk, left_margin, right_margin = get_chunk_with_margin(
563576
recording_segment, start_frame, end_frame, None, max_margin, add_zeros=True
@@ -590,6 +603,9 @@ def _compute_peak_pipeline_chunk(segment_index, start_frame, end_frame, worker_c
590603
node_output = node.compute(traces_chunk, *node_input_args)
591604
pipeline_outputs[node] = node_output
592605

606+
if skip_after_n_peaks_per_worker is not None and isinstance(node, PeakSource):
607+
worker_ctx["num_peaks"] += node_output[0].size
608+
593609
# propagate the output
594610
pipeline_outputs_tuple = tuple()
595611
for node in nodes:

src/spikeinterface/core/tests/test_node_pipeline.py

Lines changed: 34 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -83,7 +83,7 @@ def test_run_node_pipeline(cache_folder_creation):
8383
extremum_channel_inds = get_template_extremum_channel(sorting_analyzer, peak_sign="neg", outputs="index")
8484

8585
peaks = sorting_to_peaks(sorting, extremum_channel_inds, spike_peak_dtype)
86-
print(peaks.size)
86+
# print(peaks.size)
8787

8888
peak_retriever = PeakRetriever(recording, peaks)
8989
# this test when no spikes in last chunks
@@ -191,6 +191,37 @@ def test_run_node_pipeline(cache_folder_creation):
191191
unpickled_node = pickle.loads(pickled_node)
192192

193193

194+
def test_skip_after_n_peaks():
195+
recording, sorting = generate_ground_truth_recording(num_channels=10, num_units=10, durations=[10.0])
196+
197+
# job_kwargs = dict(chunk_duration="0.5s", n_jobs=2, progress_bar=False)
198+
job_kwargs = dict(chunk_duration="0.5s", n_jobs=1, progress_bar=False)
199+
200+
spikes = sorting.to_spike_vector()
201+
202+
# create peaks from spikes
203+
sorting_analyzer = create_sorting_analyzer(sorting, recording, format="memory")
204+
sorting_analyzer.compute(["random_spikes", "templates"], **job_kwargs)
205+
extremum_channel_inds = get_template_extremum_channel(sorting_analyzer, peak_sign="neg", outputs="index")
206+
207+
peaks = sorting_to_peaks(sorting, extremum_channel_inds, spike_peak_dtype)
208+
# print(peaks.size)
209+
210+
node0 = PeakRetriever(recording, peaks)
211+
node1 = AmplitudeExtractionNode(recording, parents=[node0], param0=6.6, return_output=True)
212+
nodes = [node0, node1]
213+
214+
skip_after_n_peaks = 30
215+
some_amplitudes = run_node_pipeline(recording, nodes, job_kwargs, gather_mode="memory", skip_after_n_peaks=skip_after_n_peaks)
216+
217+
assert some_amplitudes.size >= skip_after_n_peaks
218+
assert some_amplitudes.size < spikes.size
219+
220+
221+
222+
194223
if __name__ == "__main__":
195-
folder = Path("./cache_folder/core")
196-
test_run_node_pipeline(folder)
224+
# folder = Path("./cache_folder/core")
225+
# test_run_node_pipeline(folder)
226+
227+
test_skip_after_n_peaks()

0 commit comments

Comments
 (0)