Skip to content

Commit de2904b

Browse files
committed
Refactor plotting for GTSTudy
1 parent d599b0b commit de2904b

4 files changed

Lines changed: 262 additions & 224 deletions

File tree

src/spikeinterface/benchmark/benchmark_base.py

Lines changed: 2 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -258,25 +258,8 @@ def get_run_times(self, case_keys=None):
258258
return df
259259

260260
def plot_run_times(self, case_keys=None):
261-
if case_keys is None:
262-
case_keys = list(self.cases.keys())
263-
run_times = self.get_run_times(case_keys=case_keys)
264-
265-
colors = self.get_colors()
266-
import matplotlib.pyplot as plt
267-
268-
fig, ax = plt.subplots()
269-
labels = []
270-
for i, key in enumerate(case_keys):
271-
labels.append(self.cases[key]["label"])
272-
rt = run_times.at[key, "run_times"]
273-
ax.bar(i, rt, width=0.8, color=colors[key])
274-
ax.set_xticks(np.arange(len(case_keys)))
275-
ax.set_xticklabels(labels, rotation=45.0)
276-
return fig
277-
278-
# ax = run_times.plot(kind="bar")
279-
# return ax.figure
261+
from .benchmark_plot_tools import plot_run_times
262+
return plot_run_times(self, case_keys=case_keys)
280263

281264
def compute_results(self, case_keys=None, verbose=False, **result_params):
282265
if case_keys is None:
Lines changed: 211 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
1+
import numpy as np
22

33

44

