Skip to content

Commit fdc92c4

Browse files
committed
Support released phylib spike selector
1 parent ba18b83 commit fdc92c4

3 files changed

Lines changed: 88 additions & 9 deletions

File tree

docs/changelog.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,8 @@ behavior they verify rather than listed separately.
4040

4141
### Fixed
4242

43+
- Start the GUI with released phylib versions that do not yet expose the
44+
disjoint-spike selection optimization hint.
4345
- Keep dataset-local view settings isolated from global GUI state. In
4446
particular, a Firing Rate time range saved or leaked from another recording
4547
no longer clips spikes in a fresh dataset.

phy/apps/base.py

Lines changed: 27 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -113,6 +113,20 @@ def _sample_spikes_evenly(spike_ids, n_spikes):
113113
return np.asarray(spike_ids[indices], dtype=np.int64)
114114

115115

116+
def _select_spikes_evenly(selector, n_spikes, cluster_ids, **kwargs):
117+
"""Use phylib's even selector when available, with a 2.7-compatible fallback."""
118+
if 'sample_evenly' in inspect.signature(selector).parameters:
119+
return selector(n_spikes, cluster_ids, sample_evenly=True, **kwargs)
120+
121+
selected = [
122+
_sample_spikes_evenly(selector(None, [cluster_id], **kwargs), n_spikes)
123+
for cluster_id in cluster_ids
124+
]
125+
if not selected:
126+
return np.array([], dtype=np.int64)
127+
return np.concatenate(selected)
128+
129+
116130
def _spike_budget_fields(per_cluster, total, max_n_clusters, background=None):
117131
"""Return dialog fields for a per-cluster budget and optional shared cap."""
118132
per_cluster_default = per_cluster if per_cluster is not None else total or 100000
@@ -318,8 +332,8 @@ def _get_waveforms_with_n_spikes(self, cluster_id, n_spikes_waveforms, current_f
318332
# Or keep spikes from a subset of the chunks for performance reasons (decompression will
319333
# happen on the fly here).
320334
else:
321-
spike_ids = self.selector(
322-
n_spikes_waveforms, [cluster_id], subset_chunks=True, sample_evenly=True
335+
spike_ids = _select_spikes_evenly(
336+
self.selector, n_spikes_waveforms, [cluster_id], subset_chunks=True
323337
)
324338

325339
# Get the best channels.
@@ -1269,13 +1283,17 @@ def spikes_per_cluster(cluster_id):
12691283
except AttributeError:
12701284
chunk_bounds = [0.0, self.model.spike_samples[-1] + 1]
12711285

1272-
self.selector = SpikeSelector(
1273-
get_spikes_per_cluster=spikes_per_cluster,
1274-
spike_times=self.model.spike_samples, # NOTE: chunk_bounds is in samples, not seconds
1275-
chunk_bounds=chunk_bounds,
1276-
n_chunks_kept=self.n_chunks_kept,
1277-
spikes_are_disjoint=True,
1278-
)
1286+
selector_kwargs = {
1287+
'get_spikes_per_cluster': spikes_per_cluster,
1288+
# NOTE: chunk_bounds is in samples, not seconds.
1289+
'spike_times': self.model.spike_samples,
1290+
'chunk_bounds': chunk_bounds,
1291+
'n_chunks_kept': self.n_chunks_kept,
1292+
}
1293+
# phylib 2.7 does not yet expose this optimization hint.
1294+
if 'spikes_are_disjoint' in inspect.signature(SpikeSelector).parameters:
1295+
selector_kwargs['spikes_are_disjoint'] = True
1296+
self.selector = SpikeSelector(**selector_kwargs)
12791297

12801298
def _cache_methods(self):
12811299
"""Cache methods as specified in `self._memcached` and `self._cached`."""

phy/apps/tests/test_base.py

Lines changed: 59 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -46,6 +46,7 @@
4646
TraceMixin,
4747
WaveformMixin,
4848
_allocate_spike_counts,
49+
_select_spikes_evenly,
4950
_spike_budget_fields,
5051
_spike_budget_values,
5152
)
@@ -215,6 +216,64 @@ def test_controller_close(tempdir):
215216
assert all(handler.stream is None for handler in handlers)
216217

217218

219+
def test_set_selector_supports_released_and_newer_phylib():
220+
controller = object.__new__(BaseController)
221+
controller.model = Bunch(spike_samples=np.arange(4))
222+
controller.supervisor = Bunch(clustering=Clustering(np.array([0, 0, 1, 1])))
223+
controller.n_chunks_kept = 2
224+
225+
class ReleasedSpikeSelector:
226+
def __init__(
227+
self,
228+
get_spikes_per_cluster=None,
229+
spike_times=None,
230+
chunk_bounds=None,
231+
n_chunks_kept=None,
232+
):
233+
self.spikes_are_disjoint = None
234+
235+
with patch('phy.apps.base.SpikeSelector', ReleasedSpikeSelector):
236+
controller._set_selector()
237+
assert controller.selector.spikes_are_disjoint is None
238+
239+
class NewerSpikeSelector(ReleasedSpikeSelector):
240+
def __init__(self, *args, spikes_are_disjoint=False, **kwargs):
241+
super().__init__(*args, **kwargs)
242+
self.spikes_are_disjoint = spikes_are_disjoint
243+
244+
with patch('phy.apps.base.SpikeSelector', NewerSpikeSelector):
245+
controller._set_selector()
246+
assert controller.selector.spikes_are_disjoint is True
247+
248+
249+
def test_select_spikes_evenly_supports_released_and_newer_phylib():
250+
spikes = {0: np.arange(5), 1: np.arange(5, 10)}
251+
252+
def released_selector(n_spikes, cluster_ids, subset_chunks=False):
253+
assert subset_chunks
254+
assert n_spikes is None
255+
return np.concatenate([spikes[cluster_id] for cluster_id in cluster_ids])
256+
257+
np.testing.assert_array_equal(
258+
_select_spikes_evenly(released_selector, 2, [0, 1], subset_chunks=True),
259+
[0, 4, 5, 9],
260+
)
261+
262+
def newer_selector(n_spikes, cluster_ids, subset_chunks=False, sample_evenly=False):
263+
assert (n_spikes, cluster_ids, subset_chunks, sample_evenly) == (
264+
2,
265+
[0, 1],
266+
True,
267+
True,
268+
)
269+
return np.array([1, 8])
270+
271+
np.testing.assert_array_equal(
272+
_select_spikes_evenly(newer_selector, 2, [0, 1], subset_chunks=True),
273+
[1, 8],
274+
)
275+
276+
218277
def test_get_firing_rate_fast_path():
219278
spike_times = np.array([0.1, 0.2, 0.4, 0.8, 1.6, 3.2])
220279
clustering = Clustering(np.array([0, 1, 1, 2, 2, 3]))

0 commit comments

Comments
 (0)