From 2937371abfd78e30174996d03008f97d2e7002c6 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 9 Jul 2026 16:12:00 +0100 Subject: [PATCH 01/16] feat: enhance single channel denoiser and waveformnode --- src/spikeinterface/core/core_tools.py | 5 + src/spikeinterface/core/node_pipeline.py | 65 +++-- .../waveforms/neural_network_denoiser.py | 246 ++++++++++++++---- .../waveforms/tests/conftest.py | 12 + .../tests/test_neural_network_denoiser.py | 36 ++- 5 files changed, 295 insertions(+), 69 deletions(-) diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index 0486f9742c..15a8d6c17e 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -773,3 +773,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 cc4d66a072..a43b228d47 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -12,7 +12,7 @@ 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: @@ -297,8 +297,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,13 +321,30 @@ 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 @@ -333,8 +352,10 @@ 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 +368,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 +389,8 @@ def __init__( parents=parents, ms_before=ms_before, ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, return_output=return_output, ) @@ -379,8 +406,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 +430,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 +454,8 @@ def __init__( parents=parents, ms_before=ms_before, ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, return_output=return_output, ) diff --git a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py index 77cada6ed3..e6be40174d 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py @@ -1,6 +1,11 @@ import json +import warnings import importlib.util +from pathlib import Path from typing import List, Optional +import numpy as np + +from huggingface_hub import list_repo_files if importlib.util.find_spec("torch") is not None: import torch @@ -11,86 +16,188 @@ HAVE_TORCH = False if importlib.util.find_spec("huggingface_hub") is not None: - from huggingface_hub import hf_hub_download + from huggingface_hub import hf_hub_download, list_repo_files - HAVE_HUGGINFACE = True + HAVE_HUGGINGFACE = True else: - HAVE_HUGGINFACE = False + HAVE_HUGGINGFACE = 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): +class SingleChannelDenoiser(WaveformsNode): + """ + 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. + """ + def __init__( - self, recording: BaseRecording, return_output: bool = True, parents: Optional[List[PipelineNode]] = None + 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, + spike_size=121, ): - 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: + assert HAVE_TORCH, "To use the SingleChannelDenoiser you need to install torch" + waveform_node = find_parent_of_type(parents, WaveformsNode) + if waveform_node 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, + waveform_node.ms_before, + waveform_node.ms_after, return_output=return_output, parents=parents, ) - - self.assert_model_and_waveform_temporal_match(waveform_extractor) - + 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") # Load model - self.denoiser = self.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 + ) - def assert_model_and_waveform_temporal_match(self, waveform_extractor: WaveformsNode): + self.assert_model_and_waveform_temporal_match( + waveform_node, model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path + ) + + def assert_model_and_waveform_temporal_match( + self, + waveform_node: WaveformsNode, + 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 = 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.sampling_frequency - # Load the model temporal parameters - repo_id = "SpikeInterface/test_repo" - subfolder = "mearec_toy_model" - filename = "params.json" + 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 - 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() + 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") - return denoiser + if model_num_samples is not None: + if model_num_samples != waveform_node.nbefore + waveform_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 - waveform_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, + ): + 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() + model_name = Path(model_path).stem + return denoiser, model_relative_path def compute(self, traces, peaks, waveforms): num_channels = waveforms.shape[2] @@ -134,3 +241,38 @@ def load(self, device="cpu"): checkpoint = torch.load(self.pretrained_path, map_location=device) self.load_state_dict(checkpoint) return self + + 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 temporal parameters from the waveform extractor + waveforms_ms_before = waveform_node.ms_before + waveforms_ms_after = waveform_node.ms_after + waveforms_sampling_frequency = waveform_node.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) 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..bd5cca41b4 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.neural_network_denoiser import SingleChannelDenoiser def test_single_channel_toy_denoiser_in_peak_pipeline(generated_recording, detected_peaks, chunk_executor_kwargs): @@ -15,9 +15,41 @@ def test_single_channel_toy_denoiser_in_peak_pipeline(generated_recording, detec 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/test_repo", + model_name="mearec_toy_denoiser", + spike_size=128, + ) 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 + waveform_extraction = ExtractDenseWaveforms(recording, nbefore=nbefore, nafter=nafter, return_output=True) + + # 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 From 11db29d5ae036962b015f8f0db5bde788e599c9d Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 9 Jul 2026 16:23:12 +0100 Subject: [PATCH 02/16] fix: remove double definition --- .../waveforms/tests/test_neural_network_denoiser.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) 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 bd5cca41b4..23dd20e577 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py @@ -8,7 +8,6 @@ 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) @@ -18,7 +17,7 @@ def test_single_channel_toy_denoiser_in_peak_pipeline(generated_recording, detec toy_denoiser = SingleChannelDenoiser( recording, parents=[peak_retriever, waveform_extraction], - repo_id="SpikeInterface/test_repo", + repo_id="SpikeInterface/waveform_denoiser", model_name="mearec_toy_denoiser", spike_size=128, ) @@ -35,7 +34,6 @@ def test_single_channel_yass_denoiser(generated_recording_30khz, detected_peaks, nbefore = 42 nafter = 79 - waveform_extraction = ExtractDenseWaveforms(recording, nbefore=nbefore, nafter=nafter, return_output=True) # Build nodes for computation peak_retriever = PeakRetriever(recording, peaks) From 5a73c442fb734604a827fce3b499217136585c29 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 9 Jul 2026 16:27:53 +0100 Subject: [PATCH 03/16] fix: model name --- .../sortingcomponents/waveforms/neural_network_denoiser.py | 2 +- .../waveforms/tests/test_neural_network_denoiser.py | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py index e6be40174d..913c0eaba3 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py @@ -55,7 +55,6 @@ def __init__( model_folder: Optional[str] = None, repo_id: Optional[str] = None, model_name: Optional[str] = None, - spike_size=121, ): assert HAVE_TORCH, "To use the SingleChannelDenoiser you need to install torch" waveform_node = find_parent_of_type(parents, WaveformsNode) @@ -73,6 +72,7 @@ def __init__( 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 = waveform_node.nbefore + waveform_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 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 23dd20e577..fdca4a04b7 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/tests/test_neural_network_denoiser.py @@ -18,8 +18,7 @@ def test_single_channel_toy_denoiser_in_peak_pipeline(generated_recording, detec recording, parents=[peak_retriever, waveform_extraction], repo_id="SpikeInterface/waveform_denoiser", - model_name="mearec_toy_denoiser", - spike_size=128, + model_name="toy_model_mearec", ) nodes = [peak_retriever, waveform_extraction, toy_denoiser] From 48784aeb358d5004677bcf2c27862545f803a1ea Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 9 Jul 2026 16:29:49 +0100 Subject: [PATCH 04/16] fix: remove duplicated function --- .../waveforms/neural_network_denoiser.py | 35 ------------------- 1 file changed, 35 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py index 913c0eaba3..9b67ac9141 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py @@ -241,38 +241,3 @@ def load(self, device="cpu"): checkpoint = torch.load(self.pretrained_path, map_location=device) self.load_state_dict(checkpoint) return self - - 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 temporal parameters from the waveform extractor - waveforms_ms_before = waveform_node.ms_before - waveforms_ms_after = waveform_node.ms_after - waveforms_sampling_frequency = waveform_node.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) From d9bf644554e114509e081ed5e667a6e38b6743e6 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Fri, 10 Jul 2026 09:45:19 +0100 Subject: [PATCH 05/16] fix: remove import --- .../sortingcomponents/waveforms/neural_network_denoiser.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py index 9b67ac9141..1f03f139d9 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py @@ -5,8 +5,6 @@ from typing import List, Optional import numpy as np -from huggingface_hub import list_repo_files - if importlib.util.find_spec("torch") is not None: import torch from torch import nn From 856fe9523f6335c71815108182a6bd0b903e97dd Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Tue, 21 Jul 2026 18:45:42 +0200 Subject: [PATCH 06/16] feat: denoising waveforms multi-method and sparse localization --- src/spikeinterface/core/node_pipeline.py | 1 + src/spikeinterface/preprocessing/motion.py | 136 +++++++--- .../clustering/random_projections.py | 2 +- .../peak_localization/base.py | 34 ++- .../peak_localization/center_of_mass.py | 9 +- .../peak_localization/grid.py | 16 +- .../peak_localization/main.py | 27 +- .../peak_localization/monopolar.py | 7 +- .../tests/test_peak_localization.py | 36 ++- .../waveforms/hanning_filter.py | 3 + .../waveforms/neural_network_denoiser.py | 241 ------------------ .../waveforms/savgol_denoiser.py | 58 ----- .../waveforms/temporal_pca.py | 71 ------ .../tests/test_neural_network_denoiser.py | 2 +- .../waveforms/tests/test_savgol_denoiser.py | 2 +- 15 files changed, 193 insertions(+), 452 deletions(-) delete mode 100644 src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py delete mode 100644 src/spikeinterface/sortingcomponents/waveforms/savgol_denoiser.py diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index c9db938ea7..9cbeda7c84 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -19,6 +19,7 @@ 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__( diff --git a/src/spikeinterface/preprocessing/motion.py b/src/spikeinterface/preprocessing/motion.py index 62bdbfd7a9..4a28362fec 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,6 +310,8 @@ def compute_motion( ] = "dredge_fast", detect_kwargs: dict = {}, select_kwargs: dict = {}, + denoise_kwargs: dict = {}, + extract_waveforms_kwargs=None, localize_peaks_kwargs: dict = {}, estimate_motion_kwargs: dict = {}, output_motion_info: bool = False, @@ -304,6 +326,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 +338,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 +350,24 @@ 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 ExtractDenseWaveforms, run_node_pipeline, PeakRetriever 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 +377,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 +400,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 +413,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 +424,63 @@ 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) - run_times = dict( - detect_peaks=t1 - t0, - select_peaks=t2 - t1, - localize_peaks=t3 - t2, + pipeline_nodes = [peaks_node] + + if extract_waveforms_kwargs is None: + extract_waveforms_kwargs = {"ms_before": 0.1, "ms_after": 0.3} + extract_dense_node = ExtractDenseWaveforms(recording, parents=[peaks_node], **extract_waveforms_kwargs) + pipeline_nodes.append(extract_dense_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_dense_node], **denoise_kwargs_without_method ) + extract_waveforms_node = denoise_node + pipeline_nodes.append(denoise_node) + pipeline_run_time_name += "denoise-localize" + else: + extract_waveforms_node = extract_dense_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_node], + 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 +528,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 +540,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 +561,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 +601,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..ba3da340b1 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/base.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/base.py @@ -1,10 +1,18 @@ +import numpy as np from spikeinterface.core.node_pipeline import ( PipelineNode, ) +from spikeinterface.core.node_pipeline import ( + find_parent_of_type, + WaveformsNode, + ExtractDenseWaveforms, + ExtractSparseWaveforms, +) from spikeinterface.core import get_channel_distances +# TODO: make this sparse and pre-instantiate neighbor mask in case of ExtractSparseWaveforms class LocalizeBase(PipelineNode): def __init__(self, recording, parents, return_output=True, radius_um=75.0): @@ -12,9 +20,29 @@ def __init__(self, recording, parents, return_output=True, radius_um=75.0): self.recording = recording self.radius_um = radius_um self.contact_locations = recording.get_channel_locations() - self.channel_distance = get_channel_distances(recording) - self.neighbours_mask = self.channel_distance <= radius_um - self._kwargs["radius_um"] = radius_um + + # 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 + if isinstance(waveform_extractor, ExtractSparseWaveforms): + self.sparse_waveforms = True + self.neighbours_mask = waveform_extractor.neighbours_mask + else: + self.sparse_waveforms = False + self.channel_distance = get_channel_distances(recording) + self.neighbours_mask = self.channel_distance <= radius_um + self._kwargs["radius_um"] = radius_um def get_dtype(self): return self._dtype + + # TODO: fix sparsity here + def get_sparse_waveform(self, waveform, chan_inds): + """Get sparse waveforms from dense waveforms""" + if self.sparse_waveforms: + return waveform + 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..e81343832c 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) 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..77d665d0e9 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/grid.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/grid.py @@ -69,14 +69,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 +89,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 +124,11 @@ 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 + + wf = self.get_sparse_waveform(waveforms[idx], np.flatnonzero(self.neighbours_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..1dc1c6b0f5 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,9 @@ def get_localization_pipeline_nodes( method_kwargs=None, ms_before=0.5, ms_after=0.5, + nbefore=None, + nafter=None, + waveform_method="dense", job_kwargs=None, ): @@ -34,10 +38,15 @@ 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 - ) + waveform_kwargs = 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 + ) method_class = peak_localization_methods[method] @@ -55,9 +64,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 +78,9 @@ def localize_peaks( method_kwargs=None, ms_before=0.5, ms_after=0.5, + nbefore=None, + nafter=None, + waveform_method="dense", pipeline_kwargs=None, verbose=False, job_kwargs=None, @@ -149,6 +161,9 @@ def localize_peaks( method_kwargs=method_kwargs, ms_before=ms_before, ms_after=ms_after, + nbefore=nbefore, + nafter=nafter, + waveform_method=waveform_method, job_kwargs=job_kwargs, ) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py b/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py index 942074ef31..4d3fa6c2d6 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) 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..9660c3a4df 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,26 @@ from spikeinterface.sortingcomponents.tests.common import make_dataset -def test_localize_peaks(): +@pytest.fixture +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 + + +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,6 +98,25 @@ def test_localize_peaks(): assert peaks.size == peak_locations.shape[0] list_locations.append(("minimize_with_log_penality_v_peak", peak_locations)) + +@pytest.mark.parametrize("method", ["center_of_mass", "grid_convolution", "monopolar_triangulation"]) +def test_localize_peaks_sparse(peaks_and_recording, method): + recording, peaks = peaks_and_recording + + job_kwargs = dict(n_jobs=1, chunk_size=10000, progress_bar=True) + + # test sparse waveforms + peak_locations = localize_peaks( + recording, + peaks, + method_kwargs=dict( + method=method, + ), + waveform_method="sparse", + job_kwargs=job_kwargs, + ) + assert peaks.size == peak_locations.shape[0] + # DEBUG # import MEArec # recgen = MEArec.load_recordings(recordings=local_path, return_h5_objects=True, 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 1f03f139d9..0000000000 --- a/src/spikeinterface/sortingcomponents/waveforms/neural_network_denoiser.py +++ /dev/null @@ -1,241 +0,0 @@ -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, WaveformsNode, find_parent_of_type -from .waveform_utils import to_temporal_representation, from_temporal_representation - - -class SingleChannelDenoiser(WaveformsNode): - """ - 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. - """ - - 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, - ): - assert HAVE_TORCH, "To use the SingleChannelDenoiser you need to install torch" - waveform_node = find_parent_of_type(parents, WaveformsNode) - if waveform_node is None: - raise TypeError(f"Model should have a {WaveformsNode.__name__} in its parents") - - super().__init__( - recording, - waveform_node.ms_before, - waveform_node.ms_after, - 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 = waveform_node.nbefore + waveform_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 - ) - - self.assert_model_and_waveform_temporal_match( - waveform_node, model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path - ) - - def assert_model_and_waveform_temporal_match( - self, - waveform_node: WaveformsNode, - 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 = waveform_node.ms_before - waveforms_ms_after = waveform_node.ms_after - waveforms_sampling_frequency = waveform_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 != waveform_node.nbefore + waveform_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 - waveform_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, - ): - 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() - 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/savgol_denoiser.py deleted file mode 100644 index 1ed9e4bffa..0000000000 --- a/src/spikeinterface/sortingcomponents/waveforms/savgol_denoiser.py +++ /dev/null @@ -1,58 +0,0 @@ -from typing import List, Optional - -from spikeinterface.core import BaseRecording -from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type - - -class SavGolDenoiser(WaveformsNode): - """ - Waveform Denoiser based on a simple Savitzky-Golay filtering - https://en.wikipedia.org/wiki/Savitzky%E2%80%93Golay_filter - - 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 - 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__( - self, - recording: BaseRecording, - return_output: bool = True, - parents: Optional[List[PipelineNode]] = None, - 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, - ) - - 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)) - - def compute(self, traces, peaks, waveforms): - # Denoise - import scipy.signal - - denoised_waveforms = scipy.signal.savgol_filter(waveforms, self.window_length, self.order, axis=1) - - return denoised_waveforms diff --git a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py index 0170038c96..36bb32cda0 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py +++ b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py @@ -246,77 +246,6 @@ def compute(self, traces: np.ndarray, peaks: np.ndarray, waveforms: np.ndarray) return projected_waveforms.astype(self.dtype, copy=False) -class TemporalPCADenoising(TemporalPCBaseNode): - """ - 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 - - """ - - def __init__( - self, - recording: BaseRecording, - parents: List[PipelineNode], - pca_model=None, - model_folder_path=None, - return_output=True, - ): - TemporalPCBaseNode.__init__( - self, - recording=recording, - parents=parents, - return_output=return_output, - pca_model=pca_model, - model_folder_path=model_folder_path, - ) - - 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_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 - - class MotionAwareTemporalPCAProjection(TemporalPCBaseNode): """ Similar to TemporalPCAProjection but also apply interpolation to revert a motion. 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 fdca4a04b7..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 SingleChannelDenoiser +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): 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, From 4ef5f1835dea338ce3bfdc252e9a757e5232c1fe Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 09:44:51 +0200 Subject: [PATCH 07/16] fix: add waveform/denoising module --- .../waveforms/denoising/__init__.py | 2 + .../waveforms/denoising/main.py | 151 +++++++++++ .../waveforms/denoising/method_list.py | 6 + .../denoising/neural_network_denoiser.py | 253 ++++++++++++++++++ .../waveforms/denoising/savgol_denoiser.py | 66 +++++ .../denoising/temporal_pca_denoiser.py | 85 ++++++ 6 files changed, 563 insertions(+) create mode 100644 src/spikeinterface/sortingcomponents/waveforms/denoising/__init__.py create mode 100644 src/spikeinterface/sortingcomponents/waveforms/denoising/main.py create mode 100644 src/spikeinterface/sortingcomponents/waveforms/denoising/method_list.py create mode 100644 src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py create mode 100644 src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py create mode 100644 src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py 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..6c8aa999c7 --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py @@ -0,0 +1,253 @@ +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, WaveformsNode, find_parent_of_type +from ..waveform_utils import to_temporal_representation, from_temporal_representation + + +class SingleChannelDenoiser(WaveformsNode): + """ + 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" + waveform_node = find_parent_of_type(parents, WaveformsNode) + if waveform_node is None: + raise TypeError(f"Model should have a {WaveformsNode.__name__} in its parents") + + super().__init__( + recording, + waveform_node.ms_before, + waveform_node.ms_after, + 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 = waveform_node.nbefore + waveform_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( + waveform_node, model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path + ) + + def assert_model_and_waveform_temporal_match( + self, + waveform_node: WaveformsNode, + 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 = waveform_node.ms_before + waveforms_ms_after = waveform_node.ms_after + waveforms_sampling_frequency = waveform_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 != waveform_node.nbefore + waveform_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 - waveform_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/denoising/savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py new file mode 100644 index 0000000000..a09352720f --- /dev/null +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py @@ -0,0 +1,66 @@ +from typing import List, Optional + +from spikeinterface.core import BaseRecording +from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type + + +class SavGolDenoiser(WaveformsNode): + """ + Waveform Denoiser based on a simple Savitzky-Golay filtering + https://en.wikipedia.org/wiki/Savitzky%E2%80%93Golay_filter + + 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 + order: int, default: 3 + The order of the filter + window_length_ms: float, default: 0.25 + 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__( + self, + recording: BaseRecording, + return_output: bool = True, + parents: Optional[List[PipelineNode]] = None, + 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, + ) + + 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)) + + def compute(self, traces, peaks, waveforms): + # Denoise + import scipy.signal + + denoised_waveforms = scipy.signal.savgol_filter(waveforms, self.window_length, self.order, axis=1) + + return denoised_waveforms 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..fca01b60d1 --- /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 ..temporal_pca import TemporalPCBaseNode, to_temporal_representation, from_temporal_representation + + +class TemporalPCADenoiser(TemporalPCBaseNode): + """ + 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, + ): + TemporalPCBaseNode.__init__( + self, + recording=recording, + parents=parents, + return_output=return_output, + 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 From 844d1d73e86451a05fc0d7c65678207a81acfaf2 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 10:45:33 +0200 Subject: [PATCH 08/16] feat: add option to localize peaks from sparse waveforms --- .../peak_localization/base.py | 31 +++++-- .../peak_localization/center_of_mass.py | 2 +- .../peak_localization/grid.py | 2 +- .../peak_localization/main.py | 18 +++- .../peak_localization/monopolar.py | 2 +- .../tests/test_peak_localization.py | 82 +++++++++++-------- 6 files changed, 90 insertions(+), 47 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/base.py b/src/spikeinterface/sortingcomponents/peak_localization/base.py index ba3da340b1..11e98671c2 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/base.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/base.py @@ -10,6 +10,7 @@ ) from spikeinterface.core import get_channel_distances +from spikeinterface.sortingcomponents import waveforms # TODO: make this sparse and pre-instantiate neighbor mask in case of ExtractSparseWaveforms @@ -20,6 +21,7 @@ def __init__(self, recording, parents, return_output=True, radius_um=75.0): self.recording = recording 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_extractor = find_parent_of_type(self.parents, WaveformsNode) @@ -27,22 +29,35 @@ def __init__(self, recording, parents, return_output=True, radius_um=75.0): 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.neighbours_mask = self.channel_distance <= radius_um if isinstance(waveform_extractor, ExtractSparseWaveforms): self.sparse_waveforms = True - self.neighbours_mask = waveform_extractor.neighbours_mask + # 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_extractor.neighbours_mask + self.neighbours_mask &= self.extraction_neighbours_mask else: self.sparse_waveforms = False - self.channel_distance = get_channel_distances(recording) - self.neighbours_mask = self.channel_distance <= radius_um - self._kwargs["radius_um"] = radius_um + self._kwargs["radius_um"] = radius_um def get_dtype(self): return self._dtype - # TODO: fix sparsity here - def get_sparse_waveform(self, waveform, chan_inds): + def get_sparse_waveform(self, waveform, chan_inds, main_chan): """Get sparse waveforms from dense waveforms""" if self.sparse_waveforms: - return waveform + # 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) + if waveform.ndim == 2: + return waveform[:, local_inds] + else: + return waveform[:, :, local_inds] else: - return waveform[:, :, chan_inds] + 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 e81343832c..de4ebbd953 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/center_of_mass.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/center_of_mass.py @@ -42,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 = self.get_sparse_waveform(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 77d665d0e9..697319775b 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/grid.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/grid.py @@ -125,7 +125,7 @@ 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 - wf = self.get_sparse_waveform(waveforms[idx], np.flatnonzero(self.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 = (wf * self.prototype).sum(axis=1) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/main.py b/src/spikeinterface/sortingcomponents/peak_localization/main.py index 1dc1c6b0f5..6edfd7a557 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/main.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/main.py @@ -28,6 +28,7 @@ def get_localization_pipeline_nodes( nbefore=None, nafter=None, waveform_method="dense", + waveform_kwargs=None, job_kwargs=None, ): @@ -38,7 +39,9 @@ def get_localization_pipeline_nodes( assert method_kwargs is not None # peak_retriever = PeakRetriever(recording, peaks) - waveform_kwargs = dict(ms_before=ms_before, ms_after=ms_after, nbefore=nbefore, nafter=nafter) + 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 @@ -47,6 +50,7 @@ def get_localization_pipeline_nodes( 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] @@ -81,6 +85,7 @@ def localize_peaks( nbefore=None, nafter=None, waveform_method="dense", + waveform_kwargs=None, pipeline_kwargs=None, verbose=False, job_kwargs=None, @@ -107,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 @@ -164,6 +179,7 @@ def localize_peaks( 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 4d3fa6c2d6..e52fbd4a36 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/monopolar.py @@ -83,7 +83,7 @@ def compute(self, traces, peaks, waveforms): chan_inds = np.flatnonzero(chan_mask) local_contact_locations = self.contact_locations[chan_inds, :] - wf = self.get_sparse_waveform(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 9660c3a4df..9674b01b9c 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py @@ -7,8 +7,7 @@ from spikeinterface.sortingcomponents.tests.common import make_dataset -@pytest.fixture -def peaks_and_recording(): +def _peaks_and_recording(): recording, _ = make_dataset() peaks = detect_peaks( @@ -21,6 +20,11 @@ def peaks_and_recording(): 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 @@ -99,11 +103,11 @@ def test_localize_peaks(peaks_and_recording): list_locations.append(("minimize_with_log_penality_v_peak", peak_locations)) -@pytest.mark.parametrize("method", ["center_of_mass", "grid_convolution", "monopolar_triangulation"]) +@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=1, chunk_size=10000, progress_bar=True) + job_kwargs = dict(n_jobs=2, chunk_size=10000, progress_bar=True) # test sparse waveforms peak_locations = localize_peaks( @@ -112,41 +116,49 @@ def test_localize_peaks_sparse(peaks_and_recording, method): method_kwargs=dict( method=method, ), - waveform_method="sparse", + waveform_method="sparse", # if method != "grid_convolution" else "dense", job_kwargs=job_kwargs, ) assert peaks.size == peak_locations.shape[0] - # 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_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") From dd03739876468623366fd260cd2f24bbf5050d0f Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 10:51:18 +0200 Subject: [PATCH 09/16] feat: sparse waveforms for motion --- src/spikeinterface/preprocessing/motion.py | 25 +++++++++++++++------- 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/src/spikeinterface/preprocessing/motion.py b/src/spikeinterface/preprocessing/motion.py index 4a28362fec..2f5844cf7f 100644 --- a/src/spikeinterface/preprocessing/motion.py +++ b/src/spikeinterface/preprocessing/motion.py @@ -311,9 +311,10 @@ def compute_motion( detect_kwargs: dict = {}, select_kwargs: dict = {}, denoise_kwargs: dict = {}, - extract_waveforms_kwargs=None, 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, @@ -362,7 +363,12 @@ def compute_motion( 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, PeakRetriever + 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 @@ -435,8 +441,11 @@ def compute_motion( if extract_waveforms_kwargs is None: extract_waveforms_kwargs = {"ms_before": 0.1, "ms_after": 0.3} - extract_dense_node = ExtractDenseWaveforms(recording, parents=[peaks_node], **extract_waveforms_kwargs) - pipeline_nodes.append(extract_dense_node) + 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"] @@ -445,13 +454,13 @@ def compute_motion( key: denoise_kwarg for key, denoise_kwarg in denoise_kwargs.items() if key != "method" } denoise_node = denoise_class( - recording, parents=[peaks_node, extract_dense_node], **denoise_kwargs_without_method + recording, parents=[peaks_node, extract_waveforms_node], **denoise_kwargs_without_method ) - extract_waveforms_node = denoise_node + extract_waveforms_for_localization = denoise_node pipeline_nodes.append(denoise_node) pipeline_run_time_name += "denoise-localize" else: - extract_waveforms_node = extract_dense_node + extract_waveforms_for_localization = extract_waveforms_node pipeline_run_time_name += "localize" # node detect + localize @@ -462,7 +471,7 @@ def compute_motion( } localize_node = method_class( recording, - parents=[peaks_node, extract_waveforms_node], + parents=[peaks_node, extract_waveforms_for_localization], return_output=True, **localize_peaks_kwargs_without_method, ) From eef8dafcd0b7ea72bf61bfe4d729f7609fe2f434 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 11:03:41 +0200 Subject: [PATCH 10/16] fix:propagate sparse_waveforms attrs from base WaveformNode --- src/spikeinterface/core/node_pipeline.py | 2 ++ .../waveforms/denoising/savgol_denoiser.py | 8 ++++---- .../waveforms/temporal_pca.py | 19 ++++++++++--------- 3 files changed, 16 insertions(+), 13 deletions(-) diff --git a/src/spikeinterface/core/node_pipeline.py b/src/spikeinterface/core/node_pipeline.py index 9cbeda7c84..14982cb69b 100644 --- a/src/spikeinterface/core/node_pipeline.py +++ b/src/spikeinterface/core/node_pipeline.py @@ -347,6 +347,7 @@ def __init__( 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): @@ -470,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/sortingcomponents/waveforms/denoising/savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py index a09352720f..402f619855 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py @@ -39,14 +39,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: + waveform_node = find_parent_of_type(parents, WaveformsNode) + if waveform_node 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, + waveform_node.ms_before, + waveform_node.ms_after, return_output=return_output, parents=parents, ) diff --git a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py index 36bb32cda0..959efe060e 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py +++ b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py @@ -28,18 +28,19 @@ def __init__( 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: + 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") super().__init__( recording, - waveform_extractor.ms_before, - waveform_extractor.ms_after, + waveform_node.ms_before, + waveform_node.ms_after, return_output=return_output, parents=parents, ) + self.sparse_waveforms = waveform_node.sparse_waveforms if pca_model is None: self.model_folder_path = model_folder_path @@ -58,18 +59,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"] From de466b622fa43787e1e54d084fd28e7b74af23b5 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 11:06:18 +0200 Subject: [PATCH 11/16] fix: add files --- .../waveforms/denoising/neural_network_denoiser.py | 1 + .../sortingcomponents/waveforms/denoising/savgol_denoiser.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py index 6c8aa999c7..0ac7d9499f 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py @@ -77,6 +77,7 @@ def __init__( return_output=return_output, parents=parents, ) + self.sparse_waveforms = waveform_node.sparse_waveforms 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: diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py index 402f619855..3465f88723 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py @@ -50,7 +50,7 @@ def __init__( return_output=return_output, parents=parents, ) - + self.sparse_waveforms = waveform_node.sparse_waveforms self.order = order waveforms_sampling_frequency = self.recording.get_sampling_frequency() self.window_length = int(window_length_ms * waveforms_sampling_frequency / 1000) From 8535e1e8c5722611fff2df7917afacf5b25d9728 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 11:09:43 +0200 Subject: [PATCH 12/16] fix: propagate sparse_waveforms attrs from base LocalizationNode --- .../sortingcomponents/peak_localization/base.py | 16 +++++++--------- 1 file changed, 7 insertions(+), 9 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/base.py b/src/spikeinterface/sortingcomponents/peak_localization/base.py index 11e98671c2..9c5f66d81d 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/base.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/base.py @@ -24,21 +24,19 @@ def __init__(self, recording, parents, return_output=True, radius_um=75.0): self.channel_distance = get_channel_distances(recording) # Find waveform extractor in the parents - waveform_extractor = find_parent_of_type(self.parents, WaveformsNode) - if waveform_extractor is None: + 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_extractor.nbefore - self.nafter = waveform_extractor.nafter + self.nbefore = waveform_node.nbefore + self.nafter = waveform_node.nafter self.neighbours_mask = self.channel_distance <= radius_um - if isinstance(waveform_extractor, ExtractSparseWaveforms): - self.sparse_waveforms = True + 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_extractor.neighbours_mask + self.extraction_neighbours_mask = waveform_node.neighbours_mask self.neighbours_mask &= self.extraction_neighbours_mask - else: - self.sparse_waveforms = False self._kwargs["radius_um"] = radius_um def get_dtype(self): From 8c6c63df7e32a114ae370fc6ef60365e1261491f Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 11:13:06 +0200 Subject: [PATCH 13/16] fix: propagate sparse_waveforms attrs in denoisers --- .../waveforms/denoising/neural_network_denoiser.py | 4 +++- .../waveforms/denoising/savgol_denoiser.py | 7 +++++-- .../sortingcomponents/waveforms/temporal_pca.py | 4 +++- 3 files changed, 11 insertions(+), 4 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py index 0ac7d9499f..424804bbdd 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py @@ -77,7 +77,6 @@ def __init__( return_output=return_output, parents=parents, ) - self.sparse_waveforms = waveform_node.sparse_waveforms 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: @@ -91,6 +90,9 @@ def __init__( self.assert_model_and_waveform_temporal_match( waveform_node, model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path ) + # Propagate waveforms node parameters + self.sparse_waveforms = waveform_node.sparse_waveforms + self.neighbours_mask = waveform_node.neighbours_mask def assert_model_and_waveform_temporal_match( self, diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py index 3465f88723..96003a2690 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py @@ -50,11 +50,14 @@ def __init__( return_output=return_output, parents=parents, ) - self.sparse_waveforms = waveform_node.sparse_waveforms - self.order = order waveforms_sampling_frequency = self.recording.get_sampling_frequency() + + self.order = order self.window_length = int(window_length_ms * waveforms_sampling_frequency / 1000) self.order = min(self.order, self.window_length - 1) + # Propagate waveforms node parameters + self.sparse_waveforms = waveform_node.sparse_waveforms + self.neighbours_mask = waveform_node.neighbours_mask self._kwargs.update(dict(order=order, window_length_ms=window_length_ms)) def compute(self, traces, peaks, waveforms): diff --git a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py index 959efe060e..08d804f754 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py +++ b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py @@ -40,7 +40,6 @@ def __init__( parents=parents, ) - self.sparse_waveforms = waveform_node.sparse_waveforms if pca_model is None: self.model_folder_path = model_folder_path @@ -62,6 +61,9 @@ def __init__( self.assert_model_and_waveform_temporal_match(waveform_node) else: self.pca_model = pca_model + # Propagate waveforms node parameters + self.sparse_waveforms = waveform_node.sparse_waveforms + self.neighbours_mask = waveform_node.neighbours_mask def assert_model_and_waveform_temporal_match(self, waveform_node: WaveformsNode): """ From 0fa97795b6d98b3485d20876d35e1a87ce404205 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 12:33:31 +0200 Subject: [PATCH 14/16] fix: smaller sparsity in extract than localize --- src/spikeinterface/preprocessing/motion.py | 5 +- .../peak_localization/base.py | 11 ++-- .../peak_localization/grid.py | 10 ++-- .../tests/test_peak_localization.py | 22 ++++++++ .../denoising/neural_network_denoiser.py | 29 ++++------ .../waveforms/denoising/savgol_denoiser.py | 14 ++--- .../denoising/temporal_pca_denoiser.py | 12 ++--- .../waveforms/temporal_pca.py | 54 ++++++++++--------- .../waveforms/tests/test_temporal_pca.py | 7 +-- .../waveforms/waveform_thresholder.py | 9 +--- .../waveforms/waveform_utils.py | 37 +++++++++++++ 11 files changed, 128 insertions(+), 82 deletions(-) diff --git a/src/spikeinterface/preprocessing/motion.py b/src/spikeinterface/preprocessing/motion.py index 2f5844cf7f..4d3e784c84 100644 --- a/src/spikeinterface/preprocessing/motion.py +++ b/src/spikeinterface/preprocessing/motion.py @@ -454,7 +454,10 @@ def compute_motion( 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], **denoise_kwargs_without_method + 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) diff --git a/src/spikeinterface/sortingcomponents/peak_localization/base.py b/src/spikeinterface/sortingcomponents/peak_localization/base.py index 9c5f66d81d..b0398c6300 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/base.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/base.py @@ -5,15 +5,12 @@ from spikeinterface.core.node_pipeline import ( find_parent_of_type, WaveformsNode, - ExtractDenseWaveforms, - ExtractSparseWaveforms, ) from spikeinterface.core import get_channel_distances from spikeinterface.sortingcomponents import waveforms -# TODO: make this sparse and pre-instantiate neighbor mask in case of ExtractSparseWaveforms class LocalizeBase(PipelineNode): def __init__(self, recording, parents, return_output=True, radius_um=75.0): @@ -49,7 +46,13 @@ def get_sparse_waveform(self, waveform, chan_inds, main_chan): # 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) + 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: diff --git a/src/spikeinterface/sortingcomponents/peak_localization/grid.py b/src/spikeinterface/sortingcomponents/peak_localization/grid.py index 697319775b..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 @@ -124,6 +117,9 @@ 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) 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 9674b01b9c..f87bd8657d 100644 --- a/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py +++ b/src/spikeinterface/sortingcomponents/peak_localization/tests/test_peak_localization.py @@ -122,6 +122,28 @@ def test_localize_peaks_sparse(peaks_and_recording, method): 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 diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py index 424804bbdd..c088e9ccdb 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/neural_network_denoiser.py @@ -21,11 +21,11 @@ HAVE_HUGGINGFACE = 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 +from spikeinterface.core.node_pipeline import PipelineNode +from ..waveform_utils import to_temporal_representation, from_temporal_representation, WaveformTransformer -class SingleChannelDenoiser(WaveformsNode): +class SingleChannelDenoiser(WaveformTransformer): """ Denoiser for temporal dimension of waveforms. It takes as input a WaveformsNode and outputs denoised waveforms. @@ -66,14 +66,9 @@ def __init__( device=None, ): assert HAVE_TORCH, "To use the SingleChannelDenoiser you need to install torch" - waveform_node = find_parent_of_type(parents, WaveformsNode) - if waveform_node is None: - raise TypeError(f"Model should have a {WaveformsNode.__name__} in its parents") super().__init__( recording, - waveform_node.ms_before, - waveform_node.ms_after, return_output=return_output, parents=parents, ) @@ -81,22 +76,18 @@ def __init__( 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 = waveform_node.nbefore + waveform_node.nafter + 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( - waveform_node, model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path + model_folder=model_folder, repo_id=repo_id, model_relative_path=model_relative_path ) - # Propagate waveforms node parameters - self.sparse_waveforms = waveform_node.sparse_waveforms - self.neighbours_mask = waveform_node.neighbours_mask def assert_model_and_waveform_temporal_match( self, - waveform_node: WaveformsNode, model_relative_path: str, model_folder: Optional[str] = None, repo_id: Optional[str] = None, @@ -105,9 +96,9 @@ def assert_model_and_waveform_temporal_match( Asserts that the model and the waveform extractor have the same temporal parameters """ # Extract temporal parameters from the waveform extractor - waveforms_ms_before = waveform_node.ms_before - waveforms_ms_after = waveform_node.ms_after - waveforms_sampling_frequency = waveform_node.recording.sampling_frequency + 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: @@ -134,7 +125,7 @@ def assert_model_and_waveform_temporal_match( model_nbefore = model_info.get("nbefore") if model_num_samples is not None: - if model_num_samples != waveform_node.nbefore + waveform_node.nafter: + 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}" ) @@ -154,7 +145,7 @@ def assert_model_and_waveform_temporal_match( 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 - waveform_node.nbefore) > 5: + 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" ) diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/savgol_denoiser.py index 96003a2690..3a871ef14e 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/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 @@ -39,14 +40,8 @@ def __init__( order: int = 3, window_length_ms: float = 0.25, ): - waveform_node = find_parent_of_type(parents, WaveformsNode) - if waveform_node is None: - raise TypeError(f"SavGolDenoiser should have a single {WaveformsNode.__name__} in its parents") - super().__init__( recording, - waveform_node.ms_before, - waveform_node.ms_after, return_output=return_output, parents=parents, ) @@ -55,9 +50,6 @@ def __init__( self.order = order self.window_length = int(window_length_ms * waveforms_sampling_frequency / 1000) self.order = min(self.order, self.window_length - 1) - # Propagate waveforms node parameters - self.sparse_waveforms = waveform_node.sparse_waveforms - self.neighbours_mask = waveform_node.neighbours_mask self._kwargs.update(dict(order=order, window_length_ms=window_length_ms)) def compute(self, traces, peaks, waveforms): diff --git a/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py b/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py index fca01b60d1..5f3e2907db 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py +++ b/src/spikeinterface/sortingcomponents/waveforms/denoising/temporal_pca_denoiser.py @@ -3,10 +3,11 @@ from spikeinterface.core import BaseRecording from spikeinterface.core.node_pipeline import PipelineNode -from ..temporal_pca import TemporalPCBaseNode, to_temporal_representation, from_temporal_representation +from ..waveform_utils import WaveformTransformer +from ..temporal_pca import TemporalPCMixin, to_temporal_representation, from_temporal_representation -class TemporalPCADenoiser(TemporalPCBaseNode): +class TemporalPCADenoiser(WaveformTransformer, TemporalPCMixin): """ A step that performs a PCA denoising on the waveforms extracted by a peak_detection function. @@ -44,11 +45,10 @@ def __init__( model_folder_path=None, return_output=True, ): - TemporalPCBaseNode.__init__( + WaveformTransformer.__init__(self, recording=recording, parents=parents, return_output=return_output) + TemporalPCMixin.__init__( self, - recording=recording, - parents=parents, - return_output=return_output, + waveform_node=self.waveforms_node, pca_model=pca_model, model_folder_path=model_folder_path, ) diff --git a/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py b/src/spikeinterface/sortingcomponents/waveforms/temporal_pca.py index 08d804f754..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_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") - - super().__init__( - recording, - waveform_node.ms_before, - waveform_node.ms_after, - return_output=return_output, - parents=parents, - ) - if pca_model is None: self.model_folder_path = model_folder_path @@ -61,9 +47,6 @@ def __init__( self.assert_model_and_waveform_temporal_match(waveform_node) else: self.pca_model = pca_model - # Propagate waveforms node parameters - self.sparse_waveforms = waveform_node.sparse_waveforms - self.neighbours_mask = waveform_node.neighbours_mask def assert_model_and_waveform_temporal_match(self, waveform_node: WaveformsNode): """ @@ -173,10 +156,33 @@ 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 TemporalPCABaseNode(PipelineNode, TemporalPCMixin): + """ + Base class for temporal PCA projection node + """ + + def __init__( + self, + recording: BaseRecording, + parents: List[PipelineNode], + pca_model=None, + model_folder_path=None, + return_output=True, + ): + 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.recording = recording -class TemporalPCAProjection(TemporalPCBaseNode): +class TemporalPCAProjection(TemporalPCABaseNode): """ A step that performs a PCA projection on the waveforms extracted by a waveforms parent node. @@ -208,7 +214,7 @@ def __init__( dtype="float32", return_output=True, ): - TemporalPCBaseNode.__init__( + TemporalPCABaseNode.__init__( self, recording=recording, parents=parents, @@ -249,7 +255,7 @@ def compute(self, traces: np.ndarray, peaks: np.ndarray, waveforms: np.ndarray) 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. @@ -285,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/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..14d79a0192 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py +++ b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py @@ -1,3 +1,40 @@ +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") + + super().__init__( + recording, + waveforms_node.ms_before, + waveforms_node.ms_after, + 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 From 93cde6148f1566817dcee6c06c3a7c5739e64f4f Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 12:36:05 +0200 Subject: [PATCH 15/16] fix: propagate args to WaveformsNode from WaveformsTransformer --- .../sortingcomponents/waveforms/waveform_utils.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py index 14d79a0192..d7b23bfff8 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py +++ b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py @@ -24,8 +24,10 @@ def __init__(self, recording: BaseRecording, return_output: bool = True, parents super().__init__( recording, - waveforms_node.ms_before, - waveforms_node.ms_after, + ms_before=waveforms_node.ms_before, + ms_after=waveforms_node.ms_after, + nbefore=waveforms_node.nbefore, + nafter=waveforms_node.nafter, return_output=return_output, parents=parents, ) From d96fc55e7fc0f6ad6a93d530c2cd2bddca0da8eb Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Wed, 22 Jul 2026 12:40:24 +0200 Subject: [PATCH 16/16] fix: propagate args to WaveformsNode from WaveformsTransformer 2 --- .../sortingcomponents/waveforms/waveform_utils.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py index d7b23bfff8..3eb0b2e1b1 100644 --- a/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py +++ b/src/spikeinterface/sortingcomponents/waveforms/waveform_utils.py @@ -22,10 +22,9 @@ def __init__(self, recording: BaseRecording, return_output: bool = True, parents 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, - ms_before=waveforms_node.ms_before, - ms_after=waveforms_node.ms_after, nbefore=waveforms_node.nbefore, nafter=waveforms_node.nafter, return_output=return_output,