Skip to content

Commit 580408d

Browse files
authored
Expose seed for spike_locations (grid) (#4717)
1 parent 9817050 commit 580408d

3 files changed

Lines changed: 13 additions & 2 deletions

File tree

src/spikeinterface/postprocessing/spike_locations.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,9 @@ class ComputeSpikeLocations(BaseSpikeVectorExtension):
2020
The localization method to use
2121
method_kwargs : dict, default: dict()
2222
Other kwargs depending on the method.
23+
seed : int or None, default: None
24+
Seed for random number generator. Used by the `grid_convolution` method to
25+
reproducibly subsample peaks when computing the prototype waveform.
2326
2427
Returns
2528
-------
@@ -48,6 +51,7 @@ def _set_params(
4851
spike_retriever_kwargs=None,
4952
method="center_of_mass",
5053
method_kwargs={},
54+
seed=None,
5155
):
5256
if spike_retriever_kwargs is None:
5357
spike_retriever_kwargs = {}
@@ -57,6 +61,7 @@ def _set_params(
5761
spike_retriever_kwargs=spike_retriever_kwargs,
5862
method=method,
5963
method_kwargs=method_kwargs,
64+
seed=seed,
6065
)
6166

6267
def _get_pipeline_nodes(self):
@@ -77,6 +82,7 @@ def _get_pipeline_nodes(self):
7782
method_kwargs=self.params["method_kwargs"],
7883
ms_before=self.params["ms_before"],
7984
ms_after=self.params["ms_after"],
85+
seed=self.params.get("seed"),
8086
)
8187
return nodes
8288

src/spikeinterface/sortingcomponents/peak_localization/main.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -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:

src/spikeinterface/sortingcomponents/peak_selection.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -106,7 +106,6 @@ def select_peak_indices(peaks, method, seed, **method_kwargs):
106106

107107
selected_indices = []
108108

109-
seed = seed if seed else None
110109
rng = np.random.default_rng(seed=seed)
111110

112111
if method == "uniform":

0 commit comments

Comments
 (0)