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()