Skip to content

Commit cbb8156

Browse files
committed
fix: zarr loading and rename numpy_folder -> binary
1 parent 67a8d2d commit cbb8156

5 files changed

Lines changed: 9 additions & 14 deletions

File tree

src/spikeinterface/core/loading.py

Lines changed: 3 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -292,23 +292,17 @@ def _load_object_from_zarr(folder_or_url, object_type, **kwargs):
292292
elif object_type == "Recording":
293293
from .zarrextractors import read_zarr_recording
294294

295-
storage_options = kwargs.get("storage_options", None)
296-
load_compression_ratio = kwargs.get("load_compression_ratio", False)
297-
recording = read_zarr_recording(
298-
folder_or_url, storage_options=storage_options, load_compression_ratio=load_compression_ratio
299-
)
295+
recording = read_zarr_recording(folder_or_url, **kwargs)
300296
return recording
301297
elif object_type == "Sorting":
302298
from .zarrextractors import read_zarr_sorting
303299

304-
storage_options = kwargs.get("storage_options", None)
305-
sorting = read_zarr_sorting(folder_or_url, storage_options=storage_options)
300+
sorting = read_zarr_sorting(folder_or_url, **kwargs)
306301
return sorting
307302
elif object_type == "Recording|Sorting":
308303
# This case shoudl deprecated soon because the read_zarr is ultra ambiguous
309304
# just testing if the zarr contains unit_ids or channel_ids but many object also contains it (see template)!!!!
310305
from .zarrextractors import read_zarr
311306

312-
storage_options = kwargs.get("storage_options", None)
313-
rec_or_sorting = read_zarr(folder_or_url, storage_options=storage_options)
307+
rec_or_sorting = read_zarr(folder_or_url, **kwargs)
314308
return rec_or_sorting

src/spikeinterface/core/tests/test_basesorting.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -78,7 +78,7 @@ def test_BaseSorting(create_cache_folder):
7878
# cache new format : numpy_folder
7979
folder = cache_folder / "simple_sorting_numpy_folder"
8080
sorting.set_property("test", np.ones(len(sorting.unit_ids)))
81-
sorting.save(folder=folder, format="numpy_folder")
81+
sorting.save(folder=folder, format="binary")
8282
sorting2 = BaseExtractor.load(folder)
8383
assert isinstance(sorting2, NumpyFolderSorting)
8484

@@ -143,7 +143,7 @@ def test_BaseSorting(create_cache_folder):
143143

144144
# test save to zarr
145145
# compressor = get_default_zarr_compressor()
146-
sorting_zarr = sorting.save(format="zarr", folder=cache_folder / "sorting")
146+
sorting_zarr = sorting.save(format="zarr", folder=cache_folder / "sorting.zarr")
147147
sorting_zarr_loaded = load(cache_folder / "sorting.zarr")
148148
# annotations is False because Zarr adds compression ratios
149149
check_sortings_equal(sorting, sorting_zarr, check_annotations=False, check_properties=True)

src/spikeinterface/core/tests/test_loading.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -102,7 +102,7 @@ def test_load_binary_recording(generate_recording_sorting, tmp_path, output_form
102102
check_recordings_equal(rec, rec_loaded)
103103

104104

105-
@pytest.mark.parametrize("output_format", ["numpy_folder", "zarr"])
105+
@pytest.mark.parametrize("output_format", ["binary", "zarr"])
106106
def test_load_binary_sorting(generate_recording_sorting, tmp_path, output_format):
107107
_, sort = generate_recording_sorting
108108
_ = sort.save(folder=tmp_path / "test_sorting", format=output_format, overwrite=True)

src/spikeinterface/core/tests/test_zarrextractors.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -60,13 +60,13 @@ def test_ZarrSortingExtractor(tmp_path):
6060
np_sorting = generate_sorting()
6161

6262
# store in root standard normal way
63-
folder = tmp_path / "zarr_sorting"
63+
folder = tmp_path / "zarr_sorting.zarr"
6464
ZarrSortingExtractor.write_sorting(np_sorting, folder)
6565
sorting = ZarrSortingExtractor(folder)
6666
sorting = load(sorting.to_dict())
6767

6868
# store the sorting in a sub group (for instance SortingResult)
69-
folder = tmp_path / "zarr_sorting_sub_group"
69+
folder = tmp_path / "zarr_sorting_sub_group.zarr"
7070
zarr_root = zarr.open(folder, mode="w")
7171
zarr_sorting_group = zarr_root.create_group("sorting")
7272
add_sorting_to_zarr_group(sorting, zarr_sorting_group)

src/spikeinterface/core/zarrextractors.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -495,6 +495,7 @@ def write_sorting(
495495
zarr_root = zarr.open(str(folder_path), mode="w", storage_options=storage_options)
496496
zarr_root.attrs["zarr_class_info"] = retrieve_importing_provenance(ZarrSortingExtractor)
497497
add_sorting_to_zarr_group(sorting, zarr_root, **kwargs)
498+
return ZarrSortingExtractor(folder_path, storage_options=storage_options)
498499

499500

500501
read_zarr_recording = define_function_from_class(source_class=ZarrRecordingExtractor, name="read_zarr_recording")

0 commit comments

Comments
 (0)