Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
31 commits
Select commit Hold shift + click to select a range
3385cf2
wip: lazy load (analyzer + extensions)
alejoe91 Jun 15, 2026
121a055
todo
alejoe91 Jun 19, 2026
6fdd194
feat: move lazy logic to load only
alejoe91 Jun 24, 2026
1b82c87
Merge branch 'main' into lazy-extensions
alejoe91 Jun 24, 2026
10dfd9f
test: add tests for lazy mode
alejoe91 Jun 24, 2026
1a1e197
Merge branch 'lazy-extensions' of github.com:alejoe91/spikeinterface …
alejoe91 Jun 24, 2026
21a6a59
feat: implement ZarrSpikeVector - a memmap like lazy spike vector fo…
alejoe91 Jun 24, 2026
e1a1e75
fix: add lazy spike vector as kwarg
alejoe91 Jun 24, 2026
6187b31
Merge branch 'main' into lazy-extensions
chrishalcrow Jun 29, 2026
8e1bcf7
Merge branch 'main' into lazy-extensions
alejoe91 Jul 2, 2026
e44a9bd
fix: conflicts
alejoe91 Jul 17, 2026
66cd624
Merge branch 'lazy-extensions' of github.com:alejoe91/spikeinterface …
alejoe91 Jul 17, 2026
b710917
fix: don't copy extension data if sorting_analyzer is lazy
alejoe91 Jul 21, 2026
6fa122a
fix: save=False in compute if analyzer is lazy
alejoe91 Jul 21, 2026
09ae857
feat: add gather to zarr in node pipeline
alejoe91 Jul 21, 2026
080b17e
fix: only delete extension folders if not lazy
alejoe91 Jul 21, 2026
1b6a82c
feat: GatherToZarr and save extension data in chunks to disk
alejoe91 Jul 21, 2026
ae1b5e4
fix: solve conflicts
alejoe91 Jul 21, 2026
d7e025b
fix: conflicts
alejoe91 Jul 21, 2026
60769a5
fix: AnalyzerExtension __del__ closes all memmaps refs
alejoe91 Jul 23, 2026
b147765
feat: extract-waveforms directly to zarr datasets
alejoe91 Jul 23, 2026
84ea3ab
fix: conflicts
alejoe91 Jul 23, 2026
03b617c
Merge branch 'gather-to-zarr' into extract-waveforms-to-zarr
alejoe91 Jul 23, 2026
4ad9dab
Merge branch 'main' of github.com:SpikeInterface/spikeinterface into …
alejoe91 Jul 23, 2026
8ddf404
fix: delete extensions
alejoe91 Jul 23, 2026
d5813cf
Merge branch 'gather-to-zarr' into extract-waveforms-to-zarr
alejoe91 Jul 28, 2026
9e88a67
fix: propagate zarr backend to ComputeWaveforms._run
alejoe91 Jul 28, 2026
566a496
fix: conflicts
alejoe91 Sep 22, 2026
22c9b05
fix: use materialize array
alejoe91 Sep 22, 2026
1f3629c
fix: expose zarr_target_chunk_bytes
alejoe91 Sep 22, 2026
3cbddbf
lint
alejoe91 Sep 22, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions src/spikeinterface/core/analyzer_extension_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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

Expand Down
1 change: 0 additions & 1 deletion src/spikeinterface/core/sortinganalyzer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)):
Expand Down
55 changes: 55 additions & 0 deletions src/spikeinterface/core/tests/test_analyzer_extension_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
54 changes: 54 additions & 0 deletions src/spikeinterface/core/tests/test_waveform_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
Loading
Loading