From f169568fb793ee15fe1471d83e91eaa208a6c9d9 Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Fri, 2 Oct 2026 10:36:01 +0200 Subject: [PATCH 1/2] perf: filter in float32 if data is not float64 --- src/spikeinterface/preprocessing/filter.py | 6 ++++++ .../preprocessing/tests/test_filter.py | 21 +++++++++++++++++++ 2 files changed, 27 insertions(+) diff --git a/src/spikeinterface/preprocessing/filter.py b/src/spikeinterface/preprocessing/filter.py index 2cf3e47c93..b7b354dce1 100644 --- a/src/spikeinterface/preprocessing/filter.py +++ b/src/spikeinterface/preprocessing/filter.py @@ -170,6 +170,12 @@ 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. + if filter_mode == "sos" and np.dtype(dtype) in (np.float32, np.int16): + coeff = np.asarray(coeff, dtype="float32") self.coeff = coeff self.filter_mode = filter_mode self.direction = direction diff --git a/src/spikeinterface/preprocessing/tests/test_filter.py b/src/spikeinterface/preprocessing/tests/test_filter.py index 3b47d0eb1a..cf226013fd 100644 --- a/src/spikeinterface/preprocessing/tests/test_filter.py +++ b/src/spikeinterface/preprocessing/tests/test_filter.py @@ -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 From 08cd3e5bb221ee888c54920ffa0079a1bebae7ee Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Fri, 2 Oct 2026 11:48:37 +0200 Subject: [PATCH 2/2] fix: test filter stability and fallback to float64 if unstable --- src/spikeinterface/preprocessing/filter.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/src/spikeinterface/preprocessing/filter.py b/src/spikeinterface/preprocessing/filter.py index b7b354dce1..0f12865003 100644 --- a/src/spikeinterface/preprocessing/filter.py +++ b/src/spikeinterface/preprocessing/filter.py @@ -174,8 +174,17 @@ def __init__( # 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 = np.asarray(coeff, dtype="float32") + 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