Skip to content

Commit 3561531

Browse files
authored
Merge branch 'main' into fix_numpy_2.0_representation
2 parents af51937 + a1dc9d7 commit 3561531

26 files changed

Lines changed: 991 additions & 524 deletions

src/spikeinterface/benchmark/benchmark_base.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -366,6 +366,39 @@ def get_units_snr(self, key):
366366
def get_result(self, key):
367367
return self.benchmarks[key].result
368368

369+
def get_pairs_by_level(self, level):
370+
"""
371+
usefull for function like plot_performance_losses() where you need to plot one pair of results
372+
This generate list of pairs for a given level.
373+
"""
374+
375+
level_index = self.levels.index(level)
376+
377+
possible_values = []
378+
for key in self.cases.keys():
379+
assert isinstance(key, tuple), "get_pairs_by_level need tuple keys"
380+
level_value = key[level_index]
381+
if level_value not in possible_values:
382+
possible_values.append(level_value)
383+
assert len(possible_values) == 2, "get_pairs_by_level() : you need exactly 2 value for this levels"
384+
385+
pairs = []
386+
for key in self.cases.keys():
387+
388+
case0 = list(key)
389+
case1 = list(key)
390+
case0[level_index] = possible_values[0]
391+
case1[level_index] = possible_values[1]
392+
case0 = tuple(case0)
393+
case1 = tuple(case1)
394+
395+
pair = (case0, case1)
396+
397+
if pair not in pairs:
398+
pairs.append(pair)
399+
400+
return pairs
401+
369402

