diff --git a/.github/workflows/run_unix.yml b/.github/workflows/run_unix.yml index 557001dae..8a9c6b201 100644 --- a/.github/workflows/run_unix.yml +++ b/.github/workflows/run_unix.yml @@ -354,6 +354,8 @@ jobs: run: sudo -H python3 $GITHUB_WORKSPACE/python_package/examples/tests/transforms.py - name: Downsampling Python run: sudo -H python3 $GITHUB_WORKSPACE/python_package/examples/tests/downsampling.py + - name: Wavelet Buffer Size Python + run: sudo -H python3 $GITHUB_WORKSPACE/python_package/examples/tests/wavelet_buffer_size.py - name: ICA Python run: sudo -H python3 $GITHUB_WORKSPACE/python_package/examples/tests/ica.py - name: CSP Python diff --git a/.github/workflows/run_windows.yml b/.github/workflows/run_windows.yml index 244c3f0e5..537773653 100644 --- a/.github/workflows/run_windows.yml +++ b/.github/workflows/run_windows.yml @@ -235,6 +235,9 @@ jobs: - name: Downsampling Python Test run: python %GITHUB_WORKSPACE%\python_package\examples\tests\downsampling.py shell: cmd + - name: Wavelet Buffer Size Python Test + run: python %GITHUB_WORKSPACE%\python_package\examples\tests\wavelet_buffer_size.py + shell: cmd - name: CSP Python Test run: python %GITHUB_WORKSPACE%\python_package\examples\tests\csp.py shell: cmd diff --git a/csharp_package/brainflow/brainflow/data_filter.cs b/csharp_package/brainflow/brainflow/data_filter.cs index b8f14f429..fcd312899 100644 --- a/csharp_package/brainflow/brainflow/data_filter.cs +++ b/csharp_package/brainflow/brainflow/data_filter.cs @@ -351,7 +351,7 @@ public static double get_heart_rate (double[] ppg_ir, double[] ppg_red, int samp /// tuple of wavelet coeffs in format [A(J) D(J) D(J-1) ..... D(1)] where J is decomposition level, A - app coeffs, D - detailed coeffs, and array with lengths for each block public static Tuple perform_wavelet_transform (double[] data, int wavelet, int decomposition_level, int extension) { - double[] wavelet_coeffs = new double[data.Length + 2 * (40 + 1)]; + double[] wavelet_coeffs = new double[data.Length + 2 * decomposition_level * (40 + 1)]; int[] lengths = new int[decomposition_level + 1]; int res = DataHandlerLibrary.perform_wavelet_transform (data, data.Length, wavelet, decomposition_level, extension, wavelet_coeffs, lengths); if (res != (int)BrainFlowExitCodes.STATUS_OK) @@ -1115,7 +1115,7 @@ public static unsafe double get_railed_percentage (double[,] data, int row_num, /// tuple of wavelet coeffs in format [A(J) D(J) D(J-1) ..... D(1)] where J is decomposition level, A - app coeffs, D - detailed coeffs, and array with lengths for each block public static unsafe Tuple perform_wavelet_transform (double[,] data, int row_num, int wavelet, int decomposition_level, int extension) { - double[] wavelet_coeffs = new double[data.Length + 2 * (40 + 1)]; + double[] wavelet_coeffs = new double[data.GetLength (1) + 2 * decomposition_level * (40 + 1)]; int[] lengths = new int[decomposition_level + 1]; int res = (int)BrainFlowExitCodes.STATUS_OK; if ((row_num < 0) || (row_num >= data.GetLength (0))) diff --git a/julia_package/brainflow/src/data_filter.jl b/julia_package/brainflow/src/data_filter.jl index 3c5e6ad42..70b7efa00 100644 --- a/julia_package/brainflow/src/data_filter.jl +++ b/julia_package/brainflow/src/data_filter.jl @@ -297,7 +297,7 @@ end end @brainflow_rethrow function perform_wavelet_transform(data, wavelet::WaveletType, decomposition_level::Integer, extension::WaveletExtensionType) - wavelet_coeffs = Vector{Float64}(undef, length(data) + 2 * (40 + 1)) + wavelet_coeffs = Vector{Float64}(undef, length(data) + 2 * decomposition_level * (40 + 1)) lengths = Vector{Cint}(undef, decomposition_level + 1) ccall((:perform_wavelet_transform, DATA_HANDLER_INTERFACE), Cint, (Ptr{Float64}, Cint, Cint, Cint, Cint, Ptr{Float64}, Ptr{Cint}), data, length(data), Int32(wavelet), Int32(decomposition_level), Int32(extension), wavelet_coeffs, lengths) diff --git a/matlab_package/brainflow/DataFilter.m b/matlab_package/brainflow/DataFilter.m index e5764ed1b..aefa98f39 100644 --- a/matlab_package/brainflow/DataFilter.m +++ b/matlab_package/brainflow/DataFilter.m @@ -186,7 +186,7 @@ function disable_data_logger() task_name = 'perform_wavelet_transform'; temp_input = libpointer('doublePtr', data); lib_name = DataFilter.load_lib(); - temp_output = libpointer('doublePtr', zeros(1, int32(size(data, 2) + 2 *(40 + 1)))); + temp_output = libpointer('doublePtr', zeros(1, int32(size(data, 2) + 2 * decomposition_level * (40 + 1)))); lenghts = libpointer('int32Ptr', zeros(1, decomposition_level + 1)); exit_code = calllib(lib_name, task_name, temp_input, size(data, 2), wavelet, decomposition_level, extension, temp_output, lenghts); DataFilter.check_ec(exit_code, task_name); diff --git a/nodejs_package/brainflow/data_filter.ts b/nodejs_package/brainflow/data_filter.ts index 66ea6a94d..b032298f4 100644 --- a/nodejs_package/brainflow/data_filter.ts +++ b/nodejs_package/brainflow/data_filter.ts @@ -491,7 +491,7 @@ export class DataFilter public static performWaveletTransform(data: number[], wavelet: WaveletTypes, decompositionLevel: number, extension: WaveletExtensionTypes): [number[], number[]] { - const waveletCoeffs = [...new Array (data.length + 2 * (40 + 1)).fill(0)]; + const waveletCoeffs = [...new Array (data.length + 2 * decompositionLevel * (40 + 1)).fill(0)]; const lengths = [...new Array (decompositionLevel + 1).fill(0)]; const res = DataHandlerDLL.getInstance().performWaveletTransform( data, data.length, wavelet, decompositionLevel, extension, waveletCoeffs, lengths); diff --git a/python_package/brainflow/data_filter.py b/python_package/brainflow/data_filter.py index d98a3dda2..3ea170406 100644 --- a/python_package/brainflow/data_filter.py +++ b/python_package/brainflow/data_filter.py @@ -878,7 +878,7 @@ def perform_wavelet_transform(cls, data, wavelet: int, decomposition_level: int, """ check_memory_layout_row_major(data, 1) - wavelet_coeffs = numpy.zeros(data.shape[0] + 2 * (40 + 1)).astype(numpy.float64) + wavelet_coeffs = numpy.zeros(data.shape[0] + 2 * decomposition_level * (40 + 1)).astype(numpy.float64) lengths = numpy.zeros(decomposition_level + 1).astype(numpy.int32) res = DataHandlerDLL.get_instance().perform_wavelet_transform(data, data.shape[0], wavelet, decomposition_level, extension_type, diff --git a/python_package/examples/tests/wavelet_buffer_size.py b/python_package/examples/tests/wavelet_buffer_size.py new file mode 100644 index 000000000..3891a5095 --- /dev/null +++ b/python_package/examples/tests/wavelet_buffer_size.py @@ -0,0 +1,51 @@ +import numpy as np + +from brainflow.data_filter import DataFilter, WaveletExtensionTypes, WaveletTypes + +# db15 has a 30 tap filter, so a level 5 decomposition needs a reasonably long signal and +# grows the coefficient array well past the old fixed 2 * (40 + 1) headroom. +DATA_LEN = 1024 +DECOMPOSITION_LEVEL = 5 + + +def main(): + rng = np.random.default_rng(3) + data = rng.standard_normal(DATA_LEN) + + output = DataFilter.perform_wavelet_transform( + np.copy(data), WaveletTypes.DB15, DECOMPOSITION_LEVEL, WaveletExtensionTypes.SYMMETRIC + ) + coeffs, lengths = output + + # The binding allocated data_len + 2 * (40 + 1) regardless of decomposition level, but + # the native side writes sum(lengths) doubles. At db15 level 5 that is 1167 against an + # allocation of 1106, so the transform wrote past the end of the array. Comparing the + # returned coefficient count against sum(lengths) is what makes the shortfall visible, + # because the wrapper slices the buffer it allocated. + assert coeffs.shape[0] == int(np.sum(lengths)), (coeffs.shape[0], int(np.sum(lengths))) + + # The headroom has to scale with the level, so an under-allocation at any level shows up. + for level in range(1, DECOMPOSITION_LEVEL + 1): + level_output = DataFilter.perform_wavelet_transform( + np.copy(data), WaveletTypes.DB15, level, WaveletExtensionTypes.SYMMETRIC + ) + level_coeffs, level_lengths = level_output + assert level_coeffs.shape[0] == int(np.sum(level_lengths)), ( + level, + level_coeffs.shape[0], + int(np.sum(level_lengths)), + ) + + # A truncated coefficient array cannot reconstruct the signal, so the round trip is the + # end-to-end check that nothing was lost. + restored = DataFilter.perform_inverse_wavelet_transform( + output, DATA_LEN, WaveletTypes.DB15, DECOMPOSITION_LEVEL, WaveletExtensionTypes.SYMMETRIC + ) + assert restored.shape[0] == DATA_LEN, restored.shape + assert np.allclose(restored, data, atol=1e-8), np.max(np.abs(restored - data)) + + print('wavelet buffer size regression passed') + + +if __name__ == '__main__': + main()