From 3385cf26aab1208b3fb832c1d7c7a8bde412b645 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Mon, 15 Jun 2026 17:35:18 +0200 Subject: [PATCH 01/13] wip: lazy load (analyzer + extensions) --- src/spikeinterface/core/sortinganalyzer.py | 56 +++++++++++++++------- src/spikeinterface/core/sortingfolder.py | 6 +-- 2 files changed, 41 insertions(+), 21 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index b5885598fe..38ef7266f8 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -217,7 +217,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. @@ -245,7 +247,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: @@ -407,7 +411,7 @@ def create( return sorting_analyzer @classmethod - def load(cls, folder, recording=None, load_extensions=True, format="auto", backend_options=None): + def load(cls, folder, recording=None, load_extensions=True, format="auto", backend_options=None, lazy=False): """ Load folder or zarr. The recording can be given if the recording location has changed. @@ -422,11 +426,11 @@ def load(cls, folder, recording=None, load_extensions=True, format="auto", backe 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 ) if not is_path_remote(str(folder)): @@ -532,15 +536,23 @@ def create_binary_folder(cls, folder, sorting, recording, sparsity, return_in_uV return cls.load_from_binary_folder(folder, recording=recording, backend_options=backend_options) @classmethod - def load_from_binary_folder(cls, folder, recording=None, backend_options=None): + def load_from_binary_folder(cls, folder, recording=None, backend_options=None, lazy=False): from .loading import load folder = Path(folder) assert folder.is_dir(), f"This folder does not exists {folder}" # load internal sorting copy in memory + 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(folder / "sorting"), with_metadata=True, copy_spike_vector=True + NumpyFolderSorting(folder / "sorting", **numpy_folder_kwargs), + with_metadata=True, + copy_spike_vector=copy_spike_vector, ) # Try to load the recording if not provided @@ -698,7 +710,7 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_in_uV, rec_att return cls.load_from_zarr(folder, recording=recording, backend_options=backend_options) @classmethod - def load_from_zarr(cls, folder, recording=None, backend_options=None): + def load_from_zarr(cls, folder, recording=None, backend_options=None, lazy=False): import zarr from .loading import load @@ -722,6 +734,8 @@ def load_from_zarr(cls, folder, recording=None, backend_options=None): ) # load internal sorting in memory + if lazy: + copy_spike_vector = False sorting = NumpySorting.from_sorting( ZarrSortingExtractor(folder, zarr_group="sorting", storage_options=storage_options), with_metadata=True, @@ -1894,7 +1908,7 @@ def get_saved_extension_names(self): return saved_extension_names - def get_extension(self, extension_name: str): + def get_extension(self, extension_name: str, lazy: bool = False): """ Get a AnalyzerExtension. If not loaded then load is automatic. @@ -1906,13 +1920,13 @@ def get_extension(self, extension_name: str): return self.extensions[extension_name] elif self.format != "memory" and self.has_extension(extension_name): - self.load_extension(extension_name) + self.load_extension(extension_name, lazy=lazy) return self.extensions[extension_name] else: return None - def load_extension(self, extension_name: str): + def load_extension(self, extension_name: str, lazy: bool = False): """ Load an extension from a folder or zarr into the `ResultSorting.extensions` dict. @@ -1920,6 +1934,8 @@ def load_extension(self, extension_name: str): ---------- extension_name : str The extension name. + lazy : bool, default: False + If True, array data are not loaded in memory, but kept as memmap/zarr arrays Returns ------- @@ -1936,7 +1952,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=lazy) self.extensions[extension_name] = extension_instance @@ -2414,20 +2430,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: @@ -2527,7 +2543,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() @@ -2550,7 +2566,8 @@ def load_data(self): # 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) + kwargs = dict(mmap_mode="r") if lazy else dict() + ext_data = np.load(ext_data_file, **kwargs) elif ext_data_file.suffix == ".csv": import pandas as pd @@ -2587,7 +2604,10 @@ def load_data(self): ext_data = ext_data_[0] else: # this load in memmory - ext_data = np.array(ext_data_) + if lazy: + ext_data = ext_data_ + else: + ext_data = np.array(ext_data_) self.set_data(ext_data_name, ext_data) if len(self.data) == 0: diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index c0d66393d2..2dba9d4465 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -24,7 +24,7 @@ class NumpyFolderSorting(BaseSorting): mode = "folder" name = "NumpyFolder" - def __init__(self, folder_path): + def __init__(self, folder_path, mmap_mode=None): folder_path = Path(folder_path) with open(folder_path / "numpysorting_info.json", "r") as f: @@ -36,7 +36,7 @@ def __init__(self, folder_path): BaseSorting.__init__(self, sampling_frequency, unit_ids) - 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)) @@ -47,7 +47,7 @@ def __init__(self, folder_path): folder_metadata = folder_path self.load_metadata_from_folder(folder_metadata) - self._kwargs = dict(folder_path=str(folder_path.absolute())) + self._kwargs = dict(folder_path=str(folder_path.absolute()), mmap_mode=mmap_mode) @staticmethod def write_sorting(sorting, save_path): From 121a0559a91bc9f3003dbf5a36b422c0f1473c35 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 18 Jun 2026 19:09:51 -0600 Subject: [PATCH 02/13] todo --- src/spikeinterface/core/zarrextractors.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index bbc797c693..7e8dd72a47 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -289,6 +289,8 @@ def __init__(self, folder_path: Path | str, storage_options: dict | None = None, BaseSorting.__init__(self, sampling_frequency, unit_ids) + # TODO: make a virtual memmap view of the spike vector or override to_spike_vector to behave like + # a memmap spikes = np.zeros(len(spikes_group["sample_index"]), dtype=minimum_spike_dtype) spikes["sample_index"] = spikes_group["sample_index"][:] spikes["unit_index"] = spikes_group["unit_index"][:] From 6fdd1947a1c27eb7aa69658c14e66a208e632324 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 24 Jun 2026 15:46:24 +0200 Subject: [PATCH 03/13] feat: move lazy logic to load only --- src/spikeinterface/core/sortinganalyzer.py | 35 +++++++++++----------- 1 file changed, 17 insertions(+), 18 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 38ef7266f8..8814c30a36 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -283,6 +283,7 @@ def __init__( sparsity: ChannelSparsity | None = None, return_in_uV: bool = True, backend_options: dict | None = None, + lazy: bool = False, ): # very fast init because checks are done in load and create self.sorting = sorting @@ -308,6 +309,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() @@ -549,6 +553,7 @@ def load_from_binary_folder(cls, folder, recording=None, backend_options=None, l else: numpy_folder_kwargs = dict() copy_spike_vector = True + sorting = NumpySorting.from_sorting( NumpyFolderSorting(folder / "sorting", **numpy_folder_kwargs), with_metadata=True, @@ -613,6 +618,7 @@ def load_from_binary_folder(cls, folder, recording=None, backend_options=None, l sparsity=sparsity, return_in_uV=return_in_uV, backend_options=backend_options, + lazy=lazy, ) sorting_analyzer.folder = folder @@ -733,13 +739,12 @@ def load_from_zarr(cls, folder, recording=None, backend_options=None, lazy=False "Please consider re-generating the SortingAnalyzer object." ) - # load internal sorting in memory - if lazy: - copy_spike_vector = False + # TODO: make a Virtual memmap of ZarrSorting spike vector + copy_spike_vector = False if lazy else True sorting = NumpySorting.from_sorting( ZarrSortingExtractor(folder, zarr_group="sorting", storage_options=storage_options), with_metadata=True, - copy_spike_vector=True, + copy_spike_vector=copy_spike_vector, ) # load recording if possible @@ -784,6 +789,7 @@ def load_from_zarr(cls, folder, recording=None, backend_options=None, lazy=False sparsity=sparsity, return_in_uV=return_in_uV, backend_options=backend_options, + lazy=lazy, ) sorting_analyzer.folder = folder @@ -1908,7 +1914,7 @@ def get_saved_extension_names(self): return saved_extension_names - def get_extension(self, extension_name: str, lazy: bool = False): + def get_extension(self, extension_name: str): """ Get a AnalyzerExtension. If not loaded then load is automatic. @@ -1920,13 +1926,13 @@ def get_extension(self, extension_name: str, lazy: bool = False): return self.extensions[extension_name] elif self.format != "memory" and self.has_extension(extension_name): - self.load_extension(extension_name, lazy=lazy) + self.load_extension(extension_name) return self.extensions[extension_name] else: return None - def load_extension(self, extension_name: str, lazy: bool = False): + def load_extension(self, extension_name: str): """ Load an extension from a folder or zarr into the `ResultSorting.extensions` dict. @@ -1934,8 +1940,6 @@ def load_extension(self, extension_name: str, lazy: bool = False): ---------- extension_name : str The extension name. - lazy : bool, default: False - If True, array data are not loaded in memory, but kept as memmap/zarr arrays Returns ------- @@ -1952,7 +1956,7 @@ def load_extension(self, extension_name: str, lazy: bool = False): if extension_class is None: return None - extension_instance = extension_class.load(self, lazy=lazy) + extension_instance = extension_class.load(self, lazy=self._lazy) self.extensions[extension_name] = extension_instance @@ -2563,9 +2567,8 @@ def load_data(self, lazy=False): 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 + # 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. kwargs = dict(mmap_mode="r") if lazy else dict() ext_data = np.load(ext_data_file, **kwargs) elif ext_data_file.suffix == ".csv": @@ -2603,11 +2606,7 @@ def load_data(self, lazy=False): elif "object" in ext_data_.attrs: ext_data = ext_data_[0] else: - # this load in memmory - if lazy: - ext_data = ext_data_ - else: - 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: From 10dfd9f3060e0343461b26cb837ca6a1b6fe6482 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 24 Jun 2026 16:59:38 +0200 Subject: [PATCH 04/13] test: add tests for lazy mode --- src/spikeinterface/core/sortinganalyzer.py | 7 ++- .../core/tests/test_sortinganalyzer.py | 58 ++++++++++++++++++- 2 files changed, 62 insertions(+), 3 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 8814c30a36..3c8dae4188 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -437,7 +437,7 @@ def load(cls, folder, recording=None, load_extensions=True, format="auto", backe folder, recording=recording, backend_options=backend_options, lazy=lazy ) - if not is_path_remote(str(folder)): + if not is_path_remote(str(folder)) and not lazy: if load_extensions: sorting_analyzer.load_all_saved_extension() @@ -1008,6 +1008,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(): diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index a9bd71b5c0..43d4b153fb 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -119,7 +119,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): @@ -201,7 +201,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(): @@ -325,6 +325,60 @@ def test_SortingAnalyzer_interleaved_probegroup(dataset): assert np.array_equal(recording.get_channel_locations(), sorting_analyzer.get_channel_locations()) +def test_load_in_lazy_mode_binary(tmp_path, dataset): + recording, sorting = dataset + + folder = tmp_path / "test_SortingAnalyzer_binary_folder" + if folder.exists(): + shutil.rmtree(folder) + + sorting_analyzer = create_sorting_analyzer( + sorting, recording, format="binary_folder", folder=folder, sparse=False, sparsity=None + ) + + sorting_analyzer.compute(["random_spikes", "templates", "spike_amplitudes"]) + # load in lazy mode and check that extension data are memmap + sorting_analyzer_lazy = load_sorting_analyzer(folder, format="auto", lazy=True) + 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, np.memmap) + 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, np.memmap) + + +def test_load_in_lazy_mode_zarr(tmp_path, dataset): + import zarr + + recording, sorting = dataset + + folder = tmp_path / "test_SortingAnalyzer_zarr_folder.zarr" + if folder.exists(): + shutil.rmtree(folder) + + sorting_analyzer = create_sorting_analyzer( + sorting, recording, format="zarr", folder=folder, sparse=False, sparsity=None + ) + + sorting_analyzer.compute(["random_spikes", "templates", "spike_amplitudes"]) + # load in lazy mode and check that extension data are zarr arrays + sorting_analyzer_lazy = load_sorting_analyzer(folder, format="auto", lazy=True) + 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, zarr.Array) + 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, zarr.Array) + + def _check_sorting_analyzers(sorting_analyzer, original_sorting, cache_folder): register_result_extension(DummyAnalyzerExtension) From 21a6a596f0d999e7c05dd521603b450a1e45f839 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 24 Jun 2026 17:18:42 +0200 Subject: [PATCH 05/13] feat: implement ZarrSpikeVector - a memmap like lazy spike vector for zarr --- src/spikeinterface/core/sortinganalyzer.py | 16 +- .../core/tests/test_sortinganalyzer.py | 11 +- src/spikeinterface/core/zarrextractors.py | 140 ++++++++++++++++-- 3 files changed, 153 insertions(+), 14 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 3c8dae4188..1588de85bb 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -739,10 +739,20 @@ def load_from_zarr(cls, folder, recording=None, backend_options=None, lazy=False "Please consider re-generating the SortingAnalyzer object." ) - # TODO: make a Virtual memmap of ZarrSorting spike vector - copy_spike_vector = False if lazy else True + 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=copy_spike_vector, ) diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index 43d4b153fb..e0411bc9cd 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -337,8 +337,11 @@ def test_load_in_lazy_mode_binary(tmp_path, dataset): ) sorting_analyzer.compute(["random_spikes", "templates", "spike_amplitudes"]) - # load in lazy mode and check that extension data are memmap + # 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(), np.memmap) + template_ext = sorting_analyzer_lazy.get_extension("templates") template_data = template_ext.data for key, value in template_data.items(): @@ -353,6 +356,7 @@ def test_load_in_lazy_mode_binary(tmp_path, dataset): def test_load_in_lazy_mode_zarr(tmp_path, dataset): import zarr + from spikeinterface.core.zarrextractors import ZarrSpikeVector recording, sorting = dataset @@ -365,8 +369,11 @@ def test_load_in_lazy_mode_zarr(tmp_path, dataset): ) sorting_analyzer.compute(["random_spikes", "templates", "spike_amplitudes"]) - # load in lazy mode and check that extension data are zarr arrays + # load in lazy mode and check that spikevector is ZarrSpikeVector andextension data are zarr arrays sorting_analyzer_lazy = load_sorting_analyzer(folder, format="auto", lazy=True) + + assert isinstance(sorting_analyzer_lazy.sorting.to_spike_vector(), ZarrSpikeVector) + template_ext = sorting_analyzer_lazy.get_extension("templates") template_data = template_ext.data for key, value in template_data.items(): diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index 7e8dd72a47..e2f1b9cd7b 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -241,6 +241,112 @@ def get_traces( return traces +class _ZarrSegmentIndex: + """Lazy segment_index array derived from segment_slices stored in zarr.""" + + def __init__(self, segment_slices: np.ndarray, n: int): + self._segment_slices = segment_slices + self._n = n + + def __len__(self) -> int: + return self._n + + def __array__(self, dtype=None): + arr = np.empty(self._n, dtype="int64") + for seg_idx, (s0, s1) in enumerate(self._segment_slices): + arr[s0:s1] = seg_idx + return arr if dtype is None else arr.astype(dtype) + + def __getitem__(self, key): + return np.asarray(self)[key] + + def __eq__(self, other): + return np.asarray(self) == other + + +class ZarrSpikeVector: + """ + Virtual structured spike vector backed by zarr arrays. + + Mimics a memmap-backed numpy structured array with fields + (sample_index, unit_index, segment_index) without loading any data + at construction time. Data is read from zarr lazily: + + * Field access (``spikes["sample_index"]``) returns the zarr array + (or a lazy segment-index object). + * Slice access (``spikes[s0:s1]``) materialises only that slice. + * ``np.asarray(spikes)`` materialises the full array. + + The zarr arrays are assumed to be stored in sorted order + (segment_index ASC, sample_index ASC, unit_index ASC), which is the + ordering guaranteed by :func:`add_sorting_to_zarr_group`. + """ + + def __init__(self, spikes_group, segment_slices: np.ndarray): + self._sample_index = spikes_group["sample_index"] + self._unit_index = spikes_group["unit_index"] + self._segment_slices = np.asarray(segment_slices, dtype="int64") + self._n = len(self._sample_index) + self.dtype = np.dtype(minimum_spike_dtype) + + @property + def size(self) -> int: + return self._n + + def __len__(self) -> int: + return self._n + + def __getitem__(self, key): + if isinstance(key, str): + if key == "sample_index": + return self._sample_index + elif key == "unit_index": + return self._unit_index + elif key == "segment_index": + return _ZarrSegmentIndex(self._segment_slices, self._n) + else: + raise KeyError(f"ZarrSpikeVector has no field {key!r}") + + if isinstance(key, (int, np.integer)): + idx = int(key) + if idx < 0: + idx += self._n + result = np.empty(1, dtype=self.dtype) + result["sample_index"][0] = self._sample_index[idx] + result["unit_index"][0] = self._unit_index[idx] + result["segment_index"][0] = int(np.searchsorted(self._segment_slices[:, 0], idx, side="right")) - 1 + return result[0] + + if isinstance(key, slice): + start, stop, step = key.indices(self._n) + n = len(range(start, stop, step)) + result = np.empty(n, dtype=self.dtype) + result["sample_index"] = self._sample_index[start:stop:step] + result["unit_index"] = self._unit_index[start:stop:step] + if step == 1: + seg_index = np.empty(n, dtype="int64") + for seg_idx, (s0, s1) in enumerate(self._segment_slices): + lo = max(start, int(s0)) - start + hi = min(stop, int(s1)) - start + if hi > lo: + seg_index[lo:hi] = seg_idx + result["segment_index"] = seg_index + else: + result["segment_index"] = _ZarrSegmentIndex(self._segment_slices, self._n)[start:stop:step] + return result + + # fallback for fancy/boolean indexing: materialise then index + return np.asarray(self)[key] + + def __array__(self, dtype=None): + arr = np.empty(self._n, dtype=self.dtype) + arr["sample_index"] = self._sample_index[:] + arr["unit_index"] = self._unit_index[:] + for seg_idx, (s0, s1) in enumerate(self._segment_slices): + arr["segment_index"][s0:s1] = seg_idx + return arr if dtype is None else arr.astype(dtype) + + class ZarrSortingExtractor(BaseSorting): """ SortingExtractor for a zarr format @@ -257,13 +363,23 @@ class ZarrSortingExtractor(BaseSorting): Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc. zarr_group : str or None, default: None Optional zarr group path to load the sorting from. This can be used when the sorting is not stored at the root, but in sub group. + lazy_spike_vector : bool, default: False + If True, the spike vector is loaded lazily. This can be useful for large sortings with many spikes. + If False, the spike vector is loaded in memory. Default: False + Returns ------- sorting : ZarrSortingExtractor The sorting Extractor """ - def __init__(self, folder_path: Path | str, storage_options: dict | None = None, zarr_group: str | None = None): + def __init__( + self, + folder_path: Path | str, + storage_options: dict | None = None, + zarr_group: str | None = None, + lazy_spike_vector: bool = False, + ): folder_path, folder_path_kwarg = resolve_zarr_path(folder_path) @@ -289,15 +405,21 @@ def __init__(self, folder_path: Path | str, storage_options: dict | None = None, BaseSorting.__init__(self, sampling_frequency, unit_ids) - # TODO: make a virtual memmap view of the spike vector or override to_spike_vector to behave like - # a memmap - spikes = np.zeros(len(spikes_group["sample_index"]), dtype=minimum_spike_dtype) - spikes["sample_index"] = spikes_group["sample_index"][:] - spikes["unit_index"] = spikes_group["unit_index"][:] - for i, (start, end) in enumerate(segment_slices_list): - spikes["segment_index"][start:end] = i - spikes = spikes[np.lexsort((spikes["unit_index"], spikes["sample_index"], spikes["segment_index"]))] + if lazy_spike_vector: + spikes = ZarrSpikeVector(spikes_group, segment_slices_list) + else: + # Materialize the spike vector in memory and sort it by (segment_index, sample_index, unit_index) + spikes = np.zeros(len(spikes_group["sample_index"]), dtype=minimum_spike_dtype) + spikes["sample_index"] = spikes_group["sample_index"][:] + spikes["unit_index"] = spikes_group["unit_index"][:] + for i, (start, end) in enumerate(segment_slices_list): + spikes["segment_index"][start:end] = i + spikes = spikes[np.lexsort((spikes["unit_index"], spikes["sample_index"], spikes["segment_index"]))] + self._cached_spike_vector = spikes + # pre-populate segment slices so _get_spike_vector_segment_slices() never + # needs to materialise the full segment_index array + self._cached_spike_vector_segment_slices = np.asarray(segment_slices_list, dtype="int64") for segment_index in range(num_segments): soring_segment = SpikeVectorSortingSegment(spikes, segment_index, unit_ids) From e1a1e75a149264acfe93e01129b7ba459387e66a Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 24 Jun 2026 17:19:26 +0200 Subject: [PATCH 06/13] fix: add lazy spike vector as kwarg --- src/spikeinterface/core/zarrextractors.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index e2f1b9cd7b..a832129fe1 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -437,7 +437,12 @@ def __init__( if annotations is not None: self.annotate(**annotations) - self._kwargs = {"folder_path": folder_path_kwarg, "storage_options": storage_options, "zarr_group": zarr_group} + self._kwargs = { + "folder_path": folder_path_kwarg, + "storage_options": storage_options, + "zarr_group": zarr_group, + "lazy_spike_vector": lazy_spike_vector, + } @staticmethod def write_sorting(sorting: BaseSorting, folder_path: str | Path, storage_options: dict | None = None, **kwargs): From b710917e2aa5f1e2acc33d98287a454f2e7d482f Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 21 Jul 2026 12:17:14 +0200 Subject: [PATCH 07/13] fix: don't copy extension data if sorting_analyzer is lazy --- src/spikeinterface/core/analyzer_extension_core.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/spikeinterface/core/analyzer_extension_core.py b/src/spikeinterface/core/analyzer_extension_core.py index 261710278a..dde1ba8620 100644 --- a/src/spikeinterface/core/analyzer_extension_core.py +++ b/src/spikeinterface/core/analyzer_extension_core.py @@ -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 From 6fa122af4af8c355dae13a1ec6d5d0c0c4ff500c Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 21 Jul 2026 12:28:08 +0200 Subject: [PATCH 08/13] fix: save=False in compute if analyzer is lazy --- src/spikeinterface/core/sortinganalyzer.py | 4 ++ .../core/tests/test_sortinganalyzer.py | 62 +++++++------------ 2 files changed, 28 insertions(+), 38 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index c86d6da08c..a15128df40 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -2197,6 +2197,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): diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index 2938391122..25aeb78c1a 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -361,65 +361,51 @@ def test_SortingAnalyzer_interleaved_probegroup(dataset): assert np.array_equal(recording.get_channel_locations(), sorting_analyzer.get_channel_locations()) -def test_load_in_lazy_mode_binary(tmp_path, dataset): +@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_binary_folder" + 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="binary_folder", folder=folder, sparse=False, sparsity=None + 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(), np.memmap) - - 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, np.memmap) - 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, np.memmap) - - -def test_load_in_lazy_mode_zarr(tmp_path, dataset): - import zarr - from spikeinterface.core.zarrextractors import ZarrSpikeVector - - recording, sorting = dataset - - folder = tmp_path / "test_SortingAnalyzer_zarr_folder.zarr" - if folder.exists(): - shutil.rmtree(folder) - - sorting_analyzer = create_sorting_analyzer( - sorting, recording, format="zarr", folder=folder, sparse=False, sparsity=None - ) - - sorting_analyzer.compute(["random_spikes", "templates", "spike_amplitudes"]) - # load in lazy mode and check that spikevector is ZarrSpikeVector andextension data are zarr arrays - sorting_analyzer_lazy = load_sorting_analyzer(folder, format="auto", lazy=True) - - assert isinstance(sorting_analyzer_lazy.sorting.to_spike_vector(), ZarrSpikeVector) + 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, zarr.Array) + 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, zarr.Array) + 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): From 09ae85722bd6f0eeb36eb3854e1ec5d31604b210 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 21 Jul 2026 13:56:14 +0200 Subject: [PATCH 09/13] feat: add gather to zarr in node pipeline --- src/spikeinterface/core/node_pipeline.py | 109 ++++++++++++++++-- src/spikeinterface/core/sortinganalyzer.py | 12 +- .../core/tests/test_node_pipeline.py | 29 +++++ 3 files changed, 140 insertions(+), 10 deletions(-) diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index 556370f1ed..79b60af923 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -1,6 +1,6 @@ -from typing import Type -import struct +from typing import Type, Literal import copy +import struct import warnings from pathlib import Path @@ -534,7 +534,7 @@ def run_node_pipeline( nodes: list[PipelineNode], job_kwargs: dict, job_name: str = "pipeline", - gather_mode: str = "memory", + gather_mode: Literal["memory", "npy", "zarr"] = "memory", gather_kwargs: dict = {}, squeeze_output: bool = True, folder: str | None = None, @@ -580,14 +580,14 @@ def run_node_pipeline( The classical job_kwargs job_name : str The name of the pipeline used for the progress_bar - gather_mode : "memory" | "npy" + gather_mode : "memory" | "npy" | "zarr" How to gather the output of the nodes. gather_kwargs : dict - Options to control the "gather engine". See GatherToMemory or GatherToNpy. + Options to control the "gather engine". See GatherToMemory, GatherToNpy or GatherToZarr. squeeze_output : bool, default True If only one output node then squeeze the tuple folder : str | Path | None - Used for gather_mode="npy" + Used for gather_mode="npy" or gather_mode="zarr" names : list of str Names of outputs. verbose : bool, default False @@ -622,6 +622,8 @@ def run_node_pipeline( gather_func = GatherToMemory() elif gather_mode == "npy": gather_func = GatherToNpy(folder, names, **gather_kwargs) + elif gather_mode == "zarr": + gather_func = GatherToZarr(folder, names, **gather_kwargs) else: raise ValueError(f"wrong gather_mode : {gather_mode}") @@ -908,5 +910,96 @@ def finalize_buffers(self, squeeze_output=False): class GatherToZarr: - pass - # Fot me (sam) this is not necessary unless someone realy really want to use + """ + Gather output of nodes into a zarr folder and then open them back as zarr arrays. + + Each returned output ("name") is stored as a zarr array in the root group. Buffers + are appended chunk by chunk along the first axis as the pipeline runs, so the full + result never has to be held in memory. This is the on-disk equivalent of + GatherToNpy but with chunking and compression. + + Parameters + ---------- + folder : str | Path + The folder where the zarr store is created. + names : list of str + Names of the outputs. One zarr array is created per name. + compressor : numcodecs codec | "default" | None, default: "default" + The compressor used for every array. If "default", the SpikeInterface default + zarr compressor is used (Blosc-zstd, level 5, bitshuffle). If None, no compression. + zarr_chunk_size : int | None, default: None + Number of rows (first axis) per zarr chunk. If None, the size of the first + gathered buffer is used. + """ + + def __init__(self, folder, names, compressor="default", zarr_chunk_size=None): + import zarr + + from spikeinterface.core.zarrextractors import get_default_zarr_compressor + + self.folder = Path(folder) + assert names is not None + self.names = names + self.zarr_chunk_size = zarr_chunk_size + + if compressor == "default": + compressor = get_default_zarr_compressor() + self.compressor = compressor + + self.zarr_root = zarr.open(str(self.folder), mode="w") + # arrays are created lazily on the first buffer so we know dtype and trailing shape + self.arrays = [None] * len(names) + + self.tuple_mode = None + + def __call__(self, res): + if res is None: + return + + if self.tuple_mode is None: + # first loop only + self.tuple_mode = isinstance(res, tuple) + if self.tuple_mode: + assert len(self.names) == len(res) + else: + assert len(self.names) == 1 + + if not self.tuple_mode: + res = (res,) + + # distribute buffers to zarr arrays + for i, name in enumerate(self.names): + buf = np.require(res[i], requirements="C") + if self.arrays[i] is None: + # first loop only : create the array with the right dtype and trailing shape + trailing_shape = buf.shape[1:] + chunk0 = self.zarr_chunk_size if self.zarr_chunk_size is not None else max(1, buf.shape[0]) + self.arrays[i] = self.zarr_root.create_dataset( + name=name, + shape=(0,) + trailing_shape, + chunks=(chunk0,) + trailing_shape, + dtype=buf.dtype, + compressor=self.compressor, + ) + self.arrays[i].append(buf, axis=0) + + def finalize_buffers(self, squeeze_output=False): + import zarr + + # consolidate metadata for faster/cleaner re-opening + zarr.consolidate_metadata(self.zarr_root.store) + + if self.tuple_mode: + outs = () + for i, name in enumerate(self.names): + outs += (self.zarr_root[name],) + + if len(outs) == 1 and squeeze_output: + # when tuple size == 1 then remove the tuple + return outs[0] + else: + # always a tuple even of size 1 + return outs + else: + # only one array + return self.zarr_root[self.names[0]] diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index a15128df40..420408e39c 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -2296,7 +2296,9 @@ def compute_one_extension(self, extension_name, save=True, verbose=False, **kwar self.extensions[extension_name] = extension_instance return extension_instance - def compute_several_extensions(self, extensions, save=True, verbose=False, **job_kwargs): + def compute_several_extensions( + self, extensions, save=True, verbose=False, gather_mode="memory", gather_kwargs=None, **job_kwargs + ): """ Compute several extensions @@ -2312,6 +2314,11 @@ def compute_several_extensions(self, extensions, save=True, verbose=False, **job It the extension can be saved then it is saved. If not then the extension will only live in memory as long as the object is deleted. save=False is convenient to try some parameters without changing an already saved extension. + gather_mode : "memory" | "numpy", default: "memory" + Gather mode for node_pipeline extensions. If "memory", the results are gathered in memory. + If "numpy", the results are gathered in numpy arrays. + gather_kwargs : dict | None, default: None + Additional keyword arguments for the gather function. If None, default gather kwargs are used. Returns ------- @@ -2401,7 +2408,8 @@ def compute_several_extensions(self, extensions, save=True, verbose=False, **job all_nodes, job_kwargs=job_kwargs, job_name=job_name, - gather_mode="memory", + gather_mode=gather_mode, + gather_kwargs=gather_kwargs, squeeze_output=False, verbose=verbose, ) diff --git a/src/spikeinterface/core/tests/test_node_pipeline.py b/src/spikeinterface/core/tests/test_node_pipeline.py index ab809b5d4d..d0f62e9a6c 100644 --- a/src/spikeinterface/core/tests/test_node_pipeline.py +++ b/src/spikeinterface/core/tests/test_node_pipeline.py @@ -180,6 +180,35 @@ def test_run_node_pipeline(cache_folder_creation): assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms2) assert np.array_equal(denoised_waveforms_rms2, denoised_waveforms_rms3) + # gather zarr mode + import zarr + + zarr_folder = cache_folder / f"pipeline_folder_{loop}.zarr" + if zarr_folder.is_dir(): + shutil.rmtree(zarr_folder) + output = run_node_pipeline( + recording, + nodes, + job_kwargs, + gather_mode="zarr", + folder=zarr_folder, + names=["amplitudes", "waveforms_rms", "denoised_waveforms_rms"], + ) + amplitudes_z, waveforms_rms_z, denoised_waveforms_rms_z = output + + # values must match the memory gather + assert np.array_equal(amplitudes, amplitudes_z[:]) + assert np.array_equal(waveforms_rms, waveforms_rms_z[:]) + assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_z[:]) + + # arrays must be persisted on disk and re-openable + zarr_root = zarr.open(str(zarr_folder), mode="r") + for name in ("amplitudes", "waveforms_rms", "denoised_waveforms_rms"): + assert name in zarr_root + assert np.array_equal(amplitudes, zarr_root["amplitudes"][:]) + assert np.array_equal(waveforms_rms, zarr_root["waveforms_rms"][:]) + assert np.array_equal(denoised_waveforms_rms, zarr_root["denoised_waveforms_rms"][:]) + # Test pickle mechanism for node in nodes: import pickle From 080b17eed260287f3bd4130e3979733431a9480c Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 21 Jul 2026 13:57:57 +0200 Subject: [PATCH 10/13] fix: only delete extension folders if not lazy --- src/spikeinterface/core/sortinganalyzer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index a15128df40..5801ba2834 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -2516,7 +2516,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() From 1b6a82c36a95b1fe8caebe3695768dec9e292684 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 21 Jul 2026 15:20:05 +0200 Subject: [PATCH 11/13] feat: GatherToZarr and save extension data in chunks to disk --- .../core/analyzer_extension_core.py | 47 ++-- src/spikeinterface/core/node_pipeline.py | 208 +++++++++++++++--- src/spikeinterface/core/sortinganalyzer.py | 78 ++++++- .../core/tests/test_node_pipeline.py | 102 +++++++++ .../core/tests/test_sortinganalyzer.py | 142 ++++++++++++ .../postprocessing/amplitude_scalings.py | 1 + 6 files changed, 520 insertions(+), 58 deletions(-) diff --git a/src/spikeinterface/core/analyzer_extension_core.py b/src/spikeinterface/core/analyzer_extension_core.py index dde1ba8620..85dcf4b890 100644 --- a/src/spikeinterface/core/analyzer_extension_core.py +++ b/src/spikeinterface/core/analyzer_extension_core.py @@ -1532,13 +1532,23 @@ def _set_params(self, **kwargs): def _run(self, verbose=False, **job_kwargs): from spikeinterface.core.node_pipeline import run_node_pipeline - # TODO: should we save directly to npy in binary_folder format / or to zarr? - # if self.sorting_analyzer.format == "binary_folder": - # gather_mode = "npy" - # extension_folder = self.sorting_analyzer.folder / "extenstions" / self.extension_name - # gather_kwargs = {"folder": extension_folder} - gather_mode = "memory" - gather_kwargs = {} + # gather results directly to the final on-disk location (one npy file / zarr dataset per + # nodepipeline variable) to avoid an extra in-memory copy. This is only done when we are + # actually saving to a disk format (see AnalyzerExtension.run()); otherwise gather in memory. + gather_to_disk = getattr(self, "_save_to_disk", False) and self.format in ("binary_folder", "zarr") + if gather_to_disk: + extension_folder = self.sorting_analyzer.folder / "extensions" / self.extension_name + names = self.nodepipeline_variables + if self.format == "binary_folder": + gather_mode = "npy" + folder = [extension_folder / f"{name}.npy" for name in names] + else: + gather_mode = "zarr" + folder = [extension_folder / name for name in names] + else: + gather_mode = "memory" + folder = None + names = None job_kwargs = fix_job_kwargs(job_kwargs) nodes = self.get_pipeline_nodes() @@ -1548,7 +1558,8 @@ def _run(self, verbose=False, **job_kwargs): job_kwargs=job_kwargs, job_name=self.extension_name, gather_mode=gather_mode, - gather_kwargs=gather_kwargs, + folder=folder, + names=names, verbose=False, ) if isinstance(data, tuple): @@ -1599,6 +1610,11 @@ def _get_data(self, outputs="numpy", concatenated=False, return_data_name=None, ), f"return_data_name {return_data_name} not in nodepipeline_variables {self.nodepipeline_variables}" all_data = self.data[return_data_name] + # data gathered directly into a zarr store (e.g. by the node pipeline) is kept as a + # zarr.Array handle. On a non-lazy analyzer we materialize it to a numpy array (this mirrors + # the non-lazy load convention). A memmap is an np.ndarray subclass so it is left untouched. + if not self.sorting_analyzer._lazy and not isinstance(all_data, np.ndarray): + all_data = np.asarray(all_data) keep_mask = None if periods is not None: keep_mask = select_sorting_periods_mask( @@ -1659,7 +1675,9 @@ def _select_units_extension_data(self, unit_ids): new_data = dict() for data_name in self.nodepipeline_variables: if self.data.get(data_name) is not None: - new_data[data_name] = self.data[data_name][keep_spike_mask] + # np.asarray materializes a zarr.Array to numpy (needed for boolean-mask indexing) + # while leaving a numpy array / memmap untouched (the mask then reads only the subset) + new_data[data_name] = np.asarray(self.data[data_name])[keep_spike_mask] return new_data @@ -1670,15 +1688,18 @@ def _merge_extension_data( for data_name in self.nodepipeline_variables: if self.data.get(data_name) is not None: if keep_mask is None: - new_data[data_name] = self.data[data_name].copy() + # return an independent copy (also materializes a zarr.Array to numpy) + new_data[data_name] = np.array(self.data[data_name]) else: - new_data[data_name] = self.data[data_name][keep_mask] + # np.asarray leaves a numpy array / memmap untouched; the mask reads only the subset + new_data[data_name] = np.asarray(self.data[data_name])[keep_mask] return new_data def _split_extension_data(self, split_units, new_unit_ids, new_sorting_analyzer, verbose=False, **job_kwargs): - # splitting only changes random spikes assignments - return self.data.copy() + # splitting only changes random spikes assignments, the per-spike data is unchanged. + # materialize to numpy in case data is a zarr.Array so it can be saved to the new store. + return {data_name: np.array(data) for data_name, data in self.data.items()} def _update_data_after_merge_or_split(old_analyzer, new_analyzer, old_arr, new_sub_arr, new_unit_ids): diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index 79b60af923..4fea3bf0c8 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -537,7 +537,7 @@ def run_node_pipeline( gather_mode: Literal["memory", "npy", "zarr"] = "memory", gather_kwargs: dict = {}, squeeze_output: bool = True, - folder: str | None = None, + folder: str | Path | list | None = None, names: list[str] | None = None, verbose: bool = False, skip_after_n_peaks: int | None = None, @@ -586,8 +586,10 @@ def run_node_pipeline( Options to control the "gather engine". See GatherToMemory, GatherToNpy or GatherToZarr. squeeze_output : bool, default True If only one output node then squeeze the tuple - folder : str | Path | None - Used for gather_mode="npy" or gather_mode="zarr" + folder : str | Path | list | None + Used for gather_mode="npy" or gather_mode="zarr". Either a single folder (one file/array + per name is created inside it) or a list of explicit per-output destinations + (see GatherToNpy and GatherToZarr). names : list of str Names of outputs. verbose : bool, default False @@ -809,24 +811,54 @@ class GatherToNpy: * speculate on a header length (1024) * accumulate in C order the buffer * create the npy v1.0 header at the end with the correct shape and dtype + + Parameters + ---------- + folder : str | Path | list of (str | Path) | None + Where to write the npy files. Two modes: + + * a single folder (str | Path) : one ``.npy`` file is created per name inside it. + * a list of file paths : explicit destination ``.npy`` file per output. This is useful + to gather directly to a final location (e.g. an extension folder). When `names` is + None, they are derived from the file stems. + names : list of str | None + Names of the outputs. Can be None when `folder` is a list of file paths. + npy_header_size : int, default: 1024 + The reserved header size for the npy files. + exist_ok : bool, default: False + Whether the `folder` is allowed to already exist. Only used when `folder` is a single folder. """ - def __init__(self, folder, names, npy_header_size=1024, exist_ok=False): - self.folder = Path(folder) - self.folder.mkdir(parents=True, exist_ok=exist_ok) - assert names is not None - self.names = names + def __init__(self, folder=None, names=None, npy_header_size=1024, exist_ok=False): self.npy_header_size = npy_header_size + if isinstance(folder, (list, tuple)): + # explicit destination file per output + self.file_paths = [Path(p) for p in folder] + if names is None: + names = [file_path.stem for file_path in self.file_paths] + assert len(self.file_paths) == len(names), "`folder` (list of files) must have the same length as `names`" + self.names = names + self.folder = None + # make sure parent folders exist + for file_path in self.file_paths: + file_path.parent.mkdir(parents=True, exist_ok=True) + else: + assert folder is not None, "`folder` must be given" + assert names is not None, "`names` must be given when `folder` is a single folder" + self.folder = Path(folder) + self.folder.mkdir(parents=True, exist_ok=exist_ok) + self.names = names + self.file_paths = [self.folder / (name + ".npy") for name in names] + self.tuple_mode = None self.files = [] self.dtypes = [] self.shapes0 = [] self.final_shapes = [] - for name in names: - filename = self.folder / (name + ".npy") - f = open(filename, "wb+") + for file_path in self.file_paths: + f = open(file_path, "wb+") f.seek(npy_header_size) self.files.append(f) self.dtypes.append(None) @@ -864,7 +896,7 @@ def finalize_buffers(self, squeeze_output=False): f.close() for i, name in enumerate(self.names): - filename = self.folder / (name + ".npy") + filename = self.file_paths[i] shape = (self.shapes0[i],) if self.final_shapes[i] is not None: @@ -894,7 +926,7 @@ def finalize_buffers(self, squeeze_output=False): if self.tuple_mode: outs = () for i, name in enumerate(self.names): - filename = self.folder / (name + ".npy") + filename = self.file_paths[i] outs += (np.load(filename, mmap_mode="r"),) if len(outs) == 1 and squeeze_output: @@ -905,7 +937,7 @@ def finalize_buffers(self, squeeze_output=False): return outs else: # only one file - filename = self.folder / (self.names[0] + ".npy") + filename = self.file_paths[0] return np.load(filename, mmap_mode="r") @@ -920,35 +952,99 @@ class GatherToZarr: Parameters ---------- - folder : str | Path - The folder where the zarr store is created. - names : list of str - Names of the outputs. One zarr array is created per name. + folder : str | Path | list of (str | Path | zarr.Array) | None + Where to write the zarr arrays. Two modes: + + * a single folder (str | Path) : a fresh zarr store is created there and one zarr + array is created per name in its root. + * a list of explicit destinations : buffers are appended directly to these datasets + instead of creating a fresh store. This is useful to gather directly to a final + location (e.g. a SortingAnalyzer extension group). Each item can be: + + - a path pointing inside a zarr store, e.g. + ``"my-analyzer.zarr/extensions/spike_amplitudes/amplitudes"``. The store root + (the ``.zarr`` part) is opened in append mode and the dataset is created on the + fly (with the right dtype and trailing shape) once the first buffer arrives, or + - a pre-created (and resizable) ``zarr.Array``, empty on the first axis + (shape ``(0, *trailing)``) with the correct trailing shape and dtype. + + When `names` is None, they are derived from the datasets' basenames. + names : list of str | None + Names of the outputs. Can be None when `folder` is a list of explicit destinations. compressor : numcodecs codec | "default" | None, default: "default" The compressor used for every array. If "default", the SpikeInterface default zarr compressor is used (Blosc-zstd, level 5, bitshuffle). If None, no compression. + Ignored for destinations that are already a ``zarr.Array`` (they keep their own compressor). zarr_chunk_size : int | None, default: None - Number of rows (first axis) per zarr chunk. If None, the size of the first - gathered buffer is used. + Number of rows (first axis) per zarr chunk. If None, it is computed automatically + so that each chunk is about `zarr_target_chunk_bytes` (see below), which gives a + sensible chunk size regardless of the dtype and trailing shape. Ignored for + destinations that are already a ``zarr.Array``. + zarr_target_chunk_bytes : int, default: 10485760 (10 MiB) + Target (uncompressed) size in bytes of one zarr chunk, used to compute the number of + rows per chunk when `zarr_chunk_size` is None. Ignored for destinations that are + already a ``zarr.Array``. """ - def __init__(self, folder, names, compressor="default", zarr_chunk_size=None): + def __init__( + self, + folder=None, + names=None, + compressor="default", + zarr_chunk_size=None, + zarr_target_chunk_bytes=10 * 1024 * 1024, + ): import zarr from spikeinterface.core.zarrextractors import get_default_zarr_compressor - self.folder = Path(folder) - assert names is not None - self.names = names self.zarr_chunk_size = zarr_chunk_size + self.zarr_target_chunk_bytes = zarr_target_chunk_bytes if compressor == "default": compressor = get_default_zarr_compressor() self.compressor = compressor - self.zarr_root = zarr.open(str(self.folder), mode="w") - # arrays are created lazily on the first buffer so we know dtype and trailing shape - self.arrays = [None] * len(names) + if isinstance(folder, (list, tuple)): + # append to explicit destinations (paths inside a store or pre-created zarr.Arrays) + num_datasets = len(folder) + self.arrays = [None] * num_datasets + # per output : (root_group, internal_path) used to lazily create the dataset, + # or None when the dataset is already a zarr.Array + self._create_specs = [None] * num_datasets + derived_names = [] + root_cache = {} + for i, dataset in enumerate(folder): + if isinstance(dataset, (str, Path)): + store_path, internal_path = _split_zarr_store_path(dataset) + store_path = str(store_path) + if store_path not in root_cache: + # append mode : do not wipe an existing store (e.g. an analyzer) + root_cache[store_path] = zarr.open(store_path, mode="a") + self._create_specs[i] = (root_cache[store_path], internal_path) + derived_names.append(internal_path.split("/")[-1]) + else: + # already a zarr.Array + self.arrays[i] = dataset + derived_names.append(dataset.basename) + if names is None: + names = derived_names + assert num_datasets == len(names), "`folder` (list of datasets) must have the same length as `names`" + self.names = names + self.folder = None + self.zarr_root = None + # we do not own the store so we must not consolidate/close it + self._owns_store = False + else: + assert folder is not None, "`folder` must be given" + assert names is not None, "`names` must be given when `folder` is a single folder" + self.names = names + self.folder = Path(folder) + self.zarr_root = zarr.open(str(self.folder), mode="w") + # arrays are created lazily on the first buffer so we know dtype and trailing shape + self.arrays = [None] * len(names) + self._create_specs = [(self.zarr_root, name) for name in names] + self._owns_store = True self.tuple_mode = None @@ -972,27 +1068,33 @@ def __call__(self, res): buf = np.require(res[i], requirements="C") if self.arrays[i] is None: # first loop only : create the array with the right dtype and trailing shape + root, internal_path = self._create_specs[i] trailing_shape = buf.shape[1:] - chunk0 = self.zarr_chunk_size if self.zarr_chunk_size is not None else max(1, buf.shape[0]) - self.arrays[i] = self.zarr_root.create_dataset( - name=name, + if self.zarr_chunk_size is not None: + chunk0 = self.zarr_chunk_size + else: + # pick the number of rows per chunk to target ~zarr_target_chunk_bytes per chunk + row_nbytes = int(np.prod(trailing_shape, dtype="int64")) * buf.dtype.itemsize + chunk0 = max(1, self.zarr_target_chunk_bytes // max(1, row_nbytes)) + self.arrays[i] = root.create_dataset( + name=internal_path, shape=(0,) + trailing_shape, chunks=(chunk0,) + trailing_shape, dtype=buf.dtype, compressor=self.compressor, + overwrite=True, ) self.arrays[i].append(buf, axis=0) def finalize_buffers(self, squeeze_output=False): import zarr - # consolidate metadata for faster/cleaner re-opening - zarr.consolidate_metadata(self.zarr_root.store) + if self._owns_store: + # consolidate metadata for faster/cleaner re-opening + zarr.consolidate_metadata(self.zarr_root.store) if self.tuple_mode: - outs = () - for i, name in enumerate(self.names): - outs += (self.zarr_root[name],) + outs = tuple(self.arrays) if len(outs) == 1 and squeeze_output: # when tuple size == 1 then remove the tuple @@ -1002,4 +1104,38 @@ def finalize_buffers(self, squeeze_output=False): return outs else: # only one array - return self.zarr_root[self.names[0]] + return self.arrays[0] + + +def _split_zarr_store_path(dataset_path): + """ + Split a path pointing inside a zarr store into (store_path, internal_path). + + The store is the portion of the path up to and including the component ending with + ``.zarr``. The remaining components form the internal (group/dataset) path. + + Example + ------- + >>> _split_zarr_store_path("a/b/my-analyzer.zarr/extensions/spike_amplitudes/amplitudes") + (PosixPath('a/b/my-analyzer.zarr'), 'extensions/spike_amplitudes/amplitudes') + """ + parts = Path(dataset_path).parts + zarr_index = None + for index, part in enumerate(parts): + if part.endswith(".zarr"): + zarr_index = index + break + if zarr_index is None: + raise ValueError( + f"Could not find a '.zarr' store in path '{dataset_path}'. " + "Provide a path like '.zarr/'." + ) + internal_parts = parts[zarr_index + 1 :] + if len(internal_parts) == 0: + raise ValueError( + f"No dataset path inside the zarr store found in '{dataset_path}'. " + "Provide a path like '.zarr/'." + ) + store_path = Path(*parts[: zarr_index + 1]) + internal_path = "/".join(internal_parts) + return store_path, internal_path diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 420408e39c..fc32eced13 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -40,7 +40,6 @@ from .sortingfolder import NumpyFolderSorting from .zarrextractors import get_default_zarr_compressor, ZarrSortingExtractor, super_zarr_open, _write_object_array from .node_pipeline import run_node_pipeline -from .globals import get_global_job_kwargs # high level function @@ -2297,7 +2296,7 @@ def compute_one_extension(self, extension_name, save=True, verbose=False, **kwar return extension_instance def compute_several_extensions( - self, extensions, save=True, verbose=False, gather_mode="memory", gather_kwargs=None, **job_kwargs + self, extensions, save=True, verbose=False, gather_mode=None, gather_kwargs=None, **job_kwargs ): """ Compute several extensions @@ -2314,9 +2313,9 @@ def compute_several_extensions( It the extension can be saved then it is saved. If not then the extension will only live in memory as long as the object is deleted. save=False is convenient to try some parameters without changing an already saved extension. - gather_mode : "memory" | "numpy", default: "memory" - Gather mode for node_pipeline extensions. If "memory", the results are gathered in memory. - If "numpy", the results are gathered in numpy arrays. + gather_mode : "memory" | "numpy" | "zarr" | None, default: None + Gather mode for node_pipeline extensions. If None, the results are gathered using the same format + as the SortingAnalyzer ("memory" for "memory", "npy" for "binary_folder", "zarr" for "zarr"). gather_kwargs : dict | None, default: None Additional keyword arguments for the gather function. If None, default gather kwargs are used. @@ -2384,14 +2383,40 @@ def compute_several_extensions( result_routage = [] extension_instances = {} + # decide how the pipeline results are gathered. + # when we save to a disk format we gather directly to the extension final location + # (npy files or zarr datasets) to avoid an extra in-memory copy. Otherwise (memory + # format, save=False or read-only analyzer) we gather in memory. + save_to_disk = save and not self.is_read_only() + if gather_mode is None: + if save_to_disk and self.format == "binary_folder": + gather_mode = "npy" + elif save_to_disk and self.format == "zarr": + gather_mode = "zarr" + else: + gather_mode = "memory" + if gather_kwargs is None: + gather_kwargs = {} + + # for disk gather modes we build one destination per output variable + gather_folder = [] if gather_mode in ("npy", "zarr") else None + gather_names = [] if gather_mode in ("npy", "zarr") else None + for extension_name, extension_params in extensions_with_pipeline.items(): extension_class = get_extension_class(extension_name) assert ( self.has_recording() or self.has_temporary_recording() ), f"Extension {extension_name} requires the recording" + extension_folder = self.folder / "extensions" / extension_name if gather_folder is not None else None for variable_name in extension_class.nodepipeline_variables: result_routage.append((extension_name, variable_name)) + if gather_mode == "npy": + gather_folder.append(extension_folder / f"{variable_name}.npy") + gather_names.append(variable_name) + elif gather_mode == "zarr": + gather_folder.append(extension_folder / variable_name) + gather_names.append(variable_name) extension_instance = extension_class(self) extension_instance.set_params(save=save, **extension_params) @@ -2400,6 +2425,13 @@ def compute_several_extensions( nodes = extension_instance.get_pipeline_nodes() all_nodes.extend(nodes) + # reset and save params before running so the pipeline can gather directly into the + # (freshly created) extension folders/groups (mirrors AnalyzerExtension.run()) + if save_to_disk: + for extension_instance in extension_instances.values(): + extension_instance._save_params() + extension_instance._save_importing_provenance() + job_name = "Compute : " + " + ".join(extensions_with_pipeline.keys()) t_start = perf_counter() @@ -2409,6 +2441,8 @@ def compute_several_extensions( job_kwargs=job_kwargs, job_name=job_name, gather_mode=gather_mode, + folder=gather_folder, + names=gather_names, gather_kwargs=gather_kwargs, squeeze_output=False, verbose=verbose, @@ -2425,7 +2459,18 @@ def compute_several_extensions( for extension_name, extension_instance in extension_instances.items(): self.extensions[extension_name] = extension_instance - if save: + if save_to_disk: + # params/provenance already saved above (before the run). Here we only persist + # the run info and the data. For disk gather modes the data is already written + # to its final location so _save_data() is a no-op for those variables. + extension_instance._save_run_info() + extension_instance._save_data() + if self.format == "zarr": + import zarr + + zarr.consolidate_metadata(self._get_zarr_root().store) + elif save: + # memory format or read-only : keep the previous behavior extension_instance.save() for extension_name, extension_params in extensions_post_pipeline.items(): @@ -3234,18 +3279,24 @@ def split( return new_extension def run(self, save=True, **kwargs): - if save and not self.sorting_analyzer.is_read_only(): + save_to_disk = save and not self.sorting_analyzer.is_read_only() + if save_to_disk: # NB: this call to _save_params() also resets the folder or zarr group self._save_params() self._save_importing_provenance() + # let _run() know whether it may gather results directly to the final on-disk location. + # This is only valid when we actually save to a disk format (the folder/group has just + # been reset above). Otherwise _run() must gather in memory. + self._save_to_disk = save_to_disk + t_start = perf_counter() self._run(**kwargs) t_end = perf_counter() self.run_info["runtime_s"] = t_end - t_start self.run_info["run_completed"] = True - if save and not self.sorting_analyzer.is_read_only(): + if save_to_disk: self._save_run_info() self._save_data() if self.format == "zarr": @@ -3304,6 +3355,8 @@ def _save_data(self): except: raise Exception(f"Could not save {ext_data_name} as extension data") elif self.format == "zarr": + import zarr + saving_options = self.sorting_analyzer._backend_options.get("saving_options", {}) extension_group = self._get_zarr_extension_group(mode="r+") @@ -3312,6 +3365,10 @@ def _save_data(self): saving_options["compressor"] = get_default_zarr_compressor() for ext_data_name, ext_data in self.data.items(): + if isinstance(ext_data, zarr.Array): + # the data was gathered directly into the extension group (e.g. by the node + # pipeline), so it is already in its final location : nothing to copy + continue if ext_data_name in extension_group: del extension_group[ext_data_name] if isinstance(ext_data, (dict, list)): @@ -3319,7 +3376,10 @@ def _save_data(self): _write_object_array(extension_group, ext_data_name, ext_data_, codec="json") extension_group[ext_data_name].attrs["dict"] = True elif isinstance(ext_data, np.ndarray): - extension_group.create_dataset(name=ext_data_name, data=ext_data, **saving_options) + # only save the array if the dataset does not already exist, since it is created directly + # by the run_node_pipeline() function in the case of nodepipeline extensions + if ext_data_name not in extension_group: + extension_group.create_dataset(name=ext_data_name, data=ext_data, **saving_options) elif HAS_PANDAS and isinstance(ext_data, pd.DataFrame): df_group = extension_group.create_group(ext_data_name) # first we save the index diff --git a/src/spikeinterface/core/tests/test_node_pipeline.py b/src/spikeinterface/core/tests/test_node_pipeline.py index d0f62e9a6c..831d481005 100644 --- a/src/spikeinterface/core/tests/test_node_pipeline.py +++ b/src/spikeinterface/core/tests/test_node_pipeline.py @@ -209,6 +209,60 @@ def test_run_node_pipeline(cache_folder_creation): assert np.array_equal(waveforms_rms, zarr_root["waveforms_rms"][:]) assert np.array_equal(denoised_waveforms_rms, zarr_root["denoised_waveforms_rms"][:]) + # gather npy mode with an explicit list of file paths (final location) + npy_files_folder = cache_folder / f"pipeline_npy_files_{loop}" + if npy_files_folder.is_dir(): + shutil.rmtree(npy_files_folder) + npy_files = [ + npy_files_folder / "amp.npy", + npy_files_folder / "sub" / "rms.npy", + npy_files_folder / "denoised_rms.npy", + ] + output = run_node_pipeline( + recording, + nodes, + job_kwargs, + gather_mode="npy", + folder=npy_files, + ) + amplitudes_f, waveforms_rms_f, denoised_waveforms_rms_f = output + for npy_file in npy_files: + assert npy_file.is_file() + assert np.array_equal(amplitudes, amplitudes_f) + assert np.array_equal(waveforms_rms, waveforms_rms_f) + assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_f) + + # gather zarr mode with an explicit list of dataset paths, created on the fly + # inside an existing store (final location, e.g. an analyzer extension group) + datasets_store = cache_folder / f"pipeline_zarr_datasets_{loop}.zarr" + if datasets_store.is_dir(): + shutil.rmtree(datasets_store) + # pre-existing store that must not be wiped + root = zarr.open(str(datasets_store), mode="w") + root.attrs["preexisting"] = True + dataset_paths = [ + datasets_store / "extensions" / "amplitudes", + datasets_store / "extensions" / "waveforms_rms", + datasets_store / "extensions" / "denoised_waveforms_rms", + ] + output = run_node_pipeline( + recording, + nodes, + job_kwargs, + gather_mode="zarr", + folder=dataset_paths, + ) + amplitudes_d, waveforms_rms_d, denoised_waveforms_rms_d = output + assert np.array_equal(amplitudes, amplitudes_d[:]) + assert np.array_equal(waveforms_rms, waveforms_rms_d[:]) + assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_d[:]) + # data must be persisted at the passed final location and the store not wiped + root_reopen = zarr.open(str(datasets_store), mode="r") + assert root_reopen.attrs.get("preexisting", False) + assert np.array_equal(amplitudes, root_reopen["extensions"]["amplitudes"][:]) + assert np.array_equal(waveforms_rms, root_reopen["extensions"]["waveforms_rms"][:]) + assert np.array_equal(denoised_waveforms_rms, root_reopen["extensions"]["denoised_waveforms_rms"][:]) + # Test pickle mechanism for node in nodes: import pickle @@ -217,6 +271,54 @@ def test_run_node_pipeline(cache_folder_creation): unpickled_node = pickle.loads(pickled_node) +def test_gather_to_zarr_chunking(tmp_path): + # the zarr chunk size along the first axis must be picked from a byte target (not from the + # size of the first gathered buffer), so it stays sensible for billions of spikes and never + # collapses to 1 row per chunk on a quiet first chunk. + recording, sorting = generate_ground_truth_recording(num_channels=8, num_units=5, durations=[20.0], seed=7) + + # small chunks + n_jobs>1 so the first non-empty buffer is small (would give chunk0==1 with the + # old first-buffer heuristic) + job_kwargs = dict(chunk_duration="0.1s", n_jobs=2, progress_bar=False) + + spikes = sorting.to_spike_vector() + peaks = np.zeros(spikes.size, dtype=spike_peak_dtype) + peaks["sample_index"] = spikes["sample_index"] + peaks["segment_index"] = spikes["segment_index"] + + peak_retriever = PeakRetriever(recording, peaks) + ms_before, ms_after = 0.5, 1.0 + dense_waveforms = ExtractDenseWaveforms( + recording, parents=[peak_retriever], ms_before=ms_before, ms_after=ms_after, return_output=True + ) + nodes = [peak_retriever, dense_waveforms] + + # default byte target (10 MiB) + target_bytes = 10 * 1024 * 1024 + waveforms = run_node_pipeline( + recording, nodes, job_kwargs, gather_mode="zarr", folder=tmp_path / "wfs.zarr", names=["waveforms"] + ) + nbefore_after = dense_waveforms.nbefore + dense_waveforms.nafter + row_nbytes = nbefore_after * recording.get_num_channels() * np.dtype("float32").itemsize + expected_chunk0 = max(1, target_bytes // row_nbytes) + assert waveforms.chunks[0] == expected_chunk0 + assert waveforms.chunks[0] > 1 + assert waveforms.chunks[1:] == waveforms.shape[1:] + + # explicit override is respected + waveforms2 = run_node_pipeline( + recording, + nodes, + job_kwargs, + gather_mode="zarr", + folder=tmp_path / "wfs2.zarr", + names=["waveforms"], + gather_kwargs={"zarr_chunk_size": 1234}, + ) + assert waveforms2.chunks[0] == 1234 + assert np.array_equal(waveforms[:], waveforms2[:]) + + def test_skip_after_n_peaks_and_recording_slices(): recording, sorting = generate_ground_truth_recording(num_channels=10, num_units=10, durations=[10.0], seed=2205) diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index 25aeb78c1a..9c1cf01210 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -792,6 +792,148 @@ def test_runtime_dependencies(dataset): assert not sorting_analyzer.has_extension("dummy_pipeline") +def _compute_reference_pipeline_data(dataset): + """Compute the dummy_pipeline extension in memory to use as a reference.""" + recording, sorting = dataset + analyzer = create_sorting_analyzer(sorting, recording, format="memory", sparse=False, sparsity=None) + analyzer.compute(["random_spikes", "templates"]) + analyzer.compute({"dummy_pipeline": {"param0": 5.5}}) + return analyzer.get_extension("dummy_pipeline").get_data() + + +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) +def test_compute_pipeline_extension_gather_to_disk(tmp_path, dataset, format): + """ + When computing node-pipeline extensions on a disk-backed analyzer, the results are gathered + directly to their final location (npy files for binary_folder, zarr datasets for zarr) instead + of being accumulated in memory and copied afterwards. This test checks that: + * the auto gather_mode selection matches the analyzer format + * the data is written in place and kept as a memmap / zarr.Array (no extra copy) + * the values match a plain in-memory computation and survive a reload + * recomputing (overwriting) works + """ + import zarr + + register_result_extension(DummyPipelineAnalyzerExtension) + recording, sorting = dataset + + amp_ref = _compute_reference_pipeline_data(dataset) + + if format == "memory": + folder = None + elif format == "binary_folder": + folder = tmp_path / "analyzer" + else: + folder = tmp_path / "analyzer.zarr" + + analyzer = create_sorting_analyzer(sorting, recording, format=format, folder=folder, sparse=False, sparsity=None) + analyzer.compute(["random_spikes", "templates"]) + analyzer.compute({"dummy_pipeline": {"param0": 5.5}}) + + ext = analyzer.get_extension("dummy_pipeline") + assert np.array_equal(ext.get_data(), amp_ref) + + if format == "binary_folder": + # written directly to the final npy file and kept as a memmap (not re-copied by _save_data) + amp_file = folder / "extensions" / "dummy_pipeline" / "amp.npy" + assert amp_file.is_file() + assert isinstance(ext.data["amp"], np.memmap) + elif format == "zarr": + # written directly as a zarr dataset in the extension group + root = analyzer._get_zarr_root(mode="r") + assert "amp" in root["extensions"]["dummy_pipeline"] + assert isinstance(ext.data["amp"], zarr.Array) + + if format != "memory": + # data must survive a reload from disk + analyzer_reloaded = load_sorting_analyzer(folder) + assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + + # recompute (overwrite) must not corrupt or leave stale data behind + analyzer.compute({"dummy_pipeline": {"param0": 5.5}}) + analyzer_reloaded = load_sorting_analyzer(folder) + assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + + +@pytest.mark.parametrize("format", ["binary_folder", "zarr"]) +def test_compute_pipeline_extension_save_false(tmp_path, dataset, format): + """ + With save=False on a disk-backed analyzer, node-pipeline extensions are computed in memory + and nothing is written to disk. + """ + register_result_extension(DummyPipelineAnalyzerExtension) + recording, sorting = dataset + + folder = tmp_path / ("analyzer" if format == "binary_folder" else "analyzer.zarr") + analyzer = create_sorting_analyzer(sorting, recording, format=format, folder=folder, sparse=False, sparsity=None) + analyzer.compute(["random_spikes", "templates"]) + analyzer.compute({"dummy_pipeline": {"param0": 5.5}}, save=False) + + # in memory the extension is available + assert analyzer.has_extension("dummy_pipeline") + + # but nothing was written to disk + analyzer_reloaded = load_sorting_analyzer(folder) + assert not analyzer_reloaded.has_extension("dummy_pipeline") + + +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) +def test_compute_one_pipeline_extension_gather_to_disk(tmp_path, dataset, format): + """ + Same as test_compute_pipeline_extension_gather_to_disk but through compute_one_extension + (i.e. computing a single node-pipeline extension via a string input), which uses + BaseSpikeVectorExtension._run() to gather directly to disk. + """ + import zarr + + register_result_extension(DummyPipelineAnalyzerExtension) + recording, sorting = dataset + + amp_ref = _compute_reference_pipeline_data(dataset) + + if format == "memory": + folder = None + elif format == "binary_folder": + folder = tmp_path / "analyzer" + else: + folder = tmp_path / "analyzer.zarr" + + analyzer = create_sorting_analyzer(sorting, recording, format=format, folder=folder, sparse=False, sparsity=None) + analyzer.compute(["random_spikes", "templates"]) + # single string -> compute_one_extension -> BaseSpikeVectorExtension._run + analyzer.compute("dummy_pipeline", param0=5.5) + + ext = analyzer.get_extension("dummy_pipeline") + assert np.array_equal(ext.get_data(), amp_ref) + + if format == "binary_folder": + assert (folder / "extensions" / "dummy_pipeline" / "amp.npy").is_file() + assert isinstance(ext.data["amp"], np.memmap) + elif format == "zarr": + root = analyzer._get_zarr_root(mode="r") + assert "amp" in root["extensions"]["dummy_pipeline"] + assert isinstance(ext.data["amp"], zarr.Array) + + if format != "memory": + # data must survive a reload and recompute (overwrite) must work + analyzer_reloaded = load_sorting_analyzer(folder) + assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + + analyzer.compute("dummy_pipeline", param0=5.5) + analyzer_reloaded = load_sorting_analyzer(folder) + assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + + # save=False on a disk analyzer: computed in memory, nothing written to disk + folder2 = tmp_path / ("analyzer_nosave" + (".zarr" if format == "zarr" else "")) + analyzer2 = create_sorting_analyzer( + sorting, recording, format=format, folder=folder2, sparse=False, sparsity=None + ) + analyzer2.compute(["random_spikes", "templates"]) + analyzer2.compute("dummy_pipeline", param0=5.5, save=False) + assert analyzer2.has_extension("dummy_pipeline") + assert not load_sorting_analyzer(folder2).has_extension("dummy_pipeline") + + def test_select_channels(dataset): recording, sorting = dataset sorting_analyzer = create_sorting_analyzer(sorting, recording, format="memory", sparse=False, sparsity=None) diff --git a/src/spikeinterface/postprocessing/amplitude_scalings.py b/src/spikeinterface/postprocessing/amplitude_scalings.py index 7870c82be1..6befe8c9a9 100644 --- a/src/spikeinterface/postprocessing/amplitude_scalings.py +++ b/src/spikeinterface/postprocessing/amplitude_scalings.py @@ -1,6 +1,7 @@ import numpy as np from spikeinterface.core import ChannelSparsity +from spikeinterface.core.core_tools import ms_to_samples from spikeinterface.core.template_tools import get_dense_templates_array, _get_nbefore from spikeinterface.core.sortinganalyzer import register_result_extension from spikeinterface.core.analyzer_extension_core import BaseSpikeVectorExtension From 60769a59362103c56c188ea90ef44ce9217cf765 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 23 Jul 2026 11:32:44 +0200 Subject: [PATCH 12/13] fix: AnalyzerExtension __del__ closes all memmaps refs --- src/spikeinterface/core/sortinganalyzer.py | 23 ++++++++++++++ .../core/tests/test_sortinganalyzer.py | 30 +++++++++---------- 2 files changed, 37 insertions(+), 16 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 942ad6575f..9167389254 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -2968,6 +2968,22 @@ def __init__(self, sorting_analyzer): self.run_info = self._default_run_info_dict() self.data = dict() + def __del__(self): + # Close any open memmap file handles held in `data` (e.g. when an extension gathers or + # loads its data as a memmap). On Windows an open memmap prevents deleting the underlying + # file, so releasing the handles here allows the extension folder to be removed when the + # extension is dropped (e.g. on recompute). Best-effort: __del__ must never raise. + data = self.__dict__.get("data", None) + if not data: + return + for value in data.values(): + mmap = getattr(value, "_mmap", None) if isinstance(value, np.memmap) else None + if mmap is not None: + try: + mmap.close() + except Exception: + pass + def _default_run_info_dict(self): return dict(run_completed=False, runtime_s=None) @@ -3460,6 +3476,13 @@ def _reset_extension_folder(self): Delete the extension in a folder (binary or zarr) and create an empty one. """ if self.format == "binary_folder": + # Drop the analyzer's reference to a previously computed extension of the same name so + # it can be garbage collected. Its __del__ releases any open file handles (e.g. a memmap + # kept in `data` when gathering directly to npy) which, on Windows, would otherwise + # prevent deleting the folder below. + old_extension = self.sorting_analyzer.extensions.pop(self.extension_name, None) + del old_extension + extension_folder = self._get_binary_extension_folder() if extension_folder.is_dir(): shutil.rmtree(extension_folder) diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index 9c1cf01210..e8f3ad1213 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -830,29 +830,28 @@ def test_compute_pipeline_extension_gather_to_disk(tmp_path, dataset, format): analyzer.compute(["random_spikes", "templates"]) analyzer.compute({"dummy_pipeline": {"param0": 5.5}}) - ext = analyzer.get_extension("dummy_pipeline") - assert np.array_equal(ext.get_data(), amp_ref) + # NB: do not keep a local reference to the extension (or to `ext.data["amp"]`) across the + # recompute below: on Windows an open memmap on amp.npy would prevent deleting the folder. + assert np.array_equal(analyzer.get_extension("dummy_pipeline").get_data(), amp_ref) if format == "binary_folder": # written directly to the final npy file and kept as a memmap (not re-copied by _save_data) amp_file = folder / "extensions" / "dummy_pipeline" / "amp.npy" assert amp_file.is_file() - assert isinstance(ext.data["amp"], np.memmap) + assert isinstance(analyzer.get_extension("dummy_pipeline").data["amp"], np.memmap) elif format == "zarr": # written directly as a zarr dataset in the extension group root = analyzer._get_zarr_root(mode="r") assert "amp" in root["extensions"]["dummy_pipeline"] - assert isinstance(ext.data["amp"], zarr.Array) + assert isinstance(analyzer.get_extension("dummy_pipeline").data["amp"], zarr.Array) if format != "memory": # data must survive a reload from disk - analyzer_reloaded = load_sorting_analyzer(folder) - assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + assert np.array_equal(load_sorting_analyzer(folder).get_extension("dummy_pipeline").get_data(), amp_ref) # recompute (overwrite) must not corrupt or leave stale data behind analyzer.compute({"dummy_pipeline": {"param0": 5.5}}) - analyzer_reloaded = load_sorting_analyzer(folder) - assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + assert np.array_equal(load_sorting_analyzer(folder).get_extension("dummy_pipeline").get_data(), amp_ref) @pytest.mark.parametrize("format", ["binary_folder", "zarr"]) @@ -903,25 +902,24 @@ def test_compute_one_pipeline_extension_gather_to_disk(tmp_path, dataset, format # single string -> compute_one_extension -> BaseSpikeVectorExtension._run analyzer.compute("dummy_pipeline", param0=5.5) - ext = analyzer.get_extension("dummy_pipeline") - assert np.array_equal(ext.get_data(), amp_ref) + # NB: do not keep a local reference to the extension (or to `ext.data["amp"]`) across the + # recompute below: on Windows an open memmap on amp.npy would prevent deleting the folder. + assert np.array_equal(analyzer.get_extension("dummy_pipeline").get_data(), amp_ref) if format == "binary_folder": assert (folder / "extensions" / "dummy_pipeline" / "amp.npy").is_file() - assert isinstance(ext.data["amp"], np.memmap) + assert isinstance(analyzer.get_extension("dummy_pipeline").data["amp"], np.memmap) elif format == "zarr": root = analyzer._get_zarr_root(mode="r") assert "amp" in root["extensions"]["dummy_pipeline"] - assert isinstance(ext.data["amp"], zarr.Array) + assert isinstance(analyzer.get_extension("dummy_pipeline").data["amp"], zarr.Array) if format != "memory": # data must survive a reload and recompute (overwrite) must work - analyzer_reloaded = load_sorting_analyzer(folder) - assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + assert np.array_equal(load_sorting_analyzer(folder).get_extension("dummy_pipeline").get_data(), amp_ref) analyzer.compute("dummy_pipeline", param0=5.5) - analyzer_reloaded = load_sorting_analyzer(folder) - assert np.array_equal(analyzer_reloaded.get_extension("dummy_pipeline").get_data(), amp_ref) + assert np.array_equal(load_sorting_analyzer(folder).get_extension("dummy_pipeline").get_data(), amp_ref) # save=False on a disk analyzer: computed in memory, nothing written to disk folder2 = tmp_path / ("analyzer_nosave" + (".zarr" if format == "zarr" else "")) From 8ddf40467afc5c512e239469f396b11489568f4d Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 23 Jul 2026 13:33:04 +0200 Subject: [PATCH 13/13] fix: delete extensions --- src/spikeinterface/core/sortinganalyzer.py | 24 ++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 0e4f94ebc3..5fd99623e3 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -2969,10 +2969,16 @@ def __init__(self, sorting_analyzer): self.data = dict() def __del__(self): + # Ensure open memmap file handles are released when the extension is garbage collected. + # Best-effort: __del__ must never raise. + self._release_data_file_handles() + + def _release_data_file_handles(self): # Close any open memmap file handles held in `data` (e.g. when an extension gathers or # loads its data as a memmap). On Windows an open memmap prevents deleting the underlying - # file, so releasing the handles here allows the extension folder to be removed when the - # extension is dropped (e.g. on recompute). Best-effort: __del__ must never raise. + # file, so releasing the handles allows the extension folder to be removed (e.g. on + # recompute). This closes the file handle but leaves the (now unusable) array object and any + # non-memmap data in `data` untouched. data = self.__dict__.get("data", None) if not data: return @@ -3478,12 +3484,14 @@ def _reset_extension_folder(self): Delete the extension in a folder (binary or zarr) and create an empty one. """ if self.format == "binary_folder": - # Drop the analyzer's reference to a previously computed extension of the same name so - # it can be garbage collected. Its __del__ releases any open file handles (e.g. a memmap - # kept in `data` when gathering directly to npy) which, on Windows, would otherwise - # prevent deleting the folder below. - old_extension = self.sorting_analyzer.extensions.pop(self.extension_name, None) - del old_extension + # Release open file handles (e.g. a memmap kept in `data` when gathering/extracting + # directly to npy) of a previously computed extension of the same name. On Windows an + # open memmap would otherwise prevent deleting the folder below. We deliberately do NOT + # remove it from `self.sorting_analyzer.extensions` nor clear its data: some extensions + # (e.g. metrics) read the previous extension back during their computation. + old_extension = self.sorting_analyzer.extensions.get(self.extension_name, None) + if old_extension is not None and old_extension is not self: + old_extension._release_data_file_handles() extension_folder = self._get_binary_extension_folder() if extension_folder.is_dir():