Skip to content

Commit 605b7b4

Browse files
committed
Fix saving analyzer directly to remote storage
1 parent f4d6d36 commit 605b7b4

1 file changed

Lines changed: 41 additions & 23 deletions

File tree

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 41 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -124,12 +124,14 @@ def create_sorting_analyzer(
124124
"""
125125
if format != "memory":
126126
if format == "zarr":
127-
folder = clean_zarr_folder_name(folder)
128-
if Path(folder).is_dir():
129-
if not overwrite:
130-
raise ValueError(f"Folder already exists {folder}! Use overwrite=True to overwrite it.")
131-
else:
132-
shutil.rmtree(folder)
127+
if not is_path_remote(folder):
128+
folder = clean_zarr_folder_name(folder)
129+
if not is_path_remote(folder):
130+
if Path(folder).is_dir():
131+
if not overwrite:
132+
raise ValueError(f"Folder already exists {folder}! Use overwrite=True to overwrite it.")
133+
else:
134+
shutil.rmtree(folder)
133135

134136
# handle sparsity
135137
if sparsity is not None:
@@ -249,6 +251,9 @@ def __repr__(self) -> str:
249251
nchan = self.get_num_channels()
250252
nunits = self.get_num_units()
251253
txt = f"{clsname}: {nchan} channels - {nunits} units - {nseg} segments - {self.format}"
254+
if self.format != "memory":
255+
if is_path_remote(str(self.folder)):
256+
txt += f" (remote)"
252257
if self.is_sparse():
253258
txt += " - sparse"
254259
if self.has_recording():
@@ -311,7 +316,8 @@ def create(
311316
)
312317
elif format == "zarr":
313318
assert folder is not None, "For format='zarr' folder must be provided"
314-
folder = clean_zarr_folder_name(folder)
319+
if not is_path_remote(folder):
320+
folder = clean_zarr_folder_name(folder)
315321
sorting_analyzer = cls.create_zarr(
316322
folder,
317323
sorting,
@@ -349,12 +355,7 @@ def load(cls, folder, recording=None, load_extensions=True, format="auto", backe
349355
folder, recording=recording, backend_options=backend_options
350356
)
351357

352-
if is_path_remote(str(folder)):
353-
sorting_analyzer.folder = folder
354-
# in this case we only load extensions when needed
355-
else:
356-
sorting_analyzer.folder = Path(folder)
357-
358+
if not is_path_remote(str(folder)):
358359
if load_extensions:
359360
sorting_analyzer.load_all_saved_extension()
360361

@@ -537,12 +538,16 @@ def load_from_binary_folder(cls, folder, recording=None, backend_options=None):
537538
def _get_zarr_root(self, mode="r+"):
538539
import zarr
539540

540-
# if is_path_remote(str(self.folder)):
541-
# mode = "r"
541+
assert mode in ("r+", "a", "r"), "mode must be 'r+', 'a' or 'r'"
542+
542543
storage_options = self._backend_options.get("storage_options", {})
543544
# we open_consolidated only if we are in read mode
544545
if mode in ("r+", "a"):
545-
zarr_root = zarr.open(str(self.folder), mode=mode, storage_options=storage_options)
546+
try:
547+
zarr_root = zarr.open(str(self.folder), mode=mode, storage_options=storage_options)
548+
except Exception as e:
549+
# this could happen in remote mode, and it's a way to check if the folder is still there
550+
zarr_root = zarr.open_consolidated(self.folder, mode=mode, storage_options=storage_options)
546551
else:
547552
zarr_root = zarr.open_consolidated(self.folder, mode=mode, storage_options=storage_options)
548553
return zarr_root
@@ -554,10 +559,14 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
554559
import numcodecs
555560
from .zarrextractors import add_sorting_to_zarr_group
556561

557-
folder = clean_zarr_folder_name(folder)
558-
559-
if folder.is_dir():
560-
raise ValueError(f"Folder already exists {folder}")
562+
if is_path_remote(folder):
563+
remote = True
564+
else:
565+
remote = False
566+
if not remote:
567+
folder = clean_zarr_folder_name(folder)
568+
if folder.is_dir():
569+
raise ValueError(f"Folder already exists {folder}")
561570

562571
backend_options = {} if backend_options is None else backend_options
563572
storage_options = backend_options.get("storage_options", {})
@@ -572,8 +581,9 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
572581
zarr_root.attrs["settings"] = check_json(settings)
573582

574583
# the recording
584+
relative_to = folder if not remote else None
575585
if recording is not None:
576-
rec_dict = recording.to_dict(relative_to=folder, recursive=True)
586+
rec_dict = recording.to_dict(relative_to=relative_to, recursive=True)
577587
if recording.check_serializability("json"):
578588
# zarr_root.create_dataset("recording", data=rec_dict, object_codec=numcodecs.JSON())
579589
zarr_rec = np.array([check_json(rec_dict)], dtype=object)
@@ -589,7 +599,7 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
589599
warnings.warn("Recording not provided, instntiating SortingAnalyzer in recordingless mode.")
590600

591601
# sorting provenance
592-
sort_dict = sorting.to_dict(relative_to=folder, recursive=True)
602+
sort_dict = sorting.to_dict(relative_to=relative_to, recursive=True)
593603
if sorting.check_serializability("json"):
594604
zarr_sort = np.array([check_json(sort_dict)], dtype=object)
595605
zarr_root.create_dataset("sorting_provenance", data=zarr_sort, object_codec=numcodecs.JSON())
@@ -1106,7 +1116,15 @@ def copy(self):
11061116
def is_read_only(self) -> bool:
11071117
if self.format == "memory":
11081118
return False
1109-
return not os.access(self.folder, os.W_OK)
1119+
elif self.format == "binary_folder":
1120+
return not os.access(self.folder, os.W_OK)
1121+
else:
1122+
if not is_path_remote(str(self.folder)):
1123+
return not os.access(self.folder, os.W_OK)
1124+
else:
1125+
# in this case we don't know if the file is read only so an error
1126+
# will be raised if we try to save/append
1127+
return False
11101128

11111129
## map attribute and property zone
11121130

0 commit comments

Comments
 (0)