Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions src/spikeinterface/preprocessing/filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,21 @@ def __init__(
direction="forward-backward",
):
BasePreprocessorSegment.__init__(self, parent_recording_segment)
# scipy computes in the precision of the coefficients, so float64 sos coefficients make
# it filter in float64 (twice the memory of float32) whatever the output dtype. For these
# outputs float32 sos coefficients are accurate enough (error far below int16
# quantization). The "ba" form is numerically fragile, so it stays float64.
# Very low cutoff frequencies (e.g. 0.5 Hz at 28 kHz) can make float32 sos coefficients
# numerically unstable (sum(a)==0 → pole at z=1). Validate before committing.
if filter_mode == "sos" and np.dtype(dtype) in (np.float32, np.int16):
coeff_f32 = np.asarray(coeff, dtype="float32")
from scipy.signal import sosfilt_zi

try:
sosfilt_zi(coeff_f32)
coeff = coeff_f32
except ValueError:
pass # keep float64 — filter is numerically unstable at float32 precision
self.coeff = coeff
self.filter_mode = filter_mode
self.direction = direction
Expand Down
21 changes: 21 additions & 0 deletions src/spikeinterface/preprocessing/tests/test_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,27 @@ def test_filter_opencl():
# plt.show()


def test_filter_float32_coefficients():
# float32 and <=16-bit integer outputs filter with float32 sos coefficients (half the memory
# of float64); the result stays far below int16 quantization from the float64 computation
recording = generate_recording(durations=[1.0], num_channels=4)
rec32 = bandpass_filter(recording, freq_min=300.0, freq_max=6000.0, dtype="float32")
rec64 = bandpass_filter(recording, freq_min=300.0, freq_max=6000.0, dtype="float64")
recording_int16 = recording.astype("int16")
rec16 = bandpass_filter(recording_int16, freq_min=300.0, freq_max=6000.0, dtype="int16")
rec_i32 = bandpass_filter(recording_int16, freq_min=300.0, freq_max=6000.0, dtype="int32")
assert rec32._recording_segments[0].coeff.dtype == np.float32
assert rec16._recording_segments[0].coeff.dtype == np.float32
assert rec64._recording_segments[0].coeff.dtype == np.float64
assert rec_i32._recording_segments[0].coeff.dtype == np.float64
traces32, traces64 = rec32.get_traces(), rec64.get_traces()
assert traces32.dtype == np.float32
assert np.max(np.abs(traces32 - traces64)) < 1e-3 * np.std(traces64)
# integer output: rounding can only flip by one unit, for values sitting at .5
traces16, traces_i32 = rec16.get_traces().astype(int), rec_i32.get_traces().astype(int)
assert np.max(np.abs(traces16 - traces_i32)) <= 1


if __name__ == "__main__":
import tempfile
from pathlib import Path
Expand Down
Loading