Conversation
h-mayorquin
left a comment
There was a problem hiding this comment.
I think the idea works but there are two points. First, this changes the output a bit and second, the stability check is a bit brittle.
First, why it changes values. Float32 moves each filtered value by a tiny amount (below 0.001 int16 units for the default band) and when a value sits at .5 that is enough to round to a different int16. So some samples differ by 1 from main. It should not have any consequences but the output is not the same as before and I think it is worth flagging.
Second, why it is brittle. sosfilt_zi only throws an error when the rounding to float32 lands exactly on sum(a) == 0. A filter that is close to unstable but misses that exact value goes through to float32.
I ran the check in this PR on the bandpass_filter filter at 30 kHz with freq_max=300 and a few values of freq_min, on synthetic int16 traces and against the same filter in float64:
| band | what this PR does | max error / std | distance of the largest pole from 1 |
|---|---|---|---|
| 0.1-300 Hz | float32 | 2.6e+16 | 6.5e-06 |
| 0.5-300 Hz | float32 | 14 | 3.2e-05 |
| 1-300 Hz | float64 fallback | 0 | 6.4e-05 |
| 2-300 Hz | float32 | 0.5 | 1.3e-04 |
| 5-300 Hz | float32 | 0.053 | 3.1e-04 |
| 10-300 Hz | float32 | 0.046 | 6.1e-04 |
| 300-6000 Hz | float32 | 2.5e-05 | 1.8e-02 |
1 Hz falls back but 0.1 Hz blows up and 0.5 Hz gives an error 14 times the std of the signal. These need ignore_low_freq_error=True but that is what we tell people to do in the LFP how-to.
I think the check should be on how far the poles of the float64 filter are from the unit circle (the last column):
if filter_mode == "sos" and np.dtype(dtype) in (np.float32, np.int16):
# float32 rounding moves the poles, so only use it when they are far from the unit circle
max_pole_radius = max(np.abs(np.roots(section[3:])).max() for section in coeff)
if 1.0 - max_pole_radius > 1e-2:
coeff = np.asarray(coeff, dtype="float32")With 1e-2 the default band stays in float32 and all the low cutoffs above go to float64.
Now, for something completly different. From my memory the last time I looked into this msot of the memory use comes from the scipy filtering routines themselves. Actually, now I checked and the forward-backward ones (sosfiltfilt, filtfilt) hold three copies of the input at their peak: the padded input plus a float copy for each pass. I think that the best performance improvement here is to actually chunk across channels. Internally the scipy routine loops across channels anyway (in Cython) so we are trading a reduction in peak memory for a python loop over the blocks. Note that the filtering inside the loop is not done in python (that is what usually makes a loop slow), only the loop itself is and that costs microseconds. I filtered a 10 s chunk of 384 int16 channels at 30 kHz with the default bandpass, calling sosfiltfilt on blocks of channels and writing each block into a preallocated int16 output (the 384 row is the code path of main and this PR). Peak memory is measured with tracemalloc on top of the input chunk, the time is the best of three runs and the difference is in int16 units against main:
| coefficients | channels per block | peak memory | time | max difference vs main |
|---|---|---|---|---|
| float64 | 384 (main) | 1,978 MiB | 2.31 s | 0 |
| float64 | 64 | 696 MiB | 2.36 s | 0 |
| float64 | 32 | 458 MiB | 2.23 s | 0 |
| float64 | 8 | 279 MiB | 2.00 s | 0 |
| float32 | 384 (this PR) | 1,099 MiB | 1.98 s | 1 |
| float32 | 64 | 476 MiB | 2.13 s | 1 |
| float32 | 32 | 348 MiB | 1.97 s | 1 |
| float32 | 8 | 252 MiB | 1.85 s | 1 |
Of the things above the first is a request (float32 with the pole check) and the second is a suggestion. If you agree we can do the channel blocks here or I can take care of this on a separate PR.
The two changes combined reduce the peak memory of filtering a 10 s chunk of 384 channels by 82% (1,978 to 348 MiB), whereas float32 alone reduces it by 44% and the channel blocks alone by 77%. What do you think?
This came up while testing memory usage of shards/chunks in various operations. The memory allcoation for a standard filter + CMR seemed bloated.
A simple way to reduce this by a lot is to make fitler coefficients float32 (when the input data is not float64). In this case, the filtered traces will be float32 instrad of 64, effectively halving the peak RAM usage.
Peak memory of
bandpass_filteron one 10 s block (384 channels, int16 output, SI's default5th-order bandpass 300–6000 Hz), real Neuropixels 1.0 data:
What this means for Zarr
What it means with writing to zarr v3 (#4260) in shards, bandpass + CMR → float32 zarr, WavPack, 8 workers:
So we can target larger shards for most operations.
E.g. 20s shards on a 16CPU / 64GB machine would peak at ~47 GB RAM