Skip to content

Commit f0cd3cb

Browse files
authored
Merge branch 'main' into support_python_3.13
2 parents 97ee37f + cabe66e commit f0cd3cb

4 files changed

Lines changed: 207 additions & 108 deletions

File tree

src/spikeinterface/widgets/sorting_summary.py

Lines changed: 57 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@
22

33
import numpy as np
44

5+
import warnings
6+
57
from .base import BaseWidget, to_attr
68

79
from .amplitudes import AmplitudesWidget
@@ -14,6 +16,9 @@
1416
from ..core import SortingAnalyzer
1517

1618

19+
_default_displayed_unit_properties = ["firing_rate", "num_spikes", "x", "y", "amplitude_median", "snr", "rp_violation"]
20+
21+
1722
class SortingSummaryWidget(BaseWidget):
1823
"""
1924
Plots spike sorting summary.
@@ -42,14 +47,24 @@ class SortingSummaryWidget(BaseWidget):
4247
label_choices : list or None, default: None
4348
List of labels to be added to the curation table
4449
(sortingview backend)
45-
unit_table_properties : list or None, default: None
50+
displayed_unit_properties : list or None, default: None
4651
List of properties to be added to the unit table.
4752
These may be drawn from the sorting extractor, and, if available,
48-
the quality_metrics and template_metrics extensions of the SortingAnalyzer.
53+
the quality_metrics/template_metrics/unit_locations extensions of the SortingAnalyzer.
4954
See all properties available with sorting.get_property_keys(), and, if available,
5055
analyzer.get_extension("quality_metrics").get_data().columns and
5156
analyzer.get_extension("template_metrics").get_data().columns.
52-
(sortingview backend)
57+
extra_unit_properties : dict or None, default: None
58+
A dict with extra units properties to display.
59+
curation_dict : dict or None, default: None
60+
When curation is True, optionaly the viewer can get a previous 'curation_dict'
61+
to continue/check previous curations on this analyzer.
62+
In this case label_definitions must be None beacuse it is already included in the curation_dict.
63+
(spikeinterface_gui backend)
64+
label_definitions : dict or None, default: None
65+
When curation is True, optionaly the user can provide a label_definitions dict.
66+
This replaces the label_choices in the curation_format.
67+
(spikeinterface_gui backend)
5368
"""
5469

