Skip to content

Commit 9238023

Browse files
authored
Merge pull request #3349 from jonahpearl/zarr_folder_suffix
Fix zarr folder suffix handling
2 parents 007b6ef + c992ca6 commit 9238023

3 files changed

Lines changed: 27 additions & 11 deletions

File tree

src/spikeinterface/core/base.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from .globals import get_global_tmp_folder, is_set_global_tmp_folder
1818
from .core_tools import (
1919
check_json,
20+
clean_zarr_folder_name,
2021
is_dict_extractor,
2122
SIJsonEncoder,
2223
make_paths_relative,
@@ -1061,9 +1062,7 @@ def save_to_zarr(
10611062
print(f"Use zarr_path={zarr_path}")
10621063
else:
10631064
if storage_options is None:
1064-
folder = Path(folder)
1065-
if folder.suffix != ".zarr":
1066-
folder = folder.parent / f"{folder.stem}.zarr"
1065+
folder = clean_zarr_folder_name(folder)
10671066
if folder.is_dir() and overwrite:
10681067
shutil.rmtree(folder)
10691068
zarr_path = folder

src/spikeinterface/core/core_tools.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -153,6 +153,13 @@ def check_json(dictionary: dict) -> dict:
153153
return json.loads(json_string)
154154

155155

156+
def clean_zarr_folder_name(folder):
157+
folder = Path(folder)
158+
if folder.suffix != ".zarr":
159+
folder = folder.parent / f"{folder.stem}.zarr"
160+
return folder
161+
162+
156163
def add_suffix(file_path, possible_suffix):
157164
file_path = Path(file_path)
158165
if isinstance(possible_suffix, str):

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 18 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@
2323

2424
from .base import load_extractor
2525
from .recording_tools import check_probe_do_not_overlap, get_rec_attributes, do_recording_attributes_match
26-
from .core_tools import check_json, retrieve_importing_provenance, is_path_remote
26+
from .core_tools import check_json, retrieve_importing_provenance, is_path_remote, clean_zarr_folder_name
2727
from .sorting_tools import generate_unit_ids_for_merge_group, _get_ids_after_merging
2828
from .job_tools import split_job_kwargs
2929
from .numpyextractors import NumpySorting
@@ -111,6 +111,8 @@ def create_sorting_analyzer(
111111
sparsity off (or give external sparsity) like this.
112112
"""
113113
if format != "memory":
114+
if format == "zarr":
115+
folder = clean_zarr_folder_name(folder)
114116
if Path(folder).is_dir():
115117
if not overwrite:
116118
raise ValueError(f"Folder already exists {folder}! Use overwrite=True to overwrite it.")
@@ -162,6 +164,8 @@ def load_sorting_analyzer(folder, load_extensions=True, format="auto"):
162164
The loaded SortingAnalyzer
163165
164166
"""
167+
if format == "zarr":
168+
folder = clean_zarr_folder_name(folder)
165169
return SortingAnalyzer.load(folder, load_extensions=load_extensions, format=format)
166170

167171

@@ -269,6 +273,8 @@ def create(
269273
sorting_analyzer = cls.load_from_binary_folder(folder, recording=recording)
270274
sorting_analyzer.folder = Path(folder)
271275
elif format == "zarr":
276+
assert folder is not None, "For format='zarr' folder must be provided"
277+
folder = clean_zarr_folder_name(folder)
272278
cls.create_zarr(folder, sorting, recording, sparsity, return_scaled, rec_attributes=None)
273279
sorting_analyzer = cls.load_from_zarr(folder, recording=recording)
274280
sorting_analyzer.folder = Path(folder)
@@ -487,10 +493,7 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
487493
import zarr
488494
import numcodecs
489495

490-
folder = Path(folder)
491-
# force zarr sufix
492-
if folder.suffix != ".zarr":
493-
folder = folder.parent / f"{folder.stem}.zarr"
496+
folder = clean_zarr_folder_name(folder)
494497

495498
if folder.is_dir():
496499
raise ValueError(f"Folder already exists {folder}")
@@ -768,9 +771,7 @@ def _save_or_select_or_merge(
768771

769772
elif format == "zarr":
770773
assert folder is not None, "For format='zarr' folder must be provided"
771-
folder = Path(folder)
772-
if folder.suffix != ".zarr":
773-
folder = folder.parent / f"{folder.stem}.zarr"
774+
folder = clean_zarr_folder_name(folder)
774775
SortingAnalyzer.create_zarr(
775776
folder, sorting_provenance, recording, sparsity, self.return_scaled, self.rec_attributes
776777
)
@@ -829,6 +830,8 @@ def save_as(self, format="memory", folder=None) -> "SortingAnalyzer":
829830
format : "memory" | "binary_folder" | "zarr", default: "memory"
830831
The new backend format to use
831832
"""
833+
if format == "zarr":
834+
folder = clean_zarr_folder_name(folder)
832835
return self._save_or_select_or_merge(format=format, folder=folder)
833836

834837
def select_units(self, unit_ids, format="memory", folder=None) -> "SortingAnalyzer":
@@ -854,6 +857,8 @@ def select_units(self, unit_ids, format="memory", folder=None) -> "SortingAnalyz
854857
The newly create sorting_analyzer with the selected units
855858
"""
856859
# TODO check that unit_ids are in same order otherwise many extension do handle it properly!!!!
860+
if format == "zarr":
861+
folder = clean_zarr_folder_name(folder)
857862
return self._save_or_select_or_merge(format=format, folder=folder, unit_ids=unit_ids)
858863

859864
def remove_units(self, remove_unit_ids, format="memory", folder=None) -> "SortingAnalyzer":
@@ -880,6 +885,8 @@ def remove_units(self, remove_unit_ids, format="memory", folder=None) -> "Sortin
880885
"""
881886
# TODO check that unit_ids are in same order otherwise many extension do handle it properly!!!!
882887
unit_ids = self.unit_ids[~np.isin(self.unit_ids, remove_unit_ids)]
888+
if format == "zarr":
889+
folder = clean_zarr_folder_name(folder)
883890
return self._save_or_select_or_merge(format=format, folder=folder, unit_ids=unit_ids)
884891

885892
def merge_units(
@@ -938,6 +945,9 @@ def merge_units(
938945
The newly create `SortingAnalyzer` with the selected units
939946
"""
940947

948+
if format == "zarr":
949+
folder = clean_zarr_folder_name(folder)
950+
941951
assert merging_mode in ["soft", "hard"], "Merging mode should be either soft or hard"
942952

943953
if len(merge_unit_groups) == 0:

0 commit comments

Comments
 (0)