diff --git a/src/spikeinterface/postprocessing/template_similarity.py b/src/spikeinterface/postprocessing/template_similarity.py index 3db365b95d..be91eccfcb 100644 --- a/src/spikeinterface/postprocessing/template_similarity.py +++ b/src/spikeinterface/postprocessing/template_similarity.py @@ -296,13 +296,11 @@ def _compute_similarity_matrix_numba( elif method == "cosine": metric = 2 + overlapping_i_list = typed.List() overlapping_j_list = typed.List() active_channels_list = typed.List() for src_unit in range(num_templates): - overlapping_ids = typed.List() - overlapping_chs = typed.List() - start = src_unit if same_array else 0 for tgt_unit in range(start, other_num_templates): @@ -328,11 +326,9 @@ def _compute_similarity_matrix_numba( ch = np.arange(num_channels, dtype=np.uint16) if len(ch) > 0: - overlapping_ids.append(np.uint16(tgt_unit)) - overlapping_chs.append(ch) - - overlapping_j_list.append(overlapping_ids) - active_channels_list.append(overlapping_chs) + overlapping_i_list.append(np.int64(src_unit)) + overlapping_j_list.append(np.int64(tgt_unit)) + active_channels_list.append(ch) for count in range(len(shift_loop)): shift = shift_loop[count] @@ -340,44 +336,40 @@ def _compute_similarity_matrix_numba( src_sliced = templates_array[:, num_shifts : num_samples - num_shifts] tgt_sliced = other_templates_array[:, num_shifts + shift : num_samples - num_shifts + shift] - for i in prange(num_templates): - i_ = np.int64(i) - src_template = src_sliced[i_] - overlapping_ids = overlapping_j_list[i_] - overlapping_chs = active_channels_list[i_] - - for pair_idx in range(len(overlapping_ids)): - j = np.uint16(overlapping_ids[pair_idx]) - ch = overlapping_chs[pair_idx] - - src_ch = src_template[:, ch] - tgt_ch = tgt_sliced[j][:, ch] - - if metric == 0: - # l1 - norm_i = np.sum(np.abs(src_ch)) - norm_j = np.sum(np.abs(tgt_ch)) - dist = np.sum(np.abs(src_ch - tgt_ch)) - distances[count, i, j] = dist / (norm_i + norm_j) - - elif metric == 1: - # l2 - norm_i = sqrt(np.sum(src_ch**2)) - norm_j = sqrt(np.sum(tgt_ch**2)) - dist = sqrt(np.sum((src_ch - tgt_ch) ** 2)) - distances[count, i, j] = dist / (norm_i + norm_j) - - elif metric == 2: - # cosine - dot = np.sum(src_ch * tgt_ch) - norm_i = sqrt(np.sum(src_ch**2)) - norm_j = sqrt(np.sum(tgt_ch**2)) - denom = norm_i * norm_j - if denom > 0.0: - distances[count, i, j] = 1.0 - dot / denom - - if same_array: - distances[count, j, i] = distances[count, i, j] + for pair_idx in prange(len(overlapping_j_list)): + pair_idx_ = np.int64(pair_idx) + i = overlapping_i_list[pair_idx_] + j = overlapping_j_list[pair_idx_] + active_channels = active_channels_list[pair_idx_] + + src_ch = src_sliced[i][:, active_channels] + tgt_ch = tgt_sliced[j][:, active_channels] + + if metric == 0: + # l1 + norm_i = np.sum(np.abs(src_ch)) + norm_j = np.sum(np.abs(tgt_ch)) + dist = np.sum(np.abs(src_ch - tgt_ch)) + distances[count, i, j] = dist / (norm_i + norm_j) + + elif metric == 1: + # l2 + norm_i = sqrt(np.sum(src_ch**2)) + norm_j = sqrt(np.sum(tgt_ch**2)) + dist = sqrt(np.sum((src_ch - tgt_ch) ** 2)) + distances[count, i, j] = dist / (norm_i + norm_j) + + elif metric == 2: + # cosine + dot = np.sum(src_ch * tgt_ch) + norm_i = sqrt(np.sum(src_ch**2)) + norm_j = sqrt(np.sum(tgt_ch**2)) + denom = norm_i * norm_j + if denom > 0.0: + distances[count, i, j] = 1.0 - dot / denom + + if same_array: + distances[count, j, i] = distances[count, i, j] if same_array and shift != 0: distances[num_shifts_both_sides - count - 1] = distances[count].T diff --git a/src/spikeinterface/postprocessing/tests/test_template_similarity.py b/src/spikeinterface/postprocessing/tests/test_template_similarity.py index 9fa7a73fec..fd963fa013 100644 --- a/src/spikeinterface/postprocessing/tests/test_template_similarity.py +++ b/src/spikeinterface/postprocessing/tests/test_template_similarity.py @@ -129,6 +129,74 @@ def test_equal_results_numba(params): assert np.allclose(result_numpy, result_numba, 1e-3) +def test_equal_results_numba_asymmetric_sparse(): + rng = np.random.default_rng(seed=2205) + templates_array = rng.random(size=(2, 20, 8), dtype=np.float32) + other_templates_array = rng.random(size=(7, 20, 8), dtype=np.float32) + sparsity_mask = rng.random((2, 8)) > 0.4 + other_sparsity_mask = rng.random((7, 8)) > 0.4 + sparsity_mask[0] = np.array([True, True, False, False, False, False, False, False]) + other_sparsity_mask[0] = np.array([False, False, True, True, False, False, False, False]) + + for method in ("cosine", "l1", "l2"): + for support in ("dense", "union", "intersection"): + result_numba = _compute_similarity_matrix_numba( + templates_array, + other_templates_array, + num_shifts=3, + method=method, + sparsity_mask=sparsity_mask, + other_sparsity_mask=other_sparsity_mask, + support=support, + ) + result_numpy = _compute_similarity_matrix_numpy( + templates_array, + other_templates_array, + num_shifts=3, + method=method, + sparsity_mask=sparsity_mask, + other_sparsity_mask=other_sparsity_mask, + support=support, + ) + + np.testing.assert_allclose(result_numba, result_numpy, rtol=1e-5, atol=1e-5) + + +def test_equal_results_numba_same_array_sparse(): + # templates_array is passed as both sides so same_array=True, exercising the + # pair-list restriction to tgt_unit >= src_unit and the diagonal mirroring, + # which the asymmetric test above cannot reach. + rng = np.random.default_rng(seed=2205) + templates_array = rng.random(size=(6, 20, 8), dtype=np.float32) + sparsity_mask = rng.random((6, 8)) > 0.4 + sparsity_mask[0] = np.array([True, True, False, False, False, False, False, False]) + sparsity_mask[1] = np.array([False, False, True, True, False, False, False, False]) + + for method in ("cosine", "l1", "l2"): + for support in ("dense", "union", "intersection"): + for num_shifts in (0, 3): + result_numba = _compute_similarity_matrix_numba( + templates_array, + templates_array, + num_shifts=num_shifts, + method=method, + sparsity_mask=sparsity_mask, + other_sparsity_mask=sparsity_mask, + support=support, + ) + result_numpy = _compute_similarity_matrix_numpy( + templates_array, + templates_array, + num_shifts=num_shifts, + method=method, + sparsity_mask=sparsity_mask, + other_sparsity_mask=sparsity_mask, + support=support, + ) + + np.testing.assert_allclose(result_numba, result_numpy, rtol=1e-5, atol=1e-5) + + if __name__ == "__main__": from spikeinterface.postprocessing.tests.common_extension_tests import get_dataset from spikeinterface.core import estimate_sparsity