Skip to content

Commit 68a8691

Browse files
committed
Common function between benchlark_clustering/benchmark_matching/benchmark_sorter
1 parent de2904b commit 68a8691

4 files changed

Lines changed: 59 additions & 74 deletions

File tree

src/spikeinterface/benchmark/benchmark_clustering.py

Lines changed: 10 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -160,49 +160,19 @@ def get_count_units(self, case_keys=None, well_detected_score=None, redundant_sc
160160

161161
return count_units
162162

163-
def plot_unit_counts(self, case_keys=None, figsize=None, **extra_kwargs):
164-
from spikeinterface.widgets.widget_list import plot_study_unit_counts
163+
# plotting by methods
164+
def plot_unit_counts(self, **kwargs):
165+
from .benchmark_plot_tools import plot_unit_counts
166+
return plot_unit_counts(self, **kwargs)
165167

166-
plot_study_unit_counts(self, case_keys, figsize=figsize, **extra_kwargs)
168+
def plot_agreement_matrix(self, **kwargs):
169+
from .benchmark_plot_tools import plot_agreement_matrix
170+
return plot_agreement_matrix(self, **kwargs)
167171

168-
def plot_agreements(self, case_keys=None, figsize=(15, 15)):
169-
if case_keys is None:
170-
case_keys = list(self.cases.keys())
171-
import pylab as plt
172-
173-
fig, axs = plt.subplots(ncols=len(case_keys), nrows=1, figsize=figsize, squeeze=False)
174-
175-
for count, key in enumerate(case_keys):
176-
ax = axs[0, count]
177-
ax.set_title(self.cases[key]["label"])
178-
plot_agreement_matrix(self.get_result(key)["gt_comparison"], ax=ax)
179-
180-
return fig
181-
182-
def plot_performances_vs_snr(self, case_keys=None, figsize=(15, 15)):
183-
if case_keys is None:
184-
case_keys = list(self.cases.keys())
185-
import pylab as plt
186-
187-
fig, axes = plt.subplots(ncols=1, nrows=3, figsize=figsize)
188-
189-
for count, k in enumerate(("accuracy", "recall", "precision")):
172+
def plot_performances_vs_snr(self, **kwargs):
173+
from .benchmark_plot_tools import plot_performances_vs_snr
174+
return plot_performances_vs_snr(self, **kwargs)
190175

191-
ax = axes[count]
192-
for key in case_keys:
193-
label = self.cases[key]["label"]
194-
195-
analyzer = self.get_sorting_analyzer(key)
196-
metrics = analyzer.get_extension("quality_metrics").get_data()
197-
x = metrics["snr"].values
198-
y = self.get_result(key)["gt_comparison"].get_performance()[k].values
199-
ax.scatter(x, y, marker=".", label=label)
200-
ax.set_title(k)
201-
202-
if count == 2:
203-
ax.legend()
204-
205-
return fig
206176

207177
def plot_error_metrics(self, metric="cosine", case_keys=None, figsize=(15, 5)):
208178

src/spikeinterface/benchmark/benchmark_matching.py

Lines changed: 6 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -61,42 +61,14 @@ def create_benchmark(self, key):
6161
benchmark = MatchingBenchmark(recording, gt_sorting, params)
6262
return benchmark
6363

64-
def plot_agreements(self, case_keys=None, figsize=None):
65-
if case_keys is None:
66-
case_keys = list(self.cases.keys())
67-
import pylab as plt
68-
69-
fig, axs = plt.subplots(ncols=len(case_keys), nrows=1, figsize=figsize, squeeze=False)
70-
71-
for count, key in enumerate(case_keys):
72-
ax = axs[0, count]
73-
ax.set_title(self.cases[key]["label"])
74-
plot_agreement_matrix(self.get_result(key)["gt_comparison"], ax=ax)
64+
def plot_agreement_matrix(self, **kwargs):
65+
from .benchmark_plot_tools import plot_agreement_matrix
66+
return plot_agreement_matrix(self, **kwargs)
7567

76-
def plot_performances_vs_snr(self, case_keys=None, figsize=None, metrics=["accuracy", "recall", "precision"]):
77-
if case_keys is None:
78-
case_keys = list(self.cases.keys())
79-
80-
import matplotlib.pyplot as plt
81-
fig, axs = plt.subplots(ncols=1, nrows=len(metrics), figsize=figsize, squeeze=False)
82-
83-
for count, k in enumerate(metrics):
68+
def plot_performances_vs_snr(self, **kwargs):
69+
from .benchmark_plot_tools import plot_performances_vs_snr
70+
return plot_performances_vs_snr(self, **kwargs)
8471

85-
ax = axs[count, 0]
86-
for key in case_keys:
87-
label = self.cases[key]["label"]
88-
89-
analyzer = self.get_sorting_analyzer(key)
90-
metrics = analyzer.get_extension("quality_metrics").get_data()
91-
x = metrics["snr"].values
92-
y = self.get_result(key)["gt_comparison"].get_performance()[k].values
93-
ax.scatter(x, y, marker=".", label=label)
94-
ax.set_title(k)
95-
96-
if count == 2:
97-
ax.legend()
98-
99-
return fig
10072

10173
def plot_collisions(self, case_keys=None, figsize=None):
10274
if case_keys is None:

src/spikeinterface/benchmark/benchmark_plot_tools.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -217,3 +217,31 @@ def plot_agreement_matrix(study, ordered=True, case_keys=None):
217217
ax.set_yticks([])
218218
ax.set_xticks([])
219219

220+
221+
def plot_performances_vs_snr(study, case_keys=None, figsize=None, metrics=["accuracy", "recall", "precision"]):
222+
import matplotlib.pyplot as plt
223+
224+
if case_keys is None:
225+
case_keys = list(study.cases.keys())
226+
227+
fig, axs = plt.subplots(ncols=1, nrows=len(metrics), figsize=figsize, squeeze=False)
228+
229+
for count, k in enumerate(metrics):
230+
231+
ax = axs[count, 0]
232+
for key in case_keys:
233+
label = study.cases[key]["label"]
234+
235+
analyzer = study.get_sorting_analyzer(key)
236+
metrics = analyzer.get_extension("quality_metrics").get_data()
237+
x = metrics["snr"].values
238+
y = study.get_result(key)["gt_comparison"].get_performance()[k].values
239+
ax.scatter(x, y, marker=".", label=label)
240+
ax.set_title(k)
241+
242+
ax.set_ylim(0, 1.05)
243+
244+
if count == 2:
245+
ax.legend()
246+
247+
return fig

src/spikeinterface/benchmark/benchmark_sorter.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -123,3 +123,18 @@ def get_count_units(self, case_keys=None, well_detected_score=None, redundant_sc
123123

124124
return count_units
125125

126+
# plotting as methods
127+
def plot_unit_counts(self, **kwargs):
128+
from .benchmark_plot_tools import plot_unit_counts
129+
return plot_unit_counts(self, **kwargs)
130+
131+
def plot_performances(self, **kwargs):
132+
from .benchmark_plot_tools import plot_performances
133+
return plot_performances(self, **kwargs)
134+
135+
def plot_agreement_matrix(self, **kwargs):
136+
from .benchmark_plot_tools import plot_agreement_matrix
137+
return plot_agreement_matrix(self, **kwargs)
138+
139+
140+

0 commit comments

Comments
 (0)