From 6a4c490dc3ae2969a8c09d355dcc8caf8bd90708 Mon Sep 17 00:00:00 2001 From: Samuel Garcia Date: Thu, 9 Jul 2026 15:52:05 +0200 Subject: [PATCH 01/18] mode the Base.save to Sorting.save() and Recording.save() move saving logic to extractor classes --- src/spikeinterface/core/base.py | 135 +++++++++++------- src/spikeinterface/core/baserecording.py | 114 ++++++++------- .../core/baserecordingsnippets.py | 10 +- src/spikeinterface/core/basesorting.py | 2 +- src/spikeinterface/core/binaryfolder.py | 67 ++++++++- src/spikeinterface/core/npyfoldersnippets.py | 2 +- src/spikeinterface/core/numpyextractors.py | 14 +- src/spikeinterface/core/sortingfolder.py | 19 ++- 8 files changed, 245 insertions(+), 118 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index 8be962ec6f..fa6bf46b05 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -480,6 +480,7 @@ def to_dict( self, include_annotations: bool = False, include_properties: bool = False, + include_extra_metadata: bool = True, relative_to: str | Path | None = None, folder_metadata=None, recursive: bool = False, @@ -609,7 +610,8 @@ def to_dict( folder_metadata = Path(folder_metadata).resolve().absolute().relative_to(relative_to) dump_dict["folder_metadata"] = str(folder_metadata) - self._extra_metadata_to_dict(dump_dict) + if include_extra_metadata: + self._extra_metadata_to_dict(dump_dict) return dump_dict @@ -640,15 +642,17 @@ def from_dict(dictionary: dict, base_folder: Path | str | None = None) -> "BaseE folder_metadata = Path(folder_metadata) if dictionary.get("relative_paths", False): folder_metadata = base_folder / folder_metadata - extractor.load_metadata_from_folder(folder_metadata) + extractor._load_metadata_from_folder(folder_metadata) return extractor - def load_metadata_from_folder(self, folder_metadata): - # hack to load probe for recording - folder_metadata = Path(folder_metadata) + def _load_metadata_from_folder(self, folder): + """ + Load properties from a 'properties' subfolder previously saved as npy. + """ + folder = Path(folder) # load properties - prop_folder = folder_metadata / "properties" + prop_folder = folder / "properties" if prop_folder.is_dir(): for prop_file in prop_folder.iterdir(): if prop_file.suffix == ".npy": @@ -656,10 +660,13 @@ def load_metadata_from_folder(self, folder_metadata): key = prop_file.stem self.set_property(key, values) - self._extra_metadata_from_folder(folder_metadata) + self._extra_metadata_from_folder(folder) - def save_metadata_to_folder(self, folder_metadata): - self._extra_metadata_to_folder(folder_metadata) + def _save_metadata_to_folder(self, folder_metadata): + """ + Save properties into a 'properties' subfolder using npy format. + """ + # self._extra_metadata_to_folder(folder_metadata) # save properties prop_folder = folder_metadata / "properties" @@ -778,6 +785,7 @@ def dump_to_json( file_path: str | Path | None = None, relative_to: str | Path | bool | None = None, folder_metadata: str | Path | None = None, + include_extra_metadata: bool = True, ) -> None: """ Dump recording extractor to json file. @@ -803,6 +811,7 @@ def dump_to_json( dump_dict = self.to_dict( include_annotations=True, include_properties=False, + include_extra_metadata=include_extra_metadata, relative_to=relative_to, folder_metadata=folder_metadata, recursive=True, @@ -891,13 +900,13 @@ def _save(self, folder, **save_kwargs): # this is internally call by cache(...) main function raise NotImplementedError - def _extra_metadata_from_folder(self, folder): - # This implemented in BaseRecording for probe - pass + # def _extra_metadata_from_folder(self, folder): + # # This implemented in BaseRecording for probe + # pass - def _extra_metadata_to_folder(self, folder): - # This implemented in BaseRecording for probe - pass + # def _extra_metadata_to_folder(self, folder): + # # This implemented in BaseRecording for probe + # pass def _extra_metadata_from_dict(self, dump_dict): # This implemented in BaseRecording for probe @@ -907,48 +916,62 @@ def _extra_metadata_to_dict(self, dump_dict): # This implemented in BaseRecording for probe pass - def save(self, **kwargs) -> "BaseExtractor": - """ - Save a SpikeInterface object. + # def save(self, **kwargs) -> "BaseExtractor": + # """ + # Save a SpikeInterface object. + + # Parameters + # ---------- + # kwargs: Keyword arguments for saving. + # * format: "memory", "zarr", or "binary" (for recording) / "memory" or "numpy_folder" or "npz_folder" for sorting. + # In case format is not memory, the recording is saved to a folder. See format specific functions for + # more info (`save_to_memory()`, `save_to_folder()`, `save_to_zarr()`) + # * folder: if provided, the folder path where the object is saved + # * name: if provided and folder is not given, the name of the folder in the global temporary + # folder (use set_global_tmp_folder() to change this folder) where the object is saved. + # If folder and name are not given, the object is saved in the global temporary folder with + # a random string + # * dump_ext: "json" or "pkl", default "json" (if format is "folder") + # * verbose: if True output is verbose + # * **save_kwargs: additional kwargs format-dependent and job kwargs for recording + # {} + + # Returns + # ------- + # loaded_extractor: BaseRecording or BaseSorting + # The reference to the saved object after it is loaded back + # """ + # format = kwargs.get("format", None) + # if format == "memory": + # loaded_extractor = self.save_to_memory(**kwargs) + # elif format == "zarr": + # loaded_extractor = self.save_to_zarr(**kwargs) + # else: + # loaded_extractor = self.save_to_folder(**kwargs) + # return loaded_extractor + + # save.__doc__ = save.__doc__.format(_shared_job_kwargs_doc) - Parameters - ---------- - kwargs: Keyword arguments for saving. - * format: "memory", "zarr", or "binary" (for recording) / "memory" or "numpy_folder" or "npz_folder" for sorting. - In case format is not memory, the recording is saved to a folder. See format specific functions for - more info (`save_to_memory()`, `save_to_folder()`, `save_to_zarr()`) - * folder: if provided, the folder path where the object is saved - * name: if provided and folder is not given, the name of the folder in the global temporary - folder (use set_global_tmp_folder() to change this folder) where the object is saved. - If folder and name are not given, the object is saved in the global temporary folder with - a random string - * dump_ext: "json" or "pkl", default "json" (if format is "folder") - * verbose: if True output is verbose - * **save_kwargs: additional kwargs format-dependent and job kwargs for recording - {} + def save_to_memory(self, sharedmem=True, **save_kwargs) -> "BaseExtractor": + warnings.warn("save_to_memory() should be save(format='memory')", FutureWarning) + return self.save(format="memory", sharedmem=sharedmem, **save_kwargs) - Returns - ------- - loaded_extractor: BaseRecording or BaseSorting - The reference to the saved object after it is loaded back - """ - format = kwargs.get("format", None) - if format == "memory": - loaded_extractor = self.save_to_memory(**kwargs) - elif format == "zarr": - loaded_extractor = self.save_to_zarr(**kwargs) - else: - loaded_extractor = self.save_to_folder(**kwargs) - return loaded_extractor + # save_kwargs.pop("format", None) - save.__doc__ = save.__doc__.format(_shared_job_kwargs_doc) + # cached = self._save(format="memory", sharedmem=sharedmem, **save_kwargs) + # self.copy_metadata(cached) + # return cached - def save_to_memory(self, sharedmem=True, **save_kwargs) -> "BaseExtractor": - save_kwargs.pop("format", None) + def _save_provenance_to_folder(self, folder): + provenance_file_path = folder / f"provenance.json" + if self.check_serializability("json"): + self.dump_to_json(file_path=provenance_file_path, relative_to=folder) + elif self.check_serializability("pickle"): + provenance_file = folder / f"provenance.pkl" + self.dump_to_pickle(provenance_file, relative_to=folder) + else: + warnings.warn("The extractor is not serializable to file. The provenance will not be saved.") - cached = self._save(format="memory", sharedmem=sharedmem, **save_kwargs) - self.copy_metadata(cached) - return cached # TODO rename to saveto_binary_folder def save_to_folder( @@ -1010,6 +1033,14 @@ def save_to_folder( If the folder already exists and `overwrite` is False. """ + warnings.warn("save_to_folder() should be recording.save(format='binray') " + "or sorting.save(format='numpy_folder') " + "This ambiguous method should not be used anymore!!", + FutureWarning) + # we keep the default format for recording and sorting like in old version + return self.save(folder=folder, verbose=verbose, **save_kwargs) + + if folder is None: cache_folder = get_global_tmp_folder() if name is None: diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index c61d602026..d23d062253 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -318,63 +318,71 @@ def get_data(self, start_frame: int, end_frame: int, segment_index: int | None = def get_shape(self, segment_index: int | None = None) -> tuple[int, ...]: return (self.get_num_samples(segment_index=segment_index), self.get_num_channels()) - def _save(self, format="binary", verbose: bool = False, **save_kwargs): + def save(self, format="binary", verbose: bool = False, **save_kwargs): kwargs, job_kwargs = split_job_kwargs(save_kwargs) if format == "binary": - from .time_series_tools import write_binary - from .binaryrecordingextractor import BinaryRecordingExtractor - from .binaryfolder import BinaryFolderRecording + # from .time_series_tools import write_binary + # from .binaryrecordingextractor import BinaryRecordingExtractor + # from .binaryfolder import BinaryFolderRecording + + # folder = kwargs["folder"] + # file_paths = [folder / f"traces_cached_seg{i}.raw" for i in range(self.get_num_segments())] + # dtype = kwargs.get("dtype", None) or self.get_dtype() + # t_starts = self._get_t_starts() + + # write_binary(self, file_paths=file_paths, dtype=dtype, verbose=verbose, **job_kwargs) + + # # This is created so it can be saved as json because the `BinaryFolderRecording` requires it loading + # # See the __init__ of `BinaryFolderRecording` + + # binary_rec = BinaryRecordingExtractor( + # file_paths=file_paths, + # sampling_frequency=self.get_sampling_frequency(), + # num_channels=self.get_num_channels(), + # dtype=dtype, + # t_starts=t_starts, + # channel_ids=self.get_channel_ids(), + # time_axis=0, + # file_offset=0, + # is_filtered=self.is_filtered(), + # gain_to_uV=self.get_channel_gains(), + # offset_to_uV=self.get_channel_offsets(), + # ) + # binary_rec.dump(folder / "binary.json", relative_to=folder) + # cached = BinaryFolderRecording(folder_path=folder) + + # # timestamps are not saved in binary, so we have to set them explicitly + # for segment_index in range(self.get_num_segments()): + # if self.has_time_vector(segment_index): + # # the use of get_times is preferred since timestamps are converted to array + # time_vector = self.get_times(segment_index=segment_index) + # cached.set_times(time_vector, segment_index=segment_index) - folder = kwargs["folder"] - file_paths = [folder / f"traces_cached_seg{i}.raw" for i in range(self.get_num_segments())] - dtype = kwargs.get("dtype", None) or self.get_dtype() - t_starts = self._get_t_starts() - - write_binary(self, file_paths=file_paths, dtype=dtype, verbose=verbose, **job_kwargs) - - # This is created so it can be saved as json because the `BinaryFolderRecording` requires it loading - # See the __init__ of `BinaryFolderRecording` - - binary_rec = BinaryRecordingExtractor( - file_paths=file_paths, - sampling_frequency=self.get_sampling_frequency(), - num_channels=self.get_num_channels(), - dtype=dtype, - t_starts=t_starts, - channel_ids=self.get_channel_ids(), - time_axis=0, - file_offset=0, - is_filtered=self.is_filtered(), - gain_to_uV=self.get_channel_gains(), - offset_to_uV=self.get_channel_offsets(), - ) - binary_rec.dump(folder / "binary.json", relative_to=folder) - cached = BinaryFolderRecording(folder_path=folder) - # timestamps are not saved in binary, so we have to set them explicitly - for segment_index in range(self.get_num_segments()): - if self.has_time_vector(segment_index): - # the use of get_times is preferred since timestamps are converted to array - time_vector = self.get_times(segment_index=segment_index) - cached.set_times(time_vector, segment_index=segment_index) + from .binaryfolder import BinaryFolderRecording + BinaryFolderRecording.write_recording(self, folder=kwargs["folder"], + dtype=kwargs.get("dtype", None), **job_kwargs) + cached = BinaryFolderRecording(folder_path=kwargs["folder"]) elif format == "memory": if kwargs.get("sharedmem", True): from .numpyextractors import SharedMemoryRecording - cached = SharedMemoryRecording.from_recording(self, **job_kwargs) + cached = SharedMemoryRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs) else: from spikeinterface.core import NumpyRecording - cached = NumpyRecording.from_recording(self, **job_kwargs) + cached = NumpyRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs) + + # self.copy_metadata(cached) - # timestamps are not saved in memory, so we have to set them explicitly - for segment_index in range(self.get_num_segments()): - if self.has_time_vector(segment_index): - # the use of get_times is preferred since timestamps are converted to array - time_vector = self.get_times(segment_index=segment_index) - cached.set_times(time_vector, segment_index=segment_index) + # # timestamps are not saved in memory, so we have to set them explicitly + # for segment_index in range(self.get_num_segments()): + # if self.has_time_vector(segment_index): + # # the use of get_times is preferred since timestamps are converted to array + # time_vector = self.get_times(segment_index=segment_index) + # cached.set_times(time_vector, segment_index=segment_index) elif format == "zarr": from .zarrextractors import ZarrRecordingExtractor @@ -387,10 +395,6 @@ def _save(self, format="binary", verbose: bool = False, **save_kwargs): cached = ZarrRecordingExtractor(zarr_path, storage_options) # timestamps are saved and restored in zarr, so no need to set them explicitly - elif format == "nwb": - # TODO implement a format based on zarr - raise NotImplementedError - else: raise ValueError(f"format {format} not supported") @@ -407,15 +411,15 @@ def _extra_metadata_from_folder(self, folder): time_vector = np.load(time_file, mmap_mode="r") rs.time_vector = time_vector - def _extra_metadata_to_folder(self, folder): - super()._extra_metadata_to_folder(folder) + # def _extra_metadata_to_folder(self, folder): + # super()._extra_metadata_to_folder(folder) - # save time vector if any - for segment_index, rs in enumerate(self.segments): - d = rs.get_times_kwargs() - time_vector = d["time_vector"] - if time_vector is not None: - np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) + # # save time vector if any + # for segment_index, rs in enumerate(self.segments): + # d = rs.get_times_kwargs() + # time_vector = d["time_vector"] + # if time_vector is not None: + # np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) def select_channels(self, channel_ids: list | np.ndarray | tuple) -> "BaseRecording": """ diff --git a/src/spikeinterface/core/baserecordingsnippets.py b/src/spikeinterface/core/baserecordingsnippets.py index c69d08297f..f26ce683a2 100644 --- a/src/spikeinterface/core/baserecordingsnippets.py +++ b/src/spikeinterface/core/baserecordingsnippets.py @@ -339,11 +339,11 @@ def _extra_metadata_from_folder(self, folder): if "contact_vector" in self.get_property_keys(): self.delete_property("contact_vector") - def _extra_metadata_to_folder(self, folder): - # save probe - if self.has_probe(): - probegroup = self.get_probegroup() - write_probeinterface(folder / "probegroup.json", probegroup) + # def _extra_metadata_to_folder(self, folder): + # # save probe + # if self.has_probe(): + # probegroup = self.get_probegroup() + # write_probeinterface(folder / "probegroup.json", probegroup) def _extra_metadata_from_dict(self, dump_dict): # load probe and hanlde backward-compatibility with legacy "contact_vector"/"location" property diff --git a/src/spikeinterface/core/basesorting.py b/src/spikeinterface/core/basesorting.py index f953a54703..86cc313f25 100644 --- a/src/spikeinterface/core/basesorting.py +++ b/src/spikeinterface/core/basesorting.py @@ -500,7 +500,7 @@ def get_times( else: return None - def _save(self, format="numpy_folder", **save_kwargs): + def save(self, format="numpy_folder", **save_kwargs): """ This function replaces the old CachesortingExtractor, but enables more engines for caching a results. diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index e9986193a3..a62b5dab06 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -3,8 +3,9 @@ import numpy as np -from probeinterface import read_probeinterface +from probeinterface import read_probeinterface, write_probeinterface +from spikeinterface.core import BaseRecording from .binaryrecordingextractor import BinaryRecordingExtractor from .core_tools import define_function_from_class, make_paths_absolute @@ -86,6 +87,8 @@ def __init__(self, folder_path): if probegroup is not None: self._probegroup = probegroup + + # self._load_metadata_from_folder(folder_path) # Load time vectors if any for segment_index, rs in enumerate(self.segments): @@ -112,5 +115,67 @@ def get_binary_description(self): ) return d + @staticmethod + def write_recording( + recording: BaseRecording, folder: str | Path, dtype=None, verbose=False, **job_kwargs + ): + from .time_series_tools import write_binary + from .binaryrecordingextractor import BinaryRecordingExtractor + from .binaryfolder import BinaryFolderRecording + + folder = Path(folder) + + file_paths = [folder / f"traces_cached_seg{i}.raw" for i in range(recording.get_num_segments())] + if dtype is None: + dtype = recording.get_dtype() + t_starts = recording._get_t_starts() + + write_binary(recording, file_paths=file_paths, dtype=dtype, verbose=verbose, **job_kwargs) + + recording._save_metadata_to_folder(folder) + recording._save_provenance_to_folder(folder) + + if recording.has_probe(): + probegroup = recording.get_probegroup() + write_probeinterface(folder / "probegroup.json", probegroup) + + + # This is created so it can be saved as json because the `BinaryFolderRecording` requires it loading + # See the __init__ + + binary_rec = BinaryRecordingExtractor( + file_paths=file_paths, + sampling_frequency=recording.get_sampling_frequency(), + num_channels=recording.get_num_channels(), + dtype=dtype, + t_starts=t_starts, + channel_ids=recording.get_channel_ids(), + time_axis=0, + file_offset=0, + is_filtered=recording.is_filtered(), + gain_to_uV=recording.get_channel_gains(), + offset_to_uV=recording.get_channel_offsets(), + ) + binary_rec.dump(folder / "binary.json", relative_to=folder) + + for segment_index, rs in enumerate(recording.segments): + d = rs.get_times_kwargs() + time_vector = d["time_vector"] + if time_vector is not None: + np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) + + # make the si_folder file to make the load() easier + cached = BinaryFolderRecording(folder_path=folder) + si_folder_path = folder / f"si_folder.json" + cached.dump_to_json(file_path=si_folder_path, relative_to=folder, include_extra_metadata=False) + + + # # timestamps are not saved in binary, so we have to set them explicitly + # for segment_index in range(recording.get_num_segments()): + # if recording.has_time_vector(segment_index): + # # the use of get_times is preferred since timestamps are converted to array + # time_vector = recording.get_times(segment_index=segment_index) + # cached.set_times(time_vector, segment_index=segment_index) + read_binary_folder = define_function_from_class(source_class=BinaryFolderRecording, name="read_binary_folder") diff --git a/src/spikeinterface/core/npyfoldersnippets.py b/src/spikeinterface/core/npyfoldersnippets.py index 1e465d827a..c0ebd1b81b 100644 --- a/src/spikeinterface/core/npyfoldersnippets.py +++ b/src/spikeinterface/core/npyfoldersnippets.py @@ -42,7 +42,7 @@ def __init__(self, folder_path): NpySnippetsExtractor.__init__(self, **d["kwargs"]) folder_metadata = folder_path - self.load_metadata_from_folder(folder_metadata) + self._load_metadata_from_folder(folder_metadata) self._kwargs = dict(folder_path=str(Path(folder_path).absolute())) self._bin_kwargs = d["kwargs"] diff --git a/src/spikeinterface/core/numpyextractors.py b/src/spikeinterface/core/numpyextractors.py index 07788853e0..6b3175bd7a 100644 --- a/src/spikeinterface/core/numpyextractors.py +++ b/src/spikeinterface/core/numpyextractors.py @@ -77,7 +77,7 @@ def __init__(self, traces_list, sampling_frequency, t_starts=None, channel_ids=N } @staticmethod - def from_recording(source_recording, **job_kwargs): + def from_recording(source_recording, with_metadata=True, with_time_vector=False, **job_kwargs): traces_list, shms = write_memory_recording(source_recording, dtype=None, **job_kwargs) t_starts = source_recording._get_t_starts() @@ -97,6 +97,18 @@ def from_recording(source_recording, **job_kwargs): t_starts=t_starts, channel_ids=source_recording.channel_ids, ) + + if with_metadata: + source_recording.copy_metadata(recording) + + if with_time_vector: + for segment_index in range(source_recording.get_num_segments()): + if source_recording.has_time_vector(segment_index): + # the use of get_times is preferred since timestamps are converted to array + time_vector = source_recording.get_times(segment_index=segment_index) + recording.set_times(time_vector, segment_index=segment_index) + + return recording diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index c0d66393d2..73ed70d4ce 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -45,7 +45,7 @@ def __init__(self, folder_path): self._cached_spike_vector = self.spikes folder_metadata = folder_path - self.load_metadata_from_folder(folder_metadata) + self._load_metadata_from_folder(folder_metadata) self._kwargs = dict(folder_path=str(folder_path.absolute())) @@ -66,6 +66,15 @@ def write_sorting(sorting, save_path): info_file.write_text(json.dumps(d), encoding="utf8") np.save(save_path / "spikes.npy", sorting.to_spike_vector()) + sorting._save_metadata_to_folder(save_path) + sorting._save_provenance_to_folder(save_path) + + # make the si_folder file to make the load() easier + cached = NumpyFolderSorting(folder_path=save_path) + si_folder_path = save_path / f"si_folder.json" + cached.dump_to_json(file_path=si_folder_path, relative_to=save_path, include_extra_metadata=False) + + class NpzFolderSorting(NpzSortingExtractor): """ @@ -108,7 +117,7 @@ def __init__(self, folder_path): NpzSortingExtractor.__init__(self, **d["kwargs"]) folder_metadata = folder_path - self.load_metadata_from_folder(folder_metadata) + self._load_metadata_from_folder(folder_metadata) self._kwargs = dict(folder_path=str(folder_path.absolute())) self._npz_kwargs = d["kwargs"] @@ -122,9 +131,15 @@ def write_sorting(sorting, save_path): if npz_file.exists(): raise ValueError("NpzFolderSorting.write_sorting the folder already contains sorting_cached.npz") NpzSortingExtractor.write_sorting(sorting, npz_file) + sorting._save_metadata_to_folder(save_path) cached = NpzSortingExtractor(npz_file) cached.dump(save_path / "npz.json", relative_to=save_path) + # make the si_folder file to make the load() easier + cached = NpzFolderSorting(folder_path=save_path) + si_folder_path = save_path / f"si_folder.json" + cached.dump_to_json(file_path=si_folder_path, relative_to=save_path, include_extra_metadata=False) + read_numpy_sorting_folder = define_function_from_class( source_class=NumpyFolderSorting, name="read_numpy_sorting_folder" From 2e386b56b62113a7a2fc3ff79003f7d3d69ff928 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 9 Jul 2026 13:54:29 +0000 Subject: [PATCH 02/18] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/spikeinterface/core/base.py | 12 ++++++------ src/spikeinterface/core/baserecording.py | 13 ++++++++----- src/spikeinterface/core/binaryfolder.py | 14 +++++--------- src/spikeinterface/core/numpyextractors.py | 3 +-- src/spikeinterface/core/sortingfolder.py | 1 - 5 files changed, 20 insertions(+), 23 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index fa6bf46b05..735546c382 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -972,7 +972,6 @@ def _save_provenance_to_folder(self, folder): else: warnings.warn("The extractor is not serializable to file. The provenance will not be saved.") - # TODO rename to saveto_binary_folder def save_to_folder( self, @@ -1033,14 +1032,15 @@ def save_to_folder( If the folder already exists and `overwrite` is False. """ - warnings.warn("save_to_folder() should be recording.save(format='binray') " - "or sorting.save(format='numpy_folder') " - "This ambiguous method should not be used anymore!!", - FutureWarning) + warnings.warn( + "save_to_folder() should be recording.save(format='binray') " + "or sorting.save(format='numpy_folder') " + "This ambiguous method should not be used anymore!!", + FutureWarning, + ) # we keep the default format for recording and sorting like in old version return self.save(folder=folder, verbose=verbose, **save_kwargs) - if folder is None: cache_folder = get_global_tmp_folder() if name is None: diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index d23d062253..b1aa3fa20c 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -359,22 +359,25 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): # time_vector = self.get_times(segment_index=segment_index) # cached.set_times(time_vector, segment_index=segment_index) - from .binaryfolder import BinaryFolderRecording - BinaryFolderRecording.write_recording(self, folder=kwargs["folder"], - dtype=kwargs.get("dtype", None), **job_kwargs) + + BinaryFolderRecording.write_recording( + self, folder=kwargs["folder"], dtype=kwargs.get("dtype", None), **job_kwargs + ) cached = BinaryFolderRecording(folder_path=kwargs["folder"]) elif format == "memory": if kwargs.get("sharedmem", True): from .numpyextractors import SharedMemoryRecording - cached = SharedMemoryRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs) + cached = SharedMemoryRecording.from_recording( + self, with_metadata=True, with_time_vector=True, **job_kwargs + ) else: from spikeinterface.core import NumpyRecording cached = NumpyRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs) - + # self.copy_metadata(cached) # # timestamps are not saved in memory, so we have to set them explicitly diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index a62b5dab06..9ce69f75d6 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -87,7 +87,7 @@ def __init__(self, folder_path): if probegroup is not None: self._probegroup = probegroup - + # self._load_metadata_from_folder(folder_path) # Load time vectors if any @@ -116,9 +116,7 @@ def get_binary_description(self): return d @staticmethod - def write_recording( - recording: BaseRecording, folder: str | Path, dtype=None, verbose=False, **job_kwargs - ): + def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, verbose=False, **job_kwargs): from .time_series_tools import write_binary from .binaryrecordingextractor import BinaryRecordingExtractor from .binaryfolder import BinaryFolderRecording @@ -139,9 +137,8 @@ def write_recording( probegroup = recording.get_probegroup() write_probeinterface(folder / "probegroup.json", probegroup) - # This is created so it can be saved as json because the `BinaryFolderRecording` requires it loading - # See the __init__ + # See the __init__ binary_rec = BinaryRecordingExtractor( file_paths=file_paths, @@ -157,7 +154,7 @@ def write_recording( offset_to_uV=recording.get_channel_offsets(), ) binary_rec.dump(folder / "binary.json", relative_to=folder) - + for segment_index, rs in enumerate(recording.segments): d = rs.get_times_kwargs() time_vector = d["time_vector"] @@ -169,13 +166,12 @@ def write_recording( si_folder_path = folder / f"si_folder.json" cached.dump_to_json(file_path=si_folder_path, relative_to=folder, include_extra_metadata=False) - # # timestamps are not saved in binary, so we have to set them explicitly # for segment_index in range(recording.get_num_segments()): # if recording.has_time_vector(segment_index): # # the use of get_times is preferred since timestamps are converted to array # time_vector = recording.get_times(segment_index=segment_index) - # cached.set_times(time_vector, segment_index=segment_index) + # cached.set_times(time_vector, segment_index=segment_index) read_binary_folder = define_function_from_class(source_class=BinaryFolderRecording, name="read_binary_folder") diff --git a/src/spikeinterface/core/numpyextractors.py b/src/spikeinterface/core/numpyextractors.py index 6b3175bd7a..454d97098d 100644 --- a/src/spikeinterface/core/numpyextractors.py +++ b/src/spikeinterface/core/numpyextractors.py @@ -100,7 +100,7 @@ def from_recording(source_recording, with_metadata=True, with_time_vector=False, if with_metadata: source_recording.copy_metadata(recording) - + if with_time_vector: for segment_index in range(source_recording.get_num_segments()): if source_recording.has_time_vector(segment_index): @@ -108,7 +108,6 @@ def from_recording(source_recording, with_metadata=True, with_time_vector=False, time_vector = source_recording.get_times(segment_index=segment_index) recording.set_times(time_vector, segment_index=segment_index) - return recording diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index 73ed70d4ce..e3849e5d33 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -75,7 +75,6 @@ def write_sorting(sorting, save_path): cached.dump_to_json(file_path=si_folder_path, relative_to=save_path, include_extra_metadata=False) - class NpzFolderSorting(NpzSortingExtractor): """ NpzFolderSorting is the old internal format used in spikeinterface (<=0.98.0) From 14c780a030da24c58304f6e5c209cabbf7ea2b70 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 8 Sep 2026 10:56:08 +0000 Subject: [PATCH 03/18] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/spikeinterface/core/base.py | 1 - src/spikeinterface/core/baserecording.py | 1 - src/spikeinterface/core/basesnippets.py | 1 - src/spikeinterface/core/basesorting.py | 1 - src/spikeinterface/core/binaryfolder.py | 8 +++++++- src/spikeinterface/core/core_tools.py | 3 +-- src/spikeinterface/core/npyfoldersnippets.py | 11 +++++++---- src/spikeinterface/core/sortingfolder.py | 8 +++++++- .../core/tests/test_noise_levels_propagation.py | 2 +- src/spikeinterface/core/tests/test_sortinganalyzer.py | 1 + 10 files changed, 24 insertions(+), 13 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index 83a2269824..109ade1531 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -653,7 +653,6 @@ def from_dict(dictionary: dict, base_folder: Path | str | None = None) -> "BaseE # load_properties_from_binary_folder(folder_metadata, self) return extractor - def clone(self) -> "BaseExtractor": """ Clones an existing extractor into a new instance. diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 30d5386f9d..5503d1c327 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -329,7 +329,6 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): self, folder=kwargs["folder"], dtype=kwargs.get("dtype", None), **job_kwargs ) - elif format == "memory": if kwargs.get("sharedmem", True): from .numpyextractors import SharedMemoryRecording diff --git a/src/spikeinterface/core/basesnippets.py b/src/spikeinterface/core/basesnippets.py index 6bcbebdf3f..df8fc0cbaf 100644 --- a/src/spikeinterface/core/basesnippets.py +++ b/src/spikeinterface/core/basesnippets.py @@ -9,7 +9,6 @@ from .core_tools import save_properties_to_binary_folder - # snippets segments? diff --git a/src/spikeinterface/core/basesorting.py b/src/spikeinterface/core/basesorting.py index d73ea0cffe..9fdd256bd5 100644 --- a/src/spikeinterface/core/basesorting.py +++ b/src/spikeinterface/core/basesorting.py @@ -502,7 +502,6 @@ def get_times( else: return None - def save(self, format="numpy_folder", **save_kwargs): """ Save a sorting object to disk in a specified format. diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index 14dfa20f40..5af197af07 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -9,7 +9,13 @@ from spikeinterface.core import BaseRecording from .binaryrecordingextractor import BinaryRecordingExtractor -from .core_tools import define_function_from_class, make_paths_absolute, load_properties_from_binary_folder, save_properties_to_binary_folder, save_extractor_provenance +from .core_tools import ( + define_function_from_class, + make_paths_absolute, + load_properties_from_binary_folder, + save_properties_to_binary_folder, + save_extractor_provenance, +) class BinaryFolderRecording(BinaryRecordingExtractor): diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index fbd77e5113..80769c4714 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -806,8 +806,6 @@ def slice_rows(array: np.ndarray | zarr.Array, row_indices: np.ndarray | list) - return array[row_indices, ...] - - def load_properties_from_binary_folder(folder: str | Path, extractor: "BaseExtractor") -> dict: """ Load properties from a folder properties as .npy files and return sets them @@ -848,6 +846,7 @@ def save_properties_to_binary_folder(folder: str | Path, extractor: "BaseExtract values = extractor.get_property(key) np.save(folder / f"{key}.npy", values, allow_pickle=True) + def save_extractor_provenance(folder: str | Path, extractor: "BaseExtractor"): folder = Path(folder) if extractor.check_serializability("json"): diff --git a/src/spikeinterface/core/npyfoldersnippets.py b/src/spikeinterface/core/npyfoldersnippets.py index c702ee486b..9e026904ea 100644 --- a/src/spikeinterface/core/npyfoldersnippets.py +++ b/src/spikeinterface/core/npyfoldersnippets.py @@ -6,8 +6,12 @@ from probeinterface import read_probeinterface, write_probeinterface from .npysnippetsextractor import NpySnippetsExtractor -from .core_tools import define_function_from_class, make_paths_absolute, load_properties_from_binary_folder, save_properties_to_binary_folder - +from .core_tools import ( + define_function_from_class, + make_paths_absolute, + load_properties_from_binary_folder, + save_properties_to_binary_folder, +) class NpyFolderSnippets(NpySnippetsExtractor): @@ -54,7 +58,7 @@ def __init__(self, folder_path): self._kwargs = dict(folder_path=str(Path(folder_path).absolute())) self._bin_kwargs = d["kwargs"] - + @staticmethod def write_snippets(snippets, folder, dtype=None): @@ -87,7 +91,6 @@ def write_snippets(snippets, folder, dtype=None): probegroup = snippets.get_probegroup() write_probeinterface(folder / "probegroup.json", probegroup) - cached = NpyFolderSnippets(folder_path=folder) # important backward compatibility : annoations are handled (sadly) only is this file # so we need to set then here (sad hack) diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index d499d1a2a1..6a9752aa48 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -6,7 +6,13 @@ from .basesorting import BaseSorting, SpikeVectorSortingSegment from .npzsortingextractor import NpzSortingExtractor -from .core_tools import define_function_from_class, make_paths_absolute, load_properties_from_binary_folder, save_properties_to_binary_folder, save_extractor_provenance +from .core_tools import ( + define_function_from_class, + make_paths_absolute, + load_properties_from_binary_folder, + save_properties_to_binary_folder, + save_extractor_provenance, +) class NumpyFolderSorting(BaseSorting): diff --git a/src/spikeinterface/core/tests/test_noise_levels_propagation.py b/src/spikeinterface/core/tests/test_noise_levels_propagation.py index 690290af3e..98e7a929fd 100644 --- a/src/spikeinterface/core/tests/test_noise_levels_propagation.py +++ b/src/spikeinterface/core/tests/test_noise_levels_propagation.py @@ -11,7 +11,7 @@ def test_skip_noise_levels_propagation(create_cache_folder): rec = generate_recording(durations=[5.0], num_channels=4) rec.set_property("test", ["1", "2", "3", "4"]) - rec = rec.save(folder=create_cache_folder/"rec_saved_noise") + rec = rec.save(folder=create_cache_folder / "rec_saved_noise") noise_level_raw = get_noise_levels(rec, return_in_uV=False, method="mad") assert "noise_level_mad_raw" in rec.get_property_keys() diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index 0616c9f998..9b360d3ec3 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -985,6 +985,7 @@ def test_main_channel_from_templates_sparse_recordingless(): if __name__ == "__main__": import tempfile from pathlib import Path + tmp_path = Path(tempfile.mkdtemp()) / "test_SortingAnalyzer" dataset = get_dataset() From b4d184c803f152f15bff4ccbedbd72a16ee3a6c7 Mon Sep 17 00:00:00 2001 From: Samuel Garcia Date: Tue, 8 Sep 2026 16:33:17 +0200 Subject: [PATCH 04/18] Fix some rec.save() without folder --- .github/scripts/test_kilosort4_ci.py | 2 +- .../comparison/tests/test_multisortingcomparison.py | 10 ++++++---- src/spikeinterface/core/baserecording.py | 5 +++++ src/spikeinterface/core/binaryfolder.py | 10 +++++----- src/spikeinterface/core/loading.py | 2 +- src/spikeinterface/core/sortingfolder.py | 1 + src/spikeinterface/preprocessing/tests/test_filter.py | 11 ++++++++--- .../preprocessing/tests/test_resample.py | 2 +- src/spikeinterface/sorters/basesorter.py | 2 +- 9 files changed, 29 insertions(+), 16 deletions(-) diff --git a/.github/scripts/test_kilosort4_ci.py b/.github/scripts/test_kilosort4_ci.py index 75bac8f03b..b9ebaffb21 100644 --- a/.github/scripts/test_kilosort4_ci.py +++ b/.github/scripts/test_kilosort4_ci.py @@ -481,7 +481,7 @@ def test_use_binary_file(self, tmp_path): from the recording. """ recording = self._get_ground_truth_recording() - recording_bin = recording.save() + recording_bin = recording.save(folder=tmp_path / "recording_for_ks4") # run with SI wrapper sorting_ks4 = si.run_sorter( diff --git a/src/spikeinterface/comparison/tests/test_multisortingcomparison.py b/src/spikeinterface/comparison/tests/test_multisortingcomparison.py index 6fad43bde0..afec6a5e4f 100644 --- a/src/spikeinterface/comparison/tests/test_multisortingcomparison.py +++ b/src/spikeinterface/comparison/tests/test_multisortingcomparison.py @@ -22,19 +22,20 @@ def setup_module(tmp_path_factory): return multicomparison_folder -def make_sorting(times1, labels1, times2, labels2, times3, labels3): +def make_sorting(times1, labels1, times2, labels2, times3, labels3, sorting_folder): sampling_frequency = 30000.0 sorting1 = NumpySorting.from_samples_and_labels([times1], [labels1], sampling_frequency) sorting2 = NumpySorting.from_samples_and_labels([times2], [labels2], sampling_frequency) sorting3 = NumpySorting.from_samples_and_labels([times3], [labels3], sampling_frequency) - sorting1 = sorting1.save() - sorting2 = sorting2.save() - sorting3 = sorting3.save() + sorting1 = sorting1.save(folder=sorting_folder/"sorting1") + sorting2 = sorting2.save(folder=sorting_folder/"sorting2") + sorting3 = sorting3.save(folder=sorting_folder/"sorting3") return sorting1, sorting2, sorting3 def test_compare_multiple_sorters(setup_module): multicomparison_folder = setup_module + sorting_folder = multicomparison_folder.parent # simple match sorting1, sorting2, sorting3 = make_sorting( [100, 200, 300, 400, 500, 600, 700, 800, 900], @@ -43,6 +44,7 @@ def test_compare_multiple_sorters(setup_module): [0, 1, 2, 0, 1, 2, 0, 1, 2, 3, 3, 4, 4], [101, 201, 301, 400, 500, 600, 700, 800, 900, 1000, 1100, 2000, 3000, 3100, 3200, 3300], [0, 1, 2, 0, 1, 2, 0, 1, 2, 3, 3, 4, 4, 5, 5, 5], + sorting_folder ) msc = compare_multiple_sorters([sorting1, sorting2, sorting3], verbose=True) msc_shuffle = compare_multiple_sorters([sorting3, sorting1, sorting2]) diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 30d5386f9d..0a0bf86c9a 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -322,6 +322,8 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): kwargs, job_kwargs = split_job_kwargs(save_kwargs) if format == "binary": + if "folder" not in kwargs: + raise ValueError("Missing folder in recording.save(folder='...')") from .binaryfolder import BinaryFolderRecording @@ -352,6 +354,9 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): # cached.set_times(time_vector, segment_index=segment_index) elif format == "zarr": + if "folder" not in kwargs: + raise ValueError("Missing folder in recording.save(folder='...')") + from .zarrextractors import ZarrRecordingExtractor folder_path = kwargs["folder"] diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index 14dfa20f40..c53fbfdca8 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -148,11 +148,11 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve ) binary_rec.dump(folder / "binary.json", relative_to=folder) - for segment_index, rs in enumerate(recording.segments): - d = rs.get_times_kwargs() - time_vector = d["time_vector"] - if time_vector is not None: - np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) + # for segment_index, rs in enumerate(recording.segments): + # d = rs.get_times_kwargs() + # time_vector = d["time_vector"] + # if time_vector is not None: + # np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) # make the si_folder file to make the load() easier cached = BinaryFolderRecording(folder_path=folder) diff --git a/src/spikeinterface/core/loading.py b/src/spikeinterface/core/loading.py index 93ac7c2a56..40e3530be7 100644 --- a/src/spikeinterface/core/loading.py +++ b/src/spikeinterface/core/loading.py @@ -196,7 +196,7 @@ def _guess_object_from_local_folder(folder): # before the SortingAnlazer, it was WaveformExtractor (v<0.101) return "WaveformExtractor" elif (folder / f"si_folder.json").is_file(): - # In later versions (0.94 Date: Tue, 8 Sep 2026 14:34:11 +0000 Subject: [PATCH 05/18] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../comparison/tests/test_multisortingcomparison.py | 8 ++++---- src/spikeinterface/core/baserecording.py | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/src/spikeinterface/comparison/tests/test_multisortingcomparison.py b/src/spikeinterface/comparison/tests/test_multisortingcomparison.py index afec6a5e4f..72f39f74ea 100644 --- a/src/spikeinterface/comparison/tests/test_multisortingcomparison.py +++ b/src/spikeinterface/comparison/tests/test_multisortingcomparison.py @@ -27,9 +27,9 @@ def make_sorting(times1, labels1, times2, labels2, times3, labels3, sorting_fold sorting1 = NumpySorting.from_samples_and_labels([times1], [labels1], sampling_frequency) sorting2 = NumpySorting.from_samples_and_labels([times2], [labels2], sampling_frequency) sorting3 = NumpySorting.from_samples_and_labels([times3], [labels3], sampling_frequency) - sorting1 = sorting1.save(folder=sorting_folder/"sorting1") - sorting2 = sorting2.save(folder=sorting_folder/"sorting2") - sorting3 = sorting3.save(folder=sorting_folder/"sorting3") + sorting1 = sorting1.save(folder=sorting_folder / "sorting1") + sorting2 = sorting2.save(folder=sorting_folder / "sorting2") + sorting3 = sorting3.save(folder=sorting_folder / "sorting3") return sorting1, sorting2, sorting3 @@ -44,7 +44,7 @@ def test_compare_multiple_sorters(setup_module): [0, 1, 2, 0, 1, 2, 0, 1, 2, 3, 3, 4, 4], [101, 201, 301, 400, 500, 600, 700, 800, 900, 1000, 1100, 2000, 3000, 3100, 3200, 3300], [0, 1, 2, 0, 1, 2, 0, 1, 2, 3, 3, 4, 4, 5, 5, 5], - sorting_folder + sorting_folder, ) msc = compare_multiple_sorters([sorting1, sorting2, sorting3], verbose=True) msc_shuffle = compare_multiple_sorters([sorting3, sorting1, sorting2]) diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 541e20fece..0aa3ea8820 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -355,7 +355,7 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): elif format == "zarr": if "folder" not in kwargs: raise ValueError("Missing folder in recording.save(folder='...')") - + from .zarrextractors import ZarrRecordingExtractor folder_path = kwargs["folder"] From c78c84ae9e2f9e8599c7e4d2020003fd36d6ad9b Mon Sep 17 00:00:00 2001 From: Samuel Garcia Date: Tue, 8 Sep 2026 17:00:28 +0200 Subject: [PATCH 06/18] Create an explicit annoations.json file for BinaryFolderRecording and NumpyFolderSorting --- .gitignore | 4 +++ src/spikeinterface/core/binaryfolder.py | 31 ++++++++-------- src/spikeinterface/core/core_tools.py | 45 ++++++++++++++++++++++++ src/spikeinterface/core/sortingfolder.py | 18 +++++++--- 4 files changed, 78 insertions(+), 20 deletions(-) diff --git a/.gitignore b/.gitignore index 2baa4b4f92..1bf2060a46 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,10 @@ spikeinterface/widgets/tests/mearec_test/* +# for sam +dev_*.ipynb +dev_*.py + # Vscode .vscode/* diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index 5295348390..210ed75276 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -15,6 +15,8 @@ load_properties_from_binary_folder, save_properties_to_binary_folder, save_extractor_provenance, + save_annotations_to_folder, + load_annotations_from_folder, ) @@ -52,6 +54,7 @@ def __init__(self, folder_path): # Load properties load_properties_from_binary_folder(folder_path / "properties", self) + load_annotations_from_folder(folder_path, self) # Load the probegroup probe_file = folder_path / "probegroup.json" @@ -131,6 +134,9 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve save_properties_to_binary_folder(folder / "properties", recording) save_extractor_provenance(folder, recording) + # new in version 0.105.0, before that annotations were handle by "si_folder.json" file + save_annotations_to_folder(folder, recording) + if recording.has_probe(): probegroup = recording.get_probegroup() @@ -154,27 +160,20 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve ) binary_rec.dump(folder / "binary.json", relative_to=folder) - # for segment_index, rs in enumerate(recording.segments): - # d = rs.get_times_kwargs() - # time_vector = d["time_vector"] - # if time_vector is not None: - # np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) + # TODO alessio : remove this, it is needed to pass tests + # save times + for segment_index, rs in enumerate(recording.segments): + d = rs.get_times_kwargs() + time_vector = d["time_vector"] + if time_vector is not None: + np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) + - # make the si_folder file to make the load() easier + # make the si_folder file to make the load() easier until version 0.105.0 cached = BinaryFolderRecording(folder_path=folder) - # important backward compatibility : annoations are handled (sadly) only is this file - # so we need to set then here (sad hack) - cached._annotations = deepcopy({k: recording._annotations[k] for k in recording._annotations.keys()}) si_folder_path = folder / f"si_folder.json" cached.dump_to_json(file_path=si_folder_path, relative_to=folder, include_extra_metadata=False) - # # timestamps are not saved in binary, so we have to set them explicitly - # for segment_index in range(recording.get_num_segments()): - # if recording.has_time_vector(segment_index): - # # the use of get_times is preferred since timestamps are converted to array - # time_vector = recording.get_times(segment_index=segment_index) - # cached.set_times(time_vector, segment_index=segment_index) - return cached diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index 80769c4714..d76ab4ce2e 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -846,6 +846,51 @@ def save_properties_to_binary_folder(folder: str | Path, extractor: "BaseExtract values = extractor.get_property(key) np.save(folder / f"{key}.npy", values, allow_pickle=True) +def save_annotations_to_folder(folder: str | Path, extractor: "BaseExtractor"): + """ + Save annotaions in json format from an extractor (recording or sorting). + This used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0 + + Parameters + ---------- + folder : str or Path + The folder where the properties will be saved as .npy files. + extractor : BaseExtractor + The extractor from which the annotations will be saved. + """ + folder = Path(folder) + (folder / "annotations").write_text(json.dumps(extractor._annotations, indent=4), encoding="utf8") + +def load_annotations_from_folder(folder: str | Path, extractor: "BaseExtractor"): + """ + Save annotaions in json format from an extractor (recording or sorting). + This used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0 + + Parameters + ---------- + folder : str or Path + The folder where the properties will be saved as .npy files. + extractor : BaseExtractor + The extractor from which the annotations will be saved. + """ + folder = Path(folder) + annotations_file = folder / "annotations" + if annotations_file.exists(): + with open(annotations_file, "r") as f: + annotations = json.load(f) + extractor._annotations.update(annotations) + else: + # this was before 0.105.0 + si_folder_json = folder / "si_folder.json" + if si_folder_json.is_file(): + with open(si_folder_json, "r") as f: + si_folder_dict = json.load(f) + if "annotations" in si_folder_dict: + annotations = si_folder_dict["annotations"] + extractor._annotations.update(annotations) + + + def save_extractor_provenance(folder: str | Path, extractor: "BaseExtractor"): folder = Path(folder) diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index 630eb0daf8..5d1a587418 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -1,6 +1,7 @@ from pathlib import Path import json from copy import deepcopy +import warnings import numpy as np @@ -11,6 +12,8 @@ make_paths_absolute, load_properties_from_binary_folder, save_properties_to_binary_folder, + save_annotations_to_folder, + load_annotations_from_folder, save_extractor_provenance, ) @@ -52,6 +55,7 @@ def __init__(self, folder_path, mmap_mode: str | None = None): self._cached_spike_vector = self.spikes load_properties_from_binary_folder(folder_path / "properties", self) + load_annotations_from_folder(folder_path, self) self._kwargs = dict(folder_path=str(folder_path.absolute()), mmap_mode=mmap_mode) @@ -74,12 +78,11 @@ def write_sorting(sorting, save_path): save_properties_to_binary_folder(save_path / "properties", sorting) save_extractor_provenance(save_path, sorting) + # new in version 0.105.0, before that annotations were handle by "si_folder.json" file + save_annotations_to_folder(save_path, sorting) - # make the si_folder file to make the load() easier + # make the si_folder file to make the load() easier until version 0.105.0 cached = NumpyFolderSorting(folder_path=save_path) - # important backward compatibility : annoations are handled (sadly) only is this file - # so we need to set then here (sad hack) - cached._annotations = deepcopy({k: sorting._annotations[k] for k in sorting._annotations.keys()}) si_folder_path = save_path / f"si_folder.json" cached.dump_to_json(file_path=si_folder_path, relative_to=save_path, include_extra_metadata=False) @@ -127,12 +130,19 @@ def __init__(self, folder_path): NpzSortingExtractor.__init__(self, **d["kwargs"]) load_properties_from_binary_folder(folder_path / "properties", self) + load_annotations_from_folder(folder_path, self) self._kwargs = dict(folder_path=str(folder_path.absolute())) self._npz_kwargs = d["kwargs"] @staticmethod def write_sorting(sorting, save_path): + warnings.warn( + "`NpzFolderSorting.write_sorting()` is deprecated and will be removed in 0.106.0. The NpzFolderSorting() read will stay for a while", + category=FutureWarning, + stacklevel=2, + ) + save_path = Path(save_path) save_path.mkdir(parents=True, exist_ok=True) From 653bfff0f24d4be2f44e041468872bab4767d68f Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 8 Sep 2026 15:01:25 +0000 Subject: [PATCH 07/18] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/spikeinterface/core/binaryfolder.py | 2 -- src/spikeinterface/core/core_tools.py | 4 ++-- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index 210ed75276..a462ad806f 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -137,7 +137,6 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve # new in version 0.105.0, before that annotations were handle by "si_folder.json" file save_annotations_to_folder(folder, recording) - if recording.has_probe(): probegroup = recording.get_probegroup() write_probeinterface(folder / "probegroup.json", probegroup) @@ -168,7 +167,6 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve if time_vector is not None: np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) - # make the si_folder file to make the load() easier until version 0.105.0 cached = BinaryFolderRecording(folder_path=folder) si_folder_path = folder / f"si_folder.json" diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index d76ab4ce2e..8e24027064 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -846,6 +846,7 @@ def save_properties_to_binary_folder(folder: str | Path, extractor: "BaseExtract values = extractor.get_property(key) np.save(folder / f"{key}.npy", values, allow_pickle=True) + def save_annotations_to_folder(folder: str | Path, extractor: "BaseExtractor"): """ Save annotaions in json format from an extractor (recording or sorting). @@ -861,6 +862,7 @@ def save_annotations_to_folder(folder: str | Path, extractor: "BaseExtractor"): folder = Path(folder) (folder / "annotations").write_text(json.dumps(extractor._annotations, indent=4), encoding="utf8") + def load_annotations_from_folder(folder: str | Path, extractor: "BaseExtractor"): """ Save annotaions in json format from an extractor (recording or sorting). @@ -890,8 +892,6 @@ def load_annotations_from_folder(folder: str | Path, extractor: "BaseExtractor") extractor._annotations.update(annotations) - - def save_extractor_provenance(folder: str | Path, extractor: "BaseExtractor"): folder = Path(folder) if extractor.check_serializability("json"): From 8ed0644f97f6db6e82af9aa949b444756cafcc25 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 8 Sep 2026 18:07:51 +0200 Subject: [PATCH 08/18] fix: save() functions in tests --- .../preprocessing/tests/test_clip.py | 19 +++++++++++-------- .../tests/test_common_reference.py | 12 ++++++------ .../preprocessing/tests/test_filter.py | 16 +++++++++++----- .../tests/test_normalize_scale.py | 6 +++--- .../preprocessing/tests/test_rectify.py | 4 ++-- .../preprocessing/tests/test_whiten.py | 2 +- .../preprocessing/tests/test_zero_padding.py | 4 ++-- 7 files changed, 36 insertions(+), 27 deletions(-) diff --git a/src/spikeinterface/preprocessing/tests/test_clip.py b/src/spikeinterface/preprocessing/tests/test_clip.py index 96020692a1..2e6d30d9d7 100644 --- a/src/spikeinterface/preprocessing/tests/test_clip.py +++ b/src/spikeinterface/preprocessing/tests/test_clip.py @@ -5,14 +5,14 @@ import numpy as np -def test_clip(): +def test_clip(create_cache_folder): rec = generate_recording() rec0 = clip(rec, a_min=-2, a_max=3.0) - rec0.save(verbose=False) + rec0.save(folder=create_cache_folder / "rec_clip0", verbose=False) rec1 = clip(rec, a_min=-1.5) - rec1.save(verbose=False) + rec1.save(folder=create_cache_folder / "rec_clip1", verbose=False) traces0 = rec0.get_traces(segment_index=0, channel_ids=["1"]) assert traces0.shape[1] == 1 @@ -25,14 +25,14 @@ def test_clip(): assert np.all(-1.5 <= traces1[1]) -def test_blank_saturation(): +def test_blank_saturation(create_cache_folder): rec = generate_recording() rec0 = blank_saturation(rec, abs_threshold=3.0) - rec0.save(verbose=False) + rec0.save(folder=create_cache_folder / "rec_sat0", verbose=False) rec1 = blank_saturation(rec, quantile_threshold=0.01, direction="both", chunk_size=10000) - rec1.save(verbose=False) + rec1.save(folder=create_cache_folder / "rec_sat1", verbose=False) traces0 = rec0.get_traces(segment_index=0, channel_ids=["1"]) assert traces0.shape[1] == 1 @@ -46,5 +46,8 @@ def test_blank_saturation(): if __name__ == "__main__": - test_clip() - test_blank_saturation() + import tempfile + + cache_folder = tempfile.TemporaryDirectory() + test_clip(create_cache_folder=cache_folder.name) + test_blank_saturation(create_cache_folder=cache_folder.name) diff --git a/src/spikeinterface/preprocessing/tests/test_common_reference.py b/src/spikeinterface/preprocessing/tests/test_common_reference.py index 37e658f41e..139213056f 100644 --- a/src/spikeinterface/preprocessing/tests/test_common_reference.py +++ b/src/spikeinterface/preprocessing/tests/test_common_reference.py @@ -20,7 +20,7 @@ def recording(): return _generate_test_recording() -def test_common_reference(recording): +def test_common_reference(recording, create_cache_folder): # Test simple case rec_cmr = common_reference(recording, reference="global", operator="median") rec_cmr_ref = common_reference(recording, reference="global", operator="median", ref_channel_ids=["a", "b", "c"]) @@ -47,11 +47,11 @@ def test_common_reference(recording): assert np.allclose(traces[:, 1], rec_local_car.get_traces()[:, 1] + np.mean(traces[:, [3]], axis=1), atol=0.01) # Saving tests - rec_cmr.save(verbose=False) - rec_car.save(verbose=False) - rec_sin.save(verbose=False) - rec_local_cmr.save(verbose=False) - rec_local_car.save(verbose=False) + rec_cmr.save(folder=create_cache_folder / "rec_cmr", verbose=False) + rec_car.save(folder=create_cache_folder / "rec_car", verbose=False) + rec_sin.save(folder=create_cache_folder / "rec_sin", verbose=False) + rec_local_cmr.save(folder=create_cache_folder / "rec_local_cmr", verbose=False) + rec_local_car.save(folder=create_cache_folder / "rec_local_car", verbose=False) def test_common_reference_channel_slicing(recording): diff --git a/src/spikeinterface/preprocessing/tests/test_filter.py b/src/spikeinterface/preprocessing/tests/test_filter.py index 478df3f3b6..3b47d0eb1a 100644 --- a/src/spikeinterface/preprocessing/tests/test_filter.py +++ b/src/spikeinterface/preprocessing/tests/test_filter.py @@ -147,13 +147,15 @@ def test_filter(create_cache_folder): rec2 = bandpass_filter(rec, freq_min=300.0, freq_max=6000.0) # compute by chunk - rec2_cached0 = rec2.save(chunk_size=100000, verbose=False, progress_bar=True) + rec2_cached0 = rec2.save( + folder=create_cache_folder / "rec2_cached0", chunk_size=100000, verbose=False, progress_bar=True + ) # compute by chunkf with joblib - rec2_cached1 = rec2.save(total_memory="10k", n_jobs=4, verbose=True) + rec2_cached1 = rec2.save(folder=create_cache_folder / "rec2_cached1", total_memory="10k", n_jobs=4, verbose=True) # compute once - rec2_cached2 = rec2.save(verbose=False) + rec2_cached2 = rec2.save(folder=create_cache_folder / "rec2_cached2", verbose=False) trace0 = rec2.get_traces(segment_index=0) trace1 = rec2_cached1.get_traces(segment_index=0) @@ -171,7 +173,9 @@ def test_filter(create_cache_folder): rec5 = filter(rec, coeff=coeff, filter_mode="sos", margin_ms=5.0) # compute by chunk - rec5_cached0 = rec5.save(chunk_size=100000, verbose=False, progress_bar=True) + rec5_cached0 = rec5.save( + folder=create_cache_folder / "rec5_cached0", chunk_size=100000, verbose=False, progress_bar=True + ) trace50 = rec5.get_traces(segment_index=0) trace51 = rec5_cached0.get_traces(segment_index=0) @@ -180,7 +184,9 @@ def test_filter(create_cache_folder): # reflect padding test rec6 = bandpass_filter(rec, freq_min=300.0, freq_max=6000.0, add_reflect_padding=True) - rec6_cached = rec6.save(chunk_size=150000, verbose=False, progress_bar=True) + rec6_cached = rec6.save( + folder=create_cache_folder / "rec6_cached", chunk_size=150000, verbose=False, progress_bar=True + ) trace0 = rec6.get_traces(segment_index=0) trace1 = rec6_cached.get_traces(segment_index=0) diff --git a/src/spikeinterface/preprocessing/tests/test_normalize_scale.py b/src/spikeinterface/preprocessing/tests/test_normalize_scale.py index b366fb2df7..a304a18f18 100644 --- a/src/spikeinterface/preprocessing/tests/test_normalize_scale.py +++ b/src/spikeinterface/preprocessing/tests/test_normalize_scale.py @@ -5,17 +5,17 @@ from spikeinterface.preprocessing import normalize_by_quantile, scale, center, zscore -def test_normalize_by_quantile(): +def test_normalize_by_quantile(create_cache_folder): rec = generate_recording() rec2 = normalize_by_quantile(rec, mode="by_channel") - rec2.save(verbose=False) + rec2.save(folder=create_cache_folder / "rec2", verbose=False) traces = rec2.get_traces(segment_index=0, channel_ids=["1"]) assert traces.shape[1] == 1 rec2 = normalize_by_quantile(rec, mode="pool_channel") - rec2.save(verbose=False) + rec2.save(folder=create_cache_folder / "rec2_pool", verbose=False) # import matplotlib.pyplot as plt # from spikeinterface.widgets import plot_traces diff --git a/src/spikeinterface/preprocessing/tests/test_rectify.py b/src/spikeinterface/preprocessing/tests/test_rectify.py index f99d9f5ef9..bfd998d294 100644 --- a/src/spikeinterface/preprocessing/tests/test_rectify.py +++ b/src/spikeinterface/preprocessing/tests/test_rectify.py @@ -2,11 +2,11 @@ from spikeinterface.preprocessing import rectify -def test_rectify(): +def test_rectify(create_cache_folder): rec = generate_recording() rec2 = rectify(rec) - rec2.save(verbose=False) + rec2.save(folder=create_cache_folder / "rec2", verbose=False) traces = rec2.get_traces(segment_index=0, channel_ids=["1"]) assert traces.shape[1] == 1 diff --git a/src/spikeinterface/preprocessing/tests/test_whiten.py b/src/spikeinterface/preprocessing/tests/test_whiten.py index f96a2f7925..9f9780b366 100644 --- a/src/spikeinterface/preprocessing/tests/test_whiten.py +++ b/src/spikeinterface/preprocessing/tests/test_whiten.py @@ -419,7 +419,7 @@ def test_whiten_general(self, create_cache_folder): np.sum(W == 0) == 6 rec2 = whiten(rec) - rec2.save(verbose=False) + rec2.save(folder=cache_folder / "rec2", verbose=False) # test dtype rec_int = scale(rec2, dtype="int16") diff --git a/src/spikeinterface/preprocessing/tests/test_zero_padding.py b/src/spikeinterface/preprocessing/tests/test_zero_padding.py index f0eb443cde..b6fcb00f03 100644 --- a/src/spikeinterface/preprocessing/tests/test_zero_padding.py +++ b/src/spikeinterface/preprocessing/tests/test_zero_padding.py @@ -8,13 +8,13 @@ from spikeinterface.preprocessing.zero_channel_pad import TracePaddedRecording -def test_zero_padding_channel(): +def test_zero_padding_channel(create_cache_folder): num_original_channels = 4 num_padded_channels = num_original_channels + 8 rec = generate_recording(num_channels=num_original_channels, durations=[10]) rec2 = zero_channel_pad(rec, num_channels=num_padded_channels) - rec2.save(verbose=False) + rec2.save(folder=create_cache_folder / "rec2", verbose=False) print(rec2) From a89048285f68b45b6c550551b7fa3be34c65b637 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 9 Sep 2026 10:59:26 +0200 Subject: [PATCH 09/18] fix: tests and add some TODOs --- src/spikeinterface/core/base.py | 235 ++++++++++-------- src/spikeinterface/core/baserecording.py | 5 + src/spikeinterface/core/tests/test_loading.py | 2 +- .../preprocessing/tests/test_pipeline.py | 2 +- src/spikeinterface/sortingcomponents/tools.py | 10 +- 5 files changed, 138 insertions(+), 116 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index 109ade1531..d2b9b90723 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -877,12 +877,12 @@ def save_to_folder( Legacy method. The 'new' way is : - * recording.save(format='binray', folder=...) + * recording.save(format='binary', folder=...) * sorting.save(format='numpy_folder', folder=...) """ warnings.warn( - "save_to_folder() should be recording.save(format='binray') " + "save_to_folder() should be recording.save(format='binary') " "or sorting.save(format='numpy_folder') " "This ambiguous method should not be used anymore!!", FutureWarning, @@ -890,51 +890,51 @@ def save_to_folder( # we keep the default format for recording and sorting like in old version return self.save(folder=folder, verbose=verbose, **save_kwargs) - if folder is None: - cache_folder = get_global_tmp_folder() - if name is None: - name = "".join(random.choices(string.ascii_uppercase + string.digits, k=8)) - folder = cache_folder / name - if verbose: - print(f"Use cache_folder={folder}") - else: - folder = cache_folder / name - if not is_set_global_tmp_folder(): - if verbose: - print(f"Use cache_folder={folder}") - else: - folder = Path(folder) - if overwrite and folder.is_dir(): - import shutil - - shutil.rmtree(folder) - - assert not folder.exists(), f"folder {folder} already exists, choose another name or use overwrite=True" - folder.mkdir(parents=True, exist_ok=False) - - # dump provenance - provenance_file_path = folder / f"provenance.json" - if self.check_serializability("json"): - self.dump_to_json(file_path=provenance_file_path, relative_to=folder) - elif self.check_serializability("pickle"): - provenance_file = folder / f"provenance.pkl" - self.dump_to_pickle(provenance_file, relative_to=folder) - else: - warnings.warn("The extractor is not serializable to file. The provenance will not be saved.") - - # save data (done the subclass) - self.save_metadata_to_folder(folder) - cached = self._save(folder=folder, verbose=verbose, **save_kwargs) - cached.load_metadata_from_folder(folder) - - # copy properties/ - self.copy_metadata(cached) + # if folder is None: + # cache_folder = get_global_tmp_folder() + # if name is None: + # name = "".join(random.choices(string.ascii_uppercase + string.digits, k=8)) + # folder = cache_folder / name + # if verbose: + # print(f"Use cache_folder={folder}") + # else: + # folder = cache_folder / name + # if not is_set_global_tmp_folder(): + # if verbose: + # print(f"Use cache_folder={folder}") + # else: + # folder = Path(folder) + # if overwrite and folder.is_dir(): + # import shutil + + # shutil.rmtree(folder) + + # assert not folder.exists(), f"folder {folder} already exists, choose another name or use overwrite=True" + # folder.mkdir(parents=True, exist_ok=False) + + # # dump provenance + # provenance_file_path = folder / f"provenance.json" + # if self.check_serializability("json"): + # self.dump_to_json(file_path=provenance_file_path, relative_to=folder) + # elif self.check_serializability("pickle"): + # provenance_file = folder / f"provenance.pkl" + # self.dump_to_pickle(provenance_file, relative_to=folder) + # else: + # warnings.warn("The extractor is not serializable to file. The provenance will not be saved.") + + # # save data (done the subclass) + # self.save_metadata_to_folder(folder) + # cached = self._save(folder=folder, verbose=verbose, **save_kwargs) + # cached.load_metadata_from_folder(folder) + + # # copy properties/ + # self.copy_metadata(cached) - # Dump the extractor to json file - si_folder_path = folder / f"si_folder.json" - cached.dump_to_json(file_path=si_folder_path, relative_to=folder) + # # Dump the extractor to json file + # si_folder_path = folder / f"si_folder.json" + # cached.dump_to_json(file_path=si_folder_path, relative_to=folder) - return cached + # return cached def save_to_zarr( self, @@ -947,74 +947,91 @@ def save_to_zarr( **save_kwargs, ): """ - Save extractor to zarr. - - Parameters - ---------- - name: str or None, default: None - Name of the subfolder in get_global_tmp_folder() - If "name" is given, "folder" must be None - folder: str, Path, or None, default: None - The folder used to save the zarr output. If the folder does not have a ".zarr" suffix, - it will be automatically appended - overwrite: bool, default: False - If True, the folder is removed if it already exists - storage_options: dict or None, default: None - Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc. - For cloud storage locations, this should not be None (in case of default values, use an empty dict) - channel_chunk_size: int or None, default: None - Channels per chunk (only for BaseRecording) - compressor: numcodecs.Codec or None, default: None - Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used - filters: list[numcodecs.Codec] or None, default: None - Global filters for zarr (global) - compressor_by_dataset: dict or None, default: None - Optional compressor per dataset: - - traces - - times - If None, the global compressor is used - filters_by_dataset: dict or None, default: None - Optional filters per dataset: - - traces - - times - If None, the global filters are used - verbose: bool, default: True - If True, the output is verbose - auto_cast_uint: bool, default: True - If True, unsigned integers are cast to signed integers to avoid issues with zarr (only for BaseRecording) + Legacy method. - Returns - ------- - cached: ZarrExtractor - Saved copy of the extractor. + The 'new' way is : + * recording.save(format='zarr', folder=...) + * sorting.save(format='zarr', folder=...) """ - from .zarrextractors import read_zarr - save_kwargs.pop("format", None) + warnings.warn( + "save_to_zarr() should be recording.save(format='zarr') " + "or sorting.save(format='zarr') " + "This ambiguous method should not be used anymore!!", + FutureWarning, + ) + # we keep the default format for recording and sorting like in old version + return self.save(folder=folder, format="zarr", verbose=verbose, **save_kwargs) + + # """ + # Save extractor to zarr. + + # Parameters + # ---------- + # name: str or None, default: None + # Name of the subfolder in get_global_tmp_folder() + # If "name" is given, "folder" must be None + # folder: str, Path, or None, default: None + # The folder used to save the zarr output. If the folder does not have a ".zarr" suffix, + # it will be automatically appended + # overwrite: bool, default: False + # If True, the folder is removed if it already exists + # storage_options: dict or None, default: None + # Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc. + # For cloud storage locations, this should not be None (in case of default values, use an empty dict) + # channel_chunk_size: int or None, default: None + # Channels per chunk (only for BaseRecording) + # compressor: numcodecs.Codec or None, default: None + # Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used + # filters: list[numcodecs.Codec] or None, default: None + # Global filters for zarr (global) + # compressor_by_dataset: dict or None, default: None + # Optional compressor per dataset: + # - traces + # - times + # If None, the global compressor is used + # filters_by_dataset: dict or None, default: None + # Optional filters per dataset: + # - traces + # - times + # If None, the global filters are used + # verbose: bool, default: True + # If True, the output is verbose + # auto_cast_uint: bool, default: True + # If True, unsigned integers are cast to signed integers to avoid issues with zarr (only for BaseRecording) + + # Returns + # ------- + # cached: ZarrExtractor + # Saved copy of the extractor. + # """ + # from .zarrextractors import read_zarr + + # save_kwargs.pop("format", None) + + # if folder is None: + # cache_folder = get_global_tmp_folder() + # if name is None: + # name = "".join(random.choices(string.ascii_uppercase + string.digits, k=8)) + # zarr_path = (cache_folder / name).with_suffix(".zarr") + # if verbose: + # print(f"Saving to zarr_path={zarr_path}") + # else: + # if storage_options is None: # save locally (not cloud storage) + # folder = clean_zarr_folder_name(folder) + # if folder.is_dir() and overwrite: + # shutil.rmtree(folder) + # zarr_path = folder + + # if not is_path_remote(zarr_path): + # assert not zarr_path.exists(), f"Path {zarr_path} already exists, choose another name" + # save_kwargs["zarr_path"] = zarr_path + # save_kwargs["storage_options"] = storage_options + # save_kwargs["channel_chunk_size"] = channel_chunk_size + # cached = self._save(format="zarr", verbose=verbose, **save_kwargs) + # cached = read_zarr(zarr_path) - if folder is None: - cache_folder = get_global_tmp_folder() - if name is None: - name = "".join(random.choices(string.ascii_uppercase + string.digits, k=8)) - zarr_path = (cache_folder / name).with_suffix(".zarr") - if verbose: - print(f"Saving to zarr_path={zarr_path}") - else: - if storage_options is None: # save locally (not cloud storage) - folder = clean_zarr_folder_name(folder) - if folder.is_dir() and overwrite: - shutil.rmtree(folder) - zarr_path = folder - - if not is_path_remote(zarr_path): - assert not zarr_path.exists(), f"Path {zarr_path} already exists, choose another name" - save_kwargs["zarr_path"] = zarr_path - save_kwargs["storage_options"] = storage_options - save_kwargs["channel_chunk_size"] = channel_chunk_size - cached = self._save(format="zarr", verbose=verbose, **save_kwargs) - cached = read_zarr(zarr_path) - - return cached + # return cached def _load_extractor_from_dict(dic) -> "BaseExtractor": diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 0aa3ea8820..3e3f93b947 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -319,8 +319,13 @@ def get_shape(self, segment_index: int | None = None) -> tuple[int, ...]: return (self.get_num_samples(segment_index=segment_index), self.get_num_channels()) def save(self, format="binary", verbose: bool = False, **save_kwargs): + """ + TODO: each object.save should have extensive docstring with all the options and examples + """ kwargs, job_kwargs = split_job_kwargs(save_kwargs) + # TODO: add overwrite option to binary/zarr save + if format == "binary": if "folder" not in kwargs: raise ValueError("Missing folder in recording.save(folder='...')") diff --git a/src/spikeinterface/core/tests/test_loading.py b/src/spikeinterface/core/tests/test_loading.py index 979113a2aa..257895fcc2 100644 --- a/src/spikeinterface/core/tests/test_loading.py +++ b/src/spikeinterface/core/tests/test_loading.py @@ -196,7 +196,7 @@ def test_load_aggregate_recording_from_json(generate_recording_sorting, tmp_path aggregated_rec = aggregate_channels(list_of_recs) recording_path = tmp_path / "aggregated_recording" - aggregated_rec.save_to_folder(folder=recording_path) + aggregated_rec.save(folder=recording_path) loaded_rec = load(recording_path / "provenance.json", base_folder=recording_path) assert np.all(loaded_rec.get_property("group") == recording.get_property("group")) diff --git a/src/spikeinterface/preprocessing/tests/test_pipeline.py b/src/spikeinterface/preprocessing/tests/test_pipeline.py index 37376781b7..3133358fda 100644 --- a/src/spikeinterface/preprocessing/tests/test_pipeline.py +++ b/src/spikeinterface/preprocessing/tests/test_pipeline.py @@ -176,7 +176,7 @@ def test_loading_provenance(create_cache_folder): # when several run seed=2205, ) - pp_rec.save_to_folder(folder=cache_folder) + pp_rec.save(folder=cache_folder) loaded_pp_dict = get_preprocessing_dict_from_file(cache_folder / "provenance.pkl") diff --git a/src/spikeinterface/sortingcomponents/tools.py b/src/spikeinterface/sortingcomponents/tools.py index 9836faaaad..c7e4bdc8ac 100644 --- a/src/spikeinterface/sortingcomponents/tools.py +++ b/src/spikeinterface/sortingcomponents/tools.py @@ -412,7 +412,7 @@ def cache_preprocessing( if total_memory is None: mem_ok = _check_cache_memory(recording, memory_limit, total_memory) if mem_ok: - recording = recording.save_to_memory(format="memory", shared=True, **job_kwargs) + recording = recording.save(format="memory", sharedmem=True, **job_kwargs) else: import warnings @@ -421,11 +421,11 @@ def cache_preprocessing( elif mode == "folder": assert folder is not None, "cache_preprocessing(): folder must be given" - recording = recording.save_to_folder(folder=folder, **job_kwargs) + recording = recording.save(folder=folder, format="binary", **job_kwargs) cache_info["folder"] = folder elif mode == "zarr": assert folder is not None, "cache_preprocessing(): folder must be given" - recording = recording.save_to_zarr(folder=folder, **job_kwargs) + recording = recording.save(folder=folder, format="zarr", **job_kwargs) cache_info["folder"] = folder elif mode == "no-cache": recording = recording @@ -433,11 +433,11 @@ def cache_preprocessing( mem_ok = _check_cache_memory(recording, memory_limit, total_memory) if mem_ok: # first try memory first - recording = recording.save_to_memory(format="memory", shared=True, **job_kwargs) + recording = recording.save(format="memory", sharedmem=True, **job_kwargs) cache_info["mode"] = "memory" elif folder is not None: # then try folder - recording = recording.save_to_folder(folder=folder, **job_kwargs) + recording = recording.save(folder=folder, format="binary", **job_kwargs) cache_info["mode"] = "folder" cache_info["folder"] = folder else: From 67a8d2d88ae44702956a88e021e9bde746d4c00f Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 9 Sep 2026 11:58:55 +0200 Subject: [PATCH 10/18] refactor: cleanup save functions and add detailed docs --- src/spikeinterface/core/base.py | 56 +---- src/spikeinterface/core/baserecording.py | 100 +++++--- .../core/baserecordingsnippets.py | 25 -- src/spikeinterface/core/basesnippets.py | 27 ++- src/spikeinterface/core/basesorting.py | 94 ++++++-- src/spikeinterface/core/binaryfolder.py | 41 ++-- src/spikeinterface/core/sortingfolder.py | 32 +-- src/spikeinterface/core/time_series.py | 2 +- src/spikeinterface/core/zarrextractors.py | 214 +++++------------- 9 files changed, 260 insertions(+), 331 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index d2b9b90723..d33ca7f72e 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -855,12 +855,6 @@ def save_to_memory(self, sharedmem=True, **save_kwargs) -> "BaseExtractor": warnings.warn("save_to_memory() should be save(format='memory')", FutureWarning) return self.save(format="memory", sharedmem=sharedmem, **save_kwargs) - # save_kwargs.pop("format", None) - - # cached = self._save(format="memory", sharedmem=sharedmem, **save_kwargs) - # self.copy_metadata(cached) - # return cached - def save(self): # Need to be implemented in Recording and Sorting raise NotImplementedError() @@ -878,64 +872,18 @@ def save_to_folder( The 'new' way is : * recording.save(format='binary', folder=...) - * sorting.save(format='numpy_folder', folder=...) + * sorting.save(format='binary', folder=...) """ warnings.warn( "save_to_folder() should be recording.save(format='binary') " - "or sorting.save(format='numpy_folder') " + "or sorting.save(format='binary') " "This ambiguous method should not be used anymore!!", FutureWarning, ) # we keep the default format for recording and sorting like in old version return self.save(folder=folder, verbose=verbose, **save_kwargs) - # if folder is None: - # cache_folder = get_global_tmp_folder() - # if name is None: - # name = "".join(random.choices(string.ascii_uppercase + string.digits, k=8)) - # folder = cache_folder / name - # if verbose: - # print(f"Use cache_folder={folder}") - # else: - # folder = cache_folder / name - # if not is_set_global_tmp_folder(): - # if verbose: - # print(f"Use cache_folder={folder}") - # else: - # folder = Path(folder) - # if overwrite and folder.is_dir(): - # import shutil - - # shutil.rmtree(folder) - - # assert not folder.exists(), f"folder {folder} already exists, choose another name or use overwrite=True" - # folder.mkdir(parents=True, exist_ok=False) - - # # dump provenance - # provenance_file_path = folder / f"provenance.json" - # if self.check_serializability("json"): - # self.dump_to_json(file_path=provenance_file_path, relative_to=folder) - # elif self.check_serializability("pickle"): - # provenance_file = folder / f"provenance.pkl" - # self.dump_to_pickle(provenance_file, relative_to=folder) - # else: - # warnings.warn("The extractor is not serializable to file. The provenance will not be saved.") - - # # save data (done the subclass) - # self.save_metadata_to_folder(folder) - # cached = self._save(folder=folder, verbose=verbose, **save_kwargs) - # cached.load_metadata_from_folder(folder) - - # # copy properties/ - # self.copy_metadata(cached) - - # # Dump the extractor to json file - # si_folder_path = folder / f"si_folder.json" - # cached.dump_to_json(file_path=si_folder_path, relative_to=folder) - - # return cached - def save_to_zarr( self, name=None, diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 3e3f93b947..81375757ad 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -320,20 +320,79 @@ def get_shape(self, segment_index: int | None = None) -> tuple[int, ...]: def save(self, format="binary", verbose: bool = False, **save_kwargs): """ - TODO: each object.save should have extensive docstring with all the options and examples + Save a `BaseRecording` object to a specified format: + + * "binary" + * "zarr" + * "memory" + + Parameters + ---------- + format : str, default: "binary" + The format to save the recording in. Options are: + - "binary": Saves the recording in binary format. + - "zarr": Saves the recording in Zarr format. + - "memory": Saves the recording in memory (shared memory or numpy array). + verbose : bool, default: False + If True, prints additional information during the save process. + **save_kwargs : dict + Additional keyword arguments specific to the chosen format. + All formats support job_kwargs for parallel processing + (see `si.get_global_job_kwargs()` for default values). + + * "binary" format: + - folder : str or Path + The folder where the binary files will be saved. + - overwrite : bool, default: False + If True, existing files in the folder will be overwritten. + - dtype : str, optional + The data type to use for saving the recording. If not provided, the recording's dtype + will be used. + * "zarr" format: + - folder : str or Path + The folder where the Zarr files will be saved. + - overwrite: bool, default: False + If True, the folder is removed if it already exists + - storage_options: dict or None, default: None + Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc. + For cloud storage locations, this should not be None (in case of default values, use an empty dict) + - channel_chunk_size: int or None, default: None + Channels per chunk (only for BaseRecording) + - compressor: numcodecs.Codec or None, default: None + Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used + - filters: list[numcodecs.Codec] or None, default: None + Global filters for zarr (global) + - compressor_by_dataset: dict or None, default: None + Optional compressor per dataset: + - traces + - times + If None, the global compressor is used + - filters_by_dataset: dict or None, default: None + Optional filters per dataset: + - traces + - times + If None, the global filters are used + * "memory" format: + - sharedmem : bool, default: True + If True, the recording is saved in shared memory. If False, it is saved as + a numpy array in memory. + + Returns + ------- + BaseRecording + The saved recording object in the specified format. """ kwargs, job_kwargs = split_job_kwargs(save_kwargs) - # TODO: add overwrite option to binary/zarr save - if format == "binary": if "folder" not in kwargs: raise ValueError("Missing folder in recording.save(folder='...')") from .binaryfolder import BinaryFolderRecording + folder = kwargs.pop("folder") cached = BinaryFolderRecording.write_recording( - self, folder=kwargs["folder"], dtype=kwargs.get("dtype", None), **job_kwargs + self, folder_path=folder, verbose=verbose, **kwargs, **job_kwargs ) elif format == "memory": @@ -348,49 +407,22 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): cached = NumpyRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs) - # self.copy_metadata(cached) - - # # timestamps are not saved in memory, so we have to set them explicitly - # for segment_index in range(self.get_num_segments()): - # if self.has_time_vector(segment_index): - # # the use of get_times is preferred since timestamps are converted to array - # time_vector = self.get_times(segment_index=segment_index) - # cached.set_times(time_vector, segment_index=segment_index) - elif format == "zarr": if "folder" not in kwargs: raise ValueError("Missing folder in recording.save(folder='...')") + folder_path = kwargs.pop("folder") from .zarrextractors import ZarrRecordingExtractor - folder_path = kwargs["folder"] - if isinstance(folder_path, Path) and folder_path.suffix != "zarr": - # automatically add the zarr suffix - folder_path = folder_path.with_suffix(".zarr") - - storage_options = kwargs.pop("storage_options", None) - ZarrRecordingExtractor.write_recording( - self, folder_path, storage_options, verbose=verbose, **kwargs, **job_kwargs + cached = ZarrRecordingExtractor.write_recording( + self, folder_path=folder_path, verbose=verbose, **kwargs, **job_kwargs ) - cached = ZarrRecordingExtractor(folder_path, storage_options) - # timestamps are saved and restored in zarr, so no need to set them explicitly else: raise ValueError(f"format {format} not supported") return cached - def _extra_metadata_from_folder(self, folder): - # load probe - super()._extra_metadata_from_folder(folder) - - # load time vector if any - for segment_index, rs in enumerate(self.segments): - time_file = folder / f"times_cached_seg{segment_index}.npy" - if time_file.is_file(): - time_vector = np.load(time_file, mmap_mode="r") - rs._time_vector = time_vector - def select_channels(self, channel_ids: list | np.ndarray | tuple) -> "BaseRecording": """ Returns a new recording object with a subset of channels. diff --git a/src/spikeinterface/core/baserecordingsnippets.py b/src/spikeinterface/core/baserecordingsnippets.py index 6c679e0da2..11eb909b59 100644 --- a/src/spikeinterface/core/baserecordingsnippets.py +++ b/src/spikeinterface/core/baserecordingsnippets.py @@ -320,31 +320,6 @@ def _extra_metadata_copy(self, other): if self._probegroup is not None: other._probegroup = self._probegroup.copy() - def _extra_metadata_from_folder(self, folder): - # load probe from folder - # Note: we don't need any fix for legacy probegroups, since the - # set_probegroup() method will handle the device_channel_indices - # sorting and global contact order - folder = Path(folder) - probe_file = folder / "probegroup.json" - legacy_probe_file = folder / "probe.json" - if probe_file.is_file(): - probegroup = read_probeinterface(probe_file) - self.set_probegroup(probegroup) - elif legacy_probe_file.is_file(): - probegroup = read_probeinterface(legacy_probe_file) - self.set_probegroup(probegroup) - - # remove "contact_vector" property if present as it is not needed anymore - if "contact_vector" in self.get_property_keys(): - self.delete_property("contact_vector") - - # def _extra_metadata_to_folder(self, folder): - # # save probe - # if self.has_probe(): - # probegroup = self.get_probegroup() - # write_probeinterface(folder / "probegroup.json", probegroup) - def _extra_metadata_from_dict(self, dump_dict): # load probe and handle backward-compatibility with legacy "contact_vector"/"location" property if "probegroup" in dump_dict: diff --git a/src/spikeinterface/core/basesnippets.py b/src/spikeinterface/core/basesnippets.py index df8fc0cbaf..5b38c47791 100644 --- a/src/spikeinterface/core/basesnippets.py +++ b/src/spikeinterface/core/basesnippets.py @@ -1,5 +1,3 @@ -from .base import BaseSegment -from .baserecordingsnippets import BaseRecordingSnippets import numpy as np from warnings import warn @@ -7,9 +5,8 @@ from pathlib import Path -from .core_tools import save_properties_to_binary_folder - -# snippets segments? +from .base import BaseSegment +from .baserecordingsnippets import BaseRecordingSnippets class BaseSnippets(BaseRecordingSnippets): @@ -21,6 +18,12 @@ class BaseSnippets(BaseRecordingSnippets): _main_features = [] def __init__(self, sampling_frequency: float, nbefore: int | None, snippet_len: int, channel_ids: list, dtype): + warn( + "`BaseSnippets` is deprecated and will be removed in version 0.106.0." + "Only continuous recordings with `BaseRecording` will be supported.", + FutureWarning, + stacklevel=2, + ) BaseRecordingSnippets.__init__( self, channel_ids=channel_ids, sampling_frequency=sampling_frequency, dtype=dtype ) @@ -213,9 +216,11 @@ def _select_segments(self, segment_indices): def save(self, format="npy", **save_kwargs): """ - At the moment only "npy" and "memory" avaiable: - """ + Save a `BaseSnippets` object to a specified format: + * "npy" + * "memory" + """ if format == "npy": from spikeinterface.core.npyfoldersnippets import NpyFolderSnippets @@ -240,14 +245,12 @@ def save(self, format="npy", **save_kwargs): nbefore=self.nbefore, channel_ids=self.channel_ids, ) - + if self.has_probe(): + probegroup = self.get_probegroup() + cached.set_probegroup(probegroup) else: raise ValueError(f"format {format} not supported") - if self.has_probe(): - probegroup = self.get_probegroup() - cached.set_probegroup(probegroup) - return cached def get_times(self): diff --git a/src/spikeinterface/core/basesorting.py b/src/spikeinterface/core/basesorting.py index 9fdd256bd5..3414c81987 100644 --- a/src/spikeinterface/core/basesorting.py +++ b/src/spikeinterface/core/basesorting.py @@ -502,45 +502,77 @@ def get_times( else: return None - def save(self, format="numpy_folder", **save_kwargs): + def save(self, format="binary", **save_kwargs): """ - Save a sorting object to disk in a specified format. + Save a `BaseSorting` object to a specified format: - Note - ---- - This function replaces the old CacheSortingExtractor, but enables more engines + * "binary" - old "numpy_folder" + * "zarr" + * "memory" + * "npz_folder" - deprecated - for caching a results. + Parameters + ---------- + format : str, default: "binary" + The format to save the sorting in. Options are: + - "binary"/"numpy_folder": Saves the sorting in a binary numpy folder format. + - "zarr": Saves the sorting in Zarr format. + - "memory": Saves the sorting in memory (shared memory or numpy array). + - "npz_folder": Saves the sorting in a deprecated npz folder format. + verbose : bool, default: False + If True, prints additional information during the save process. + **save_kwargs : dict + Additional keyword arguments specific to the chosen format. + + * "binary" format: + - folder : str or Path + The folder where the binary files will be saved. + - overwrite : bool, default: False + If True, existing files in the folder will be overwritten. + * "zarr" format: + - folder : str or Path + The folder where the Zarr files will be saved. + - overwrite: bool, default: False + If True, the folder is removed if it already exists + - storage_options: dict or None, default: None + Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc. + For cloud storage locations, this should not be None (in case of default values, use an empty dict) + - compressor: numcodecs.Codec or None, default: None + Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used + * "memory" format: + - sharedmem : bool, default: True + If True, the recording is saved in shared memory. If False, it is saved as + a numpy array in memory. - Since v0.98.0 "numpy_folder" is used by defult. - From v0.96.0 to 0.97.0 "npz_folder" was the default. + Returns + ------- + BaseSorting + The saved sorting object in the specified format. """ if format == "numpy_folder": + warnings.warn( + "The 'numpy_folder' is renamed to 'binary' and will be removed in 0.106.0. " + "Please use 'binary' instead.", + FutureWarning, + stacklevel=2, + ) + format = "binary" + + if format == "binary": from .sortingfolder import NumpyFolderSorting + if "folder" not in save_kwargs: + raise ValueError("For 'binary' format, 'folder' must be specified in save_kwargs.") folder = save_kwargs.pop("folder") - NumpyFolderSorting.write_sorting(self, folder) - cached = NumpyFolderSorting(folder) + cached = NumpyFolderSorting.write_sorting(self, folder_path=folder, **save_kwargs) elif format == "zarr": from .zarrextractors import ZarrSortingExtractor - folder_path = save_kwargs.pop("folder") - if isinstance(folder_path, Path) and folder_path.suffix != "zarr": - # automatically add the zarr suffix - folder_path = folder_path.with_suffix(".zarr") - - storage_options = save_kwargs.pop("storage_options", None) - ZarrSortingExtractor.write_sorting(self, folder_path, storage_options, **save_kwargs) - cached = ZarrSortingExtractor(folder_path, storage_options) - - elif format == "npz_folder": - from .sortingfolder import NpzFolderSorting - + if "folder" not in save_kwargs: + raise ValueError("For 'zarr' format, 'folder' must be specified in save_kwargs.") folder = save_kwargs.pop("folder") - NpzFolderSorting.write_sorting(self, folder) - cached = NpzFolderSorting(folder_path=folder) - + cached = ZarrSortingExtractor.write_sorting(self, folder, **save_kwargs) elif format == "memory": if save_kwargs.get("sharedmem", True): from .numpyextractors import SharedMemorySorting @@ -550,6 +582,18 @@ def save(self, format="numpy_folder", **save_kwargs): from .numpyextractors import NumpySorting cached = NumpySorting.from_sorting(self, with_metadata=True) + elif format == "npz_folder": + warnings.warn( + "The 'npz_folder' format is deprecated and will be removed in 0.106.0. " + "Please use 'numpy_folder' instead.", + FutureWarning, + stacklevel=True, + ) + from .sortingfolder import NpzFolderSorting + + folder = save_kwargs.pop("folder") + NpzFolderSorting.write_sorting(self, folder) + cached = NpzFolderSorting(folder_path=folder) else: raise ValueError(f"Format {format} not supported") diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index a462ad806f..464959e467 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -2,6 +2,7 @@ import json from copy import deepcopy +import shutil import numpy as np @@ -117,29 +118,41 @@ def get_binary_description(self): return d @staticmethod - def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, verbose=False, **job_kwargs): + def write_recording( + recording: BaseRecording, + folder_path: str | Path, + verbose: bool = False, + overwrite: bool = False, + dtype=None, + **job_kwargs, + ): from .time_series_tools import write_binary from .binaryrecordingextractor import BinaryRecordingExtractor from .binaryfolder import BinaryFolderRecording - folder = Path(folder) - folder.mkdir(exist_ok=False) - - file_paths = [folder / f"traces_cached_seg{i}.raw" for i in range(recording.get_num_segments())] + folder_path = Path(folder_path) + if folder_path.is_dir(): + if not overwrite: + raise FileExistsError(f"Folder {folder_path} already exists. Use overwrite=True to overwrite it.") + else: + shutil.rmtree(folder_path) + folder_path.mkdir(exist_ok=False, parents=True) + + file_paths = [folder_path / f"traces_cached_seg{i}.raw" for i in range(recording.get_num_segments())] if dtype is None: dtype = recording.get_dtype() t_starts = recording._get_t_starts() write_binary(recording, file_paths=file_paths, dtype=dtype, verbose=verbose, **job_kwargs) - save_properties_to_binary_folder(folder / "properties", recording) - save_extractor_provenance(folder, recording) + save_properties_to_binary_folder(folder_path / "properties", recording) + save_extractor_provenance(folder_path, recording) # new in version 0.105.0, before that annotations were handle by "si_folder.json" file - save_annotations_to_folder(folder, recording) + save_annotations_to_folder(folder_path, recording) if recording.has_probe(): probegroup = recording.get_probegroup() - write_probeinterface(folder / "probegroup.json", probegroup) + write_probeinterface(folder_path / "probegroup.json", probegroup) # This is created so it can be saved as json because the `BinaryFolderRecording` requires it loading # See the __init__ @@ -157,7 +170,7 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve gain_to_uV=recording.get_channel_gains(), offset_to_uV=recording.get_channel_offsets(), ) - binary_rec.dump(folder / "binary.json", relative_to=folder) + binary_rec.dump(folder_path / "binary.json", relative_to=folder_path) # TODO alessio : remove this, it is needed to pass tests # save times @@ -165,12 +178,12 @@ def write_recording(recording: BaseRecording, folder: str | Path, dtype=None, ve d = rs.get_times_kwargs() time_vector = d["time_vector"] if time_vector is not None: - np.save(folder / f"times_cached_seg{segment_index}.npy", time_vector) + np.save(folder_path / f"times_cached_seg{segment_index}.npy", time_vector) # make the si_folder file to make the load() easier until version 0.105.0 - cached = BinaryFolderRecording(folder_path=folder) - si_folder_path = folder / f"si_folder.json" - cached.dump_to_json(file_path=si_folder_path, relative_to=folder, include_extra_metadata=False) + cached = BinaryFolderRecording(folder_path=folder_path) + si_folder_path = folder_path / f"si_folder.json" + cached.dump_to_json(file_path=si_folder_path, relative_to=folder_path, include_extra_metadata=False) return cached diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index 5d1a587418..f9fc81361d 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -1,6 +1,7 @@ from pathlib import Path import json from copy import deepcopy +import shutil import warnings import numpy as np @@ -60,31 +61,34 @@ def __init__(self, folder_path, mmap_mode: str | None = None): self._kwargs = dict(folder_path=str(folder_path.absolute()), mmap_mode=mmap_mode) @staticmethod - def write_sorting(sorting, save_path): + def write_sorting(sorting, folder_path, overwrite: bool = False): # the folder can already exists but not contaning numpysorting_info.json - save_path = Path(save_path) - save_path.mkdir(parents=True, exist_ok=True) - - info_file = save_path / "numpysorting_info.json" - if info_file.exists(): - raise ValueError("NumpyFolderSorting.write_sorting the folder already contains numpysorting_info.json") + folder_path = Path(folder_path) + if folder_path.is_dir(): + if not overwrite: + raise ValueError("NumpyFolderSorting.write_sorting the folder already exists") + else: + shutil.rmtree(folder_path) + folder_path.mkdir(parents=True, exist_ok=True) + + info_file = folder_path / "numpysorting_info.json" d = { "sampling_frequency": float(sorting.get_sampling_frequency()), "unit_ids": sorting.unit_ids.tolist(), "num_segments": sorting.get_num_segments(), } info_file.write_text(json.dumps(d), encoding="utf8") - np.save(save_path / "spikes.npy", sorting.to_spike_vector()) + np.save(folder_path / "spikes.npy", sorting.to_spike_vector()) - save_properties_to_binary_folder(save_path / "properties", sorting) - save_extractor_provenance(save_path, sorting) + save_properties_to_binary_folder(folder_path / "properties", sorting) + save_extractor_provenance(folder_path, sorting) # new in version 0.105.0, before that annotations were handle by "si_folder.json" file - save_annotations_to_folder(save_path, sorting) + save_annotations_to_folder(folder_path, sorting) # make the si_folder file to make the load() easier until version 0.105.0 - cached = NumpyFolderSorting(folder_path=save_path) - si_folder_path = save_path / f"si_folder.json" - cached.dump_to_json(file_path=si_folder_path, relative_to=save_path, include_extra_metadata=False) + cached = NumpyFolderSorting(folder_path=folder_path) + si_folder_path = folder_path / f"si_folder.json" + cached.dump_to_json(file_path=si_folder_path, relative_to=folder_path, include_extra_metadata=False) return cached diff --git a/src/spikeinterface/core/time_series.py b/src/spikeinterface/core/time_series.py index 640a256d24..ec5bee656b 100644 --- a/src/spikeinterface/core/time_series.py +++ b/src/spikeinterface/core/time_series.py @@ -14,7 +14,7 @@ # The backing store depends on how the recording was created/loaded: # - np.ndarray : set_times() (writeable, in-memory) # - np.memmap : BinaryFolderRecording load via np.load(..., mmap_mode="r") -# -- *read-only* ; see BaseRecording._extra_metadata_from_folder +# -- *read-only* ; see BinaryFolderRecording.__init__ # - zarr.Array : ZarrRecordingExtractor load # -- *read-only* ; see ZarrRecordingExtractor.__init__ # Code reading `._time_vector` must not assume it is writeable (see `shift_times`). diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index 6cdc1c9fde..73de374c9d 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -1,3 +1,4 @@ +import shutil import warnings from pathlib import Path @@ -9,9 +10,9 @@ from .base import minimum_spike_dtype, _get_class_from_string from .baserecording import BaseRecording, BaseRecordingSegment from .basesorting import BaseSorting, SpikeVectorSortingSegment -from .core_tools import define_function_from_class, check_json, retrieve_importing_provenance from .job_tools import split_job_kwargs -from .core_tools import is_path_remote +from .core_tools import define_function_from_class, check_json, retrieve_importing_provenance, is_path_remote +from .time_series_tools import _write_time_series_to_zarr def super_zarr_open(folder_path: str | Path, mode: str = "r", storage_options: dict | None = None): @@ -236,11 +237,17 @@ def __init__( @staticmethod def write_recording( - recording: BaseRecording, folder_path: str | Path, storage_options: dict | None = None, **kwargs + recording: BaseRecording, + folder_path: str | Path, + overwrite: bool = False, + storage_options: dict | None = None, + **kwargs, ): + folder_path = create_zarr_path_for_write(folder_path, overwrite=overwrite) zarr_root = zarr.open(str(folder_path), mode="w", storage_options=storage_options) zarr_root.attrs["zarr_class_info"] = retrieve_importing_provenance(ZarrRecordingExtractor) add_recording_to_zarr_group(recording, zarr_root, **kwargs) + return ZarrRecordingExtractor(folder_path, storage_options=storage_options) class ZarrRecordingSegment(BaseRecordingSegment): @@ -474,10 +481,17 @@ def __init__( } @staticmethod - def write_sorting(sorting: BaseSorting, folder_path: str | Path, storage_options: dict | None = None, **kwargs): + def write_sorting( + sorting: BaseSorting, + folder_path: str | Path, + overwrite: bool = False, + storage_options: dict | None = None, + **kwargs, + ): """ Write a sorting extractor to zarr format. """ + folder_path = create_zarr_path_for_write(folder_path, overwrite=overwrite) zarr_root = zarr.open(str(folder_path), mode="w", storage_options=storage_options) zarr_root.attrs["zarr_class_info"] = retrieve_importing_provenance(ZarrSortingExtractor) add_sorting_to_zarr_group(sorting, zarr_root, **kwargs) @@ -531,7 +545,7 @@ def resolve_zarr_path(folder_path: str | Path): Parameters ---------- """ - if str(folder_path).startswith("s3:") or str(folder_path).startswith("gcs:"): + if is_path_remote(folder_path): # cloud location, no need to resolve return folder_path, folder_path else: @@ -540,6 +554,27 @@ def resolve_zarr_path(folder_path: str | Path): return folder_path, folder_path_kwarg +def create_zarr_path_for_write(folder_path: str | Path, overwrite: bool = False): + """ + Resolve a path to a zarr folder for writing. + + Parameters + ---------- + folder_path : str or Path + Path to the zarr root file + """ + if not is_path_remote(folder_path): + folder_path = Path(folder_path) + folder_path = folder_path.with_suffix(".zarr") + if folder_path.is_dir(): + if not overwrite: + raise FileExistsError(f"Folder {folder_path} already exists. Use overwrite=True to overwrite it.") + else: + shutil.rmtree(folder_path) + folder_path.mkdir(exist_ok=False, parents=True) + return folder_path + + def _write_object_array( group, name: str, @@ -667,7 +702,6 @@ def add_sorting_to_zarr_group(sorting: BaseSorting, zarr_group: zarr.Group, **kw add_properties_and_annotations(zarr_group, sorting) -# Recording def add_recording_to_zarr_group(recording: BaseRecording, zarr_group: zarr.Group, verbose=False, dtype=None, **kwargs): zarr_kwargs, job_kwargs = split_job_kwargs(kwargs) @@ -681,6 +715,14 @@ def add_recording_to_zarr_group(recording: BaseRecording, zarr_group: zarr.Group zarr_group.attrs["num_segments"] = int(recording.get_num_segments()) zarr_group.create_dataset(name="channel_ids", data=recording.get_channel_ids(), compressor=None) dataset_paths = [f"traces_seg{i}" for i in range(recording.get_num_segments())] + dataset_timestamps_paths: list | None = None + if any(recording.has_time_vector(i) for i in range(recording.get_num_segments())): + dataset_timestamps_paths = [] + for i in range(recording.get_num_segments()): + if recording.has_time_vector(i): + dataset_timestamps_paths.append(f"times_seg{i}") + else: + dataset_timestamps_paths.append(None) dtype = recording.get_dtype() if dtype is None else dtype channel_chunk_size = zarr_kwargs.get("channel_chunk_size", None) @@ -691,159 +733,27 @@ def add_recording_to_zarr_group(recording: BaseRecording, zarr_group: zarr.Group compressor_traces = compressor_by_dataset.get("traces", global_compressor) filters_traces = filters_by_dataset.get("traces", global_filters) - add_traces_to_zarr( - recording=recording, + compressor_times = compressor_by_dataset.get("times", global_compressor) + filters_times = filters_by_dataset.get("times", global_filters) + + _write_time_series_to_zarr( + time_series=recording, zarr_group=zarr_group, dataset_paths=dataset_paths, - compressor=compressor_traces, - filters=filters_traces, + dataset_timestamps_paths=dataset_timestamps_paths, + compressor_data=compressor_traces, + filters_data=filters_traces, dtype=dtype, - channel_chunk_size=channel_chunk_size, - verbose=verbose, + extra_chunks=(channel_chunk_size,), + compressor_times=compressor_times, + filters_times=filters_times, + verbose=False, **job_kwargs, ) - # save probe + # Save probegroup if recording.has_probe(): probegroup = recording.get_probegroup() zarr_group.attrs["probegroup"] = check_json(probegroup.to_dict(array_as_list=True)) - - # save time vector if any - t_starts = np.zeros(recording.get_num_segments(), dtype="float64") * np.nan - for segment_index, rs in enumerate(recording.segments): - d = rs.get_times_kwargs() - time_vector = d["time_vector"] - - compressor_times = compressor_by_dataset.get("times", global_compressor) - filters_times = filters_by_dataset.get("times", global_filters) - - if time_vector is not None: - _ = zarr_group.create_dataset( - name=f"times_seg{segment_index}", - data=time_vector, - filters=filters_times, - compressor=compressor_times, - ) - elif d["t_start"] is not None: - t_starts[segment_index] = d["t_start"] - - if np.any(~np.isnan(t_starts)): - zarr_group.create_dataset(name="t_starts", data=t_starts, compressor=None) - + # Add properties and annotations add_properties_and_annotations(zarr_group, recording) - - -def add_traces_to_zarr( - recording, - zarr_group, - dataset_paths, - channel_chunk_size=None, - dtype=None, - compressor=None, - filters=None, - verbose=False, - **job_kwargs, -): - """ - Save the trace of a recording extractor in several zarr format. - - Parameters - ---------- - recording : RecordingExtractor - The recording extractor object to be saved in .dat format - zarr_group : zarr.Group - The zarr group to add traces to - dataset_paths : list - List of paths to traces datasets in the zarr group - channel_chunk_size : int or None, default: None (chunking in time only) - Channels per chunk - dtype : dtype, default: None - Type of the saved data - compressor : zarr compressor or None, default: None - Zarr compressor - filters : list, default: None - List of zarr filters - verbose : bool, default: False - If True, output is verbose (when chunks are used) - {} - """ - from .job_tools import ( - ensure_chunk_size, - fix_job_kwargs, - TimeSeriesChunkExecutor, - ) - - assert dataset_paths is not None, "Provide 'file_path'" - - if not isinstance(dataset_paths, list): - dataset_paths = [dataset_paths] - assert len(dataset_paths) == recording.get_num_segments() - - if dtype is None: - dtype = recording.get_dtype() - - job_kwargs = fix_job_kwargs(job_kwargs) - chunk_size = ensure_chunk_size(recording, **job_kwargs) - - # create zarr datasets files - zarr_datasets = [] - for segment_index in range(recording.get_num_segments()): - num_frames = recording.get_num_samples(segment_index) - num_channels = recording.get_num_channels() - dset_name = dataset_paths[segment_index] - shape = (num_frames, num_channels) - dset = zarr_group.create_dataset( - name=dset_name, - shape=shape, - chunks=(chunk_size, channel_chunk_size), - dtype=dtype, - filters=filters, - compressor=compressor, - ) - zarr_datasets.append(dset) - # synchronizer=zarr.ThreadSynchronizer()) - - # use executor (loop or workers) - func = _write_zarr_chunk - init_func = _init_zarr_worker - init_args = (recording, zarr_datasets, dtype) - executor = TimeSeriesChunkExecutor( - recording, func, init_func, init_args, verbose=verbose, job_name="write_zarr_recording", **job_kwargs - ) - executor.run() - - -# used by write_zarr_recording + TimeSeriesChunkExecutor -def _init_zarr_worker(recording, zarr_datasets, dtype): - import zarr - - # create a local dict per worker - worker_ctx = {} - worker_ctx["recording"] = recording - worker_ctx["zarr_datasets"] = zarr_datasets - worker_ctx["dtype"] = np.dtype(dtype) - - return worker_ctx - - -# used by write_zarr_recording + TimeSeriesChunkExecutor -def _write_zarr_chunk(segment_index, start_frame, end_frame, worker_ctx): - import gc - - # recover variables of the worker - recording = worker_ctx["recording"] - dtype = worker_ctx["dtype"] - zarr_dataset = worker_ctx["zarr_datasets"][segment_index] - - # apply function - traces = recording.get_traces( - start_frame=start_frame, - end_frame=end_frame, - segment_index=segment_index, - ) - traces = traces.astype(dtype) - zarr_dataset[start_frame:end_frame, :] = traces - - # fix memory leak by forcing garbage collection - del traces - gc.collect() From cbb8156dc568faebe7794cf4f7db5674b2ae4b3d Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 9 Sep 2026 12:31:24 +0200 Subject: [PATCH 11/18] fix: zarr loading and rename numpy_folder -> binary --- src/spikeinterface/core/loading.py | 12 +++--------- src/spikeinterface/core/tests/test_basesorting.py | 4 ++-- src/spikeinterface/core/tests/test_loading.py | 2 +- src/spikeinterface/core/tests/test_zarrextractors.py | 4 ++-- src/spikeinterface/core/zarrextractors.py | 1 + 5 files changed, 9 insertions(+), 14 deletions(-) diff --git a/src/spikeinterface/core/loading.py b/src/spikeinterface/core/loading.py index 40e3530be7..4f92ca96bc 100644 --- a/src/spikeinterface/core/loading.py +++ b/src/spikeinterface/core/loading.py @@ -292,23 +292,17 @@ def _load_object_from_zarr(folder_or_url, object_type, **kwargs): elif object_type == "Recording": from .zarrextractors import read_zarr_recording - storage_options = kwargs.get("storage_options", None) - load_compression_ratio = kwargs.get("load_compression_ratio", False) - recording = read_zarr_recording( - folder_or_url, storage_options=storage_options, load_compression_ratio=load_compression_ratio - ) + recording = read_zarr_recording(folder_or_url, **kwargs) return recording elif object_type == "Sorting": from .zarrextractors import read_zarr_sorting - storage_options = kwargs.get("storage_options", None) - sorting = read_zarr_sorting(folder_or_url, storage_options=storage_options) + sorting = read_zarr_sorting(folder_or_url, **kwargs) return sorting elif object_type == "Recording|Sorting": # This case shoudl deprecated soon because the read_zarr is ultra ambiguous # just testing if the zarr contains unit_ids or channel_ids but many object also contains it (see template)!!!! from .zarrextractors import read_zarr - storage_options = kwargs.get("storage_options", None) - rec_or_sorting = read_zarr(folder_or_url, storage_options=storage_options) + rec_or_sorting = read_zarr(folder_or_url, **kwargs) return rec_or_sorting diff --git a/src/spikeinterface/core/tests/test_basesorting.py b/src/spikeinterface/core/tests/test_basesorting.py index 9def88dba6..75687c54e6 100644 --- a/src/spikeinterface/core/tests/test_basesorting.py +++ b/src/spikeinterface/core/tests/test_basesorting.py @@ -78,7 +78,7 @@ def test_BaseSorting(create_cache_folder): # cache new format : numpy_folder folder = cache_folder / "simple_sorting_numpy_folder" sorting.set_property("test", np.ones(len(sorting.unit_ids))) - sorting.save(folder=folder, format="numpy_folder") + sorting.save(folder=folder, format="binary") sorting2 = BaseExtractor.load(folder) assert isinstance(sorting2, NumpyFolderSorting) @@ -143,7 +143,7 @@ def test_BaseSorting(create_cache_folder): # test save to zarr # compressor = get_default_zarr_compressor() - sorting_zarr = sorting.save(format="zarr", folder=cache_folder / "sorting") + sorting_zarr = sorting.save(format="zarr", folder=cache_folder / "sorting.zarr") sorting_zarr_loaded = load(cache_folder / "sorting.zarr") # annotations is False because Zarr adds compression ratios check_sortings_equal(sorting, sorting_zarr, check_annotations=False, check_properties=True) diff --git a/src/spikeinterface/core/tests/test_loading.py b/src/spikeinterface/core/tests/test_loading.py index 257895fcc2..9f1e5c4578 100644 --- a/src/spikeinterface/core/tests/test_loading.py +++ b/src/spikeinterface/core/tests/test_loading.py @@ -102,7 +102,7 @@ def test_load_binary_recording(generate_recording_sorting, tmp_path, output_form check_recordings_equal(rec, rec_loaded) -@pytest.mark.parametrize("output_format", ["numpy_folder", "zarr"]) +@pytest.mark.parametrize("output_format", ["binary", "zarr"]) def test_load_binary_sorting(generate_recording_sorting, tmp_path, output_format): _, sort = generate_recording_sorting _ = sort.save(folder=tmp_path / "test_sorting", format=output_format, overwrite=True) diff --git a/src/spikeinterface/core/tests/test_zarrextractors.py b/src/spikeinterface/core/tests/test_zarrextractors.py index cc0c60721e..7e58898a6e 100644 --- a/src/spikeinterface/core/tests/test_zarrextractors.py +++ b/src/spikeinterface/core/tests/test_zarrextractors.py @@ -60,13 +60,13 @@ def test_ZarrSortingExtractor(tmp_path): np_sorting = generate_sorting() # store in root standard normal way - folder = tmp_path / "zarr_sorting" + folder = tmp_path / "zarr_sorting.zarr" ZarrSortingExtractor.write_sorting(np_sorting, folder) sorting = ZarrSortingExtractor(folder) sorting = load(sorting.to_dict()) # store the sorting in a sub group (for instance SortingResult) - folder = tmp_path / "zarr_sorting_sub_group" + folder = tmp_path / "zarr_sorting_sub_group.zarr" zarr_root = zarr.open(folder, mode="w") zarr_sorting_group = zarr_root.create_group("sorting") add_sorting_to_zarr_group(sorting, zarr_sorting_group) diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index 73de374c9d..4846172df3 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -495,6 +495,7 @@ def write_sorting( zarr_root = zarr.open(str(folder_path), mode="w", storage_options=storage_options) zarr_root.attrs["zarr_class_info"] = retrieve_importing_provenance(ZarrSortingExtractor) add_sorting_to_zarr_group(sorting, zarr_root, **kwargs) + return ZarrSortingExtractor(folder_path, storage_options=storage_options) read_zarr_recording = define_function_from_class(source_class=ZarrRecordingExtractor, name="read_zarr_recording") From 4dc62cb23895c02d7b4af9d3407ff87d2db8a4dc Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 9 Sep 2026 13:20:12 +0200 Subject: [PATCH 12/18] fix: remove folder_metadata, add time_kwargs to dump_dict if time_info_modified --- .github/scripts/serialization/objects.py | 50 ++++++++++--- .../serialization/serialize_objects.py | 4 ++ .../workflows/cross_version_serialization.yml | 22 ++---- .../benchmark/benchmark_base.py | 4 +- src/spikeinterface/core/base.py | 51 +++++--------- src/spikeinterface/core/baserecording.py | 22 ++++++ src/spikeinterface/core/binaryfolder.py | 42 +++++++---- .../core/binaryrecordingextractor.py | 68 ++++++++++++++---- src/spikeinterface/core/core_tools.py | 15 ++-- src/spikeinterface/core/loading.py | 2 +- src/spikeinterface/core/numpyextractors.py | 4 +- src/spikeinterface/core/sortingfolder.py | 14 +++- .../core/tests/test_baserecording.py | 70 ++++++++++++++++--- .../core/tests/test_basesnippets.py | 13 ++-- .../core/tests/test_basesorting.py | 14 ++-- src/spikeinterface/core/time_series.py | 33 ++++++--- src/spikeinterface/core/time_series_tools.py | 61 +++++++++++----- .../preprocessing/basepreprocessor.py | 23 ++++++ 18 files changed, 360 insertions(+), 152 deletions(-) diff --git a/.github/scripts/serialization/objects.py b/.github/scripts/serialization/objects.py index ce3d215260..78fb3aed0b 100644 --- a/.github/scripts/serialization/objects.py +++ b/.github/scripts/serialization/objects.py @@ -29,10 +29,13 @@ "json": ".json", "pickle": ".pkl", "binary": "_binary", + "binary_parallel": "_binary_parallel", "numpy_folder": "_numpy_folder", "zarr": ".zarr", + "zarr_parallel": "_parallel.zarr", } +DEFAULT_DURATION = 2.0 # seconds, for recordings and sortings # --- json (recipe) entries: moved class + a second class ------------------------- @@ -46,7 +49,12 @@ def _build_noise_generator_recording(): except ImportError: from spikeinterface.core.generate import NoiseGeneratorRecording - return NoiseGeneratorRecording(num_channels=4, sampling_frequency=30000.0, durations=[1.0, 1.5], seed=0) + return NoiseGeneratorRecording( + num_channels=4, + sampling_frequency=30000.0, + durations=[DEFAULT_DURATION, DEFAULT_DURATION + 0.5], + seed=0, + ) def _check_noise_generator_recording(rec): @@ -59,7 +67,7 @@ def _check_noise_generator_recording(rec): def _build_mock_recording(): from spikeinterface.core import generate_recording - return generate_recording(num_channels=4, durations=[1.0], sampling_frequency=30000.0, seed=0) + return generate_recording(num_channels=4, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) def _check_mock_recording(rec): @@ -77,7 +85,7 @@ def _build_recording_with_properties(): import numpy as np from spikeinterface.core import generate_recording - rec = generate_recording(num_channels=4, durations=[1.0], sampling_frequency=30000.0, seed=0) + rec = generate_recording(num_channels=4, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) rec.set_property("quality", np.array(["good", "good", "bad", "good"])) rec.annotate(experimenter="test") return rec @@ -94,7 +102,7 @@ def _build_recording_with_probe(): from probeinterface import generate_linear_probe from spikeinterface.core import generate_recording - rec = generate_recording(num_channels=8, durations=[1.0], sampling_frequency=30000.0, seed=0) + rec = generate_recording(num_channels=8, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) probe = generate_linear_probe(num_elec=8) probe.set_device_channel_indices(np.arange(8)) if parse(si_version) <= parse("0.105.0"): @@ -104,6 +112,20 @@ def _build_recording_with_probe(): return rec_with_probe +def _build_recording_with_timestamps(): + rec = _build_mock_recording() + times = rec.get_times(segment_index=0) + 100 + rec.set_times(times, segment_index=0) + return rec + + +def _check_recording_with_timestamps(rec): + import numpy as np + expected_times = np.arange(int(DEFAULT_DURATION * 30000)) / 30000.0 + 100 + times = rec.get_times(segment_index=0) + assert np.allclose(times, expected_times) + + def _check_recording_with_probe(rec): import numpy as np @@ -121,7 +143,7 @@ def _build_recording_with_interleaved_probes(): from probeinterface import ProbeGroup, generate_linear_probe from spikeinterface.core import generate_recording - rec = generate_recording(num_channels=8, durations=[1.0], sampling_frequency=30000.0, seed=0) + rec = generate_recording(num_channels=8, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) probe0 = generate_linear_probe(num_elec=4) probe1 = generate_linear_probe(num_elec=4) probe1.move([100.0, 0.0]) @@ -164,7 +186,7 @@ def _build_preprocessed_chain(): from spikeinterface.core import generate_recording from spikeinterface.preprocessing import common_reference, scale - rec = generate_recording(num_channels=4, durations=[1.0], sampling_frequency=30000.0, seed=0) + rec = generate_recording(num_channels=4, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) # Two nested scipy-free preprocessing wrappers (scale then common_reference): this # exercises recursive parent reload without pulling scipy into the environments. return common_reference(scale(rec, gain=2.0)) @@ -181,7 +203,7 @@ def _check_preprocessed_chain(rec): def _build_sorting(): from spikeinterface.core import generate_sorting - return generate_sorting(num_units=5, sampling_frequency=30000.0, durations=[1.0]) + return generate_sorting(num_units=5, sampling_frequency=30000.0, durations=[DEFAULT_DURATION]) def _check_sorting(sorting): @@ -195,7 +217,7 @@ def _build_sorting_with_properties(): import numpy as np from spikeinterface.core import generate_sorting - sorting = generate_sorting(num_units=4, sampling_frequency=30000.0, durations=[1.0]) + sorting = generate_sorting(num_units=4, sampling_frequency=30000.0, durations=[DEFAULT_DURATION]) sorting.set_property("quality", np.array(["good", "good", "bad", "good"])) sorting.annotate(experimenter="test") return sorting @@ -224,19 +246,25 @@ def _check_sorting_with_properties(sorting): "id": "recording_with_properties", "build": _build_recording_with_properties, "check": _check_recording_with_properties, - "formats": ["binary", "zarr"], + "formats": ["binary", "zarr", "binary_parallel", "zarr_parallel"], + }, + { + "id": "recording_with_timestamps", + "build": _build_recording_with_timestamps, + "check": _check_recording_with_timestamps, + "formats": ["binary", "zarr", "binary_parallel", "zarr_parallel"], }, { "id": "recording_with_probe", "build": _build_recording_with_probe, "check": _check_recording_with_probe, - "formats": ["binary", "zarr"], + "formats": ["binary", "zarr", "binary_parallel", "zarr_parallel"], }, { "id": "recording_with_interleaved_probes", "build": _build_recording_with_interleaved_probes, "check": _check_recording_with_interleaved_probes, - "formats": ["binary", "zarr"], + "formats": ["binary", "zarr", "binary_parallel", "zarr_parallel"], }, { "id": "preprocessed_chain", diff --git a/.github/scripts/serialization/serialize_objects.py b/.github/scripts/serialization/serialize_objects.py index 44a68ca26e..c4e9326f08 100644 --- a/.github/scripts/serialization/serialize_objects.py +++ b/.github/scripts/serialization/serialize_objects.py @@ -33,10 +33,14 @@ obj.dump_to_pickle(dest) elif fmt == "binary": obj.save(folder=dest, format="binary", overwrite=True) + elif fmt == "binary_parallel": + obj.save(folder=dest, format="binary", overwrite=True, n_jobs=2) elif fmt == "numpy_folder": obj.save(folder=dest, format="numpy_folder", overwrite=True) elif fmt == "zarr": obj.save(folder=dest, format="zarr", overwrite=True) + elif fmt == "zarr_parallel": + obj.save(folder=dest, format="zarr", overwrite=True, n_jobs=2) print(f" wrote {dest.name} ({fmt})") print(f"Fixtures written to: {out_dir.resolve()}") diff --git a/.github/workflows/cross_version_serialization.yml b/.github/workflows/cross_version_serialization.yml index 4712a3976d..4bafaed4e1 100644 --- a/.github/workflows/cross_version_serialization.yml +++ b/.github/workflows/cross_version_serialization.yml @@ -6,9 +6,8 @@ name: Cross-version serialization # # Delivery model: Live generation. Fixtures are regenerated from a real old install each # run in an isolated uv environment (nothing committed). The version list is computed -# from PyPI (latest patch of each minor at or above MIN_VERSION_TO_TEST): on a pull -# request only the latest released minor is tested; the full matrix runs weekly and on -# manual dispatch. +# from PyPI (latest patch of each minor at or above MIN_VERSION_TO_TEST), and the full +# matrix is tested on every pull request and on manual dispatch. on: pull_request: @@ -20,11 +19,9 @@ on: # watches core only; a break introduced in generation/ or preprocessing/ would not # trigger it. paths: - - "src/spikeinterface/core/**" - - ".github/scripts/serialization/**" - - ".github/workflows/cross_version_serialization.yml" - schedule: - - cron: "0 4 * * 0" # weekly, Sunday 04:00 UTC + - 'src/spikeinterface/core/**' + - '.github/scripts/serialization/**' + - '.github/workflows/cross_version_serialization.yml' workflow_dispatch: concurrency: @@ -33,8 +30,8 @@ concurrency: jobs: # Compute which released versions the matrix below tests loading from, exposed as the - # job output `list`. Latest patch of each minor at or above the floor on schedule / - # dispatch; only the latest minor on a pull request (keeps per-PR runs cheap). + # job output `list`: the latest patch of each minor at or above the floor, tested on + # every trigger. versions: name: Select versions runs-on: ubuntu-latest @@ -46,7 +43,6 @@ jobs: python-version: "3.10" - id: set env: - GITHUB_EVENT_NAME: ${{ github.event_name }} # Support floor: oldest minor to test. Below this, installs fail on the CI # Python (numcodecs dependency rot in 0.100/0.101). Raise when 0.102 rots. MIN_VERSION_TO_TEST: "0.102" # We will change this when making breaking changes (hopefully not too often) @@ -67,10 +63,6 @@ jobs: candidates = [v for v in available if not v.is_prerelease and (v.major, v.minor) >= (floor.major, floor.minor)] minor_releases = [max(group) for _, group in groupby(candidates, key=lambda v: (v.major, v.minor))] - # on a pull request, test only the latest minor - if os.environ.get('GITHUB_EVENT_NAME') == 'pull_request': - minor_releases = minor_releases[-1:] - print(json.dumps([str(v) for v in minor_releases])) ") diff --git a/src/spikeinterface/benchmark/benchmark_base.py b/src/spikeinterface/benchmark/benchmark_base.py index 47ec6a70cd..2bd55791a7 100644 --- a/src/spikeinterface/benchmark/benchmark_base.py +++ b/src/spikeinterface/benchmark/benchmark_base.py @@ -190,7 +190,7 @@ def create(cls, study_folder, datasets={}, cases={}, levels=None): # sortings are pickled + saved as NumpyFolderSorting # gt_sorting.dump_to_pickle(study_folder / f"datasets/gt_sortings/{key}.pickle") - # gt_sorting.save(format="numpy_folder", folder=study_folder / f"datasets/gt_sortings/{key}") + # gt_sorting.save(format="binary", folder=study_folder / f"datasets/gt_sortings/{key}") # analyzer path (local or external) (study_folder / "analyzers_path.json").write_text(json.dumps(analyzers_path, indent=4), encoding="utf8") @@ -651,7 +651,7 @@ def _save_keys(self, saved_keys, folder): with open(folder / f"{k}.pickle", mode="wb") as f: pickle.dump(self.result[k], f) elif format == "sorting": - self.result[k].save(folder=folder / k, format="numpy_folder", overwrite=True) + self.result[k].save(folder=folder / k, format="binary", overwrite=True) elif format == "Motion": self.result[k].save(folder=folder / k) elif format == "zarr_templates": diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index d33ca7f72e..76a50d27dd 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -488,7 +488,6 @@ def to_dict( include_properties: bool = False, include_extra_metadata: bool = True, relative_to: str | Path | None = None, - folder_metadata=None, recursive: bool = False, ) -> dict: """ @@ -516,9 +515,6 @@ def to_dict( If provided, file and folder paths will be made relative to this path, enabling portability in folder formats such as the waveform extractor, by default None. - folder_metadata : str | Path | None, default: None - Path to a folder containing additional metadata files (e.g., probe information in BaseRecording) - in numpy `npy` format, by default None. recursive : bool, default: False If True, recursively apply `to_dict` to dictionaries within the kwargs, by default False. @@ -539,13 +535,11 @@ def to_dict( "relative_paths": , "annotations": , "properties": , - "folder_metadata": } Notes ----- - The `relative_to` argument only has an effect if `recursive` is set to True. - - The `folder_metadata` argument will be made relative to `relative_to` if both are specified. - The `version` field in the resulting dictionary reflects the version of the module from which the extractor class originates. - The full class attribute above is the full import of the class, e.g. @@ -562,9 +556,8 @@ def to_dict( to_dict_kwargs = dict( include_annotations=include_annotations, include_properties=include_properties, - # make_paths_relative() will make the recusrivity later: + # make_paths_relative() will make the recursivity later: relative_to=None, - folder_metadata=folder_metadata, recursive=recursive, ) @@ -588,7 +581,9 @@ def to_dict( dump_dict["annotations"] = self._annotations else: # include only main annotations - dump_dict["annotations"] = {k: self._annotations.get(k, None) for k in self._main_annotations} + dump_dict["annotations"] = { + k: self._annotations.get(k) for k in self._main_annotations if self._annotations.get(k) is not None + } if include_properties: dump_dict["properties"] = self._properties @@ -611,11 +606,6 @@ def to_dict( # warnings.warn("Try to BaseExtractor.to_dict() using relative_to but there is no common folder") dump_dict["relative_paths"] = False - if folder_metadata is not None: - if relative_to is not None: - folder_metadata = Path(folder_metadata).resolve().absolute().relative_to(relative_to) - dump_dict["folder_metadata"] = str(folder_metadata) - if include_extra_metadata: self._extra_metadata_to_dict(dump_dict) @@ -644,13 +634,6 @@ def from_dict(dictionary: dict, base_folder: Path | str | None = None) -> "BaseE dictionary = make_paths_absolute(dictionary, base_folder) extractor = _load_extractor_from_dict(dictionary) - # TODO sam : need to check but normally this is not usefull anymore - # folder_metadata = dictionary.get("folder_metadata", None) - # if folder_metadata is not None: - # folder_metadata = Path(folder_metadata) - # if dictionary.get("relative_paths", False): - # folder_metadata = base_folder / folder_metadata - # load_properties_from_binary_folder(folder_metadata, self) return extractor def clone(self) -> "BaseExtractor": @@ -711,7 +694,7 @@ def _get_file_path(file_path: str | Path, extensions: Sequence) -> Path: ) return file_path - def dump(self, file_path: str | Path, relative_to=None, folder_metadata=None) -> None: + def dump(self, file_path: str | Path, relative_to=None) -> None: """ Dumps extractor to json or pickle @@ -724,9 +707,9 @@ def dump(self, file_path: str | Path, relative_to=None, folder_metadata=None) -> This means that file and folder paths in extractor objects kwargs are changed to be relative rather than absolute. """ if str(file_path).endswith(".json"): - self.dump_to_json(file_path, relative_to=relative_to, folder_metadata=folder_metadata) + self.dump_to_json(file_path, relative_to=relative_to) elif str(file_path).endswith(".pkl") or str(file_path).endswith(".pickle"): - self.dump_to_pickle(file_path, relative_to=relative_to, folder_metadata=folder_metadata) + self.dump_to_pickle(file_path, relative_to=relative_to) else: raise ValueError("Dump: file must .json or .pkl") @@ -734,8 +717,9 @@ def dump_to_json( self, file_path: str | Path | None = None, relative_to: str | Path | bool | None = None, - folder_metadata: str | Path | None = None, include_extra_metadata: bool = True, + include_properties: bool = False, + include_annotations: bool = True, ) -> None: """ Dump recording extractor to json file. @@ -748,8 +732,12 @@ def dump_to_json( relative_to: str, Path, True or None If not None, files and folders are serialized relative to this path. If True, the relative folder is the parent folder. This means that file and folder paths in extractor objects kwargs are changed to be relative rather than absolute. - folder_metadata: str, Path, or None - Folder with files containing additional information (e.g. probe in BaseRecording) and properties + include_extra_metadata: bool + If True, extra metadata is included in the json file. This is useful for saving probe + include_properties: bool + If True, all properties are dumped + include_annotations: bool + If True, all annotations are dumped """ assert self.check_serializability("json"), "The extractor is not json serializable" @@ -759,11 +747,10 @@ def dump_to_json( relative_to = relative_to.resolve().absolute() dump_dict = self.to_dict( - include_annotations=True, - include_properties=False, + include_annotations=include_annotations, + include_properties=include_properties, include_extra_metadata=include_extra_metadata, relative_to=relative_to, - folder_metadata=folder_metadata, recursive=True, ) file_path = self._get_file_path(file_path, [".json"]) @@ -778,7 +765,6 @@ def dump_to_pickle( file_path: str | Path | None = None, relative_to: str | Path | bool | None = None, include_properties: bool = True, - folder_metadata: str | Path | None = None, ): """ Dump recording extractor to a pickle file. @@ -793,8 +779,6 @@ def dump_to_pickle( This means that file and folder paths in extractor objects kwargs are changed to be relative rather than absolute. include_properties: bool If True, all properties are dumped - folder_metadata: str, Path, or None - Folder with files containing additional information (e.g. probe in BaseRecording) and properties. """ assert self.check_serializability("pickle"), "The extractor is not serializable to file with pickle" @@ -810,7 +794,6 @@ def dump_to_pickle( dump_dict = self.to_dict( include_annotations=True, include_properties=include_properties, - folder_metadata=folder_metadata, relative_to=relative_to, recursive=recursive, ) diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 81375757ad..1b30c75b06 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -423,6 +423,28 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): return cached + def _extra_metadata_to_dict(self, dump_dict): + super()._extra_metadata_to_dict(dump_dict) + + # Add times_kwargs if the recording has been modified in memory (e.g. by set_times / shift_times / reset_times) + if self._time_info_modified: + dump_dict["times_kwargs"] = [] + for segment_index in range(self.get_num_segments()): + times_kwargs = self.segments[segment_index].get_times_kwargs() + dump_dict["times_kwargs"].append(times_kwargs) + + def _extra_metadata_from_dict(self, dump_dict): + super()._extra_metadata_from_dict(dump_dict) + + if "times_kwargs" in dump_dict: + # When serializing, dump timestamps information because this could have been + # set in memory + times_kwargs_list = dump_dict["times_kwargs"] + for segment_index, times_kwargs in enumerate(times_kwargs_list): + self.segments[segment_index]._sampling_frequency = times_kwargs["sampling_frequency"] + self.segments[segment_index]._t_start = times_kwargs["t_start"] + self.segments[segment_index]._time_vector = times_kwargs["time_vector"] + def select_channels(self, channel_ids: list | np.ndarray | tuple) -> "BaseRecording": """ Returns a new recording object with a subset of channels. diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index 464959e467..c35b722557 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -141,9 +141,23 @@ def write_recording( file_paths = [folder_path / f"traces_cached_seg{i}.raw" for i in range(recording.get_num_segments())] if dtype is None: dtype = recording.get_dtype() - t_starts = recording._get_t_starts() - - write_binary(recording, file_paths=file_paths, dtype=dtype, verbose=verbose, **job_kwargs) + # Check if there are any time vectors + t_starts = recording.get_segment_t_starts() + if recording.has_any_time_vector(): + file_timestamps_paths = [ + folder_path / f"times_cached_seg{i}.raw" for i in range(recording.get_num_segments()) + ] + else: + file_timestamps_paths = None + + write_binary( + recording, + file_paths=file_paths, + file_timestamps_paths=file_timestamps_paths, + dtype=dtype, + verbose=verbose, + **job_kwargs, + ) save_properties_to_binary_folder(folder_path / "properties", recording) save_extractor_provenance(folder_path, recording) @@ -156,9 +170,9 @@ def write_recording( # This is created so it can be saved as json because the `BinaryFolderRecording` requires it loading # See the __init__ - binary_rec = BinaryRecordingExtractor( file_paths=file_paths, + file_timestamps_paths=file_timestamps_paths, sampling_frequency=recording.get_sampling_frequency(), num_channels=recording.get_num_channels(), dtype=dtype, @@ -172,18 +186,18 @@ def write_recording( ) binary_rec.dump(folder_path / "binary.json", relative_to=folder_path) - # TODO alessio : remove this, it is needed to pass tests - # save times - for segment_index, rs in enumerate(recording.segments): - d = rs.get_times_kwargs() - time_vector = d["time_vector"] - if time_vector is not None: - np.save(folder_path / f"times_cached_seg{segment_index}.npy", time_vector) - - # make the si_folder file to make the load() easier until version 0.105.0 + # Create the si_folder file to make the load() easier until version 0.105.0 + # All properties, annotations, and probe information are already saved in the folder, + # so we don't need to include them in the si_folder.json cached = BinaryFolderRecording(folder_path=folder_path) si_folder_path = folder_path / f"si_folder.json" - cached.dump_to_json(file_path=si_folder_path, relative_to=folder_path, include_extra_metadata=False) + cached.dump_to_json( + file_path=si_folder_path, + relative_to=folder_path, + include_properties=False, + include_annotations=False, + include_extra_metadata=False, + ) return cached diff --git a/src/spikeinterface/core/binaryrecordingextractor.py b/src/spikeinterface/core/binaryrecordingextractor.py index 059526a028..37f4ba3467 100644 --- a/src/spikeinterface/core/binaryrecordingextractor.py +++ b/src/spikeinterface/core/binaryrecordingextractor.py @@ -38,6 +38,8 @@ class BinaryRecordingExtractor(BaseRecording): The offset to apply to the traces is_filtered : bool or None, default: None If True, the recording is assumed to be filtered. If None, is_filtered is not set. + file_timestamps_paths : str or Path or list, default: None + Path to the binary file containing timestamps for each segment. If None, timestamps are not loaded Notes ----- @@ -51,17 +53,18 @@ class BinaryRecordingExtractor(BaseRecording): def __init__( self, - file_paths, - sampling_frequency, - dtype, + file_paths: str | Path | list[str | Path], + sampling_frequency: float, + dtype: str | np.dtype, num_channels: int | None = None, - t_starts=None, - channel_ids=None, - time_axis=0, - file_offset=0, - gain_to_uV=None, - offset_to_uV=None, - is_filtered=None, + t_starts: list[float] | None = None, + channel_ids: list[str | int] | None = None, + time_axis: int = 0, + file_offset: int = 0, + gain_to_uV: float | np.ndarray | None = None, + offset_to_uV: float | np.ndarray | None = None, + is_filtered: bool | None = None, + file_timestamps_paths: str | Path | list[str | Path] | None = None, ): if channel_ids is None: @@ -82,6 +85,12 @@ def __init__( assert len(t_starts) == len(file_path_list), "t_starts must be a list of the same size as file_paths" t_starts = [float(t_start) for t_start in t_starts] + if file_timestamps_paths is not None: + if isinstance(file_timestamps_paths, list): + file_timestamps_paths = [Path(p) for p in file_timestamps_paths] + else: + file_timestamps_paths = [Path(file_timestamps_paths)] + dtype = np.dtype(dtype) for i, file_path in enumerate(file_path_list): @@ -89,8 +98,20 @@ def __init__( t_start = None else: t_start = t_starts[i] + if file_timestamps_paths is None: + file_timestamps_path = None + else: + file_timestamps_path = file_timestamps_paths[i] + rec_segment = BinaryRecordingSegment( - file_path, sampling_frequency, t_start, num_channels, dtype, time_axis, file_offset + file_path, + sampling_frequency, + t_start, + num_channels, + dtype, + time_axis, + file_offset, + file_timestamps_path, ) self.add_recording_segment(rec_segment) @@ -115,6 +136,9 @@ def __init__( "gain_to_uV": gain_to_uV, "offset_to_uV": offset_to_uV, "is_filtered": is_filtered, + "file_timestamps_paths": ( + [str(Path(e).absolute()) for e in file_timestamps_paths] if file_timestamps_paths is not None else None + ), } @classmethod @@ -163,7 +187,17 @@ def get_binary_description(self): class BinaryRecordingSegment(BaseRecordingSegment): - def __init__(self, file_path, sampling_frequency, t_start, num_channels, dtype, time_axis, file_offset): + def __init__( + self, + file_path, + sampling_frequency, + t_start, + num_channels, + dtype, + time_axis, + file_offset, + file_timestamps_path=None, + ): BaseRecordingSegment.__init__(self, sampling_frequency=sampling_frequency, t_start=t_start) self.num_channels = num_channels self.dtype = np.dtype(dtype) @@ -174,6 +208,9 @@ def __init__(self, file_path, sampling_frequency, t_start, num_channels, dtype, self.bytes_per_sample = self.num_channels * self.dtype.itemsize self.data_size_in_bytes = Path(file_path).stat().st_size - file_offset self.num_samples = self.data_size_in_bytes // self.bytes_per_sample + self.file_timestamps_path = file_timestamps_path + if file_timestamps_path is not None: + self._time_vector = np.memmap(file_timestamps_path, dtype="float64", mode="r", shape=(self.num_samples,)) def get_num_samples(self) -> int: """Returns the number of samples in this signal block @@ -240,6 +277,13 @@ def __del__(self): warnings.warn(f"Error closing file handle in BinaryRecordingSegment: {e}") pass + if self._time_vector is not None: + try: + self._time_vector._mmap.close() # Close the underlying mmap object + del self._time_vector + except Exception as e: + pass + # For backward compatibility (old good time) BinDatRecordingExtractor = BinaryRecordingExtractor diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index 8e24027064..0edb2a7ea4 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -849,34 +849,35 @@ def save_properties_to_binary_folder(folder: str | Path, extractor: "BaseExtract def save_annotations_to_folder(folder: str | Path, extractor: "BaseExtractor"): """ - Save annotaions in json format from an extractor (recording or sorting). + Save annotations in json format from an extractor (recording or sorting). This used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0 Parameters ---------- folder : str or Path - The folder where the properties will be saved as .npy files. + The folder where the annotations will be saved as a json file. extractor : BaseExtractor The extractor from which the annotations will be saved. """ folder = Path(folder) - (folder / "annotations").write_text(json.dumps(extractor._annotations, indent=4), encoding="utf8") + print(f"Saving annotations to {folder / 'annotations.json'}: {extractor._annotations}") + (folder / "annotations.json").write_text(json.dumps(extractor._annotations, indent=4), encoding="utf8") def load_annotations_from_folder(folder: str | Path, extractor: "BaseExtractor"): """ - Save annotaions in json format from an extractor (recording or sorting). + Load annotations in json format from an extractor (recording or sorting). This used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0 Parameters ---------- folder : str or Path - The folder where the properties will be saved as .npy files. + The folder where the annotations will be loaded from. extractor : BaseExtractor - The extractor from which the annotations will be saved. + The extractor to which the annotations will be added. """ folder = Path(folder) - annotations_file = folder / "annotations" + annotations_file = folder / "annotations.json" if annotations_file.exists(): with open(annotations_file, "r") as f: annotations = json.load(f) diff --git a/src/spikeinterface/core/loading.py b/src/spikeinterface/core/loading.py index 4f92ca96bc..9c06e95300 100644 --- a/src/spikeinterface/core/loading.py +++ b/src/spikeinterface/core/loading.py @@ -245,7 +245,7 @@ def _load_object_from_folder(folder, object_type: str, **kwargs): f = folder / f"cached.{dump_ext}" if f.is_file(): si_file = f - return BaseExtractor.load(si_file, base_folder=folder) + return load(si_file, base_folder=folder) elif object_type.startswith("Group"): diff --git a/src/spikeinterface/core/numpyextractors.py b/src/spikeinterface/core/numpyextractors.py index f5ebc7b15c..b58a98eddb 100644 --- a/src/spikeinterface/core/numpyextractors.py +++ b/src/spikeinterface/core/numpyextractors.py @@ -80,7 +80,7 @@ def __init__(self, traces_list, sampling_frequency, t_starts=None, channel_ids=N def from_recording(source_recording, with_metadata=True, with_time_vector=False, **job_kwargs): traces_list, shms = write_memory_recording(source_recording, dtype=None, **job_kwargs) - t_starts = source_recording._get_t_starts() + t_starts = source_recording.get_segment_t_starts() if shms[0] is not None: # if the computation was done in parallel then traces_list is shared array @@ -219,7 +219,7 @@ def __del__(self): def from_recording(source_recording, with_metadata=True, with_time_vector=False, **job_kwargs): traces_list, shms = write_memory_recording(source_recording, buffer_type="sharedmem", **job_kwargs) - t_starts = source_recording._get_t_starts() + t_starts = source_recording.get_segment_t_starts() recording = SharedMemoryRecording( shm_names=[shm.name for shm in shms], diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index f9fc81361d..7916f24b2e 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -28,7 +28,7 @@ class NumpyFolderSorting(BaseSorting): * a "numpysorting_info.json" containing sampling_frequency, unit_ids and num_segments * a metadata folder for units properties. - It is created with the function: `sorting.save(folder="/myfolder", format="numpy_folder")` + It is created with the function: `sorting.save(folder="/myfolder", format="binary")` """ @@ -85,10 +85,18 @@ def write_sorting(sorting, folder_path, overwrite: bool = False): # new in version 0.105.0, before that annotations were handle by "si_folder.json" file save_annotations_to_folder(folder_path, sorting) - # make the si_folder file to make the load() easier until version 0.105.0 + # Create the si_folder file to make the load() easier until version 0.105.0 + # All properties, annotations, and probe information are already saved in the folder, + # so we don't need to include them in the si_folder.json cached = NumpyFolderSorting(folder_path=folder_path) si_folder_path = folder_path / f"si_folder.json" - cached.dump_to_json(file_path=si_folder_path, relative_to=folder_path, include_extra_metadata=False) + cached.dump_to_json( + file_path=si_folder_path, + relative_to=folder_path, + include_extra_metadata=False, + include_properties=False, + include_annotations=False, + ) return cached diff --git a/src/spikeinterface/core/tests/test_baserecording.py b/src/spikeinterface/core/tests/test_baserecording.py index 7debd337c3..809df7abbe 100644 --- a/src/spikeinterface/core/tests/test_baserecording.py +++ b/src/spikeinterface/core/tests/test_baserecording.py @@ -6,6 +6,7 @@ import json import pickle from pathlib import Path +import platform import pytest import numpy as np from numpy.testing import assert_raises @@ -17,7 +18,7 @@ NumpyRecording, load, get_default_zarr_compressor, - aggregate_channels, + load, ) from spikeinterface.core.base import BaseExtractor from spikeinterface.core.testing import check_recordings_equal @@ -101,14 +102,14 @@ def test_BaseRecording(create_cache_folder): # dump/load json rec.dump_to_json(cache_folder / "test_BaseRecording.json") - rec2 = BaseExtractor.load(cache_folder / "test_BaseRecording.json") + rec2 = load(cache_folder / "test_BaseRecording.json") rec3 = load(cache_folder / "test_BaseRecording.json") check_recordings_equal(rec, rec2, return_in_uV=False, check_annotations=True, check_properties=False) check_recordings_equal(rec, rec3, return_in_uV=False, check_annotations=True, check_properties=False) # dump/load pickle rec.dump_to_pickle(cache_folder / "test_BaseRecording.pkl") - rec2 = BaseExtractor.load(cache_folder / "test_BaseRecording.pkl") + rec2 = load(cache_folder / "test_BaseRecording.pkl") rec3 = load(cache_folder / "test_BaseRecording.pkl") check_recordings_equal(rec, rec2, return_in_uV=False, check_annotations=True, check_properties=True) check_recordings_equal(rec, rec3, return_in_uV=False, check_annotations=True, check_properties=True) @@ -120,12 +121,12 @@ def test_BaseRecording(create_cache_folder): # dump/load json - relative to rec.dump_to_json(cache_folder / "test_BaseRecording_rel.json", relative_to=cache_folder) - rec2 = BaseExtractor.load(cache_folder / "test_BaseRecording_rel.json", base_folder=cache_folder) + rec2 = load(cache_folder / "test_BaseRecording_rel.json", base_folder=cache_folder) rec3 = load(cache_folder / "test_BaseRecording_rel.json", base_folder=cache_folder) # dump/load relative=True rec.dump_to_json(cache_folder / "test_BaseRecording_rel_true.json", relative_to=True) - rec2 = BaseExtractor.load(cache_folder / "test_BaseRecording_rel_true.json", base_folder=True) + rec2 = load(cache_folder / "test_BaseRecording_rel_true.json", base_folder=True) rec3 = load(cache_folder / "test_BaseRecording_rel_true.json", base_folder=True) check_recordings_equal(rec, rec2, return_in_uV=False, check_annotations=True) check_recordings_equal(rec, rec3, return_in_uV=False, check_annotations=True) @@ -137,12 +138,12 @@ def test_BaseRecording(create_cache_folder): # dump/load pkl - relative to rec.dump_to_pickle(cache_folder / "test_BaseRecording_rel.pkl", relative_to=cache_folder) - rec2 = BaseExtractor.load(cache_folder / "test_BaseRecording_rel.pkl", base_folder=cache_folder) + rec2 = load(cache_folder / "test_BaseRecording_rel.pkl", base_folder=cache_folder) rec3 = load(cache_folder / "test_BaseRecording_rel.pkl", base_folder=cache_folder) # dump/load relative=True rec.dump_to_pickle(cache_folder / "test_BaseRecording_rel_true.pkl", relative_to=True) - rec2 = BaseExtractor.load(cache_folder / "test_BaseRecording_rel_true.pkl", base_folder=True) + rec2 = load(cache_folder / "test_BaseRecording_rel_true.pkl", base_folder=True) rec3 = load(cache_folder / "test_BaseRecording_rel_true.pkl", base_folder=True) check_recordings_equal(rec, rec2, return_in_uV=False, check_annotations=True) check_recordings_equal(rec, rec3, return_in_uV=False, check_annotations=True) @@ -155,7 +156,7 @@ def test_BaseRecording(create_cache_folder): # cache to binary folder = cache_folder / "simple_recording" rec.save(format="binary", folder=folder) - rec2 = BaseExtractor.load(folder) + rec2 = load(folder) assert "quality" in rec2.get_property_keys() values = rec2.get_property("quality") assert values[0] == 1.0 @@ -166,7 +167,7 @@ def test_BaseRecording(create_cache_folder): assert np.array_equal(groups, [0, 0, 1]) # but also possible - rec3 = BaseExtractor.load(cache_folder / "simple_recording") + rec3 = load(cache_folder / "simple_recording") # cache to memory rec4 = rec3.save(format="memory", shared=False) @@ -494,6 +495,57 @@ def test_time_slice_with_time_vector(): assert np.allclose(sliced_recording_times.get_traces(), sliced_recording_frames.get_traces()) +@pytest.mark.parametrize( + "mp_context", + [ + pytest.param( + "fork", marks=pytest.mark.skipif(platform.system() != "Linux", reason="fork only supported on Linux") + ), + pytest.param( + "forkserver", + marks=pytest.mark.skipif(platform.system() != "Linux", reason="forkserver only supported on Linux"), + ), + "spawn", + ], +) +def test_save_load_binary_with_time_vector(create_cache_folder, mp_context): + cache_folder = create_cache_folder + rec = generate_recording(durations=[5.0], num_channels=3, sampling_frequency=10_000.0) + times = rec.get_times(segment_index=0) + 100.0 + + # Set time vector + rec.set_times(times=times, segment_index=0, with_warning=False) + # Save + rec_saved = rec.save(folder=cache_folder / f"recording_with_time_vector_{mp_context}", format="binary") + assert np.allclose(rec.get_times(segment_index=0), rec_saved.get_times(segment_index=0)) + + # Save + rec_saved_par = rec.save( + folder=cache_folder / f"recording_with_time_vector_par_{mp_context}", + format="binary", + n_jobs=2, + mp_context=mp_context, + ) + assert np.allclose(rec.get_times(segment_index=0), rec_saved_par.get_times(segment_index=0)) + + # Now reset_times and save again, to check that the time vector is not saved + rec_saved.reset_times() + rec_saved_no_time_vector = rec_saved.save( + folder=cache_folder / f"recording_without_time_vector_{mp_context}", format="binary" + ) + assert not rec_saved_no_time_vector.has_time_vector(segment_index=0) + + # Now make sure the same happens if we save in parallel with multiple jobs, which requires pickling/unpickling + # the recording object + rec_saved_no_time_vector_par = rec_saved.save( + folder=cache_folder / f"recording_without_time_vector_par_{mp_context}", + format="binary", + n_jobs=2, + mp_context=mp_context, + ) + assert not rec_saved_no_time_vector_par.has_time_vector(segment_index=0) + + if __name__ == "__main__": import tempfile diff --git a/src/spikeinterface/core/tests/test_basesnippets.py b/src/spikeinterface/core/tests/test_basesnippets.py index 8f163cc4fa..585672a1ea 100644 --- a/src/spikeinterface/core/tests/test_basesnippets.py +++ b/src/spikeinterface/core/tests/test_basesnippets.py @@ -9,8 +9,7 @@ from numpy.testing import assert_raises from probeinterface import Probe -from spikeinterface.core import generate_snippets -from spikeinterface.core import NumpySnippets, load +from spikeinterface.core import generate_snippets, load, NumpySnippets from spikeinterface.core.npysnippetsextractor import NpySnippetsExtractor from spikeinterface.core.base import BaseExtractor @@ -94,12 +93,12 @@ def test_BaseSnippets(create_cache_folder): # dump/load json snippets.dump_to_json(cache_folder / "test_BaseSnippets.json") - snippets2 = BaseExtractor.load(cache_folder / "test_BaseSnippets.json") + snippets2 = load(cache_folder / "test_BaseSnippets.json") snippets3 = load(cache_folder / "test_BaseSnippets.json") # dump/load pickle snippets.dump_to_pickle(cache_folder / "test_BaseSnippets.pkl") - snippets2 = BaseExtractor.load(cache_folder / "test_BaseSnippets.pkl") + snippets2 = load(cache_folder / "test_BaseSnippets.pkl") snippets3 = load(cache_folder / "test_BaseSnippets.pkl") # dump/load dict - relative @@ -109,14 +108,14 @@ def test_BaseSnippets(create_cache_folder): # dump/load json snippets.dump_to_json(cache_folder / "test_BaseSnippets_rel.json", relative_to=cache_folder) - snippets2 = BaseExtractor.load(cache_folder / "test_BaseSnippets_rel.json", base_folder=cache_folder) + snippets2 = load(cache_folder / "test_BaseSnippets_rel.json", base_folder=cache_folder) snippets3 = load(cache_folder / "test_BaseSnippets_rel.json", base_folder=cache_folder) # cache to npy folder = cache_folder / "simple_snippets" print(folder) snippets.save(format="npy", folder=folder) - snippets2 = BaseExtractor.load(folder) + snippets2 = load(folder) assert "quality" in snippets2.get_property_keys() values = snippets2.get_property("quality") assert values[0] == 1.0 @@ -127,7 +126,7 @@ def test_BaseSnippets(create_cache_folder): assert np.array_equal(groups, [0, 0, 1]) # but also possible - snippets3 = BaseExtractor.load(cache_folder / "simple_snippets") + snippets3 = load(cache_folder / "simple_snippets") # cache to memory snippets4 = snippets3.save(format="memory") diff --git a/src/spikeinterface/core/tests/test_basesorting.py b/src/spikeinterface/core/tests/test_basesorting.py index 75687c54e6..4e72e80877 100644 --- a/src/spikeinterface/core/tests/test_basesorting.py +++ b/src/spikeinterface/core/tests/test_basesorting.py @@ -56,14 +56,14 @@ def test_BaseSorting(create_cache_folder): # dump/load json sorting.dump_to_json(cache_folder / "test_BaseSorting.json") - sorting2 = BaseExtractor.load(cache_folder / "test_BaseSorting.json") + sorting2 = load(cache_folder / "test_BaseSorting.json") sorting3 = load(cache_folder / "test_BaseSorting.json") check_sortings_equal(sorting, sorting2, check_annotations=True, check_properties=False) check_sortings_equal(sorting, sorting3, check_annotations=True, check_properties=False) # dump/load pickle sorting.dump_to_pickle(cache_folder / "test_BaseSorting.pkl") - sorting2 = BaseExtractor.load(cache_folder / "test_BaseSorting.pkl") + sorting2 = load(cache_folder / "test_BaseSorting.pkl") sorting3 = load(cache_folder / "test_BaseSorting.pkl") check_sortings_equal(sorting, sorting2, check_annotations=True, check_properties=True) check_sortings_equal(sorting, sorting3, check_annotations=True, check_properties=True) @@ -72,18 +72,18 @@ def test_BaseSorting(create_cache_folder): folder = cache_folder / "simple_sorting_npz_folder" sorting.set_property("test", np.ones(len(sorting.unit_ids))) sorting.save(folder=folder, format="npz_folder") - sorting2 = BaseExtractor.load(folder) + sorting2 = load(folder) assert isinstance(sorting2, NpzFolderSorting) - # cache new format : numpy_folder - folder = cache_folder / "simple_sorting_numpy_folder" + # cache new format : binary + folder = cache_folder / "simple_sorting_binary" sorting.set_property("test", np.ones(len(sorting.unit_ids))) sorting.save(folder=folder, format="binary") - sorting2 = BaseExtractor.load(folder) + sorting2 = load(folder) assert isinstance(sorting2, NumpyFolderSorting) # but also possible - sorting3 = BaseExtractor.load(folder) + sorting3 = load(folder) check_sortings_equal(sorting, sorting2, check_annotations=True, check_properties=True) check_sortings_equal(sorting, sorting3, check_annotations=True, check_properties=True) diff --git a/src/spikeinterface/core/time_series.py b/src/spikeinterface/core/time_series.py index ec5bee656b..0d6cbc6e42 100644 --- a/src/spikeinterface/core/time_series.py +++ b/src/spikeinterface/core/time_series.py @@ -33,6 +33,8 @@ class TimeSeries(ABC): """ _preferred_mp_context = None + # Flag to indicate whether time info has been modified in-memory (e.g. by set_times or shift_times). + _time_info_modified = False def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -216,9 +218,7 @@ def has_time_vector(self, segment_index: Optional[int] = None): True if the recording has time vectors, False otherwise """ segment_index = self._check_segment_index(segment_index) - rs = self.segments[segment_index] - d = rs.get_times_kwargs() - return d["time_vector"] is not None + return self.segments[segment_index].has_time_vector() def set_times(self, times, segment_index=None, with_warning=True): """Set times for a recording segment. @@ -248,6 +248,7 @@ def set_times(self, times, segment_index=None, with_warning=True): "times are not always propagated across preprocessing" "Use this carefully!" ) + self._time_info_modified = True def reset_times(self): """ @@ -262,6 +263,7 @@ def reset_times(self): rs._time_vector = None rs._t_start = None rs._sampling_frequency = self.sampling_frequency + self._time_info_modified = True def shift_times(self, shift: int | float, segment_index: int | None = None) -> None: """ @@ -297,6 +299,7 @@ def shift_times(self, shift: int | float, segment_index: int | None = None) -> N else: new_start_time = 0 + shift if rs._t_start is None else rs._t_start + shift rs._t_start = new_start_time + self._time_info_modified = True def sample_index_to_time(self, sample_ind, segment_index=None): """ @@ -361,7 +364,7 @@ def get_total_duration(self) -> float: duration = sum([self.get_duration(segment_index) for segment_index in range(self.get_num_segments())]) return duration - def _get_t_starts(self): + def get_segment_t_starts(self): # handle t_starts t_starts = [] for rs in self.segments: @@ -372,14 +375,11 @@ def _get_t_starts(self): t_starts = None return t_starts - def _get_time_vectors(self): - time_vectors = [] + def has_any_time_vector(self): for rs in self.segments: - d = rs.get_times_kwargs() - time_vectors.append(d["time_vector"]) - if all(time_vector is None for time_vector in time_vectors): - time_vectors = None - return time_vectors + if rs.has_time_vector(): + return True + return False def _searchsorted_right_lazy(time_vector: TimeVector, time_s: float | np.ndarray) -> np.int64 | np.ndarray: @@ -521,6 +521,17 @@ def time_to_sample_index(self, time_s): return sample_index + def has_time_vector(self) -> bool: + """ + Returns whether the segment has a time vector. + + Returns + ------- + bool + True if the segment has a time vector, False otherwise. + """ + return self._time_vector is not None + def get_num_samples(self) -> int: """Returns the number of samples in this signal segment diff --git a/src/spikeinterface/core/time_series_tools.py b/src/spikeinterface/core/time_series_tools.py index a5c64563ec..44b6925bb5 100644 --- a/src/spikeinterface/core/time_series_tools.py +++ b/src/spikeinterface/core/time_series_tools.py @@ -70,12 +70,20 @@ def write_binary( file_path_dict = {segment_index: file_path for segment_index, file_path in enumerate(file_path_list)} if file_timestamps_paths is not None: - file_timestamps_path_dict = { - segment_index: file_path for segment_index, file_path in enumerate(file_timestamps_paths) - } + file_timestamps_path_list = ( + [file_timestamps_paths] if not isinstance(file_timestamps_paths, list) else file_timestamps_paths + ) + if len(file_timestamps_path_list) != num_segments: + raise ValueError( + "'file_timestamps_paths' must be a list of the same size as the number of segments in the time_series" + ) else: - file_timestamps_path_dict = None - for segment_index, file_path in file_path_dict.items(): + file_timestamps_path_list = [None] * num_segments + + file_path_dict = {} + file_timestamps_path_dict = {} + for segment_index, file_path in enumerate(file_path_list): + file_path_dict[segment_index] = file_path num_samples = time_series.get_num_samples(segment_index=segment_index) data_size_bytes = sample_size_bytes * num_samples file_size_bytes = data_size_bytes + byte_offset @@ -86,14 +94,13 @@ def write_binary( file.seek(file_size_bytes - 1) file.write(b"\0") - if file_timestamps_path_dict is not None: - file_timestamps_path = file_timestamps_path_dict[segment_index] + file_timestamps_path = file_timestamps_path_list[segment_index] + if file_timestamps_path is not None and time_series.has_time_vector(segment_index=segment_index): + file_timestamps_path_dict[segment_index] = file_timestamps_path with open(file_timestamps_path, "wb+") as file: file.seek(num_samples * 8 - 1) # 8 bytes for float64 timestamps file.write(b"\0") - assert Path(file_path).is_file() - # use executor (loop or workers) func = _write_binary_chunk init_func = _init_binary_worker @@ -114,13 +121,22 @@ def _init_binary_worker(time_series, file_path_dict, dtype, byte_offset, file_ti file_dict = {segment_index: open(file_path, "rb+") for segment_index, file_path in file_path_dict.items()} worker_ctx["file_dict"] = file_dict - worker_ctx["file_timestamps_dict"] = file_timestamps_path_dict + if file_timestamps_path_dict is not None: + file_timestamps_dict = { + segment_index: open(file_timestamps_path, "rb+") + for segment_index, file_timestamps_path in file_timestamps_path_dict.items() + } + worker_ctx["file_timestamps_dict"] = file_timestamps_dict + else: + worker_ctx["file_timestamps_dict"] = None return worker_ctx # used by write_binary + TimeSeriesChunkExecutor def _write_binary_chunk(segment_index, start_frame, end_frame, worker_ctx): + import gc + # recover variables of the worker time_series = worker_ctx["time_series"] dtype = worker_ctx["dtype"] @@ -139,15 +155,26 @@ def _write_binary_chunk(segment_index, start_frame, end_frame, worker_ctx): file.write(data.data) # flush is important!! file.flush() + del data if file_timestamps_dict is not None: - file_timestamps = file_timestamps_dict[segment_index] - timestamps = time_series.get_times(start_frame=start_frame, end_frame=end_frame, segment_index=segment_index) - timestamps = timestamps.astype("float64", order="c", copy=False) - timestamp_byte_offset = start_frame * 8 # 8 bytes for float64 - file.seek(timestamp_byte_offset) - file.write(timestamps.data) - file.flush() + # Some segments might not have timestamps to save + if segment_index in file_timestamps_dict: + file_timestamps = file_timestamps_dict[segment_index] + timestamps = time_series.get_times( + start_frame=start_frame, end_frame=end_frame, segment_index=segment_index + ) + timestamps = timestamps.astype("float64", order="c", copy=False) + timestamp_byte_offset = start_frame * 8 # 8 bytes for float64 + file_timestamps.seek(timestamp_byte_offset) + file_timestamps.write(timestamps.data) + file_timestamps.flush() + del timestamps + + # fix memory leak by forcing garbage collection (same issue as _write_zarr_chunk, + # e.g. reading compressed zarr chunks leaves reference cycles that the generational + # GC doesn't clear promptly in a tight chunk loop) + gc.collect() write_binary.__doc__ = write_binary.__doc__.format(_shared_job_kwargs_doc) diff --git a/src/spikeinterface/preprocessing/basepreprocessor.py b/src/spikeinterface/preprocessing/basepreprocessor.py index 64d57d3637..d28de0b34f 100644 --- a/src/spikeinterface/preprocessing/basepreprocessor.py +++ b/src/spikeinterface/preprocessing/basepreprocessor.py @@ -34,3 +34,26 @@ def get_num_samples(self): def get_traces(self, start_frame, end_frame, channel_indices): raise NotImplementedError + + # Preprocessors never change the frame numbering (no offset, no resampling), so time + # handling is a pure pass-through to the parent segment. Delegating live (instead of + # relying on the time_vector/t_start copied into __init__ above) lets any lazy/offset-aware + # override further up the chain (e.g. FrameSliceRecordingSegment after a frame_slice) keep + # working without materializing a full time_vector every time this segment is reconstructed. + def get_times(self, start_frame=None, end_frame=None): + return self.parent_recording_segment.get_times(start_frame=start_frame, end_frame=end_frame) + + def get_start_time(self): + return self.parent_recording_segment.get_start_time() + + def get_end_time(self): + return self.parent_recording_segment.get_end_time() + + def sample_index_to_time(self, sample_ind): + return self.parent_recording_segment.sample_index_to_time(sample_ind) + + def time_to_sample_index(self, time_s): + return self.parent_recording_segment.time_to_sample_index(time_s) + + def get_times_kwargs(self): + return self.parent_recording_segment.get_times_kwargs() From 4909ec1ce6d01593eedee5a4c353c895a3b138ec Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 16 Sep 2026 11:35:15 +0200 Subject: [PATCH 13/18] revert numpy_folder change --- src/spikeinterface/core/base.py | 4 ++-- src/spikeinterface/core/basesorting.py | 23 ++++++------------- src/spikeinterface/core/sortingfolder.py | 2 +- .../core/tests/test_basesorting.py | 2 +- 4 files changed, 11 insertions(+), 20 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index 76a50d27dd..e0ff3e7370 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -855,12 +855,12 @@ def save_to_folder( The 'new' way is : * recording.save(format='binary', folder=...) - * sorting.save(format='binary', folder=...) + * sorting.save(format='numpy_folder', folder=...) """ warnings.warn( "save_to_folder() should be recording.save(format='binary') " - "or sorting.save(format='binary') " + "or sorting.save(format='numpy_folder') " "This ambiguous method should not be used anymore!!", FutureWarning, ) diff --git a/src/spikeinterface/core/basesorting.py b/src/spikeinterface/core/basesorting.py index 3414c81987..57fdf9417b 100644 --- a/src/spikeinterface/core/basesorting.py +++ b/src/spikeinterface/core/basesorting.py @@ -502,20 +502,20 @@ def get_times( else: return None - def save(self, format="binary", **save_kwargs): + def save(self, format="numpy_folder", **save_kwargs): """ Save a `BaseSorting` object to a specified format: - * "binary" - old "numpy_folder" + * "numpy_folder" * "zarr" * "memory" * "npz_folder" - deprecated Parameters ---------- - format : str, default: "binary" + format : str, default: "numpy_folder" The format to save the sorting in. Options are: - - "binary"/"numpy_folder": Saves the sorting in a binary numpy folder format. + - "numpy_folder": Saves the sorting in a binary numpy folder format. - "zarr": Saves the sorting in Zarr format. - "memory": Saves the sorting in memory (shared memory or numpy array). - "npz_folder": Saves the sorting in a deprecated npz folder format. @@ -524,9 +524,9 @@ def save(self, format="binary", **save_kwargs): **save_kwargs : dict Additional keyword arguments specific to the chosen format. - * "binary" format: + * "numpy_folder" format: - folder : str or Path - The folder where the binary files will be saved. + The folder where the files will be saved. - overwrite : bool, default: False If True, existing files in the folder will be overwritten. * "zarr" format: @@ -550,19 +550,10 @@ def save(self, format="binary", **save_kwargs): The saved sorting object in the specified format. """ if format == "numpy_folder": - warnings.warn( - "The 'numpy_folder' is renamed to 'binary' and will be removed in 0.106.0. " - "Please use 'binary' instead.", - FutureWarning, - stacklevel=2, - ) - format = "binary" - - if format == "binary": from .sortingfolder import NumpyFolderSorting if "folder" not in save_kwargs: - raise ValueError("For 'binary' format, 'folder' must be specified in save_kwargs.") + raise ValueError("For 'numpy_folder' format, 'folder' must be specified in save_kwargs.") folder = save_kwargs.pop("folder") cached = NumpyFolderSorting.write_sorting(self, folder_path=folder, **save_kwargs) diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index 7916f24b2e..79098b6dcc 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -28,7 +28,7 @@ class NumpyFolderSorting(BaseSorting): * a "numpysorting_info.json" containing sampling_frequency, unit_ids and num_segments * a metadata folder for units properties. - It is created with the function: `sorting.save(folder="/myfolder", format="binary")` + It is created with the function: `sorting.save(folder="/myfolder", format="numpy_folder")` """ diff --git a/src/spikeinterface/core/tests/test_basesorting.py b/src/spikeinterface/core/tests/test_basesorting.py index 4e72e80877..bdcc0dbd3e 100644 --- a/src/spikeinterface/core/tests/test_basesorting.py +++ b/src/spikeinterface/core/tests/test_basesorting.py @@ -78,7 +78,7 @@ def test_BaseSorting(create_cache_folder): # cache new format : binary folder = cache_folder / "simple_sorting_binary" sorting.set_property("test", np.ones(len(sorting.unit_ids))) - sorting.save(folder=folder, format="binary") + sorting.save(folder=folder, format="numpy_folder") sorting2 = load(folder) assert isinstance(sorting2, NumpyFolderSorting) From 51c6b513abb3f63d59aa467b097982e0158e4317 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 16 Sep 2026 11:46:42 +0200 Subject: [PATCH 14/18] fix: make TimeSeries a class to instantiate self._time_info_modified --- src/spikeinterface/core/baserecording.py | 1 + src/spikeinterface/core/tests/test_loading.py | 2 +- src/spikeinterface/core/time_series.py | 6 ++++-- 3 files changed, 6 insertions(+), 3 deletions(-) diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 1b30c75b06..3ddefb53b8 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -41,6 +41,7 @@ def __init__(self, sampling_frequency: float, channel_ids: list, dtype): BaseRecordingSnippets.__init__( self, channel_ids=channel_ids, sampling_frequency=sampling_frequency, dtype=dtype ) + TimeSeries.__init__(self) # initialize main annotation and properties self.annotate(is_filtered=False) diff --git a/src/spikeinterface/core/tests/test_loading.py b/src/spikeinterface/core/tests/test_loading.py index 9f1e5c4578..257895fcc2 100644 --- a/src/spikeinterface/core/tests/test_loading.py +++ b/src/spikeinterface/core/tests/test_loading.py @@ -102,7 +102,7 @@ def test_load_binary_recording(generate_recording_sorting, tmp_path, output_form check_recordings_equal(rec, rec_loaded) -@pytest.mark.parametrize("output_format", ["binary", "zarr"]) +@pytest.mark.parametrize("output_format", ["numpy_folder", "zarr"]) def test_load_binary_sorting(generate_recording_sorting, tmp_path, output_format): _, sort = generate_recording_sorting _ = sort.save(folder=tmp_path / "test_sorting", format=output_format, overwrite=True) diff --git a/src/spikeinterface/core/time_series.py b/src/spikeinterface/core/time_series.py index 0d6cbc6e42..b3dfe89857 100644 --- a/src/spikeinterface/core/time_series.py +++ b/src/spikeinterface/core/time_series.py @@ -33,8 +33,10 @@ class TimeSeries(ABC): """ _preferred_mp_context = None - # Flag to indicate whether time info has been modified in-memory (e.g. by set_times or shift_times). - _time_info_modified = False + + def __init__(self): + # Flag to indicate whether time info has been modified in-memory (e.g. by set_times or shift_times). + self._time_info_modified = False def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) From a2301eccbfd4b5dfed6402e08c49e7044c6dc729 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 16 Sep 2026 11:59:44 +0200 Subject: [PATCH 15/18] refactor: rename save property functions and fix docstrings --- src/spikeinterface/core/base.py | 2 +- src/spikeinterface/core/binaryfolder.py | 8 ++-- src/spikeinterface/core/core_tools.py | 14 +++--- .../core/frameslicerecording.py | 47 ++++++++++++++++++- src/spikeinterface/core/npyfoldersnippets.py | 8 ++-- src/spikeinterface/core/sortingfolder.py | 12 ++--- 6 files changed, 68 insertions(+), 23 deletions(-) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index e0ff3e7370..7fa7fb24c4 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -24,7 +24,7 @@ make_paths_absolute, check_paths_relative, retrieve_importing_provenance, - load_properties_from_binary_folder, + load_properties_from_folder, ) from .job_tools import _shared_job_kwargs_doc diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index c35b722557..c0090719ca 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -13,8 +13,8 @@ from .core_tools import ( define_function_from_class, make_paths_absolute, - load_properties_from_binary_folder, - save_properties_to_binary_folder, + load_properties_from_folder, + save_properties_to_folder, save_extractor_provenance, save_annotations_to_folder, load_annotations_from_folder, @@ -54,7 +54,7 @@ def __init__(self, folder_path): BinaryRecordingExtractor.__init__(self, **d["kwargs"]) # Load properties - load_properties_from_binary_folder(folder_path / "properties", self) + load_properties_from_folder(folder_path / "properties", self) load_annotations_from_folder(folder_path, self) # Load the probegroup @@ -159,7 +159,7 @@ def write_recording( **job_kwargs, ) - save_properties_to_binary_folder(folder_path / "properties", recording) + save_properties_to_folder(folder_path / "properties", recording) save_extractor_provenance(folder_path, recording) # new in version 0.105.0, before that annotations were handle by "si_folder.json" file save_annotations_to_folder(folder_path, recording) diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index 5c3a948609..083f9738b9 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -813,7 +813,7 @@ def slice_rows(array: np.ndarray | zarr.Array, row_indices: np.ndarray | list) - return array[row_indices, ...] -def load_properties_from_binary_folder(folder: str | Path, extractor: "BaseExtractor") -> dict: +def load_properties_from_folder(folder: str | Path, extractor: "BaseExtractor") -> dict: """ Load properties from a folder properties as .npy files and return sets them as properties to the extractor. @@ -836,7 +836,7 @@ def load_properties_from_binary_folder(folder: str | Path, extractor: "BaseExtra extractor.set_property(key, values) -def save_properties_to_binary_folder(folder: str | Path, extractor: "BaseExtractor"): +def save_properties_to_folder(folder: str | Path, extractor: "BaseExtractor"): """ Save properties from an extractor to a folder as .npy files. @@ -856,8 +856,8 @@ def save_properties_to_binary_folder(folder: str | Path, extractor: "BaseExtract def save_annotations_to_folder(folder: str | Path, extractor: "BaseExtractor"): """ - Save annotations in json format from an extractor (recording or sorting). - This used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0 + Save BaseExtractor annotations to annotations.json in the provided folder. + This is used for `BinaryFolderRecording` and `NumpyFolderSorting` since version 0.105.0. Parameters ---------- @@ -867,14 +867,14 @@ def save_annotations_to_folder(folder: str | Path, extractor: "BaseExtractor"): The extractor from which the annotations will be saved. """ folder = Path(folder) - print(f"Saving annotations to {folder / 'annotations.json'}: {extractor._annotations}") (folder / "annotations.json").write_text(json.dumps(extractor._annotations, indent=4), encoding="utf8") def load_annotations_from_folder(folder: str | Path, extractor: "BaseExtractor"): """ - Load annotations in json format from an extractor (recording or sorting). - This used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0 + Load annotations from annotations.json in the provided folder. + This is used for BinaryFolderRecording and NumpyFolderSorting since version 0.105.0. + If the file doesn't exist, it will try to load annotations from the si_folder.json "annotations" field. Parameters ---------- diff --git a/src/spikeinterface/core/frameslicerecording.py b/src/spikeinterface/core/frameslicerecording.py index a3136583db..bc7913e505 100644 --- a/src/spikeinterface/core/frameslicerecording.py +++ b/src/spikeinterface/core/frameslicerecording.py @@ -81,10 +81,13 @@ def __init__(self, parent_recording_segment, start_frame, end_frame): d = d.copy() if d["time_vector"] is None: d["t_start"] = parent_recording_segment.sample_index_to_time(start_frame) + parent_has_time_vector = False else: - d["time_vector"] = d["time_vector"][start_frame:end_frame] + d["time_vector"] = None + parent_has_time_vector = True BaseRecordingSegment.__init__(self, **d) self._parent_recording_segment = parent_recording_segment + self._parent_has_time_vector = parent_has_time_vector self.start_frame = start_frame self.end_frame = end_frame @@ -98,3 +101,45 @@ def get_traces(self, start_frame, end_frame, channel_indices): start_frame=parent_start, end_frame=parent_end, channel_indices=channel_indices ) return traces + + # Times methods below mirror get_traces(): defer to the parent segment with an offset + # instead of slicing self._time_vector directly. Slicing here would force reading the + # whole window right away (e.g. a zarr.Array fetch+decompress), and since this segment + # is rebuilt from scratch once per worker process in a multiprocessing job, that cost + # would be paid again in every worker instead of being read lazily per chunk. + def get_times(self, start_frame=None, end_frame=None): + if self._parent_has_time_vector: + start_frame = int(start_frame) if start_frame is not None else 0 + end_frame = int(end_frame) if end_frame is not None else self.get_num_samples() + return self._parent_recording_segment.get_times( + start_frame=self.start_frame + start_frame, end_frame=self.start_frame + end_frame + ) + else: + return super().get_times(start_frame=start_frame, end_frame=end_frame) + + def get_start_time(self) -> float: + if self._parent_has_time_vector: + return self._parent_recording_segment.sample_index_to_time(self.start_frame) + else: + return super().get_start_time() + + def get_end_time(self) -> float: + if self._parent_has_time_vector: + return self._parent_recording_segment.sample_index_to_time(self.end_frame - 1) + else: + return super().get_end_time() + + def sample_index_to_time(self, sample_ind): + if self._parent_has_time_vector: + return self._parent_recording_segment.sample_index_to_time(self.start_frame + sample_ind) + else: + return super().sample_index_to_time(sample_ind) + + def time_to_sample_index(self, time_s): + if self._parent_has_time_vector: + return self._parent_recording_segment.time_to_sample_index(time_s) - self.start_frame + else: + return super().time_to_sample_index(time_s) + + def has_time_vector(self) -> bool: + return self._parent_has_time_vector diff --git a/src/spikeinterface/core/npyfoldersnippets.py b/src/spikeinterface/core/npyfoldersnippets.py index 9e026904ea..31dee04078 100644 --- a/src/spikeinterface/core/npyfoldersnippets.py +++ b/src/spikeinterface/core/npyfoldersnippets.py @@ -9,8 +9,8 @@ from .core_tools import ( define_function_from_class, make_paths_absolute, - load_properties_from_binary_folder, - save_properties_to_binary_folder, + load_properties_from_folder, + save_properties_to_folder, ) @@ -54,7 +54,7 @@ def __init__(self, folder_path): if probe_file.is_file(): self._probegroup = read_probeinterface(probe_file) - load_properties_from_binary_folder(folder_path / "properties", self) + load_properties_from_folder(folder_path / "properties", self) self._kwargs = dict(folder_path=str(Path(folder_path).absolute())) self._bin_kwargs = d["kwargs"] @@ -85,7 +85,7 @@ def write_snippets(snippets, folder, dtype=None): ) cached.dump(folder / "npy.json", relative_to=folder) - save_properties_to_binary_folder(folder / "properties", snippets) + save_properties_to_folder(folder / "properties", snippets) if snippets.has_probe(): probegroup = snippets.get_probegroup() diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index 79098b6dcc..3a068e0384 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -11,8 +11,8 @@ from .core_tools import ( define_function_from_class, make_paths_absolute, - load_properties_from_binary_folder, - save_properties_to_binary_folder, + load_properties_from_folder, + save_properties_to_folder, save_annotations_to_folder, load_annotations_from_folder, save_extractor_provenance, @@ -55,7 +55,7 @@ def __init__(self, folder_path, mmap_mode: str | None = None): # important trick : the cache is already spikes vector self._cached_spike_vector = self.spikes - load_properties_from_binary_folder(folder_path / "properties", self) + load_properties_from_folder(folder_path / "properties", self) load_annotations_from_folder(folder_path, self) self._kwargs = dict(folder_path=str(folder_path.absolute()), mmap_mode=mmap_mode) @@ -80,7 +80,7 @@ def write_sorting(sorting, folder_path, overwrite: bool = False): info_file.write_text(json.dumps(d), encoding="utf8") np.save(folder_path / "spikes.npy", sorting.to_spike_vector()) - save_properties_to_binary_folder(folder_path / "properties", sorting) + save_properties_to_folder(folder_path / "properties", sorting) save_extractor_provenance(folder_path, sorting) # new in version 0.105.0, before that annotations were handle by "si_folder.json" file save_annotations_to_folder(folder_path, sorting) @@ -141,7 +141,7 @@ def __init__(self, folder_path): NpzSortingExtractor.__init__(self, **d["kwargs"]) - load_properties_from_binary_folder(folder_path / "properties", self) + load_properties_from_folder(folder_path / "properties", self) load_annotations_from_folder(folder_path, self) self._kwargs = dict(folder_path=str(folder_path.absolute())) @@ -162,7 +162,7 @@ def write_sorting(sorting, save_path): if npz_file.exists(): raise ValueError("NpzFolderSorting.write_sorting the folder already contains sorting_cached.npz") NpzSortingExtractor.write_sorting(sorting, npz_file) - save_properties_to_binary_folder(save_path / "properties", sorting) + save_properties_to_folder(save_path / "properties", sorting) cached = NpzSortingExtractor(npz_file) cached.dump(save_path / "npz.json", relative_to=save_path) From 443f80a7e3be836528489ac3e7c0e58e733e737d Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 16 Sep 2026 12:01:23 +0200 Subject: [PATCH 16/18] feat: add main_ids to dump_dict --- src/spikeinterface/core/base.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index 7fa7fb24c4..8dda3817a6 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -576,6 +576,7 @@ def to_dict( dump_dict = retrieve_importing_provenance(self.__class__) dump_dict["kwargs"] = kwargs + dump_dict["main_ids"] = self._main_ids if include_annotations: dump_dict["annotations"] = self._annotations From f0b76371d3ceb517bc5156f3cf219d12a34273d1 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 16 Sep 2026 12:20:31 +0200 Subject: [PATCH 17/18] fix: benchmark base --- src/spikeinterface/benchmark/benchmark_base.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/spikeinterface/benchmark/benchmark_base.py b/src/spikeinterface/benchmark/benchmark_base.py index 2bd55791a7..47ec6a70cd 100644 --- a/src/spikeinterface/benchmark/benchmark_base.py +++ b/src/spikeinterface/benchmark/benchmark_base.py @@ -190,7 +190,7 @@ def create(cls, study_folder, datasets={}, cases={}, levels=None): # sortings are pickled + saved as NumpyFolderSorting # gt_sorting.dump_to_pickle(study_folder / f"datasets/gt_sortings/{key}.pickle") - # gt_sorting.save(format="binary", folder=study_folder / f"datasets/gt_sortings/{key}") + # gt_sorting.save(format="numpy_folder", folder=study_folder / f"datasets/gt_sortings/{key}") # analyzer path (local or external) (study_folder / "analyzers_path.json").write_text(json.dumps(analyzers_path, indent=4), encoding="utf8") @@ -651,7 +651,7 @@ def _save_keys(self, saved_keys, folder): with open(folder / f"{k}.pickle", mode="wb") as f: pickle.dump(self.result[k], f) elif format == "sorting": - self.result[k].save(folder=folder / k, format="binary", overwrite=True) + self.result[k].save(folder=folder / k, format="numpy_folder", overwrite=True) elif format == "Motion": self.result[k].save(folder=folder / k) elif format == "zarr_templates": From 874542ea8b217faa4c8d8fac8301fc244d99b4e0 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 16 Sep 2026 12:22:36 +0200 Subject: [PATCH 18/18] fix: resample test --- src/spikeinterface/preprocessing/resample.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/spikeinterface/preprocessing/resample.py b/src/spikeinterface/preprocessing/resample.py index 737454ce00..74a044d368 100644 --- a/src/spikeinterface/preprocessing/resample.py +++ b/src/spikeinterface/preprocessing/resample.py @@ -149,8 +149,8 @@ def __init__( # Compute time_vector or t_start, following the pattern from DecimateRecordingSegment. # Do not use BasePreprocessorSegment because we have to reset the sampling rate! - if parent_recording_segment._time_vector is not None: - parent_tv = np.asarray(parent_recording_segment._time_vector) + if parent_recording_segment.has_time_vector(): + parent_tv = np.asarray(parent_recording_segment.get_times()) # Detect gaps in the parent time vector. # A true gap means at least one dropped sample, so dt >= 2 * expected_dt.