-
Notifications
You must be signed in to change notification settings - Fork 280
Optimize SpikeRetriever and RAM usage for extensions
#4809
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
feeccf0
1871f5b
1f1bdd2
ea80b95
90043c5
4e8e9f9
2c19568
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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,40 @@ 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.spike_sample_indices = np.asarray(self.spike_vector["sample_index"]) | ||
| self.sorting = sorting | ||
| self._peaks = None | ||
|
|
||
| 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)) | ||
| # For mono-segment, we avoid an extra slice, otherwise make them a tuple for slicing | ||
| 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] | ||
|
|
||
| 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): | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Who uses this? I could not find a caller in the repo. Accessing it brings back the full allocation that this PR removes. It also calls Should we remove it? |
||
| 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 +261,34 @@ 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.segment_slices is not None: | ||
| sl = self.segment_slices[segment_index] | ||
| sample_indices_in_segment = self.spike_sample_indices[sl] | ||
| else: | ||
| sample_indices_in_segment = self.spike_sample_indices | ||
| if self.include_spikes_in_margin: | ||
| i0, i1 = np.searchsorted( | ||
| peaks_in_segment["sample_index"], [start_frame - max_margin, end_frame + max_margin] | ||
| ) | ||
| i0, i1 = np.searchsorted(sample_indices_in_segment, [start_frame - max_margin, end_frame + max_margin]) | ||
| else: | ||
| i0, i1 = np.searchsorted(peaks_in_segment["sample_index"], [start_frame, end_frame]) | ||
| i0, i1 = np.searchsorted(sample_indices_in_segment, [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] | ||
| 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 +391,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 +523,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) | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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")], | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I am wary of the float64 to float32 change. |
||
| "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()) | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Can you read
self.spike_vector["sample_index"]directly inget_peak_sliceinstead of storing it.np.searchsortedworks on the field view without a copy. Plus, when the node is pickled for spawn workers the view is pickled as its own array so every worker gets the sample indices twice