Skip to content

Commit 508de1e

Browse files
authored
Improve merging and iterative merging (#3487)
1 parent 2c6e800 commit 508de1e

29 files changed

Lines changed: 1036 additions & 188 deletions

doc/modules/curation.rst

Lines changed: 43 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -120,31 +120,31 @@ the unit (among the redundant ones), with a better template alignment.
120120
Auto-merging units
121121
^^^^^^^^^^^^^^^^^^
122122

123-
The :py:func:`~spikeinterface.curation.get_potential_auto_merge` function returns a list of potential merges.
123+
The :py:func:`~spikeinterface.curation.compute_merge_unit_groups` function returns a list of potential merges.
124124
The list of potential merges can be then applied to the sorting output.
125-
:py:func:`~spikeinterface.curation.get_potential_auto_merge` has many internal tricks and steps to identify potential
125+
:py:func:`~spikeinterface.curation.compute_merge_unit_groups` has many internal tricks and steps to identify potential
126126
merges. It offers multiple "presets" and the flexibility to apply individual steps, with different parameters.
127127
**Read the function documentation carefully and do not apply it blindly!**
128128

129129

130130
.. code-block:: python
131131
132132
from spikeinterface import create_sorting_analyzer
133-
from spikeinterface.curation import get_potential_auto_merge
133+
from spikeinterface.curation import compute_merge_unit_groups
134134
135135
analyzer = create_sorting_analyzer(sorting=sorting, recording=recording)
136136
137137
# some extensions are required
138138
analyzer.compute(["random_spikes", "templates", "template_similarity", "correlograms"])
139139
140140
# merges is a list of unit pairs, with unit_ids to be merged.
141-
merge_unit_pairs = get_potential_auto_merge(
141+
merge_unit_pairs = compute_merge_unit_groups(
142142
analyzer=analyzer,
143143
preset="similarity_correlograms",
144144
)
145145
# with resolve_graph=True, merges_resolved is a list of merge groups,
146146
# which can contain more than two units
147-
merge_unit_groups = get_potential_auto_merge(
147+
merge_unit_groups = compute_merge_unit_groups(
148148
analyzer=analyzer,
149149
preset="similarity_correlograms",
150150
resolve_graph=True
@@ -153,6 +153,44 @@ merges. It offers multiple "presets" and the flexibility to apply individual ste
153153
# here we apply the merges
154154
analyzer_merged = analyzer.merge_units(merge_unit_groups=merge_unit_groups)
155155
156+
There is also the convenient :py:func:`~spikeinterface.curation.auto_merge_units` function that combines the
157+
:py:func:`~spikeinterface.curation.compute_merge_unit_groups` and :py:func:`~spikeinterface.core.SortingAnalyzer.merge_units` functions.
158+
This is a high level function that allows you to apply either one or several presets/lists of steps in one go. For example, let's
159+
assume you want to apply the "x_contamination" preset, but iteratively and with slightly different parameters: first,
160+
you want to focus on the templates that are very similar, according to their template similarities, before
161+
considering those that might be more distant. Such a greedy and iterative scheme has been proved to be less
162+
prone to wrong merges. To do so, you'll need to do the following:
163+
164+
.. code-block:: python
165+
166+
from spikeinterface import create_sorting_analyzer
167+
from spikeinterface.curation import auto_merge_units
168+
169+
analyzer = create_sorting_analyzer(sorting=sorting, recording=recording)
170+
171+
# some extensions are required
172+
analyzer.compute(["random_spikes", "templates", "template_similarity", "correlograms"])
173+
analyzer.compute("unit_locations", method="monopolar_triangulation")
174+
175+
template_diff_thresh = [0.05, 0.15, 0.25]
176+
presets = ["x_contaminations"] * len(template_diff_thresh)
177+
steps_params = [
178+
{"template_similarity": {"template_diff_thresh": i}}
179+
for i in template_diff_thresh
180+
]
181+
182+
analyzer_merged = auto_merge_units(
183+
analyzer,
184+
presets=presets,
185+
steps_params=steps_params,
186+
recursive=True,
187+
**job_kwargs,
188+
)
189+
190+
The extra keyword ``recursive`` specifies that for each presets/sequences of steps, merges are performed
191+
until no further merges are possible. The ``job_kwargs`` are the parameters for the parallelization.
192+
**Be careful that the merges can not be reverted, so be sure to not erase your analyzer and create a new variable**
193+
156194

157195
Manual curation
158196
---------------

src/spikeinterface/benchmark/benchmark_base.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111

1212
from spikeinterface.core import SortingAnalyzer
13-
13+
from spikeinterface.core.job_tools import fix_job_kwargs, split_job_kwargs
1414
from spikeinterface import load, create_sorting_analyzer, load_sorting_analyzer
1515
from spikeinterface.widgets import get_some_colors
1616

@@ -219,7 +219,7 @@ def run(self, case_keys=None, keep=True, verbose=False, **job_kwargs):
219219
for key in job_keys:
220220
benchmark = self.create_benchmark(key)
221221
t0 = time.perf_counter()
222-
benchmark.run()
222+
benchmark.run(**job_kwargs)
223223
t1 = time.perf_counter()
224224
self.benchmarks[key] = benchmark
225225
bench_folder = self.folder / "results" / self.key_to_str(key)
@@ -264,6 +264,7 @@ def plot_run_times(self, case_keys=None):
264264
return plot_run_times(self, case_keys=case_keys)
265265

266266
def compute_results(self, case_keys=None, verbose=False, **result_params):
267+
267268
if case_keys is None:
268269
case_keys = list(self.cases.keys())
269270

@@ -309,7 +310,7 @@ def get_templates(self, key, operator="average"):
309310
templates = sorting_analyzer.get_extenson("templates").get_data(operator=operator)
310311
return templates
311312

312-
def compute_metrics(self, case_keys=None, metric_names=["snr", "firing_rate"], force=False):
313+
def compute_metrics(self, case_keys=None, metric_names=["snr", "firing_rate"], force=False, **job_kwargs):
313314
if case_keys is None:
314315
case_keys = self.cases.keys()
315316

@@ -329,7 +330,7 @@ def compute_metrics(self, case_keys=None, metric_names=["snr", "firing_rate"], f
329330
sorting_analyzer = self.get_sorting_analyzer(key)
330331
qm_ext = sorting_analyzer.get_extension("quality_metrics")
331332
if qm_ext is None or force:
332-
qm_ext = sorting_analyzer.compute("quality_metrics", metric_names=metric_names)
333+
qm_ext = sorting_analyzer.compute("quality_metrics", metric_names=metric_names, **job_kwargs)
333334

334335
# TODO remove this metics CSV file!!!!
335336
metrics = qm_ext.get_data()

src/spikeinterface/benchmark/benchmark_clustering.py

Lines changed: 12 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111

1212
import numpy as np
13-
13+
from spikeinterface.core.job_tools import fix_job_kwargs, split_job_kwargs
1414
from .benchmark_base import Benchmark, BenchmarkStudy
1515
from spikeinterface.core.sortinganalyzer import create_sorting_analyzer
1616
from spikeinterface.core.template_tools import get_template_extremum_channel
@@ -36,6 +36,8 @@ def run(self, **job_kwargs):
3636
self.result["peak_labels"] = peak_labels
3737

3838
def compute_result(self, **result_params):
39+
result_params, job_kwargs = split_job_kwargs(result_params)
40+
job_kwargs = fix_job_kwargs(job_kwargs)
3941
self.noise = self.result["peak_labels"] < 0
4042
spikes = self.gt_sorting.to_spike_vector()
4143
self.result["sliced_gt_sorting"] = NumpySorting(
@@ -47,8 +49,11 @@ def compute_result(self, **result_params):
4749
gt_unit_locations = self.gt_sorting.get_property("gt_unit_locations")
4850
if gt_unit_locations is None:
4951
print("'gt_unit_locations' is not a property of the sorting so compute it")
50-
gt_analyzer = create_sorting_analyzer(self.gt_sorting, self.recording, format="memory", sparse=True)
51-
gt_analyzer.compute(["random_spikes", "templates"])
52+
gt_analyzer = create_sorting_analyzer(
53+
self.gt_sorting, self.recording, format="memory", sparse=True, **job_kwargs
54+
)
55+
gt_analyzer.compute("random_spikes")
56+
gt_analyzer.compute("templates", **job_kwargs)
5257
ext = gt_analyzer.compute("unit_locations", method="monopolar_triangulation")
5358
gt_unit_locations = ext.get_data()
5459
self.gt_sorting.set_property("gt_unit_locations", gt_unit_locations)
@@ -64,17 +69,17 @@ def compute_result(self, **result_params):
6469
)
6570

6671
sorting_analyzer = create_sorting_analyzer(
67-
self.result["sliced_gt_sorting"], self.recording, format="memory", sparse=False
72+
self.result["sliced_gt_sorting"], self.recording, format="memory", sparse=False, **job_kwargs
6873
)
6974
sorting_analyzer.compute("random_spikes")
70-
ext = sorting_analyzer.compute("templates")
75+
ext = sorting_analyzer.compute("templates", **job_kwargs)
7176
self.result["sliced_gt_templates"] = ext.get_data(outputs="Templates")
7277

7378
sorting_analyzer = create_sorting_analyzer(
74-
self.result["clustering"], self.recording, format="memory", sparse=False
79+
self.result["clustering"], self.recording, format="memory", sparse=False, **job_kwargs
7580
)
7681
sorting_analyzer.compute("random_spikes")
77-
ext = sorting_analyzer.compute("templates")
82+
ext = sorting_analyzer.compute("templates", **job_kwargs)
7883
self.result["clustering_templates"] = ext.get_data(outputs="Templates")
7984

8085
_run_key_saved = [("peak_labels", "npy")]

src/spikeinterface/benchmark/benchmark_matching.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -133,10 +133,10 @@ def get_count_units(self, case_keys=None, well_detected_score=None, redundant_sc
133133

134134
return count_units
135135

136-
def plot_unit_counts(self, case_keys=None, figsize=None):
137-
from spikeinterface.widgets.widget_list import plot_study_unit_counts
136+
def plot_unit_counts(self, case_keys=None, **kwargs):
137+
from .benchmark_plot_tools import plot_unit_counts
138138

139-
plot_study_unit_counts(self, case_keys, figsize=figsize)
139+
return plot_unit_counts(self, case_keys, **kwargs)
140140

141141
def plot_unit_losses(self, before, after, metric=["accuracy"], figsize=None):
142142
import matplotlib.pyplot as plt
Lines changed: 173 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,173 @@
1+
from __future__ import annotations
2+
3+
from spikeinterface.curation.auto_merge import auto_merge_units
4+
from spikeinterface.comparison import compare_sorter_to_ground_truth
5+
from spikeinterface.core.sortinganalyzer import create_sorting_analyzer
6+
from spikeinterface.widgets import (
7+
plot_unit_templates,
8+
plot_amplitudes,
9+
plot_crosscorrelograms,
10+
)
11+
12+
import numpy as np
13+
from .benchmark_base import Benchmark, BenchmarkStudy
14+
15+
16+
class MergingBenchmark(Benchmark):
17+
18+
def __init__(self, recording, splitted_sorting, params, gt_sorting, splitted_cells=None):
19+
self.recording = recording
20+
self.splitted_sorting = splitted_sorting
21+
self.gt_sorting = gt_sorting
22+
self.splitted_cells = splitted_cells
23+
self.method_kwargs = params["method_kwargs"]
24+
self.result = {}
25+
26+
def run(self, **job_kwargs):
27+
sorting_analyzer = create_sorting_analyzer(
28+
self.splitted_sorting, self.recording, format="memory", sparse=True, **job_kwargs
29+
)
30+
# sorting_analyzer.compute(['random_spikes', 'templates'])
31+
# sorting_analyzer.compute('template_similarity', max_lag_ms=0.1, method="l2", **job_kwargs)
32+
merged_analyzer, self.result["merged_pairs"], self.result["merges"], self.result["outs"] = auto_merge_units(
33+
sorting_analyzer, extra_outputs=True, **self.method_kwargs, **job_kwargs
34+
)
35+
36+
self.result["sorting"] = merged_analyzer.sorting
37+
38+
def compute_result(self, **result_params):
39+
sorting = self.result["sorting"]
40+
comp = compare_sorter_to_ground_truth(self.gt_sorting, sorting, exhaustive_gt=True)
41+
self.result["gt_comparison"] = comp
42+
43+
_run_key_saved = [("sorting", "sorting"), ("merges", "pickle"), ("merged_pairs", "pickle"), ("outs", "pickle")]
44+
_result_key_saved = [("gt_comparison", "pickle")]
45+
46+
47+
class MergingStudy(BenchmarkStudy):
48+
49+
benchmark_class = MergingBenchmark
50+
51+
def create_benchmark(self, key):
52+
dataset_key = self.cases[key]["dataset"]
53+
recording, gt_sorting = self.datasets[dataset_key]
54+
params = self.cases[key]["params"]
55+
init_kwargs = self.cases[key]["init_kwargs"]
56+
benchmark = MergingBenchmark(recording, gt_sorting, params, **init_kwargs)
57+
return benchmark
58+
59+
def get_count_units(self, case_keys=None, well_detected_score=None, redundant_score=None, overmerged_score=None):
60+
import pandas as pd
61+
62+
if case_keys is None:
63+
case_keys = list(self.cases.keys())
64+
65+
if isinstance(case_keys[0], str):
66+
index = pd.Index(case_keys, name=self.levels)
67+
else:
68+
index = pd.MultiIndex.from_tuples(case_keys, names=self.levels)
69+
70+
columns = ["num_gt", "num_sorter", "num_well_detected"]
71+
comp = self.get_result(case_keys[0])["gt_comparison"]
72+
if comp.exhaustive_gt:
73+
columns.extend(["num_false_positive", "num_redundant", "num_overmerged", "num_bad"])
74+
count_units = pd.DataFrame(index=index, columns=columns, dtype=int)
75+
76+
for key in case_keys:
77+
comp = self.get_result(key)["gt_comparison"]
78+
assert comp is not None, "You need to do study.run_comparisons() first"
79+
80+
gt_sorting = comp.sorting1
81+
sorting = comp.sorting2
82+
83+
count_units.loc[key, "num_gt"] = len(gt_sorting.get_unit_ids())
84+
count_units.loc[key, "num_sorter"] = len(sorting.get_unit_ids())
85+
count_units.loc[key, "num_well_detected"] = comp.count_well_detected_units(well_detected_score)
86+
87+
if comp.exhaustive_gt:
88+
count_units.loc[key, "num_redundant"] = comp.count_redundant_units(redundant_score)
89+
count_units.loc[key, "num_overmerged"] = comp.count_overmerged_units(overmerged_score)
90+
count_units.loc[key, "num_false_positive"] = comp.count_false_positive_units(redundant_score)
91+
count_units.loc[key, "num_bad"] = comp.count_bad_units()
92+
93+
return count_units
94+
95+
def plot_agreement_matrix(self, **kwargs):
96+
from .benchmark_plot_tools import plot_agreement_matrix
97+
98+
return plot_agreement_matrix(self, **kwargs)
99+
100+
def plot_unit_counts(self, case_keys=None, **kwargs):
101+
from .benchmark_plot_tools import plot_unit_counts
102+
103+
return plot_unit_counts(self, case_keys, **kwargs)
104+
105+
def get_splitted_pairs(self, case_key):
106+
return self.benchmarks[case_key].splitted_cells
107+
108+
def get_splitted_pairs_index(self, case_key, pair):
109+
for count, i in enumerate(self.benchmarks[case_key].splitted_cells):
110+
if i == pair:
111+
return count
112+
113+
def plot_splitted_amplitudes(self, case_key, pair_index=0, backend="ipywidgets"):
114+
analyzer = self.get_sorting_analyzer(case_key)
115+
if analyzer.get_extension("spike_amplitudes") is None:
116+
analyzer.compute(["spike_amplitudes"])
117+
plot_amplitudes(analyzer, unit_ids=self.get_splitted_pairs(case_key)[pair_index], backend=backend)
118+
119+
def plot_splitted_correlograms(self, case_key, pair_index=0, backend="ipywidgets"):
120+
analyzer = self.get_sorting_analyzer(case_key)
121+
if analyzer.get_extension("correlograms") is None:
122+
analyzer.compute(["correlograms"])
123+
if analyzer.get_extension("template_similarity") is None:
124+
analyzer.compute(["template_similarity"])
125+
plot_crosscorrelograms(analyzer, unit_ids=self.get_splitted_pairs(case_key)[pair_index])
126+
127+
def plot_splitted_templates(self, case_key, pair_index=0, backend="ipywidgets"):
128+
analyzer = self.get_sorting_analyzer(case_key)
129+
if analyzer.get_extension("spike_amplitudes") is None:
130+
analyzer.compute(["spike_amplitudes"])
131+
plot_unit_templates(analyzer, unit_ids=self.get_splitted_pairs(case_key)[pair_index], backend=backend)
132+
133+
def plot_potential_merges(self, case_key, min_snr=None, backend="ipywidgets"):
134+
analyzer = self.get_sorting_analyzer(case_key)
135+
mylist = self.get_splitted_pairs(case_key)
136+
137+
if analyzer.get_extension("spike_amplitudes") is None:
138+
analyzer.compute(["spike_amplitudes"])
139+
if analyzer.get_extension("correlograms") is None:
140+
analyzer.compute(["correlograms"])
141+
142+
if min_snr is not None:
143+
select_from = analyzer.sorting.unit_ids
144+
if analyzer.get_extension("noise_levels") is None:
145+
analyzer.compute("noise_levels")
146+
if analyzer.get_extension("quality_metrics") is None:
147+
analyzer.compute("quality_metrics", metric_names=["snr"])
148+
149+
snr = analyzer.get_extension("quality_metrics").get_data()["snr"].values
150+
select_from = select_from[snr > min_snr]
151+
mylist_selection = []
152+
for i in mylist:
153+
if (i[0] in select_from) or (i[1] in select_from):
154+
mylist_selection += [i]
155+
mylist = mylist_selection
156+
157+
from spikeinterface.widgets import plot_potential_merges
158+
159+
plot_potential_merges(analyzer, mylist, backend=backend)
160+
161+
def plot_performed_merges(self, case_key, backend="ipywidgets"):
162+
analyzer = self.get_sorting_analyzer(case_key)
163+
164+
if analyzer.get_extension("spike_amplitudes") is None:
165+
analyzer.compute(["spike_amplitudes"])
166+
if analyzer.get_extension("correlograms") is None:
167+
analyzer.compute(["correlograms"])
168+
169+
all_merges = list(self.benchmarks[case_key].result["merged_pairs"].values())
170+
171+
from spikeinterface.widgets import plot_potential_merges
172+
173+
plot_potential_merges(analyzer, all_merges, backend=backend)

src/spikeinterface/benchmark/benchmark_motion_estimation.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,7 +86,7 @@ def __init__(
8686
def run(self, **job_kwargs):
8787
p = self.params
8888

89-
noise_levels = get_noise_levels(self.recording, return_scaled=False)
89+
noise_levels = get_noise_levels(self.recording, return_scaled=False, **job_kwargs)
9090

9191
t0 = time.perf_counter()
9292
peaks = detect_peaks(self.recording, noise_levels=noise_levels, **p["detect_kwargs"], **job_kwargs)

0 commit comments

Comments
 (0)