Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 18 additions & 12 deletions src/spikeinterface/postprocessing/amplitude_scalings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {}
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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
-------
Expand All @@ -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 = [
Expand All @@ -449,18 +456,17 @@ 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(
sparsity_mask,
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


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down Expand Up @@ -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)])
Expand Down
Loading