diff --git a/src/spikeinterface/postprocessing/amplitude_scalings.py b/src/spikeinterface/postprocessing/amplitude_scalings.py index 9564f64141..be5967ad1b 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, + 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): """ 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 + The indices of `spikes` in `spikes_within_margin`. Providing these indices avoids + searching `spikes_within_margin` once for every spike. Returns ------- @@ -422,11 +430,7 @@ 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] - + 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 = [ spike["sample_index"] - delta_collision_samples, @@ -449,6 +453,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 +461,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..d6faabb616 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,38 @@ 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]]) + + class TestAmplitudeScalingsExtension(AnalyzerExtensionCommonTestSuite): @pytest.mark.parametrize("params", [dict(handle_collisions=True), dict(handle_collisions=False)])