Skip to content

Commit ae19c2a

Browse files
authored
Adding Hanning filtering for waveforms pipeline node
1 parent 8a3df4f commit ae19c2a

3 files changed

Lines changed: 95 additions & 3 deletions

File tree

src/spikeinterface/sortingcomponents/clustering/circus.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from spikeinterface.core.recording_tools import get_noise_levels, get_channel_distances
2121
from spikeinterface.sortingcomponents.peak_selection import select_peaks
2222
from spikeinterface.sortingcomponents.waveforms.temporal_pca import TemporalPCAProjection
23+
from spikeinterface.sortingcomponents.waveforms.hanning_filter import HanningFilter
2324
from spikeinterface.core.template import Templates
2425
from spikeinterface.core.sparsity import compute_sparsity
2526
from spikeinterface.sortingcomponents.tools import remove_empty_templates
@@ -101,6 +102,12 @@ def main_function(cls, recording, peaks, params, job_kwargs=dict()):
101102
valid = np.argmax(np.abs(wfs), axis=1) == nbefore
102103
wfs = wfs[valid]
103104

105+
# Perform Hanning filtering
106+
hanning_before = np.hanning(2 * nbefore)
107+
hanning_after = np.hanning(2 * nafter)
108+
hanning = np.concatenate((hanning_before[:nbefore], hanning_after[nafter:]))
109+
wfs *= hanning
110+
104111
from sklearn.decomposition import TruncatedSVD
105112

106113
tsvd = TruncatedSVD(params["n_svd"][0])
@@ -134,11 +141,13 @@ def main_function(cls, recording, peaks, params, job_kwargs=dict()):
134141
radius_um=radius_um,
135142
)
136143

137-
node2 = TemporalPCAProjection(
138-
recording, parents=[node0, node1], return_output=True, model_folder_path=model_folder
144+
node2 = HanningFilter(recording, parents=[node0, node1], return_output=False)
145+
146+
node3 = TemporalPCAProjection(
147+
recording, parents=[node0, node2], return_output=True, model_folder_path=model_folder
139148
)
140149

141-
pipeline_nodes = [node0, node1, node2]
150+
pipeline_nodes = [node0, node1, node2, node3]
142151

143152
if len(params["recursive_kwargs"]) == 0:
144153
from sklearn.decomposition import PCA
Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
import pytest
2+
3+
4+
from spikeinterface.sortingcomponents.waveforms.hanning_filter import HanningFilter
5+
6+
from spikeinterface.core.node_pipeline import (
7+
PeakRetriever,
8+
ExtractDenseWaveforms,
9+
run_node_pipeline,
10+
)
11+
12+
13+
def test_hanning_filter(generated_recording, detected_peaks, chunk_executor_kwargs):
14+
recording = generated_recording
15+
peaks = detected_peaks
16+
17+
# Parameters
18+
ms_before = 1.0
19+
ms_after = 1.0
20+
21+
# Node initialization
22+
peak_retriever = PeakRetriever(recording, peaks)
23+
24+
extract_waveforms = ExtractDenseWaveforms(
25+
recording=recording, parents=[peak_retriever], ms_before=ms_before, ms_after=ms_after, return_output=True
26+
)
27+
28+
hanning_filter = HanningFilter(recording=recording, parents=[peak_retriever, extract_waveforms])
29+
pipeline_nodes = [peak_retriever, extract_waveforms, hanning_filter]
30+
31+
# Extract projected waveforms and compare
32+
waveforms, denoised_waveforms = run_node_pipeline(recording, nodes=pipeline_nodes, job_kwargs=chunk_executor_kwargs)
33+
assert waveforms.shape == denoised_waveforms.shape
Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,50 @@
1+
from __future__ import annotations
2+
3+
4+
from typing import List, Optional
5+
import numpy as np
6+
from spikeinterface.core import BaseRecording
7+
from spikeinterface.core.node_pipeline import PipelineNode, WaveformsNode, find_parent_of_type
8+
9+
10+
class HanningFilter(WaveformsNode):
11+
"""
12+
Hanning Filtering to remove border effects while extracting waveforms
13+
14+
Parameters
15+
----------
16+
recording: BaseRecording
17+
The recording extractor object
18+
return_output: bool, default: True
19+
Whether to return output from this node
20+
parents: list of PipelineNodes, default: None
21+
The parent nodes of this node
22+
"""
23+
24+
def __init__(
25+
self,
26+
recording: BaseRecording,
27+
return_output: bool = True,
28+
parents: Optional[List[PipelineNode]] = None,
29+
):
30+
waveform_extractor = find_parent_of_type(parents, WaveformsNode)
31+
if waveform_extractor is None:
32+
raise TypeError(f"HanningFilter should have a single {WaveformsNode.__name__} in its parents")
33+
34+
super().__init__(
35+
recording,
36+
waveform_extractor.ms_before,
37+
waveform_extractor.ms_after,
38+
return_output=return_output,
39+
parents=parents,
40+
)
41+
42+
hanning_before = np.hanning(2 * self.nbefore)
43+
hanning_after = np.hanning(2 * self.nafter)
44+
hanning = np.concatenate((hanning_before[: self.nbefore], hanning_after[self.nafter :]))
45+
self.hanning = hanning[:, None]
46+
self._kwargs.update(dict())
47+
48+
def compute(self, traces, peaks, waveforms):
49+
denoised_waveforms = waveforms * self.hanning
50+
return denoised_waveforms

0 commit comments

Comments
 (0)