Skip to content

Commit 4034a2c

Browse files
committed
Merge branch 'main' into improve_inter_sample_shift_docstring
2 parents 1065902 + 418bb86 commit 4034a2c

10 files changed

Lines changed: 353 additions & 80 deletions

File tree

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2092,6 +2092,13 @@ def load_data(self):
20922092
import pandas as pd
20932093

20942094
ext_data = pd.read_csv(ext_data_file, index_col=0)
2095+
# we need to cast the index to the unit id dtype (int or str)
2096+
unit_ids = self.sorting_analyzer.unit_ids
2097+
if ext_data.shape[0] == unit_ids.size:
2098+
# we force dtype to be the same as unit_ids
2099+
if ext_data.index.dtype != unit_ids.dtype:
2100+
ext_data.index = ext_data.index.astype(unit_ids.dtype)
2101+
20952102
elif ext_data_file.suffix == ".pkl":
20962103
with ext_data_file.open("rb") as f:
20972104
ext_data = pickle.load(f)

src/spikeinterface/core/template.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -205,6 +205,7 @@ def to_sparse(self, sparsity):
205205
unit_ids=self.unit_ids,
206206
probe=self.probe,
207207
check_for_consistent_sparsity=self.check_for_consistent_sparsity,
208+
is_scaled=self.is_scaled,
208209
)
209210

210211
def get_one_template_dense(self, unit_index):

src/spikeinterface/sorters/internal/spyking_circus2.py

Lines changed: 77 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -6,12 +6,16 @@
66
import numpy as np
77

88
from spikeinterface.core import NumpySorting
9-
from spikeinterface.core.job_tools import fix_job_kwargs
9+
from spikeinterface.core.job_tools import fix_job_kwargs, split_job_kwargs
1010
from spikeinterface.core.recording_tools import get_noise_levels
1111
from spikeinterface.core.template import Templates
1212
from spikeinterface.core.waveform_tools import estimate_templates
1313
from spikeinterface.preprocessing import common_reference, whiten, bandpass_filter, correct_motion
14-
from spikeinterface.sortingcomponents.tools import cache_preprocessing
14+
from spikeinterface.sortingcomponents.tools import (
15+
cache_preprocessing,
16+
get_prototype_and_waveforms_from_recording,
17+
get_shuffled_recording_slices,
18+
)
1519
from spikeinterface.core.basesorting import minimum_spike_dtype
1620
from spikeinterface.core.sparsity import compute_sparsity
1721
from spikeinterface.core.sortinganalyzer import create_sorting_analyzer
@@ -26,7 +30,7 @@ class Spykingcircus2Sorter(ComponentsBasedSorter):
2630
_default_params = {
2731
"general": {"ms_before": 2, "ms_after": 2, "radius_um": 75},
2832
"sparsity": {"method": "snr", "amplitude_mode": "peak_to_peak", "threshold": 0.25},
29-
"filtering": {"freq_min": 150, "freq_max": 7000, "ftype": "bessel", "filter_order": 2},
33+
"filtering": {"freq_min": 150, "freq_max": 7000, "ftype": "bessel", "filter_order": 2, "margin_ms": 10},
3034
"whitening": {"mode": "local", "regularize": False},
3135
"detection": {"peak_sign": "neg", "detect_threshold": 4},
3236
"selection": {
@@ -53,6 +57,7 @@ class Spykingcircus2Sorter(ComponentsBasedSorter):
5357
"cache_preprocessing": {"mode": "memory", "memory_limit": 0.5, "delete_cache": True},
5458
"multi_units_only": False,
5559
"job_kwargs": {"n_jobs": 0.5},
60+
"seed": 42,
5661
"debug": False,
5762
}
5863

@@ -74,18 +79,21 @@ class Spykingcircus2Sorter(ComponentsBasedSorter):
7479
"merging": "A dictionary to specify the final merging param to group cells after template matching (get_potential_auto_merge)",
7580
"motion_correction": "A dictionary to be provided if motion correction has to be performed (dense probe only)",
7681
"apply_preprocessing": "Boolean to specify whether circus 2 should preprocess the recording or not. If yes, then high_pass filtering + common\
77-
median reference + zscore",
82+
median reference + whitening",
83+
"apply_motion_correction": "Boolean to specify whether circus 2 should apply motion correction to the recording or not",
84+
"matched_filtering": "Boolean to specify whether circus 2 should detect peaks via matched filtering (slightly slower)",
7885
"cache_preprocessing": "How to cache the preprocessed recording. Mode can be memory, file, zarr, with extra arguments. In case of memory (default), \
7986
memory_limit will control how much RAM can be used. In case of folder or zarr, delete_cache controls if cache is cleaned after sorting",
8087
"multi_units_only": "Boolean to get only multi units activity (i.e. one template per electrode)",
8188
"job_kwargs": "A dictionary to specify how many jobs and which parameters they should used",
89+
"seed": "An int to control how chunks are shuffled while detecting peaks",
8290
"debug": "Boolean to specify if internal data structures made during the sorting should be kept for debugging",
8391
}
8492

