66import numpy as np
77
88from 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
1010from spikeinterface .core .recording_tools import get_noise_levels
1111from spikeinterface .core .template import Templates
1212from spikeinterface .core .waveform_tools import estimate_templates
1313from 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+ )
1519from spikeinterface .core .basesorting import minimum_spike_dtype
1620from spikeinterface .core .sparsity import compute_sparsity
1721from 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