5570
def __init__(
@@ -60,11 +75,24 @@ def __init__(
6075
max_amplitudes_per_unit=None,
6176
min_similarity_for_correlograms=0.2,
6277
curation=False,
63-
unit_table_properties=None,
78+
displayed_unit_properties=None,
79+
extra_unit_properties=None,
6480
label_choices=None,
81+
curation_dict=None,
82+
label_definitions=None,
6583
backend=None,
84+
unit_table_properties=None,
6685
**backend_kwargs,
6786
):
87+
88+
if unit_table_properties is not None:
89+
warnings.warn(
90+
"plot_sorting_summary() : unit_table_properties is deprecated, use displayed_unit_properties instead",
91+
category=DeprecationWarning,
92+
stacklevel=2,
93+
)
94+
displayed_unit_properties = unit_table_properties
95+
6896
sorting_analyzer = self.ensure_sorting_analyzer(sorting_analyzer)
6997
self.check_extensions(
7098
sorting_analyzer, ["correlograms", "spike_amplitudes", "unit_locations", "template_similarity"]
@@ -74,18 +102,29 @@ def __init__(
74102
if unit_ids is None:
75103
unit_ids = sorting.get_unit_ids()
76104

77-
plot_data = dict(
105+
if curation_dict is not None and label_definitions is not None:
106+
raise ValueError("curation_dict and label_definitions are mutualy exclusive, they cannot be not None both")
107+
108+
if displayed_unit_properties is None:
109+
displayed_unit_properties = list(_default_displayed_unit_properties)
110+
if extra_unit_properties is not None:
111+
displayed_unit_properties += list(extra_unit_properties.keys())
112+
113+
data_plot = dict(
78114
sorting_analyzer=sorting_analyzer,
79115
unit_ids=unit_ids,
80116
sparsity=sparsity,
81117
min_similarity_for_correlograms=min_similarity_for_correlograms,
82-
unit_table_properties=unit_table_properties,
118+
displayed_unit_properties=displayed_unit_properties,
119+
extra_unit_properties=extra_unit_properties,
83120
curation=curation,
84121
label_choices=label_choices,
85122
max_amplitudes_per_unit=max_amplitudes_per_unit,
123+
curation_dict=curation_dict,
124+
label_definitions=label_definitions,
86125
)
87126

88-
BaseWidget.__init__(self, plot_data, backend=backend, **backend_kwargs)
127+
BaseWidget.__init__(self, data_plot, backend=backend, **backend_kwargs)
89128

90129
def plot_sortingview(self, data_plot, **backend_kwargs):
91130
import sortingview.views as vv
@@ -156,7 +195,7 @@ def plot_sortingview(self, data_plot, **backend_kwargs):
156195

157196
# unit ids
158197
v_units_table = generate_unit_table_view(
159-
dp.sorting_analyzer, dp.unit_table_properties, similarity_scores=similarity_scores
198+
dp.sorting_analyzer, dp.displayed_unit_properties, similarity_scores=similarity_scores
160199
)
161200

162201
if dp.curation:
@@ -190,9 +229,14 @@ def plot_sortingview(self, data_plot, **backend_kwargs):
190229
def plot_spikeinterface_gui(self, data_plot, **backend_kwargs):
191230
sorting_analyzer = data_plot["sorting_analyzer"]
192231

193-
import spikeinterface_gui
232+
from spikeinterface_gui import run_mainwindow
194233

195-
app = spikeinterface_gui.mkQApp()
196-
win = spikeinterface_gui.MainWindow(sorting_analyzer, curation=data_plot["curation"])
197-
win.show()
198-
app.exec_()
234+
run_mainwindow(
235+
sorting_analyzer,
236+
with_traces=True,
237+
curation=data_plot["curation"],
238+
curation_dict=data_plot["curation_dict"],
239+
label_definitions=data_plot["label_definitions"],
240+
extra_unit_properties=data_plot["extra_unit_properties"],
241+
displayed_unit_properties=data_plot["displayed_unit_properties"],
242+
)

src/spikeinterface/widgets/tests/test_widgets.py

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,9 @@ def setUpClass(cls):
7373
spike_amplitudes=dict(),
7474
unit_locations=dict(),
7575
spike_locations=dict(),
76-
quality_metrics=dict(metric_names=["snr", "isi_violation", "num_spikes", "amplitude_cutoff"]),
76+
quality_metrics=dict(
77+
metric_names=["snr", "isi_violation", "num_spikes", "firing_rate", "amplitude_cutoff"]
78+
),
7779
template_metrics=dict(),
7880
correlograms=dict(),
7981
template_similarity=dict(),
@@ -531,26 +533,37 @@ def test_plot_sorting_summary(self):
531533
possible_backends = list(sw.SortingSummaryWidget.get_possible_backends())
532534
for backend in possible_backends:
533535
if backend not in self.skip_backends:
534-
sw.plot_sorting_summary(self.sorting_analyzer_dense, backend=backend, **self.backend_kwargs[backend])
535-
sw.plot_sorting_summary(self.sorting_analyzer_sparse, backend=backend, **self.backend_kwargs[backend])
536+
sw.plot_sorting_summary(
537+
self.sorting_analyzer_dense,
538+
displayed_unit_properties=[],
539+
backend=backend,
540+
**self.backend_kwargs[backend],
541+
)
542+
sw.plot_sorting_summary(
543+
self.sorting_analyzer_sparse,
544+
displayed_unit_properties=[],
545+
backend=backend,
546+
**self.backend_kwargs[backend],
547+
)
536548
sw.plot_sorting_summary(
537549
self.sorting_analyzer_sparse,
538550
sparsity=self.sparsity_strict,
551+
displayed_unit_properties=[],
539552
backend=backend,
540553
**self.backend_kwargs[backend],
541554
)
542-
# add unit_properties
555+
# select unit_properties
543556
sw.plot_sorting_summary(
544557
self.sorting_analyzer_sparse,
545-
unit_table_properties=["firing_rate", "snr"],
558+
displayed_unit_properties=["firing_rate", "snr"],
546559
backend=backend,
547560
**self.backend_kwargs[backend],
548561
)
549562
# adding a missing property should raise a warning
550563
with self.assertWarns(UserWarning):
551564
sw.plot_sorting_summary(
552565
self.sorting_analyzer_sparse,
553-
unit_table_properties=["missing_property"],
566+
displayed_unit_properties=["missing_property"],
554567
backend=backend,
555568
**self.backend_kwargs[backend],
556569
)
@@ -688,9 +701,9 @@ def test_plot_motion_info(self):
688701
# mytest.test_plot_unit_presence()
689702
# mytest.test_plot_peak_activity()
690703
# mytest.test_plot_multicomparison()
691-
# mytest.test_plot_sorting_summary()
704+
mytest.test_plot_sorting_summary()
692705
# mytest.test_plot_motion()
693-
mytest.test_plot_motion_info()
694-
plt.show()
706+
# mytest.test_plot_motion_info()
707+
# plt.show()
695708

696709
# TestWidgets.tearDownClass()

src/spikeinterface/widgets/utils.py

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -243,3 +243,93 @@ def array_to_image(
243243
output_image = np.frombuffer(image.tobytes(), dtype=np.uint8).reshape(output_image.shape)
244244

245245
return output_image
246+
247+
248+
def make_units_table_from_sorting(sorting, units_table=None):
249+
"""
250+
Make a DataFrame from sorting properties.
251+
Only for properties with ndim=1
252+
253+
Parameters
254+
----------
255+
sorting : Sorting
256+
The Sorting object
257+
units_table : None | pd.DataFrame
258+
Optionally a existing dataframe.
259+
260+
Returns
261+
-------
262+
units_table : pd.DataFrame
263+
Table containing all columns.
264+
"""
265+
266+
if units_table is None:
267+
import pandas as pd
268+
269+
units_table = pd.DataFrame(index=sorting.unit_ids)
270+
271+
for col in sorting.get_property_keys():
272+
values = sorting.get_property(col)
273+
if values.dtype.kind in "iuUSfb" and values.ndim == 1:
274+
units_table.loc[:, col] = values
275+
276+
return units_table
277+
278+
279+
def make_units_table_from_analyzer(
280+
analyzer,
281+
extra_properties=None,
282+
):
283+
"""
284+
Make a DataFrame by aggregating :
285+
* quality metrics
286+
* template metrics
287+
* unit_position
288+
* sorting properties
289+
* extra columns
290+
291+
This used in sortingview and spikeinterface-gui to display the units table in a flexible way.
292+
293+
Parameters
294+
----------
295+
sorting_analyzer : SortingAnalyzer
296+
The SortingAnalyzer object
297+
extra_properties : None | dict
298+
Extra columns given as dict.
299+
300+
Returns
301+
-------
302+
units_table : pd.DataFrame
303+
Table containing all columns.
304+
"""
305+
import pandas as pd
306+
307+
all_df = []
308+
309+
if analyzer.get_extension("unit_locations") is not None:
310+
locs = analyzer.get_extension("unit_locations").get_data()
311+
df = pd.DataFrame(locs[:, :2], columns=["x", "y"], index=analyzer.unit_ids)
312+
all_df.append(df)
313+
314+
if analyzer.get_extension("quality_metrics") is not None:
315+
df = analyzer.get_extension("quality_metrics").get_data()
316+
all_df.append(df)
317+
318+
if analyzer.get_extension("template_metrics") is not None:
319+
df = analyzer.get_extension("template_metrics").get_data()
320+
all_df.append(df)
321+
322+
if len(all_df) > 0:
323+
units_table = pd.concat(all_df, axis=1)
324+
else:
325+
units_table = pd.DataFrame(index=analyzer.unit_ids)
326+
327+
make_units_table_from_sorting(analyzer.sorting, units_table=units_table)
328+
329+
if extra_properties is not None:
330+
for col, values in extra_properties.items():
331+
# the ndim = 1 is important because we need column only for the display in gui.
332+
if values.dtype.kind in "iuUSfb" and values.ndim == 1:
333+
units_table.loc[:, col] = values
334+
335+
return units_table

0 commit comments

Comments
 (0)