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
84 changes: 38 additions & 46 deletions src/spikeinterface/postprocessing/template_similarity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):

Expand All @@ -328,56 +326,50 @@ 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]

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