diff --git a/src/spikeinterface/core/basesorting.py b/src/spikeinterface/core/basesorting.py index 36a53e51dc..fdd15e08e6 100644 --- a/src/spikeinterface/core/basesorting.py +++ b/src/spikeinterface/core/basesorting.py @@ -149,6 +149,44 @@ def get_total_duration(self) -> float: ), "This methods requires an associated recording. Call self.register_recording() first." return self._recording.get_total_duration() + def search_cached_spikes_sorted( + self, + indices: list[int], + segment_index: int | None = None, + ): + """ + Search sample indices (frames) in the cached spike vector of one segment. + + Equivalent to `np.searchsorted(segment_sample_index, indices, side="left")`. + Sortings with a lazy spike vector override this to avoid materialising it. + + Parameters + ---------- + indices : list[int] + The sample indices (frames) to search. + segment_index : int | None, default: None + The segment to search. Can be None for mono-segment sortings. + + Returns + ------- + positions : np.ndarray + The insertion positions, relative to the start of the segment in the spike vector. + """ + if self._cached_spike_vector is None: + self._compute_and_cache_spike_vector() + spikes = self._cached_spike_vector + if not isinstance(spikes, np.ndarray): # np.memmap is an ndarray + # np.searchsorted would materialise a lazy vector on every call + raise TypeError( + f"{type(self).__name__} holds a lazy spike vector ({type(spikes).__name__}) " + "and must override search_cached_spikes_sorted()" + ) + if segment_index is None: + assert self.get_num_segments() == 1, "segment_index is required for multi-segment sortings" + segment_index = 0 + start, stop = self._get_spike_vector_segment_slices()[segment_index] + return np.searchsorted(spikes["sample_index"][start:stop], indices) + def get_unit_spike_train( self, unit_id: str | int, @@ -1083,11 +1121,16 @@ def _get_spike_vector_segment_slices(self): if self._cached_spike_vector_segment_slices is None: # compute the, this is needed when spikevector is loaded from format and not computed num_seg = self.get_num_segments() - slices = np.searchsorted(self._cached_spike_vector["segment_index"], np.arange(num_seg + 1)) - self._cached_spike_vector_segment_slices = np.zeros((num_seg, 2), dtype="int64") - for seg_index in range(num_seg): - self._cached_spike_vector_segment_slices[seg_index, 0] = slices[seg_index] - self._cached_spike_vector_segment_slices[seg_index, 1] = slices[seg_index + 1] + if num_seg == 1: + self._cached_spike_vector_segment_slices = np.array( + [[0, self._cached_spike_vector.size]], dtype="int64" + ) + else: + slices = np.searchsorted(self._cached_spike_vector["segment_index"], np.arange(num_seg + 1)) + self._cached_spike_vector_segment_slices = np.zeros((num_seg, 2), dtype="int64") + for seg_index in range(num_seg): + self._cached_spike_vector_segment_slices[seg_index, 0] = slices[seg_index] + self._cached_spike_vector_segment_slices[seg_index, 1] = slices[seg_index + 1] return self._cached_spike_vector_segment_slices def to_reordered_spike_vector( diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index 6ecf66918d..7978ada949 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -51,7 +51,11 @@ def __init__( parents = [parents] self.parents = parents - self._kwargs = dict() + self._kwargs = dict( + time_series=time_series, + return_output=return_output, + parents=parents, + ) def get_margin(self): # can optionally be overwritten @@ -110,10 +114,15 @@ def __init__(self, recording, peaks): self.peaks = peaks # precompute segment slice - self.segment_slices = [] - for segment_index in range(recording.get_num_segments()): - i0, i1 = np.searchsorted(peaks["segment_index"], [segment_index, segment_index + 1]) - self.segment_slices.append(slice(i0, i1)) + if recording.get_num_segments() > 1: + self.segment_slices = [] + for segment_index in range(recording.get_num_segments()): + i0, i1 = np.searchsorted(peaks["segment_index"], [segment_index, segment_index + 1]) + self.segment_slices.append(slice(i0, i1)) + else: + self.segment_slices = None + + self._kwargs.update(dict(peaks=peaks)) def get_margin(self): return 0 @@ -122,15 +131,22 @@ def get_dtype(self): return base_peak_dtype def get_peak_slice(self, segment_index, start_frame, end_frame, max_margin): - sl = self.segment_slices[segment_index] - peaks_in_segment = self.peaks[sl] + if self.segment_slices is not None: + sl = self.segment_slices[segment_index] + peaks_in_segment = self.peaks[sl] + else: + peaks_in_segment = self.peaks i0, i1 = np.searchsorted(peaks_in_segment["sample_index"], [start_frame, end_frame]) return i0, i1 def compute(self, traces, start_frame, end_frame, segment_index, max_margin, peak_slice): # get local peaks - sl = self.segment_slices[segment_index] - peaks_in_segment = self.peaks[sl] + if self.segment_slices is not None: + sl = self.segment_slices[segment_index] + peaks_in_segment = self.peaks[sl] + else: + peaks_in_segment = self.peaks + # i0, i1 = np.searchsorted(peaks_in_segment["sample_index"], [start_frame, end_frame]) i0, i1 = peak_slice local_peaks = peaks_in_segment[i0:i1] @@ -195,7 +211,6 @@ def __init__( category=FutureWarning, stacklevel=2, ) - self._dtype = spike_peak_dtype self.include_spikes_in_margin = include_spikes_in_margin @@ -204,19 +219,43 @@ def __init__( main_channel_ids = sorting.get_property("main_channel_id") assert main_channel_ids is not None, "SpikeRetriever needs the sorting to have `main_channel_id`s." - main_channel_indices = recording.ids_to_indices(main_channel_ids) - self.peaks = sorting_to_peaks(sorting, main_channel_indices, self._dtype) + self.main_channel_indices = recording.ids_to_indices(main_channel_ids) + self.spike_vector, segment_slices = sorting.to_spike_vector(return_slices=True) + self.sorting = sorting + self._peaks = None + + # get_peak_slice() returns positions relative to the segment start + if sorting.get_num_segments() == 1: + self.segment_slices = None + else: + self.segment_slices = [slice(int(s0), int(s1)) for s0, s1 in segment_slices] + + # build any lazy search index now, in the parent, so forked workers inherit it + # instead of each building its own + sorting.search_cached_spikes_sorted([0], segment_index=0) if not channel_from_template: channel_distance = get_channel_distances(recording) self.neighbours_mask = channel_distance <= radius_um self.peak_sign = peak_sign - # precompute segment slice - self.segment_slices = [] - for segment_index in range(recording.get_num_segments()): - i0, i1 = np.searchsorted(self.peaks["segment_index"], [segment_index, segment_index + 1]) - self.segment_slices.append(slice(i0, i1)) + self._kwargs.update( + dict( + sorting=sorting, + channel_from_template=channel_from_template, + extremum_channel_inds=extremum_channel_inds, + radius_um=radius_um, + peak_sign=peak_sign, + include_spikes_in_margin=include_spikes_in_margin, + ) + ) + + @property + def peaks(self): + if self._peaks is not None: + return self._peaks + self._peaks = sorting_to_peaks(self.sorting, self.main_channel_indices) + return self._peaks def get_margin(self): return 0 @@ -225,26 +264,32 @@ def get_dtype(self): return self._dtype def get_peak_slice(self, segment_index, start_frame, end_frame, max_margin): - sl = self.segment_slices[segment_index] - peaks_in_segment = self.peaks[sl] if self.include_spikes_in_margin: - i0, i1 = np.searchsorted( - peaks_in_segment["sample_index"], [start_frame - max_margin, end_frame + max_margin] - ) + indices = [start_frame - max_margin, end_frame + max_margin] else: - i0, i1 = np.searchsorted(peaks_in_segment["sample_index"], [start_frame, end_frame]) - return i0, i1 + indices = [start_frame, end_frame] + return self.sorting.search_cached_spikes_sorted( + indices=indices, + segment_index=segment_index, + ) def compute(self, traces, start_frame, end_frame, segment_index, max_margin, peak_slice): # get local peaks - sl = self.segment_slices[segment_index] - peaks_in_segment = self.peaks[sl] i0, i1 = peak_slice + if self.segment_slices is not None: + s0 = self.segment_slices[segment_index].start + spikes = self.spike_vector[s0 + i0 : s0 + i1] + else: + spikes = self.spike_vector[i0:i1] - local_peaks = peaks_in_segment[i0:i1] + local_peaks = np.zeros(spikes.size, dtype=self._dtype) + local_peaks["sample_index"] = spikes["sample_index"] + local_peaks["channel_index"] = self.main_channel_indices[spikes["unit_index"]] + local_peaks["amplitude"] = 0.0 + local_peaks["segment_index"] = spikes["segment_index"] + local_peaks["unit_index"] = spikes["unit_index"] # make sample index local to traces - local_peaks = local_peaks.copy() local_peaks["sample_index"] -= start_frame - max_margin # handle flag for margin @@ -347,6 +392,15 @@ def __init__( self.nafter = ms_to_samples(ms_after, sampling_frequency) self.neighbours_mask = None + self._kwargs.update( + dict( + ms_before=ms_before, + ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, + ) + ) + class ExtractDenseWaveforms(WaveformsNode): def __init__( @@ -470,6 +524,13 @@ def __init__( self.neighbours_mask = self.channel_distance <= radius_um self.max_num_chans = np.max(np.sum(self.neighbours_mask, axis=1)) + self._kwargs.update( + dict( + radius_um=radius_um, + sparsity_mask=sparsity_mask, + ) + ) + def get_margin(self): return max(self.nbefore, self.nafter) diff --git a/src/spikeinterface/core/sortinganalyzer.py b/src/spikeinterface/core/sortinganalyzer.py index 1c6604516b..42cfd41792 100644 --- a/src/spikeinterface/core/sortinganalyzer.py +++ b/src/spikeinterface/core/sortinganalyzer.py @@ -1171,22 +1171,13 @@ def load_from_zarr( settings = cls._handle_backward_compatibility_settings_pre_init(settings) # Load sorting (in memory or lazy) + zarr_sorting = ZarrSortingExtractor( + folder, zarr_group="sorting", storage_options=storage_options, lazy_spike_vector=lazy + ) if lazy: - copy_spike_vector = False - lazy_spike_vector = True + sorting = zarr_sorting else: - copy_spike_vector = True - lazy_spike_vector = False - sorting = NumpySorting.from_sorting( - ZarrSortingExtractor( - folder, - zarr_group="sorting", - storage_options=storage_options, - lazy_spike_vector=lazy_spike_vector, - ), - with_metadata=True, - copy_spike_vector=copy_spike_vector, - ) + sorting = NumpySorting.from_sorting(zarr_sorting, with_metadata=True, copy_spike_vector=True) # Load recording (if available) if recording is None: diff --git a/src/spikeinterface/core/tests/test_zarrextractors.py b/src/spikeinterface/core/tests/test_zarrextractors.py index 6278541698..8c3a08da0f 100644 --- a/src/spikeinterface/core/tests/test_zarrextractors.py +++ b/src/spikeinterface/core/tests/test_zarrextractors.py @@ -1,6 +1,7 @@ import pytest from pathlib import Path +import numpy as np import zarr from spikeinterface.core import ( @@ -10,6 +11,7 @@ ) from spikeinterface.core.zarrextractors import ( ZarrRecordingExtractor, + ZarrSampleIndexSearch, ZarrSortingExtractor, add_sorting_to_zarr_group, get_default_zarr_compressor, @@ -78,7 +80,64 @@ def test_ZarrSortingExtractor(tmp_path): sorting = load(sorting.to_dict()) +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 + # the (tiny) chunk boundaries + segments = [np.sort(rng.integers(0, 40, size=101)), np.sort(rng.integers(0, 25, size=58))] + sample_index = np.concatenate(segments) + bounds = np.cumsum([0] + [len(s) for s in segments]) + z = zarr.open(tmp_path / "sample_index.zarr", mode="w", shape=sample_index.shape, chunks=(7,), dtype="int64") + z[:] = sample_index + + values = np.arange(-3, 45) + for chunk_firsts in (sample_index[::7], None): + search = ZarrSampleIndexSearch(z, chunk_firsts) + for start, stop in zip(bounds[:-1], bounds[1:]): + expected = np.searchsorted(sample_index[start:stop], values, side="left") + np.testing.assert_array_equal(search.searchsorted(values, start, stop), expected) + # empty range + np.testing.assert_array_equal(search.searchsorted([5], 10, 10), [0]) + + +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" + ZarrSortingExtractor.write_sorting(sorting, folder) + # re-store sample_index in small chunks, so that the search crosses chunk boundaries + # and segments start in the middle of a chunk + 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) + assert spikes_group["sample_index"].nchunks > 3 + + in_ram = ZarrSortingExtractor(folder) + lazy = ZarrSortingExtractor(folder, lazy_spike_vector=True) + assert type(lazy.to_spike_vector()).__name__ == "ZarrSpikeVector" + + rng = np.random.default_rng(1) + num_samples = int(sorting.to_spike_vector()["sample_index"].max()) + 100 + for segment_index in range(sorting.get_num_segments()): + frames = np.sort(rng.integers(-10, num_samples, size=200)) + expected = in_ram.search_cached_spikes_sorted(frames, segment_index=segment_index) + np.testing.assert_array_equal(lazy.search_cached_spikes_sorted(frames, segment_index=segment_index), expected) + + # stores written before the chunk index existed rebuild it from the data + del zarr.open(folder, mode="a")["spikes/sample_index_chunk_firsts"] + old = ZarrSortingExtractor(folder, lazy_spike_vector=True) + frames = np.arange(-5, num_samples, 37) + for segment_index in range(sorting.get_num_segments()): + np.testing.assert_array_equal( + old.search_cached_spikes_sorted(frames, segment_index=segment_index), + in_ram.search_cached_spikes_sorted(frames, segment_index=segment_index), + ) + + if __name__ == "__main__": tmp_path = Path("tmp") test_zarr_compression_options(tmp_path) test_ZarrSortingExtractor(tmp_path) + test_ZarrSampleIndexSearch(tmp_path) + test_ZarrSortingExtractor_lazy_search(tmp_path) diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index 06f887871b..8a5a9b2d8f 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -457,6 +457,74 @@ def copy(self): return np.copy(np.asarray(self)) +class ZarrSampleIndexSearch: + """ + `np.searchsorted` on a zarr `sample_index` array without materialising it. + + The array is sorted within each segment and stored in fixed-size zarr chunks. The + search keeps the first value of every chunk (8 B per chunk), picks the one chunk that + can hold the answer, and decodes only that chunk. The last decoded chunk is cached, so + consecutive searches in the same region decode nothing. + + Parameters + ---------- + sample_index : zarr.Array + The 1D `sample_index` array of a spike vector. + chunk_firsts : np.ndarray | None, default: None + `sample_index[::chunk_length]`, as written by `add_sorting_to_zarr_group`. + If None (stores written before this index existed), it is read from the array, + which decodes every chunk once. + """ + + def __init__(self, sample_index, chunk_firsts=None): + self._sample_index = sample_index + self._num_spikes = sample_index.shape[0] + self._chunk_length = sample_index.chunks[0] + if chunk_firsts is None: + starts = np.arange(0, self._num_spikes, self._chunk_length) + chunk_firsts = sample_index.get_orthogonal_selection(starts) if starts.size else [] + self._chunk_firsts = np.asarray(chunk_firsts, dtype="int64") + self._cached_chunk = (-1, None) # (chunk index, decoded values) + + def _get_chunk(self, chunk_index): + # read the (index, values) pair once into a local and replace it as a whole: threads + # sharing this object (pool_engine="thread") then always see a consistent pair, and a + # race can only cost an extra decode + cached = self._cached_chunk + if cached[0] != chunk_index: + start = chunk_index * self._chunk_length + stop = min(start + self._chunk_length, self._num_spikes) + cached = (chunk_index, self._sample_index[start:stop]) + self._cached_chunk = cached + return cached[1] + + def searchsorted(self, values, start, stop): + """ + Equivalent to `np.searchsorted(sample_index[start:stop], values, side="left")`. + + Returns positions relative to `start`. `start:stop` must be a range over which + `sample_index` is sorted, i.e. one segment. + """ + values = np.atleast_1d(np.asarray(values)) + out = np.zeros(values.size, dtype="int64") + if stop <= start: + return out + first_chunk = start // self._chunk_length + last_chunk = (stop - 1) // self._chunk_length + # first values of the chunks after `first_chunk`, all of which lie inside start:stop + later_firsts = self._chunk_firsts[first_chunk + 1 : last_chunk + 1] + for i, value in enumerate(values): + # the answer lies in the last chunk whose first value is < value: side="left" keeps + # runs of equal values that cross a chunk boundary in the earlier chunk + chunk_index = first_chunk + int(np.searchsorted(later_firsts, value, side="left")) + chunk_start = chunk_index * self._chunk_length + lo = max(chunk_start, start) + hi = min(chunk_start + self._chunk_length, stop) + block = self._get_chunk(chunk_index)[lo - chunk_start : hi - chunk_start] + out[i] = lo + int(np.searchsorted(block, value, side="left")) - start + return out + + class ZarrSortingExtractor(BaseSorting): """ SortingExtractor for a zarr format @@ -527,7 +595,9 @@ def __init__( # 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"]))] - + 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 @@ -556,6 +626,28 @@ def __init__( "lazy_spike_vector": lazy_spike_vector, } + def search_cached_spikes_sorted( + self, + indices: list[int], + segment_index: int | None = None, + ): + if not self._lazy_spike_vector: + return super().search_cached_spikes_sorted( + indices=indices, + segment_index=segment_index, + ) + if segment_index is None: + assert self.get_num_segments() == 1, "segment_index is required for multi-segment sortings" + segment_index = 0 + if self._sample_index_search is None: + # the chunk index is written with the sorting; older stores rebuild it from the data + chunk_firsts = None + if "sample_index_chunk_firsts" in self._spikes_group: + chunk_firsts = self._spikes_group["sample_index_chunk_firsts"][:] + self._sample_index_search = ZarrSampleIndexSearch(self._spikes_group["sample_index"], chunk_firsts) + start, stop = self._cached_spike_vector_segment_slices[segment_index] + return self._sample_index_search.searchsorted(indices, int(start), int(stop)) + @staticmethod def write_sorting( sorting: BaseSorting, @@ -785,6 +877,15 @@ def add_sorting_to_zarr_group( segment_slices.append([i0, i1]) spikes_group.create_dataset(name="segment_slices", data=segment_slices, compressor=None) + # 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( + name="sample_index_chunk_firsts", + data=np.asarray(spikes["sample_index"][::chunk_length], dtype="int64"), + compressor=None, + ) + add_properties_and_annotations(zarr_group, sorting) diff --git a/src/spikeinterface/postprocessing/amplitude_scalings.py b/src/spikeinterface/postprocessing/amplitude_scalings.py index d28d2ab1c0..69a1420638 100644 --- a/src/spikeinterface/postprocessing/amplitude_scalings.py +++ b/src/spikeinterface/postprocessing/amplitude_scalings.py @@ -178,13 +178,12 @@ def __init__( PipelineNode.__init__(self, recording, parents=parents, return_output=return_output) self.return_in_uV = return_in_uV if return_in_uV and recording.has_scaleable_traces(): - self._dtype = np.float32 self._gains = recording.get_channel_gains() self._offsets = recording.get_channel_offsets() else: - self._dtype = recording.get_dtype() self._gains = None self._offsets = None + self._dtype = np.float32 spike_retriever = find_parent_of_type(parents, SpikeRetriever) assert isinstance( spike_retriever, SpikeRetriever @@ -268,7 +267,7 @@ def compute(self, traces, peaks): collisions = {} # compute the scaling for each spike - scalings = np.zeros(len(local_spikes), dtype=float) + scalings = np.zeros(len(local_spikes), dtype=self._dtype) spike_collision_mask = np.zeros(len(local_spikes), dtype=bool) for spike_index, spike in enumerate(local_spikes): diff --git a/src/spikeinterface/postprocessing/unit_locations.py b/src/spikeinterface/postprocessing/unit_locations.py index 4d913ab4f9..7cd94d7ee9 100644 --- a/src/spikeinterface/postprocessing/unit_locations.py +++ b/src/spikeinterface/postprocessing/unit_locations.py @@ -7,10 +7,10 @@ # this dict is for peak location dtype_localize_by_method = { - "center_of_mass": [("x", "float64"), ("y", "float64")], - "grid_convolution": [("x", "float64"), ("y", "float64"), ("z", "float64")], - "peak_channel": [("x", "float64"), ("y", "float64")], - "monopolar_triangulation": [("x", "float64"), ("y", "float64"), ("z", "float64"), ("alpha", "float64")], + "center_of_mass": [("x", "float32"), ("y", "float32")], + "grid_convolution": [("x", "float32"), ("y", "float32"), ("z", "float32")], + "peak_channel": [("x", "float32"), ("y", "float32")], + "monopolar_triangulation": [("x", "float32"), ("y", "float32"), ("z", "float32"), ("alpha", "float32")], } possible_localization_methods = list(dtype_localize_by_method.keys()) diff --git a/src/spikeinterface/preprocessing/detect_artifacts.py b/src/spikeinterface/preprocessing/detect_artifacts.py index c29199ae3b..ab19c49edc 100644 --- a/src/spikeinterface/preprocessing/detect_artifacts.py +++ b/src/spikeinterface/preprocessing/detect_artifacts.py @@ -126,6 +126,15 @@ def __init__( else: self.diff_threshold_unscaled = None + self._kwargs.update( + dict( + saturation_threshold_uV=saturation_threshold_uV, + diff_threshold_uV=diff_threshold_uV, + proportion=proportion, + signed=signed, + ) + ) + def get_margin(self) -> int: """Return the number of margin samples required on each side of a chunk.""" return 0 diff --git a/src/spikeinterface/sortingcomponents/waveforms/features_from_peaks.py b/src/spikeinterface/sortingcomponents/waveforms/features_from_peaks.py index 1695f2cc59..0c7ffe7447 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/features_from_peaks.py +++ b/src/spikeinterface/sortingcomponents/waveforms/features_from_peaks.py @@ -91,8 +91,8 @@ def __init__( self.all_channels = all_channels self.peak_sign = peak_sign - self._kwargs.update(dict(all_channels=all_channels, peak_sign=peak_sign)) self._dtype = recording.get_dtype() + self._kwargs.update(dict(all_channels=all_channels, peak_sign=peak_sign)) def get_dtype(self): return self._dtype @@ -125,8 +125,8 @@ def __init__( self.channel_distance = get_channel_distances(recording) self.neighbours_mask = self.channel_distance <= radius_um self.all_channels = all_channels - self._kwargs.update(dict(radius_um=radius_um, all_channels=all_channels)) self._dtype = recording.get_dtype() + self._kwargs.update(dict(radius_um=radius_um, all_channels=all_channels)) def get_dtype(self): return self._dtype @@ -168,6 +168,7 @@ def __init__( self.radius_um = radius_um self.sparse = sparse self.noise_threshold = noise_threshold + self._dtype = recording.get_dtype() self._kwargs.update( dict( projections=projections, @@ -177,7 +178,6 @@ def __init__( feature=feature, ) ) - self._dtype = recording.get_dtype() def get_dtype(self): return self._dtype