@@ -7,3 +7,213 @@ def _simpleaxis(ax):
77
ax.spines["right"].set_visible(False)
88
ax.get_xaxis().tick_bottom()
99
ax.get_yaxis().tick_left()
10+
11+
12+
def plot_run_times(study, case_keys=None):
13+
"""
14+
Plot run times for a BenchmarkStudy.
15+
16+
Parameters
17+
----------
18+
study : SorterStudy
19+
A study object.
20+
case_keys : list or None
21+
A selection of cases to plot, if None, then all.
22+
"""
23+
import matplotlib.pyplot as plt
24+
25+
if case_keys is None:
26+
case_keys = list(study.cases.keys())
27+
28+
run_times = study.get_run_times(case_keys=case_keys)
29+
30+
colors = study.get_colors()
31+
32+
33+
fig, ax = plt.subplots()
34+
labels = []
35+
for i, key in enumerate(case_keys):
36+
labels.append(study.cases[key]["label"])
37+
rt = run_times.at[key, "run_times"]
38+
ax.bar(i, rt, width=0.8, color=colors[key])
39+
ax.set_xticks(np.arange(len(case_keys)))
40+
ax.set_xticklabels(labels, rotation=45.0)
41+
return fig
42+
43+
44+
def plot_unit_counts(study, case_keys=None):
45+
"""
46+
Plot unit counts for a study: "num_well_detected", "num_false_positive", "num_redundant", "num_overmerged"
47+
48+
Parameters
49+
----------
50+
study : SorterStudy
51+
A study object.
52+
case_keys : list or None
53+
A selection of cases to plot, if None, then all.
54+
"""
55+
import matplotlib.pyplot as plt
56+
from spikeinterface.widgets.utils import get_some_colors
57+
58+
if case_keys is None:
59+
case_keys = list(study.cases.keys())
60+
61+
62+
count_units = study.get_count_units(case_keys=case_keys)
63+
64+
fig, ax = plt.subplots()
65+
66+
columns = count_units.columns.tolist()
67+
columns.remove("num_gt")
68+
columns.remove("num_sorter")
69+
70+
ncol = len(columns)
71+
72+
colors = get_some_colors(columns, color_engine="auto", map_name="hot")
73+
colors["num_well_detected"] = "green"
74+
75+
xticklabels = []
76+
for i, key in enumerate(case_keys):
77+
for c, col in enumerate(columns):
78+
x = i + 1 + c / (ncol + 1)
79+
y = count_units.loc[key, col]
80+
if not "well_detected" in col:
81+
y = -y
82+
83+
if i == 0:
84+
label = col.replace("num_", "").replace("_", " ").title()
85+
else:
86+
label = None
87+
88+
ax.bar([x], [y], width=1 / (ncol + 2), label=label, color=colors[col])
89+
90+
xticklabels.append(study.cases[key]["label"])
91+
92+
ax.set_xticks(np.arange(len(case_keys)) + 1)
93+
ax.set_xticklabels(xticklabels)
94+
ax.legend()
95+
96+
return fig
97+
98+
def plot_performances(study, mode="ordered", performance_names=("accuracy", "precision", "recall"), case_keys=None):
99+
"""
100+
Plot performances over case for a study.
101+
102+
Parameters
103+
----------
104+
study : GroundTruthStudy
105+
A study object.
106+
mode : "ordered" | "snr" | "swarm", default: "ordered"
107+
Which plot mode to use:
108+
109+
* "ordered": plot performance metrics vs unit indices ordered by decreasing accuracy
110+
* "snr": plot performance metrics vs snr
111+
* "swarm": plot performance metrics as a swarm plot (see seaborn.swarmplot for details)
112+
performance_names : list or tuple, default: ("accuracy", "precision", "recall")
113+
Which performances to plot ("accuracy", "precision", "recall")
114+
case_keys : list or None
115+
A selection of cases to plot, if None, then all.
116+
"""
117+
import matplotlib.pyplot as plt
118+
import pandas as pd
119+
import seaborn as sns
120+
121+
if case_keys is None:
122+
case_keys = list(study.cases.keys())
123+
124+
perfs=study.get_performance_by_unit(case_keys=case_keys)
125+
colors = study.get_colors()
126+
127+
128+
if mode in ("ordered", "snr"):
129+
num_axes = len(performance_names)
130+
fig, axs = plt.subplots(ncols=num_axes)
131+
else:
132+
fig, ax = plt.subplots()
133+
134+
if mode == "ordered":
135+
for count, performance_name in enumerate(performance_names):
136+
ax = axs.flatten()[count]
137+
for key in case_keys:
138+
label = study.cases[key]["label"]
139+
val = perfs.xs(key).loc[:, performance_name].values
140+
val = np.sort(val)[::-1]
141+
ax.plot(val, label=label, c=colors[key])
142+
ax.set_title(performance_name)
143+
if count == len(performance_names) - 1:
144+
ax.legend(bbox_to_anchor=(0.05, 0.05), loc="lower left", framealpha=0.8)
145+
146+
elif mode == "snr":
147+
metric_name = mode
148+
for count, performance_name in enumerate(performance_names):
149+
ax = axs.flatten()[count]
150+
151+
max_metric = 0
152+
for key in case_keys:
153+
x = study.get_metrics(key).loc[:, metric_name].values
154+
y = perfs.xs(key).loc[:, performance_name].values
155+
label = study.cases[key]["label"]
156+
ax.scatter(x, y, s=10, label=label, color=colors[key])
157+
max_metric = max(max_metric, np.max(x))
158+
ax.set_title(performance_name)
159+
ax.set_xlim(0, max_metric * 1.05)
160+
ax.set_ylim(0, 1.05)
161+
if count == 0:
162+
ax.legend(loc="lower right")
163+
164+
elif mode == "swarm":
165+
levels = perfs.index.names
166+
df = pd.melt(
167+
perfs.reset_index(),
168+
id_vars=levels,
169+
var_name="Metric",
170+
value_name="Score",
171+
value_vars=performance_names,
172+
)
173+
df["x"] = df.apply(lambda r: " ".join([r[col] for col in levels]), axis=1)
174+
sns.swarmplot(data=df, x="x", y="Score", hue="Metric", dodge=True, ax=ax)
175+
176+
177+
def plot_agreement_matrix(study, ordered=True, case_keys=None):
178+
"""
179+
Plot agreement matri ces for cases in a study.
180+
181+
Parameters
182+
----------
183+
study : GroundTruthStudy
184+
A study object.
185+
case_keys : list or None
186+
A selection of cases to plot, if None, then all.
187+
ordered : bool
188+
Order units with best agreement scores.
189+
This enable to see agreement on a diagonal.
190+
"""
191+
192+
import matplotlib.pyplot as plt
193+
from spikeinterface.widgets import AgreementMatrixWidget
194+
195+
if case_keys is None:
196+
case_keys = list(study.cases.keys())
197+
198+
199+
num_axes = len(case_keys)
200+
fig, axs = plt.subplots(ncols=num_axes)
201+
202+
for count, key in enumerate(case_keys):
203+
ax = axs.flatten()[count]
204+
comp = study.get_result(key)["gt_comparison"]
205+
206+
unit_ticks = len(comp.sorting1.unit_ids) <= 16
207+
count_text = len(comp.sorting1.unit_ids) <= 16
208+
209+
AgreementMatrixWidget(
210+
comp, ordered=ordered, count_text=count_text, unit_ticks=unit_ticks, backend="matplotlib", ax=ax
211+
)
212+
label = study.cases[key]["label"]
213+
ax.set_xlabel(label)
214+
215+
if count > 0:
216+
ax.set_ylabel(None)
217+
ax.set_yticks([])
218+
ax.set_xticks([])
219+

src/spikeinterface/benchmark/tests/test_benchmark_sorter.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -63,10 +63,10 @@ def test_SorterStudy(setup_module):
6363
print(study)
6464

6565
# # this run the sorters
66-
study.run()
66+
# study.run()
6767

6868
# # this run comparisons
69-
study.compute_results()
69+
# study.compute_results()
7070
print(study)
7171

7272
# this is from the base class

0 commit comments

Comments
 (0)