Skip to content
Merged
Show file tree
Hide file tree
Changes from 15 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 10 additions & 16 deletions src/spikeinterface/core/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
clean_zarr_folder_name,
is_dict_extractor,
SIJsonEncoder,
is_path_remote,
make_paths_relative,
make_paths_absolute,
check_paths_relative,
Expand Down Expand Up @@ -643,7 +644,7 @@ def from_dict(dictionary: dict, base_folder: Path | str | None = None) -> "BaseE
extractor.load_metadata_from_folder(folder_metadata)
return extractor

def load_metadata_from_folder(self, folder_metadata):
def load_metadata_from_folder(self, folder_metadata: str | Path):
# hack to load probe for recording
folder_metadata = Path(folder_metadata)

Expand All @@ -658,11 +659,11 @@ def load_metadata_from_folder(self, folder_metadata):

self._extra_metadata_from_folder(folder_metadata)

def save_metadata_to_folder(self, folder_metadata):
def save_metadata_to_folder(self, folder_metadata: str | Path):
self._extra_metadata_to_folder(folder_metadata)

# save properties
prop_folder = folder_metadata / "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)
Expand Down Expand Up @@ -1084,24 +1085,17 @@ def save_to_zarr(
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 / f"{name}.zarr"
if verbose:
print(f"Use zarr_path={zarr_path}")
else:
Comment thread
ecobost marked this conversation as resolved.
zarr_path = cache_folder / f"{name}.zarr"
if not is_set_global_tmp_folder():
if verbose:
print(f"Use zarr_path={zarr_path}")
zarr_path = (cache_folder / name).with_suffix(".zarr")
if verbose:
print(f"Saving to zarr_path={zarr_path}")
else:
if storage_options is None:
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
else:
zarr_path = folder
zarr_path = folder

if isinstance(zarr_path, Path):
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
Expand Down
35 changes: 17 additions & 18 deletions src/spikeinterface/core/basesorting.py
Original file line number Diff line number Diff line change
Expand Up @@ -344,7 +344,7 @@ def register_recording(self, recording, check_spike_frames: bool = True):
warnings.warn(
"Some spikes exceed the recording's duration! "
"Removing these excess spikes with `spikeinterface.curation.remove_excess_spikes()` "
"Might be necessary for further postprocessing."
"might be necessary for further postprocessing."
)
self._recording = recording
# Copy the recording's start times into the sorting segments. This way,
Expand Down Expand Up @@ -500,14 +500,16 @@ def get_times(
else:
return None

def _save(self, format="numpy_folder", **save_kwargs):
"""
This function replaces the old CachesortingExtractor, but enables more engines
def _save(self, format: str = "numpy_folder", **save_kwargs):
"""Save a sorting object to disk in a specified format.

Note
----
This function replaces the old CacheSortingExtractor, but enables more engines
for caching a results.

Since v0.98.0 "numpy_folder" is used by defult.
From v0.96.0 to 0.97.0 "npz_folder" was the default.

"""
if format == "numpy_folder":
from .sortingfolder import NumpyFolderSorting
Expand All @@ -516,10 +518,6 @@ def _save(self, format="numpy_folder", **save_kwargs):
NumpyFolderSorting.write_sorting(self, folder)
cached = NumpyFolderSorting(folder)

if self.has_recording():
warnings.warn("The registered recording will not be persistent on disk, but only available in memory")
cached.register_recording(self._recording)

elif format == "zarr":
from .zarrextractors import ZarrSortingExtractor

Expand All @@ -528,21 +526,13 @@ def _save(self, format="numpy_folder", **save_kwargs):
ZarrSortingExtractor.write_sorting(self, zarr_path, storage_options, **save_kwargs)
cached = ZarrSortingExtractor(zarr_path, storage_options)

if self.has_recording():
warnings.warn("The registered recording will not be persistent on disk, but only available in memory")
cached.register_recording(self._recording)

elif format == "npz_folder":
from .sortingfolder import NpzFolderSorting

folder = save_kwargs.pop("folder")
NpzFolderSorting.write_sorting(self, folder)
cached = NpzFolderSorting(folder_path=folder)

if self.has_recording():
warnings.warn("The registered recording will not be persistent on disk, but only available in memory")
cached.register_recording(self._recording)

elif format == "memory":
if save_kwargs.get("sharedmem", True):
from .numpyextractors import SharedMemorySorting
Expand All @@ -553,7 +543,16 @@ def _save(self, format="numpy_folder", **save_kwargs):

cached = NumpySorting.from_sorting(self)
else:
raise ValueError(f"format {format} not supported")
raise ValueError(f"Format {format} not supported")

# Re-register the recording if saving to disk (not memory)
if self.has_recording() and format != "memory":
warnings.warn(
"The recording registered to this sorting object will not be saved to disk"
Comment thread
alejoe91 marked this conversation as resolved.
Outdated
"Reloading the sorting later will not include the recording"
)
cached.register_recording(self._recording)

return cached

def get_unit_property(self, unit_id, key):
Expand Down
8 changes: 3 additions & 5 deletions src/spikeinterface/core/core_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,11 +207,9 @@ def check_json(dictionary: dict) -> dict:
return json.loads(json_string)


def clean_zarr_folder_name(folder):
def clean_zarr_folder_name(folder: str | Path) -> Path:
folder = Path(folder)
if folder.suffix != ".zarr":
folder = folder.parent / f"{folder.stem}.zarr"
return folder
return folder if folder.suffix == ".zarr" else folder.with_suffix(".zarr")


def add_suffix(file_path, possible_suffix):
Expand Down Expand Up @@ -772,7 +770,7 @@ def is_path_remote(path: str | Path) -> bool:

Returns
-------
bool
is_remote: bool
Whether the path is a remote path.
"""
return "s3://" in str(path) or "gcs://" in str(path)
Expand Down
4 changes: 2 additions & 2 deletions src/spikeinterface/core/globals.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
temp_folder_set = False


def get_global_tmp_folder():
def get_global_tmp_folder() -> Path:
"""
Get the global path temporary folder.
"""
Expand All @@ -30,7 +30,7 @@ def get_global_tmp_folder():
return temp_folder


def set_global_tmp_folder(folder):
def set_global_tmp_folder(folder: str | Path):
"""
Set the global path temporary folder.
"""
Expand Down
Loading
Loading