Skip to content

Commit 753cf23

Browse files
committed
Update AmplitudesWidget
1 parent 755bf18 commit 753cf23

1 file changed

Lines changed: 18 additions & 165 deletions

File tree

src/spikeinterface/widgets/amplitudes.py

Lines changed: 18 additions & 165 deletions
Original file line numberDiff line numberDiff line change
@@ -3,13 +3,16 @@
33
import numpy as np
44
from warnings import warn
55

6+
from .rasters import BaseRasterWidget
67
from .base import BaseWidget, to_attr
78
from .utils import get_some_colors
89

910
from ..core.sortinganalyzer import SortingAnalyzer
1011

12+
from spikeinterface.core import SortingAnalyzer
1113

12-
class AmplitudesWidget(BaseWidget):
14+
15+
class AmplitudesWidget(BaseRasterWidget):
1316
"""
1417
Plots spike amplitudes
1518
@@ -19,6 +22,9 @@ class AmplitudesWidget(BaseWidget):
1922
The input waveform extractor
2023
unit_ids : list or None, default: None
2124
List of unit ids
25+
unit_colors : dict | None, default: None
26+
Dict of colors with key : unit, value : color. If None, uses `get_some_colors` to
27+
generate colors.
2228
segment_index : int or None, default: None
2329
The segment index (or None if mono-segment)
2430
max_spikes_per_unit : int or None, default: None
@@ -61,9 +67,6 @@ def __init__(
6167
if unit_ids is None:
6268
unit_ids = sorting.unit_ids
6369

64-
if unit_colors is None:
65-
unit_colors = get_some_colors(sorting.unit_ids)
66-
6770
if sorting.get_num_segments() > 1:
6871
if segment_index is None:
6972
warn("More than one segment available! Using segment_index 0")
@@ -73,13 +76,11 @@ def __init__(
7376
amplitudes_segment = amplitudes[segment_index]
7477
total_duration = sorting_analyzer.get_num_samples(segment_index) / sorting_analyzer.sampling_frequency
7578

76-
spiketrains_segment = {}
77-
for i, unit_id in enumerate(sorting.unit_ids):
78-
times = sorting.get_unit_spike_train(unit_id, segment_index=segment_index)
79-
times = times / sorting.get_sampling_frequency()
80-
spiketrains_segment[unit_id] = times
79+
all_spiketrains = {
80+
unit_id: sorting.get_unit_spike_train(unit_id, segment_index=0, return_times=True)
81+
for unit_id in sorting.unit_ids
82+
}
8183

82-
all_spiketrains = spiketrains_segment
8384
all_amplitudes = amplitudes_segment
8485
if max_spikes_per_unit is not None:
8586
spiketrains_to_plot = dict()
@@ -97,167 +98,19 @@ def __init__(
9798
spiketrains_to_plot = all_spiketrains
9899
amplitudes_to_plot = all_amplitudes
99100

100-
if plot_histograms and bins is None:
101-
bins = 100
102-
103101
plot_data = dict(
104-
sorting_analyzer=sorting_analyzer,
105-
amplitudes=amplitudes_to_plot,
106-
unit_ids=unit_ids,
102+
spike_train_data=spiketrains_to_plot,
103+
y_axis_data=amplitudes_to_plot,
107104
unit_colors=unit_colors,
108-
spiketrains=spiketrains_to_plot,
109-
total_duration=total_duration,
110105
plot_histograms=plot_histograms,
111106
bins=bins,
107+
unit_ids=unit_ids,
112108
hide_unit_selector=hide_unit_selector,
113109
plot_legend=plot_legend,
110+
y_label="Amplitude",
114111
)
115112

116-
BaseWidget.__init__(self, plot_data, backend=backend, **backend_kwargs)
117-
118-
def plot_matplotlib(self, data_plot, **backend_kwargs):
119-
import matplotlib.pyplot as plt
120-
from .utils_matplotlib import make_mpl_figure
121-
122-
dp = to_attr(data_plot)
123-
124-
if backend_kwargs["axes"] is not None:
125-
axes = backend_kwargs["axes"]
126-
if dp.plot_histograms:
127-
assert np.asarray(axes).size == 2
128-
else:
129-
assert np.asarray(axes).size == 1
130-
elif backend_kwargs["ax"] is not None:
131-
assert not dp.plot_histograms
132-
else:
133-
if dp.plot_histograms:
134-
backend_kwargs["num_axes"] = 2
135-
backend_kwargs["ncols"] = 2
136-
else:
137-
backend_kwargs["num_axes"] = None
138-
139-
self.figure, self.axes, self.ax = make_mpl_figure(**backend_kwargs)
140-
141-
scatter_ax = self.axes.flatten()[0]
142-
143-
for unit_id in dp.unit_ids:
144-
spiketrains = dp.spiketrains[unit_id]
145-
amps = dp.amplitudes[unit_id]
146-
scatter_ax.scatter(spiketrains, amps, color=dp.unit_colors[unit_id], s=3, alpha=1, label=unit_id)
147-
148-
if dp.plot_histograms:
149-
if dp.bins is None:
150-
bins = int(len(spiketrains) / 30)
151-
else:
152-
bins = dp.bins
153-
ax_hist = self.axes.flatten()[1]
154-
# this is super slow, using plot and np.histogram is really much faster (and nicer!)
155-
# ax_hist.hist(amps, bins=bins, orientation="horizontal", color=dp.unit_colors[unit_id], alpha=0.8)
156-
count, bins = np.histogram(amps, bins=bins)
157-
ax_hist.plot(count, bins[:-1], color=dp.unit_colors[unit_id], alpha=0.8)
158-
159-
if dp.plot_histograms:
160-
ax_hist = self.axes.flatten()[1]
161-
ax_hist.set_ylim(scatter_ax.get_ylim())
162-
ax_hist.axis("off")
163-
# self.figure.tight_layout()
164-
165-
if dp.plot_legend:
166-
if hasattr(self, "legend") and self.legend is not None:
167-
self.legend.remove()
168-
self.legend = self.figure.legend(
169-
loc="upper center", bbox_to_anchor=(0.5, 1.0), ncol=5, fancybox=True, shadow=True
170-
)
171-
172-
scatter_ax.set_xlim(0, dp.total_duration)
173-
scatter_ax.set_xlabel("Times [s]")
174-
scatter_ax.set_ylabel(f"Amplitude")
175-
scatter_ax.spines["top"].set_visible(False)
176-
scatter_ax.spines["right"].set_visible(False)
177-
self.figure.subplots_adjust(bottom=0.1, top=0.9, left=0.1)
178-
179-
def plot_ipywidgets(self, data_plot, **backend_kwargs):
180-
import matplotlib.pyplot as plt
181-
182-
# import ipywidgets.widgets as widgets
183-
import ipywidgets.widgets as W
184-
from IPython.display import display
185-
from .utils_ipywidgets import check_ipywidget_backend, UnitSelector
186-
187-
check_ipywidget_backend()
188-
189-
self.next_data_plot = data_plot.copy()
190-
191-
cm = 1 / 2.54
192-
analyzer = data_plot["sorting_analyzer"]
193-
194-
width_cm = backend_kwargs["width_cm"]
195-
height_cm = backend_kwargs["height_cm"]
196-
197-
ratios = [0.15, 0.85]
198-
199-
with plt.ioff():
200-
output = W.Output()
201-
with output:
202-
self.figure = plt.figure(figsize=((ratios[1] * width_cm) * cm, height_cm * cm))
203-
plt.show()
204-
205-
self.unit_selector = UnitSelector(analyzer.unit_ids)
206-
self.unit_selector.value = list(analyzer.unit_ids)[:1]
207-
208-
self.checkbox_histograms = W.Checkbox(
209-
value=data_plot["plot_histograms"],
210-
description="hist",
211-
)
212-
213-
left_sidebar = W.VBox(
214-
children=[
215-
self.unit_selector,
216-
self.checkbox_histograms,
217-
],
218-
layout=W.Layout(align_items="center", width="100%", height="100%"),
219-
)
220-
221-
self.widget = W.AppLayout(
222-
center=self.figure.canvas,
223-
left_sidebar=left_sidebar,
224-
pane_widths=ratios + [0],
225-
)
226-
227-
# a first update
228-
self._full_update_plot()
229-
230-
self.unit_selector.observe(self._update_plot, names="value", type="change")
231-
self.checkbox_histograms.observe(self._full_update_plot, names="value", type="change")
232-
233-
if backend_kwargs["display"]:
234-
display(self.widget)
235-
236-
def _full_update_plot(self, change=None):
237-
self.figure.clear()
238-
data_plot = self.next_data_plot
239-
data_plot["unit_ids"] = self.unit_selector.value
240-
data_plot["plot_histograms"] = self.checkbox_histograms.value
241-
data_plot["plot_legend"] = False
242-
243-
backend_kwargs = dict(figure=self.figure, axes=None, ax=None)
244-
self.plot_matplotlib(data_plot, **backend_kwargs)
245-
self._update_plot()
246-
247-
def _update_plot(self, change=None):
248-
for ax in self.axes.flatten():
249-
ax.clear()
250-
251-
data_plot = self.next_data_plot
252-
data_plot["unit_ids"] = self.unit_selector.value
253-
data_plot["plot_histograms"] = self.checkbox_histograms.value
254-
data_plot["plot_legend"] = False
255-
256-
backend_kwargs = dict(figure=None, axes=self.axes, ax=None)
257-
self.plot_matplotlib(data_plot, **backend_kwargs)
258-
259-
self.figure.canvas.draw()
260-
self.figure.canvas.flush_events()
113+
BaseRasterWidget.__init__(self, **plot_data, backend=backend, **backend_kwargs)
261114

262115
def plot_sortingview(self, data_plot, **backend_kwargs):
263116
import sortingview.views as vv
@@ -270,8 +123,8 @@ def plot_sortingview(self, data_plot, **backend_kwargs):
270123
sa_items = [
271124
vv.SpikeAmplitudesItem(
272125
unit_id=u,
273-
spike_times_sec=dp.spiketrains[u].astype("float32"),
274-
spike_amplitudes=dp.amplitudes[u].astype("float32"),
126+
spike_times_sec=dp.spike_train_data[u].astype("float32"),
127+
spike_amplitudes=dp.y_axis_data[u].astype("float32"),
275128
)
276129
for u in unit_ids
277130
]

0 commit comments

Comments
 (0)