8593
sorter_description = """Spyking Circus 2 is a rewriting of Spyking Circus, within the SpikeInterface framework
8694
It uses a more conservative clustering algorithm (compared to Spyking Circus), which is less prone to hallucinate units and/or find noise.
8795
In addition, it also uses a full Orthogonal Matching Pursuit engine to reconstruct the traces, leading to more spikes
88-
being discovered."""
96+
being discovered. The code is much faster and memory efficient, inheriting from all the preprocessing possibilities of spikeinterface"""
8997

9098
@classmethod
9199
def get_sorter_version(cls):
@@ -114,7 +122,7 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
114122
from spikeinterface.sortingcomponents.clustering import find_cluster_from_peaks
115123
from spikeinterface.sortingcomponents.matching import find_spikes_from_templates
116124
from spikeinterface.sortingcomponents.tools import remove_empty_templates
117-
from spikeinterface.sortingcomponents.tools import get_prototype_spike, check_probe_for_drift_correction
125+
from spikeinterface.sortingcomponents.tools import check_probe_for_drift_correction
118126

119127
job_kwargs = fix_job_kwargs(params["job_kwargs"])
120128
job_kwargs.update({"progress_bar": verbose})
@@ -131,10 +139,14 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
131139
## First, we are filtering the data
132140
filtering_params = params["filtering"].copy()
133141
if params["apply_preprocessing"]:
142+
if verbose:
143+
print("Preprocessing the recording (bandpass filtering + CMR + whitening)")
134144
recording_f = bandpass_filter(recording, **filtering_params, dtype="float32")
135145
if num_channels > 1:
136146
recording_f = common_reference(recording_f)
137147
else:
148+
if verbose:
149+
print("Skipping preprocessing (whitening only)")
138150
recording_f = recording
139151
recording_f.annotate(is_filtered=True)
140152

@@ -157,12 +169,14 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
157169
# TODO add , regularize=True chen ready
158170
whitening_kwargs = params["whitening"].copy()
159171
whitening_kwargs["dtype"] = "float32"
160-
whitening_kwargs["radius_um"] = radius_um
172+
whitening_kwargs["regularize"] = whitening_kwargs.get("regularize", False)
161173
if num_channels == 1:
162174
whitening_kwargs["regularize"] = False
175+
if whitening_kwargs["regularize"]:
176+
whitening_kwargs["regularize_kwargs"] = {"method": "LedoitWolf"}
163177

164178
recording_w = whiten(recording_f, **whitening_kwargs)
165-
noise_levels = get_noise_levels(recording_w, return_scaled=False)
179+
noise_levels = get_noise_levels(recording_w, return_scaled=False, **job_kwargs)
166180

167181
if recording_w.check_serializability("json"):
168182
recording_w.dump(sorter_output_folder / "preprocessed_recording.json", relative_to=None)
@@ -173,42 +187,69 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
173187

174188
## Then, we are detecting peaks with a locally_exclusive method
175189
detection_params = params["detection"].copy()
176-
detection_params.update(job_kwargs)
177-
178-
detection_params["radius_um"] = detection_params.get("radius_um", 50)
190+
selection_params = params["selection"].copy()
191+
detection_params["radius_um"] = radius_um
179192
detection_params["exclude_sweep_ms"] = exclude_sweep_ms
180193
detection_params["noise_levels"] = noise_levels
181194

182195
fs = recording_w.get_sampling_frequency()
183196
nbefore = int(ms_before * fs / 1000.0)
184197
nafter = int(ms_after * fs / 1000.0)
185198

199+
skip_peaks = not params["multi_units_only"] and selection_params.get("method", "uniform") == "uniform"
200+
max_n_peaks = selection_params["n_peaks_per_channel"] * num_channels
201+
n_peaks = max(selection_params["min_n_peaks"], max_n_peaks)
202+
203+
if params["debug"]:
204+
clustering_folder = sorter_output_folder / "clustering"
205+
clustering_folder.mkdir(parents=True, exist_ok=True)
206+
np.save(clustering_folder / "noise_levels.npy", noise_levels)
207+
186208
if params["matched_filtering"]:
187-
peaks = detect_peaks(recording_w, "locally_exclusive", **detection_params, skip_after_n_peaks=5000)
188-
prototype = get_prototype_spike(recording_w, peaks, ms_before, ms_after, **job_kwargs)
209+
prototype, waveforms, _ = get_prototype_and_waveforms_from_recording(
210+
recording_w,
211+
n_peaks=10000,
212+
ms_before=ms_before,
213+
ms_after=ms_after,
214+
seed=params["seed"],
215+
**detection_params,
216+
**job_kwargs,
217+
)
189218
detection_params["prototype"] = prototype
190219
detection_params["ms_before"] = ms_before
191-
peaks = detect_peaks(recording_w, "matched_filtering", **detection_params)
220+
if params["debug"]:
221+
np.save(clustering_folder / "waveforms.npy", waveforms)
222+
np.save(clustering_folder / "prototype.npy", prototype)
223+
if skip_peaks:
224+
detection_params["skip_after_n_peaks"] = n_peaks
225+
detection_params["recording_slices"] = get_shuffled_recording_slices(
226+
recording_w, seed=params["seed"], **job_kwargs
227+
)
228+
peaks = detect_peaks(recording_w, "matched_filtering", **detection_params, **job_kwargs)
192229
else:
193-
peaks = detect_peaks(recording_w, "locally_exclusive", **detection_params)
230+
waveforms = None
231+
if skip_peaks:
232+
detection_params["skip_after_n_peaks"] = n_peaks
233+
detection_params["recording_slices"] = get_shuffled_recording_slices(
234+
recording_w, seed=params["seed"], **job_kwargs
235+
)
236+
peaks = detect_peaks(recording_w, "locally_exclusive", **detection_params, **job_kwargs)
194237

195-
if verbose:
196-
print("We found %d peaks in total" % len(peaks))
238+
if not skip_peaks and verbose:
239+
print("Found %d peaks in total" % len(peaks))
197240

198241
if params["multi_units_only"]:
199242
sorting = NumpySorting.from_peaks(peaks, sampling_frequency, unit_ids=recording_w.unit_ids)
200243
else:
201244
## We subselect a subset of all the peaks, by making the distributions os SNRs over all
202245
## channels as flat as possible
203246
selection_params = params["selection"]
204-
selection_params["n_peaks"] = min(len(peaks), selection_params["n_peaks_per_channel"] * num_channels)
205-
selection_params["n_peaks"] = max(selection_params["min_n_peaks"], selection_params["n_peaks"])
206-
247+
selection_params["n_peaks"] = n_peaks
207248
selection_params.update({"noise_levels": noise_levels})
208249
selected_peaks = select_peaks(peaks, **selection_params)
209250

210251
if verbose:
211-
print("We kept %d peaks for clustering" % len(selected_peaks))
252+
print("Kept %d peaks for clustering" % len(selected_peaks))
212253

213254
## We launch a clustering (using hdbscan) relying on positions and features extracted on
214255
## the fly from the snippets
@@ -218,10 +259,13 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
218259
clustering_params["radius_um"] = radius_um
219260
clustering_params["waveforms"]["ms_before"] = ms_before
220261
clustering_params["waveforms"]["ms_after"] = ms_after
262+
clustering_params["few_waveforms"] = waveforms
221263
clustering_params["noise_levels"] = noise_levels
222-
clustering_params["ms_before"] = exclude_sweep_ms
223-
clustering_params["ms_after"] = exclude_sweep_ms
264+
clustering_params["ms_before"] = ms_before
265+
clustering_params["ms_after"] = ms_after
266+
clustering_params["verbose"] = verbose
224267
clustering_params["tmp_folder"] = sorter_output_folder / "clustering"
268+
clustering_params["noise_threshold"] = detection_params.get("detect_threshold", 4)
225269

226270
legacy = clustering_params.get("legacy", True)
227271

@@ -246,12 +290,8 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
246290
unit_ids = np.arange(len(np.unique(labeled_peaks["unit_index"])))
247291
sorting = NumpySorting(labeled_peaks, sampling_frequency, unit_ids=unit_ids)
248292

249-
clustering_folder = sorter_output_folder / "clustering"
250-
clustering_folder.mkdir(parents=True, exist_ok=True)
251-
252-
if not params["debug"]:
253-
shutil.rmtree(clustering_folder)
254-
else:
293+
if params["debug"]:
294+
np.save(clustering_folder / "peak_labels", peak_labels)
255295
np.save(clustering_folder / "labels", labels)
256296
np.save(clustering_folder / "peaks", selected_peaks)
257297

@@ -294,7 +334,7 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
294334
np.save(fitting_folder / "spikes", spikes)
295335

296336
if verbose:
297-
print("We found %d spikes" % len(spikes))
337+
print("Found %d spikes" % len(spikes))
298338

299339
## And this is it! We have a spyking circus
300340
sorting = np.zeros(spikes.size, dtype=minimum_spike_dtype)
@@ -334,10 +374,10 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
334374
sorting.save(folder=curation_folder)
335375
# np.save(fitting_folder / "amplitudes", guessed_amplitudes)
336376

337-
sorting = final_cleaning_circus(recording_w, sorting, templates, **merging_params)
377+
sorting = final_cleaning_circus(recording_w, sorting, templates, merging_params, **job_kwargs)
338378

339379
if verbose:
340-
print(f"Final merging, keeping {len(sorting.unit_ids)} units")
380+
print(f"Kept {len(sorting.unit_ids)} units after final merging")
341381

342382
folder_to_delete = None
343383
cache_mode = params["cache_preprocessing"].get("mode", "memory")
@@ -376,17 +416,18 @@ def create_sorting_analyzer_with_templates(sorting, recording, templates, remove
376416
return sa
377417

378418

379-
def final_cleaning_circus(recording, sorting, templates, **merging_kwargs):
419+
def final_cleaning_circus(recording, sorting, templates, merging_kwargs, **job_kwargs):
380420

381421
from spikeinterface.core.sorting_tools import apply_merges_to_sorting
382422

383423
sa = create_sorting_analyzer_with_templates(sorting, recording, templates)
384424

385-
sa.compute("unit_locations", method="monopolar_triangulation")
425+
sa.compute("unit_locations", method="monopolar_triangulation", **job_kwargs)
386426
similarity_kwargs = merging_kwargs.pop("similarity_kwargs", {})
387-
sa.compute("template_similarity", **similarity_kwargs)
427+
sa.compute("template_similarity", **similarity_kwargs, **job_kwargs)
388428
correlograms_kwargs = merging_kwargs.pop("correlograms_kwargs", {})
389-
sa.compute("correlograms", **correlograms_kwargs)
429+
sa.compute("correlograms", **correlograms_kwargs, **job_kwargs)
430+
390431
auto_merge_kwargs = merging_kwargs.pop("auto_merge", {})
391432
merges = get_potential_auto_merge(sa, resolve_graph=True, **auto_merge_kwargs)
392433
sorting = apply_merges_to_sorting(sa.sorting, merges)

0 commit comments

Comments
 (0)