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
44 changes: 23 additions & 21 deletions src/spikeinterface/core/sorting_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -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)

Expand Down
32 changes: 32 additions & 0 deletions src/spikeinterface/core/tests/test_sorting_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Expand Down
Loading