diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index ad64323909..a0a415ed18 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -779,3 +779,8 @@ def is_path_remote(path: str | Path) -> bool: def ms_to_samples(ms: float, sampling_frequency: float) -> int: """Convert a duration in milliseconds to the nearest number of samples.""" return round(ms * sampling_frequency / 1000.0) + + +def samples_to_ms(samples: int, sampling_frequency: float) -> float: + """Convert a duration in samples to milliseconds.""" + return samples / sampling_frequency * 1000.0 diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index 556370f1ed..14982cb69b 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -12,13 +12,14 @@ 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 +from spikeinterface.core.core_tools import ms_to_samples, samples_to_ms class PipelineNode: # If False (general case) then compute(traces_chunk, *node_input_args) # If True then compute(traces_chunk, start_frame, end_frame, segment_index, max_margin, *node_input_args) + name = None _compute_has_extended_signature = False def __init__( @@ -297,8 +298,10 @@ class WaveformsNode(PipelineNode): def __init__( self, recording: BaseRecording, - ms_before: float, - ms_after: float, + ms_before: float | None = None, + ms_after: float | None = None, + nbefore: int | None = None, + nafter: int | None = None, parents: list[PipelineNode] | None = None, return_output: bool = False, ): @@ -319,22 +322,42 @@ def __init__( return_output : bool, default: False Whether or not the output of the node is returned by the pipeline """ + if ms_before is None and nbefore is None: + raise ValueError("Either ms_before or nbefore must be provided.") + if ms_after is None and nafter is None: + raise ValueError("Either ms_after or nafter must be provided.") + if ms_before is not None and nbefore is not None: + raise ValueError("Only one of ms_before or nbefore should be provided.") + if ms_after is not None and nafter is not None: + raise ValueError("Only one of ms_after or nafter should be provided.") PipelineNode.__init__(self, recording, parents=parents, return_output=return_output) self.recording = recording - self.ms_before = ms_before - self.ms_after = ms_after - self.nbefore = ms_to_samples(ms_before, recording.get_sampling_frequency()) - self.nafter = ms_to_samples(ms_after, recording.get_sampling_frequency()) + sampling_frequency = recording.sampling_frequency + if nbefore is not None: + self.nbefore = nbefore + self.ms_before = samples_to_ms(nbefore, sampling_frequency) + else: + self.ms_before = ms_before + self.nbefore = ms_to_samples(ms_before, sampling_frequency) + if nafter is not None: + self.nafter = nafter + self.ms_after = samples_to_ms(nafter, sampling_frequency) + else: + self.ms_after = ms_after + self.nafter = ms_to_samples(ms_after, sampling_frequency) self.neighbours_mask = None + self.sparse_waveforms = False class ExtractDenseWaveforms(WaveformsNode): def __init__( self, recording: BaseRecording, - ms_before: float, - ms_after: float, + ms_before: float | None = None, + ms_after: float | None = None, + nbefore: int | None = None, + nafter: int | None = None, parents: list[PipelineNode] | None = None, return_output: bool = False, ): @@ -347,10 +370,14 @@ def __init__( ---------- recording : BaseRecording The recording object. - ms_before : float + ms_before : float | None The number of milliseconds to include before the peak of the spike - ms_after : float + ms_after : float | None The number of milliseconds to include after the peak of the spike + nbefore : int | None, default: None + The number of samples to include before the peak of the spike + nafter : int | None, default: None + The number of samples to include after the peak of the spike parents : list[PipelineNode] | None, default: None Pass parents nodes to perform a previous computation return_output : bool, default: False @@ -364,6 +391,8 @@ def __init__( parents=parents, ms_before=ms_before, ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, return_output=return_output, ) @@ -379,8 +408,10 @@ class ExtractSparseWaveforms(WaveformsNode): def __init__( self, recording: BaseRecording, - ms_before: float, - ms_after: float, + ms_before: float | None = None, + ms_after: float | None = None, + nbefore: int | None = None, + nafter: int | None = None, parents: list[PipelineNode] | None = None, return_output: bool = False, radius_um: float = 100.0, @@ -401,10 +432,14 @@ def __init__( ---------- recording : BaseRecording The recording object - ms_before : float + ms_before : float | None The number of milliseconds to include before the peak of the spike - ms_after : float + ms_after : float | None The number of milliseconds to include after the peak of the spike + nbefore : int | None, default: None + The number of samples to include before the peak of the spike + nafter : int | None, default: None + The number of samples to include after the peak of the spike parents : list[PipelineNode] | None, default: None Pass parents nodes to perform a previous computation return_output : bool, default: False @@ -421,6 +456,8 @@ def __init__( parents=parents, ms_before=ms_before, ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, return_output=return_output, ) @@ -434,6 +471,7 @@ def __init__( self.radius_um = radius_um self.neighbours_mask = self.channel_distance <= radius_um self.max_num_chans = np.max(np.sum(self.neighbours_mask, axis=1)) + self.sparse_waveforms = True def get_margin(self): return max(self.nbefore, self.nafter) diff --git a/src/spikeinterface/preprocessing/motion.py b/src/spikeinterface/preprocessing/motion.py index 62bdbfd7a9..4d3e784c84 100644 --- a/src/spikeinterface/preprocessing/motion.py +++ b/src/spikeinterface/preprocessing/motion.py @@ -12,7 +12,6 @@ from spikeinterface.core import get_noise_levels, fix_job_kwargs from spikeinterface.core.job_tools import _shared_job_kwargs_doc from spikeinterface.core.core_tools import SIJsonEncoder -from spikeinterface.core.job_tools import _shared_job_kwargs_doc from spikeinterface.core import BaseRecording motion_options_preset = { @@ -27,6 +26,7 @@ radius_um=80.0, ), "select_kwargs": dict(), + "denoise_kwargs": dict(), "localize_peaks_kwargs": dict( method="monopolar_triangulation", ), @@ -50,6 +50,7 @@ ), "localize_peaks_kwargs": dict(method="monopolar_triangulation"), "select_kwargs": dict(), + "denoise_kwargs": dict(), "estimate_motion_kwargs": dict(method="medicine"), "interpolate_motion_kwargs": dict(), }, @@ -64,6 +65,7 @@ radius_um=80.0, ), "select_kwargs": dict(), + "denoise_kwargs": dict(), "localize_peaks_kwargs": dict( method="grid_convolution", ), @@ -91,6 +93,7 @@ radius_um=80.0, ), "select_kwargs": dict(), + "denoise_kwargs": dict(), "localize_peaks_kwargs": dict(method="monopolar_triangulation"), "estimate_motion_kwargs": dict(method="decentralized", direction="y", rigid=False), "interpolate_motion_kwargs": dict( @@ -107,6 +110,7 @@ radius_um=80.0, ), "select_kwargs": dict(), + "denoise_kwargs": dict(), "localize_peaks_kwargs": dict(method="grid_convolution"), "estimate_motion_kwargs": dict(method="decentralized", direction="y", rigid=False), "interpolate_motion_kwargs": dict( @@ -124,6 +128,7 @@ radius_um=75.0, ), "select_kwargs": dict(), + "denoise_kwargs": dict(), # "localize_peaks_kwargs": dict(method="grid_convolution"), "localize_peaks_kwargs": dict(method="center_of_mass"), "estimate_motion_kwargs": dict(method="dredge_ap", bin_s=5.0, rigid=True), @@ -142,6 +147,7 @@ radius_um=50, ), "select_kwargs": dict(), + "denoise_kwargs": dict(), "localize_peaks_kwargs": dict( method="grid_convolution", weight_method={"mode": "gaussian_2d", "sigma_list_um": np.linspace(5, 25, 5)}, @@ -163,6 +169,7 @@ "": { "detect_kwargs": {}, "select_kwargs": {}, + "denoise_kwargs": {}, "localize_peaks_kwargs": {}, "estimate_motion_kwargs": {}, "interpolate_motion_kwargs": {}, @@ -177,6 +184,7 @@ def _get_default_motion_params(): params = dict() from spikeinterface.sortingcomponents.peak_detection import detect_peak_methods + from spikeinterface.sortingcomponents.waveforms.denoising import denoising_methods from spikeinterface.sortingcomponents.peak_localization import peak_localization_methods from spikeinterface.sortingcomponents.motion.motion_estimation import estimate_motion_methods, estimate_motion @@ -190,6 +198,15 @@ def _get_default_motion_params(): # no design by subclass params["select_kwargs"] = dict() + params["denoise_kwargs"] = dict() + for method_name, method_class in denoising_methods.items(): + sig = inspect.signature(method_class.__init__) + p = {k: v.default for k, v in sig.parameters.items() if k != "self" and v.default != inspect.Parameter.empty} + p.pop("parents", None) + p.pop("return_output", None) + p.pop("return_tensor", None) + params["denoise_kwargs"][method_name] = p + params["localize_peaks_kwargs"] = dict() for method_name, method_class in peak_localization_methods.items(): sig = inspect.signature(method_class.__init__) @@ -258,15 +275,18 @@ def get_motion_parameters_preset(preset): return params -def _update_motion_kwargs(preset, detect_kwargs, select_kwargs, localize_peaks_kwargs, estimate_motion_kwargs): +def _update_motion_kwargs( + preset, detect_kwargs, select_kwargs, denoise_kwargs, localize_peaks_kwargs, estimate_motion_kwargs +): params = motion_options_preset[preset] detect_kwargs = dict(params["detect_kwargs"], **detect_kwargs) select_kwargs = dict(params["select_kwargs"], **select_kwargs) + denoise_kwargs = dict(params["denoise_kwargs"], **denoise_kwargs) localize_peaks_kwargs = dict(params["localize_peaks_kwargs"], **localize_peaks_kwargs) estimate_motion_kwargs = dict(params["estimate_motion_kwargs"], **estimate_motion_kwargs) - return detect_kwargs, select_kwargs, localize_peaks_kwargs, estimate_motion_kwargs + return detect_kwargs, select_kwargs, denoise_kwargs, localize_peaks_kwargs, estimate_motion_kwargs def _update_interpolation_kwargs(preset, interpolation_kwargs): @@ -290,8 +310,11 @@ def compute_motion( ] = "dredge_fast", detect_kwargs: dict = {}, select_kwargs: dict = {}, + denoise_kwargs: dict = {}, localize_peaks_kwargs: dict = {}, estimate_motion_kwargs: dict = {}, + extract_waveforms_method: Literal["dense", "sparse"] = "dense", + extract_waveforms_kwargs=None, output_motion_info: bool = False, folder: str | Path | None = None, overwrite: bool = False, @@ -304,6 +327,7 @@ def compute_motion( This function has some intermediate steps that can be controlled one by one with parameters: * detect peaks * (optional) sub-sample peaks to speed up the localization + * (optional) denoise waveforms * localize peaks * estimate the motion @@ -315,6 +339,7 @@ def compute_motion( * :py:func:`~spikeinterface.sortingcomponents.peak_detection.detect_peaks` * :py:func:`~spikeinterface.sortingcomponents.peak_selection.select_peaks` + * :py:func:`~spikeinterface.sortingcomponents.waveforms.denoising.denoise_waveforms` * :py:func:`~spikeinterface.sortingcomponents.peak_localization.localize_peaks` * :py:func:`~spikeinterface.sortingcomponents.motion.motion.estimate_motion` @@ -326,20 +351,29 @@ def compute_motion( Returns ======= + motion : Motion + The motion object that contains the estimated motion. motion_info : dict A dictionary containing a motion objects, peaks, peak locations, run_times and the parameters used to compute these. + Only returned if `output_motion_info=True`. """ # local import are important because "sortingcomponents" is not important by default from spikeinterface.sortingcomponents.peak_detection import detect_peaks, detect_peak_methods from spikeinterface.sortingcomponents.peak_selection import select_peaks + from spikeinterface.sortingcomponents.waveforms.denoising import denoising_methods from spikeinterface.sortingcomponents.peak_localization import localize_peaks, peak_localization_methods - from spikeinterface.core.node_pipeline import ExtractDenseWaveforms, run_node_pipeline + from spikeinterface.core.node_pipeline import ( + PeakRetriever, + ExtractSparseWaveforms, + ExtractDenseWaveforms, + run_node_pipeline, + ) from spikeinterface.sortingcomponents.motion.motion_estimation import estimate_motion, estimate_motion_methods # get preset params and update if necessary - detect_kwargs, select_kwargs, localize_peaks_kwargs, estimate_motion_kwargs = _update_motion_kwargs( - preset, detect_kwargs, select_kwargs, localize_peaks_kwargs, estimate_motion_kwargs + detect_kwargs, select_kwargs, denoise_kwargs, localize_peaks_kwargs, estimate_motion_kwargs = _update_motion_kwargs( + preset, detect_kwargs, select_kwargs, denoise_kwargs, localize_peaks_kwargs, estimate_motion_kwargs ) job_kwargs = fix_job_kwargs(job_kwargs) @@ -349,6 +383,7 @@ def compute_motion( preset=preset, detect_kwargs=detect_kwargs, select_kwargs=select_kwargs, + denoise_kwargs=denoise_kwargs, localize_peaks_kwargs=localize_peaks_kwargs, estimate_motion_kwargs=estimate_motion_kwargs, job_kwargs=job_kwargs, @@ -371,6 +406,8 @@ def compute_motion( no_selection_kwargs = len(select_kwargs) == 0 + run_times = dict() + pipeline_run_time_name = "" if no_selection_kwargs: # maybe do this directly in the folder when not None, but might be slow on external storage gather_mode = "memory" @@ -382,35 +419,8 @@ def compute_motion( } if method_class.need_noise_levels: detect_kwargs_without_method["noise_levels"] = noise_levels - - node0 = method_class(recording, **detect_kwargs_without_method) - - node1 = ExtractDenseWaveforms(recording, parents=[node0], ms_before=0.1, ms_after=0.3) - - # node detect + localize - method = localize_peaks_kwargs["method"] - method_class = peak_localization_methods[method] - localize_peaks_kwargs_without_method = { - key: localize_peaks_kwarg for key, localize_peaks_kwarg in localize_peaks_kwargs.items() if key != "method" - } - node2 = method_class( - recording, parents=[node0, node1], return_output=True, **localize_peaks_kwargs_without_method - ) - pipeline_nodes = [node0, node1, node2] - t0 = time.perf_counter() - peaks, peak_locations = run_node_pipeline( - recording, - pipeline_nodes, - job_kwargs, - job_name="detect and localize", - gather_mode=gather_mode, - gather_kwargs=None, - squeeze_output=False, - folder=None, - names=None, - ) - t1 = time.perf_counter() - run_times = dict(detect_and_localize=t1 - t0) + peaks_node = method_class(recording, **detect_kwargs_without_method) + pipeline_run_time_name += "detect-" else: # localization is done after select_peaks() pipeline_nodes = None @@ -420,17 +430,69 @@ def compute_motion( method_kwargs["noise_levels"] = noise_levels peaks = detect_peaks(recording, method_kwargs=method_kwargs, pipeline_nodes=None, job_kwargs=job_kwargs) t1 = time.perf_counter() + run_times["detect"] = t1 - t0 # select some peaks peaks = select_peaks(peaks, **select_kwargs, **job_kwargs) t2 = time.perf_counter() - peak_locations = localize_peaks(recording, peaks, method_kwargs=localize_peaks_kwargs, job_kwargs=job_kwargs) - t3 = time.perf_counter() + run_times["select"] = t2 - t1 + peaks_node = PeakRetriever(recording, peaks) + + pipeline_nodes = [peaks_node] - run_times = dict( - detect_peaks=t1 - t0, - select_peaks=t2 - t1, - localize_peaks=t3 - t2, + if extract_waveforms_kwargs is None: + extract_waveforms_kwargs = {"ms_before": 0.1, "ms_after": 0.3} + if extract_waveforms_method == "sparse": + extract_waveforms_node = ExtractSparseWaveforms(recording, parents=[peaks_node], **extract_waveforms_kwargs) + else: + extract_waveforms_node = ExtractDenseWaveforms(recording, parents=[peaks_node], **extract_waveforms_kwargs) + pipeline_nodes.append(extract_waveforms_node) + + if denoise_kwargs is not None and len(denoise_kwargs) > 0: + denoise_method = denoise_kwargs["method"] + denoise_class = denoising_methods[denoise_method] + denoise_kwargs_without_method = { + key: denoise_kwarg for key, denoise_kwarg in denoise_kwargs.items() if key != "method" + } + denoise_node = denoise_class( + recording, + parents=[peaks_node, extract_waveforms_node], + return_output=False, + **denoise_kwargs_without_method, ) + extract_waveforms_for_localization = denoise_node + pipeline_nodes.append(denoise_node) + pipeline_run_time_name += "denoise-localize" + else: + extract_waveforms_for_localization = extract_waveforms_node + pipeline_run_time_name += "localize" + + # node detect + localize + method = localize_peaks_kwargs["method"] + method_class = peak_localization_methods[method] + localize_peaks_kwargs_without_method = { + key: localize_peaks_kwarg for key, localize_peaks_kwarg in localize_peaks_kwargs.items() if key != "method" + } + localize_node = method_class( + recording, + parents=[peaks_node, extract_waveforms_for_localization], + return_output=True, + **localize_peaks_kwargs_without_method, + ) + pipeline_nodes.append(localize_node) + t0 = time.perf_counter() + peaks, peak_locations = run_node_pipeline( + recording, + pipeline_nodes, + job_kwargs, + job_name=pipeline_run_time_name, + gather_mode=gather_mode, + gather_kwargs=None, + squeeze_output=False, + folder=None, + names=None, + ) + t1 = time.perf_counter() + run_times = {pipeline_run_time_name: t1 - t0} t0 = time.perf_counter() try: @@ -478,6 +540,7 @@ def correct_motion( overwrite: bool = False, detect_kwargs: dict = {}, select_kwargs: dict = {}, + denoise_kwargs: dict = {}, localize_peaks_kwargs: dict = {}, estimate_motion_kwargs: dict = {}, interpolate_motion_kwargs: dict = {}, @@ -489,6 +552,7 @@ def correct_motion( This function has some intermediate steps that can be controlled one by one with parameters: * detect peaks * (optional) sub-sample peaks to speed up the localization + * (optional) denoise waveforms * localize peaks * estimate the motion * create and return a `InterpolateMotionRecording` recording object @@ -509,6 +573,7 @@ def correct_motion( * :py:func:`~spikeinterface.sortingcomponents.peak_detection.detect_peaks` * :py:func:`~spikeinterface.sortingcomponents.peak_selection.select_peaks` + * :py:func:`~spikeinterface.sortingcomponents.waveforms.denoising.denoise_waveforms` * :py:func:`~spikeinterface.sortingcomponents.peak_localization.localize_peaks` * :py:func:`~spikeinterface.sortingcomponents.motion.motion.estimate_motion` * :py:func:`~spikeinterface.sortingcomponents.motion.motion.interpolate_motion` @@ -548,6 +613,7 @@ def correct_motion( overwrite=overwrite, detect_kwargs=detect_kwargs, select_kwargs=select_kwargs, + denoise_kwargs=denoise_kwargs, localize_peaks_kwargs=localize_peaks_kwargs, estimate_motion_kwargs=estimate_motion_kwargs, output_motion_info=True, diff --git a/src/spikeinterface/sortingcomponents/clustering/random_projections.py b/src/spikeinterface/sortingcomponents/clustering/random_projections.py index cd451f4d07..8871f18bfe 100644 --- a/src/spikeinterface/sortingcomponents/clustering/random_projections.py +++ b/src/spikeinterface/sortingcomponents/clustering/random_projections.py @@ -15,7 +15,7 @@ from spikeinterface.core.core_tools import ms_to_samples from spikeinterface.core.waveform_tools import estimate_templates from spikeinterface.sortingcomponents.clustering.merging_tools import merge_peak_labels_from_templates -from spikeinterface.sortingcomponents.waveforms.savgol_denoiser import SavGolDenoiser +from spikeinterface.sortingcomponents.waveforms.denoising.savgol_denoiser import SavGolDenoiser from spikeinterface.sortingcomponents.waveforms.features_from_peaks import RandomProjectionsFeature from spikeinterface.core.template import Templates from spikeinterface.core.node_pipeline import ( diff --git a/src/spikeinterface/sortingcomponents/peak_localization/base.py b/src/spikeinterface/sortingcomponents/peak_localization/base.py index 5e6a16fe4c..b0398c6300 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/base.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/base.py @@ -1,8 +1,14 @@ +import numpy as np from spikeinterface.core.node_pipeline import ( PipelineNode, ) +from spikeinterface.core.node_pipeline import ( + find_parent_of_type, + WaveformsNode, +) from spikeinterface.core import get_channel_distances +from spikeinterface.sortingcomponents import waveforms class LocalizeBase(PipelineNode): @@ -13,8 +19,46 @@ def __init__(self, recording, parents, return_output=True, radius_um=75.0): self.radius_um = radius_um self.contact_locations = recording.get_channel_locations() self.channel_distance = get_channel_distances(recording) + + # Find waveform extractor in the parents + waveform_node = find_parent_of_type(self.parents, WaveformsNode) + if waveform_node is None: + raise TypeError(f"{self.name} should have a single {WaveformsNode.__name__} in its parents") + self.nbefore = waveform_node.nbefore + self.nafter = waveform_node.nafter + self.neighbours_mask = self.channel_distance <= radius_um + self.sparse_waveforms = waveform_node.sparse_waveforms + if self.sparse_waveforms: + # waveforms only exist for channels within the extractor's own sparsity, + # so radius_um can only narrow that neighborhood down, never extend it + self.extraction_neighbours_mask = waveform_node.neighbours_mask + self.neighbours_mask &= self.extraction_neighbours_mask self._kwargs["radius_um"] = radius_um def get_dtype(self): return self._dtype + + def get_sparse_waveform(self, waveform, chan_inds, main_chan): + """Get sparse waveforms from dense waveforms""" + if self.sparse_waveforms: + # sparse waveforms are stored contiguously (zero-padded) following the + # extractor's own sparsity mask for main_chan, so chan_inds (a subset of + # that sparsity) must be mapped to its position among the stored channels + extraction_chan_inds = np.flatnonzero(self.extraction_neighbours_mask[main_chan]) + local_inds = np.searchsorted(extraction_chan_inds, chan_inds).clip(max=len(extraction_chan_inds) - 1) + if not np.array_equal(extraction_chan_inds[local_inds], chan_inds): + raise ValueError( + "Requested channels fall outside the sparsity radius used to extract the waveforms: " + "the localization radius_um (plus, for grid_convolution, margin_um) must not exceed " + "the waveform extraction radius_um." + ) + if waveform.ndim == 2: + return waveform[:, local_inds] + else: + return waveform[:, :, local_inds] + else: + if waveform.ndim == 2: + return waveform[:, chan_inds] + else: + return waveform[:, :, chan_inds] diff --git a/src/spikeinterface/sortingcomponents/peak_localization/center_of_mass.py b/src/spikeinterface/sortingcomponents/peak_localization/center_of_mass.py index e2ab1feeb1..de4ebbd953 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/center_of_mass.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/center_of_mass.py @@ -32,13 +32,6 @@ def __init__(self, recording, parents, return_output=True, radius_um=75.0, featu assert feature in ["ptp", "mean", "energy", "peak_voltage"], f"{feature} is not a valid feature" self.feature = feature - - # Find waveform extractor in the parents - waveform_extractor = find_parent_of_type(self.parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"{self.name} should have a single {WaveformsNode.__name__} in its parents") - - self.nbefore = waveform_extractor.nbefore self._kwargs.update(dict(feature=feature)) def compute(self, traces, peaks, waveforms): @@ -49,7 +42,7 @@ def compute(self, traces, peaks, waveforms): (chan_inds,) = np.nonzero(self.neighbours_mask[main_chan]) local_contact_locations = self.contact_locations[chan_inds, :] - wf = waveforms[idx][:, :, chan_inds] + wf = self.get_sparse_waveform(waveforms[idx], chan_inds, main_chan) if self.feature == "ptp": wf_data = np.ptp(wf, axis=1) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/grid.py b/src/spikeinterface/sortingcomponents/peak_localization/grid.py index c55b6b712e..fa1f741e57 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/grid.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/grid.py @@ -1,13 +1,6 @@ import numpy as np import warnings - -from spikeinterface.core.node_pipeline import ( - find_parent_of_type, - PipelineNode, - WaveformsNode, -) - from .base import LocalizeBase from spikeinterface.postprocessing.unit_locations import dtype_localize_by_method @@ -69,14 +62,7 @@ def __init__( self.peak_sign = peak_sign self.percentile = 100 - percentile assert 0 <= self.percentile <= 100, "Percentile should be in [0, 100]" - contact_locations = recording.get_channel_locations() - # Find waveform extractor in the parents - waveform_extractor = find_parent_of_type(self.parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"{self.name} should have a single {WaveformsNode.__name__} in its parents") - - self.nbefore = waveform_extractor.nbefore - self.nafter = waveform_extractor.nafter + self.weight_method = weight_method fs = self.recording.get_sampling_frequency() @@ -96,7 +82,7 @@ def __init__( self.nearest_template_mask, self.z_factors, ) = get_grid_convolution_templates_and_weights( - contact_locations, + self.contact_locations, self.radius_um, self.upsampling_um, self.margin_um, @@ -131,8 +117,14 @@ def compute(self, traces, peaks, waveforms): num_templates = np.sum(nearest_mask) channel_mask = np.sum(self.weights_sparsity_mask[:, :, nearest_mask], axis=(0, 2)) > 0 + if self.sparse_waveforms: + # channels outside the waveform extractor's own sparsity have no data at all + channel_mask &= self.extraction_neighbours_mask[main_chan] + + wf = self.get_sparse_waveform(waveforms[idx], np.flatnonzero(channel_mask), main_chan) + sub_w = self.weights[:, channel_mask, :][:, :, nearest_mask] - global_products = (waveforms[idx][:, :, channel_mask] * self.prototype).sum(axis=1) + global_products = (wf * self.prototype).sum(axis=1) dot_products = np.zeros((nb_weights, num_spikes, num_templates), dtype=np.float32) for count in range(nb_weights): diff --git a/src/spikeinterface/sortingcomponents/peak_localization/main.py b/src/spikeinterface/sortingcomponents/peak_localization/main.py index 71fb810eda..6edfd7a557 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/main.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/main.py @@ -12,6 +12,7 @@ PeakRetriever, SpikeRetriever, ExtractDenseWaveforms, + ExtractSparseWaveforms, ) @@ -24,6 +25,10 @@ def get_localization_pipeline_nodes( method_kwargs=None, ms_before=0.5, ms_after=0.5, + nbefore=None, + nafter=None, + waveform_method="dense", + waveform_kwargs=None, job_kwargs=None, ): @@ -34,10 +39,18 @@ def get_localization_pipeline_nodes( assert method_kwargs is not None # peak_retriever = PeakRetriever(recording, peaks) - - extract_dense_waveforms = ExtractDenseWaveforms( - recording, parents=[peak_source], ms_before=ms_before, ms_after=ms_after, return_output=False - ) + if waveform_kwargs is None: + waveform_kwargs = dict() + waveform_kwargs.update(dict(ms_before=ms_before, ms_after=ms_after, nbefore=nbefore, nafter=nafter)) + if waveform_method == "dense": + extract_waveforms = ExtractDenseWaveforms( + recording, parents=[peak_source], return_output=False, **waveform_kwargs + ) + else: + extract_waveforms = ExtractSparseWaveforms( + recording, parents=[peak_source], return_output=False, **waveform_kwargs + ) + print("Sparsity radius for waveform extraction:", extract_waveforms.radius_um) method_class = peak_localization_methods[method] @@ -55,9 +68,9 @@ def get_localization_pipeline_nodes( recording, peaks=peak_source.peaks, ms_before=ms_before, ms_after=ms_after, job_kwargs=job_kwargs ) - localization_nodes = method_class(recording, parents=[peak_source, extract_dense_waveforms], **method_kwargs) + localization_nodes = method_class(recording, parents=[peak_source, extract_waveforms], **method_kwargs) - pipeline_nodes = [peak_source, extract_dense_waveforms, localization_nodes] + pipeline_nodes = [peak_source, extract_waveforms, localization_nodes] return pipeline_nodes @@ -69,6 +82,10 @@ def localize_peaks( method_kwargs=None, ms_before=0.5, ms_after=0.5, + nbefore=None, + nafter=None, + waveform_method="dense", + waveform_kwargs=None, pipeline_kwargs=None, verbose=False, job_kwargs=None, @@ -95,6 +112,16 @@ def localize_peaks( The number of milliseconds to include before the peak of the spike ms_after : float The number of milliseconds to include after the peak of the spike + nbefore : int | None + The number of samples to include before the peak of the spike. If None, it is + computed from ms_before and the sampling frequency of the recording. + nafter : int | None + The number of samples to include after the peak of the spike. If None, it is + computed from ms_after and the sampling frequency of the recording. + waveform_method : str, default: "dense" + The method to use for extracting waveforms. Can be "dense" or "sparse + waveform_kwargs : dict + Params specific of the waveform extraction method. pipeline_kwargs : dict Dict transmited to run_node_pipelines to handle fine details like : gather_mode/folder/skip_after_n_peaks/recording_slices @@ -149,6 +176,10 @@ def localize_peaks( method_kwargs=method_kwargs, ms_before=ms_before, ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, + waveform_method=waveform_method, + waveform_kwargs=waveform_kwargs, job_kwargs=job_kwargs, ) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py b/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py index 942074ef31..e52fbd4a36 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py @@ -60,11 +60,6 @@ def __init__( self.optimizer = optimizer self.feature = feature - waveform_extractor = find_parent_of_type(self.parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"{self.name} should have a single {WaveformsNode.__name__} in its parents") - - self.nbefore = waveform_extractor.nbefore if enforce_decrease: self.enforce_decrease_radial_parents = make_radial_order_parents( self.contact_locations, self.neighbours_mask @@ -88,7 +83,7 @@ def compute(self, traces, peaks, waveforms): chan_inds = np.flatnonzero(chan_mask) local_contact_locations = self.contact_locations[chan_inds, :] - wf = waveforms[i, :][:, chan_inds] + wf = self.get_sparse_waveform(waveforms[i, :], chan_inds, peak["channel_index"]) if self.feature == "ptp": wf_data = np.ptp(wf, axis=0) elif self.feature == "energy": diff --git a/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py b/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py index 9db0588f6d..f87bd8657d 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py @@ -7,19 +7,30 @@ from spikeinterface.sortingcomponents.tests.common import make_dataset -def test_localize_peaks(): +def _peaks_and_recording(): recording, _ = make_dataset() - # job_kwargs = dict(n_jobs=2, chunk_size=10000, progress_bar=True) - job_kwargs = dict(n_jobs=1, chunk_size=10000, progress_bar=True) - peaks = detect_peaks( recording, method="locally_exclusive", method_kwargs=dict(peak_sign="neg", detect_threshold=5, exclude_sweep_ms=1.0), - job_kwargs=job_kwargs, + job_kwargs=dict(n_jobs=1, chunk_size=10000, progress_bar=True), ) + return recording, peaks + + +@pytest.fixture +def peaks_and_recording(): + return _peaks_and_recording() + + +def test_localize_peaks(peaks_and_recording): + recording, peaks = peaks_and_recording + + # job_kwargs = dict(n_jobs=2, chunk_size=10000, progress_bar=True) + job_kwargs = dict(n_jobs=1, chunk_size=10000, progress_bar=True) + list_locations = [] peak_locations = localize_peaks(recording, peaks, method="center_of_mass", job_kwargs=job_kwargs) @@ -91,36 +102,85 @@ def test_localize_peaks(): assert peaks.size == peak_locations.shape[0] list_locations.append(("minimize_with_log_penality_v_peak", peak_locations)) - # DEBUG - # import MEArec - # recgen = MEArec.load_recordings(recordings=local_path, return_h5_objects=True, - # check_suffix=False, - # load=['recordings', 'spiketrains', 'channel_positions'], - # load_waveforms=False) - # soma_positions = np.zeros((len(recgen.spiketrains), 3), dtype='float32') - # for i, st in enumerate(recgen.spiketrains): - # soma_positions[i, :] = st.annotations['soma_position'] - # import matplotlib.pyplot as plt - # import spikeinterface.widgets as sw - # from probeinterface.plotting import plot_probe - # for title, peak_locations in list_locations: - # probe = recording.get_probe() - # fig, axs = plt.subplots(ncols=2, sharey=True) - # ax = axs[0] - # ax.set_title(title) - # plot_probe(probe, ax=ax) - # ax.scatter(peak_locations['x'], peak_locations['y'], color='k', s=1, alpha=0.5) - # ax.set_xlabel('x') - # ax.set_ylabel('y') - # #MEArec is "yz" in 2D - # ax.scatter(soma_positions[:, 1], soma_positions[:, 2], color='g', s=20, marker='*') - # ax = axs[1] - # if 'z' in peak_locations.dtype.fields: - # ax.scatter(peak_locations['z'], peak_locations['y'], color='k', s=1, alpha=0.5) - # ax.set_xlabel('z') - # ax.set_title(title) - # plt.show() + +@pytest.mark.parametrize("method", ["center_of_mass", "monopolar_triangulation", "grid_convolution"]) +def test_localize_peaks_sparse(peaks_and_recording, method): + recording, peaks = peaks_and_recording + + job_kwargs = dict(n_jobs=2, chunk_size=10000, progress_bar=True) + + # test sparse waveforms + peak_locations = localize_peaks( + recording, + peaks, + method_kwargs=dict( + method=method, + ), + waveform_method="sparse", # if method != "grid_convolution" else "dense", + job_kwargs=job_kwargs, + ) + assert peaks.size == peak_locations.shape[0] + + +@pytest.mark.parametrize("method", ["center_of_mass", "monopolar_triangulation", "grid_convolution"]) +def test_localize_sparse_narrow(peaks_and_recording, method): + """Test that a smaller sparsity in waveforms than localization is handled""" + recording, peaks = peaks_and_recording + + job_kwargs = dict(n_jobs=2, chunk_size=10000, progress_bar=True) + + # test sparse waveforms + peak_locations = localize_peaks( + recording, + peaks, + method_kwargs=dict( + method=method, + radius_um=150, # larger than waveform radius + ), + waveform_method="sparse", # if method != "grid_convolution" else "dense", + waveform_kwargs=dict(radius_um=50), # smaller than localization radius + job_kwargs=job_kwargs, + ) + assert peaks.size == peak_locations.shape[0] + + +@pytest.mark.parametrize("method", ["center_of_mass", "monopolar_triangulation", "grid_convolution"]) +def test_sparse_and_dense_are_close(peaks_and_recording, method): + recording, peaks = peaks_and_recording + + job_kwargs = dict(n_jobs=2, chunk_size=10000, progress_bar=True) + + # test sparse waveforms + radius_um = 150.0 + peak_locations_sparse = localize_peaks( + recording, + peaks, + method_kwargs=dict( + method=method, + ), + waveform_method="sparse", + waveform_kwargs=dict(radius_um=radius_um), + job_kwargs=job_kwargs, + ) + peak_locations_dense = localize_peaks( + recording, + peaks, + method_kwargs=dict( + method=method, + ), + waveform_method="dense", + job_kwargs=job_kwargs, + ) + # Allow a 2um tolerance for the difference between sparse and dense localization results + np.testing.assert_allclose(peak_locations_sparse["x"], peak_locations_dense["x"], rtol=0.01, atol=1) + np.testing.assert_allclose(peak_locations_sparse["y"], peak_locations_dense["y"], rtol=0.01, atol=1) + if "z" in peak_locations_sparse.dtype.names: + np.testing.assert_allclose(peak_locations_sparse["z"], peak_locations_dense["z"], rtol=0.01, atol=1) if __name__ == "__main__": - test_localize_peaks() + import pytest + + # run the is close test only for center of mass + peaks_and_recording_obj = _peaks_and_recording() + test_sparse_and_dense_are_close(peaks_and_recording_obj, method="center_of_mass") diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/__init__.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/__init__.py new file mode 100644 index 0000000000..8a7ff6f45c --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/__init__.py @@ -0,0 +1,2 @@ +from .method_list import denoising_methods +from .main import denoise_waveforms, get_denoising_pipeline_nodes diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/main.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/main.py new file mode 100644 index 0000000000..e1ac9c7f2d --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/main.py @@ -0,0 +1,151 @@ +import warnings +from typing import Literal +import numpy as np + +from .method_list import denoising_methods + +from spikeinterface.sortingcomponents.tools import make_multi_method_doc +from spikeinterface.core.job_tools import split_job_kwargs, fix_job_kwargs + +from spikeinterface.core.node_pipeline import ( + run_node_pipeline, + PipelineNode, + PeakRetriever, + ExtractDenseWaveforms, + ExtractSparseWaveforms, +) + + +# This method is used both by localize_peaks() and compute_spike_locations() +# message to pierre yger : do not remove this function any more, please +def get_denoising_pipeline_nodes( + recording, + peak_source, + method="center_of_mass", + method_kwargs=None, + ms_before=0.5, + ms_after=0.5, + nbefore=None, + nafter=None, + waveform_kwargs=None, + waveform_method: Literal["dense", "sparse"] = "dense", + job_kwargs=None, +) -> list[PipelineNode]: + assert method in denoising_methods, f"Method {method} is not supported. Choose from {denoising_methods.keys()}" + + assert method_kwargs is not None + + waveform_kwargs = waveform_kwargs or {} + waveform_kwargs.update({"ms_before": ms_before, "ms_after": ms_after, "nbefore": nbefore, "nafter": nafter}) + if waveform_method == "dense": + waveforms_node = ExtractDenseWaveforms(recording, parents=[peak_source], return_output=False, **waveform_kwargs) + else: + waveforms_node = ExtractSparseWaveforms( + recording, parents=[peak_source], return_output=False, **waveform_kwargs + ) + + method_class = denoising_methods[method] + denoising = method_class(recording, parents=[peak_source, waveforms_node], **method_kwargs) + pipeline_nodes = [peak_source, waveforms_node, denoising] + + return pipeline_nodes + + +def denoise_waveforms( + recording, + peaks, + method=None, + method_kwargs=None, + ms_before=0.5, + ms_after=0.5, + nbefore=None, + nafter=None, + waveform_method: Literal["dense", "sparse"] = "dense", + waveform_kwargs=None, + pipeline_kwargs=None, + verbose=False, + job_kwargs=None, +) -> np.ndarray: + """Denoise waveforms using the specified method. + + Parameters + ---------- + recording : RecordingExtractor + The recording extractor object. + peaks : array + Peaks array, as returned by detect_peaks() in "compact_numpy" way. + method : str + The denoising method to use. See `denoising_methods` for available methods. + method_kwargs : dict + Params specific of the method. + ms_before : float + The number of milliseconds to include before the peak of the spike + ms_after : float + The number of milliseconds to include after the peak of the spike + pipeline_kwargs : dict + Dict transmited to run_node_pipelines to handle fine details + like : gather_mode/folder/skip_after_n_peaks/recording_slices + verbose : Bool, default: False + If True, output is verbose + job_kwargs : dict | None, default None + A job kwargs dict. If None or empty dict, then the global one is used. + + {method_doc} + + Returns + ------- + denoised_waveforms : np.ndarray + Denoised waveforms of shape (n_spikes, n_channels, n_samples) + """ + if method_kwargs is None: + method_kwargs = dict() + + if "method" in method_kwargs: + # for flexibility the caller can put method inside method_kwargs + assert method is None + method_kwargs = method_kwargs.copy() + method = method_kwargs.pop("method") + + if method is None: + warnings.warn("localize_peaks() method should be explicitly given, nicely 'center_of_mass' is used") + method = "center_of_mass" + + job_kwargs = fix_job_kwargs(job_kwargs) + + assert method in denoising_methods, f"Method {method} is not supported. Choose from {denoising_methods.keys()}" + + peak_source = PeakRetriever(recording, peaks) + + pipeline_nodes = get_denoising_pipeline_nodes( + recording, + peak_source, + method=method, + method_kwargs=method_kwargs, + ms_before=ms_before, + ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, + waveform_method=waveform_method, + waveform_kwargs=waveform_kwargs, + job_kwargs=job_kwargs, + ) + + if pipeline_kwargs is None: + pipeline_kwargs = dict() + + job_name = f"denoise waveforms ({method})" + peak_locations = run_node_pipeline( + recording, + pipeline_nodes, + job_kwargs, + job_name=job_name, + squeeze_output=True, + verbose=verbose, + **pipeline_kwargs, + ) + + return peak_locations + + +method_doc = make_multi_method_doc(list(denoising_methods.values())) +denoise_waveforms.__doc__ = denoise_waveforms.__doc__.format(method_doc=method_doc) diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/method_list.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/method_list.py new file mode 100644 index 0000000000..2435692ac8 --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/method_list.py @@ -0,0 +1,6 @@ +from .neural_network_denoiser import SingleChannelDenoiser +from .savgol_denoiser import SavGolDenoiser +from .temporal_pca_denoiser import TemporalPCADenoiser + +_methods_list = [SingleChannelDenoiser, SavGolDenoiser, TemporalPCADenoiser] +denoising_methods = {m.name: m for m in _methods_list} diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py new file mode 100644 index 0000000000..c088e9ccdb --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py @@ -0,0 +1,247 @@ +import json +import warnings +import importlib.util +from pathlib import Path +from typing import List, Optional +import numpy as np + +if importlib.util.find_spec("torch") is not None: + import torch + from torch import nn + + HAVE_TORCH = True +else: + HAVE_TORCH = False + +if importlib.util.find_spec("huggingface_hub") is not None: + from huggingface_hub import hf_hub_download, list_repo_files + + HAVE_HUGGINGFACE = True +else: + HAVE_HUGGINGFACE = False + +from spikeinterface.core import BaseRecording +from spikeinterface.core.node_pipeline import PipelineNode +from ..waveform_utils import to_temporal_representation, from_temporal_representation, WaveformTransformer + + +class SingleChannelDenoiser(WaveformTransformer): + """ + Denoiser for temporal dimension of waveforms. It takes as input a WaveformsNode and outputs denoised waveforms. + + Parameters + ---------- + recording: BaseRecording + The recording object. + return_output: bool, default True + Whether to return the output of the node. + parents: list of PipelineNode, optional + The parent nodes of this node. Must include a WaveformsNode. + model_folder: str, optional + Path to a folder containing the model .pt file and a .json file with temporal parameters + repo_id: str, optional + Huggingface repo id to download the model from. Must contain a .pt file and a .json file with temporal parameters + model_name: str, optional + Name of the model to use. If there are multiple .pt files in the model_folder, this specifies which one to use. + """ + + name = "single_channel_denoiser" + params_doc = """ + model_folder: str, optional + Path to a folder containing the model .pt file and a .json file with temporal parameters + repo_id: str, optional + Huggingface repo id to download the model from. Must contain a .pt file and a .json file with temporal parameters + model_name: str, optional + Name of the model to use. If there are multiple .pt files in the model_folder, this specifies which one to use. + """ + + def __init__( + self, + recording: BaseRecording, + return_output: bool = True, + parents: Optional[List[PipelineNode]] = None, + model_folder: Optional[str] = None, + repo_id: Optional[str] = None, + model_name: Optional[str] = None, + device=None, + ): + assert HAVE_TORCH, "To use the SingleChannelDenoiser you need to install torch" + + super().__init__( + recording, + return_output=return_output, + parents=parents, + ) + if model_folder is None and repo_id is None: + raise ValueError("You need to specify either model_folder or repo_id") + if model_folder is not None and repo_id is not None: + raise ValueError("You cannot specify both model_folder and repo_id") + spike_size = self.waveforms_node.nbefore + self.waveforms_node.nafter + # Load model + self.denoiser, model_relative_path = self.load_model( + model_folder=model_folder, repo_id=repo_id, model_name=model_name, spike_size=spike_size, device=device + ) + + self.assert_model_and_waveform_temporal_match( + model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path + ) + + def assert_model_and_waveform_temporal_match( + self, + model_relative_path: str, + model_folder: Optional[str] = None, + repo_id: Optional[str] = None, + ): + """ + Asserts that the model and the waveform extractor have the same temporal parameters + """ + # Extract temporal parameters from the waveform extractor + waveforms_ms_before = self.waveforms_node.ms_before + waveforms_ms_after = self.waveforms_node.ms_after + waveforms_sampling_frequency = self.waveforms_node.recording.sampling_frequency + + json_file_path = None + if model_folder is not None: + json_file_path = Path(model_folder) / str(model_relative_path).replace(".pt", ".json") + else: + try: + filename = str(model_relative_path).replace(".pt", ".json") + json_file_path = hf_hub_download(repo_id=repo_id, filename=filename) + except Exception as e: + warnings.warn(f"Could not download json file from repo {repo_id}. Model might misbehave") + + if json_file_path is None or not Path(json_file_path).exists(): + warnings.warn(f"Could not find json file for model {model_relative_path}. Model might misbehave") + return + + # Load the json file in the json_file_path_variable + with open(json_file_path, "r") as json_file: + model_info = json.load(json_file) + + model_ms_before = model_info.get("ms_before") + model_ms_after = model_info.get("ms_after") + model_sampling_frequency = model_info.get("sampling_frequency") + model_num_samples = model_info.get("num_samples") + model_nbefore = model_info.get("nbefore") + + if model_num_samples is not None: + if model_num_samples != self.waveforms_node.nbefore + self.waveforms_node.nafter: + raise ValueError( + f"Model num_samples {model_num_samples} does not match waveform extractor num_samples {waveform_node.num_samples}" + ) + if model_ms_before is not None: + if abs(model_ms_before - waveforms_ms_before) > 0.1: + raise ValueError( + f"Difference between model ms_before {model_ms_before} and waveform extractor ms_before {waveforms_ms_before} is too large" + ) + if model_ms_after is not None: + if abs(model_ms_after - waveforms_ms_after) > 0.1: + raise ValueError( + f"Difference between model ms_after {model_ms_after} and waveform extractor ms_after {waveforms_ms_after} is too large" + ) + if model_sampling_frequency is not None: + if not np.isclose(model_sampling_frequency, waveforms_sampling_frequency, rtol=1e-3): + raise ValueError( + f"Difference between sampling_frequency {model_sampling_frequency} does not match waveform extractor sampling_frequency {waveforms_sampling_frequency}" + ) + if model_nbefore is not None: + if abs(model_nbefore - self.waveforms_node.nbefore) > 5: + raise ValueError( + f"Difference between model nbefore {model_nbefore} and waveform extractor nbefore {waveform_node.nbefore} is too large" + ) + + def load_model( + self, + model_folder: Optional[str] = None, + repo_id: Optional[str] = None, + model_name: Optional[str] = None, + spike_size: int = 121, + device: Optional[str] = None, + ): + if model_folder is not None: + pt_files = [f for f in Path(model_folder).iterdir("") if f.suffix == ".pt"] + if len(pt_files) == 1: + model_path = pt_files[0] + else: + if model_name is not None: + raise ValueError(f"Multiple models found in {model_folder}. Please specify model_name") + assert ( + model_name is not None + ), "If there are multiple .pt files in the repo, you need to specify the model_name" + filename = [f for f in pt_files if model_name in f] + if len(filename) == 0: + raise ValueError(f"Model {model_name} not found in repo {repo_id}") + elif len(filename) > 1: + raise ValueError(f"Multiple models found for {model_name} in repo {repo_id}: {filename}") + else: + model_path = filename[0] + model_relative_path = model_path.relative_to(model_folder) + else: + assert HAVE_HUGGINGFACE, "To download models from Huggingface you need to install huggingface_hub" + + repo_filenames = list_repo_files(repo_id=repo_id) + + pt_files = [f for f in repo_filenames if f.endswith(".pt")] + if len(pt_files) == 1: + filename = pt_files[0] + else: + assert ( + model_name is not None + ), "If there are multiple .pt files in the repo, you need to specify the model_name" + filename = [f for f in pt_files if model_name in f] + if len(filename) == 0: + raise ValueError(f"Model {model_name} not found in repo {repo_id}") + elif len(filename) > 1: + raise ValueError(f"Multiple models found for {model_name} in repo {repo_id}: {filename}") + else: + filename = filename[0] + model_path = hf_hub_download(repo_id=repo_id, filename=filename) + model_relative_path = filename + + denoiser = SingleChannel1dCNNDenoiser(pretrained_path=model_path, spike_size=spike_size) + denoiser = denoiser.load(device=device) + model_name = Path(model_path).stem + return denoiser, model_relative_path + + def compute(self, traces, peaks, waveforms): + num_channels = waveforms.shape[2] + + # Collapse channels and transform to torch tensor + temporal_waveforms = to_temporal_representation(waveforms) + temporal_waveforms_tensor = torch.from_numpy(temporal_waveforms).float() + + # Denoise + denoised_temporal_waveforms = self.denoiser(temporal_waveforms_tensor).detach().numpy() + + # Reconstruct representation with channels + denoised_waveforms = from_temporal_representation(denoised_temporal_waveforms, num_channels) + + return denoised_waveforms + + +if HAVE_TORCH: + + class SingleChannel1dCNNDenoiser(nn.Module): + def __init__(self, pretrained_path=None, n_filters=[16, 8], filter_sizes=[5, 11], spike_size=121): + super().__init__() + + out_channels_conv1, out_channels_conv_2 = n_filters + kernel_size_conv1, kernel_size_conv2 = filter_sizes + self.conv1 = nn.Sequential(nn.Conv1d(1, out_channels_conv1, kernel_size_conv1), nn.ReLU()) + self.conv2 = nn.Sequential(nn.Conv1d(out_channels_conv1, out_channels_conv_2, kernel_size_conv2), nn.ReLU()) + n_input_feat = out_channels_conv_2 * (spike_size - kernel_size_conv1 - kernel_size_conv2 + 2) + self.out = nn.Linear(n_input_feat, spike_size) + self.pretrained_path = pretrained_path + + def forward(self, x): + x = x[:, None] + x = self.conv1(x) + x = self.conv2(x) + x = x.view(x.shape[0], -1) + x = self.out(x) + return x + + def load(self, device="cpu"): + checkpoint = torch.load(self.pretrained_path, map_location=device) + self.load_state_dict(checkpoint) + return self diff --git a/src/spikeinterface/sortingcomponents/waveforms/savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py similarity index 74% rename from src/spikeinterface/sortingcomponents/waveforms/savgol_denoiser.py rename to src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py index 1ed9e4bffa..3a871ef14e 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/savgol_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py @@ -1,10 +1,11 @@ from typing import List, Optional from spikeinterface.core import BaseRecording -from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type +from spikeinterface.core.node_pipeline import PipelineNode +from ..waveform_utils import WaveformTransformer -class SavGolDenoiser(WaveformsNode): +class SavGolDenoiser(WaveformTransformer): """ Waveform Denoiser based on a simple Savitzky-Golay filtering https://en.wikipedia.org/wiki/Savitzky%E2%80%93Golay_filter @@ -18,9 +19,17 @@ class SavGolDenoiser(WaveformsNode): parents: list of PipelineNodes, default: None The parent nodes of this node order: int, default: 3 - the order of the filter + The order of the filter window_length_ms: float, default: 0.25 - the temporal duration of the filter in ms + tThe temporal duration of the filter in ms + """ + + name = "savgol_denoiser" + params_doc = """ + order: int, default: 3 + The order of the filter + window_length_ms: float, default: 0.25 + The temporal duration of the filter in ms """ def __init__( @@ -31,20 +40,14 @@ def __init__( order: int = 3, window_length_ms: float = 0.25, ): - waveform_extractor = find_parent_of_type(parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"SavGolDenoiser should have a single {WaveformsNode.__name__} in its parents") - super().__init__( recording, - waveform_extractor.ms_before, - waveform_extractor.ms_after, return_output=return_output, parents=parents, ) + waveforms_sampling_frequency = self.recording.get_sampling_frequency() self.order = order - waveforms_sampling_frequency = self.recording.get_sampling_frequency() self.window_length = int(window_length_ms * waveforms_sampling_frequency / 1000) self.order = min(self.order, self.window_length - 1) self._kwargs.update(dict(order=order, window_length_ms=window_length_ms)) diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py new file mode 100644 index 0000000000..5f3e2907db --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py @@ -0,0 +1,85 @@ +import numpy as np + + +from spikeinterface.core import BaseRecording +from spikeinterface.core.node_pipeline import PipelineNode +from ..waveform_utils import WaveformTransformer +from ..temporal_pca import TemporalPCMixin, to_temporal_representation, from_temporal_representation + + +class TemporalPCADenoiser(WaveformTransformer, TemporalPCMixin): + """ + A step that performs a PCA denoising on the waveforms extracted by a peak_detection function. + + This class needs a model_folder_path with a trained model. A model can be trained with the + static method TemporalPCAProjection.fit(). + + Parameters + ---------- + recording : BaseRecording + The recording object + parents: list + The parent nodes of this node. This should contain a mechanism to extract waveforms + pca_model: sklearn model | None + The already fitted sklearn model instead of model_folder_path + model_folder_path : str | Path | None + If pca_model is None, the path to the folder containing the pca model and the training metadata. + return_output: bool, default: True + use false to suppress the output of this node in the pipeline + + """ + + name = "temporal_pca_denoising" + params_doc = """ + pca_model: sklearn model, optional + The already fitted PCA model. + model_folder_path: str | Path, optional + Path to a folder containing the trained PCA model and the training metadata. + """ + + def __init__( + self, + recording: BaseRecording, + parents: list[PipelineNode], + pca_model=None, + model_folder_path=None, + return_output=True, + ): + WaveformTransformer.__init__(self, recording=recording, parents=parents, return_output=return_output) + TemporalPCMixin.__init__( + self, + waveform_node=self.waveforms_node, + pca_model=pca_model, + model_folder_path=model_folder_path, + ) + + def compute(self, traces: np.ndarray, peaks: np.ndarray, waveforms: np.ndarray) -> np.ndarray: + """ + Denoises the waveforms using the PCA model trained in the fit method or loaded from the model_folder_path. + + Parameters + ---------- + traces : np.ndarray + The traces of the recording. + peaks : np.ndarray + The peaks resulting from a peak_detection step. + waveforms : np.ndarray + Waveforms extracted from the recording using a WavefomExtractor node. + + Returns + ------- + np.ndarray + The projected waveforms. + + """ + num_channels = waveforms.shape[2] + + if waveforms.shape[0] > 0: + temporal_waveform = to_temporal_representation(waveforms) + projected_temporal_waveforms = self.pca_model.transform(temporal_waveform) + temporal_denoised_waveforms = self.pca_model.inverse_transform(projected_temporal_waveforms) + denoised_waveforms = from_temporal_representation(temporal_denoised_waveforms, num_channels) + else: + denoised_waveforms = np.zeros_like(waveforms) + + return denoised_waveforms diff --git a/src/spikeinterface/sortingcomponents/waveforms/hanning_filter.py b/src/spikeinterface/sortingcomponents/waveforms/hanning_filter.py index c6d1070e6d..5deedac216 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/hanning_filter.py +++ b/src/spikeinterface/sortingcomponents/waveforms/hanning_filter.py @@ -18,6 +18,9 @@ class HanningFilter(WaveformsNode): The parent nodes of this node """ + name = "hanning_filter" + params_doc = "" + def __init__( self, recording: BaseRecording, diff --git a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py deleted file mode 100644 index 77cada6ed3..0000000000 --- a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py +++ /dev/null @@ -1,136 +0,0 @@ -import json -import importlib.util -from typing import List, Optional - -if importlib.util.find_spec("torch") is not None: - import torch - from torch import nn - - HAVE_TORCH = True -else: - HAVE_TORCH = False - -if importlib.util.find_spec("huggingface_hub") is not None: - from huggingface_hub import hf_hub_download - - HAVE_HUGGINFACE = True -else: - HAVE_HUGGINFACE = False - -from spikeinterface.core import BaseRecording -from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type -from .waveform_utils import to_temporal_representation, from_temporal_representation - - -class SingleChannelToyDenoiser(WaveformsNode): - def __init__( - self, recording: BaseRecording, return_output: bool = True, parents: Optional[List[PipelineNode]] = None - ): - assert HAVE_TORCH, "To use the SingleChannelToyDenoiser you need to install torch" - waveform_extractor = find_parent_of_type(parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"Model should have a {WaveformsNode.__name__} in its parents") - - super().__init__( - recording, - waveform_extractor.ms_before, - waveform_extractor.ms_after, - return_output=return_output, - parents=parents, - ) - - self.assert_model_and_waveform_temporal_match(waveform_extractor) - - # Load model - self.denoiser = self.load_model() - - def assert_model_and_waveform_temporal_match(self, waveform_extractor: WaveformsNode): - """ - Asserts that the model and the waveform extractor have the same temporal parameters - """ - # Extract temporal parameters from the waveform extractor - waveforms_ms_before = waveform_extractor.ms_before - waveforms_ms_after = waveform_extractor.ms_after - waveforms_sampling_frequency = waveform_extractor.recording.get_sampling_frequency() - - # Load the model temporal parameters - repo_id = "SpikeInterface/test_repo" - subfolder = "mearec_toy_model" - filename = "params.json" - - json_file_path = hf_hub_download(repo_id=repo_id, subfolder=subfolder, filename=filename) - # Load the json file in the json_file_path_variable - with open(json_file_path, "r") as json_file: - peak_interval_dict = json.load(json_file) - - model_ms_before = peak_interval_dict["ms_before"] - model_ms_after = peak_interval_dict["ms_after"] - model_sampling_frequency = peak_interval_dict["sampling_frequency"] - - ms_before_mismatch = waveforms_ms_before != model_ms_before - ms_after_missmatch = waveforms_ms_after != model_ms_after - sampling_frequency_mismatch = waveforms_sampling_frequency != model_sampling_frequency - if ms_before_mismatch or ms_after_missmatch or sampling_frequency_mismatch: - exception_string = ( - "Model and waveforms mismatch \n" - f"{model_ms_before=} and {waveforms_ms_after=} \n" - f"{model_ms_after=} and {waveforms_ms_after=} \n" - f"{model_sampling_frequency=} and {waveforms_sampling_frequency=} \n" - ) - raise ValueError(exception_string) - - def load_model(self): - assert HAVE_HUGGINFACE, "To download models from Hugginface you need to install huggingface_hub" - - repo_id = "SpikeInterface/test_repo" - subfolder = "mearec_toy_model" - filename = "toy_model_marec.pt" - - model_path = hf_hub_download(repo_id=repo_id, subfolder=subfolder, filename=filename) - denoiser = SingleChannel1dCNNDenoiser(pretrained_path=model_path, spike_size=128) - denoiser = denoiser.load() - - return denoiser - - def compute(self, traces, peaks, waveforms): - num_channels = waveforms.shape[2] - - # Collapse channels and transform to torch tensor - temporal_waveforms = to_temporal_representation(waveforms) - temporal_waveforms_tensor = torch.from_numpy(temporal_waveforms).float() - - # Denoise - denoised_temporal_waveforms = self.denoiser(temporal_waveforms_tensor).detach().numpy() - - # Reconstruct representation with channels - denoised_waveforms = from_temporal_representation(denoised_temporal_waveforms, num_channels) - - return denoised_waveforms - - -if HAVE_TORCH: - - class SingleChannel1dCNNDenoiser(nn.Module): - def __init__(self, pretrained_path=None, n_filters=[16, 8], filter_sizes=[5, 11], spike_size=121): - super().__init__() - - out_channels_conv1, out_channels_conv_2 = n_filters - kernel_size_conv1, kernel_size_conv2 = filter_sizes - self.conv1 = nn.Sequential(nn.Conv1d(1, out_channels_conv1, kernel_size_conv1), nn.ReLU()) - self.conv2 = nn.Sequential(nn.Conv1d(out_channels_conv1, out_channels_conv_2, kernel_size_conv2), nn.ReLU()) - n_input_feat = out_channels_conv_2 * (spike_size - kernel_size_conv1 - kernel_size_conv2 + 2) - self.out = nn.Linear(n_input_feat, spike_size) - self.pretrained_path = pretrained_path - - def forward(self, x): - x = x[:, None] - x = self.conv1(x) - x = self.conv2(x) - x = x.view(x.shape[0], -1) - x = self.out(x) - return x - - def load(self, device="cpu"): - checkpoint = torch.load(self.pretrained_path, map_location=device) - self.load_state_dict(checkpoint) - return self diff --git a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py index 0170038c96..b136d4e2cf 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py +++ b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py @@ -14,32 +14,18 @@ from .waveform_utils import to_temporal_representation, from_temporal_representation -class TemporalPCBaseNode(WaveformsNode): +class TemporalPCMixin: def __init__( self, - recording: BaseRecording, - parents: List[PipelineNode], + waveform_node, pca_model=None, model_folder_path=None, - return_output=True, ): """ Base class for PCA projection nodes. Contains the logic of the fit method that should be inherited by all the child classess. The child should implement a compute method that does a specific operation (e.g. project, denoise, etc) """ - waveform_extractor = find_parent_of_type(parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"TemporalPCA should have a single {WaveformsNode.__name__} in its parents") - - super().__init__( - recording, - waveform_extractor.ms_before, - waveform_extractor.ms_after, - return_output=return_output, - parents=parents, - ) - if pca_model is None: self.model_folder_path = model_folder_path @@ -58,18 +44,18 @@ def __init__( with open(params_path, "rb") as f: self.params = json.load(f) - self.assert_model_and_waveform_temporal_match(waveform_extractor) + self.assert_model_and_waveform_temporal_match(waveform_node) else: self.pca_model = pca_model - def assert_model_and_waveform_temporal_match(self, waveform_extractor: WaveformsNode): + def assert_model_and_waveform_temporal_match(self, waveform_node: WaveformsNode): """ Asserts that the model and the waveform extractor have the same temporal parameters """ # Extract the first waveform extractor in the parents - waveforms_ms_before = waveform_extractor.ms_before - waveforms_ms_after = waveform_extractor.ms_after - waveforms_sampling_frequency = waveform_extractor.recording.get_sampling_frequency() + waveforms_ms_before = waveform_node.ms_before + waveforms_ms_after = waveform_node.ms_after + waveforms_sampling_frequency = waveform_node.recording.get_sampling_frequency() model_ms_before = self.params["ms_before"] model_ms_after = self.params["ms_after"] @@ -170,30 +156,12 @@ def fit( return model_folder_path -TemporalPCBaseNode.fit.__doc__ = TemporalPCBaseNode.fit.__doc__.format(_shared_job_kwargs_doc) +TemporalPCMixin.fit.__doc__ = TemporalPCMixin.fit.__doc__.format(_shared_job_kwargs_doc) -class TemporalPCAProjection(TemporalPCBaseNode): +class TemporalPCABaseNode(PipelineNode, TemporalPCMixin): """ - A step that performs a PCA projection on the waveforms extracted by a waveforms parent node. - - This class needs a model_folder_path with a trained model. A model can be trained with the - static method TemporalPCAProjection.fit(). - - - Parameters - ---------- - recording : BaseRecording - The recording object - parents: list - The parent nodes of this node. This should contain a mechanism to extract waveforms - pca_model: sklearn model | None - The already fitted sklearn model instead of model_folder_path - model_folder_path : str | Path | None - If pca_model is None, the path to the folder containing the pca model and the training metadata. - return_output: bool, default: True - use false to suppress the output of this node in the pipeline - + Base class for temporal PCA projection node """ def __init__( @@ -202,57 +170,26 @@ def __init__( parents: List[PipelineNode], pca_model=None, model_folder_path=None, - dtype="float32", return_output=True, ): - TemporalPCBaseNode.__init__( - self, - recording=recording, - parents=parents, - return_output=return_output, - pca_model=pca_model, - model_folder_path=model_folder_path, + PipelineNode.__init__(self, recording, parents=parents, return_output=return_output) + waveform_node = find_parent_of_type(parents, WaveformsNode) + if waveform_node is None: + raise TypeError(f"TemporalPCA should have a single {WaveformsNode.__name__} in its parents") + TemporalPCMixin.__init__( + self, waveform_node=waveform_node, model_folder_path=model_folder_path, pca_model=pca_model ) - self.n_components = self.pca_model.n_components - self.dtype = np.dtype(dtype) - - def compute(self, traces: np.ndarray, peaks: np.ndarray, waveforms: np.ndarray) -> np.ndarray: - """ - Projects the waveforms using the PCA model trained in the fit method or loaded from the model_folder_path. - - Parameters - ---------- - traces : np.ndarray - The traces of the recording. - peaks : np.ndarray - The peaks resulting from a peak_detection step. - waveforms : np.ndarray - Waveforms extracted from the recording using a WavefomExtractor node. - - Returns - ------- - np.ndarray - The projected waveforms. - - """ - - num_channels = waveforms.shape[2] - if waveforms.shape[0] > 0: - temporal_waveforms = to_temporal_representation(waveforms) - projected_temporal_waveforms = self.pca_model.transform(temporal_waveforms) - projected_waveforms = from_temporal_representation(projected_temporal_waveforms, num_channels) - else: - projected_waveforms = np.zeros((0, self.n_components, num_channels), dtype=self.dtype) - return projected_waveforms.astype(self.dtype, copy=False) + self.recording = recording -class TemporalPCADenoising(TemporalPCBaseNode): +class TemporalPCAProjection(TemporalPCABaseNode): """ - A step that performs a PCA denoising on the waveforms extracted by a peak_detection function. + A step that performs a PCA projection on the waveforms extracted by a waveforms parent node. This class needs a model_folder_path with a trained model. A model can be trained with the static method TemporalPCAProjection.fit(). + Parameters ---------- recording : BaseRecording @@ -274,9 +211,10 @@ def __init__( parents: List[PipelineNode], pca_model=None, model_folder_path=None, + dtype="float32", return_output=True, ): - TemporalPCBaseNode.__init__( + TemporalPCABaseNode.__init__( self, recording=recording, parents=parents, @@ -284,6 +222,8 @@ def __init__( pca_model=pca_model, model_folder_path=model_folder_path, ) + self.n_components = self.pca_model.n_components + self.dtype = np.dtype(dtype) def compute(self, traces: np.ndarray, peaks: np.ndarray, waveforms: np.ndarray) -> np.ndarray: """ @@ -304,20 +244,18 @@ def compute(self, traces: np.ndarray, peaks: np.ndarray, waveforms: np.ndarray) The projected waveforms. """ - num_channels = waveforms.shape[2] + num_channels = waveforms.shape[2] if waveforms.shape[0] > 0: - temporal_waveform = to_temporal_representation(waveforms) - projected_temporal_waveforms = self.pca_model.transform(temporal_waveform) - temporal_denoised_waveforms = self.pca_model.inverse_transform(projected_temporal_waveforms) - denoised_waveforms = from_temporal_representation(temporal_denoised_waveforms, num_channels) + temporal_waveforms = to_temporal_representation(waveforms) + projected_temporal_waveforms = self.pca_model.transform(temporal_waveforms) + projected_waveforms = from_temporal_representation(projected_temporal_waveforms, num_channels) else: - denoised_waveforms = np.zeros_like(waveforms) - - return denoised_waveforms + projected_waveforms = np.zeros((0, self.n_components, num_channels), dtype=self.dtype) + return projected_waveforms.astype(self.dtype, copy=False) -class MotionAwareTemporalPCAProjection(TemporalPCBaseNode): +class MotionAwareTemporalPCAProjection(TemporalPCABaseNode): """ Similar to TemporalPCAProjection but also apply interpolation to revert a motion. @@ -353,7 +291,7 @@ def __init__( dtype="float32", return_output=True, ): - TemporalPCBaseNode.__init__( + TemporalPCABaseNode.__init__( self, recording=recording, parents=parents, diff --git a/src/spikeinterface/sortingcomponents/waveforms/tests/conftest.py b/src/spikeinterface/sortingcomponents/waveforms/tests/conftest.py index f3ab58d87b..fdf74420ac 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/tests/conftest.py +++ b/src/spikeinterface/sortingcomponents/waveforms/tests/conftest.py @@ -22,6 +22,18 @@ def generated_recording(): return recording +@pytest.fixture(scope="module") +def generated_recording_30khz(): + recording, sorting = generate_ground_truth_recording( + durations=[10.0], + sampling_frequency=30000.0, + num_channels=32, + num_units=10, + seed=2205, + ) + return recording + + @pytest.fixture(scope="module") def detected_peaks(generated_recording, chunk_executor_kwargs): recording = generated_recording diff --git a/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py index 6ae694c083..6feab6889d 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py @@ -1,5 +1,5 @@ from spikeinterface.core.node_pipeline import run_node_pipeline, PeakRetriever, ExtractDenseWaveforms -from spikeinterface.sortingcomponents.waveforms.neural_network_denoiser import SingleChannelToyDenoiser +from spikeinterface.sortingcomponents.waveforms.denoising.neural_network_denoiser import SingleChannelDenoiser def test_single_channel_toy_denoiser_in_peak_pipeline(generated_recording, detected_peaks, chunk_executor_kwargs): @@ -8,16 +8,45 @@ def test_single_channel_toy_denoiser_in_peak_pipeline(generated_recording, detec ms_before = 2.0 ms_after = 2.0 - waveform_extraction = ExtractDenseWaveforms(recording, ms_before=ms_before, ms_after=ms_after, return_output=True) # Build nodes for computation peak_retriever = PeakRetriever(recording, peaks) waveform_extraction = ExtractDenseWaveforms( recording, parents=[peak_retriever], ms_before=ms_before, ms_after=ms_after, return_output=True ) - toy_denoiser = SingleChannelToyDenoiser(recording, parents=[peak_retriever, waveform_extraction]) + toy_denoiser = SingleChannelDenoiser( + recording, + parents=[peak_retriever, waveform_extraction], + repo_id="SpikeInterface/waveform_denoiser", + model_name="toy_model_mearec", + ) nodes = [peak_retriever, waveform_extraction, toy_denoiser] waveforms, denoised_waveforms = run_node_pipeline(recording, nodes=nodes, job_kwargs=chunk_executor_kwargs) assert waveforms.shape == denoised_waveforms.shape + + +def test_single_channel_yass_denoiser(generated_recording_30khz, detected_peaks, chunk_executor_kwargs): + recording = generated_recording_30khz + peaks = detected_peaks + + nbefore = 42 + nafter = 79 + + # Build nodes for computation + peak_retriever = PeakRetriever(recording, peaks) + waveform_extraction = ExtractDenseWaveforms( + recording, parents=[peak_retriever], nbefore=nbefore, nafter=nafter, return_output=True + ) + yass_denoiser = SingleChannelDenoiser( + recording, + parents=[peak_retriever, waveform_extraction], + repo_id="spikeinterface/waveform_denoiser", + model_name="yass_ibl", + ) + + nodes = [peak_retriever, waveform_extraction, yass_denoiser] + waveforms, denoised_waveforms = run_node_pipeline(recording, nodes=nodes, job_kwargs=chunk_executor_kwargs) + + assert waveforms.shape == denoised_waveforms.shape diff --git a/src/spikeinterface/sortingcomponents/waveforms/tests/test_savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/tests/test_savgol_denoiser.py index 651b681078..924d57d493 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/tests/test_savgol_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/tests/test_savgol_denoiser.py @@ -1,7 +1,7 @@ import pytest -from spikeinterface.sortingcomponents.waveforms.savgol_denoiser import SavGolDenoiser +from spikeinterface.sortingcomponents.waveforms.denoising.savgol_denoiser import SavGolDenoiser from spikeinterface.core.node_pipeline import ( PeakRetriever, diff --git a/src/spikeinterface/sortingcomponents/waveforms/tests/test_temporal_pca.py b/src/spikeinterface/sortingcomponents/waveforms/tests/test_temporal_pca.py index dc01ab3431..230f650475 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/tests/test_temporal_pca.py +++ b/src/spikeinterface/sortingcomponents/waveforms/tests/test_temporal_pca.py @@ -1,7 +1,8 @@ import pytest -from spikeinterface.sortingcomponents.waveforms.temporal_pca import TemporalPCAProjection, TemporalPCADenoising +from spikeinterface.sortingcomponents.waveforms.temporal_pca import TemporalPCAProjection +from spikeinterface.sortingcomponents.waveforms.denoising.method_list import TemporalPCADenoiser from spikeinterface.core.node_pipeline import ( PeakRetriever, ExtractDenseWaveforms, @@ -67,7 +68,7 @@ def test_pca_denoising(generated_recording, detected_peaks, model_path_of_traine extract_waveforms = ExtractDenseWaveforms( recording=recording, parents=[peak_retriever], ms_before=ms_before, ms_after=ms_after, return_output=True ) - pca_denoising = TemporalPCADenoising( + pca_denoising = TemporalPCADenoiser( recording=recording, model_folder_path=model_folder_path, parents=[peak_retriever, extract_waveforms] ) pipeline_nodes = [peak_retriever, extract_waveforms, pca_denoising] @@ -97,7 +98,7 @@ def test_pca_denoising_sparse(generated_recording, detected_peaks, model_path_of radius_um=radius_um, return_output=True, ) - pca_denoising = TemporalPCADenoising( + pca_denoising = TemporalPCADenoiser( recording=recording, model_folder_path=model_folder_path, parents=[peak_retriever, extract_waveforms] ) pipeline_nodes = [peak_retriever, extract_waveforms, pca_denoising] diff --git a/src/spikeinterface/sortingcomponents/waveforms/waveform_thresholder.py b/src/spikeinterface/sortingcomponents/waveforms/waveform_thresholder.py index ec223d0047..cefe3750e3 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/waveform_thresholder.py +++ b/src/spikeinterface/sortingcomponents/waveforms/waveform_thresholder.py @@ -5,9 +5,10 @@ from spikeinterface.core import BaseRecording, get_noise_levels from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type +from .waveform_utils import WaveformTransformer -class WaveformThresholder(WaveformsNode): +class WaveformThresholder(WaveformTransformer): """ A node that performs waveform thresholding based on a selected feature. @@ -47,14 +48,8 @@ def __init__( random_chunk_kwargs: dict = {}, operator: callable = operator.le, ): - waveform_extractor = find_parent_of_type(parents, WaveformsNode) - if waveform_extractor is None: - raise TypeError(f"SavGolDenoiser should have a single {WaveformsNode.__name__} in its parents") - super().__init__( recording, - waveform_extractor.ms_before, - waveform_extractor.ms_after, return_output=return_output, parents=parents, ) diff --git a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py index a30514eb1d..3eb0b2e1b1 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py +++ b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py @@ -1,3 +1,41 @@ +from spikeinterface.core import BaseRecording +from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type + + +class WaveformTransformer(WaveformsNode): + """ + Base class for waveform transformers. It is a WaveformsNode that takes waveforms as input and returns transformed waveforms. + It can be used to apply any transformation to the waveforms, such as denoising, filtering, etc. + + Parameters + ---------- + recording: BaseRecording + The recording extractor object + return_output: bool, default: True + Whether to return output from this node + parents: list of PipelineNodes, default: None + The parent nodes of this node. This should contain a mechanism to extract waveforms + """ + + def __init__(self, recording: BaseRecording, return_output: bool = True, parents: list[PipelineNode] = None): + waveforms_node = find_parent_of_type(parents, WaveformsNode) + if waveforms_node is None: + raise TypeError(f"{self.__class__.__name__} should have a single {WaveformsNode.__name__} in its parents") + + # We only propagate nbefore/nafter since the WaveformNode already made the conversion + super().__init__( + recording, + nbefore=waveforms_node.nbefore, + nafter=waveforms_node.nafter, + return_output=return_output, + parents=parents, + ) + self.waveforms_node = waveforms_node + # Propagate waveforms node parameters + self.sparse_waveforms = waveforms_node.sparse_waveforms + self.neighbours_mask = waveforms_node.neighbours_mask + + def to_temporal_representation(waveforms): """ Transform waveforms to temporal representation. Collapses the channel dimension (spatial) leaving only