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/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/.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/.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/comparison/tests/test_multisortingcomparison.py b/src/spikeinterface/comparison/tests/test_multisortingcomparison.py index 6fad43bde0..72f39f74ea 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/base.py b/src/spikeinterface/core/base.py index 3c03d8a861..8dda3817a6 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -24,6 +24,7 @@ make_paths_absolute, check_paths_relative, retrieve_importing_provenance, + load_properties_from_folder, ) from .job_tools import _shared_job_kwargs_doc @@ -485,8 +486,8 @@ 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, ) -> dict: """ @@ -514,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. @@ -537,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. @@ -560,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, ) @@ -581,12 +576,15 @@ 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 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 @@ -609,12 +607,8 @@ 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) - - self._extra_metadata_to_dict(dump_dict) + if include_extra_metadata: + self._extra_metadata_to_dict(dump_dict) return dump_dict @@ -640,38 +634,8 @@ def from_dict(dictionary: dict, base_folder: Path | str | None = None) -> "BaseE assert base_folder is not None, "When relative_paths=True, need to provide base_folder" dictionary = make_paths_absolute(dictionary, base_folder) extractor = _load_extractor_from_dict(dictionary) - 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 - extractor.load_metadata_from_folder(folder_metadata) - return extractor - - def load_metadata_from_folder(self, folder_metadata: str | Path): - # hack to load probe for recording - folder_metadata = Path(folder_metadata) - # load properties - prop_folder = folder_metadata / "properties" - if prop_folder.is_dir(): - for prop_file in prop_folder.iterdir(): - if prop_file.suffix == ".npy": - values = np.load(prop_file, allow_pickle=True) - key = prop_file.stem - self.set_property(key, values) - - self._extra_metadata_from_folder(folder_metadata) - - def save_metadata_to_folder(self, folder_metadata: str | Path): - self._extra_metadata_to_folder(folder_metadata) - - # save properties - prop_folder = Path(folder_metadata) / "properties" - prop_folder.mkdir(parents=True, exist_ok=False) - for key in self.get_property_keys(): - values = self.get_property(key) - np.save(prop_folder / (key + ".npy"), values) + return extractor def clone(self) -> "BaseExtractor": """ @@ -731,7 +695,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 @@ -744,9 +708,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") @@ -754,7 +718,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. @@ -767,8 +733,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" @@ -778,10 +748,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"]) @@ -796,7 +766,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. @@ -811,8 +780,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" @@ -828,7 +795,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, ) @@ -859,71 +825,24 @@ def __reduce__(self): intialization_args = (self.to_dict(),) return (instance_constructor, intialization_args) - def _save(self, folder, **save_kwargs): - # This implemented in BaseRecording or baseSorting - # 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_to_folder(self, folder): - # This implemented in BaseRecording for probe - pass - def _extra_metadata_from_dict(self, dump_dict): + # Hook for subclass (quite bad design) # This implemented in BaseRecording for probe pass def _extra_metadata_to_dict(self, dump_dict): + # Hook for subclass (quite bad design) # This implemented in BaseRecording for probe pass - 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) - def save_to_memory(self, sharedmem=True, **save_kwargs) -> "BaseExtractor": - save_kwargs.pop("format", None) + warnings.warn("save_to_memory() should be save(format='memory')", FutureWarning) + return self.save(format="memory", sharedmem=sharedmem, **save_kwargs) - 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() - # TODO rename to saveto_binary_folder def save_to_folder( self, name: str | None = None, @@ -933,101 +852,21 @@ def save_to_folder( **save_kwargs, ): """ - Save the extractor and its data to a folder. - - This method extracts trace data, saves it to a file (using a memory-mapped approach), - and stores both the original extractor's provenance - and the extractor's metadata in JSON format. - - The folder's final location and name can be specified in a couple of ways ways: - - 1. Explicitly providing the full path: - ``` - extractor.save_to_folder(folder="/path/to/save/") - ``` + Legacy method. - 2. Providing a subfolder name, with the base folder being determined automatically: - ``` - extractor.save_to_folder(name="my_extractor_data") - ``` - In this case, the data is saved in a subfolder named "my_extractor_data" - within the global temporary folder (set using `set_global_tmp_folder`). If no - global temporary folder is set, one will be generated automatically. - - 3. If neither `name` nor `folder` is provided, a random name will be generated - for the subfolder within the global temporary folder. - - Parameters - ---------- - name : str or Path, optional - The name of the subfolder within the global temporary folder. If `folder` - is provided, this argument must be None. - folder : str or Path, optional - The full path of the folder where the data should be saved. If `name` is - provided, this argument must be None. - overwrite : bool, default: False - If True, an existing folder at the specified path will be deleted before saving. - verbose : bool, default: True - If True, print information about the cache folder being used. - **save_kwargs - Additional keyword arguments to be passed to the underlying save method. - - Returns - ------- - cached_extractor - A saved copy of the extractor in the specified format. - - Raises - ------ - AssertionError - If the folder already exists and `overwrite` is False. + The 'new' way is : + * recording.save(format='binary', folder=...) + * sorting.save(format='numpy_folder', folder=...) """ - 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 + warnings.warn( + "save_to_folder() should be recording.save(format='binary') " + "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) def save_to_zarr( self, @@ -1040,74 +879,91 @@ def save_to_zarr( **save_kwargs, ): """ - Save extractor to zarr. + Legacy method. - 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. + The 'new' way is : + * recording.save(format='zarr', folder=...) + * sorting.save(format='zarr', folder=...) """ - 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) - - return cached + 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) + + # return cached def _load_extractor_from_dict(dic) -> "BaseExtractor": diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 5fc7d37505..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) @@ -318,104 +319,132 @@ 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): + """ + 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) if format == "binary": - from .time_series_tools import write_binary - from .binaryrecordingextractor import BinaryRecordingExtractor + if "folder" not in kwargs: + raise ValueError("Missing folder in recording.save(folder='...')") + 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(), + folder = kwargs.pop("folder") + cached = BinaryFolderRecording.write_recording( + self, folder_path=folder, verbose=verbose, **kwargs, **job_kwargs ) - 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) 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) - - # 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) + cached = NumpyRecording.from_recording(self, with_metadata=True, with_time_vector=True, **job_kwargs) 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 - zarr_path = kwargs.pop("zarr_path") - storage_options = kwargs.pop("storage_options") - ZarrRecordingExtractor.write_recording( - self, zarr_path, storage_options, verbose=verbose, **kwargs, **job_kwargs + cached = ZarrRecordingExtractor.write_recording( + self, folder_path=folder_path, verbose=verbose, **kwargs, **job_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") 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 _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) + 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": """ diff --git a/src/spikeinterface/core/baserecordingsnippets.py b/src/spikeinterface/core/baserecordingsnippets.py index 66c6c5614b..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 bb7704d7ec..5b38c47791 100644 --- a/src/spikeinterface/core/basesnippets.py +++ b/src/spikeinterface/core/basesnippets.py @@ -1,9 +1,12 @@ -from .base import BaseSegment -from .baserecordingsnippets import BaseRecordingSnippets import numpy as np from warnings import warn -# snippets segments? +from copy import deepcopy + +from pathlib import Path + +from .base import BaseSegment +from .baserecordingsnippets import BaseRecordingSnippets class BaseSnippets(BaseRecordingSnippets): @@ -15,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 ) @@ -188,9 +197,6 @@ def get_snippets_from_frames( return self.get_snippets(indices, channel_ids=channel_ids, return_in_uV=return_in_uV) - def _save(self, format="binary", **save_kwargs): - raise NotImplementedError - def select_channels(self, channel_ids: list | np.ndarray | tuple) -> "BaseSnippets": from .channelslice import ChannelSliceSnippets @@ -208,36 +214,20 @@ def _select_segments(self, segment_indices): return SelectSegmentSnippets(self, segment_indices=segment_indices) - def _save(self, format="npy", **save_kwargs): - """ - At the moment only "npy" and "memory" avaiable: + def save(self, format="npy", **save_kwargs): """ + Save a `BaseSnippets` object to a specified format: + * "npy" + * "memory" + """ if format == "npy": - from spikeinterface.core.npysnippetsextractor import NpySnippetsExtractor + from spikeinterface.core.npyfoldersnippets import NpyFolderSnippets folder = save_kwargs["folder"] - file_paths = [folder / f"traces_cached_seg{i}.npy" for i in range(self.get_num_segments())] - dtype = save_kwargs.get("dtype", None) - if dtype is None: - dtype = self.get_dtype() - - from spikeinterface.core.npysnippetsextractor import NpySnippetsExtractor - - NpySnippetsExtractor.write_snippets(snippets=self, file_paths=file_paths, dtype=dtype) - cached = NpySnippetsExtractor( - file_paths=file_paths, - sampling_frequency=self.get_sampling_frequency(), - channel_ids=self.get_channel_ids(), - nbefore=self.nbefore, - gain_to_uV=self.get_channel_gains(), - offset_to_uV=self.get_channel_offsets(), - ) - cached.dump(folder / "npy.json", relative_to=folder) - - from spikeinterface.core.npyfoldersnippets import NpyFolderSnippets + folder = Path(folder) - cached = NpyFolderSnippets(folder_path=folder) + cached = NpyFolderSnippets.write_snippets(self, folder, dtype=save_kwargs.get("dtype", None)) elif format == "memory": snippets_list = [] @@ -255,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 ec285c5e6b..79e07d8cb4 100644 --- a/src/spikeinterface/core/basesorting.py +++ b/src/spikeinterface/core/basesorting.py @@ -4,6 +4,8 @@ import numpy as np +from pathlib import Path + from .base import BaseExtractor, BaseSegment, minimum_spike_dtype from .waveform_tools import has_exceeding_spikes @@ -507,48 +509,89 @@ def get_times( else: return None - def _save(self, format: str = "numpy_folder", **save_kwargs): - """Save a sorting object to disk in a specified format. + def save(self, format="numpy_folder", **save_kwargs): + """ + Save a `BaseSorting` object to a specified format: + + * "numpy_folder" + * "zarr" + * "memory" + * "npz_folder" - deprecated - Note - ---- - This function replaces the old CacheSortingExtractor, but enables more engines - for caching a results. + Parameters + ---------- + format : str, default: "numpy_folder" + The format to save the sorting in. Options are: + - "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. + + * "numpy_folder" format: + - folder : str or Path + The folder where the 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": from .sortingfolder import NumpyFolderSorting + if "folder" not in save_kwargs: + raise ValueError("For 'numpy_folder' 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 - zarr_path = save_kwargs.pop("zarr_path") - storage_options = save_kwargs.pop("storage_options") - ZarrSortingExtractor.write_sorting(self, zarr_path, storage_options, **save_kwargs) - cached = ZarrSortingExtractor(zarr_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 - cached = SharedMemorySorting.from_sorting(self) + cached = SharedMemorySorting.from_sorting(self, with_metadata=True) else: from .numpyextractors import NumpySorting - cached = NumpySorting.from_sorting(self) + 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 ee86c4dc7c..c0090719ca 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -1,12 +1,24 @@ from pathlib import Path import json +from copy import deepcopy +import shutil + 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 +from .core_tools import ( + define_function_from_class, + make_paths_absolute, + load_properties_from_folder, + save_properties_to_folder, + save_extractor_provenance, + save_annotations_to_folder, + load_annotations_from_folder, +) class BinaryFolderRecording(BinaryRecordingExtractor): @@ -42,15 +54,8 @@ def __init__(self, folder_path): BinaryRecordingExtractor.__init__(self, **d["kwargs"]) # Load properties - prop_folder = folder_path / "properties" - if prop_folder.is_dir(): - for prop_file in prop_folder.iterdir(): - if prop_file.suffix == ".npy": - values = np.load(prop_file, allow_pickle=True) - key = prop_file.stem - if key == "contact_vector": - continue - self.set_property(key, values) + load_properties_from_folder(folder_path / "properties", self) + load_annotations_from_folder(folder_path, self) # Load the probegroup probe_file = folder_path / "probegroup.json" @@ -112,5 +117,89 @@ def get_binary_description(self): ) return d + @staticmethod + 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 = 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() + # 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_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) + + if recording.has_probe(): + probegroup = recording.get_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__ + 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, + 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_path / "binary.json", relative_to=folder_path) + + # 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_properties=False, + include_annotations=False, + include_extra_metadata=False, + ) + + return cached + read_binary_folder = define_function_from_class(source_class=BinaryFolderRecording, name="read_binary_folder") 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 b53625096c..083f9738b9 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -1,3 +1,4 @@ +import warnings from pathlib import Path, WindowsPath from collections import namedtuple from collections.abc import Generator, Callable @@ -810,3 +811,102 @@ def slice_rows(array: np.ndarray | zarr.Array, row_indices: np.ndarray | list) - return array.oindex[row_indices] else: return array[row_indices, ...] + + +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. + + Parameters + ---------- + folder : str or Path + The folder containing the properties as .npy files. + extractor : BaseExtractor + The extractor to which the properties will be set. + """ + folder = Path(folder) + if folder.is_dir(): + for prop_file in folder.iterdir(): + if prop_file.suffix == ".npy": + values = np.load(prop_file, allow_pickle=True) + key = prop_file.stem + if key == "contact_vector": + continue + extractor.set_property(key, values) + + +def save_properties_to_folder(folder: str | Path, extractor: "BaseExtractor"): + """ + Save properties from an extractor to a folder as .npy files. + + Parameters + ---------- + folder : str or Path + The folder where the properties will be saved as .npy files. + extractor : BaseExtractor + The extractor from which the properties will be saved. + """ + folder = Path(folder) + folder.mkdir(exist_ok=True) + for key in extractor.get_property_keys(): + 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 BaseExtractor annotations to annotations.json in the provided folder. + This is used for `BinaryFolderRecording` and `NumpyFolderSorting` since version 0.105.0. + + Parameters + ---------- + folder : str or Path + 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.json").write_text(json.dumps(extractor._annotations, indent=4), encoding="utf8") + + +def load_annotations_from_folder(folder: str | Path, extractor: "BaseExtractor"): + """ + 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 + ---------- + folder : str or Path + The folder where the annotations will be loaded from. + extractor : BaseExtractor + The extractor to which the annotations will be added. + """ + folder = Path(folder) + annotations_file = folder / "annotations.json" + 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) + if extractor.check_serializability("json"): + provenance_file_path = folder / f"provenance.json" + extractor.dump_to_json(file_path=provenance_file_path, relative_to=folder) + elif extractor.check_serializability("pickle"): + provenance_file = folder / f"provenance.pkl" + extractor.dump_to_pickle(provenance_file, relative_to=folder) + else: + warnings.warn("The extractor is not serializable to file. The provenance will not be saved.") 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/loading.py b/src/spikeinterface/core/loading.py index 93ac7c2a56..9c06e95300 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 1 @@ -348,13 +349,18 @@ def test_get_best_job_kwargs(): if __name__ == "__main__": + import tempfile + from pathlib import Path + + tmp_path = Path(tempfile.mkdtemp()) + # test_divide_segment_into_chunks() - # test_ensure_n_jobs() + test_ensure_n_jobs(tmp_path) # test_ensure_chunk_size() # test_ChunkExecutor() # test_fix_job_kwargs() # test_split_job_kwargs() - test_worker_index() + # test_worker_index() # test_get_best_job_kwargs() # quick_becnhmark() 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/core/tests/test_noise_levels_propagation.py b/src/spikeinterface/core/tests/test_noise_levels_propagation.py index bfa6b7b6f0..98e7a929fd 100644 --- a/src/spikeinterface/core/tests/test_noise_levels_propagation.py +++ b/src/spikeinterface/core/tests/test_noise_levels_propagation.py @@ -8,10 +8,10 @@ import numpy as np -def test_skip_noise_levels_propagation(): +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() + 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() @@ -37,4 +37,9 @@ def test_skip_noise_levels_propagation(): if __name__ == "__main__": - test_skip_noise_levels_propagation() + import tempfile + from pathlib import Path + + tmp_path = Path(tempfile.mkdtemp()) + + test_skip_noise_levels_propagation(tmp_path) diff --git a/src/spikeinterface/core/tests/test_numpy_extractors.py b/src/spikeinterface/core/tests/test_numpy_extractors.py index 6eb9918b66..daaf99d72a 100644 --- a/src/spikeinterface/core/tests/test_numpy_extractors.py +++ b/src/spikeinterface/core/tests/test_numpy_extractors.py @@ -160,7 +160,7 @@ def test_NumpyEvent(): if __name__ == "__main__": # test_NumpyRecording() - test_SharedMemoryRecording() - # test_NumpySorting() + # test_SharedMemoryRecording() + test_NumpySorting() # test_SharedMemorySorting() # test_NumpyEvent() diff --git a/src/spikeinterface/core/tests/test_recording_tools.py b/src/spikeinterface/core/tests/test_recording_tools.py index 477e4f04aa..13d68e3b45 100644 --- a/src/spikeinterface/core/tests/test_recording_tools.py +++ b/src/spikeinterface/core/tests/test_recording_tools.py @@ -149,7 +149,6 @@ def test_write_memory_recording(): recording = MockRecording( num_channels=2, durations=[10.325, 3.5], sampling_frequency=30_000, strategy="tile_pregenerated" ) - recording = recording.save() # write with loop traces_list, shms = write_memory_recording(recording, dtype=None, verbose=True, n_jobs=1) diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index ab838b7a40..7c1464c175 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -324,7 +324,7 @@ def test_load_without_runtime_info(tmp_path, dataset): def test_SortingAnalyzer_tmp_recording(dataset): recording, sorting = dataset - recording_cached = recording.save(mode="memory") + recording_cached = recording.save(format="memory") sorting_analyzer = create_sorting_analyzer(sorting, recording, format="memory", sparse=False, sparsity=None) sorting_analyzer.set_temporary_recording(recording_cached) @@ -1101,12 +1101,16 @@ def test_merge_units_main_channel_id_disagreement(): if __name__ == "__main__": - tmp_path = Path("test_SortingAnalyzer") + import tempfile + from pathlib import Path + + tmp_path = Path(tempfile.mkdtemp()) / "test_SortingAnalyzer" + dataset = get_dataset() - test_SortingAnalyzer_memory(tmp_path, dataset) - test_SortingAnalyzer_binary_folder(tmp_path, dataset) - test_SortingAnalyzer_zarr(tmp_path, dataset) + # test_SortingAnalyzer_memory(tmp_path, dataset) + # test_SortingAnalyzer_binary_folder(tmp_path, dataset) + # test_SortingAnalyzer_zarr(tmp_path, dataset) test_SortingAnalyzer_tmp_recording(dataset) - test_extension() - test_extension_params() - test_runtime_dependencies(dataset) + # test_extension() + # test_extension_params() + # test_runtime_dependencies(dataset) diff --git a/src/spikeinterface/core/tests/test_time_series_tools.py b/src/spikeinterface/core/tests/test_time_series_tools.py index 4c6ba6b105..50d91734c3 100644 --- a/src/spikeinterface/core/tests/test_time_series_tools.py +++ b/src/spikeinterface/core/tests/test_time_series_tools.py @@ -139,7 +139,7 @@ def test_write_memory_recording(): recording = MockRecording( num_channels=2, durations=[10.325, 3.5], sampling_frequency=30_000, strategy="tile_pregenerated" ) - recording = recording.save() + # recording = recording.save() # write with loop traces_list, shms = write_memory(recording, dtype=None, verbose=True, n_jobs=1) 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/time_series.py b/src/spikeinterface/core/time_series.py index 640a256d24..b3dfe89857 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`). @@ -34,6 +34,10 @@ class TimeSeries(ABC): _preferred_mp_context = None + 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) if not issubclass(cls, BaseExtractor): @@ -216,9 +220,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 +250,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 +265,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 +301,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 +366,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 +377,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 +523,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/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index 640d10e1c8..448dbd94a7 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): @@ -547,13 +554,21 @@ 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) + return ZarrSortingExtractor(folder_path, storage_options=storage_options) read_zarr_recording = define_function_from_class(source_class=ZarrRecordingExtractor, name="read_zarr_recording") @@ -605,7 +620,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: @@ -614,6 +629,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, @@ -741,7 +777,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) @@ -755,6 +790,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) @@ -765,159 +808,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() 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() 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. 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 1dfc6702b9..3b47d0eb1a 100644 --- a/src/spikeinterface/preprocessing/tests/test_filter.py +++ b/src/spikeinterface/preprocessing/tests/test_filter.py @@ -140,20 +140,22 @@ def _get_filter_options(self): } -def test_filter(): +def test_filter(create_cache_folder): rec = generate_recording() - rec = rec.save() + rec = rec.save(folder=create_cache_folder / "test_filter_recording") 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(): 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(): # 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) @@ -221,4 +227,9 @@ def test_filter_opencl(): if __name__ == "__main__": - test_filter() + import tempfile + from pathlib import Path + + tmp_path = Path(tempfile.mkdtemp()) + + test_filter(tmp_path) 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_pipeline.py b/src/spikeinterface/preprocessing/tests/test_pipeline.py index 78fbbe4867..5f5c298430 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_list = get_preprocessing_list_from_file(cache_folder / "provenance.pkl") 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) diff --git a/src/spikeinterface/sorters/basesorter.py b/src/spikeinterface/sorters/basesorter.py index 0ae0f3645b..c0cd231441 100644 --- a/src/spikeinterface/sorters/basesorter.py +++ b/src/spikeinterface/sorters/basesorter.py @@ -150,7 +150,7 @@ def initialize_folder(cls, recording, output_folder, verbose, remove_existing_fo recording.dump(output_folder / "spikeinterface_recording.pickle", relative_to=output_folder) else: raise RuntimeError( - "This recording is not serializable and so can not be sorted. Consider `recording.save()` to save a " + "This recording is not serializable and so can not be sorted. Consider `recording.save(folder=...)` to save a " "compatible binary file." ) diff --git a/src/spikeinterface/sortingcomponents/tools.py b/src/spikeinterface/sortingcomponents/tools.py index 2a76b083d8..b445513a6c 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: