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+
0 commit comments