2424
2525class FilterRecording (BasePreprocessor ):
2626 """
27- Generic filter class based on:
28-
29- * scipy.signal.iirfilter
30- * scipy.signal.filtfilt or scipy.signal.sosfilt
27+ A generic filter class based on:
28+ For filter coefficient generation:
29+ * scipy.signal.iirfilter
30+ For filter application:
31+ * scipy.signal.filtfilt or scipy.signal.sosfiltfilt when direction = "forward-backward"
32+ * scipy.signal.lfilter or scipy.signal.sosfilt when direction = "forward" or "backward"
3133
3234 BandpassFilterRecording is built on top of it.
3335
@@ -56,6 +58,11 @@ class FilterRecording(BasePreprocessor):
5658 - numerator/denominator : ("ba")
5759 ftype : str, default: "butter"
5860 Filter type for `scipy.signal.iirfilter` e.g. "butter", "cheby1".
61+ direction : "forward" | "backward" | "forward-backward", default: "forward-backward"
62+ Direction of filtering:
63+ - "forward" - filter is applied to the timeseries in one direction, creating phase shifts
64+ - "backward" - the timeseries is reversed, the filter is applied and filtered timeseries reversed again. Creates phase shifts in the opposite direction to "forward"
65+ - "forward-backward" - Applies the filter in the forward and backward direction, resulting in zero-phase filtering. Note this doubles the effective filter order.
5966
6067 Returns
6168 -------
@@ -75,6 +82,7 @@ def __init__(
7582 add_reflect_padding = False ,
7683 coeff = None ,
7784 dtype = None ,
85+ direction = "forward-backward" ,
7886 ):
7987 import scipy .signal
8088
@@ -106,7 +114,13 @@ def __init__(
106114 for parent_segment in recording ._recording_segments :
107115 self .add_recording_segment (
108116 FilterRecordingSegment (
109- parent_segment , filter_coeff , filter_mode , margin , dtype , add_reflect_padding = add_reflect_padding
117+ parent_segment ,
118+ filter_coeff ,
119+ filter_mode ,
120+ margin ,
121+ dtype ,
122+ add_reflect_padding = add_reflect_padding ,
123+ direction = direction ,
110124 )
111125 )
112126
@@ -121,14 +135,25 @@ def __init__(
121135 margin_ms = margin_ms ,
122136 add_reflect_padding = add_reflect_padding ,
123137 dtype = dtype .str ,
138+ direction = direction ,
124139 )
125140
126141
127142class FilterRecordingSegment (BasePreprocessorSegment ):
128- def __init__ (self , parent_recording_segment , coeff , filter_mode , margin , dtype , add_reflect_padding = False ):
143+ def __init__ (
144+ self ,
145+ parent_recording_segment ,
146+ coeff ,
147+ filter_mode ,
148+ margin ,
149+ dtype ,
150+ add_reflect_padding = False ,
151+ direction = "forward-backward" ,
152+ ):
129153 BasePreprocessorSegment .__init__ (self , parent_recording_segment )
130154 self .coeff = coeff
131155 self .filter_mode = filter_mode
156+ self .direction = direction
132157 self .margin = margin
133158 self .add_reflect_padding = add_reflect_padding
134159 self .dtype = dtype
@@ -150,11 +175,24 @@ def get_traces(self, start_frame, end_frame, channel_indices):
150175
151176 import scipy .signal
152177
153- if self .filter_mode == "sos" :
154- filtered_traces = scipy .signal .sosfiltfilt (self .coeff , traces_chunk , axis = 0 )
155- elif self .filter_mode == "ba" :
156- b , a = self .coeff
157- filtered_traces = scipy .signal .filtfilt (b , a , traces_chunk , axis = 0 )
178+ if self .direction == "forward-backward" :
179+ if self .filter_mode == "sos" :
180+ filtered_traces = scipy .signal .sosfiltfilt (self .coeff , traces_chunk , axis = 0 )
181+ elif self .filter_mode == "ba" :
182+ b , a = self .coeff
183+ filtered_traces = scipy .signal .filtfilt (b , a , traces_chunk , axis = 0 )
184+ else :
185+ if self .direction == "backward" :
186+ traces_chunk = np .flip (traces_chunk , axis = 0 )
187+
188+ if self .filter_mode == "sos" :
189+ filtered_traces = scipy .signal .sosfilt (self .coeff , traces_chunk , axis = 0 )
190+ elif self .filter_mode == "ba" :
191+ b , a = self .coeff
192+ filtered_traces = scipy .signal .lfilter (b , a , traces_chunk , axis = 0 )
193+
194+ if self .direction == "backward" :
195+ filtered_traces = np .flip (filtered_traces , axis = 0 )
158196
159197 if right_margin > 0 :
160198 filtered_traces = filtered_traces [left_margin :- right_margin , :]
@@ -289,6 +327,73 @@ def __init__(self, recording, freq=3000, q=30, margin_ms=5.0, dtype=None):
289327notch_filter = define_function_from_class (source_class = NotchFilterRecording , name = "notch_filter" )
290328highpass_filter = define_function_from_class (source_class = HighpassFilterRecording , name = "highpass_filter" )
291329
330+
331+ def causal_filter (
332+ recording ,
333+ direction = "forward" ,
334+ band = [300.0 , 6000.0 ],
335+ btype = "bandpass" ,
336+ filter_order = 5 ,
337+ ftype = "butter" ,
338+ filter_mode = "sos" ,
339+ margin_ms = 5.0 ,
340+ add_reflect_padding = False ,
341+ coeff = None ,
342+ dtype = None ,
343+ ):
344+ """
345+ Generic causal filter built on top of the filter function.
346+
347+ Parameters
348+ ----------
349+ recording : Recording
350+ The recording extractor to be re-referenced
351+ direction : "forward" | "backward", default: "forward"
352+ Direction of causal filter. The "backward" option flips the traces in time before applying the filter
353+ and then flips them back.
354+ band : float or list, default: [300.0, 6000.0]
355+ If float, cutoff frequency in Hz for "highpass" filter type
356+ If list. band (low, high) in Hz for "bandpass" filter type
357+ btype : "bandpass" | "highpass", default: "bandpass"
358+ Type of the filter
359+ margin_ms : float, default: 5.0
360+ Margin in ms on border to avoid border effect
361+ coeff : array | None, default: None
362+ Filter coefficients in the filter_mode form.
363+ dtype : dtype or None, default: None
364+ The dtype of the returned traces. If None, the dtype of the parent recording is used
365+ add_reflect_padding : Bool, default False
366+ If True, uses a left and right margin during calculation.
367+ filter_order : order
368+ The order of the filter for `scipy.signal.iirfilter`
369+ filter_mode : "sos" | "ba", default: "sos"
370+ Filter form of the filter coefficients for `scipy.signal.iirfilter`:
371+ - second-order sections ("sos")
372+ - numerator/denominator : ("ba")
373+ ftype : str, default: "butter"
374+ Filter type for `scipy.signal.iirfilter` e.g. "butter", "cheby1".
375+
376+ Returns
377+ -------
378+ filter_recording : FilterRecording
379+ The causal-filtered recording extractor object
380+ """
381+ assert direction in ["forward" , "backward" ], "Direction must be either 'forward' or 'backward'"
382+ return filter (
383+ recording = recording ,
384+ direction = direction ,
385+ band = band ,
386+ btype = btype ,
387+ filter_order = filter_order ,
388+ ftype = ftype ,
389+ filter_mode = filter_mode ,
390+ margin_ms = margin_ms ,
391+ add_reflect_padding = add_reflect_padding ,
392+ coeff = coeff ,
393+ dtype = dtype ,
394+ )
395+
396+
292397bandpass_filter .__doc__ = bandpass_filter .__doc__ .format (_common_filter_docs )
293398highpass_filter .__doc__ = highpass_filter .__doc__ .format (_common_filter_docs )
294399
0 commit comments