Skip to content

Commit 67a8d2d

Browse files
committed
refactor: cleanup save functions and add detailed docs
1 parent a890482 commit 67a8d2d

9 files changed

Lines changed: 260 additions & 331 deletions

File tree

src/spikeinterface/core/base.py

Lines changed: 2 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -855,12 +855,6 @@ def save_to_memory(self, sharedmem=True, **save_kwargs) -> "BaseExtractor":
855855
warnings.warn("save_to_memory() should be save(format='memory')", FutureWarning)
856856
return self.save(format="memory", sharedmem=sharedmem, **save_kwargs)
857857

858-
# save_kwargs.pop("format", None)
859-
860-
# cached = self._save(format="memory", sharedmem=sharedmem, **save_kwargs)
861-
# self.copy_metadata(cached)
862-
# return cached
863-
864858
def save(self):
865859
# Need to be implemented in Recording and Sorting
866860
raise NotImplementedError()
@@ -878,64 +872,18 @@ def save_to_folder(
878872
879873
The 'new' way is :
880874
* recording.save(format='binary', folder=...)
881-
* sorting.save(format='numpy_folder', folder=...)
875+
* sorting.save(format='binary', folder=...)
882876
"""
883877

884878
warnings.warn(
885879
"save_to_folder() should be recording.save(format='binary') "
886-
"or sorting.save(format='numpy_folder') "
880+
"or sorting.save(format='binary') "
887881
"This ambiguous method should not be used anymore!!",
888882
FutureWarning,
889883
)
890884
# we keep the default format for recording and sorting like in old version
891885
return self.save(folder=folder, verbose=verbose, **save_kwargs)
892886

893-
# if folder is None:
894-
# cache_folder = get_global_tmp_folder()
895-
# if name is None:
896-
# name = "".join(random.choices(string.ascii_uppercase + string.digits, k=8))
897-
# folder = cache_folder / name
898-
# if verbose:
899-
# print(f"Use cache_folder={folder}")
900-
# else:
901-
# folder = cache_folder / name
902-
# if not is_set_global_tmp_folder():
903-
# if verbose:
904-
# print(f"Use cache_folder={folder}")
905-
# else:
906-
# folder = Path(folder)
907-
# if overwrite and folder.is_dir():
908-
# import shutil
909-
910-
# shutil.rmtree(folder)
911-
912-
# assert not folder.exists(), f"folder {folder} already exists, choose another name or use overwrite=True"
913-
# folder.mkdir(parents=True, exist_ok=False)
914-
915-
# # dump provenance
916-
# provenance_file_path = folder / f"provenance.json"
917-
# if self.check_serializability("json"):
918-
# self.dump_to_json(file_path=provenance_file_path, relative_to=folder)
919-
# elif self.check_serializability("pickle"):
920-
# provenance_file = folder / f"provenance.pkl"
921-
# self.dump_to_pickle(provenance_file, relative_to=folder)
922-
# else:
923-
# warnings.warn("The extractor is not serializable to file. The provenance will not be saved.")
924-
925-
# # save data (done the subclass)
926-
# self.save_metadata_to_folder(folder)
927-
# cached = self._save(folder=folder, verbose=verbose, **save_kwargs)
928-
# cached.load_metadata_from_folder(folder)
929-
930-
# # copy properties/
931-
# self.copy_metadata(cached)
932-
933-
# # Dump the extractor to json file
934-
# si_folder_path = folder / f"si_folder.json"
935-
# cached.dump_to_json(file_path=si_folder_path, relative_to=folder)
936-
937-
# return cached
938-
939887
def save_to_zarr(
940888
self,
941889
name=None,

src/spikeinterface/core/baserecording.py

Lines changed: 66 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -320,20 +320,79 @@ def get_shape(self, segment_index: int | None = None) -> tuple[int, ...]:
320320

321321
def save(self, format="binary", verbose: bool = False, **save_kwargs):
322322
"""
323-
TODO: each object.save should have extensive docstring with all the options and examples
323+
Save a `BaseRecording` object to a specified format:
324+
325+
* "binary"
326+
* "zarr"
327+
* "memory"
328+
329+
Parameters
330+
----------
331+
format : str, default: "binary"
332+
The format to save the recording in. Options are:
333+
- "binary": Saves the recording in binary format.
334+
- "zarr": Saves the recording in Zarr format.
335+
- "memory": Saves the recording in memory (shared memory or numpy array).
336+
verbose : bool, default: False
337+
If True, prints additional information during the save process.
338+
**save_kwargs : dict
339+
Additional keyword arguments specific to the chosen format.
340+
All formats support job_kwargs for parallel processing
341+
(see `si.get_global_job_kwargs()` for default values).
342+
343+
* "binary" format:
344+
- folder : str or Path
345+
The folder where the binary files will be saved.
346+
- overwrite : bool, default: False
347+
If True, existing files in the folder will be overwritten.
348+
- dtype : str, optional
349+
The data type to use for saving the recording. If not provided, the recording's dtype
350+
will be used.
351+
* "zarr" format:
352+
- folder : str or Path
353+
The folder where the Zarr files will be saved.
354+
- overwrite: bool, default: False
355+
If True, the folder is removed if it already exists
356+
- storage_options: dict or None, default: None
357+
Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc.
358+
For cloud storage locations, this should not be None (in case of default values, use an empty dict)
359+
- channel_chunk_size: int or None, default: None
360+
Channels per chunk (only for BaseRecording)
361+
- compressor: numcodecs.Codec or None, default: None
362+
Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used
363+
- filters: list[numcodecs.Codec] or None, default: None
364+
Global filters for zarr (global)
365+
- compressor_by_dataset: dict or None, default: None
366+
Optional compressor per dataset:
367+
- traces
368+
- times
369+
If None, the global compressor is used
370+
- filters_by_dataset: dict or None, default: None
371+
Optional filters per dataset:
372+
- traces
373+
- times
374+
If None, the global filters are used
375+
* "memory" format:
376+
- sharedmem : bool, default: True
377+
If True, the recording is saved in shared memory. If False, it is saved as
378+
a numpy array in memory.
379+
380+
Returns
381+
-------
382+
BaseRecording
383+
The saved recording object in the specified format.
324384
"""
325385
kwargs, job_kwargs = split_job_kwargs(save_kwargs)
326386

327-
# TODO: add overwrite option to binary/zarr save
328-
329387
if format == "binary":
330388
if "folder" not in kwargs:
331389
raise ValueError("Missing folder in recording.save(folder='...')")
332390

333391
from .binaryfolder import BinaryFolderRecording
334392

393+
folder = kwargs.pop("folder")
335394
cached = BinaryFolderRecording.write_recording(
336-
self, folder=kwargs["folder"], dtype=kwargs.get("dtype", None), **job_kwargs
395+
self, folder_path=folder, verbose=verbose, **kwargs, **job_kwargs
337396
)
338397

339398
elif format == "memory":
@@ -348,49 +407,22 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs):
348407

349408
cached = NumpyRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs)
350409

351-
# self.copy_metadata(cached)
352-
353-
# # timestamps are not saved in memory, so we have to set them explicitly
354-
# for segment_index in range(self.get_num_segments()):
355-
# if self.has_time_vector(segment_index):
356-
# # the use of get_times is preferred since timestamps are converted to array
357-
# time_vector = self.get_times(segment_index=segment_index)
358-
# cached.set_times(time_vector, segment_index=segment_index)
359-
360410
elif format == "zarr":
361411
if "folder" not in kwargs:
362412
raise ValueError("Missing folder in recording.save(folder='...')")
413+
folder_path = kwargs.pop("folder")
363414

364415
from .zarrextractors import ZarrRecordingExtractor
365416

366-
folder_path = kwargs["folder"]
367-
if isinstance(folder_path, Path) and folder_path.suffix != "zarr":
368-
# automatically add the zarr suffix
369-
folder_path = folder_path.with_suffix(".zarr")
370-
371-
storage_options = kwargs.pop("storage_options", None)
372-
ZarrRecordingExtractor.write_recording(
373-
self, folder_path, storage_options, verbose=verbose, **kwargs, **job_kwargs
417+
cached = ZarrRecordingExtractor.write_recording(
418+
self, folder_path=folder_path, verbose=verbose, **kwargs, **job_kwargs
374419
)
375-
cached = ZarrRecordingExtractor(folder_path, storage_options)
376-
# timestamps are saved and restored in zarr, so no need to set them explicitly
377420

378421
else:
379422
raise ValueError(f"format {format} not supported")
380423

381424
return cached
382425

383-
def _extra_metadata_from_folder(self, folder):
384-
# load probe
385-
super()._extra_metadata_from_folder(folder)
386-
387-
# load time vector if any
388-
for segment_index, rs in enumerate(self.segments):
389-
time_file = folder / f"times_cached_seg{segment_index}.npy"
390-
if time_file.is_file():
391-
time_vector = np.load(time_file, mmap_mode="r")
392-
rs._time_vector = time_vector
393-
394426
def select_channels(self, channel_ids: list | np.ndarray | tuple) -> "BaseRecording":
395427
"""
396428
Returns a new recording object with a subset of channels.

src/spikeinterface/core/baserecordingsnippets.py

Lines changed: 0 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -320,31 +320,6 @@ def _extra_metadata_copy(self, other):
320320
if self._probegroup is not None:
321321
other._probegroup = self._probegroup.copy()
322322

323-
def _extra_metadata_from_folder(self, folder):
324-
# load probe from folder
325-
# Note: we don't need any fix for legacy probegroups, since the
326-
# set_probegroup() method will handle the device_channel_indices
327-
# sorting and global contact order
328-
folder = Path(folder)
329-
probe_file = folder / "probegroup.json"
330-
legacy_probe_file = folder / "probe.json"
331-
if probe_file.is_file():
332-
probegroup = read_probeinterface(probe_file)
333-
self.set_probegroup(probegroup)
334-
elif legacy_probe_file.is_file():
335-
probegroup = read_probeinterface(legacy_probe_file)
336-
self.set_probegroup(probegroup)
337-
338-
# remove "contact_vector" property if present as it is not needed anymore
339-
if "contact_vector" in self.get_property_keys():
340-
self.delete_property("contact_vector")
341-
342-
# def _extra_metadata_to_folder(self, folder):
343-
# # save probe
344-
# if self.has_probe():
345-
# probegroup = self.get_probegroup()
346-
# write_probeinterface(folder / "probegroup.json", probegroup)
347-
348323
def _extra_metadata_from_dict(self, dump_dict):
349324
# load probe and handle backward-compatibility with legacy "contact_vector"/"location" property
350325
if "probegroup" in dump_dict:

src/spikeinterface/core/basesnippets.py

Lines changed: 15 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,12 @@
1-
from .base import BaseSegment
2-
from .baserecordingsnippets import BaseRecordingSnippets
31
import numpy as np
42
from warnings import warn
53

64
from copy import deepcopy
75

86
from pathlib import Path
97

10-
from .core_tools import save_properties_to_binary_folder
11-
12-
# snippets segments?
8+
from .base import BaseSegment
9+
from .baserecordingsnippets import BaseRecordingSnippets
1310

1411

1512
class BaseSnippets(BaseRecordingSnippets):
@@ -21,6 +18,12 @@ class BaseSnippets(BaseRecordingSnippets):
2118
_main_features = []
2219

2320
def __init__(self, sampling_frequency: float, nbefore: int | None, snippet_len: int, channel_ids: list, dtype):
21+
warn(
22+
"`BaseSnippets` is deprecated and will be removed in version 0.106.0."
23+
"Only continuous recordings with `BaseRecording` will be supported.",
24+
FutureWarning,
25+
stacklevel=2,
26+
)
2427
BaseRecordingSnippets.__init__(
2528
self, channel_ids=channel_ids, sampling_frequency=sampling_frequency, dtype=dtype
2629
)
@@ -213,9 +216,11 @@ def _select_segments(self, segment_indices):
213216

214217
def save(self, format="npy", **save_kwargs):
215218
"""
216-
At the moment only "npy" and "memory" avaiable:
217-
"""
219+
Save a `BaseSnippets` object to a specified format:
218220
221+
* "npy"
222+
* "memory"
223+
"""
219224
if format == "npy":
220225
from spikeinterface.core.npyfoldersnippets import NpyFolderSnippets
221226

@@ -240,14 +245,12 @@ def save(self, format="npy", **save_kwargs):
240245
nbefore=self.nbefore,
241246
channel_ids=self.channel_ids,
242247
)
243-
248+
if self.has_probe():
249+
probegroup = self.get_probegroup()
250+
cached.set_probegroup(probegroup)
244251
else:
245252
raise ValueError(f"format {format} not supported")
246253

247-
if self.has_probe():
248-
probegroup = self.get_probegroup()
249-
cached.set_probegroup(probegroup)
250-
251254
return cached
252255

253256
def get_times(self):

0 commit comments

Comments
 (0)