370403
class Benchmark:
371404
"""

src/spikeinterface/benchmark/benchmark_clustering.py

Lines changed: 15 additions & 98 deletions
Original file line numberDiff line numberDiff line change
@@ -181,6 +181,21 @@ def plot_performances_vs_snr(self, **kwargs):
181181

182182
return plot_performances_vs_snr(self, **kwargs)
183183

184+
def plot_performances_comparison(self, *args, **kwargs):
185+
from .benchmark_plot_tools import plot_performances_comparison
186+
187+
return plot_performances_comparison(self, *args, **kwargs)
188+
189+
def plot_performance_losses(self, *args, **kwargs):
190+
from .benchmark_plot_tools import plot_performance_losses
191+
192+
return plot_performance_losses(self, *args, **kwargs)
193+
194+
def plot_performances_vs_depth_and_snr(self, *args, **kwargs):
195+
from .benchmark_plot_tools import plot_performances_vs_depth_and_snr
196+
197+
return plot_performances_vs_depth_and_snr(self, *args, **kwargs)
198+
184199
def plot_error_metrics(self, metric="cosine", case_keys=None, figsize=(15, 5)):
185200

186201
if case_keys is None:
@@ -351,104 +366,6 @@ def plot_metrics_vs_depth_and_snr(self, metric="agreement", case_keys=None, figs
351366

352367
return fig
353368

354-
def plot_unit_losses(self, cases_before, cases_after, metric="agreement", figsize=None):
355-
356-
fig, axs = plt.subplots(ncols=len(cases_before), nrows=1, figsize=figsize)
357-
358-
for count, (case_before, case_after) in enumerate(zip(cases_before, cases_after)):
359-
360-
ax = axs[count]
361-
dataset_key = self.cases[case_before]["dataset"]
362-
_, gt_sorting1 = self.datasets[dataset_key]
363-
positions = gt_sorting1.get_property("gt_unit_locations")
364-
365-
analyzer = self.get_sorting_analyzer(case_before)
366-
metrics_before = analyzer.get_extension("quality_metrics").get_data()
367-
x = metrics_before["snr"].values
368-
369-
y_before = self.get_result(case_before)["gt_comparison"].get_performance()[metric].values
370-
y_after = self.get_result(case_after)["gt_comparison"].get_performance()[metric].values
371-
ax.set_ylabel("depth (um)")
372-
ax.set_ylabel("snr")
373-
if count > 0:
374-
ax.set_ylabel("")
375-
ax.set_yticks([], [])
376-
im = ax.scatter(positions[:, 1], x, c=(y_after - y_before), cmap="coolwarm")
377-
im.set_clim(-1, 1)
378-
# fig.colorbar(im, ax=ax)
379-
# ax.set_title(k)
380-
381-
fig.subplots_adjust(right=0.85)
382-
cbar_ax = fig.add_axes([0.9, 0.1, 0.025, 0.75])
383-
cbar = fig.colorbar(im, cax=cbar_ax, label=metric)
384-
# cbar.set_clim(-1, 1)
385-
386-
return fig
387-
388-
def plot_comparison_clustering(
389-
self,
390-
case_keys=None,
391-
performance_names=["accuracy", "recall", "precision"],
392-
colors=["g", "b", "r"],
393-
ylim=(-0.1, 1.1),
394-
figsize=None,
395-
):
396-
397-
if case_keys is None:
398-
case_keys = list(self.cases.keys())
399-
import pylab as plt
400-
401-
num_methods = len(case_keys)
402-
fig, axs = plt.subplots(ncols=num_methods, nrows=num_methods, figsize=(10, 10))
403-
for i, key1 in enumerate(case_keys):
404-
for j, key2 in enumerate(case_keys):
405-
if len(axs.shape) > 1:
406-
ax = axs[i, j]
407-
else:
408-
ax = axs[j]
409-
comp1 = self.get_result(key1)["gt_comparison"]
410-
comp2 = self.get_result(key2)["gt_comparison"]
411-
if i <= j:
412-
for performance, color in zip(performance_names, colors):
413-
perf1 = comp1.get_performance()[performance]
414-
perf2 = comp2.get_performance()[performance]
415-
ax.plot(perf2, perf1, ".", label=performance, color=color)
416-
417-
ax.plot([0, 1], [0, 1], "k--", alpha=0.5)
418-
ax.set_ylim(ylim)
419-
ax.set_xlim(ylim)
420-
ax.spines[["right", "top"]].set_visible(False)
421-
ax.set_aspect("equal")
422-
423-
label1 = self.cases[key1]["label"]
424-
label2 = self.cases[key2]["label"]
425-
if j == i:
426-
ax.set_ylabel(f"{label1}")
427-
else:
428-
ax.set_yticks([])
429-
if i == j:
430-
ax.set_xlabel(f"{label2}")
431-
else:
432-
ax.set_xticks([])
433-
if i == num_methods - 1 and j == num_methods - 1:
434-
patches = []
435-
import matplotlib.patches as mpatches
436-
437-
for color, name in zip(colors, performance_names):
438-
patches.append(mpatches.Patch(color=color, label=name))
439-
ax.legend(handles=patches, bbox_to_anchor=(1.05, 1), loc="upper left", borderaxespad=0.0)
440-
else:
441-
ax.spines["bottom"].set_visible(False)
442-
ax.spines["left"].set_visible(False)
443-
ax.spines["top"].set_visible(False)
444-
ax.spines["right"].set_visible(False)
445-
ax.set_xticks([])
446-
ax.set_yticks([])
447-
448-
plt.tight_layout(h_pad=0, w_pad=0)
449-
450-
return fig
451-
452369
def plot_some_over_merged(self, case_keys=None, overmerged_score=0.05, max_units=5, figsize=None):
453370
if case_keys is None:
454371
case_keys = list(self.cases.keys())

src/spikeinterface/benchmark/benchmark_matching.py

Lines changed: 14 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
from __future__ import annotations
22

3+
import warnings
4+
35
from spikeinterface.sortingcomponents.matching import find_spikes_from_templates
46
from spikeinterface.core import NumpySorting
57
from spikeinterface.comparison import CollisionGTComparison, compare_sorter_to_ground_truth
@@ -77,6 +79,11 @@ def plot_performances_comparison(self, **kwargs):
7779

7880
return plot_performances_comparison(self, **kwargs)
7981

82+
def plot_performances_vs_depth_and_snr(self, *args, **kwargs):
83+
from .benchmark_plot_tools import plot_performances_vs_depth_and_snr
84+
85+
return plot_performances_vs_depth_and_snr(self, *args, **kwargs)
86+
8087
def plot_collisions(self, case_keys=None, figsize=None):
8188
if case_keys is None:
8289
case_keys = list(self.cases.keys())
@@ -138,39 +145,13 @@ def plot_unit_counts(self, case_keys=None, **kwargs):
138145

139146
return plot_unit_counts(self, case_keys, **kwargs)
140147

141-
def plot_unit_losses(self, before, after, metric=["accuracy"], figsize=None):
142-
import matplotlib.pyplot as plt
143-
144-
fig, axs = plt.subplots(ncols=1, nrows=len(metric), figsize=figsize, squeeze=False)
145-
146-
for count, k in enumerate(metric):
148+
def plot_unit_losses(self, *args, **kwargs):
149+
from .benchmark_plot_tools import plot_performance_losses
147150

148-
ax = axs[0, count]
151+
warnings.warn("plot_unit_losses() is now plot_performance_losses()")
152+
return plot_performance_losses(self, *args, **kwargs)
149153

150-
label = self.cases[after]["label"]
154+
def plot_performance_losses(self, *args, **kwargs):
155+
from .benchmark_plot_tools import plot_performance_losses
151156

152-
positions = self.get_result(before)["gt_comparison"].sorting1.get_property("gt_unit_locations")
153-
154-
analyzer = self.get_sorting_analyzer(before)
155-
metrics_before = analyzer.get_extension("quality_metrics").get_data()
156-
x = metrics_before["snr"].values
157-
158-
y_before = self.get_result(before)["gt_comparison"].get_performance()[k].values
159-
y_after = self.get_result(after)["gt_comparison"].get_performance()[k].values
160-
# if count < 2:
161-
# ax.set_xticks([], [])
162-
# elif count == 2:
163-
ax.set_xlabel("depth (um)")
164-
im = ax.scatter(positions[:, 1], x, c=(y_after - y_before), cmap="coolwarm")
165-
fig.colorbar(im, ax=ax, label=k)
166-
im.set_clim(-1, 1)
167-
ax.set_title(k)
168-
ax.set_ylabel("snr")
169-
170-
# fig.subplots_adjust(right=0.85)
171-
# cbar_ax = fig.add_axes([0.9, 0.1, 0.025, 0.75])
172-
# cbar = fig.colorbar(im, cax=cbar_ax, label=metric)
173-
174-
# if count == 2:
175-
# ax.legend()
176-
return fig
157+
return plot_performance_losses(self, *args, **kwargs)

0 commit comments

Comments
 (0)