diff --git a/.github/scripts/serialization/objects.py b/.github/scripts/serialization/objects.py index 78fb3aed0b..d65856a498 100644 --- a/.github/scripts/serialization/objects.py +++ b/.github/scripts/serialization/objects.py @@ -17,10 +17,12 @@ the full state (traces or spike trains, properties, annotations, probe) to disk. Targets the on-disk encoding axis: property and annotation preservation, and the probe representation. "binary" is recording-only, hence "numpy_folder" for sortings. -""" -from packaging.version import parse -from spikeinterface import __version__ as si_version +Additionally, the `check_extra_data` kwarg can be used to store additional information in a JSON file alongside the +fixture, which is then passed to the check function. This is used for the SortingAnalyzer extension data keys, +which are computed and stored in the fixture JSON to verify that they are reloaded correctly. + +""" # Filename suffix per format, relative to the fixtures dir (folder formats use a suffix # rather than an extension). Both the generator and the loader build a fixture path as @@ -32,6 +34,7 @@ "binary_parallel": "_binary_parallel", "numpy_folder": "_numpy_folder", "zarr": ".zarr", + "binary_folder": "_binary_folder", "zarr_parallel": "_parallel.zarr", } @@ -57,7 +60,7 @@ def _build_noise_generator_recording(): ) -def _check_noise_generator_recording(rec): +def _check_noise_generator_recording(rec, check_extra_data=None): assert type(rec).__name__ == "NoiseGeneratorRecording", type(rec).__name__ assert rec.get_num_channels() == 4 assert rec.get_num_segments() == 2 @@ -70,7 +73,7 @@ def _build_mock_recording(): return generate_recording(num_channels=4, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) -def _check_mock_recording(rec): +def _check_mock_recording(rec, check_extra_data=None): from spikeinterface.core import BaseRecording assert isinstance(rec, BaseRecording), type(rec).__name__ @@ -91,7 +94,7 @@ def _build_recording_with_properties(): return rec -def _check_recording_with_properties(rec): +def _check_recording_with_properties(rec, check_extra_data=None): assert rec.get_num_channels() == 4 assert list(rec.get_property("quality")) == ["good", "good", "bad", "good"] assert rec.get_annotation("experimenter") == "test" @@ -105,11 +108,9 @@ def _build_recording_with_probe(): rec = generate_recording(num_channels=8, durations=[DEFAULT_DURATION], sampling_frequency=30000.0, seed=0) probe = generate_linear_probe(num_elec=8) probe.set_device_channel_indices(np.arange(8)) - if parse(si_version) <= parse("0.105.0"): - rec_with_probe = rec.set_probe(probe, in_place=False) # old API returns a new recording; portable across versions - else: - rec_with_probe = rec.set_probe(probe) # new API returns a new recording; portable across versions - return rec_with_probe + # old API (in_place=False by default) returns a new recording; new API is always in place and returns None + rec_with_probe = rec.set_probe(probe) + return rec_with_probe if rec_with_probe is not None else rec def _build_recording_with_timestamps(): @@ -119,14 +120,14 @@ def _build_recording_with_timestamps(): return rec -def _check_recording_with_timestamps(rec): +def _check_recording_with_timestamps(rec, check_extra_data=None): import numpy as np expected_times = np.arange(int(DEFAULT_DURATION * 30000)) / 30000.0 + 100 times = rec.get_times(segment_index=0) assert np.allclose(times, expected_times) -def _check_recording_with_probe(rec): +def _check_recording_with_probe(rec, check_extra_data=None): import numpy as np assert rec.get_num_channels() == 8 @@ -152,14 +153,12 @@ def _build_recording_with_interleaved_probes(): probegroup.add_probe(probe1) probegroup.set_global_device_channel_indices([0, 2, 4, 6, 1, 3, 5, 7]) # Interleave the two probes' channels: channel i alternates between probe0 and probe1. - if parse(si_version) <= parse("0.105.0"): - rec_with_probe = rec.set_probegroup(probegroup, in_place=False) # old API returns a new recording; portable across versions - else: - rec_with_probe = rec.set_probegroup(probegroup) # new API returns a new recording; portable across versions - return rec_with_probe + # old API (in_place=False by default) returns a new recording; new API is always in place and returns None + rec_with_probe = rec.set_probegroup(probegroup) + return rec_with_probe if rec_with_probe is not None else rec -def _check_recording_with_interleaved_probes(rec): +def _check_recording_with_interleaved_probes(rec, check_extra_data=None): import numpy as np assert rec.get_num_channels() == 8 @@ -192,7 +191,7 @@ def _build_preprocessed_chain(): return common_reference(scale(rec, gain=2.0)) -def _check_preprocessed_chain(rec): +def _check_preprocessed_chain(rec, check_extra_data=None): # The outer wrapper and the recursive parent chain must both reload (the kwargs # embed the parent recording dict, so this exercises recursive deserialization). assert type(rec).__name__ == "CommonReferenceRecording", type(rec).__name__ @@ -206,7 +205,7 @@ def _build_sorting(): return generate_sorting(num_units=5, sampling_frequency=30000.0, durations=[DEFAULT_DURATION]) -def _check_sorting(sorting): +def _check_sorting(sorting, check_extra_data=None): assert sorting.get_num_units() == 5, sorting.get_num_units() assert sorting.get_num_segments() == 1 spike_train = sorting.get_unit_spike_train(sorting.unit_ids[0], segment_index=0) @@ -223,12 +222,41 @@ def _build_sorting_with_properties(): return sorting -def _check_sorting_with_properties(sorting): +def _check_sorting_with_properties(sorting, check_extra_data=None): assert sorting.get_num_units() == 4 assert list(sorting.get_property("quality")) == ["good", "good", "bad", "good"] assert sorting.get_annotation("experimenter") == "test" +def _build_sorting_analyzer_with_extensions(): + from spikeinterface.core import generate_ground_truth_recording, create_sorting_analyzer + recording, sorting = generate_ground_truth_recording(durations=[10, 5], num_channels=16, num_units=5, seed=0) + analyzer = create_sorting_analyzer(sorting, recording) + extensions = analyzer.get_computable_extensions() + analyzer.compute(extensions, n_jobs=-1) + # for SortingAnalyzer, we also save a JSON file with extension as + # keys and date entries as values to check that everything is reloaded correctly + check_extra_data = {} + for ext_name in analyzer.extensions: + ext = analyzer.get_extension(ext_name) + check_extra_data[ext_name] = list(ext.data.keys()) + return analyzer, check_extra_data + + +def _check_sorting_analyzer_with_extensions(analyzer, check_extra_data=None): + extensions = analyzer.get_saved_extension_names() + for extension in extensions: + extension = analyzer.get_extension(extension) # just check it loads without error + ext_data = extension.get_data() + assert ext_data is not None, f"extension {extension} returned None data" + if check_extra_data is not None: + # Check that the extension data keys match what was saved in the fixture JSON. + for ext_name, expected_keys in check_extra_data.items(): + ext = analyzer.get_extension(ext_name) + actual_keys = set(ext.data.keys()) + assert actual_keys == set(expected_keys), f"extension {ext_name} keys mismatch: {actual_keys} != {set(expected_keys)}" + + OBJECTS = [ { "id": "noise_generator_recording", @@ -284,4 +312,10 @@ def _check_sorting_with_properties(sorting): "check": _check_sorting_with_properties, "formats": ["numpy_folder", "zarr"], }, + { + "id": "sorting_analyzer_with_extensions", + "build": _build_sorting_analyzer_with_extensions, + "check": _check_sorting_analyzer_with_extensions, # no check; the extensions are computed and stored, but the analyzer itself is not serialized + "formats": ["binary_folder", "zarr"], + }, ] diff --git a/.github/scripts/serialization/serialize_objects.py b/.github/scripts/serialization/serialize_objects.py index c4e9326f08..9e9ebe9bfe 100644 --- a/.github/scripts/serialization/serialize_objects.py +++ b/.github/scripts/serialization/serialize_objects.py @@ -12,9 +12,13 @@ """ import sys +import json from pathlib import Path -import spikeinterface +from spikeinterface import __version__ as si_version +from spikeinterface import SortingAnalyzer +from spikeinterface.core.base import BaseExtractor + sys.path.insert(0, str(Path(__file__).parent)) from objects import OBJECTS, FIXTURE_SUFFIX # noqa: E402 @@ -22,25 +26,43 @@ out_dir = Path(sys.argv[1]) if len(sys.argv) > 1 else Path("serialization_fixtures") out_dir.mkdir(parents=True, exist_ok=True) # do not rmtree an arbitrary path; overwrite in place -print(f"Generating serialization fixtures with spikeinterface {spikeinterface.__version__}") +print(f"Generating serialization fixtures with spikeinterface {si_version}") for entry in OBJECTS: - obj = entry["build"]() + objs = entry["build"]() + if isinstance(objs, tuple) and len(objs) == 2 and isinstance(objs[1], dict): + obj, check_extra_data = objs + else: + obj = objs + check_extra_data = None for fmt in entry["formats"]: dest = out_dir / f"{entry['id']}{FIXTURE_SUFFIX[fmt]}" - if fmt == "json": - obj.dump_to_json(dest) - elif fmt == "pickle": - obj.dump_to_pickle(dest) - elif fmt == "binary": - obj.save(folder=dest, format="binary", overwrite=True) - elif fmt == "binary_parallel": - obj.save(folder=dest, format="binary", overwrite=True, n_jobs=2) - elif fmt == "numpy_folder": - obj.save(folder=dest, format="numpy_folder", overwrite=True) - elif fmt == "zarr": - obj.save(folder=dest, format="zarr", overwrite=True) - elif fmt == "zarr_parallel": - obj.save(folder=dest, format="zarr", overwrite=True, n_jobs=2) + if isinstance(obj, BaseExtractor): + if fmt == "json": + obj.dump_to_json(dest) + elif fmt == "pickle": + obj.dump_to_pickle(dest) + elif fmt == "binary": + obj.save(folder=dest, format="binary", overwrite=True) + elif fmt == "binary_parallel": + obj.save(folder=dest, format="binary", overwrite=True, n_jobs=2) + elif fmt == "numpy_folder": + obj.save(folder=dest, format="numpy_folder", overwrite=True) + elif fmt == "zarr": + obj.save(folder=dest, format="zarr", overwrite=True) + elif fmt == "zarr_parallel": + obj.save(folder=dest, format="zarr", overwrite=True, n_jobs=2) + elif isinstance(obj, SortingAnalyzer): + if dest.is_dir(): + import shutil + shutil.rmtree(dest) + obj.save_as(folder=dest, format=fmt) + else: + raise TypeError(f"{entry['id']}: build() returned unsupported object {type(obj).__name__}") print(f" wrote {dest.name} ({fmt})") + if check_extra_data is not None: + json_dest = out_dir / f"{entry['id']}.json" + with open(json_dest, "w") as f: + json.dump(check_extra_data, f) + print(f" wrote {json_dest.name} (extra data)") print(f"Fixtures written to: {out_dir.resolve()}") diff --git a/.github/scripts/serialization/test_cross_version_compatibility.py b/.github/scripts/serialization/test_cross_version_compatibility.py index cd485f7a1e..bdb7d019ba 100644 --- a/.github/scripts/serialization/test_cross_version_compatibility.py +++ b/.github/scripts/serialization/test_cross_version_compatibility.py @@ -9,6 +9,7 @@ """ import os +import json from pathlib import Path import pytest @@ -27,4 +28,8 @@ def test_load_old_serialized_object(entry, fmt): fixture = FIXTURES_DIR / f"{entry['id']}{FIXTURE_SUFFIX[fmt]}" assert fixture.exists(), f"missing fixture {fixture}" obj = load(fixture) - entry["check"](obj) + check_extra_data = None + if (FIXTURES_DIR / f"{entry['id']}.json").exists(): + with open(FIXTURES_DIR / f"{entry['id']}.json", "r") as f: + check_extra_data = json.load(f) + entry["check"](obj, check_extra_data=check_extra_data) diff --git a/.github/workflows/all-tests.yml b/.github/workflows/all-tests.yml index e4a2888161..97d64a9b3f 100644 --- a/.github/workflows/all-tests.yml +++ b/.github/workflows/all-tests.yml @@ -25,7 +25,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ["3.10", "3.14"] # Lower and higher versions we support + python-version: ["3.12", "3.14"] # Lower and higher versions we support os: [macos-latest, windows-latest, ubuntu-latest] steps: - uses: actions/checkout@v6 diff --git a/.github/workflows/caches_cron_job.yml b/.github/workflows/caches_cron_job.yml index f904b4ba99..38d0f5bcbe 100644 --- a/.github/workflows/caches_cron_job.yml +++ b/.github/workflows/caches_cron_job.yml @@ -15,10 +15,10 @@ jobs: steps: - uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: "3.12" - uses: astral-sh/setup-uv@v7 with: - python-version: "3.10" + python-version: "3.12" enable-cache: true ignore-nothing-to-cache: true # some runs do not require building a cache (ie macOS) this says this is okay - name: Create the directory to store the data diff --git a/.github/workflows/core-test.yml b/.github/workflows/core-test.yml new file mode 100644 index 0000000000..b65606588b --- /dev/null +++ b/.github/workflows/core-test.yml @@ -0,0 +1,48 @@ +name: Testing core + +on: + pull_request: + types: [synchronize, opened, reopened] + branches: + - main + +concurrency: # Cancel previous workflows on the same pull request + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + build-and-test: + name: Test on ${{ matrix.os }} OS + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: ["ubuntu-latest", "macos-latest", "windows-latest"] + steps: + - uses: actions/checkout@v6 + - uses: actions/setup-python@v6 + with: + python-version: "3.12" + - uses: astral-sh/setup-uv@v7 + with: + python-version: "3.12" + enable-cache: true + - name: Install dependencies + run: | + git config --global user.email "CI@example.com" + git config --global user.name "CI Almighty" + uv pip install --system -e . --group test-core + - name: Test core with pytest + run: | + pytest -m "core" -vv -ra --durations=0 --durations-min=0.001 | tee report.txt; test $? -eq 0 || exit 1 + shell: bash # Necessary for pipeline to work on windows + - name: Build test summary + run: | + uv pip install --system pandas + uv pip install --system tabulate + echo "# Timing profile of core tests in ${{matrix.os}}" >> $GITHUB_STEP_SUMMARY + # Outputs markdown summary to standard output + python ./.github/scripts/build_job_summary.py report.txt >> $GITHUB_STEP_SUMMARY + cat $GITHUB_STEP_SUMMARY + rm report.txt + shell: bash # Necessary for pipeline to work on windows diff --git a/.github/workflows/cross_version_serialization.yml b/.github/workflows/cross_version_serialization.yml index 4bafaed4e1..b7d7b88c61 100644 --- a/.github/workflows/cross_version_serialization.yml +++ b/.github/workflows/cross_version_serialization.yml @@ -40,7 +40,7 @@ jobs: steps: - uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: "3.12" - id: set env: # Support floor: oldest minor to test. Below this, installs fail on the CI @@ -63,7 +63,8 @@ jobs: candidates = [v for v in available if not v.is_prerelease and (v.major, v.minor) >= (floor.major, floor.minor)] minor_releases = [max(group) for _, group in groupby(candidates, key=lambda v: (v.major, v.minor))] - print(json.dumps([str(v) for v in minor_releases])) + # always round-trip against the branch itself as well + print(json.dumps([str(v) for v in minor_releases] + ['dev'])) ") echo "list=$versions_to_test" >> "$GITHUB_OUTPUT" @@ -85,21 +86,34 @@ jobs: - name: Set up Python uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: "3.12" - name: Set up uv uses: astral-sh/setup-uv@v7 with: - python-version: "3.10" + python-version: "3.12" enable-cache: false - name: Generate fixtures with spikeinterface ${{ matrix.si-version }} + env: + SI_VERSION: ${{ matrix.si-version }} run: | - uv run --isolated --no-project --python 3.10 \ - --with "spikeinterface[core]==${{ matrix.si-version }}" \ + # "dev" means the checked-out branch; anything else is a release from PyPI. + # [full] is needed because the SortingAnalyzer fixture computes every + # extension (scikit-learn, numba, ...). + if [ "$SI_VERSION" = "dev" ]; then + spec=".[full]" + else + spec="spikeinterface[full]==$SI_VERSION" + fi + uv run --isolated --no-project --python 3.12 \ + --with "$spec" \ + --with "numpy<2.5" \ python .github/scripts/serialization/serialize_objects.py "$SI_SERIALIZATION_FIXTURES_DIR" - name: Load fixtures with the branch run: | - uv pip install --system -e . --group test-core + # [full] here too: loading the analyzer extensions unpickles e.g. the + # scikit-learn PCA models written at generation time + uv pip install --system -e ".[full]" "numpy<2.5" --group test-core pytest .github/scripts/serialization/test_cross_version_compatibility.py -v diff --git a/.github/workflows/deepinterpolation.yml b/.github/workflows/deepinterpolation.yml index 1d8d246575..c5a3af31b9 100644 --- a/.github/workflows/deepinterpolation.yml +++ b/.github/workflows/deepinterpolation.yml @@ -1,10 +1,8 @@ name: Testing deepinterpolation +# Manual only — deepinterpolation requires Python 3.12, incompatible with 3.12+ required by Zarr 3.0.0+ on: - pull_request: - types: [synchronize, opened, reopened] - branches: - - main + workflow_dispatch: concurrency: # Cancel previous workflows on the same pull request group: ${{ github.workflow }}-${{ github.ref }} @@ -22,10 +20,10 @@ jobs: - uses: actions/checkout@v6 - uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: "3.12" - uses: astral-sh/setup-uv@v7 with: - python-version: "3.10" + python-version: "3.12" enable-cache: false - name: Get changed files id: changed-files diff --git a/.github/workflows/test_containers_docker.yml b/.github/workflows/test_containers_docker.yml index 65bb2f90e9..1b794d4eb7 100644 --- a/.github/workflows/test_containers_docker.yml +++ b/.github/workflows/test_containers_docker.yml @@ -14,10 +14,10 @@ jobs: - uses: actions/checkout@v6 - uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: '3.12' - uses: astral-sh/setup-uv@v7 with: - python-version: "3.10" + python-version: '3.12' enable-cache: true - name: Python version run: python --version diff --git a/.github/workflows/test_containers_singularity.yml b/.github/workflows/test_containers_singularity.yml index c39b98357b..184c50aeff 100644 --- a/.github/workflows/test_containers_singularity.yml +++ b/.github/workflows/test_containers_singularity.yml @@ -15,10 +15,10 @@ jobs: - uses: actions/checkout@v6 - uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: '3.12' - uses: astral-sh/setup-uv@v7 with: - python-version: "3.10" + python-version: '3.12' enable-cache: true - uses: eWaterCycle/setup-singularity@v7 with: diff --git a/.github/workflows/test_imports.yml b/.github/workflows/test_imports.yml index d3d5b5dbb1..d2bd1fadcf 100644 --- a/.github/workflows/test_imports.yml +++ b/.github/workflows/test_imports.yml @@ -22,10 +22,10 @@ jobs: - uses: actions/checkout@v6 - uses: actions/setup-python@v6 with: - python-version: "3.10" + python-version: "3.12" - uses: astral-sh/setup-uv@v7 with: - python-version: "3.10" + python-version: "3.12" enable-cache: true - name: Install Spikeinterface with only core dependencies run: | diff --git a/conftest.py b/conftest.py index e4a08c170a..6c4bfbef3a 100644 --- a/conftest.py +++ b/conftest.py @@ -47,16 +47,10 @@ def pytest_collection_modifyitems(config, items): This function marks (in the pytest sense) the tests according to their name and file_path location Marking them in turn allows the tests to be run by using the pytest -m marker_name option. """ - - from spikeinterface.core.core_tools import _is_zarr_write_supported - rootdir = Path(config.rootdir) modules_location = rootdir / "src" / "spikeinterface" for item in items: # TODO: remove once writing to zarr is supported with zarr>=3 - if item.get_closest_marker("requires_zarr_write") and not _is_zarr_write_supported(): - item.add_marker(pytest.mark.skip(reason="Writing to zarr is not supported yet with zarr>=3")) - if config.getoption("--mp-context") is not None and item.name == "test_global_job_kwargs": item.add_marker(pytest.mark.skip(reason="--mp-context changes the default job kwargs")) diff --git a/doc/modules/core.rst b/doc/modules/core.rst index 297f22e945..d8e7e42ed7 100644 --- a/doc/modules/core.rst +++ b/doc/modules/core.rst @@ -644,8 +644,8 @@ used when creating a :py:class:`~spikeinterface.core.SortingAnalyzer` which caus .. _save_load: -Saving, loading, and compression --------------------------------- +Saving and loading +------------------ The Base SpikeInterface objects (:py:class:`~spikeinterface.core.BaseRecording`, :py:class:`~spikeinterface.core.BaseSorting`, and @@ -678,6 +678,9 @@ This saving/loading features enables us to store SpikeInterface objects efficien # save sorting to NPZ sorting_saved = sorting.save(folder="sorting") +Compression +----------- + **NOTE:** the Zarr format by default applies data compression with :code:`Blosc.Zstandard` codec with BIT shuffling. Any other Zarr-compatible compressors and filters can be applied using the :code:`compressor` and :code:`filters` arguments. For example, in this case we apply `LZMA `_ @@ -692,9 +695,22 @@ and use a `Delta `_ filte filters = [Delta(dtype="int16")] recording_custom_comp = recording.save(folder="recording", format="zarr", - compressor=compressor, filters=filters, + compressors=compressor, filters=filters, **job_kwargs) +Zarr stores each dataset as a grid of **chunks**: the unit of compression and of reading, so a chunk is always +decompressed as a whole. When saving a recording, the chunk size along time is set by the job parameters +(e.g. :code:`chunk_duration="1s"`). Zarr v3 also supports **shards**, which group many chunks in a single file/object. +Chunks are still compressed and read independently, but sharding drastically reduces the number of files, which is +useful on cluster file systems and cloud storage. A shard must be a multiple of the chunk size, which is set with +:code:`shard_factor` (shard size = :code:`shard_factor` x chunk size). For example, 1 s chunks and 20 s shards: + +.. code-block:: python + + job_kwargs = dict(n_jobs=8, chunk_duration="1s") + recording_sharded = recording.save(folder="recording", format="zarr", + shard_factor=20, **job_kwargs) + Parallel processing and job_kwargs ---------------------------------- diff --git a/pyproject.toml b/pyproject.toml index 32442a7e72..7b3a1fa385 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -7,7 +7,7 @@ authors = [ ] description = "Python toolkit for analysis, visualization, and comparison of spike sorting output" readme = "README.md" -requires-python = ">=3.10" +requires-python = ">=3.12" classifiers = [ "Programming Language :: Python :: 3 :: Only", "License :: OSI Approved :: MIT License", @@ -24,17 +24,16 @@ dependencies = [ "numpy>=2.3.2;python_version>='3.14'", "threadpoolctl>=3.0.0", "tqdm", - "zarr>=2.18,<3;python_version<'3.14'", - "zarr>=3;python_version>='3.14'", # Zarr writing is not yet supported with zarr 3 + "zarr>=3,<4", # for release we need pypi, so this needs to be commented "neo @ git+https://github.com/NeuralEnsemble/python-neo.git", - "probeinterface @ git+https://github.com/SpikeInterface/probeinterface.git", + # for testing we need probeinterface Zarr v3 branch, so this needs to be commented + #"probeinterface @ git+https://github.com/SpikeInterface/probeinterface.git", + "probeinterface @ git+https://github.com/alejoe91/probeinterface.git@zarrv3", # "neo>=0.14.5", # "probeinterface>=0.4.0", "packaging", "pydantic", - "numcodecs<0.16.0;python_version<'3.14'", # For supporting zarr < 3 - "numcodecs>=0.16.5;python_version>='3.14'", ] [build-system] @@ -69,7 +68,7 @@ changelog = "https://spikeinterface.readthedocs.io/en/latest/whatisnew.html" extractors = [ "MEArec>=1.8", "pynwb>=2.6.0", - "hdmf-zarr>=0.11.0", + "hdmf-zarr>=0.14.0", "pyedflib>=0.1.30", "lxml", # lxml for neuroscope "scipy", @@ -85,9 +84,9 @@ streaming_extractors = [ "fsspec", "aiohttp", "requests", - "hdmf-zarr>=0.11.0", + "hdmf-zarr>=0.14.0", "remfile", - "s3fs" + "s3fs>=2025.7.0" ] @@ -367,7 +366,6 @@ markers = [ "widgets", "sortingcomponents", "streaming_extractors: extractors that require streaming such as ross and fsspec", - "requires_zarr_write: tests that write to zarr, skipped with zarr>=3 where writing is not supported yet", ] filterwarnings =[ 'ignore:.*distutils Version classes are deprecated.*:DeprecationWarning', diff --git a/src/spikeinterface/core/base.py b/src/spikeinterface/core/base.py index 7fa7fb24c4..9b5a874e58 100644 --- a/src/spikeinterface/core/base.py +++ b/src/spikeinterface/core/base.py @@ -873,7 +873,6 @@ def save_to_zarr( folder=None, overwrite=False, storage_options=None, - channel_chunk_size=None, verbose=True, **save_kwargs, ): @@ -894,76 +893,6 @@ def save_to_zarr( # we keep the default format for recording and sorting like in old version return self.save(folder=folder, format="zarr", verbose=verbose, **save_kwargs) - # """ - # Save extractor to zarr. - - # Parameters - # ---------- - # name: str or None, default: None - # Name of the subfolder in get_global_tmp_folder() - # If "name" is given, "folder" must be None - # folder: str, Path, or None, default: None - # The folder used to save the zarr output. If the folder does not have a ".zarr" suffix, - # it will be automatically appended - # overwrite: bool, default: False - # If True, the folder is removed if it already exists - # storage_options: dict or None, default: None - # Storage options for zarr `store`. E.g., if "s3://" or "gcs://" they can provide authentication methods, etc. - # For cloud storage locations, this should not be None (in case of default values, use an empty dict) - # channel_chunk_size: int or None, default: None - # Channels per chunk (only for BaseRecording) - # compressor: numcodecs.Codec or None, default: None - # Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used - # filters: list[numcodecs.Codec] or None, default: None - # Global filters for zarr (global) - # compressor_by_dataset: dict or None, default: None - # Optional compressor per dataset: - # - traces - # - times - # If None, the global compressor is used - # filters_by_dataset: dict or None, default: None - # Optional filters per dataset: - # - traces - # - times - # If None, the global filters are used - # verbose: bool, default: True - # If True, the output is verbose - # auto_cast_uint: bool, default: True - # If True, unsigned integers are cast to signed integers to avoid issues with zarr (only for BaseRecording) - - # Returns - # ------- - # cached: ZarrExtractor - # Saved copy of the extractor. - # """ - # from .zarrextractors import read_zarr - - # save_kwargs.pop("format", None) - - # if folder is None: - # 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 / name).with_suffix(".zarr") - # if verbose: - # print(f"Saving to zarr_path={zarr_path}") - # else: - # 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 - - # 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 - # save_kwargs["channel_chunk_size"] = channel_chunk_size - # cached = self._save(format="zarr", verbose=verbose, **save_kwargs) - # cached = read_zarr(zarr_path) - - # return cached - def _load_extractor_from_dict(dic) -> "BaseExtractor": """ diff --git a/src/spikeinterface/core/baserecording.py b/src/spikeinterface/core/baserecording.py index 772c026ccc..c2769f4846 100644 --- a/src/spikeinterface/core/baserecording.py +++ b/src/spikeinterface/core/baserecording.py @@ -5,10 +5,11 @@ import numpy as np from probeinterface import read_probeinterface, write_probeinterface +from .base import BaseExtractor from .time_series import TimeSeriesSegment, TimeSeries from .baserecordingsnippets import BaseRecordingSnippets from .core_tools import convert_bytes_to_str, convert_seconds_to_str -from .job_tools import split_job_kwargs +from .job_tools import split_job_kwargs, _shared_job_kwargs_doc class BaseRecording(BaseRecordingSnippets, TimeSeries): @@ -346,6 +347,19 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): For cloud storage locations, this should not be None (in case of default values, use an empty dict) - channel_chunk_size: int or None, default: None Channels per chunk (only for BaseRecording) + - chunks: tuple | None, default: None + # TOOD: modify this + Chunks for the traces dataset. If None, chunking is applied to the time dimension only and it is + determined by the job_kwargs chunks ("chunk_size" or "chunk_duration"). + If `chunks` is not None, it needs to be a tuple of length 2 with the chunk size for the time and channel + dimensions respectively and `channel_chunk_size` should not be specified. + - shard_factor: int | tuple[int] | None, default: None + If integer, the shard size will be set to chunk_size * shard_factor in the first dimension (time). + If tuple, the shard_factor to be applied to each dimension. Note that `shard_factor` cannot + be specified together with `shards`. + -shards: tuple | None, default: None + Number of shard size. If None, no sharding is done. Note that shards dimensions need to be larger than + chunk dimensions (if chunks is not None) and that sharding is only done on the first dimension. - compressor: numcodecs.Codec or None, default: None Global compressor. If None, Blosc-zstd, level 5, with bit shuffle is used - filters: list[numcodecs.Codec] or None, default: None @@ -364,6 +378,7 @@ def save(self, format="binary", verbose: bool = False, **save_kwargs): - times If None, the global filters are used + * "memory" format: - sharedmem : bool, default: True If True, the recording is saved in shared memory. If False, it is saved as diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index 58d53dca91..519771144c 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -763,27 +763,6 @@ def measure_memory_allocation(measure_in_process: bool = True) -> float: return memory -def _is_zarr_write_supported() -> bool: - """ - Whether writing to zarr is supported with the installed zarr, which is not yet the case for zarr>=3. - """ - return int(zarr.__version__.split(".")[0]) < 3 - - -def _check_zarr_write_is_supported() -> None: - """ - Raise an informative error when writing to zarr with zarr>=3, which is not supported yet. - - zarr>=3 is what gets installed on Python 3.14, where reading existing zarr folders works but writing does not. - """ - if not _is_zarr_write_supported(): - raise NotImplementedError( - f"Writing to zarr is not supported yet with zarr {zarr.__version__}, which is the version installed " - "on Python 3.14. Use Python 3.13 or lower to save in zarr format, or save in another format such as " - "'binary_folder'." - ) - - def is_path_remote(path: str | Path) -> bool: """ Returns True if the path is a remote path (e.g., s3:// or gcs://). diff --git a/src/spikeinterface/core/loading.py b/src/spikeinterface/core/loading.py index 9c06e95300..8cc679d95e 100644 --- a/src/spikeinterface/core/loading.py +++ b/src/spikeinterface/core/loading.py @@ -270,11 +270,15 @@ def _guess_object_from_zarr(zarr_folder): return _guess_object_from_dict(spikeinterface_info) # here it is the old fashion and a bit ambiguous - if "templates_array" in zarr_root.keys(): + # get_zarr_group_keys is used to be compatible with both zarr v2 and zarr v3 groups + from .zarr_tools import get_zarr_group_keys + + root_keys = get_zarr_group_keys(zarr_root) + if "templates_array" in root_keys: return "Templates" - elif "channel_ids" in zarr_root.keys() and "unit_ids" not in zarr_root.keys(): + elif "channel_ids" in root_keys and "unit_ids" not in root_keys: return "Recording" - elif "unit_ids" in zarr_root.keys() and "channel_ids" not in zarr_root.keys(): + elif "unit_ids" in root_keys and "channel_ids" not in root_keys: return "Sorting" diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index c66d68a2e4..afa786e8e8 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -12,7 +12,7 @@ from spikeinterface.core import BaseRecording, get_chunk_with_margin from spikeinterface.core.job_tools import TimeSeriesChunkExecutor, fix_job_kwargs, _shared_job_kwargs_doc from spikeinterface.core import get_channel_distances -from spikeinterface.core.core_tools import ms_to_samples, samples_to_ms, _check_zarr_write_is_supported +from spikeinterface.core.core_tools import ms_to_samples, samples_to_ms class PipelineNode: @@ -1098,7 +1098,6 @@ def __init__( from spikeinterface.core.zarrextractors import get_default_zarr_compressor - _check_zarr_write_is_supported() if compressor == "default": compressor = get_default_zarr_compressor() self.compressor = compressor @@ -1183,12 +1182,12 @@ def __call__(self, res): # pick the number of rows per chunk to target ~zarr_target_chunk_bytes per chunk row_nbytes = int(np.prod(trailing_shape, dtype="int64")) * buf.dtype.itemsize chunk0 = max(1, self.zarr_target_chunk_bytes[name] // max(1, row_nbytes)) - self.arrays[i_name] = root.create_dataset( + self.arrays[i_name] = root.create_array( name=internal_path, shape=(0,) + trailing_shape, chunks=(chunk0,) + trailing_shape, dtype=buf.dtype, - compressor=self.compressor, + compressors=self.compressor, overwrite=True, ) self.arrays[i_name].append(buf, axis=0) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index b702237445..31cc78c5c1 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -27,7 +27,6 @@ retrieve_importing_provenance, is_path_remote, clean_zarr_folder_name, - _check_zarr_write_is_supported, ) from .sorting_tools import ( generate_unit_ids_for_merge_group, @@ -40,7 +39,14 @@ from .numpyextractors import NumpySorting from .sparsity import ChannelSparsity, estimate_sparsity from .sortingfolder import NumpyFolderSorting -from .zarrextractors import get_default_zarr_compressor, ZarrSortingExtractor, super_zarr_open, _write_object_array +from .zarrextractors import get_default_zarr_compressor, ZarrSortingExtractor, super_zarr_open +from .zarr_tools import ( + iterate_zarr_group, + get_zarr_attr_or_legacy_object, + is_sklearn_estimator, + save_sklearn_model_to_zarr_group, + load_sklearn_model_from_zarr_group, +) from .node_pipeline import run_node_pipeline # Typing hints @@ -1014,10 +1020,9 @@ def load_from_binary_folder( def _get_zarr_root(self, mode="r+"): assert mode in ("r+", "a", "r"), "mode must be 'r+', 'a' or 'r'" - if mode != "r": - _check_zarr_write_is_supported() storage_options = self._backend_options.get("storage_options", {}) + zarr_root = super_zarr_open(self.folder, mode=mode, storage_options=storage_options) return zarr_root @@ -1037,7 +1042,6 @@ def create_zarr( ) -> "SortingAnalyzer": from .zarrextractors import add_sorting_to_zarr_group, ZarrSortingExtractor - _check_zarr_write_is_supported() is_remote = is_path_remote(folder) if not is_remote: folder = clean_zarr_folder_name(folder) @@ -1049,8 +1053,13 @@ def create_zarr( storage_options = backend_options.get("storage_options", {}) saving_options = backend_options.get("saving_options", {}) + if not is_remote: + storage_options_kwargs = {} + else: + storage_options_kwargs = {"storage_options": storage_options} + # Create zarr root group (and subgroups) - zarr_root = zarr.open(folder, mode="w", storage_options=storage_options) + zarr_root = zarr.open(folder, mode="w", **storage_options_kwargs) sorting_group = zarr_root.create_group("sorting") # for sorting output recording_info_group = zarr_root.create_group("recording_info") # rec_attributes and probe group zarr_root.create_group("extensions") # used later @@ -1070,23 +1079,17 @@ def create_zarr( # Save sparsity if sparsity is not None: - zarr_root.create_dataset("sparsity_mask", data=sparsity.mask, **saving_options) + zarr_root.create_array("sparsity_mask", data=sparsity.mask, **saving_options) # Dump recording provenance if recording is not None: rec_dict = recording.to_dict(relative_to=relative_to, recursive=True) if recording.check_serializability("json"): - _write_object_array(zarr_root, "recording", check_json(rec_dict), codec="json") - elif recording.check_serializability("pickle"): - try: - _write_object_array(zarr_root, "recording", rec_dict, codec="pickle") - except: - warnings.warn( - "Failed to serialize recording with Pickle Codec! " - "The recording link will be lost for future load" - ) + # In zarr v3, store JSON-serializable data in attributes instead of using object_codec + zarr_root.attrs["recording"] = check_json(rec_dict) else: warnings.warn("The Recording is not serializable! The recording link will be lost for future load") + else: assert rec_attributes is not None, "recording or rec_attributes must be provided" warnings.warn("Recording not provided, instantiating SortingAnalyzer in recordingless mode.") @@ -1139,6 +1142,10 @@ def load_from_zarr( backend_options = {} if backend_options is None else backend_options storage_options = backend_options.get("storage_options", {}) + if not is_path_remote(str(folder)): + storage_options_kwargs = {} + else: + storage_options_kwargs = storage_options # Open root group zarr_root = super_zarr_open(str(folder), mode="r", storage_options=storage_options) @@ -1149,7 +1156,7 @@ def load_from_zarr( si_info = zarr_root.attrs["spikeinterface_info"] if parse(si_info["version"]) < parse("0.101.1"): try: - zarr_root_a = zarr.open(str(folder), mode="a", storage_options=storage_options) + zarr_root_a = zarr.open(str(folder), mode="a", **storage_options_kwargs) zarr.consolidate_metadata(zarr_root_a.store) except: warnings.warn( @@ -1185,9 +1192,9 @@ def load_from_zarr( # Load recording (if available) if recording is None: - rec_field = zarr_root.get("recording") - if rec_field is not None: - rec_dict = rec_field[0] + # In zarr v3, recording is stored in attributes (in zarr v2 it was an object array) + rec_dict = get_zarr_attr_or_legacy_object(zarr_root, "recording") + if rec_dict is not None: try: recording = load(rec_dict, base_folder=folder) except: @@ -1299,7 +1306,7 @@ def set_sorting_property( if key in zarr_root["sorting"]["properties"]: zarr_root["sorting"]["properties"][key][:] = prop_values else: - zarr_root["sorting"]["properties"].create_dataset(name=key, data=prop_values, compressor=None) + zarr_root["sorting"]["properties"].create_array(name=key, data=prop_values, compressors=None) # IMPORTANT: we need to re-consolidate the zarr store! zarr.consolidate_metadata(zarr_root.store) @@ -2172,12 +2179,15 @@ def get_sorting_provenance(self): elif self.format == "zarr": zarr_root = self._get_zarr_root(mode="r") sorting_provenance = None - sort_dict = zarr_root["sorting"].attrs.get("provenance", None) - if sort_dict is None: + # In zarr v3, sorting_provenance is stored in attributes (in zarr v2 it was an object array) + sort_dict = get_zarr_attr_or_legacy_object(zarr_root, "provenance") + if sort_dict is not None: # Legacy - if "sorting_provenance" in zarr_root.keys(): - sort_dict = zarr_root["sorting_provenance"][0] + sort_dict = get_zarr_attr_or_legacy_object(zarr_root, "sorting_provenance") if sort_dict is not None: + # try-except here is because it's not required to be able + # to load the sorting provenance, as the user might have deleted + # the original sorting folder try: sorting_provenance = load(sort_dict, base_folder=self.folder) except: @@ -2623,10 +2633,16 @@ def get_saved_extension_names(self): elif self.format == "zarr": zarr_root = self._get_zarr_root(mode="r") - if "extensions" in zarr_root.keys(): + # Avoid iterating zarr_root.keys() because legacy v2 stores may contain + # object-dtype arrays (e.g. "recording", "sorting_provenance") that zarr v3 + # cannot parse, causing ValueError on enumeration. + try: extension_group = zarr_root["extensions"] - for extension_name in extension_group.keys(): - if "params" in extension_group[extension_name].attrs.keys(): + except KeyError: + extension_group = None + if extension_group is not None: + for extension_name, extension in iterate_zarr_group(extension_group): + if "params" in extension.attrs.keys(): saved_extension_names.append(extension_name) else: @@ -3369,20 +3385,26 @@ def load_data(self, lazy=False): self.set_data(ext_data_name, ext_data) elif self.format == "zarr": extension_group = self._get_zarr_extension_group(mode="r") - for ext_data_name in extension_group.keys(): - ext_data_ = extension_group[ext_data_name] - if "dict" in ext_data_.attrs: - ext_data = ext_data_[0] + # iterate_zarr_group is used to be compatible with both zarr v2 (saved by + # spikeinterface < 0.105) and zarr v3 groups + for ext_data_name, ext_data_ in iterate_zarr_group(extension_group): + # In zarr v3, check if it's a group with dict_data attribute + if "dict_data" in ext_data_.attrs: + ext_data = ext_data_.attrs["dict_data"] + elif "sklearn_model" in ext_data_.attrs: + ext_data = load_sklearn_model_from_zarr_group(ext_data_) elif "dataframe" in ext_data_.attrs: import pandas as pd index = ext_data_["index"] ext_data = pd.DataFrame(index=index) - for col in ext_data_.keys(): + for col, col_data in iterate_zarr_group(ext_data_): if col != "index": - ext_data.loc[:, col] = ext_data_[col][:] + ext_data.loc[:, col] = col_data[:] ext_data = ext_data.convert_dtypes() - elif "object" in ext_data_.attrs: + elif "object" in ext_data_.attrs or "dict" in ext_data_.attrs: + # "dict" is the zarr v2 (spikeinterface < 0.105) flag for dict/list data, + # which was saved as a length-1 object array if lazy: continue ext_data = ext_data_[0] @@ -3484,9 +3506,10 @@ def run(self, save=True, **kwargs): self._save_data() if self.format == "zarr": - zarr.consolidate_metadata(self.sorting_analyzer._get_zarr_root().store) + zarr.consolidate_metadata(self.sorting_analyzer._get_zarr_root(mode="r+").store) def save(self): + self._reset_extension_folder() self._save_params() self._save_importing_provenance() self._save_run_info() @@ -3494,7 +3517,7 @@ def save(self): if self.format == "zarr": - zarr.consolidate_metadata(self.sorting_analyzer._get_zarr_root().store) + zarr.consolidate_metadata(self.sorting_analyzer._get_zarr_root(mode="r+").store) def _save_data(self): if self.format == "memory": @@ -3541,8 +3564,11 @@ def _save_data(self): extension_group = self._get_zarr_extension_group(mode="r+") # if compression is not externally given, we use the default - if "compressor" not in saving_options: - saving_options["compressor"] = get_default_zarr_compressor() + if "compressors" not in saving_options and "compressor" not in saving_options: + saving_options["compressors"] = get_default_zarr_compressor() + if "compressor" in saving_options: + saving_options["compressors"] = [saving_options["compressor"]] + del saving_options["compressor"] for ext_data_name, ext_data in self.data.items(): if isinstance(ext_data, zarr.Array): @@ -3551,35 +3577,40 @@ def _save_data(self): continue if ext_data_name in extension_group: del extension_group[ext_data_name] - if isinstance(ext_data, (dict, list)): - ext_data_ = check_json(ext_data) - _write_object_array(extension_group, ext_data_name, ext_data_, codec="json") - extension_group[ext_data_name].attrs["dict"] = True + if isinstance(ext_data, dict): + # In zarr v3, store dict in a subgroup with attributes + dict_group = extension_group.create_group(ext_data_name) + dict_group.attrs["dict_data"] = check_json(ext_data) elif isinstance(ext_data, np.ndarray): # only save the array if the dataset does not already exist, since it is created directly # by the run_node_pipeline() function in the case of nodepipeline extensions if ext_data_name not in extension_group: - extension_group.create_dataset(name=ext_data_name, data=ext_data, **saving_options) + extension_group.create_array(name=ext_data_name, data=ext_data, **saving_options) elif HAS_PANDAS and isinstance(ext_data, pd.DataFrame): df_group = extension_group.create_group(ext_data_name) # first we save the index indices = ext_data.index.to_numpy() if indices.dtype.kind == "O": indices = indices.astype(str) - df_group.create_dataset(name="index", data=indices) + df_group.create_array(name="index", data=indices) for col in ext_data.columns: col_data = ext_data[col].to_numpy() if col_data.dtype.kind == "O": col_data = col_data.astype(str) - df_group.create_dataset(name=col, data=col_data) + df_group.create_array(name=col, data=col_data) df_group.attrs["dataframe"] = True - else: - # any object + elif is_sklearn_estimator(ext_data): + # sklearn models (e.g. the PCA models) are saved as arrays + attributes, + # so that no pickle is needed (in zarr v2 they were pickled object arrays) try: - _write_object_array(extension_group, ext_data_name, ext_data, codec="pickle") - except: - raise Exception(f"Could not save {ext_data_name} as extension data") - extension_group[ext_data_name].attrs["object"] = True + save_sklearn_model_to_zarr_group(extension_group, ext_data_name, ext_data, **saving_options) + except Exception as e: + warnings.warn(f"Could not save the {ext_data_name} model to zarr, skipping: {e}") + if ext_data_name in extension_group: + del extension_group[ext_data_name] + else: + # any other object + warnings.warn(f"Data type of {ext_data_name} not supported for zarr saving, skipping.") def _reset_extension_folder(self): """ @@ -3673,8 +3704,6 @@ def set_params(self, save=True, **params): def _save_params(self): params_to_save = self.params.copy() - self._reset_extension_folder() - # TODO make sparsity local Result specific # if "sparsity" in params_to_save and params_to_save["sparsity"] is not None: # assert isinstance( diff --git a/src/spikeinterface/core/template.py b/src/spikeinterface/core/template.py index 0f0922e10b..a170dd4e02 100644 --- a/src/spikeinterface/core/template.py +++ b/src/spikeinterface/core/template.py @@ -321,22 +321,20 @@ def add_templates_to_zarr_group(self, zarr_group: "zarr.Group") -> None: to optimize read/write operations for individual units. """ - from .core_tools import _check_zarr_write_is_supported - - _check_zarr_write_is_supported() - # Saves one chunk per unit - arrays_chunk = (1, None, None) - zarr_group.create_dataset("templates_array", data=self.templates_array, chunks=arrays_chunk) - zarr_group.create_dataset("channel_ids", data=self.channel_ids) - zarr_group.create_dataset("unit_ids", data=self.unit_ids) + # In zarr v3, chunks must be a full tuple with actual dimensions + num_units, num_samples, num_channels = self.templates_array.shape + arrays_chunk = (1, num_samples, num_channels) + zarr_group.create_array("templates_array", data=self.templates_array, chunks=arrays_chunk) + zarr_group.create_array("channel_ids", data=self.channel_ids) + zarr_group.create_array("unit_ids", data=self.unit_ids) zarr_group.attrs["sampling_frequency"] = self.sampling_frequency zarr_group.attrs["nbefore"] = self.nbefore zarr_group.attrs["is_in_uV"] = self.is_in_uV if self.sparsity_mask is not None: - zarr_group.create_dataset("sparsity_mask", data=self.sparsity_mask) + zarr_group.create_array("sparsity_mask", data=self.sparsity_mask) if self.probe is not None: probe_group = zarr_group.create_group("probe") @@ -357,9 +355,6 @@ def to_zarr(self, folder_path: str | Path) -> None: """ import zarr - from .core_tools import _check_zarr_write_is_supported - - _check_zarr_write_is_supported() zarr_group = zarr.open_group(folder_path, mode="w") self.add_templates_to_zarr_group(zarr_group) diff --git a/src/spikeinterface/core/tests/test_analyzer_extension_core.py b/src/spikeinterface/core/tests/test_analyzer_extension_core.py index 8436d7c7a3..6423ce11ce 100644 --- a/src/spikeinterface/core/tests/test_analyzer_extension_core.py +++ b/src/spikeinterface/core/tests/test_analyzer_extension_core.py @@ -11,9 +11,8 @@ from spikeinterface.core.sortinganalyzer import _extension_children, _get_children_dependencies import numpy as np -from spikeinterface.core.core_tools import _is_zarr_write_supported -analyzer_formats = ("memory", "binary_folder", "zarr") if _is_zarr_write_supported() else ("memory", "binary_folder") +analyzer_formats = ("memory", "binary_folder", "zarr") def get_sorting_analyzer(cache_folder, format="memory", sparse=True): @@ -75,9 +74,7 @@ def _check_result_extension(sorting_analyzer, extension_name, cache_folder): # print(k, arr.shape) -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize( "sparse", [ @@ -96,7 +93,7 @@ def test_ComputeRandomSpikes(format, sparse, create_cache_folder): print("Checking results") _check_result_extension(sorting_analyzer, "random_spikes", cache_folder) - print("Delering extension") + print("Deleting extension") sorting_analyzer.delete_extension("random_spikes") print("Re-computing random spikes") @@ -107,9 +104,7 @@ def test_ComputeRandomSpikes(format, sparse, create_cache_folder): _check_result_extension(sorting_analyzer, "random_spikes", cache_folder) -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize("sparse", [True, False]) def test_ComputeWaveforms(format, sparse, create_cache_folder): cache_folder = create_cache_folder @@ -122,7 +117,6 @@ def test_ComputeWaveforms(format, sparse, create_cache_folder): _check_result_extension(sorting_analyzer, "waveforms", cache_folder) -@pytest.mark.requires_zarr_write @pytest.mark.parametrize("sparse", [False, True]) def test_ComputeWaveforms_consistent_across_formats(create_cache_folder, sparse): # Computing waveforms (and templates) on memory / binary_folder / zarr analyzers with the same @@ -178,9 +172,7 @@ def test_ComputeWaveforms_consistent_across_formats(create_cache_folder, sparse) assert np.array_equal(zarr_waveforms, wfs_mem) -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize("sparse", [True, False]) def test_ComputeTemplates(format, sparse, create_cache_folder): cache_folder = create_cache_folder @@ -269,9 +261,7 @@ def test_ComputeTemplates(format, sparse, create_cache_folder): _check_result_extension(sorting_analyzer, "templates", cache_folder) -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize("sparse", [True, False]) def test_ComputeNoiseLevels(format, sparse, create_cache_folder): cache_folder = create_cache_folder diff --git a/src/spikeinterface/core/tests/test_baserecording.py b/src/spikeinterface/core/tests/test_baserecording.py index 8fd38ce4bf..c20850a3e1 100644 --- a/src/spikeinterface/core/tests/test_baserecording.py +++ b/src/spikeinterface/core/tests/test_baserecording.py @@ -24,7 +24,6 @@ from spikeinterface.core.testing import check_recordings_equal from spikeinterface.core import generate_recording -from spikeinterface.core.core_tools import _is_zarr_write_supported def test_BaseRecording(create_cache_folder): @@ -365,34 +364,32 @@ def test_BaseRecording(create_cache_folder): rec_2d = rec_3d.planarize(axes="zy") assert np.allclose(rec_2d.get_channel_locations(), locations_3d[:, [2, 1]]) - # TODO: remove once writing to zarr is supported with zarr>=3 - if _is_zarr_write_supported(): - # test save to zarr - compressor = get_default_zarr_compressor() - rec_zarr = rec2.save(format="zarr", folder=cache_folder / "recording", compressor=compressor) - rec_zarr_loaded = load(cache_folder / "recording.zarr") - # annotations is False because Zarr adds compression ratios - check_recordings_equal(rec2, rec_zarr, return_in_uV=False, check_annotations=False, check_properties=True) - check_recordings_equal( - rec_zarr, rec_zarr_loaded, return_in_uV=False, check_annotations=False, check_properties=True - ) - for annotation_name in rec2.get_annotation_keys(): - assert rec2.get_annotation(annotation_name) == rec_zarr.get_annotation(annotation_name) - assert rec2.get_annotation(annotation_name) == rec_zarr_loaded.get_annotation(annotation_name) - - rec_zarr2 = rec2.save( - format="zarr", folder=cache_folder / "recording_channel_chunk", compressor=compressor, channel_chunk_size=2 - ) - rec_zarr2_loaded = load(cache_folder / "recording_channel_chunk.zarr") - - # annotations is False because Zarr adds compression ratios - check_recordings_equal(rec2, rec_zarr2, return_in_uV=False, check_annotations=False, check_properties=True) - check_recordings_equal( - rec_zarr2, rec_zarr2_loaded, return_in_uV=False, check_annotations=False, check_properties=True - ) - for annotation_name in rec2.get_annotation_keys(): - assert rec2.get_annotation(annotation_name) == rec_zarr2.get_annotation(annotation_name) - assert rec2.get_annotation(annotation_name) == rec_zarr2_loaded.get_annotation(annotation_name) + # test save to zarr + compressor = get_default_zarr_compressor() + rec_zarr = rec2.save(format="zarr", folder=cache_folder / "recording", compressors=compressor) + rec_zarr_loaded = load(cache_folder / "recording.zarr") + # annotations is False because Zarr adds compression ratios + check_recordings_equal(rec2, rec_zarr, return_in_uV=False, check_annotations=False, check_properties=True) + check_recordings_equal( + rec_zarr, rec_zarr_loaded, return_in_uV=False, check_annotations=False, check_properties=True + ) + for annotation_name in rec2.get_annotation_keys(): + assert rec2.get_annotation(annotation_name) == rec_zarr.get_annotation(annotation_name) + assert rec2.get_annotation(annotation_name) == rec_zarr_loaded.get_annotation(annotation_name) + + rec_zarr2 = rec2.save( + format="zarr", folder=cache_folder / "recording_channel_chunk", compressors=compressor, channel_chunk_size=2 + ) + rec_zarr2_loaded = load(cache_folder / "recording_channel_chunk.zarr") + + # annotations is False because Zarr adds compression ratios + check_recordings_equal(rec2, rec_zarr2, return_in_uV=False, check_annotations=False, check_properties=True) + check_recordings_equal( + rec_zarr2, rec_zarr2_loaded, return_in_uV=False, check_annotations=False, check_properties=True + ) + for annotation_name in rec2.get_annotation_keys(): + assert rec2.get_annotation(annotation_name) == rec_zarr2.get_annotation(annotation_name) + assert rec2.get_annotation(annotation_name) == rec_zarr2_loaded.get_annotation(annotation_name) def test_json_pickle_equivalence(create_cache_folder): diff --git a/src/spikeinterface/core/tests/test_basesorting.py b/src/spikeinterface/core/tests/test_basesorting.py index ec7c2cefb6..09c82a9ecb 100644 --- a/src/spikeinterface/core/tests/test_basesorting.py +++ b/src/spikeinterface/core/tests/test_basesorting.py @@ -25,13 +25,37 @@ from spikeinterface.core.base import BaseExtractor, minimum_spike_dtype, unit_period_dtype from spikeinterface.core.basesorting import LEXSORT_UNIT_COMPACT from spikeinterface.core.testing import check_sorted_arrays_equal, check_sortings_equal -from spikeinterface.core.core_tools import _is_zarr_write_supported + + +def _make_sorting_with_shuffled_ties(num_units, num_segments, seed=42): + """Build a NumpySorting whose cotemporal spikes are in arbitrary unit_index order. + + A spike vector is only guaranteed to be segment-blocked and sample_index-ascending within each + segment; the unit_index order among spikes sharing a sample_index is unspecified (see #4606). + Building via `NumpySorting.from_unit_dict` happens to produce unit-ascending ties, so it + can't test the shuffled tie case. + """ + rng = np.random.default_rng(seed) + num_spikes = 2_000 + + # A sample range far smaller than num_spikes, so cotemporal spikes are abundant -- including + # repeats of the same (segment, sample, unit), the tie that np.lexsort itself cannot break. + spikes = np.empty(num_spikes, dtype=minimum_spike_dtype) + spikes["sample_index"] = rng.integers(0, 200, size=num_spikes) + spikes["unit_index"] = rng.integers(0, num_units, size=num_spikes) + spikes["segment_index"] = rng.integers(0, num_segments, size=num_spikes) + + # Order by segment then sample, breaking ties randomly rather than by unit_index. + spikes = spikes[np.lexsort((rng.random(num_spikes), spikes["sample_index"], spikes["segment_index"]))] + + sorting = NumpySorting(spikes, 30_000.0, np.arange(num_units)) + assert sorting.get_num_segments() == num_segments + return sorting def test_BaseSorting(create_cache_folder): - cache_folder = create_cache_folder num_seg = 2 - file_path = cache_folder / "test_BaseSorting.npz" + file_path = create_cache_folder / "test_BaseSorting.npz" file_path.parent.mkdir(exist_ok=True) create_sorting_npz(num_seg, file_path) @@ -58,28 +82,28 @@ def test_BaseSorting(create_cache_folder): check_sortings_equal(sorting, sorting3, check_annotations=True, check_properties=True) # dump/load json - sorting.dump_to_json(cache_folder / "test_BaseSorting.json") - sorting2 = load(cache_folder / "test_BaseSorting.json") - sorting3 = load(cache_folder / "test_BaseSorting.json") + sorting.dump_to_json(create_cache_folder / "test_BaseSorting.json") + sorting2 = load(create_cache_folder / "test_BaseSorting.json") + sorting3 = load(create_cache_folder / "test_BaseSorting.json") check_sortings_equal(sorting, sorting2, check_annotations=True, check_properties=False) check_sortings_equal(sorting, sorting3, check_annotations=True, check_properties=False) # dump/load pickle - sorting.dump_to_pickle(cache_folder / "test_BaseSorting.pkl") - sorting2 = load(cache_folder / "test_BaseSorting.pkl") - sorting3 = load(cache_folder / "test_BaseSorting.pkl") + sorting.dump_to_pickle(create_cache_folder / "test_BaseSorting.pkl") + sorting2 = load(create_cache_folder / "test_BaseSorting.pkl") + sorting3 = load(create_cache_folder / "test_BaseSorting.pkl") check_sortings_equal(sorting, sorting2, check_annotations=True, check_properties=True) check_sortings_equal(sorting, sorting3, check_annotations=True, check_properties=True) # cache old format : npz_folder - folder = cache_folder / "simple_sorting_npz_folder" + folder = create_cache_folder / "simple_sorting_npz_folder" sorting.set_property("test", np.ones(len(sorting.unit_ids))) sorting.save(folder=folder, format="npz_folder") sorting2 = load(folder) assert isinstance(sorting2, NpzFolderSorting) # cache new format : binary - folder = cache_folder / "simple_sorting_binary" + folder = create_cache_folder / "simple_sorting_binary" sorting.set_property("test", np.ones(len(sorting.unit_ids))) sorting.save(folder=folder, format="numpy_folder") sorting2 = load(folder) @@ -144,44 +168,46 @@ def test_BaseSorting(create_cache_folder): del sorting6 del sorting5 - # TODO: remove once writing to zarr is supported with zarr>=3 - if _is_zarr_write_supported(): - # test save to zarr - # compressor = get_default_zarr_compressor() - sorting_zarr = sorting.save(format="zarr", folder=cache_folder / "sorting.zarr") - sorting_zarr_loaded = load(cache_folder / "sorting.zarr") - # annotations is False because Zarr adds compression ratios - check_sortings_equal(sorting, sorting_zarr, check_annotations=False, check_properties=True) - check_sortings_equal(sorting_zarr, sorting_zarr_loaded, check_annotations=False, check_properties=True) - for annotation_name in sorting.get_annotation_keys(): - assert sorting.get_annotation(annotation_name) == sorting_zarr.get_annotation(annotation_name) - assert sorting.get_annotation(annotation_name) == sorting_zarr_loaded.get_annotation(annotation_name) - - -def _make_sorting_with_shuffled_ties(num_units, num_segments, seed=42): - """Build a NumpySorting whose cotemporal spikes are in arbitrary unit_index order. - - A spike vector is only guaranteed to be segment-blocked and sample_index-ascending within each - segment; the unit_index order among spikes sharing a sample_index is unspecified (see #4606). - Building via `NumpySorting.from_unit_dict` happens to produce unit-ascending ties, so it - can't test the shuffled tie case. - """ - rng = np.random.default_rng(seed) - num_spikes = 2_000 - - # A sample range far smaller than num_spikes, so cotemporal spikes are abundant -- including - # repeats of the same (segment, sample, unit), the tie that np.lexsort itself cannot break. - spikes = np.empty(num_spikes, dtype=minimum_spike_dtype) - spikes["sample_index"] = rng.integers(0, 200, size=num_spikes) - spikes["unit_index"] = rng.integers(0, num_units, size=num_spikes) - spikes["segment_index"] = rng.integers(0, num_segments, size=num_spikes) - - # Order by segment then sample, breaking ties randomly rather than by unit_index. - spikes = spikes[np.lexsort((rng.random(num_spikes), spikes["sample_index"], spikes["segment_index"]))] - - sorting = NumpySorting(spikes, 30_000.0, np.arange(num_units)) - assert sorting.get_num_segments() == num_segments - return sorting + # test save to zarr + # compressor = get_default_zarr_compressor() + sorting_zarr = sorting.save(format="zarr", folder=create_cache_folder / "sorting.zarr") + sorting_zarr_loaded = load(create_cache_folder / "sorting.zarr") + # annotations is False because Zarr adds compression ratios + check_sortings_equal(sorting, sorting_zarr, check_annotations=False, check_properties=True) + check_sortings_equal(sorting_zarr, sorting_zarr_loaded, check_annotations=False, check_properties=True) + for annotation_name in sorting.get_annotation_keys(): + assert sorting.get_annotation(annotation_name) == sorting_zarr.get_annotation(annotation_name) + assert sorting.get_annotation(annotation_name) == sorting_zarr_loaded.get_annotation(annotation_name) + + +def test_zarr_save_with_sharding(create_cache_folder): + sorting = generate_sorting(durations=[600, 600]) + # 1 MB target chunk size for Zarr + zarr_target_bytes = 1 * 1024 * 1024 + shard_factor = 3 + sorting_zarr = sorting.save( + format="zarr", + folder=create_cache_folder / "sorting_sharded.zarr", + target_chunk_size_bytes=zarr_target_bytes, + shard_factor=shard_factor, + ) + sorting_zarr_loaded = load(create_cache_folder / "sorting_sharded.zarr") + check_sortings_equal(sorting, sorting_zarr, check_annotations=False, check_properties=True) + check_sortings_equal(sorting_zarr, sorting_zarr_loaded, check_annotations=False, check_properties=True) + + # check that chunks and shards are correctly set + spikes_group = sorting_zarr._root["spikes"] + for field in spikes_group: + if field not in ("segment_slices", "sample_index_chunk_firsts"): + array = spikes_group[field] + assert array.chunks is not None + assert array.shards is not None + assert array.shards[0] == array.chunks[0] * shard_factor + assert array.chunks[0] * array.dtype.itemsize <= zarr_target_bytes + print(f"Field: {field}, Chunks: {array.chunks}, Shards: {array.shards}") + print( + f"Field: {field}, Chunk size in bytes: {array.chunks[0] * array.dtype.itemsize} - Target: {zarr_target_bytes}" + ) @pytest.mark.parametrize("use_numba", [True, False], ids=["numba", "numpy"]) @@ -418,9 +444,9 @@ def test_select_periods(): from pathlib import Path with tempfile.TemporaryDirectory() as tmpdirname: - cache_folder = Path(tmpdirname) + create_cache_folder = Path(tmpdirname) - test_BaseSorting(cache_folder) + test_BaseSorting(create_cache_folder) test_npy_sorting() test_empty_sorting() test_select_periods() diff --git a/src/spikeinterface/core/tests/test_loading.py b/src/spikeinterface/core/tests/test_loading.py index 364c5bd910..8d57a3412f 100644 --- a/src/spikeinterface/core/tests/test_loading.py +++ b/src/spikeinterface/core/tests/test_loading.py @@ -90,7 +90,7 @@ def generate_motion_object(): return motion -@pytest.mark.parametrize("output_format", ["binary", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)]) +@pytest.mark.parametrize("output_format", ["binary", "zarr"]) def test_load_binary_recording(generate_recording_sorting, tmp_path, output_format): rec, _ = generate_recording_sorting _ = rec.save(folder=tmp_path / "test_recording", format=output_format, overwrite=True) @@ -103,7 +103,7 @@ def test_load_binary_recording(generate_recording_sorting, tmp_path, output_form check_recordings_equal(rec, rec_loaded) -@pytest.mark.parametrize("output_format", ["numpy_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)]) +@pytest.mark.parametrize("output_format", ["numpy_folder", "zarr"]) def test_load_binary_sorting(generate_recording_sorting, tmp_path, output_format): _, sort = generate_recording_sorting _ = sort.save(folder=tmp_path / "test_sorting", format=output_format, overwrite=True) @@ -141,9 +141,7 @@ def test_load_ext_extractors(generate_recording_sorting, tmp_path, extension): check_sortings_equal(sort, sort_loaded, check_properties=False) -@pytest.mark.parametrize( - "output_format", ["binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("output_format", ["binary_folder", "zarr"]) def test_load_sorting_analyzer(generate_sorting_analyzer, tmp_path, output_format): analyzer = generate_sorting_analyzer _ = analyzer.save_as(folder=tmp_path / "analyzer", format=output_format) @@ -160,7 +158,6 @@ def test_load_sorting_analyzer(generate_sorting_analyzer, tmp_path, output_forma assert ext in analyzer_loaded.extensions -@pytest.mark.requires_zarr_write def test_load_templates(tmp_path, generate_templates_object): templates = generate_templates_object templates_dict = templates.to_dict() diff --git a/src/spikeinterface/core/tests/test_node_pipeline.py b/src/spikeinterface/core/tests/test_node_pipeline.py index 533966d703..a721fc3367 100644 --- a/src/spikeinterface/core/tests/test_node_pipeline.py +++ b/src/spikeinterface/core/tests/test_node_pipeline.py @@ -16,7 +16,6 @@ ExtractDenseWaveforms, sorting_to_peaks, ) -from spikeinterface.core.core_tools import _is_zarr_write_supported class AmplitudeExtractionNode(PipelineNode): @@ -181,36 +180,34 @@ def test_run_node_pipeline(cache_folder_creation): assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms2) assert np.array_equal(denoised_waveforms_rms2, denoised_waveforms_rms3) - # TODO: remove once writing to zarr is supported with zarr>=3 - if _is_zarr_write_supported(): - # gather zarr mode - import zarr - - zarr_folder = cache_folder / f"pipeline_folder_{loop}.zarr" - if zarr_folder.is_dir(): - shutil.rmtree(zarr_folder) - output = run_node_pipeline( - recording, - nodes, - job_kwargs, - gather_mode="zarr", - dest=zarr_folder, - names=["amplitudes", "waveforms_rms", "denoised_waveforms_rms"], - ) - amplitudes_z, waveforms_rms_z, denoised_waveforms_rms_z = output - - # values must match the memory gather - assert np.array_equal(amplitudes, amplitudes_z[:]) - assert np.array_equal(waveforms_rms, waveforms_rms_z[:]) - assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_z[:]) - - # arrays must be persisted on disk and re-openable - zarr_root = zarr.open(str(zarr_folder), mode="r") - for name in ("amplitudes", "waveforms_rms", "denoised_waveforms_rms"): - assert name in zarr_root - assert np.array_equal(amplitudes, zarr_root["amplitudes"][:]) - assert np.array_equal(waveforms_rms, zarr_root["waveforms_rms"][:]) - assert np.array_equal(denoised_waveforms_rms, zarr_root["denoised_waveforms_rms"][:]) + # gather zarr mode + import zarr + + zarr_folder = cache_folder / f"pipeline_folder_{loop}.zarr" + if zarr_folder.is_dir(): + shutil.rmtree(zarr_folder) + output = run_node_pipeline( + recording, + nodes, + job_kwargs, + gather_mode="zarr", + dest=zarr_folder, + names=["amplitudes", "waveforms_rms", "denoised_waveforms_rms"], + ) + amplitudes_z, waveforms_rms_z, denoised_waveforms_rms_z = output + + # values must match the memory gather + assert np.array_equal(amplitudes, amplitudes_z[:]) + assert np.array_equal(waveforms_rms, waveforms_rms_z[:]) + assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_z[:]) + + # arrays must be persisted on disk and re-openable + zarr_root = zarr.open(str(zarr_folder), mode="r") + for name in ("amplitudes", "waveforms_rms", "denoised_waveforms_rms"): + assert name in zarr_root + assert np.array_equal(amplitudes, zarr_root["amplitudes"][:]) + assert np.array_equal(waveforms_rms, zarr_root["waveforms_rms"][:]) + assert np.array_equal(denoised_waveforms_rms, zarr_root["denoised_waveforms_rms"][:]) # gather npy mode with an explicit list of file paths (final location) npy_files_folder = cache_folder / f"pipeline_npy_files_{loop}" @@ -235,38 +232,36 @@ def test_run_node_pipeline(cache_folder_creation): assert np.array_equal(waveforms_rms, waveforms_rms_f) assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_f) - # TODO: remove once writing to zarr is supported with zarr>=3 - if _is_zarr_write_supported(): - # gather zarr mode with an explicit list of dataset paths, created on the fly - # inside an existing store (final location, e.g. an analyzer extension group) - datasets_store = cache_folder / f"pipeline_zarr_datasets_{loop}.zarr" - if datasets_store.is_dir(): - shutil.rmtree(datasets_store) - # pre-existing store that must not be wiped - root = zarr.open(str(datasets_store), mode="w") - root.attrs["preexisting"] = True - dataset_paths = [ - datasets_store / "extensions" / "amplitudes", - datasets_store / "extensions" / "waveforms_rms", - datasets_store / "extensions" / "denoised_waveforms_rms", - ] - output = run_node_pipeline( - recording, - nodes, - job_kwargs, - gather_mode="zarr", - dest=dataset_paths, - ) - amplitudes_d, waveforms_rms_d, denoised_waveforms_rms_d = output - assert np.array_equal(amplitudes, amplitudes_d[:]) - assert np.array_equal(waveforms_rms, waveforms_rms_d[:]) - assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_d[:]) - # data must be persisted at the passed final location and the store not wiped - root_reopen = zarr.open(str(datasets_store), mode="r") - assert root_reopen.attrs.get("preexisting", False) - assert np.array_equal(amplitudes, root_reopen["extensions"]["amplitudes"][:]) - assert np.array_equal(waveforms_rms, root_reopen["extensions"]["waveforms_rms"][:]) - assert np.array_equal(denoised_waveforms_rms, root_reopen["extensions"]["denoised_waveforms_rms"][:]) + # gather zarr mode with an explicit list of dataset paths, created on the fly + # inside an existing store (final location, e.g. an analyzer extension group) + datasets_store = cache_folder / f"pipeline_zarr_datasets_{loop}.zarr" + if datasets_store.is_dir(): + shutil.rmtree(datasets_store) + # pre-existing store that must not be wiped + root = zarr.open(str(datasets_store), mode="w") + root.attrs["preexisting"] = True + dataset_paths = [ + datasets_store / "extensions" / "amplitudes", + datasets_store / "extensions" / "waveforms_rms", + datasets_store / "extensions" / "denoised_waveforms_rms", + ] + output = run_node_pipeline( + recording, + nodes, + job_kwargs, + gather_mode="zarr", + dest=dataset_paths, + ) + amplitudes_d, waveforms_rms_d, denoised_waveforms_rms_d = output + assert np.array_equal(amplitudes, amplitudes_d[:]) + assert np.array_equal(waveforms_rms, waveforms_rms_d[:]) + assert np.array_equal(denoised_waveforms_rms, denoised_waveforms_rms_d[:]) + # data must be persisted at the passed final location and the store not wiped + root_reopen = zarr.open(str(datasets_store), mode="r") + assert root_reopen.attrs.get("preexisting", False) + assert np.array_equal(amplitudes, root_reopen["extensions"]["amplitudes"][:]) + assert np.array_equal(waveforms_rms, root_reopen["extensions"]["waveforms_rms"][:]) + assert np.array_equal(denoised_waveforms_rms, root_reopen["extensions"]["denoised_waveforms_rms"][:]) # Test pickle mechanism for node in nodes: @@ -276,7 +271,6 @@ def test_run_node_pipeline(cache_folder_creation): unpickled_node = pickle.loads(pickled_node) -@pytest.mark.requires_zarr_write def test_gather_to_zarr_chunking(tmp_path): # the zarr chunk size along the first axis must be picked from a byte target (not from the # size of the first gathered buffer), so it stays sensible for billions of spikes and never @@ -325,7 +319,6 @@ def test_gather_to_zarr_chunking(tmp_path): assert np.array_equal(waveforms[:], waveforms2[:]) -@pytest.mark.requires_zarr_write def test_gather_to_zarr_chunk_bytes_per_name(tmp_path): # `zarr_target_chunk_bytes` can also be a dict to use a different byte target per array recording, sorting = generate_ground_truth_recording(num_channels=8, num_units=5, durations=[20.0], seed=7) diff --git a/src/spikeinterface/core/tests/test_sortinganalyzer.py b/src/spikeinterface/core/tests/test_sortinganalyzer.py index 527e053184..9e080e720b 100644 --- a/src/spikeinterface/core/tests/test_sortinganalyzer.py +++ b/src/spikeinterface/core/tests/test_sortinganalyzer.py @@ -20,15 +20,15 @@ AnalyzerExtension, _sort_extensions_by_dependency, ) +from spikeinterface.core.zarr_tools import check_compressors_match from spikeinterface.core.analyzer_extension_core import BaseSpikeVectorExtension from spikeinterface.core.base import minimum_spike_dtype # to test basespikevectorextension with node pipeline from spikeinterface.core.node_pipeline import SpikeRetriever from spikeinterface.core.tests.test_node_pipeline import AmplitudeExtractionNode -from spikeinterface.core.core_tools import _is_zarr_write_supported -analyzer_formats = ("memory", "binary_folder", "zarr") if _is_zarr_write_supported() else ("memory", "binary_folder") +analyzer_formats = ("memory", "binary_folder", "zarr") def get_dataset(): @@ -142,11 +142,13 @@ def test_SortingAnalyzer_binary_folder(tmp_path, dataset): assert "number" in sorting_analyzer_reloded.sorting.get_property_keys() -@pytest.mark.requires_zarr_write def test_SortingAnalyzer_zarr(tmp_path, dataset): recording, sorting = dataset recording = recording.save(folder=tmp_path / "recording_zarr") + # make recording JSON serializable + recording = recording.save(folder=tmp_path / "recording_for_zarr", overwrite=True) + folder = tmp_path / "test_SortingAnalyzer_zarr.zarr" default_compressor = get_default_zarr_compressor() @@ -158,13 +160,12 @@ def test_SortingAnalyzer_zarr(tmp_path, dataset): _check_sorting_analyzers(sorting_analyzer, sorting, cache_folder=tmp_path) # check that compression is applied - assert ( - sorting_analyzer._get_zarr_root()["extensions"]["random_spikes"]["random_spikes_indices"].compressor.codec_id - == default_compressor.codec_id + check_compressors_match( + default_compressor, + sorting_analyzer._get_zarr_root()["extensions"]["random_spikes"]["random_spikes_indices"].compressors[0], ) - assert ( - sorting_analyzer._get_zarr_root()["extensions"]["templates"]["average"].compressor.codec_id - == default_compressor.codec_id + check_compressors_match( + default_compressor, sorting_analyzer._get_zarr_root()["extensions"]["templates"]["average"].compressors[0] ) # test select_units see https://github.com/SpikeInterface/spikeinterface/issues/3041 @@ -185,35 +186,33 @@ def test_SortingAnalyzer_zarr(tmp_path, dataset): sparsity=None, return_in_uV=False, overwrite=True, - backend_options={"saving_options": {"compressor": None}}, + backend_options={"saving_options": {"compressors": None}}, ) - print(sorting_analyzer_no_compression._backend_options) sorting_analyzer_no_compression.compute(["random_spikes", "templates"]) assert ( - sorting_analyzer_no_compression._get_zarr_root()["extensions"]["random_spikes"][ - "random_spikes_indices" - ].compressor - is None + len( + sorting_analyzer_no_compression._get_zarr_root()["extensions"]["random_spikes"][ + "random_spikes_indices" + ].compressors + ) + == 0 ) - assert sorting_analyzer_no_compression._get_zarr_root()["extensions"]["templates"]["average"].compressor is None + assert len(sorting_analyzer_no_compression._get_zarr_root()["extensions"]["templates"]["average"].compressors) == 0 # test a different compressor - from numcodecs import LZMA + from zarr.codecs.numcodecs import LZMA lzma_compressor = LZMA() folder = tmp_path / "test_SortingAnalyzer_zarr_lzma.zarr" sorting_analyzer_lzma = sorting_analyzer_no_compression.save_as( - format="zarr", folder=folder, backend_options={"saving_options": {"compressor": lzma_compressor}} + format="zarr", folder=folder, backend_options={"saving_options": {"compressors": lzma_compressor}} ) - assert ( - sorting_analyzer_lzma._get_zarr_root()["extensions"]["random_spikes"][ - "random_spikes_indices" - ].compressor.codec_id - == LZMA.codec_id + check_compressors_match( + lzma_compressor, + sorting_analyzer_lzma._get_zarr_root()["extensions"]["random_spikes"]["random_spikes_indices"].compressors[0], ) - assert ( - sorting_analyzer_lzma._get_zarr_root()["extensions"]["templates"]["average"].compressor.codec_id - == LZMA.codec_id + check_compressors_match( + lzma_compressor, sorting_analyzer_lzma._get_zarr_root()["extensions"]["templates"]["average"].compressors[0] ) # test set_sorting_property @@ -314,22 +313,20 @@ def test_load_without_runtime_info(tmp_path, dataset): with pytest.warns(UserWarning): sorting_analyzer = load_sorting_analyzer(folder, format="auto") - # TODO: remove once writing to zarr is supported with zarr>=3 - if _is_zarr_write_supported(): - # zarr - folder = tmp_path / "test_SortingAnalyzer_run_info.zarr" - sorting_analyzer = create_sorting_analyzer( - sorting, recording, format="zarr", folder=folder, sparse=False, sparsity=None - ) - sorting_analyzer.compute(extensions) - # remove run_info from attrs to mimic a previous version of spikeinterface - root = sorting_analyzer._get_zarr_root(mode="r+") - for ext in extensions: - del root["extensions"][ext].attrs["run_info"] - zarr.consolidate_metadata(root.store) - # should raise a warning for missing run_info - with pytest.warns(UserWarning): - sorting_analyzer = load_sorting_analyzer(folder, format="auto") + # zarr + folder = tmp_path / "test_SortingAnalyzer_run_info.zarr" + sorting_analyzer = create_sorting_analyzer( + sorting, recording, format="zarr", folder=folder, sparse=False, sparsity=None + ) + sorting_analyzer.compute(extensions) + # remove run_info from attrs to mimic a previous version of spikeinterface + root = sorting_analyzer._get_zarr_root(mode="r+") + for ext in extensions: + del root["extensions"][ext].attrs["run_info"] + zarr.consolidate_metadata(root.store) + # should raise a warning for missing run_info + with pytest.warns(UserWarning): + sorting_analyzer = load_sorting_analyzer(folder, format="auto") def test_SortingAnalyzer_tmp_recording(dataset): @@ -373,7 +370,7 @@ def test_SortingAnalyzer_interleaved_probegroup(dataset): assert np.array_equal(recording.get_channel_locations(), sorting_analyzer.get_channel_locations()) -@pytest.mark.parametrize("format", ["binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)]) +@pytest.mark.parametrize("format", ["binary_folder", "zarr"]) def test_load_in_lazy_mode(tmp_path, dataset, format): recording, sorting = dataset @@ -891,9 +888,7 @@ def _compute_reference_pipeline_data(dataset): return analyzer.get_extension("dummy_pipeline").get_data() -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize("lazy", [True, False]) def test_compute_pipeline_extension_gather_to_disk_lazy(tmp_path, dataset, format, lazy): """ @@ -958,7 +953,7 @@ def test_compute_pipeline_extension_gather_to_disk_lazy(tmp_path, dataset, forma assert np.array_equal(load_sorting_analyzer(folder).get_extension("dummy_pipeline").get_data(), amp_ref) -@pytest.mark.parametrize("format", ["binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)]) +@pytest.mark.parametrize("format", ["binary_folder", "zarr"]) def test_compute_pipeline_extension_save_false(tmp_path, dataset, format): """ With save=False on a disk-backed analyzer, node-pipeline extensions are computed in memory @@ -980,9 +975,7 @@ def test_compute_pipeline_extension_save_false(tmp_path, dataset, format): assert not analyzer_reloaded.has_extension("dummy_pipeline") -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize("lazy", [True, False]) def test_compute_one_pipeline_extension_gather_to_disk(tmp_path, dataset, format, lazy): """ diff --git a/src/spikeinterface/core/tests/test_template_class.py b/src/spikeinterface/core/tests/test_template_class.py index 20f88817cb..953527c8d5 100644 --- a/src/spikeinterface/core/tests/test_template_class.py +++ b/src/spikeinterface/core/tests/test_template_class.py @@ -115,7 +115,6 @@ def test_initialization_fail_with_dense_templates(): template = generate_test_template(template_type="sparse_with_dense_templates") -@pytest.mark.requires_zarr_write @pytest.mark.parametrize("is_in_uV", [True, False]) @pytest.mark.parametrize("template_type", ["dense", "sparse"]) def test_save_and_load_zarr(template_type, is_in_uV, tmp_path): diff --git a/src/spikeinterface/core/tests/test_time_handling.py b/src/spikeinterface/core/tests/test_time_handling.py index 4a08a1c7c8..2f3cc1a6bf 100644 --- a/src/spikeinterface/core/tests/test_time_handling.py +++ b/src/spikeinterface/core/tests/test_time_handling.py @@ -114,7 +114,7 @@ def test_has_time_vector(self, time_vector_recording): assert raw_recording.has_time_vector(segment_idx) is False assert times_recording.has_time_vector(segment_idx) is True - @pytest.mark.parametrize("mode", ["binary", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)]) + @pytest.mark.parametrize("mode", ["binary", "zarr"]) @pytest.mark.parametrize("fixture_name", ["time_vector_recording", "t_start_recording"]) def test_times_propagated_to_save_folder(self, request, fixture_name, mode, tmp_path): """ @@ -375,7 +375,7 @@ def test_save_and_load_time_shift(self, request, fixture_name, tmp_path): times_recording.get_times(segment_index=idx), loaded_recording.get_times(segment_index=idx) ) - @pytest.mark.parametrize("save_format", ["binary", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)]) + @pytest.mark.parametrize("save_format", ["binary", "zarr"]) def test_shift_times_after_load(self, request, save_format, tmp_path): """ Shift times on a recording loaded from disk as a read-only np.memmap diff --git a/src/spikeinterface/core/tests/test_unitsselectionsorting.py b/src/spikeinterface/core/tests/test_unitsselectionsorting.py index 309563df62..c240d25d88 100644 --- a/src/spikeinterface/core/tests/test_unitsselectionsorting.py +++ b/src/spikeinterface/core/tests/test_unitsselectionsorting.py @@ -248,7 +248,6 @@ def test_non_identity_selection_does_not_share(unit_ids): assert len(parent._cached_lexsorted_spike_vector) == 0 -@pytest.mark.requires_zarr_write def test_identity_selection_keeps_lazy_zarr_vector(tmp_path): """A lazy parent spike vector should stay lazy through an identity selection.""" from spikeinterface.core.zarrextractors import ZarrSortingExtractor, ZarrSpikeVector diff --git a/src/spikeinterface/core/tests/test_waveform_tools.py b/src/spikeinterface/core/tests/test_waveform_tools.py index c2750774b6..6623c92046 100644 --- a/src/spikeinterface/core/tests/test_waveform_tools.py +++ b/src/spikeinterface/core/tests/test_waveform_tools.py @@ -157,7 +157,6 @@ def test_waveform_tools(create_cache_folder): _check_all_wf_equal(list_wfs_sparse) -@pytest.mark.requires_zarr_write @pytest.mark.parametrize("sparse", [False, True]) def test_extract_waveforms_to_single_buffer_zarr(tmp_path, sparse): # the "zarr" mode writes waveforms directly to a zarr dataset. Workers return their block and @@ -211,7 +210,6 @@ def test_extract_waveforms_to_single_buffer_zarr(tmp_path, sparse): assert np.array_equal(reference, reloaded[:]) -@pytest.mark.requires_zarr_write def test_waveforms_at_segment_borders(tmp_path): # spikes near the segment borders are partially filled: samples outside of the segment are 0 from spikeinterface.core import NumpySorting diff --git a/src/spikeinterface/core/tests/test_zarrextractors.py b/src/spikeinterface/core/tests/test_zarrextractors.py index 2b388839d1..4f17cb28e4 100644 --- a/src/spikeinterface/core/tests/test_zarrextractors.py +++ b/src/spikeinterface/core/tests/test_zarrextractors.py @@ -9,6 +9,8 @@ generate_sorting, load, ) +from spikeinterface.core.testing import check_recordings_equal +from spikeinterface.core.zarr_tools import check_compressors_match from spikeinterface.core.zarrextractors import ( ZarrRecordingExtractor, ZarrSampleIndexSearch, @@ -18,51 +20,50 @@ ) -@pytest.mark.requires_zarr_write def test_zarr_compression_options(tmp_path): - from numcodecs import Blosc, Delta, FixedScaleOffset + from zarr.codecs.numcodecs import Delta, FixedScaleOffset + from zarr.codecs import BloscCodec, BloscShuffle recording = generate_recording(durations=[2]) recording.set_times(recording.get_times() + 100) # store in root standard normal way # default compressor - defaut_compressor = get_default_zarr_compressor() + default_compressor = get_default_zarr_compressor() # other compressor - other_compressor1 = Blosc(cname="zlib", clevel=3, shuffle=Blosc.NOSHUFFLE) - other_compressor2 = Blosc(cname="blosclz", clevel=8, shuffle=Blosc.AUTOSHUFFLE) + other_compressor1 = BloscCodec(cname="zlib", clevel=3, shuffle=BloscShuffle.noshuffle) + other_compressor2 = BloscCodec(cname="blosclz", clevel=8, shuffle=BloscShuffle.shuffle) # timestamps compressors / filters default_filters = None - other_filters1 = [FixedScaleOffset(scale=5, offset=2, dtype=recording.get_dtype())] + other_filters1 = [FixedScaleOffset(scale=5, offset=2, dtype=recording.get_dtype().str)] other_filters2 = [Delta(dtype="float64")] # default ZarrRecordingExtractor.write_recording(recording, tmp_path / "rec_default.zarr") rec_default = ZarrRecordingExtractor(tmp_path / "rec_default.zarr") - assert rec_default._root["traces_seg0"].compressor == defaut_compressor - assert rec_default._root["traces_seg0"].filters == default_filters - assert rec_default._root["times_seg0"].compressor == defaut_compressor - assert rec_default._root["times_seg0"].filters == default_filters + check_compressors_match(rec_default._root["traces_seg0"].compressors[0], default_compressor) + check_compressors_match(rec_default._root["times_seg0"].compressors[0], default_compressor) + check_compressors_match(rec_default._root["traces_seg0"].filters, default_filters) + check_compressors_match(rec_default._root["times_seg0"].filters, default_filters) # now with other compressor ZarrRecordingExtractor.write_recording( recording, tmp_path / "rec_other.zarr", - compressor=defaut_compressor, + compressors=default_compressor, filters=default_filters, compressor_by_dataset={"traces": other_compressor1, "times": other_compressor2}, filters_by_dataset={"traces": other_filters1, "times": other_filters2}, ) rec_other = ZarrRecordingExtractor(tmp_path / "rec_other.zarr") - assert rec_other._root["traces_seg0"].compressor == other_compressor1 - assert rec_other._root["traces_seg0"].filters == other_filters1 - assert rec_other._root["times_seg0"].compressor == other_compressor2 - assert rec_other._root["times_seg0"].filters == other_filters2 + check_compressors_match(rec_other._root["traces_seg0"].compressors[0], other_compressor1) + check_compressors_match(rec_other._root["traces_seg0"].filters, other_filters1) + check_compressors_match(rec_other._root["times_seg0"].compressors[0], other_compressor2) + check_compressors_match(rec_other._root["times_seg0"].filters, other_filters2) -@pytest.mark.requires_zarr_write def test_ZarrSortingExtractor(tmp_path): np_sorting = generate_sorting() @@ -82,6 +83,49 @@ def test_ZarrSortingExtractor(tmp_path): sorting = load(sorting.to_dict()) +def test_sharding_options(tmp_path): + recording = generate_recording(durations=[10], num_channels=20) + folder = tmp_path / "zarr_sharding.zarr" + + # explicitly specify chunks and shards + ZarrRecordingExtractor.write_recording(recording, folder, chunks=(1000, 5), shards=(5000, 10), n_jobs=2) + recording_zarr = ZarrRecordingExtractor(folder) + assert recording_zarr._root["traces_seg0"].chunks == (1000, 5) + assert recording_zarr._root["traces_seg0"].shards == (5000, 10) + check_recordings_equal(recording, recording_zarr) + + # specify shard_factor and chunk_size + folder = tmp_path / "zarr_sharding_factor.zarr" + ZarrRecordingExtractor.write_recording( + recording, folder, chunk_size=1000, channel_chunk_size=2, shard_factor=(5, 2), n_jobs=2 + ) + recording_zarr = ZarrRecordingExtractor(folder) + assert recording_zarr._root["traces_seg0"].chunks == (1000, 2) + assert recording_zarr._root["traces_seg0"].shards == (5000, 4) + check_recordings_equal(recording, recording_zarr) + + # raise error if both shards and shard_factor are provided + with pytest.raises(ValueError): + folder = tmp_path / "shards_and_shard_factor.zarr" + ZarrRecordingExtractor.write_recording( + recording, folder, chunk_size=1000, channel_chunk_size=2, shard_factor=5, shards=(5000, 10), n_jobs=2 + ) + + # raise error if shards is smaller than chunks + with pytest.raises(AssertionError): + folder = tmp_path / "shards_smaller_than_chunks.zarr" + ZarrRecordingExtractor.write_recording( + recording, folder, chunk_size=1000, channel_chunk_size=2, shards=(500, 10), n_jobs=2 + ) + + # raise error if shards is not a multiple of chunks + with pytest.raises(AssertionError): + folder = tmp_path / "shards_not_multiple_of_chunks.zarr" + ZarrRecordingExtractor.write_recording( + recording, folder, chunk_size=1000, channel_chunk_size=2, shards=(5500, 10), n_jobs=2 + ) + + def test_ZarrSampleIndexSearch(tmp_path): rng = np.random.default_rng(0) # two "segments", each sorted, with long runs of equal values so that runs cross @@ -102,7 +146,6 @@ def test_ZarrSampleIndexSearch(tmp_path): np.testing.assert_array_equal(search.searchsorted([5], 10, 10), [0]) -@pytest.mark.requires_zarr_write def test_ZarrSortingExtractor_lazy_search(tmp_path): sorting = generate_sorting(num_units=10, durations=[5.0, 3.0, 4.0], firing_rates=40.0, seed=0) folder = tmp_path / "sorting.zarr" @@ -112,8 +155,8 @@ def test_ZarrSortingExtractor_lazy_search(tmp_path): spikes_group = zarr.open(folder, mode="a")["spikes"] sample_index = spikes_group["sample_index"][:] del spikes_group["sample_index"], spikes_group["sample_index_chunk_firsts"] - spikes_group.create_dataset("sample_index", data=sample_index, chunks=(97,)) - spikes_group.create_dataset("sample_index_chunk_firsts", data=sample_index[::97], compressor=None) + spikes_group.create_array("sample_index", data=sample_index, chunks=(97,)) + spikes_group.create_array("sample_index_chunk_firsts", data=sample_index[::97], compressor=None) assert spikes_group["sample_index"].nchunks > 3 in_ram = ZarrSortingExtractor(folder) diff --git a/src/spikeinterface/core/time_series_tools.py b/src/spikeinterface/core/time_series_tools.py index 758482b965..5d40733e2f 100644 --- a/src/spikeinterface/core/time_series_tools.py +++ b/src/spikeinterface/core/time_series_tools.py @@ -297,8 +297,9 @@ def _write_time_series_to_zarr( zarr_group, dataset_paths, dataset_timestamps_paths=None, - extra_chunks=None, dtype=None, + chunks=None, + shards=None, compressor_data=None, filters_data=None, compressor_times=None, @@ -308,6 +309,7 @@ def _write_time_series_to_zarr( ): """ Save the trace of a time_series object in several zarr format. + If shard Parameters ---------- @@ -319,10 +321,10 @@ def _write_time_series_to_zarr( List of paths to traces datasets in the zarr group dataset_timestamps_paths : list or None, default: None List of paths to timestamps datasets in the zarr group. If None, timestamps are not saved. - extra_chunks : tuple or None, default: None - Extra chunking dimensions to use for the zarr dataset. - The first dimension is always time and controlled by the job_kwargs. - This is for example useful to chunk by channel, with `extra_chunks=(channel_chunk_size,)`. + chunks : tuple or None, default: None + Chunking dimensions to use for the zarr dataset. + shards : tuple or None, default: None + Sharding configuration for the zarr dataset. dtype : dtype, default: None Type of the saved data compressor_data : zarr compressor or None, default: None @@ -359,13 +361,12 @@ def _write_time_series_to_zarr( dtype = time_series.get_dtype() job_kwargs = fix_job_kwargs(job_kwargs) - chunk_size = ensure_chunk_size(time_series, **job_kwargs) - if extra_chunks is not None: - assert len(extra_chunks) == len(time_series.get_shape(0)[1:]), ( - "extra_chunks should have the same length as the number of dimensions " - "of the time_series minus one (time axis)" - ) + # place an ArrayBytesCodec passed as a compressor (e.g. WavPack) in the serializer slot + from .zarrextractors import build_codec_pipeline + + codec_kwargs_data = build_codec_pipeline(filters=filters_data, compressors=compressor_data) + codec_kwargs_times = build_codec_pipeline(filters=filters_times, compressors=compressor_times) # create zarr datasets files zarr_datasets = [] @@ -375,25 +376,27 @@ def _write_time_series_to_zarr( num_samples = time_series.get_num_samples(segment_index) dset_name = dataset_paths[segment_index] shape = time_series.get_shape(segment_index) - dset = zarr_group.create_dataset( + dset = zarr_group.create_array( name=dset_name, shape=shape, - chunks=(chunk_size,) + extra_chunks if extra_chunks is not None else (chunk_size,), + chunks=chunks, + shards=shards, dtype=dtype, - filters=filters_data, - compressor=compressor_data, + **codec_kwargs_data, ) zarr_datasets.append(dset) if dataset_timestamps_paths[segment_index] is not None: tset_name = dataset_timestamps_paths[segment_index] + chunks_times = (chunks[0],) if chunks is not None else None + shards_times = (shards[0],) if shards is not None else None zarr_timestamps_datasets.append( - zarr_group.create_dataset( + zarr_group.create_array( name=tset_name, shape=(num_samples,), - chunks=(chunk_size,), + chunks=chunks_times, + shards=shards_times, dtype="float64", - filters=filters_times, - compressor=compressor_times, + **codec_kwargs_times, ) ) else: @@ -416,7 +419,7 @@ def _write_time_series_to_zarr( t_starts[segment_index] = time_info["t_start"] if np.any(~np.isnan(t_starts)): - zarr_group.create_dataset(name="t_starts", data=t_starts, compressor=None) + zarr_group.create_array(name="t_starts", data=t_starts, compressors=None) def _init_zarr_worker(time_series, zarr_datasets, dtype, zarr_timestamps_datasets=None): diff --git a/src/spikeinterface/core/zarr_tools.py b/src/spikeinterface/core/zarr_tools.py new file mode 100644 index 0000000000..38d43de7ad --- /dev/null +++ b/src/spikeinterface/core/zarr_tools.py @@ -0,0 +1,477 @@ +import warnings + +import numpy as np +import zarr + +from spikeinterface.core.job_tools import ensure_chunk_size + +# metadata keys that are not members of a group (zarr v2 and v3) +_ZARR_METADATA_KEYS = (".zarray", ".zattrs", ".zgroup", ".zmetadata", "zarr.json") + + +def adjust_chunks_shards_and_job_kwargs( + chunks=None, extra_chunks=None, shards=None, shard_factor=None, job_kwargs=None, time_series=None +): + """ + Adjust chunks, shards, and job_kwargs for zarr storage. + + Parameters + ---------- + chunks : tuple or None + Chunking dimensions for the zarr dataset. + extra_chunks : tuple or None + Extra chunking dimensions for the zarr dataset. + shards : tuple or None + Sharding configuration for the zarr dataset. + shard_factor : int or tuple or None + Factor to determine shard sizes based on chunks. + job_kwargs : dict or None + Job-related keyword arguments, including 'chunk_size'. + time_series : TimeSeries or None + The time series object being stored. + + Returns + ------- + chunks : tuple + Adjusted chunking dimensions. + shards : tuple or None + Adjusted sharding configuration. + job_kwargs : dict + Updated job-related keyword arguments. + """ + + # Chunking and sharding + if shards is not None and shard_factor is not None: + raise ValueError("Cannot specify both 'shards' and 'shard_factor' in zarr_kwargs") + if chunks is not None and extra_chunks is not None: + raise ValueError("Cannot specify both 'chunks' and 'extra_chunks' in zarr_kwargs") + + # If not specified by chunk, we set the chunk size in the first dimension (time) to be the chunk size that we use + # for the job executor, and the chunk size in the second dimension (channels) to be either the provided + # channel_chunk_size or the total number of channels (no chunking in channels). + if chunks is not None: + job_kwargs["chunk_size"] = chunks[0] + else: + chunk_size = ensure_chunk_size(time_series, **job_kwargs) + chunks = ( + (chunk_size,) + extra_chunks if extra_chunks is not None else (chunk_size,) + time_series.get_shape(0)[1:] + ) + + if shards is not None: + assert len(shards) == len(chunks), "Shards and chunks must have the same number of dimensions" + for dim in range(len(chunks)): + assert ( + shards[dim] >= chunks[dim] and shards[dim] % chunks[dim] == 0 + ), "Shard size must be a multiple of chunk size" + # When sharding is used, chunk_size in job_kwargs is used to determine the number of samples per chunk to + # write in each job. Each process will write all chunks in a shard. + job_kwargs["chunk_size"] = shards[0] + elif shard_factor is not None: + # If shard_factor is an integer, we only apply it to the first dimension. If it's an iterable, + # we apply it to all dimensions. + if isinstance(shard_factor, (int, np.integer)): + shards = (chunks[0] * shard_factor,) + chunks[1:] + else: + if len(shard_factor) != len(chunks): + raise ValueError("shard_factor must have the same length as chunks when it is an iterable") + shards = tuple(chunks[dim] * shard_factor[dim] for dim in range(len(chunks))) + job_kwargs["chunk_size"] = shards[0] + + return chunks, shards, job_kwargs + + +def check_compressors_match(comp1, comp2, skip_typesize=True): + """ + Check if two compressor objects match. + + Parameters + ---------- + comp1 : zarr.Codec | Tuple[zarr.Codec] + The first compressor object to compare. + comp2 : zarr.Codec | Tuple[zarr.Codec] + The second compressor object to compare. + skip_typesize : bool, optional + Whether to skip the typesize check, default: True + """ + if not isinstance(comp1, (list, tuple)): + assert not isinstance(comp2, list) + comp1 = [comp1] + comp2 = [comp2] + for i in range(len(comp1)): + comp1_dict = comp1[i].to_dict() + comp2_dict = comp2[i].to_dict() + if skip_typesize: + if "typesize" in comp1_dict["configuration"]: + comp1_dict["configuration"].pop("typesize", None) + if "typesize" in comp2_dict["configuration"]: + comp2_dict["configuration"].pop("typesize", None) + assert comp1_dict == comp2_dict, f"Compressor {i} does not match: {comp1_dict} != {comp2_dict}" + + +class LegacyZarrObjectArray: + """ + Read-only stand-in for a zarr v2 object-dtype array that zarr-python >= 3 cannot open. + + spikeinterface < 0.105 saved python objects (dicts, lists, provenance dicts, ...) as + length-1 object-dtype arrays using the `numcodecs` JSON or Pickle object codecs. + zarr-python >= 3 dropped support for these dtype/codec combinations, so such arrays + are decoded "manually" with `numcodecs` and exposed through this minimal wrapper, + which mimics the small part of the `zarr.Array` API used to read them + (`attrs`, `__getitem__`, `shape`, `dtype`, `len`). + + Parameters + ---------- + values : np.ndarray + The decoded object-dtype array. + attrs : dict + The zarr attributes of the array (content of `.zattrs`). + """ + + def __init__(self, values: np.ndarray, attrs: dict): + self._values = values + self._attrs = dict(attrs) + + @property + def attrs(self) -> dict: + return self._attrs + + @property + def shape(self): + return self._values.shape + + @property + def dtype(self): + return self._values.dtype + + def __len__(self): + return len(self._values) + + def __getitem__(self, key): + return self._values[key] + + def __array__(self, dtype=None, copy=None): + if dtype is None: + return self._values + return self._values.astype(dtype) + + def __repr__(self): + return f"LegacyZarrObjectArray(shape={self.shape}, attrs={list(self._attrs.keys())})" + + +def _sync(coroutine): + """Run a zarr async coroutine synchronously (zarr-python >= 3).""" + from zarr.core.sync import sync + + return sync(coroutine) + + +def _store_get_json(store, key: str): + """Read and json-decode a metadata key from a zarr store. Return None if missing.""" + import json + + from zarr.core.buffer import default_buffer_prototype + + buffer = _sync(store.get(key, prototype=default_buffer_prototype())) + if buffer is None: + return None + return json.loads(buffer.to_bytes().decode()) + + +def _store_get_bytes(store, key: str): + """Read raw bytes from a zarr store. Return None if the key is missing.""" + from zarr.core.buffer import default_buffer_prototype + + buffer = _sync(store.get(key, prototype=default_buffer_prototype())) + if buffer is None: + return None + return buffer.to_bytes() + + +def read_legacy_zarr_object_array(store, path: str) -> LegacyZarrObjectArray: + """ + Read a zarr v2 object-dtype array written with a `numcodecs` object codec + (JSON, Pickle, MsgPack), which zarr-python >= 3 cannot open. + + Only single-chunk arrays are supported, which is what spikeinterface < 0.105 wrote. + + Parameters + ---------- + store : zarr.abc.store.Store + The store containing the array. + path : str + The path of the array inside the store. + + Returns + ------- + legacy_array : LegacyZarrObjectArray + The decoded array. + """ + import numcodecs + + path = path.strip("/") + zarray = _store_get_json(store, f"{path}/.zarray") + if zarray is None: + raise KeyError(f"No zarr v2 array metadata found at {path}") + zattrs = _store_get_json(store, f"{path}/.zattrs") or {} + + shape = tuple(zarray["shape"]) + chunks = tuple(zarray["chunks"]) + if any(chunk_size < dim for chunk_size, dim in zip(chunks, shape)): + raise NotImplementedError( + f"Legacy zarr v2 object array at {path} has more than one chunk, which is not supported" + ) + + separator = zarray.get("dimension_separator", ".") + chunk_key = separator.join(["0"] * max(len(shape), 1)) + chunk_bytes = _store_get_bytes(store, f"{path}/{chunk_key}") + if chunk_bytes is None: + # array was never written to: return an empty object array + return LegacyZarrObjectArray(np.empty(shape, dtype=object), zattrs) + + if zarray.get("compressor", None) is not None: + chunk_bytes = numcodecs.get_codec(zarray["compressor"]).decode(chunk_bytes) + values = chunk_bytes + # filters are applied on encode, so they are decoded in reverse order. + # for object arrays the object codec is the (only) filter + for filter_config in reversed(zarray.get("filters", None) or []): + values = numcodecs.get_codec(filter_config).decode(values) + values = np.asarray(values, dtype=object).reshape(shape) + + return LegacyZarrObjectArray(values, zattrs) + + +def get_zarr_attr_or_legacy_object(zarr_group, name: str): + """ + Get a value saved either in the attributes of a zarr group (spikeinterface >= 0.105) + or as a legacy zarr v2 length-1 object array with the same name + (spikeinterface < 0.105, e.g. "recording" and "sorting_provenance"). + + Parameters + ---------- + zarr_group : zarr.Group + The zarr group to read from. + name : str + The name of the attribute / legacy array. + + Returns + ------- + value : Any | None + The value, or None if it is not found (or cannot be decoded). + """ + value = zarr_group.attrs.get(name, None) + if value is not None: + return value + try: + legacy_array = read_legacy_zarr_object_array(zarr_group.store, f"{zarr_group.path}/{name}") + except Exception: + return None + if len(legacy_array) == 0: + return None + return legacy_array[0] + + +def _list_member_names(zarr_group) -> list[str]: + """List the member names of a zarr group without opening them.""" + consolidated = getattr(zarr_group.metadata, "consolidated_metadata", None) + if consolidated is not None: + return list(consolidated.metadata.keys()) + + store = zarr_group.store + if not store.supports_listing: + raise ValueError( + f"The store associated to this group ({type(store).__name__}) does not support listing, " + "so its members cannot be listed without consolidated metadata." + ) + + async def _list_dir(): + return [key async for key in store.list_dir(zarr_group.path)] + + keys = _sync(_list_dir()) + # skip zarr metadata documents and hidden files (e.g. AppleDouble "._*" files) + return sorted(key for key in keys if key not in _ZARR_METADATA_KEYS and not key.startswith(".")) + + +def iterate_zarr_group(zarr_group, skip_unreadable: bool = True): + """ + Iterate over the members (arrays and sub-groups) of a zarr group, transparently + handling zarr v2 and zarr v3 groups. + + Unlike `zarr.Group.keys()` / `zarr.Group.members()`, a single member that + zarr-python >= 3 cannot open does not make the whole iteration fail. This happens for + zarr v2 data saved by spikeinterface < 0.105, where python objects were stored as + object-dtype arrays with the `numcodecs` JSON/Pickle object codecs. Such arrays are + decoded with `numcodecs` and returned as a `LegacyZarrObjectArray`. + + Parameters + ---------- + zarr_group : zarr.Group + The zarr group to iterate over. + skip_unreadable : bool, default: True + If True, members that cannot be opened nor decoded are skipped with a warning. + If False, an error is raised instead. + + Yields + ------ + name : str + The name of the member. + member : zarr.Group | zarr.Array | LegacyZarrObjectArray + The member itself. + """ + for name in _list_member_names(zarr_group): + # note: members are retrieved outside of the yield statement, so that exceptions + # raised by the consumer of this generator are not caught here + member = None + open_error = None + try: + member = zarr_group[name] + except KeyError: + # the key is an object in the store (e.g. a stray file), not a zarr node + continue + except Exception as e: + open_error = e + + if open_error is not None: + # the member exists but zarr-python cannot open it: try the legacy object array path + try: + member = read_legacy_zarr_object_array(zarr_group.store, f"{zarr_group.path}/{name}") + except Exception as legacy_error: + if not skip_unreadable: + raise ValueError( + f"Cannot read member '{name}' of zarr group '{zarr_group.path}': " + f"{open_error}\nLegacy object array fallback failed with: {legacy_error}" + ) from open_error + warnings.warn( + f"Skipping member '{name}' of zarr group '{zarr_group.path}' because it cannot be read " + f"with zarr v{zarr.__version__}: {open_error} ({legacy_error})" + ) + continue + + yield name, member + + +def get_zarr_group_keys(zarr_group) -> list[str]: + """ + Return the names of the members of a zarr group, for both zarr v2 and zarr v3 groups. + + Contrary to `zarr.Group.keys()`, the members are not opened, so this also works for + groups containing legacy zarr v2 arrays that zarr-python >= 3 cannot open + (see `iterate_zarr_group`). + + Parameters + ---------- + zarr_group : zarr.Group + The zarr group to list. + + Returns + ------- + keys : list[str] + The names of the members of the group. + """ + return _list_member_names(zarr_group) + + +def is_sklearn_estimator(obj) -> bool: + """ + Check whether an object looks like a fitted scikit-learn estimator + (e.g. the PCA models of the "principal_components" extension). + """ + return ( + callable(getattr(obj, "get_params", None)) + and callable(getattr(obj, "set_params", None)) + and hasattr(obj, "__dict__") + ) + + +def save_sklearn_model_to_zarr_group(parent_group, name: str, model, **saving_options) -> None: + """ + Save a scikit-learn estimator in a zarr sub-group without using pickle. + + The state of the estimator (`vars(model)`) is split in two: numpy arrays are saved as + zarr arrays and the remaining (json-serializable) values are saved in the + "sklearn_model" attribute of the group, together with the class of the estimator. + See `load_sklearn_model_from_zarr_group` for the reverse operation. + + Parameters + ---------- + parent_group : zarr.Group + The zarr group in which the sub-group is created. + name : str + The name of the sub-group. + model : sklearn estimator + The estimator to save. + **saving_options : dict + Options passed to `zarr.Group.create_array` for the array parts of the state. + + Raises + ------ + ValueError + If part of the state of the estimator can neither be saved as a zarr array nor + serialized to json (e.g. a nested estimator or an arbitrary python object). + """ + import json + + state = {} + arrays = {} + for key, value in vars(model).items(): + if isinstance(value, np.ndarray): + if value.dtype.kind == "O": + raise ValueError(f"Cannot save object-dtype array '{key}' of {type(model).__name__} to zarr") + arrays[key] = value + else: + if isinstance(value, np.generic): + value = value.item() + try: + json.dumps(value) + except TypeError: + raise ValueError( + f"Cannot save attribute '{key}' of {type(model).__name__} to zarr: " + f"{type(value).__name__} is not json-serializable" + ) + state[key] = value + + model_group = parent_group.create_group(name) + for key, value in arrays.items(): + model_group.create_array(name=key, data=value, **saving_options) + + model_info = { + "class": f"{type(model).__module__}.{type(model).__qualname__}", + "state": state, + "array_keys": sorted(arrays.keys()), + } + try: + import sklearn + + model_info["sklearn_version"] = sklearn.__version__ + except ImportError: + pass + model_group.attrs["sklearn_model"] = model_info + + +def load_sklearn_model_from_zarr_group(model_group): + """ + Rebuild a scikit-learn estimator saved by `save_sklearn_model_to_zarr_group`. + + Parameters + ---------- + model_group : zarr.Group + The zarr group containing the estimator state. + + Returns + ------- + model : sklearn estimator + The estimator, with the same state it had when saved. + """ + import importlib + + model_info = model_group.attrs["sklearn_model"] + module_name, _, class_name = model_info["class"].rpartition(".") + model_class = getattr(importlib.import_module(module_name), class_name) + + # the full state is restored, so the constructor is bypassed on purpose + model = model_class.__new__(model_class) + for key, value in model_info["state"].items(): + setattr(model, key, value) + for key in model_info["array_keys"]: + setattr(model, key, np.asarray(model_group[key][...])) + + return model diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index d2a604e38b..ef95ee673e 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -10,16 +10,19 @@ from .base import minimum_spike_dtype, _get_class_from_string from .baserecording import BaseRecording, BaseRecordingSegment from .basesorting import BaseSorting, SpikeVectorSortingSegment -from .job_tools import split_job_kwargs from .core_tools import ( define_function_from_class, check_json, + is_path_remote, retrieve_importing_provenance, is_path_remote, - _check_zarr_write_is_supported, ) +from .job_tools import split_job_kwargs, fix_job_kwargs +from .zarr_tools import iterate_zarr_group, adjust_chunks_shards_and_job_kwargs from .time_series_tools import _write_time_series_to_zarr +zarr.config.set({"default_zarr_version": 3}) + def super_zarr_open(folder_path: str | Path, mode: str = "r", storage_options: dict | None = None): """ @@ -55,11 +58,11 @@ def super_zarr_open(folder_path: str | Path, mode: str = "r", storage_options: d import zarr # if mode is append or read/write, we try to open the folder with zarr.open - # since zarr.open_consolidated does not support creating new groups/datasets + # In zarr v3, we use use_consolidated parameter instead of open_consolidated if mode in ("a", "r+"): - open_funcs = (zarr.open,) + use_consolidated_options = (False,) else: - open_funcs = (zarr.open_consolidated, zarr.open) + use_consolidated_options = (True, False) # if storage_options is None, we try to open the folder with and without anonymous access # if storage_options is not None, we try to open the folder with the given storage options @@ -70,13 +73,16 @@ def super_zarr_open(folder_path: str | Path, mode: str = "r", storage_options: d root = None exception = None - if is_path_remote(folder_path): - for open_func in open_funcs: + if is_path_remote(str(folder_path)): + from zarr.storage import FsspecStore + + for use_consolidated in use_consolidated_options: if root is not None: break for storage_options in storage_options_to_test: try: - root = open_func(str(folder_path), mode=mode, storage_options=storage_options) + store = FsspecStore.from_url(str(folder_path), storage_options=storage_options) + root = zarr.open(store, mode=mode, use_consolidated=use_consolidated) break except Exception as e: exception = e @@ -84,11 +90,9 @@ def super_zarr_open(folder_path: str | Path, mode: str = "r", storage_options: d else: if not Path(folder_path).is_dir(): raise ValueError(f"Folder {folder_path} does not exist") - # zarr>=3 refuses storage_options for a local path, even an empty dict - local_kwargs = dict(storage_options=storage_options) if storage_options else dict() - for open_func in open_funcs: + for use_consolidated in use_consolidated_options: try: - root = open_func(str(folder_path), mode=mode, **local_kwargs) + root = zarr.open(str(folder_path), mode=mode, use_consolidated=use_consolidated) break except Exception as e: exception = e @@ -139,7 +143,8 @@ def __init__( assert sampling_frequency is not None, "'sampling_frequency' attiribute not found!" assert num_segments is not None, "'num_segments' attiribute not found!" - channel_ids = np.array(channel_ids) + # zarr returns vlen-utf8 as StringDType (numpy 2.0); convert via list to classic unicode array. + channel_ids = np.array(channel_ids.tolist()) dtype = self._root["traces_seg0"].dtype @@ -177,7 +182,7 @@ def __init__( if load_compression_ratio: nbytes_segment = self._root[trace_name].nbytes - nbytes_stored_segment = self._root[trace_name].nbytes_stored + nbytes_stored_segment = self._root[trace_name].nbytes_stored() if nbytes_stored_segment > 0: cr_by_segment[segment_index] = nbytes_segment / nbytes_stored_segment else: @@ -218,11 +223,15 @@ def __init__( # load properties if "properties" in self._root: prop_group = self._root["properties"] - for key in prop_group.keys(): + for key, values in iterate_zarr_group(prop_group): # Skip contact_vector property since it is not used anymore to represent probegroup if key == "contact_vector": continue - values = self._root["properties"][key] + values = values[:] + # zarr returns vlen-utf8 as StringDType (numpy 2.0); convert via list to classic unicode array. + if hasattr(values.dtype, "na_object") or values.dtype.kind == "O": + if values.size > 0 and isinstance(values.tolist()[0], str): + values = np.array(values.tolist()) self.set_property(key, values) # load annotations @@ -575,30 +584,43 @@ def __init__( unit_ids = np.array(unit_ids) assert "spikes" in self._root.keys(), "'spikes' dataset not found!" - spikes_group = self._root["spikes"] - segment_slices_list = spikes_group["segment_slices"][:] + spikes_item = self._root["spikes"] BaseSorting.__init__(self, sampling_frequency, unit_ids) - if lazy_spike_vector: - spikes = ZarrSpikeVector(spikes_group, segment_slices_list) + if isinstance(spikes_item, zarr.Group): + # Legacy format: individual field arrays + a "segment_slices" sub-array inside the group. + spikes_group = spikes_item + segment_slices_list = np.asarray(spikes_group["segment_slices"][:], dtype="int64") + + if lazy_spike_vector: + spikes = ZarrSpikeVector(spikes_group, segment_slices_list) + else: + spikes = np.zeros(spikes_group["sample_index"].shape[0], dtype=minimum_spike_dtype) + spikes["sample_index"] = spikes_group["sample_index"][:] + spikes["unit_index"] = spikes_group["unit_index"][:] + for i, (start, end) in enumerate(segment_slices_list): + spikes["segment_index"][start:end] = i else: - # Materialize the spike vector in memory and sort it by (segment_index, sample_index, unit_index) - spikes = np.zeros(spikes_group["sample_index"].shape[0], dtype=minimum_spike_dtype) - spikes["sample_index"] = spikes_group["sample_index"][:] - spikes["unit_index"] = spikes_group["unit_index"][:] - for i, (start, end) in enumerate(segment_slices_list): - spikes["segment_index"][start:end] = i - # we do not need to lexsort at init (very high cost) because there already sorted by frame before to be saved. - # In version 0.104.X this was fully lexsorted, but we don't need it anymore because it's only important in the context of SpikeVectorBased extensions in the SortingAnalyzer, which stores its own copy of the Sorting object. This makes the extension data and the spike vector always matching their order. - # spikes = spikes[np.lexsort((spikes["unit_index"], spikes["sample_index"], spikes["segment_index"]))] + # New format: a single structured zarr array; segment_slices stored as an array attribute. + # We need https://github.com/zarr-developers/zarr-python/pull/3996 released before being able to + # access the structured array lazily. Until then, we always materialise it. + spikes = spikes_item + segment_slices_list = np.asarray(spikes.attrs["segment_slices"], dtype="int64") + + # we do not need to lexsort at init (very high cost) because spikes are already sorted by frame before saving. + # In version 0.104.X this was fully lexsorted, but we don't need it anymore because it's only important in the + # context of SpikeVectorBased extensions in the SortingAnalyzer, which stores its own copy of the Sorting + # object. This makes the extension data and the spike vector always matching their order. + # spikes = spikes[np.lexsort((spikes["unit_index"], spikes["sample_index"], spikes["segment_index"]))] self._lazy_spike_vector = lazy_spike_vector self._spikes_group = spikes_group self._sample_index_search = None + self._cached_spike_vector = spikes # pre-populate segment slices so _get_spike_vector_segment_slices() never # needs to materialise the full segment_index array - self._cached_spike_vector_segment_slices = np.asarray(segment_slices_list, dtype="int64") + self._cached_spike_vector_segment_slices = segment_slices_list for segment_index in range(num_segments): soring_segment = SpikeVectorSortingSegment(spikes, segment_index, unit_ids) @@ -607,8 +629,12 @@ def __init__( # load properties if "properties" in self._root: prop_group = self._root["properties"] - for key in prop_group.keys(): - values = self._root["properties"][key] + for key, values in iterate_zarr_group(prop_group): + values = values[:] + # zarr returns vlen-utf8 as StringDType (numpy 2.0); convert via list to classic unicode array. + if hasattr(values.dtype, "na_object") or values.dtype.kind == "O": + if values.size > 0 and isinstance(values.tolist()[0], str): + values = np.array(values.tolist()) self.set_property(key, values) # load annotations @@ -731,7 +757,6 @@ def create_zarr_path_for_write(folder_path: str | Path, overwrite: bool = False) folder_path : str or Path Path to the zarr root file """ - _check_zarr_write_is_supported() if not is_path_remote(folder_path): folder_path = Path(folder_path) folder_path = folder_path.with_suffix(".zarr") @@ -744,51 +769,6 @@ def create_zarr_path_for_write(folder_path: str | Path, overwrite: bool = False) return folder_path -def _write_object_array( - group, - name: str, - data, - codec: str = "json", - overwrite: bool = True, -): - """ - Write a length-1 object-dtype array holding a Python dict/list/object. - - Centralizes the v2/v3 codec-placement difference for object blobs: under zarr-v2 - the object codec goes in ``object_codec=``; under zarr-v3 it goes in ``filters=`` - (wrapped via ``numcodecs.zarr3.*``). The helper picks the right path automatically. - - Parameters - ---------- - group : zarr.Group - The zarr group to write into. - name : str - Name of the array inside ``group``. - data : Any - The Python object to store. Wrapped into ``np.array([data], dtype=object)``. - codec : {"json", "pickle"}, default: "json" - Which object codec to use. - overwrite : bool, default: True - Whether to overwrite an existing array with the same name. - """ - import numcodecs - - if codec == "json": - codec_instance = numcodecs.JSON() - elif codec == "pickle": - codec_instance = numcodecs.Pickle() - else: - raise ValueError(f"codec must be 'json' or 'pickle', got {codec!r}") - - arr = np.array([data], dtype=object) - return group.create_dataset( - name=name, - data=arr, - object_codec=codec_instance, - overwrite=overwrite, - ) - - def get_default_zarr_compressor(clevel: int = 5): """ Return default Zarr compressor object for good preformance in int16 @@ -809,9 +789,83 @@ def get_default_zarr_compressor(clevel: int = 5): Blosc.compressor The compressor object that can be used with the save to zarr function """ - from numcodecs import Blosc + from zarr.codecs import BloscCodec, BloscShuffle + + return BloscCodec(cname="zstd", clevel=clevel, shuffle=BloscShuffle.bitshuffle) + + +def build_codec_pipeline(filters=None, compressors=None): + """ + Build zarr v3 codec kwargs from filters and compressors. - return Blosc(cname="zstd", clevel=clevel, shuffle=Blosc.BITSHUFFLE) + Classifies codecs into the three slots accepted by ``zarr.Group.create_array()``: + 1. ``filters`` — ArrayArrayCodec (e.g. Delta) + 2. ``serializer`` — ArrayBytesCodec (e.g. WavPack, BytesCodec) + 3. ``compressors``— BytesBytesCodec (e.g. BloscCodec, ZstdCodec) + + This allows callers to pass an ArrayBytesCodec (e.g. WavPack) as a + compressor and have it placed in the correct serializer slot automatically. + + Parameters + ---------- + filters : ArrayArrayCodec or list of ArrayArrayCodec or None + Codec(s) applied before serialization. + compressors : codec or list of codecs or None + Can be a mix of ArrayBytesCodec (serializer) and BytesBytesCodec + (byte-level compressors). At most one ArrayBytesCodec is allowed. + + Returns + ------- + dict + Keyword arguments to unpack into ``zarr.Group.create_array()``. + Only keys with explicit values are included; omitted keys let zarr + use its defaults. + + Raises + ------ + ValueError + If filters contain non-ArrayArrayCodec instances, if more than one + ArrayBytesCodec is provided, or if an unrecognised codec type is passed. + """ + from zarr.abc.codec import ArrayArrayCodec, ArrayBytesCodec, BytesBytesCodec + + if filters is None: + filters = [] + if not isinstance(filters, (list, tuple)): + filters = [filters] + + if compressors is None: + compressors = [] + if not isinstance(compressors, (list, tuple)): + compressors = [compressors] + + for f in filters: + if not isinstance(f, ArrayArrayCodec): + raise ValueError(f"All filters must be ArrayArrayCodec instances, got {type(f)}") + + serializers = [c for c in compressors if isinstance(c, ArrayBytesCodec)] + byte_compressors = [c for c in compressors if isinstance(c, BytesBytesCodec)] + invalid = [c for c in compressors if not isinstance(c, (ArrayBytesCodec, BytesBytesCodec))] + + if invalid: + raise ValueError( + f"Compressors must be ArrayBytesCodec or BytesBytesCodec instances, got {[type(c) for c in invalid]}" + ) + if len(serializers) > 1: + raise ValueError("Only one ArrayBytesCodec (serializer) is allowed in the codec pipeline.") + + codec_kwargs = {} + codec_kwargs["filters"] = filters + codec_kwargs["serializer"] = serializers[0] if len(serializers) == 1 else "auto" + codec_kwargs["compressors"] = byte_compressors + return codec_kwargs + + +def _has_string_fields(dtype: np.dtype) -> bool: + """Return True if dtype is or contains fixed-length unicode (U) sub-fields.""" + if dtype.names: + return any(_has_string_fields(dtype.fields[name][0]) for name in dtype.names) + return dtype.kind == "U" def add_properties_and_annotations(zarr_group: zarr.Group, recording_or_sorting: BaseRecording | BaseSorting): @@ -822,7 +876,20 @@ def add_properties_and_annotations(zarr_group: zarr.Group, recording_or_sorting: if values.dtype.kind == "O": warnings.warn(f"Property {key} not saved because it is a python Object type") continue - prop_group.create_dataset(name=key, data=values, compressor=None) + if values.dtype.names and _has_string_fields(values.dtype): + # Structured arrays with unicode sub-fields have no stable zarr v3 spec; skip them. + # Probe geometry (contact_vector) is already persisted via zarr_group.attrs["probe"]. + warnings.warn( + f"Property '{key}' not saved because it is a structured array with unicode fields, " + "which do not have a stable zarr V3 specification." + ) + continue + # Use variable-length UTF-8 (stable zarr v3 spec) for unicode arrays. + if values.dtype.kind == "U": + arr = prop_group.create_array(name=key, shape=values.shape, dtype=str, compressors=None) + arr[:] = values + else: + prop_group.create_array(name=key, data=values, compressors=None) # save annotations zarr_group.attrs["annotations"] = check_json(recording_or_sorting._annotations) @@ -841,9 +908,15 @@ def add_sorting_to_zarr_group( zarr_group : zarr.Group The zarr group kwargs : dict - Other arguments passed to the zarr compressor + Other arguments passed to the zarr writer: + + * "compressors" : zarr compressor or None, default: None + * "target_chunk_size_bytes" : tuple or None, default: None + The target chunk size for the zarr datasets. + * "shard_factor" : int or None, default: None + If given, a shard will be shard_factor * chunk size. """ - from numcodecs import Delta + from zarr.codecs import Delta if sorting.check_serializability("json"): zarr_group.attrs["provenance"] = check_json(sorting.to_dict(recursive=True, relative_to=relative_to)) @@ -853,32 +926,73 @@ def add_sorting_to_zarr_group( num_segments = sorting.get_num_segments() zarr_group.attrs["sampling_frequency"] = float(sorting.sampling_frequency) zarr_group.attrs["num_segments"] = int(num_segments) - zarr_group.create_dataset(name="unit_ids", data=sorting.unit_ids, compressor=None) + zarr_group.create_array(name="unit_ids", data=sorting.unit_ids, compressors=None) + + compressor = kwargs.get("compressors") or kwargs.get("compressor") + if compressor is None: + compressor = get_default_zarr_compressor() - compressor = kwargs.get("compressor", get_default_zarr_compressor()) + # Save the full structured spike array. The "segment_index" field is additionally stored as + # "segment_slices" (an attribute on the spikes array) which contains the start and end indices + # of spikes for each segment, to allow efficient per-segment access without scanning the array. + spikes = sorting.to_spike_vector() - # save sub fields + chunks = None + shards = None + target_chunk_size_bytes = kwargs.get("target_chunk_size_bytes") + + # We need https://github.com/zarr-developers/zarr-python/pull/3996 released before being able to + # access the structured array lazily. Until then, we always materialise it. + # For now, let's keep the old by field implementation + # if target_chunk_size_bytes is not None: + # spike_num_bytes = spikes.dtype.itemsize + # target_chunk_size = target_chunk_size_bytes // spike_num_bytes + # chunks = (target_chunk_size,) + # shard_factor = kwargs.get("shard_factor") + # if shard_factor is not None: + # shards = (target_chunk_size * shard_factor,) + # else: + # chunks = (len(spikes),) + # spikes = zarr_group.create_array( + # name="spikes", + # data=spikes, + # chunks=chunks, + # shards=shards, + # **codec_kwargs + # ) + + # Save sub fields: in this case chunks and shards are set per field spikes_group = zarr_group.create_group(name="spikes") spikes = sorting.to_spike_vector() for field in spikes.dtype.fields: if field != "segment_index": - spikes_group.create_dataset( - name=field, - data=spikes[field], - compressor=compressor, - filters=[Delta(dtype=spikes[field].dtype)], - ) - else: - segment_slices = [] - for segment_index in range(num_segments): - i0, i1 = np.searchsorted(spikes["segment_index"], [segment_index, segment_index + 1]) - segment_slices.append([i0, i1]) - spikes_group.create_dataset(name="segment_slices", data=segment_slices, compressor=None) + if target_chunk_size_bytes is not None: + spike_num_bytes = spikes[field].dtype.itemsize + target_chunk_size = target_chunk_size_bytes // spike_num_bytes + chunks = (target_chunk_size,) + shard_factor = kwargs.get("shard_factor") + if shard_factor is not None: + shards = (target_chunk_size * shard_factor,) + else: + chunks = (len(spikes),) + codec_kwargs = build_codec_pipeline(filters=Delta(dtype=spikes[field].dtype.str), compressors=compressor) + spikes_group.create_array(name=field, data=spikes[field], chunks=chunks, shards=shards, **codec_kwargs) + + segment_slices = [] + if sorting._cached_spike_vector_segment_slices is not None: + segment_slices = sorting._cached_spike_vector_segment_slices + else: + for segment_index in range(num_segments): + i0, i1 = np.searchsorted(spikes["segment_index"], [segment_index, segment_index + 1]) + segment_slices.append([i0, i1]) + segment_slices = np.array(segment_slices, dtype="int64") + spikes_group.create_array(name="segment_slices", data=segment_slices, compressors=None) + # spikes.attrs["segment_slices"] = segment_slices.tolist() # first sample_index of every zarr chunk: lets a lazy reader search sample_index # one chunk at a time (see ZarrSampleIndexSearch) instead of materialising it chunk_length = spikes_group["sample_index"].chunks[0] - spikes_group.create_dataset( + spikes_group.create_array( name="sample_index_chunk_firsts", data=np.asarray(spikes["sample_index"][::chunk_length], dtype="int64"), compressor=None, @@ -896,6 +1010,7 @@ def add_recording_to_zarr_group( **kwargs, ): zarr_kwargs, job_kwargs = split_job_kwargs(kwargs) + job_kwargs = fix_job_kwargs(job_kwargs) if recording.check_serializability("json"): zarr_group.attrs["provenance"] = check_json(recording.to_dict(recursive=True, relative_to=relative_to)) @@ -905,7 +1020,11 @@ def add_recording_to_zarr_group( # save data (done the subclass) zarr_group.attrs["sampling_frequency"] = float(recording.get_sampling_frequency()) zarr_group.attrs["num_segments"] = int(recording.get_num_segments()) - zarr_group.create_dataset(name="channel_ids", data=recording.get_channel_ids(), compressor=None) + # Use variable-length UTF-8 (stable zarr v3 spec) instead of fixed-length unicode. + channel_ids = recording.channel_ids + if channel_ids.dtype.kind in ("U", "S"): + channel_ids = channel_ids.astype("T") + arr = zarr_group.create_array(name="channel_ids", data=channel_ids, compressors=None) dataset_paths = [f"traces_seg{i}" for i in range(recording.get_num_segments())] dataset_timestamps_paths: list | None = None if any(recording.has_time_vector(i) for i in range(recording.get_num_segments())): @@ -917,16 +1036,28 @@ def add_recording_to_zarr_group( dataset_timestamps_paths.append(None) dtype = recording.get_dtype() if dtype is None else dtype - channel_chunk_size = zarr_kwargs.get("channel_chunk_size", None) - global_compressor = zarr_kwargs.pop("compressor", get_default_zarr_compressor()) + + # Compressors and filters + global_compressor = kwargs.get("compressors") or kwargs.get("compressor") + if global_compressor is None: + global_compressor = get_default_zarr_compressor() compressor_by_dataset = zarr_kwargs.pop("compressor_by_dataset", {}) global_filters = zarr_kwargs.pop("filters", None) filters_by_dataset = zarr_kwargs.pop("filters_by_dataset", {}) - compressor_traces = compressor_by_dataset.get("traces", global_compressor) filters_traces = filters_by_dataset.get("traces", global_filters) compressor_times = compressor_by_dataset.get("times", global_compressor) filters_times = filters_by_dataset.get("times", global_filters) + channel_chunk_size = zarr_kwargs.get("channel_chunk_size") + + chunks, shards, job_kwargs = adjust_chunks_shards_and_job_kwargs( + chunks=zarr_kwargs.get("chunks"), + extra_chunks=(channel_chunk_size,) if channel_chunk_size is not None else None, + shards=zarr_kwargs.get("shards"), + shard_factor=zarr_kwargs.get("shard_factor"), + job_kwargs=job_kwargs, + time_series=recording, + ) _write_time_series_to_zarr( time_series=recording, @@ -936,10 +1067,11 @@ def add_recording_to_zarr_group( compressor_data=compressor_traces, filters_data=filters_traces, dtype=dtype, - extra_chunks=(channel_chunk_size,), + chunks=chunks, + shards=shards, compressor_times=compressor_times, filters_times=filters_times, - verbose=False, + verbose=verbose, **job_kwargs, ) diff --git a/src/spikeinterface/extractors/nwbextractors.py b/src/spikeinterface/extractors/nwbextractors.py index 0d1946d39a..5aacb9a266 100644 --- a/src/spikeinterface/extractors/nwbextractors.py +++ b/src/spikeinterface/extractors/nwbextractors.py @@ -316,6 +316,29 @@ def _get_backend_from_local_file(file_path: str | Path) -> str: return backend +def _zarr_group_child_names(group): + """ + Return the names of the immediate children of a zarr group without parsing their metadata. + + zarr-python 3.x eagerly reads and validates every child's metadata when iterating + ``group.keys()``. Some arrays written by hdmf-zarr (e.g. variable-length string columns + with an integer ``fill_value``) cannot be parsed by zarr-python 3.x and make the whole + iteration fail. Listing the store directly avoids touching the children's metadata. + """ + if hasattr(group, "store_path"): # zarr v3 + from zarr.core.sync import sync + + async def _collect(): + return [key async for key in group.store.list_dir(group.path)] + + # Filter out this group's own metadata files (".zgroup", ".zattrs", "zarr.json", ...). + # list_dir does not guarantee an order, so sort for deterministic traversal (matches h5py). + names = [name for name in sync(_collect()) if not name.startswith(".") and name != "zarr.json"] + return sorted(names) + else: # zarr v2 + return list(group.keys()) + + def _find_neurodata_type_from_backend(group, path="", result=None, neurodata_type="ElectricalSeries", backend="hdf5"): """ Recursively searches for groups with the specified neurodata_type hdf5 or zarr object, @@ -325,17 +348,24 @@ def _find_neurodata_type_from_backend(group, path="", result=None, neurodata_typ import h5py group_class = h5py.Group + child_names = list(group.keys()) else: import zarr group_class = zarr.Group + child_names = _zarr_group_child_names(group) if result is None: result = [] - # zarr>=3 groups have no `items()`, `keys()` works for h5py and both zarr versions - for neurodata_name in group.keys(): - value = group[neurodata_name] + for neurodata_name in child_names: + try: + value = group[neurodata_name] + except Exception: + # Skip children whose metadata cannot be parsed (e.g. hdmf-zarr arrays with a + # fill_value that zarr-python 3.x rejects). These are never groups, so skipping + # them is safe when searching for a neurodata_type. + continue # Check if it's a group and if it has the neurodata_type if isinstance(value, group_class): current_path = f"{path}/{neurodata_name}" if path else neurodata_name @@ -350,14 +380,16 @@ def _find_neurodata_type_from_backend(group, path="", result=None, neurodata_typ def _retrieve_electrodes_indices_from_electrical_series_backend(open_file, electrical_series, backend="hdf5"): """ Retrieves the indices of the electrodes from the electrical series. - For the Zarr backend, the electrodes are stored in the electrical_series.attrs["zarr_link"]. + For the Zarr backend, the electrodes are stored in the electrical_series.attrs["_LINKS"] + or legacy electrical_series.attrs["zarr_link"]. + See https://github.com/hdmf-dev/hdmf-zarr/pull/336 """ if "electrodes" not in electrical_series: if backend == "zarr": import zarr - # links must be resolved, hdmf-zarr>=0.14 stores them under "_LINKS" instead of "zarr_link" - zarr_links = electrical_series.attrs.get("zarr_link", electrical_series.attrs.get("_LINKS")) + # links must be resolved + zarr_links = electrical_series.attrs.get("_LINKS", electrical_series.attrs.get("zarr_link", [])) electrodes_path = None for zarr_link in zarr_links: if zarr_link["name"] == "electrodes": diff --git a/src/spikeinterface/metrics/spiketrain/spiketrain_metrics.py b/src/spikeinterface/metrics/spiketrain/spiketrain_metrics.py index 79ee3e1b5e..375cd42b1c 100644 --- a/src/spikeinterface/metrics/spiketrain/spiketrain_metrics.py +++ b/src/spikeinterface/metrics/spiketrain/spiketrain_metrics.py @@ -38,7 +38,7 @@ class ComputeSpikeTrainMetrics(BaseMetricExtension): extension_name = "spiketrain_metrics" depend_on = [] - need_backward_compatibility_on_load = True + need_backward_compatibility_on_load = False metric_list = spiketrain_metrics diff --git a/src/spikeinterface/postprocessing/tests/common_extension_tests.py b/src/spikeinterface/postprocessing/tests/common_extension_tests.py index fb7770befe..da164a48b0 100644 --- a/src/spikeinterface/postprocessing/tests/common_extension_tests.py +++ b/src/spikeinterface/postprocessing/tests/common_extension_tests.py @@ -9,7 +9,6 @@ estimate_sparsity, ) from spikeinterface.core.sortinganalyzer import get_extension_class -from spikeinterface.core.core_tools import _is_zarr_write_supported extensions_which_allow_unit_ids = ["unit_locations"] extensions_with_unit_by_unit_data = ["correlograms", "template_similarity"] @@ -242,8 +241,7 @@ def run_extension_tests(self, extension_class, params): of interest with the passed parameters. Will perform tests for sparsity and format. """ - # TODO: remove once writing to zarr is supported with zarr>=3 - formats = ("memory", "binary_folder", "zarr") if _is_zarr_write_supported() else ("memory", "binary_folder") + formats = ("memory", "binary_folder", "zarr") for sparse in (True, False): for format in formats: print("sparse", sparse, format) diff --git a/src/spikeinterface/postprocessing/tests/test_multi_extensions.py b/src/spikeinterface/postprocessing/tests/test_multi_extensions.py index 6ae72e374d..7b53f24f29 100644 --- a/src/spikeinterface/postprocessing/tests/test_multi_extensions.py +++ b/src/spikeinterface/postprocessing/tests/test_multi_extensions.py @@ -133,9 +133,7 @@ def dataset_to_split(create_cache_folder): @pytest.mark.parametrize("lazy", [False, True]) @pytest.mark.parametrize("sparse", [False, True]) -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) def test_SortingAnalyzer_merge_all_extensions(dataset_to_merge, lazy, sparse, format, tmp_path): if format == "memory" and lazy: pytest.skip("lazy has no effect for format='memory' (nothing on disk to load lazily)") @@ -267,9 +265,7 @@ def test_SortingAnalyzer_merge_all_extensions(dataset_to_merge, lazy, sparse, fo @pytest.mark.parametrize("lazy", [False, True]) @pytest.mark.parametrize("sparse", [False, True]) -@pytest.mark.parametrize( - "format", ["memory", "binary_folder", pytest.param("zarr", marks=pytest.mark.requires_zarr_write)] -) +@pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) def test_SortingAnalyzer_split_all_extensions(dataset_to_split, lazy, sparse, format, tmp_path): if format == "memory" and lazy: pytest.skip("lazy has no effect for format='memory' (nothing on disk to load lazily)") diff --git a/src/spikeinterface/preprocessing/pipeline.py b/src/spikeinterface/preprocessing/pipeline.py index 0fe320681c..87c55f72ee 100644 --- a/src/spikeinterface/preprocessing/pipeline.py +++ b/src/spikeinterface/preprocessing/pipeline.py @@ -263,11 +263,9 @@ def get_preprocessing_list_from_analyzer(analyzer_folder, format="auto", backend storage_options = backend_options.get("storage_options", {}) zarr_root = super_zarr_open(str(analyzer_folder), mode="r", storage_options=storage_options) - rec_field = zarr_root.get("recording") - if rec_field is not None: - recording_dict = rec_field[0] - else: - recording_dict = {} + recording_dict = zarr_root.attrs.get("recording") + if recording_dict is None: + raise ValueError(f"Cannot find `recording` attribute in {analyzer_folder}.") preprocessing_list = _make_pipeline_list_from_recording_dict(recording_dict) diff --git a/src/spikeinterface/preprocessing/tests/test_pipeline.py b/src/spikeinterface/preprocessing/tests/test_pipeline.py index 7b53b0ab92..1ea89ec749 100644 --- a/src/spikeinterface/preprocessing/tests/test_pipeline.py +++ b/src/spikeinterface/preprocessing/tests/test_pipeline.py @@ -2,7 +2,6 @@ from spikeinterface.core.testing import check_recordings_equal from spikeinterface.core import create_sorting_analyzer -from spikeinterface.core.core_tools import _is_zarr_write_supported from spikeinterface.generation import generate_recording, generate_ground_truth_recording from spikeinterface.preprocessing import ( apply_preprocessing_pipeline, @@ -206,6 +205,7 @@ def test_loading_from_analyzer(create_cache_folder): cache_folder = create_cache_folder recording, sorting = generate_ground_truth_recording() + # Make it JSON-serializable by saving it to a folder and reloading it recording = recording.save(folder=cache_folder / "recording") preprocessing_dict = {"common_reference": {}, "highpass_filter": {"freq_min": 301.0}} @@ -219,13 +219,11 @@ def test_loading_from_analyzer(create_cache_folder): pp_recording_from_binary = apply_preprocessing_pipeline(recording, pp_list_from_binary) check_recordings_equal(pp_recording, pp_recording_from_binary) - # TODO: remove once writing to zarr is supported with zarr>=3 - if _is_zarr_write_supported(): - analyzer_zarr_folder = cache_folder / "zarr_format.zarr" - _ = create_sorting_analyzer(sorting=sorting, recording=pp_recording, format="zarr", folder=analyzer_zarr_folder) - pp_list_from_zarr = get_preprocessing_list_from_analyzer(analyzer_zarr_folder) - pp_recording_from_zarr = apply_preprocessing_pipeline(recording, pp_list_from_zarr) - check_recordings_equal(pp_recording, pp_recording_from_zarr) + analyzer_zarr_folder = cache_folder / "zarr_format.zarr" + _ = create_sorting_analyzer(sorting=sorting, recording=pp_recording, format="zarr", folder=analyzer_zarr_folder) + pp_list_from_zarr = get_preprocessing_list_from_analyzer(analyzer_zarr_folder) + pp_recording_from_zarr = apply_preprocessing_pipeline(recording, pp_list_from_zarr) + check_recordings_equal(pp_recording, pp_recording_from_zarr) def test_pipeline_recording_arg_substitution(create_cache_folder): diff --git a/src/spikeinterface/preprocessing/tests/test_scaling.py b/src/spikeinterface/preprocessing/tests/test_scaling.py index a19d116b16..27f1de8542 100644 --- a/src/spikeinterface/preprocessing/tests/test_scaling.py +++ b/src/spikeinterface/preprocessing/tests/test_scaling.py @@ -1,6 +1,6 @@ import pytest import numpy as np -from spikeinterface.core.testing_tools import generate_recording +from spikeinterface.core.generate import generate_recording from spikeinterface.preprocessing.preprocessing_classes import scale_to_uV, CenterRecording, scale_to_physical_units diff --git a/src/spikeinterface/sorters/internal/lupin.py b/src/spikeinterface/sorters/internal/lupin.py index 6d79bb036d..9cbadce7a5 100644 --- a/src/spikeinterface/sorters/internal/lupin.py +++ b/src/spikeinterface/sorters/internal/lupin.py @@ -11,7 +11,6 @@ ) from spikeinterface.core.job_tools import fix_job_kwargs -from spikeinterface.core.core_tools import _check_zarr_write_is_supported from spikeinterface.preprocessing import bandpass_filter, common_reference, zscore, whiten from spikeinterface.core.base import minimum_spike_dtype @@ -124,11 +123,6 @@ def get_sorter_version(cls): @classmethod def _run_from_folder(cls, sorter_output_folder, params, verbose): - - # TODO: remove once writing to zarr is supported with zarr>=3 - if params["save_array"]: - _check_zarr_write_is_supported() - from spikeinterface.sortingcomponents.tools import get_prototype_and_waveforms_from_recording from spikeinterface.sortingcomponents.matching import find_spikes_from_templates from spikeinterface.sortingcomponents.peak_detection import detect_peaks diff --git a/src/spikeinterface/sorters/internal/tests/test_lupin.py b/src/spikeinterface/sorters/internal/tests/test_lupin.py index c4eb788a6d..f838f9f4eb 100644 --- a/src/spikeinterface/sorters/internal/tests/test_lupin.py +++ b/src/spikeinterface/sorters/internal/tests/test_lupin.py @@ -13,7 +13,7 @@ class LupinSorterCommonTestSuite(SorterCommonTestSuite, unittest.TestCase): SorterClass = LupinSorter # TODO: remove once writing to zarr is supported with zarr>=3 (save_array writes the templates to zarr) - @pytest.mark.requires_zarr_write + def test_with_run(self): super().test_with_run() diff --git a/src/spikeinterface/sorters/internal/tests/test_tridesclous2.py b/src/spikeinterface/sorters/internal/tests/test_tridesclous2.py index cb6ed09ce4..6431532b87 100644 --- a/src/spikeinterface/sorters/internal/tests/test_tridesclous2.py +++ b/src/spikeinterface/sorters/internal/tests/test_tridesclous2.py @@ -13,7 +13,7 @@ class Tridesclous2SorterCommonTestSuite(SorterCommonTestSuite, unittest.TestCase SorterClass = Tridesclous2Sorter # TODO: remove once writing to zarr is supported with zarr>=3 (save_array writes the templates to zarr) - @pytest.mark.requires_zarr_write + def test_with_run(self): super().test_with_run() diff --git a/src/spikeinterface/sorters/internal/tridesclous2.py b/src/spikeinterface/sorters/internal/tridesclous2.py index 55d377d4db..56a506b293 100644 --- a/src/spikeinterface/sorters/internal/tridesclous2.py +++ b/src/spikeinterface/sorters/internal/tridesclous2.py @@ -15,7 +15,6 @@ ) from spikeinterface.core.job_tools import fix_job_kwargs -from spikeinterface.core.core_tools import _check_zarr_write_is_supported from spikeinterface.preprocessing import bandpass_filter, common_reference, zscore, whiten from spikeinterface.core.base import minimum_spike_dtype @@ -118,11 +117,6 @@ def get_sorter_version(cls): @classmethod def _run_from_folder(cls, sorter_output_folder, params, verbose): - - # TODO: remove once writing to zarr is supported with zarr>=3 - if params["save_array"]: - _check_zarr_write_is_supported() - from spikeinterface.sortingcomponents.matching import find_spikes_from_templates from spikeinterface.sortingcomponents.peak_detection import detect_peaks from spikeinterface.sortingcomponents.peak_selection import select_peaks diff --git a/src/spikeinterface/sorters/tests/test_launcher.py b/src/spikeinterface/sorters/tests/test_launcher.py index 35db70e677..6619431866 100644 --- a/src/spikeinterface/sorters/tests/test_launcher.py +++ b/src/spikeinterface/sorters/tests/test_launcher.py @@ -52,7 +52,6 @@ def job_list(create_cache_folder): return get_job_list(folder) -@pytest.mark.requires_zarr_write def test_run_sorter_jobs_loop(job_list): sortings = run_sorter_jobs(job_list, engine="loop", return_output=True) print(sortings) @@ -202,7 +201,6 @@ def test_run_sorter_jobs_slurm_kwargs(mocker, tmp_path, job_list): assert str(tmp_script_folder) in mock_subprocess_run.call_args_list[-1].args[0][5] -@pytest.mark.requires_zarr_write def test_run_sorter_by_property(create_cache_folder): cache_folder = create_cache_folder working_folder1 = cache_folder / "test_run_sorter_by_property_1" diff --git a/src/spikeinterface/sorters/tests/test_runsorter.py b/src/spikeinterface/sorters/tests/test_runsorter.py index 1316430e55..c07cd94c9b 100644 --- a/src/spikeinterface/sorters/tests/test_runsorter.py +++ b/src/spikeinterface/sorters/tests/test_runsorter.py @@ -24,7 +24,6 @@ def generate_recording(create_cache_folder): return _generate_recording(create_cache_folder) -@pytest.mark.requires_zarr_write @pytest.mark.xfail( platform.system() == "Windows" and parse(platform.python_version()) > parse("3.12"), reason="3rd parth threadpoolctl issue: OSError('GetModuleFileNameEx failed')",