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
135 changes: 135 additions & 0 deletions src/spikeinterface/exporters/tests/test_export_to_phy.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
import shutil

import numpy as np
import pytest

from spikeinterface.core import generate_ground_truth_recording, create_sorting_analyzer
from spikeinterface.exporters import export_to_phy

from spikeinterface.exporters.tests.common import (
Expand Down Expand Up @@ -153,6 +155,139 @@ def test_export_to_phy_metrics(sorting_analyzer_sparse_for_export, create_cache_
assert not (output_folder / f"cluster_{col_name}.tsv").is_file()


def _check_spikes_subset_files(output_folder, sorting_analyzer):
spikes_file = output_folder / "_phy_spikes_subset.spikes.npy"
channels_file = output_folder / "_phy_spikes_subset.channels.npy"
waveforms_file = output_folder / "_phy_spikes_subset.waveforms.npy"
for f in (spikes_file, channels_file, waveforms_file):
assert f.is_file()

subset_spikes = np.load(spikes_file)
subset_channels = np.load(channels_file)
subset_waveforms = np.load(waveforms_file)

waveforms_ext = sorting_analyzer.get_extension("waveforms")
n_subset = sorting_analyzer.get_extension("random_spikes").get_data().size

assert subset_spikes.shape == (n_subset,)
assert subset_channels.shape[0] == n_subset
assert subset_waveforms.shape[0] == n_subset
assert subset_waveforms.shape[1] == waveforms_ext.nbefore + waveforms_ext.nafter
assert subset_spikes.dtype == np.int64
assert subset_channels.dtype == np.int32

# spike indices are sorted and point to valid rows of spike_times.npy
spike_times = np.load(output_folder / "spike_times.npy")
assert np.all(np.diff(subset_spikes) >= 0)
assert subset_spikes.min() >= 0
assert subset_spikes.max() < spike_times.shape[0]

# channel indices are either -1 (padding) or a valid channel index
num_chans = sorting_analyzer.get_num_channels()
assert np.all((subset_channels == -1) | ((subset_channels >= 0) & (subset_channels < num_chans)))

# if sparsity is saved, each spike's channels must match its cluster's sparse channel set
# (dense exports don't write "template_ind.npy" since every unit uses all channels)
template_ind_file = output_folder / "template_ind.npy"
if template_ind_file.is_file():
spike_clusters = np.load(output_folder / "spike_clusters.npy")[:, 0]
template_ind = np.load(template_ind_file)
expected_channels = template_ind[spike_clusters[subset_spikes]]
np.testing.assert_array_equal(subset_channels, expected_channels)

# waveforms on padded (-1) channels must be zero
waveforms_by_channel = np.moveaxis(subset_waveforms, 1, 2) # (n_subset, n_channels, n_samples)
assert np.all(waveforms_by_channel[subset_channels == -1] == 0)


def test_export_to_phy_add_waveforms_sparse(sorting_analyzer_sparse_for_export, create_cache_folder):
cache_folder = create_cache_folder
output_folder = cache_folder / "phy_output_add_waveforms_sparse"
if output_folder.is_dir():
shutil.rmtree(output_folder)

sorting_analyzer = sorting_analyzer_sparse_for_export
assert sorting_analyzer.has_extension("waveforms")

export_to_phy(
sorting_analyzer,
output_folder,
compute_pc_features=False,
compute_amplitudes=False,
add_waveforms=True,
n_jobs=1,
chunk_size=10000,
progress_bar=False,
)

_check_spikes_subset_files(output_folder, sorting_analyzer)


def test_export_to_phy_add_waveforms_dense(sorting_analyzer_dense_for_export, create_cache_folder):
cache_folder = create_cache_folder
output_folder = cache_folder / "phy_output_add_waveforms_dense"
if output_folder.is_dir():
shutil.rmtree(output_folder)

sorting_analyzer = sorting_analyzer_dense_for_export
assert sorting_analyzer.has_extension("waveforms")

export_to_phy(
sorting_analyzer,
output_folder,
compute_pc_features=False,
compute_amplitudes=False,
add_waveforms=True,
n_jobs=1,
chunk_size=10000,
progress_bar=False,
)

_check_spikes_subset_files(output_folder, sorting_analyzer)


def test_export_to_phy_add_waveforms_missing_extension(create_cache_folder):
cache_folder = create_cache_folder
output_folder = cache_folder / "phy_output_add_waveforms_missing"
if output_folder.is_dir():
shutil.rmtree(output_folder)

recording, sorting = generate_ground_truth_recording(
durations=[10.0],
sampling_frequency=28000.0,
num_channels=8,
num_units=4,
generate_sorting_kwargs=dict(firing_rates=10.0, refractory_period_ms=4.0),
noise_kwargs=dict(noise_levels=5.0, strategy="on_the_fly"),
seed=2205,
)
sorting_analyzer = create_sorting_analyzer(sorting=sorting, recording=recording, format="memory", sparse=True)
# "random_spikes" is needed to compute "templates" without the "waveforms" extension
sorting_analyzer.compute("random_spikes")
sorting_analyzer.compute("templates")
sorting_analyzer.compute("template_similarity")
assert not sorting_analyzer.has_extension("waveforms")

# add_waveforms requires "waveforms"/"random_spikes" to already be computed: since
# "waveforms" is missing, it should warn and skip saving the spikes-subset files
with pytest.warns(UserWarning, match="Cannot save the phy spikes-waveforms subset"):
export_to_phy(
sorting_analyzer,
output_folder,
compute_pc_features=False,
compute_amplitudes=False,
add_waveforms=True,
n_jobs=1,
chunk_size=10000,
progress_bar=False,
)

assert not sorting_analyzer.has_extension("waveforms")
assert not (output_folder / "_phy_spikes_subset.spikes.npy").is_file()
assert not (output_folder / "_phy_spikes_subset.channels.npy").is_file()
assert not (output_folder / "_phy_spikes_subset.waveforms.npy").is_file()


if __name__ == "__main__":
sorting_analyzer_sparse = make_sorting_analyzer(sparse=True)
sorting_analyzer_group = make_sorting_analyzer(sparse=False, with_group=True)
Expand Down
61 changes: 61 additions & 0 deletions src/spikeinterface/exporters/to_phy.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ def export_to_phy(
output_folder: str | Path,
compute_pc_features: bool = True,
compute_amplitudes: bool = True,
add_waveforms: bool = False,
sparsity: Optional[ChannelSparsity] = None,
copy_binary: bool = True,
remove_if_exists: bool = False,
Expand All @@ -46,6 +47,13 @@ def export_to_phy(
If True, pc features are computed
compute_amplitudes : bool, default: True
If True, waveforms amplitudes are computed
add_waveforms : bool, default: False
If True, a subset of spike waveforms is saved in the phy "_phy_spikes_subset.*" files
(mimicking the "phy extract-waveforms" command), so that phy/phy-lib can display
individual spike waveforms without needing the raw recording. This requires the
"waveforms" and "random_spikes" SortingAnalyzer extensions to already be computed
(e.g. via `sorting_analyzer.compute(["random_spikes", "waveforms"])`); if they are not,
a warning is raised and the files are not saved.
sparsity : ChannelSparsity or None, default: None
The sparsity object
copy_binary : bool, default: True
Expand Down Expand Up @@ -235,6 +243,59 @@ def export_to_phy(
pc_feature_ind[unit_ind, : len(chan_inds)] = chan_inds
np.save(str(output_folder / "pc_feature_ind.npy"), pc_feature_ind)

if add_waveforms:
if not sorting_analyzer.has_extension("waveforms") or not sorting_analyzer.has_extension("random_spikes"):
warnings.warn(
"Cannot save the phy spikes-waveforms subset: the 'waveforms' or 'random_spikes' "
"extension is not computed."
)
else:
random_spikes_ext = sorting_analyzer.get_extension("random_spikes")
waveforms_ext = sorting_analyzer.get_extension("waveforms")

some_spikes = random_spikes_ext.get_random_spikes()
subset_spike_indices = random_spikes_ext.get_data()
subset_waveforms = waveforms_ext.get_data()

# order by absolute spike index, matching the row order of spike_times.npy
order = np.argsort(subset_spike_indices, kind="stable")
subset_spike_indices = subset_spike_indices[order]
subset_unit_index = some_spikes["unit_index"][order]
subset_waveforms = subset_waveforms[order]

max_num_channels_wf = max(len(chan_inds) for chan_inds in sparse_dict.values())
num_subset_spikes = subset_waveforms.shape[0]
num_samples_wf = subset_waveforms.shape[1]

# phy expects -1 for "no channel" padding
spike_channels = -np.ones((num_subset_spikes, max_num_channels_wf), dtype="int32")
spike_waveforms = np.zeros(
(num_subset_spikes, num_samples_wf, max_num_channels_wf), dtype=subset_waveforms.dtype
)

# if the "waveforms" extension itself is sparse, it already only stores (and orders)
# the channels of "sorting_analyzer.sparsity", which is the same as "used_sparsity"
# whenever the analyzer is sparse (see above). Otherwise, the extension is dense and
# we need to select the real channel indices for each unit.
waveforms_ext_is_sparse = waveforms_ext.sparsity is not None
for unit_id in unit_ids:
unit_index = sorting.id_to_index(unit_id)
spike_mask = subset_unit_index == unit_index
if not np.any(spike_mask):
continue
chan_inds = sparse_dict[unit_id]
spike_channels[spike_mask, : len(chan_inds)] = chan_inds
if waveforms_ext_is_sparse:
spike_waveforms[spike_mask, :, : len(chan_inds)] = subset_waveforms[spike_mask][
:, :, : len(chan_inds)
]
else:
spike_waveforms[spike_mask, :, : len(chan_inds)] = subset_waveforms[spike_mask][:, :, chan_inds]

np.save(str(output_folder / "_phy_spikes_subset.spikes.npy"), subset_spike_indices.astype("int64"))
np.save(str(output_folder / "_phy_spikes_subset.channels.npy"), spike_channels)
np.save(str(output_folder / "_phy_spikes_subset.waveforms.npy"), spike_waveforms)

# Save .tsv metadata
cluster_group = pd.DataFrame(
{"cluster_id": [i for i in range(len(unit_ids))], "group": ["unsorted"] * len(unit_ids)}
Expand Down
Loading