diff --git a/src/spikeinterface/core/analyzer_extension_core.py b/src/spikeinterface/core/analyzer_extension_core.py index c2261fb390..a782ff563e 100644 --- a/src/spikeinterface/core/analyzer_extension_core.py +++ b/src/spikeinterface/core/analyzer_extension_core.py @@ -192,6 +192,10 @@ def _run(self, verbose=False, **job_kwargs): file_path = self._get_binary_extension_folder() / "waveforms.npy" mode = "memmap" copy = False + elif self.format == "zarr": + file_path = self._get_binary_extension_folder() / "waveforms" + mode = "zarr" + copy = False else: file_path = None mode = "shared_memory" @@ -218,6 +222,8 @@ def _run(self, verbose=False, **job_kwargs): verbose=verbose, **job_kwargs, ) + if not self.sorting_analyzer._lazy and not isinstance(all_waveforms, np.ndarray): + all_waveforms = materialize_array(all_waveforms) self.data["waveforms"] = all_waveforms diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 3e80b9fab4..3640a330ff 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -2415,7 +2415,6 @@ def compute_one_extension(self, extension_name, save=True, verbose=False, **kwar else: extension_instance.run(save=save, verbose=verbose) if not self._lazy: - for variable_name in extension_instance.data.keys(): # Materialize the data if not in lazy mode if isinstance(extension_instance.data[variable_name], (np.memmap, zarr.Array)): diff --git a/src/spikeinterface/core/tests/test_analyzer_extension_core.py b/src/spikeinterface/core/tests/test_analyzer_extension_core.py index f606ca34ac..d96ccf0e37 100644 --- a/src/spikeinterface/core/tests/test_analyzer_extension_core.py +++ b/src/spikeinterface/core/tests/test_analyzer_extension_core.py @@ -115,6 +115,61 @@ def test_ComputeWaveforms(format, sparse, create_cache_folder): _check_result_extension(sorting_analyzer, "waveforms", cache_folder) +@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 + # data and fixed seeds must produce identical results, and the saved on-disk files (the npy for + # binary_folder and the zarr dataset) must match the in-memory computation. + import zarr + + cache_folder = create_cache_folder + recording, sorting = generate_ground_truth_recording( + durations=[15.0], + sampling_frequency=16000.0, + num_channels=8, + num_units=5, + seed=2406, + ) + job_kwargs = dict(n_jobs=2, chunk_duration="1s", progress_bar=False) + + analyzers = {} + for fmt in ["memory", "binary_folder", "zarr"]: + if fmt == "memory": + folder = None + elif fmt == "binary_folder": + folder = cache_folder / f"consistency_across_formats_{sparse}_binary" + else: + folder = cache_folder / f"consistency_across_formats_{sparse}.zarr" + if folder is not None and folder.exists(): + shutil.rmtree(folder) + + sorting_analyzer = create_sorting_analyzer( + sorting, recording, format=fmt, folder=folder, sparse=sparse, sparsity=None + ) + sorting_analyzer.compute("random_spikes", max_spikes_per_unit=50, seed=2205) + sorting_analyzer.compute("waveforms", **job_kwargs) + sorting_analyzer.compute("templates", operators=["average", "std"]) + analyzers[fmt] = sorting_analyzer + + # waveforms and templates must be identical across formats + wfs_mem = np.asarray(analyzers["memory"].get_extension("waveforms").get_data()) + for fmt in ["binary_folder", "zarr"]: + wfs = np.asarray(analyzers[fmt].get_extension("waveforms").get_data()) + assert np.array_equal(wfs, wfs_mem), f"waveforms differ for format {fmt}" + for operator in ["average", "std"]: + template_mem = analyzers["memory"].get_extension("templates").get_templates(operator=operator) + template_fmt = analyzers[fmt].get_extension("templates").get_templates(operator=operator) + assert np.allclose(template_fmt, template_mem), f"templates '{operator}' differ for format {fmt}" + + # the files saved on disk must match the in-memory waveforms + binary_waveforms = np.load(analyzers["binary_folder"].folder / "extensions" / "waveforms" / "waveforms.npy") + assert np.array_equal(binary_waveforms, wfs_mem) + + zarr_root = zarr.open(str(analyzers["zarr"].folder), mode="r") + zarr_waveforms = zarr_root["extensions"]["waveforms"]["waveforms"][:] + assert np.array_equal(zarr_waveforms, wfs_mem) + + @pytest.mark.parametrize("format", ["memory", "binary_folder", "zarr"]) @pytest.mark.parametrize("sparse", [True, False]) def test_ComputeTemplates(format, sparse, create_cache_folder): diff --git a/src/spikeinterface/core/tests/test_waveform_tools.py b/src/spikeinterface/core/tests/test_waveform_tools.py index 5e0350f833..1b52f63172 100644 --- a/src/spikeinterface/core/tests/test_waveform_tools.py +++ b/src/spikeinterface/core/tests/test_waveform_tools.py @@ -158,6 +158,60 @@ def test_waveform_tools(create_cache_folder): _check_all_wf_equal(list_wfs_sparse) +@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 + # the main process writes it (single writer), so parallel writes must match a reference and + # never race, even with n_jobs > 1. + import zarr + + recording, sorting = get_dataset() + sampling_frequency = recording.sampling_frequency + nbefore = ms_to_samples(3.0, sampling_frequency) + nafter = ms_to_samples(4.0, sampling_frequency) + spikes = sorting.to_spike_vector() + unit_ids = sorting.unit_ids + dtype = recording.get_dtype() + + if sparse: + sparsity_mask = np.random.RandomState(0).randint( + 0, 2, size=(unit_ids.size, recording.channel_ids.size), dtype="bool" + ) + else: + sparsity_mask = None + + common = dict(return_in_uV=False, dtype=dtype, sparsity_mask=sparsity_mask, copy=True, progress_bar=False) + + # reference computed in shared memory, single job + reference = extract_waveforms_to_single_buffer( + recording, spikes, unit_ids, nbefore, nafter, mode="shared_memory", n_jobs=1, **common + ) + + # zarr mode, single and parallel jobs, must match the reference and reload from disk + for n_jobs in (1, 2): + dataset_path = tmp_path / f"analyzer_{sparse}_{n_jobs}.zarr" / "extensions" / "waveforms" / "waveforms" + zarr_waveforms = extract_waveforms_to_single_buffer( + recording, + spikes, + unit_ids, + nbefore, + nafter, + mode="zarr", + file_path=dataset_path, + n_jobs=n_jobs, + chunk_duration="0.3s", + **common, + ) + assert isinstance(zarr_waveforms, zarr.Array) + assert zarr_waveforms.shape == reference.shape + assert np.array_equal(reference, zarr_waveforms[:]) + + # reload from disk + store_path = tmp_path / f"analyzer_{sparse}_{n_jobs}.zarr" + reloaded = zarr.open(str(store_path), mode="r")["extensions"]["waveforms"]["waveforms"] + assert np.array_equal(reference, reloaded[:]) + + def test_estimate_templates_with_accumulator(): recording, sorting = get_dataset() diff --git a/src/spikeinterface/core/waveform_tools.py b/src/spikeinterface/core/waveform_tools.py index 9fedc89324..b8f3af342a 100644 --- a/src/spikeinterface/core/waveform_tools.py +++ b/src/spikeinterface/core/waveform_tools.py @@ -420,6 +420,7 @@ def extract_waveforms_to_single_buffer( copy=True, job_name=None, verbose=False, + zarr_target_chunk_bytes: int = 10 * 1024 * 1024, **job_kwargs, ): """ @@ -458,7 +459,10 @@ def extract_waveforms_to_single_buffer( If True and the recording has scaling (gain_to_uV and offset_to_uV properties), traces are scaled to uV file_path: str or path or None, default: None - In case of memmap mode, file to save npy file + In "memmap" mode, the npy file to save the waveforms to. + In "zarr" mode, a path pointing inside a zarr store, e.g. + "my-analyzer.zarr/extensions/waveforms/waveforms" (the ".zarr" part identifies the store and + the dataset is created on the fly). dtype: numpy.dtype, default: None dtype for waveforms buffer sparsity_mask: None or array of bool, default: None @@ -470,6 +474,8 @@ def extract_waveforms_to_single_buffer( need to be referenced as long as all_waveforms will be used otherwise it might produce segmentation faults which are hard to debug. Also when copy=False the SharedMemory will need to be unlink manually if proper cleanup of resources is desired. + zarr_target_chunk_bytes: int, default: 10 * 1024 * 1024 + Target chunk size in bytes for zarr storage. {} @@ -491,8 +497,9 @@ def extract_waveforms_to_single_buffer( if mode == "shared_memory": assert file_path is None - else: + elif mode == "memmap": file_path = Path(file_path) + # for mode == "zarr", file_path is a path pointing inside a zarr store (handled below) num_spikes = spikes.size if sparsity_mask is None: @@ -502,6 +509,7 @@ def extract_waveforms_to_single_buffer( num_chans = int(np.max(np.sum(sparsity_mask, axis=1), initial=0)) # This is a numpy scalar, so we cast to int shape = (int(num_spikes), int(n_samples), int(num_chans)) + zarr_writer = None if mode == "memmap": all_waveforms = np.lib.format.open_memmap(file_path, mode="w+", dtype=dtype, shape=shape) # wf_array_info = str(file_path) @@ -516,6 +524,33 @@ def extract_waveforms_to_single_buffer( shm_name = shm.name # wf_array_info = (shm, shm_name, dtype.str, shape) wf_array_info = dict(shm=shm, shm_name=shm_name, dtype=dtype.str, shape=shape) + elif mode == "zarr": + # Create the zarr dataset up front, then fill it. Because zarr's write unit is a whole + # (compressed) chunk, parallel workers cannot safely write directly (two workers may touch + # the same boundary chunk). Instead workers return their contiguous block and the main + # process writes it (single writer) via the `zarr_writer` gather function below. + import zarr + + from .zarrextractors import get_default_zarr_compressor + from .node_pipeline import _split_zarr_store_path + + assert file_path is not None, "zarr mode requires a `file_path` pointing inside a .zarr store" + store_path, dataset_path = _split_zarr_store_path(file_path) + # chunk along the first (spike) axis so that a chunk is about zarr_target_chunk_bytes + row_nbytes = int(n_samples) * int(num_chans) * dtype.itemsize + chunk0 = max(1, zarr_target_chunk_bytes // max(1, row_nbytes)) + zarr_root = zarr.open(str(store_path), mode="a") + all_waveforms = zarr_root.create_dataset( + name=dataset_path, + shape=shape, + chunks=(chunk0, int(n_samples), int(num_chans)), + dtype=dtype, + fill_value=0, + compressor=get_default_zarr_compressor(), + overwrite=True, + ) + wf_array_info = None + zarr_writer = _ZarrSingleBufferWriter(all_waveforms) else: raise ValueError("allocate_waveforms_buffers bad mode") @@ -523,28 +558,57 @@ def extract_waveforms_to_single_buffer( if num_spikes > 0 and num_chans > 0: # and run - func = _worker_distribute_single_buffer - init_func = _init_worker_distribute_single_buffer - - init_args = ( - recording, - spikes, - wf_array_info, - nbefore, - nafter, - return_in_uV, - mode, - sparsity_mask, - ) - if job_name is None: - job_name = f"extract waveforms {mode} mono buffer" - - processor = TimeSeriesChunkExecutor( - recording, func, init_func, init_args, job_name=job_name, verbose=verbose, **job_kwargs - ) - processor.run() - - if mode == "memmap": + if mode == "zarr": + # workers return their contiguous block, the main process writes it (single writer) + func = _worker_return_single_buffer + init_func = _init_worker_return_single_buffer + init_args = ( + recording, + spikes, + nbefore, + nafter, + return_in_uV, + sparsity_mask, + dtype.str, + int(n_samples), + int(num_chans), + ) + if job_name is None: + job_name = "extract waveforms zarr mono buffer" + processor = TimeSeriesChunkExecutor( + recording, + func, + init_func, + init_args, + gather_func=zarr_writer, + job_name=job_name, + verbose=verbose, + **job_kwargs, + ) + processor.run() + else: + func = _worker_distribute_single_buffer + init_func = _init_worker_distribute_single_buffer + + init_args = ( + recording, + spikes, + wf_array_info, + nbefore, + nafter, + return_in_uV, + mode, + sparsity_mask, + ) + if job_name is None: + job_name = f"extract waveforms {mode} mono buffer" + + processor = TimeSeriesChunkExecutor( + recording, func, init_func, init_args, job_name=job_name, verbose=verbose, **job_kwargs + ) + processor.run() + + if mode in ("memmap", "zarr"): return all_waveforms elif mode == "shared_memory": if copy: @@ -650,6 +714,103 @@ def _worker_distribute_single_buffer(segment_index, start_frame, end_frame, work all_waveforms.flush() +class _ZarrSingleBufferWriter: + """ + Gather function used by `extract_waveforms_to_single_buffer` in "zarr" mode. + + Each worker returns a (start_row, block) tuple for a contiguous range of spikes. This is called + in the main process (single writer) so concurrent writes to the same zarr chunk cannot happen. + """ + + def __init__(self, zarr_array): + self.zarr_array = zarr_array + + def __call__(self, res): + if res is None: + return + start_row, block = res + self.zarr_array[start_row : start_row + block.shape[0]] = block + + +def _init_worker_return_single_buffer( + recording, spikes, nbefore, nafter, return_in_uV, sparsity_mask, dtype, n_samples, num_chans +): + worker_dict = {} + worker_dict["recording"] = recording + worker_dict["spikes"] = spikes + worker_dict["nbefore"] = nbefore + worker_dict["nafter"] = nafter + worker_dict["return_in_uV"] = return_in_uV + worker_dict["sparsity_mask"] = sparsity_mask + worker_dict["dtype"] = np.dtype(dtype) + worker_dict["n_samples"] = n_samples + worker_dict["num_chans"] = num_chans + + # prepare segment slices + segment_slices = [] + for segment_index in range(recording.get_num_segments()): + s0, s1 = np.searchsorted(spikes["segment_index"], [segment_index, segment_index + 1]) + segment_slices.append((s0, s1)) + worker_dict["segment_slices"] = segment_slices + + return worker_dict + + +# used by TimeSeriesChunkExecutor for mode="zarr": build and return a contiguous block of waveforms +# (rather than writing to a shared buffer), so the main process can write it to the zarr array. +def _worker_return_single_buffer(segment_index, start_frame, end_frame, worker_dict): + recording = worker_dict["recording"] + segment_slices = worker_dict["segment_slices"] + spikes = worker_dict["spikes"] + nbefore = worker_dict["nbefore"] + nafter = worker_dict["nafter"] + return_in_uV = worker_dict["return_in_uV"] + sparsity_mask = worker_dict["sparsity_mask"] + dtype = worker_dict["dtype"] + n_samples = worker_dict["n_samples"] + num_chans = worker_dict["num_chans"] + + seg_size = recording.get_num_samples(segment_index=segment_index) + + s0, s1 = segment_slices[segment_index] + in_seg_spikes = spikes[s0:s1] + + # take only spikes in range [start_frame, end_frame]; borders are protected by nbefore/nafter + i0, i1 = np.searchsorted( + in_seg_spikes["sample_index"], [max(start_frame, nbefore), min(end_frame, seg_size - nafter)] + ) + + if i1 <= i0: + return None + + sub_spikes = in_seg_spikes[i0:i1] + start = sub_spikes[0]["sample_index"] - nbefore + end = sub_spikes[-1]["sample_index"] + nafter + + traces = recording.get_traces( + start_frame=start, end_frame=end, segment_index=segment_index, return_in_uV=return_in_uV + ) + + onset = start + nbefore + offset = nbefore + nafter + sample_indices = sub_spikes["sample_index"] - onset + unit_indices = sub_spikes["unit_index"] + + block = np.zeros((i1 - i0, n_samples, num_chans), dtype=dtype) + for local_index, (sample_index, unit_index) in enumerate(zip(sample_indices, unit_indices)): + wf = traces[sample_index : sample_index + offset, :] + if sparsity_mask is None: + block[local_index, :, :] = wf + else: + mask = sparsity_mask[unit_index, :] + wf = wf[:, mask] + block[local_index, :, : wf.shape[1]] = wf + + # spike_indices s0 + [i0, i1) are contiguous -> write as a single slice in the main process + start_row = s0 + i0 + return (start_row, block) + + def split_waveforms_by_units(unit_ids, spikes, all_waveforms, sparsity_mask=None, folder=None): """ Split a single buffer waveforms into waveforms by units (multi buffers or multi files).