diff --git a/src/spikeinterface/core/sorting_tools.py b/src/spikeinterface/core/sorting_tools.py index 127eae9246..6761d607cd 100644 --- a/src/spikeinterface/core/sorting_tools.py +++ b/src/spikeinterface/core/sorting_tools.py @@ -487,9 +487,9 @@ def set_properties_after_merging( pre_unit_ids = sorting_pre_merge.unit_ids post_unit_ids = sorting_post_merge.unit_ids - kept_unit_ids = post_unit_ids[np.isin(post_unit_ids, pre_unit_ids)] - keep_pre_inds = sorting_pre_merge.ids_to_indices(kept_unit_ids) - keep_post_inds = sorting_post_merge.ids_to_indices(kept_unit_ids) + untouched_unit_ids = post_unit_ids[np.isin(post_unit_ids, pre_unit_ids) & ~np.isin(post_unit_ids, new_unit_ids)] + keep_pre_inds = sorting_pre_merge.ids_to_indices(untouched_unit_ids) + keep_post_inds = sorting_post_merge.ids_to_indices(untouched_unit_ids) default_missing_values = BaseExtractor.default_missing_property_values @@ -772,9 +772,12 @@ def set_properties_after_splits( pre_unit_ids = sorting_pre_split.unit_ids post_unit_ids = sorting_post_split.unit_ids - kept_unit_ids = post_unit_ids[np.isin(post_unit_ids, pre_unit_ids)] - keep_pre_inds = sorting_pre_split.ids_to_indices(kept_unit_ids) - keep_post_inds = sorting_post_split.ids_to_indices(kept_unit_ids) + all_new_split_unit_ids = [uid for group in new_unit_ids for uid in group] + untouched_unit_ids = post_unit_ids[ + np.isin(post_unit_ids, pre_unit_ids) & ~np.isin(post_unit_ids, all_new_split_unit_ids) + ] + keep_pre_inds = sorting_pre_split.ids_to_indices(untouched_unit_ids) + keep_post_inds = sorting_post_split.ids_to_indices(untouched_unit_ids) for key in prop_keys: parent_values = sorting_pre_split.get_property(key) diff --git a/src/spikeinterface/core/tests/test_sorting_tools.py b/src/spikeinterface/core/tests/test_sorting_tools.py index 4194f459b3..b8c27eab18 100644 --- a/src/spikeinterface/core/tests/test_sorting_tools.py +++ b/src/spikeinterface/core/tests/test_sorting_tools.py @@ -10,9 +10,12 @@ random_spikes_selection, spike_vector_to_indices, apply_merges_to_sorting, + apply_splits_to_sorting, _get_ids_after_merging, generate_unit_ids_for_merge_group, remap_unit_indices_in_vector, + set_properties_after_merging, + set_properties_after_splits, ) from spikeinterface.core.base import minimum_spike_dtype @@ -165,6 +168,151 @@ def test_generate_unit_ids_for_merge_group(): assert np.array_equal(new_unit_ids, ["0-5", "9-15"]) +def _make_sorting_with_properties(): + """Helper: 4-unit sorting with float and str properties.""" + times = np.array([0, 10, 20, 30, 40]) + labels = np.array(["a", "b", "c", "d", "a"]) + sorting = NumpySorting.from_samples_and_labels([times], [labels], 10_000.0, unit_ids=["a", "b", "c", "d"]) + # same value for "a" and "b", different for "c" and "d" + sorting.set_property("quality", np.array([1.0, 1.0, 2.0, 3.0])) + sorting.set_property("group", np.array(["g1", "g1", "g2", "g3"])) + return sorting + + +def test_set_properties_after_merging(): + sorting = _make_sorting_with_properties() + + # --- append strategy (baseline) --- + sorting_merged, _, _ = apply_merges_to_sorting( + sorting, [["a", "b"]], censor_ms=None, new_id_strategy="append", return_extra=True + ) + is_merged = sorting_merged.get_property("is_merged") + # "merge0" is the new merged unit; "c" and "d" are kept + merged_idx = sorting_merged.id_to_index("merge0") + kept_c_idx = sorting_merged.id_to_index("c") + kept_d_idx = sorting_merged.id_to_index("d") + assert is_merged[merged_idx] is True or bool(is_merged[merged_idx]) + assert not is_merged[kept_c_idx] + assert not is_merged[kept_d_idx] + # same quality value for "a" and "b" → propagated + quality = sorting_merged.get_property("quality") + assert quality[merged_idx] == 1.0 + + # --- take_first strategy (the bug fix) --- + sorting_merged2, _, _ = apply_merges_to_sorting( + sorting, [["a", "b"]], censor_ms=None, new_id_strategy="take_first", return_extra=True + ) + is_merged2 = sorting_merged2.get_property("is_merged") + # "a" is the new merged unit (take_first); "c" and "d" are kept + merged_idx2 = sorting_merged2.id_to_index("a") + kept_c_idx2 = sorting_merged2.id_to_index("c") + kept_d_idx2 = sorting_merged2.id_to_index("d") + assert is_merged2[merged_idx2] # was False before the fix + assert not is_merged2[kept_c_idx2] + assert not is_merged2[kept_d_idx2] + # same quality for "a" and "b" → propagated (merge logic, not raw copy) + quality2 = sorting_merged2.get_property("quality") + assert quality2[merged_idx2] == 1.0 + + # --- join strategy --- + sorting_merged3, _, _ = apply_merges_to_sorting( + sorting, [["a", "b"]], censor_ms=None, new_id_strategy="join", return_extra=True + ) + is_merged3 = sorting_merged3.get_property("is_merged") + merged_idx3 = sorting_merged3.id_to_index("a-b") + assert is_merged3[merged_idx3] + + # --- multiple merge groups --- + sorting5 = NumpySorting.from_samples_and_labels( + [np.array([0, 10, 20, 30, 40, 50])], + [np.array(["a", "b", "c", "d", "e", "f"])], + 10_000.0, + unit_ids=["a", "b", "c", "d", "e", "f"], + ) + sorting5.set_property("quality", np.array([1.0, 1.0, 2.0, 2.0, 3.0, 3.0])) + sorting_merged5 = apply_merges_to_sorting( + sorting5, [["a", "b"], ["c", "d"]], censor_ms=None, new_id_strategy="take_first" + ) + is_merged5 = sorting_merged5.get_property("is_merged") + assert is_merged5[sorting_merged5.id_to_index("a")] + assert is_merged5[sorting_merged5.id_to_index("c")] + assert not is_merged5[sorting_merged5.id_to_index("e")] + assert not is_merged5[sorting_merged5.id_to_index("f")] + + # --- different property values for merged units → default fill --- + sorting_diff = NumpySorting.from_samples_and_labels( + [np.array([0, 10, 20])], + [np.array(["a", "b", "c"])], + 10_000.0, + unit_ids=["a", "b", "c"], + ) + sorting_diff.set_property("quality", np.array([1.0, 2.0, 3.0])) # a≠b + sorting_diff_merged = apply_merges_to_sorting( + sorting_diff, [["a", "b"]], censor_ms=None, new_id_strategy="take_first" + ) + is_merged_diff = sorting_diff_merged.get_property("is_merged") + assert is_merged_diff[sorting_diff_merged.id_to_index("a")] + assert not is_merged_diff[sorting_diff_merged.id_to_index("c")] + + +def test_set_properties_after_splits(): + times = np.array([0, 10, 20, 30, 40]) + labels = np.array(["a", "b", "b", "c", "c"]) + sorting = NumpySorting.from_samples_and_labels([times], [labels], 10_000.0, unit_ids=["a", "b", "c"]) + sorting.set_property("quality", np.array([1.0, 2.0, 3.0])) + sorting.set_property("group", np.array(["g1", "g2", "g3"])) + + # --- append strategy: split "b" into two new units --- + unit_splits = {"b": [np.array([0]), np.array([1])]} + sorting_split, new_unit_ids = apply_splits_to_sorting( + sorting, unit_splits, new_id_strategy="append", return_extra=True + ) + is_split = sorting_split.get_property("is_split") + # new units for "b" → is_split=True; "a" and "c" kept → is_split=False + for new_uid in new_unit_ids[0]: + assert is_split[sorting_split.id_to_index(new_uid)] + assert not is_split[sorting_split.id_to_index("a")] + assert not is_split[sorting_split.id_to_index("c")] + # quality of "b" (2.0) propagated to both sub-units + quality = sorting_split.get_property("quality") + for new_uid in new_unit_ids[0]: + assert quality[sorting_split.id_to_index(new_uid)] == 2.0 + + # --- split strategy with str unit ids --- + sorting_split2, new_unit_ids2 = apply_splits_to_sorting( + sorting, unit_splits, new_id_strategy="split", return_extra=True + ) + is_split2 = sorting_split2.get_property("is_split") + # new units are "b-0" and "b-1" + for new_uid in new_unit_ids2[0]: + assert is_split2[sorting_split2.id_to_index(new_uid)] + assert not is_split2[sorting_split2.id_to_index("a")] + assert not is_split2[sorting_split2.id_to_index("c")] + + # --- edge case: call set_properties_after_splits directly with a new split unit id + # that overlaps with an existing pre-split unit id (defensive fix) --- + # Manually build post-split sorting: "b" was split into ["a_new", "x"] but we + # simulate a hypothetical overlap by calling set_properties_after_splits directly + # with new_unit_ids containing a pre-existing id. + times2 = np.array([0, 10, 20, 30]) + labels2 = np.array(["a", "x", "x", "c"]) + sorting_post = NumpySorting.from_samples_and_labels([times2], [labels2], 10_000.0, unit_ids=["a", "x", "c"]) + sorting_post.set_property("quality", np.empty(3)) + # "x" is a new split sub-unit of "b"; "a" and "c" are kept + # In this case none of the new split unit ids are in pre_unit_ids, so untouched_unit_ids=["a","c"] + set_properties_after_splits( + sorting_post, + sorting, + split_unit_ids=["b"], + new_unit_ids=[["x", "c"]], # "c" overlaps with pre_unit_ids — the edge case + ) + is_split3 = sorting_post.get_property("is_split") + # "c" is a new split unit here (overlapping id), so it must be is_split=True + assert is_split3[sorting_post.id_to_index("c")] + assert is_split3[sorting_post.id_to_index("x")] + assert not is_split3[sorting_post.id_to_index("a")] + + def test_remap_unit_indices_in_vector(): unit_ids = ["a", "b", "c", "d", "e"] diff --git a/src/spikeinterface/metrics/quality/misc_metrics.py b/src/spikeinterface/metrics/quality/misc_metrics.py index c8b8f33a08..304d25bdbb 100644 --- a/src/spikeinterface/metrics/quality/misc_metrics.py +++ b/src/spikeinterface/metrics/quality/misc_metrics.py @@ -1741,7 +1741,7 @@ def slidingRP_violations( test_rp_centers_mask = rp_centers > exclude_ref_period_below_ms / 1000.0 # (in seconds) # only test for refractory period durations greater than 'exclude_ref_period_below_ms' - inds_confidence90 = np.row_stack(np.where(conf_matrix[:, test_rp_centers_mask] > 0.9)) + inds_confidence90 = np.vstack(np.where(conf_matrix[:, test_rp_centers_mask] > 0.9)) if len(inds_confidence90[0]) > 0: minI = np.min(inds_confidence90[0][0])