Skip to content

Commit 60cbce4

Browse files
committed
allow for concatenated mode in get_some_projections
1 parent c5d3055 commit 60cbce4

3 files changed

Lines changed: 16 additions & 2 deletions

File tree

src/spikeinterface/postprocessing/principal_component.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -260,7 +260,11 @@ def get_some_projections(self, channel_ids=None, unit_ids=None):
260260
spike_unit_indices = some_spikes["unit_index"][selected_inds]
261261

262262
if sparsity is None:
263-
some_projections = all_projections[selected_inds, :, :][:, :, channel_indices]
263+
if self.params["mode"] == "concatenated":
264+
some_projections = all_projections[selected_inds, :]
265+
else:
266+
some_projections = all_projections[selected_inds, :, :][:, :, channel_indices]
267+
264268
else:
265269
# need re-alignement
266270
some_projections = np.zeros((selected_inds.size, num_components, channel_indices.size), dtype=dtype)

src/spikeinterface/postprocessing/tests/test_principal_component.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,13 @@ def test_mode_concatenated(self):
5050
assert pca.ndim == 2
5151
assert pca.shape[1] == n_components
5252

53+
ext_rand = sorting_analyzer.get_extension("random_spikes")
54+
num_rand_spikes = len(ext_rand.get_data())
55+
56+
some_projections = ext.get_some_projections()
57+
assert some_projections[0].shape[0] == num_rand_spikes
58+
assert some_projections[0].shape[1] == n_components
59+
5360
@pytest.mark.parametrize("sparse", [True, False])
5461
def test_get_projections(self, sparse):
5562
"""

src/spikeinterface/qualitymetrics/pca_metrics.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -157,7 +157,10 @@ def compute_pc_metrics(
157157
neighbor_channel_indices = sorting_analyzer.channel_ids_to_indices(neighbor_channel_ids)
158158

159159
labels = all_labels[np.isin(all_labels, neighbor_unit_ids)]
160-
pcs = dense_projections[np.isin(all_labels, neighbor_unit_ids)][:, :, neighbor_channel_indices]
160+
if pca_ext.params["mode"] == "concatenated":
161+
pcs = dense_projections[np.isin(all_labels, neighbor_unit_ids)]
162+
else:
163+
pcs = dense_projections[np.isin(all_labels, neighbor_unit_ids)][:, :, neighbor_channel_indices]
161164
pcs_flat = pcs.reshape(pcs.shape[0], -1)
162165

163166
func_args = (pcs_flat, labels, non_nn_metrics, unit_id, unit_ids, metric_params, max_threads_per_worker)

0 commit comments

Comments
 (0)