From 8a6a607c1d7ad0a9f5ee48bc83441db20a057450 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jes=C3=BAs=20Royeth?= Date: Tue, 15 Sep 2026 21:42:14 -0300 Subject: [PATCH] Speed up collision-aware amplitude scaling --- .../postprocessing/amplitude_scalings.py | 30 +++++++----- .../tests/test_amplitude_scalings.py | 47 ++++++++++++++++++- 2 files changed, 64 insertions(+), 13 deletions(-) diff --git a/src/spikeinterface/postprocessing/amplitude_scalings.py b/src/spikeinterface/postprocessing/amplitude_scalings.py index 9564f64141..25f22caa54 100644 --- a/src/spikeinterface/postprocessing/amplitude_scalings.py +++ b/src/spikeinterface/postprocessing/amplitude_scalings.py @@ -245,13 +245,18 @@ def compute(self, traces, peaks): # local_spikes = local_spikes_within_margin[i0:i1] local_spikes_within_margin = peaks - local_spikes = local_spikes_within_margin[~peaks["in_margin"]] + local_spike_indices = np.flatnonzero(~peaks["in_margin"]) + local_spikes = local_spikes_within_margin[local_spike_indices] # set colliding spikes apart (if needed) if handle_collisions: # local spikes with margin! collisions = find_collisions( - local_spikes, local_spikes_within_margin, delta_collision_samples, sparsity_mask + local_spikes, + local_spikes_within_margin, + delta_collision_samples, + sparsity_mask, + spike_indices=local_spike_indices, ) else: collisions = {} @@ -382,7 +387,7 @@ def _ordinary_scaling_slope(template, local_waveform): return covariance / template_variance -def find_collisions(spikes, spikes_within_margin, delta_collision_samples, sparsity_mask): +def find_collisions(spikes, spikes_within_margin, delta_collision_samples, sparsity_mask, spike_indices=None): """ Finds the collisions between spikes. @@ -413,6 +418,9 @@ def find_collisions(spikes, spikes_within_margin, delta_collision_samples, spars sparsity_mask: boolean mask A num_units x num_channels boolean array indicating whether the unit is represented on the channel. + spike_indices : np.ndarray or None, default: None + The indices of `spikes` in `spikes_within_margin`. Providing these indices avoids + searching `spikes_within_margin` once for every spike. Returns ------- @@ -422,10 +430,9 @@ def find_collisions(spikes, spikes_within_margin, delta_collision_samples, spars """ # TODO: refactor to speed-up collision_spikes_dict = {} - for spike_index, spike in enumerate(spikes): - - # find the index of the spike within spikes_within_margin - spike_index_within_margin = np.where(spikes_within_margin == spike)[0][0] + if spike_indices is None: + spike_indices = (np.where(spikes_within_margin == spike)[0][0] for spike in spikes) + for spike_index, (spike, spike_index_within_margin) in enumerate(zip(spikes, spike_indices, strict=True)): # find the spikes that fall within a temporal window around the spike peak spike_collision_window = [ @@ -449,6 +456,7 @@ def find_collisions(spikes, spikes_within_margin, delta_collision_samples, spars # Build the collusion_spikes_dict including only # spikes that overlap spatially + collision_spikes = [] for possible_overlapping_spike_index in possible_overlapping_spike_indices: if _are_units_spatially_overlapping( @@ -456,11 +464,9 @@ def find_collisions(spikes, spikes_within_margin, delta_collision_samples, spars spike["unit_index"], spikes_within_margin[possible_overlapping_spike_index]["unit_index"], ): - if spike_index not in collision_spikes_dict: - collision_spikes_dict[spike_index] = np.array([spike]) - collision_spikes_dict[spike_index] = np.concatenate( - (collision_spikes_dict[spike_index], [spikes_within_margin[possible_overlapping_spike_index]]) - ) + collision_spikes.append(spikes_within_margin[possible_overlapping_spike_index]) + if collision_spikes: + collision_spikes_dict[spike_index] = np.array([spike, *collision_spikes], dtype=spikes.dtype) return collision_spikes_dict diff --git a/src/spikeinterface/postprocessing/tests/test_amplitude_scalings.py b/src/spikeinterface/postprocessing/tests/test_amplitude_scalings.py index b9bfc09889..5bcd792f29 100644 --- a/src/spikeinterface/postprocessing/tests/test_amplitude_scalings.py +++ b/src/spikeinterface/postprocessing/tests/test_amplitude_scalings.py @@ -4,7 +4,8 @@ from spikeinterface.postprocessing.tests.common_extension_tests import AnalyzerExtensionCommonTestSuite from spikeinterface.postprocessing import ComputeAmplitudeScalings -from spikeinterface.postprocessing.amplitude_scalings import _ordinary_scaling_slope, fit_collision +from spikeinterface.core.base import spike_peak_dtype +from spikeinterface.postprocessing.amplitude_scalings import _ordinary_scaling_slope, find_collisions, fit_collision def test_ordinary_scaling_slope_float32_precision(): @@ -63,6 +64,50 @@ def test_fit_collision_recovers_positive_coefficients(): np.testing.assert_allclose(recovered, true_scalings, atol=0.05) +def test_find_collisions_with_margin_indices(monkeypatch): + dtype = spike_peak_dtype + [("in_margin", "bool")] + spikes_within_margin = np.array( + [ + (8, 0, -1.0, 0, 1, True), + (10, 0, -1.0, 0, 0, False), + (10, 1, -1.0, 0, 1, False), + (13, 2, -1.0, 0, 2, False), + (30, 1, -1.0, 0, 1, True), + ], + dtype=dtype, + ) + spike_indices = np.flatnonzero(~spikes_within_margin["in_margin"]) + spikes = spikes_within_margin[spike_indices] + sparsity_mask = np.array([[True, False], [True, True], [False, True]]) + + with monkeypatch.context() as patch_context: + patch_context.setattr(np, "where", lambda *_: pytest.fail("spike indices were searched again")) + collisions = find_collisions( + spikes, + spikes_within_margin, + delta_collision_samples=4, + sparsity_mask=sparsity_mask, + spike_indices=spike_indices, + ) + + assert set(collisions) == {0, 1, 2} + np.testing.assert_array_equal(collisions[0], spikes_within_margin[[1, 0, 2]]) + np.testing.assert_array_equal(collisions[1], spikes_within_margin[[2, 0, 1, 3]]) + np.testing.assert_array_equal(collisions[2], spikes_within_margin[[3, 2]]) + + collisions_without_indices = find_collisions( + spikes, + spikes_within_margin, + delta_collision_samples=4, + sparsity_mask=sparsity_mask, + ) + for spike_index in collisions: + np.testing.assert_array_equal(collisions_without_indices[spike_index], collisions[spike_index]) + + with pytest.raises(ValueError, match="shorter"): + find_collisions(spikes, spikes_within_margin, 4, sparsity_mask, spike_indices=spike_indices[:-1]) + + class TestAmplitudeScalingsExtension(AnalyzerExtensionCommonTestSuite): @pytest.mark.parametrize("params", [dict(handle_collisions=True), dict(handle_collisions=False)])