@@ -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
545553def _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 :
0 commit comments