Skip to content

Commit 3921806

Browse files
authored
Merge pull request #3314 from alejoe91/load-cloud-sorting-analyzer
Enable cloud-loading for analyzer Zarr
2 parents 73f6151 + 33e27b1 commit 3921806

2 files changed

Lines changed: 40 additions & 18 deletions

File tree

src/spikeinterface/core/core_tools.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -684,3 +684,20 @@ def measure_memory_allocation(measure_in_process: bool = True) -> float:
684684
memory = mem_info.total - mem_info.available
685685

686686
return memory
687+
688+
689+
def is_path_remote(path: str | Path) -> bool:
690+
"""
691+
Returns True if the path is a remote path (e.g., s3:// or gcs://).
692+
693+
Parameters
694+
----------
695+
path : str or Path
696+
The path to check.
697+
698+
Returns
699+
-------
700+
bool
701+
Whether the path is a remote path.
702+
"""
703+
return "s3://" in str(path) or "gcs://" in str(path)

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 23 additions & 18 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
26+
from .core_tools import check_json, retrieve_importing_provenance, is_path_remote
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
@@ -195,6 +195,7 @@ def __init__(
195195
format=None,
196196
sparsity=None,
197197
return_scaled=True,
198+
storage_options=None,
198199
):
199200
# very fast init because checks are done in load and create
200201
self.sorting = sorting
@@ -204,6 +205,7 @@ def __init__(
204205
self.format = format
205206
self.sparsity = sparsity
206207
self.return_scaled = return_scaled
208+
self.storage_options = storage_options
207209
# this is used to store temporary recording
208210
self._temporary_recording = None
209211

@@ -276,30 +278,34 @@ def create(
276278
return sorting_analyzer
277279

278280
@classmethod
279-
def load(cls, folder, recording=None, load_extensions=True, format="auto"):
281+
def load(cls, folder, recording=None, load_extensions=True, format="auto", storage_options=None):
280282
"""
281283
Load folder or zarr.
282284
The recording can be given if the recording location has changed.
283285
Otherwise the recording is loaded when possible.
284286
"""
285-
folder = Path(folder)
286-
assert folder.is_dir(), "Waveform folder does not exists"
287287
if format == "auto":
288288
# make better assumption and check for auto guess format
289-
if folder.suffix == ".zarr":
289+
if Path(folder).suffix == ".zarr":
290290
format = "zarr"
291291
else:
292292
format = "binary_folder"
293293

294294
if format == "binary_folder":
295295
sorting_analyzer = SortingAnalyzer.load_from_binary_folder(folder, recording=recording)
296296
elif format == "zarr":
297-
sorting_analyzer = SortingAnalyzer.load_from_zarr(folder, recording=recording)
297+
sorting_analyzer = SortingAnalyzer.load_from_zarr(
298+
folder, recording=recording, storage_options=storage_options
299+
)
298300

299-
sorting_analyzer.folder = folder
301+
if is_path_remote(str(folder)):
302+
sorting_analyzer.folder = folder
303+
# in this case we only load extensions when needed
304+
else:
305+
sorting_analyzer.folder = Path(folder)
300306

301-
if load_extensions:
302-
sorting_analyzer.load_all_saved_extension()
307+
if load_extensions:
308+
sorting_analyzer.load_all_saved_extension()
303309

304310
return sorting_analyzer
305311

@@ -470,7 +476,9 @@ def load_from_binary_folder(cls, folder, recording=None):
470476
def _get_zarr_root(self, mode="r+"):
471477
import zarr
472478

473-
zarr_root = zarr.open(self.folder, mode=mode)
479+
if is_path_remote(str(self.folder)):
480+
mode = "r"
481+
zarr_root = zarr.open(self.folder, mode=mode, storage_options=self.storage_options)
474482
return zarr_root
475483

476484
@classmethod
@@ -552,25 +560,22 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
552560
recording_info = zarr_root.create_group("extensions")
553561

554562
@classmethod
555-
def load_from_zarr(cls, folder, recording=None):
563+
def load_from_zarr(cls, folder, recording=None, storage_options=None):
556564
import zarr
557565

558-
folder = Path(folder)
559-
assert folder.is_dir(), f"This folder does not exist {folder}"
560-
561-
zarr_root = zarr.open(folder, mode="r")
566+
zarr_root = zarr.open(str(folder), mode="r", storage_options=storage_options)
562567

563568
# load internal sorting in memory
564-
# TODO propagate storage_options
565569
sorting = NumpySorting.from_sorting(
566-
ZarrSortingExtractor(folder, zarr_group="sorting"), with_metadata=True, copy_spike_vector=True
570+
ZarrSortingExtractor(folder, zarr_group="sorting", storage_options=storage_options),
571+
with_metadata=True,
572+
copy_spike_vector=True,
567573
)
568574

569575
# load recording if possible
570576
if recording is None:
571577
rec_dict = zarr_root["recording"][0]
572578
try:
573-
574579
recording = load_extractor(rec_dict, base_folder=folder)
575580
except:
576581
recording = None

0 commit comments

Comments
 (0)