From f77e7b49284daae5d466781d8946dcbb433e4d69 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jes=C3=BAs=20Royeth?= Date: Fri, 9 Oct 2026 18:45:24 -0300 Subject: [PATCH] Speed up soft merges with vectorized spike remapping --- src/spikeinterface/core/sorting_tools.py | 44 ++++++++++--------- .../core/tests/test_sorting_tools.py | 32 ++++++++++++++ 2 files changed, 55 insertions(+), 21 deletions(-) diff --git a/src/spikeinterface/core/sorting_tools.py b/src/spikeinterface/core/sorting_tools.py index f37b24046d..7158c2c9b8 100644 --- a/src/spikeinterface/core/sorting_tools.py +++ b/src/spikeinterface/core/sorting_tools.py @@ -609,12 +609,12 @@ def apply_merges_to_sorting( sorting : The new Sorting object The newly create sorting with the merged units keep_mask : numpy.array - A boolean mask, if censor_ms is not None, telling which spike from the original spike vector - has been kept, given the refractory period violations (None if censor_ms is None) + Only returned if return_extra is True. A boolean mask telling which spike from the + original spike vector has been kept. All True if censor_ms is None, otherwise False + for spikes removed for violating the refractory period. """ spikes = sorting.to_spike_vector().copy() - keep_mask = np.ones(len(spikes), dtype=bool) new_unit_ids = generate_unit_ids_for_merge_group( sorting.unit_ids, merge_unit_groups, new_unit_ids=new_unit_ids, new_id_strategy=new_id_strategy @@ -628,25 +628,21 @@ def apply_merges_to_sorting( all_unit_ids = _get_ids_after_merging(sorting.unit_ids, merge_unit_groups, new_unit_ids) all_unit_ids = list(all_unit_ids) - num_seg = sorting.get_num_segments() - segment_slices = sorting._get_spike_vector_segment_slices() - - # using this function avoids to use the mask approach and simplify a lot the algo - spike_vector_list = [spikes[s0:s1] for s0, s1 in segment_slices] - spike_indices = spike_vector_to_indices(spike_vector_list, sorting.unit_ids, absolute_index=True) - - for old_unit_id in sorting.unit_ids: - if old_unit_id in rename_ids.keys(): - new_unit_id = rename_ids[old_unit_id] - else: - new_unit_id = old_unit_id - - new_unit_index = all_unit_ids.index(new_unit_id) - for segment_index in range(num_seg): - spike_inds = spike_indices[segment_index][old_unit_id] - spikes["unit_index"][spike_inds] = new_unit_index + # The existing remap helper cannot collapse several old unit ids into one new id. + new_unit_id_to_index = {unit_id: unit_index for unit_index, unit_id in enumerate(all_unit_ids)} + old_to_new_unit_indices = np.fromiter( + (new_unit_id_to_index[rename_ids.get(unit_id, unit_id)] for unit_id in sorting.unit_ids), + dtype=spikes["unit_index"].dtype, + count=sorting.unit_ids.size, + ) if censor_ms is not None: + keep_mask = np.ones(len(spikes), dtype=bool) + num_seg = sorting.get_num_segments() + segment_slices = sorting._get_spike_vector_segment_slices() + spike_vector_list = [spikes[s0:s1] for s0, s1 in segment_slices] + spike_indices = spike_vector_to_indices(spike_vector_list, sorting.unit_ids, absolute_index=True) + rpv = int(sorting.sampling_frequency * censor_ms / 1000.0) for group_old_ids in merge_unit_groups: for segment_index in range(num_seg): @@ -658,7 +654,13 @@ def apply_merges_to_sorting( inds = np.flatnonzero(np.diff(spikes["sample_index"][group_indices]) < rpv) keep_mask[group_indices[inds + 1]] = False - spikes = spikes[keep_mask] + spikes = spikes[keep_mask] + elif return_extra: + keep_mask = np.ones(len(spikes), dtype=bool) + else: + keep_mask = None + + spikes["unit_index"] = old_to_new_unit_indices[spikes["unit_index"]] merge_sorting = NumpySorting(spikes, sorting.sampling_frequency, all_unit_ids) set_properties_after_merging(merge_sorting, sorting, merge_unit_groups, new_unit_ids=new_unit_ids) diff --git a/src/spikeinterface/core/tests/test_sorting_tools.py b/src/spikeinterface/core/tests/test_sorting_tools.py index 7ad8291073..35fcdf769c 100644 --- a/src/spikeinterface/core/tests/test_sorting_tools.py +++ b/src/spikeinterface/core/tests/test_sorting_tools.py @@ -269,6 +269,38 @@ def test_apply_merges_to_sorting(): assert np.array_equal(sorting2.unit_ids, ["a", "c"]) +def test_apply_merges_without_censor_does_not_build_spike_indices(monkeypatch): + sorting = NumpySorting.from_unit_dict( + {"a": np.array([0, 10]), "b": np.array([20, 30]), "c": np.array([40])}, + sampling_frequency=30_000.0, + ) + + def fail_if_called(*args, **kwargs): + pytest.fail("spike indices are unnecessary when censor_ms is None") + + monkeypatch.setattr("spikeinterface.core.sorting_tools.spike_vector_to_indices", fail_if_called) + merged = apply_merges_to_sorting(sorting, [["a", "b"]], censor_ms=None) + + assert np.array_equal(merged.get_unit_spike_train("c"), [40]) + assert np.array_equal(merged.get_unit_spike_train("merge0"), [0, 10, 20, 30]) + + +def test_apply_merges_without_censor_return_extra_keep_mask(): + sorting = NumpySorting.from_unit_dict( + {"a": np.array([0, 10]), "b": np.array([20, 30]), "c": np.array([40])}, + sampling_frequency=30_000.0, + ) + num_spikes = sorting.to_spike_vector().size + + merged, keep_mask, new_unit_ids = apply_merges_to_sorting(sorting, [["a", "b"]], censor_ms=None, return_extra=True) + + assert keep_mask.dtype == bool + assert keep_mask.size == num_spikes + assert keep_mask.all() + assert merged.to_spike_vector().size == num_spikes + assert new_unit_ids == ["merge0"] + + def test_get_ids_after_merging(): all_unit_ids = _get_ids_after_merging(["a", "b", "c", "d", "e"], [["a", "b"], ["d", "e"]], ["x", "d"])