Skip to content
Merged
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
2 changes: 1 addition & 1 deletion src/spikeinterface/core/analyzer_extension_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -1618,7 +1618,7 @@ def _get_data(self, outputs="numpy", concatenated=False, return_data_name=None,
sorting = self.sorting_analyzer.sorting

if outputs == "numpy":
if copy:
if copy and not self.sorting_analyzer._lazy:
return all_data.copy() # return a copy to avoid modification
else:
return all_data
Expand Down
87 changes: 65 additions & 22 deletions src/spikeinterface/core/sortinganalyzer.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,7 +357,9 @@ def create_sorting_analyzer(
return sorting_analyzer


def load_sorting_analyzer(folder, load_extensions=True, format="auto", backend_options=None) -> "SortingAnalyzer":
def load_sorting_analyzer(
folder, load_extensions=True, format="auto", backend_options=None, lazy=False
) -> "SortingAnalyzer":
"""
Load a SortingAnalyzer object from disk.

Expand Down Expand Up @@ -385,7 +387,9 @@ def load_sorting_analyzer(folder, load_extensions=True, format="auto", backend_o
The loaded SortingAnalyzer

"""
return SortingAnalyzer.load(folder, load_extensions=load_extensions, format=format, backend_options=backend_options)
return SortingAnalyzer.load(
folder, load_extensions=load_extensions, format=format, backend_options=backend_options, lazy=lazy
)


class SortingAnalyzer:
Expand Down Expand Up @@ -421,6 +425,7 @@ def __init__(
peak_sign: PeakSignType = "both",
peak_mode: PeakModeType = "extremum",
backend_options: dict | None = None,
lazy: bool = False,
):
# very fast init because checks are done in load and create
self.sorting = sorting
Expand Down Expand Up @@ -449,6 +454,9 @@ def __init__(
# (additional saving options for creating and saving datasets, e.g. compression/filters for zarr)
self._backend_options = {} if backend_options is None else backend_options

# the lazy flag is used to load the extensions in a lazy way (only when needed)
self._lazy = lazy

# extensions are not loaded at init
self.extensions = dict()

Expand Down Expand Up @@ -581,6 +589,7 @@ def load(
load_extensions: bool = True,
format: Literal["auto", "binary_folder", "zarr"] = "auto",
backend_options: dict | None = None,
lazy: bool = False,
):
"""
Load folder or zarr.
Expand All @@ -594,16 +603,16 @@ def load(

if format == "binary_folder":
sorting_analyzer = SortingAnalyzer.load_from_binary_folder(
folder, recording=recording, backend_options=backend_options
folder, recording=recording, backend_options=backend_options, lazy=lazy
)
elif format == "zarr":
sorting_analyzer = SortingAnalyzer.load_from_zarr(
folder, recording=recording, backend_options=backend_options
folder, recording=recording, backend_options=backend_options, lazy=lazy
)
else:
raise ValueError(f"SortingAnalyzer.load: wrong format {format}")

if load_extensions and not is_path_remote(folder):
if load_extensions and not lazy and not is_path_remote(folder):
sorting_analyzer.load_all_saved_extension()

return sorting_analyzer
Expand Down Expand Up @@ -885,6 +894,7 @@ def load_from_binary_folder(
folder: str | Path,
recording: BaseRecording | None = None,
backend_options: dict | None = None,
lazy: bool = False,
) -> "SortingAnalyzer":
from .loading import load

Expand Down Expand Up @@ -928,9 +938,18 @@ def load_from_binary_folder(
with open(settings_file, "w") as f:
json.dump(check_json(settings), f, indent=4)

# Load sorting (in memory)
# Load sorting (in memory or lazy)
if lazy:
numpy_folder_kwargs = dict(mmap_mode="r")
copy_spike_vector = False
else:
numpy_folder_kwargs = dict()
copy_spike_vector = True

sorting = NumpySorting.from_sorting(
NumpyFolderSorting(sorting_folder), with_metadata=True, copy_spike_vector=True
NumpyFolderSorting(folder / "sorting", **numpy_folder_kwargs),
with_metadata=True,
copy_spike_vector=copy_spike_vector,
)

# Load recording (if available)
Expand Down Expand Up @@ -970,6 +989,7 @@ def load_from_binary_folder(
peak_sign=settings["peak_sign"],
peak_mode=settings["peak_mode"],
backend_options=backend_options,
lazy=lazy,
)
sorting_analyzer.folder = folder

Expand Down Expand Up @@ -1088,6 +1108,7 @@ def load_from_zarr(
folder: str | Path,
recording: BaseRecording | None = None,
backend_options: dict | None = None,
lazy: bool = False,
) -> "SortingAnalyzer":
import zarr
from .loading import load
Expand Down Expand Up @@ -1117,11 +1138,22 @@ def load_from_zarr(
settings = zarr_root.attrs["settings"]
settings = cls._handle_backward_compatibility_settings_pre_init(settings)

# Load sorting (in memory)
# Load sorting (in memory or lazy)
if lazy:
copy_spike_vector = False
lazy_spike_vector = True
else:
copy_spike_vector = True
lazy_spike_vector = False
sorting = NumpySorting.from_sorting(
ZarrSortingExtractor(folder, zarr_group="sorting", storage_options=storage_options),
ZarrSortingExtractor(
folder,
zarr_group="sorting",
storage_options=storage_options,
lazy_spike_vector=lazy_spike_vector,
),
with_metadata=True,
copy_spike_vector=True,
copy_spike_vector=copy_spike_vector,
)

# Load recording (if available)
Expand Down Expand Up @@ -1161,6 +1193,7 @@ def load_from_zarr(
peak_sign=settings["peak_sign"],
peak_mode=settings["peak_mode"],
backend_options=backend_options,
lazy=lazy,
)
sorting_analyzer.folder = folder

Expand Down Expand Up @@ -1486,6 +1519,11 @@ def _save_or_select_or_merge_or_split(
new_sorting_analyzer : SortingAnalyzer
The newly created SortingAnalyzer object.
"""
if self._lazy:
raise ValueError(
"Cannot save, select, merge or split units when the SortingAnalyzer is lazy. "
"Please load the SortingAnalyzer with lazy=False."
)
if self.has_recording():
recording = self._recording
elif self.has_temporary_recording():
Expand Down Expand Up @@ -2212,6 +2250,10 @@ def compute(self, input, save=True, extension_params=None, verbose=False, **kwar
)

"""
if self._lazy:
# If the analyzer is lazy, we can compute extensions in memory but we won't save / overwrite any existing
# extension on disk. This is to avoid overwriting existing extensions when the analyzer is lazy.
save = False
if isinstance(input, str):
return self.compute_one_extension(extension_name=input, save=save, verbose=verbose, **kwargs)
elif isinstance(input, dict):
Expand Down Expand Up @@ -2508,7 +2550,7 @@ def load_extension(self, extension_name: str):
if extension_class is None:
return None

extension_instance = extension_class.load(self)
extension_instance = extension_class.load(self, lazy=self._lazy)

self.extensions[extension_name] = extension_instance

Expand All @@ -2527,7 +2569,7 @@ def delete_extension(self, extension_name) -> None:
"""

# delete from folder or zarr
if self.format != "memory" and self.has_extension(extension_name):
if self.format != "memory" and self.has_extension(extension_name) and not self._lazy:
# need a reload to reset the folder
ext = self.load_extension(extension_name)
ext.delete()
Expand Down Expand Up @@ -2990,20 +3032,20 @@ def _get_zarr_extension_group(self, mode="r+"):
return extension_group

@classmethod
def load(cls, sorting_analyzer):
def load(cls, sorting_analyzer, lazy=False):
ext = cls(sorting_analyzer)
ext.load_params()
ext.load_run_info()
if ext.run_info is not None:
if ext.run_info["run_completed"]:
ext.load_data()
ext.load_data(lazy=lazy)
if cls.need_backward_compatibility_on_load:
ext._handle_backward_compatibility_on_load()
if len(ext.data) > 0:
return ext
else:
# this is for back-compatibility of old analyzers
ext.load_data()
ext.load_data(lazy=lazy)
if cls.need_backward_compatibility_on_load:
ext._handle_backward_compatibility_on_load()
if len(ext.data) > 0:
Expand Down Expand Up @@ -3103,7 +3145,7 @@ def load_params(self):

self.params = params

def load_data(self):
def load_data(self, lazy=False):
ext_data = None
if self.format == "binary_folder":
extension_folder = self._get_binary_extension_folder()
Expand All @@ -3123,10 +3165,12 @@ def load_data(self):
ext_data = json.load(f)
elif ext_data_file.suffix == ".npy":
# The lazy loading of an extension is complicated because if we compute again
# and have a link to the old buffer on windows then it fails
# ext_data = np.load(ext_data_file, mmap_mode="r")
# so we go back to full loading
ext_data = np.load(ext_data_file)
# and have a link to the old buffer on windows then it fails.
# So, by default, we use full loading, but lazy can be requested on demand.
if lazy:
ext_data = np.load(ext_data_file, mmap_mode="r")
else:
ext_data = np.load(ext_data_file)
elif ext_data_file.suffix == ".csv":
import pandas as pd

Expand Down Expand Up @@ -3162,8 +3206,7 @@ def load_data(self):
elif "object" in ext_data_.attrs:
ext_data = ext_data_[0]
else:
# this load in memory
ext_data = np.array(ext_data_)
ext_data = ext_data_ if lazy else np.array(ext_data_[:])
self.set_data(ext_data_name, ext_data)

if len(self.data) == 0:
Expand Down
9 changes: 4 additions & 5 deletions src/spikeinterface/core/sortingfolder.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ class NumpyFolderSorting(BaseSorting):
mode = "folder"
name = "NumpyFolder"

def __init__(self, folder_path: str | Path):
def __init__(self, folder_path, mmap_mode: str | None = None):
folder_path = Path(folder_path)

# Load general info
Expand All @@ -37,8 +37,8 @@ def __init__(self, folder_path: str | Path):
# Init superclass
super().__init__(sampling_frequency, unit_ids)

# Load spikes vector
self.spikes = np.load(folder_path / "spikes.npy")
self.spikes = np.load(folder_path / "spikes.npy", mmap_mode=mmap_mode)

for segment_index in range(num_segments):
self.add_sorting_segment(SpikeVectorSortingSegment(self.spikes, segment_index, unit_ids))
# important trick : the cache is already spikes vector
Expand All @@ -47,8 +47,7 @@ def __init__(self, folder_path: str | Path):
# Load metadata
self.load_metadata_from_folder(folder_path)

# Save folder_path as kwargs for serialization
self._kwargs = {"folder_path": str(folder_path.absolute())}
self._kwargs = dict(folder_path=str(folder_path.absolute()), mmap_mode=mmap_mode)
Comment thread
alejoe91 marked this conversation as resolved.
Comment thread
chrishalcrow marked this conversation as resolved.

@staticmethod
def write_sorting(sorting, save_path):
Expand Down
51 changes: 49 additions & 2 deletions src/spikeinterface/core/tests/test_sortinganalyzer.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ def test_SortingAnalyzer_binary_folder(tmp_path, dataset):
assert "number" in sorting_analyzer.sorting.get_property_keys()
sorting_analyzer_reloded = load_sorting_analyzer(folder, format="auto")
assert "quality" in sorting_analyzer_reloded.sorting.get_property_keys()
assert "number" in sorting_analyzer.sorting.get_property_keys()
assert "number" in sorting_analyzer_reloded.sorting.get_property_keys()


def test_SortingAnalyzer_zarr(tmp_path, dataset):
Expand Down Expand Up @@ -213,7 +213,7 @@ def test_SortingAnalyzer_zarr(tmp_path, dataset):
assert "number" in sorting_analyzer.sorting.get_property_keys()
sorting_analyzer_reloded = load_sorting_analyzer(sorting_analyzer.folder, format="auto")
assert "quality" in sorting_analyzer_reloded.sorting.get_property_keys()
assert "number" in sorting_analyzer.sorting.get_property_keys()
assert "number" in sorting_analyzer_reloded.sorting.get_property_keys()


def test_create_by_dict():
Expand Down Expand Up @@ -361,6 +361,53 @@ def test_SortingAnalyzer_interleaved_probegroup(dataset):
assert np.array_equal(recording.get_channel_locations(), sorting_analyzer.get_channel_locations())


@pytest.mark.parametrize("format", ["binary_folder", "zarr"])
def test_load_in_lazy_mode(tmp_path, dataset, format):
recording, sorting = dataset

folder = tmp_path / "test_SortingAnalyzer_folder"
if format == "zarr":
import zarr
from spikeinterface.core.zarrextractors import ZarrSpikeVector

folder = folder.with_suffix(".zarr")
array_class = zarr.Array
spike_vector_class = ZarrSpikeVector
else:
array_class = np.memmap
spike_vector_class = np.memmap
if folder.exists():
shutil.rmtree(folder)

sorting_analyzer = create_sorting_analyzer(
sorting, recording, format=format, folder=folder, sparse=False, sparsity=None
)

sorting_analyzer.compute(["random_spikes", "templates", "spike_amplitudes"])
# load in lazy mode and check that spike vector and extension data are memmap
sorting_analyzer_lazy = load_sorting_analyzer(folder, format="auto", lazy=True)

assert isinstance(sorting_analyzer_lazy.sorting.to_spike_vector(), spike_vector_class)

template_ext = sorting_analyzer_lazy.get_extension("templates")
template_data = template_ext.data
for key, value in template_data.items():
if isinstance(value, np.ndarray):
assert isinstance(value, array_class)
spike_amplitudes_ext = sorting_analyzer_lazy.get_extension("spike_amplitudes")
spike_amplitudes_data = spike_amplitudes_ext.data
for key, value in spike_amplitudes_data.items():
if isinstance(value, np.ndarray):
assert isinstance(value, array_class)

# check that the lazy mode does not overwrite existing extensions
sorting_analyzer_lazy.compute("random_spikes", max_spikes_per_unit=10)
# reload the analyzer to check that the original extension is not overwritten
sorting_analyzer_reloaded = load_sorting_analyzer(folder, format="auto", lazy=True)
random_spikes_ext = sorting_analyzer_reloaded.get_extension("random_spikes")
assert random_spikes_ext.params["max_spikes_per_unit"] != 10


def _check_sorting_analyzers(sorting_analyzer, original_sorting, cache_folder):

register_result_extension(DummyAnalyzerExtension)
Expand Down
Loading
Loading