@@ -25,6 +25,7 @@ def get_localization_pipeline_nodes(
2525 ms_before = 0.5 ,
2626 ms_after = 0.5 ,
2727 job_kwargs = None ,
28+ seed = None ,
2829):
2930
3031 assert (
@@ -52,7 +53,7 @@ def get_localization_pipeline_nodes(
5253
5354 method_kwargs = method_kwargs .copy ()
5455 method_kwargs ["prototype" ], _ , _ = get_prototype_and_waveforms_from_peaks (
55- recording , peaks = peak_source .peaks , ms_before = ms_before , ms_after = ms_after , job_kwargs = job_kwargs
56+ recording , peaks = peak_source .peaks , ms_before = ms_before , ms_after = ms_after , job_kwargs = job_kwargs , seed = seed
5657 )
5758
5859 localization_nodes = method_class (recording , parents = [peak_source , extract_dense_waveforms ], ** method_kwargs )
@@ -72,6 +73,7 @@ def localize_peaks(
7273 pipeline_kwargs = None ,
7374 verbose = False ,
7475 job_kwargs = None ,
76+ seed = None ,
7577 ** old_kwargs ,
7678) -> np .ndarray :
7779 """Localize peak (spike) in 2D or 3D depending the method.
@@ -102,6 +104,9 @@ def localize_peaks(
102104 If True, output is verbose
103105 job_kwargs : dict | None, default None
104106 A job kwargs dict. If None or empty dict, then the global one is used.
107+ seed : int or None, default: None
108+ Seed for random number generator. Used by `grid_convolution` to reproducibly
109+ subsample peaks when computing the prototype waveform.
105110
106111 {method_doc}
107112
@@ -150,6 +155,7 @@ def localize_peaks(
150155 ms_before = ms_before ,
151156 ms_after = ms_after ,
152157 job_kwargs = job_kwargs ,
158+ seed = seed ,
153159 )
154160
155161 if pipeline_kwargs is None :
0 commit comments