Skip to main content

embedded_dsp/
transform.rs

1//! Fast Fourier Transform (FFT), Real FFT (RFFT), Discrete Cosine Transform (DCT-IV), and Bit Reversal functions.
2
3#[allow(unused_imports)]
4use crate::math::FloatMath;
5use crate::types::*;
6
7/// Bit reversal function for interleaved complex array of size `2 * n`.
8pub fn bit_reversal(data: &mut [f32], n: usize) {
9    let mut j = 0;
10    for i in 0..n {
11        if i < j {
12            data.swap(2 * i, 2 * j);
13            data.swap(2 * i + 1, 2 * j + 1);
14        }
15        let mut m = n >> 1;
16        while m >= 1 && j >= m {
17            j -= m;
18            m >>= 1;
19        }
20        j += m;
21    }
22}
23
24/// In-place Complex FFT for floating point 32-bit (`f32`).
25/// `data` is interleaved complex array of size `2 * n` (`[re0, im0, re1, im1, ...]`).
26/// `ifft_flag`: 0 for forward FFT, 1 for inverse FFT (IFFT).
27/// `bit_reverse_flag`: 1 to enable bit reversal, 0 to disable.
28pub fn cfft_f32(data: &mut [f32], n: usize, ifft_flag: u8, bit_reverse_flag: u8) {
29    if n < 2 || (n & (n - 1)) != 0 {
30        return;
31    }
32
33    if bit_reverse_flag != 0 {
34        bit_reversal(data, n);
35    }
36
37    let mut len = 2;
38    while len <= n {
39        let half_len = len / 2;
40        let angle =
41            (if ifft_flag != 0 { 2.0 } else { -2.0 }) * core::f32::consts::PI / (len as f32);
42        let w_step_re = angle.cos();
43        let w_step_im = angle.sin();
44
45        let mut i = 0;
46        while i < n {
47            let mut w_re = 1.0f32;
48            let mut w_im = 0.0f32;
49
50            for j in 0..half_len {
51                let u_idx = 2 * (i + j);
52                let v_idx = 2 * (i + j + half_len);
53
54                let u_re = data[u_idx];
55                let u_im = data[u_idx + 1];
56
57                let v_re = data[v_idx];
58                let v_im = data[v_idx + 1];
59
60                let t_re = v_re * w_re - v_im * w_im;
61                let t_im = v_re * w_im + v_im * w_re;
62
63                data[u_idx] = u_re + t_re;
64                data[u_idx + 1] = u_im + t_im;
65
66                data[v_idx] = u_re - t_re;
67                data[v_idx + 1] = u_im - t_im;
68
69                let next_w_re = w_re * w_step_re - w_im * w_step_im;
70                let next_w_im = w_re * w_step_im + w_im * w_step_re;
71                w_re = next_w_re;
72                w_im = next_w_im;
73            }
74            i += len;
75        }
76        len <<= 1;
77    }
78
79    if ifft_flag != 0 {
80        let norm = 1.0 / (n as f32);
81        for i in 0..(2 * n) {
82            data[i] *= norm;
83        }
84    }
85}
86
87/// Twiddle table length: supports radix-2 FFT sizes up to 512 (interleaved `2*n <= 1024`).
88const TWIDDLE_N: usize = 512;
89
90const fn wrap_pi(mut x: f32) -> f32 {
91    while x > core::f32::consts::PI {
92        x -= 2.0 * core::f32::consts::PI;
93    }
94    while x < -core::f32::consts::PI {
95        x += 2.0 * core::f32::consts::PI;
96    }
97    x
98}
99
100const fn cos_taylor(x: f32) -> f32 {
101    let x = wrap_pi(x);
102    let x2 = x * x;
103    let x4 = x2 * x2;
104    let x6 = x4 * x2;
105    let x8 = x4 * x4;
106    1.0 - x2 / 2.0 + x4 / 24.0 - x6 / 720.0 + x8 / 40320.0
107}
108
109const fn sin_taylor(x: f32) -> f32 {
110    let x = wrap_pi(x);
111    let x2 = x * x;
112    let x3 = x2 * x;
113    let x5 = x3 * x2;
114    let x7 = x5 * x2;
115    let x9 = x7 * x2;
116    x - x3 / 6.0 + x5 / 120.0 - x7 / 5040.0 + x9 / 362880.0
117}
118
119const fn gen_cos_q15() -> [i16; TWIDDLE_N] {
120    let mut t = [0i16; TWIDDLE_N];
121    let mut i = 0;
122    while i < TWIDDLE_N {
123        let a = (i as f32) * 2.0 * core::f32::consts::PI / TWIDDLE_N as f32;
124        let v = cos_taylor(a) * 32767.0;
125        t[i] = if v >= 32767.0 {
126            32767
127        } else if v <= -32768.0 {
128            -32768
129        } else {
130            v as i16
131        };
132        i += 1;
133    }
134    t
135}
136
137const fn gen_sin_q15() -> [i16; TWIDDLE_N] {
138    let mut t = [0i16; TWIDDLE_N];
139    let mut i = 0;
140    while i < TWIDDLE_N {
141        let a = (i as f32) * 2.0 * core::f32::consts::PI / TWIDDLE_N as f32;
142        let v = sin_taylor(a) * 32767.0;
143        t[i] = if v >= 32767.0 {
144            32767
145        } else if v <= -32768.0 {
146            -32768
147        } else {
148            v as i16
149        };
150        i += 1;
151    }
152    t
153}
154
155const COS_Q15: [i16; TWIDDLE_N] = gen_cos_q15();
156const SIN_Q15: [i16; TWIDDLE_N] = gen_sin_q15();
157
158fn twiddle_q15(k: usize, n: usize) -> (i16, i16) {
159    let idx = k.wrapping_mul(TWIDDLE_N / n) & (TWIDDLE_N - 1);
160    (COS_Q15[idx], SIN_Q15[idx])
161}
162
163fn bit_reversal_q15(data: &mut [q15], n: usize) {
164    let mut j = 0;
165    for i in 0..n {
166        if i < j {
167            data.swap(2 * i, 2 * j);
168            data.swap(2 * i + 1, 2 * j + 1);
169        }
170        let mut m = n >> 1;
171        while m >= 1 && j >= m {
172            j -= m;
173            m >>= 1;
174        }
175        j += m;
176    }
177}
178
179fn bit_reversal_q31(data: &mut [q31], n: usize) {
180    let mut j = 0;
181    for i in 0..n {
182        if i < j {
183            data.swap(2 * i, 2 * j);
184            data.swap(2 * i + 1, 2 * j + 1);
185        }
186        let mut m = n >> 1;
187        while m >= 1 && j >= m {
188            j -= m;
189            m >>= 1;
190        }
191        j += m;
192    }
193}
194
195#[inline]
196fn sat_q15(v: i32) -> q15 {
197    v.clamp(i16::MIN as i32, i16::MAX as i32) as q15
198}
199
200#[inline]
201fn sat_q31(v: i64) -> q31 {
202    v.clamp(i32::MIN as i64, i32::MAX as i64) as q31
203}
204
205/// In-place radix-2 DIT Complex FFT for Q31.
206///
207/// Each stage arithmetic-shifts right by 1 so a full-scale input does not wrap;
208/// a forward transform of length `n` is therefore scaled by about `1/n` versus
209/// [`cfft_f32`]. Inverse uses conjugated twiddles and the same per-stage shift
210/// (no extra `1/n`), so `ifft(fft(x)) ≈ x / n`.
211///
212/// `n` must be a power of two in `2..=512`. `data` is interleaved `[re, im, ...]`.
213pub fn cfft_q31(data: &mut [q31], n: usize, ifft_flag: u8, bit_reverse_flag: u8) {
214    if n < 2 || n > TWIDDLE_N || (n & (n - 1)) != 0 || data.len() < 2 * n {
215        return;
216    }
217
218    if bit_reverse_flag != 0 {
219        bit_reversal_q31(data, n);
220    }
221
222    let mut len = 2;
223    while len <= n {
224        let half_len = len / 2;
225        let mut i = 0;
226        while i < n {
227            for j in 0..half_len {
228                let (w_re_s, w_im_s) = twiddle_q15(j, len);
229                let w_re = (w_re_s as i32) << 16;
230                let mut w_im = (w_im_s as i32) << 16;
231                if ifft_flag == 0 {
232                    w_im = -w_im;
233                }
234
235                let u_idx = 2 * (i + j);
236                let v_idx = 2 * (i + j + half_len);
237                let u_re = data[u_idx] as i64;
238                let u_im = data[u_idx + 1] as i64;
239                let v_re = data[v_idx] as i64;
240                let v_im = data[v_idx + 1] as i64;
241                let wr = w_re as i64;
242                let wi = w_im as i64;
243
244                let t_re = (v_re * wr - v_im * wi) >> 31;
245                let t_im = (v_re * wi + v_im * wr) >> 31;
246
247                data[u_idx] = sat_q31((u_re + t_re) >> 1);
248                data[u_idx + 1] = sat_q31((u_im + t_im) >> 1);
249                data[v_idx] = sat_q31((u_re - t_re) >> 1);
250                data[v_idx + 1] = sat_q31((u_im - t_im) >> 1);
251            }
252            i += len;
253        }
254        len <<= 1;
255    }
256}
257
258/// In-place radix-2 DIT Complex FFT for Q15.
259///
260/// Same scaling as [`cfft_q31`]: about `1/n` per forward or inverse transform.
261/// `n` must be a power of two in `2..=512`.
262pub fn cfft_q15(data: &mut [q15], n: usize, ifft_flag: u8, bit_reverse_flag: u8) {
263    if n < 2 || n > TWIDDLE_N || (n & (n - 1)) != 0 || data.len() < 2 * n {
264        return;
265    }
266
267    if bit_reverse_flag != 0 {
268        bit_reversal_q15(data, n);
269    }
270
271    let mut len = 2;
272    while len <= n {
273        let half_len = len / 2;
274        let mut i = 0;
275        while i < n {
276            for j in 0..half_len {
277                let (w_re_s, mut w_im_s) = twiddle_q15(j, len);
278                if ifft_flag == 0 {
279                    w_im_s = w_im_s.saturating_neg();
280                }
281
282                let u_idx = 2 * (i + j);
283                let v_idx = 2 * (i + j + half_len);
284                let u_re = data[u_idx] as i32;
285                let u_im = data[u_idx + 1] as i32;
286                let v_re = data[v_idx] as i32;
287                let v_im = data[v_idx + 1] as i32;
288                let wr = w_re_s as i32;
289                let wi = w_im_s as i32;
290
291                let t_re = (v_re * wr - v_im * wi) >> 15;
292                let t_im = (v_re * wi + v_im * wr) >> 15;
293
294                data[u_idx] = sat_q15((u_re + t_re) >> 1);
295                data[u_idx + 1] = sat_q15((u_im + t_im) >> 1);
296                data[v_idx] = sat_q15((u_re - t_re) >> 1);
297                data[v_idx + 1] = sat_q15((u_im - t_im) >> 1);
298            }
299            i += len;
300        }
301        len <<= 1;
302    }
303}
304
305/// Real FFT for floating point 32-bit (`f32`).
306/// `src` has `n` real samples. `dst` receives `2 * n` complex outputs.
307pub fn rfft_f32(src: &[f32], dst: &mut [f32], n: usize, ifft_flag: u8) {
308    let len = src.len().min(n);
309    let mut c_data = [0.0f32; 1024];
310    if 2 * len > c_data.len() || dst.len() < 2 * len {
311        return;
312    }
313
314    for i in 0..len {
315        c_data[2 * i] = src[i];
316        c_data[2 * i + 1] = 0.0;
317    }
318
319    cfft_f32(&mut c_data[..2 * len], len, ifft_flag, 1);
320    dst[..2 * len].copy_from_slice(&c_data[..2 * len]);
321}
322
323/// Packed real FFT (N/2-point complex FFT of even/odd samples, then unpack).
324/// Forward only (`ifft_flag == 0`); inverse still uses a complex FFT of real+0j.
325fn packed_rfft_q15_forward(src: &[q15], dst: &mut [q15], n: usize) {
326    let m = n / 2;
327    let mut z = [0i16; 1024];
328    if 2 * m > z.len() {
329        return;
330    }
331    for k in 0..m {
332        z[2 * k] = src[2 * k];
333        z[2 * k + 1] = src[2 * k + 1];
334    }
335    cfft_q15(&mut z[..2 * m], m, 0, 1);
336
337    let z0r = z[0] as i32;
338    let z0i = z[1] as i32;
339    dst[0] = sat_q15((z0r + z0i) >> 1);
340    dst[1] = 0;
341    dst[n] = sat_q15((z0r - z0i) >> 1);
342    dst[n + 1] = 0;
343
344    for k in 1..m {
345        let zr = z[2 * k] as i32;
346        let zi = z[2 * k + 1] as i32;
347        let znr = z[2 * (m - k)] as i32;
348        let zni = z[2 * (m - k) + 1] as i32;
349
350        let xe_re = (zr + znr) >> 1;
351        let xe_im = (zi - zni) >> 1;
352        let xo_re = (zi + zni) >> 1;
353        let xo_im = (znr - zr) >> 1;
354
355        let (wr, wi_s) = twiddle_q15(k, n);
356        let wr = wr as i32;
357        let wi = -(wi_s as i32);
358
359        let t_re = (xo_re * wr - xo_im * wi) >> 15;
360        let t_im = (xo_re * wi + xo_im * wr) >> 15;
361
362        dst[2 * k] = sat_q15((xe_re + t_re) >> 1);
363        dst[2 * k + 1] = sat_q15((xe_im + t_im) >> 1);
364        dst[2 * (n - k)] = sat_q15((xe_re - t_re) >> 1);
365        dst[2 * (n - k) + 1] = sat_q15((t_im - xe_im) >> 1);
366    }
367}
368
369fn packed_rfft_q31_forward(src: &[q31], dst: &mut [q31], n: usize) {
370    let m = n / 2;
371    let mut z = [0i32; 1024];
372    if 2 * m > z.len() {
373        return;
374    }
375    for k in 0..m {
376        z[2 * k] = src[2 * k];
377        z[2 * k + 1] = src[2 * k + 1];
378    }
379    cfft_q31(&mut z[..2 * m], m, 0, 1);
380
381    let z0r = z[0] as i64;
382    let z0i = z[1] as i64;
383    dst[0] = sat_q31((z0r + z0i) >> 1);
384    dst[1] = 0;
385    dst[n] = sat_q31((z0r - z0i) >> 1);
386    dst[n + 1] = 0;
387
388    for k in 1..m {
389        let zr = z[2 * k] as i64;
390        let zi = z[2 * k + 1] as i64;
391        let znr = z[2 * (m - k)] as i64;
392        let zni = z[2 * (m - k) + 1] as i64;
393
394        let xe_re = (zr + znr) >> 1;
395        let xe_im = (zi - zni) >> 1;
396        let xo_re = (zi + zni) >> 1;
397        let xo_im = (znr - zr) >> 1;
398
399        let (wr_s, wi_s) = twiddle_q15(k, n);
400        let wr = (wr_s as i64) << 16;
401        let wi = -((wi_s as i64) << 16);
402
403        let t_re = (xo_re * wr - xo_im * wi) >> 31;
404        let t_im = (xo_re * wi + xo_im * wr) >> 31;
405
406        dst[2 * k] = sat_q31((xe_re + t_re) >> 1);
407        dst[2 * k + 1] = sat_q31((xe_im + t_im) >> 1);
408        dst[2 * (n - k)] = sat_q31((xe_re - t_re) >> 1);
409        dst[2 * (n - k) + 1] = sat_q31((t_im - xe_im) >> 1);
410    }
411}
412
413fn packed_irfft_q15(src: &[q15], dst: &mut [q15], n: usize) {
414    let m = n / 2;
415    let mut z = [0i16; 1024];
416    if 2 * m > z.len() {
417        return;
418    }
419
420    let dc = src[0] as i32;
421    let ny = src[n] as i32;
422    z[0] = sat_q15(dc + ny);
423    z[1] = sat_q15(dc - ny);
424
425    for k in 1..m {
426        let xkr = src[2 * k] as i32;
427        let xki = src[2 * k + 1] as i32;
428        let xnr = src[2 * (n - k)] as i32;
429        let xni = src[2 * (n - k) + 1] as i32;
430
431        let xe_re = xkr + xnr;
432        let xe_im = xki - xni;
433        let t_re = xkr - xnr;
434        let t_im = xki + xni;
435
436        let (wr, wi_s) = twiddle_q15(k, n);
437        let wr = wr as i32;
438        let wi = wi_s as i32;
439        let xo_re = (t_re * wr - t_im * wi) >> 15;
440        let xo_im = (t_re * wi + t_im * wr) >> 15;
441
442        z[2 * k] = sat_q15(xe_re - xo_im);
443        z[2 * k + 1] = sat_q15(xe_im + xo_re);
444    }
445
446    cfft_q15(&mut z[..2 * m], m, 1, 1);
447    for k in 0..m {
448        dst[2 * k] = sat_q15((z[2 * k] as i32) >> 1);
449        dst[2 * k + 1] = sat_q15((z[2 * k + 1] as i32) >> 1);
450    }
451}
452
453fn packed_irfft_q31(src: &[q31], dst: &mut [q31], n: usize) {
454    let m = n / 2;
455    let mut z = [0i32; 1024];
456    if 2 * m > z.len() {
457        return;
458    }
459
460    let dc = src[0] as i64;
461    let ny = src[n] as i64;
462    z[0] = sat_q31(dc + ny);
463    z[1] = sat_q31(dc - ny);
464
465    for k in 1..m {
466        let xkr = src[2 * k] as i64;
467        let xki = src[2 * k + 1] as i64;
468        let xnr = src[2 * (n - k)] as i64;
469        let xni = src[2 * (n - k) + 1] as i64;
470
471        let xe_re = xkr + xnr;
472        let xe_im = xki - xni;
473        let t_re = xkr - xnr;
474        let t_im = xki + xni;
475
476        let (wr_s, wi_s) = twiddle_q15(k, n);
477        let wr = (wr_s as i64) << 16;
478        let wi = (wi_s as i64) << 16;
479        let xo_re = (t_re * wr - t_im * wi) >> 31;
480        let xo_im = (t_re * wi + t_im * wr) >> 31;
481
482        z[2 * k] = sat_q31(xe_re - xo_im);
483        z[2 * k + 1] = sat_q31(xe_im + xo_re);
484    }
485
486    cfft_q31(&mut z[..2 * m], m, 1, 1);
487    for k in 0..m {
488        dst[2 * k] = sat_q31((z[2 * k] as i64) >> 1);
489        dst[2 * k + 1] = sat_q31((z[2 * k + 1] as i64) >> 1);
490    }
491}
492
493fn rfft_q_can_pack(len: usize, ifft_flag: u8) -> bool {
494    ifft_flag == 0 && len >= 4 && len <= TWIDDLE_N && (len & (len - 1)) == 0
495}
496
497/// Real FFT for Q31 fixed-point.
498///
499/// Forward (`ifft_flag == 0`) uses a packed N/2 complex FFT of even/odd samples
500/// (same output layout as a zero-padded [`cfft_q31`]: `2*n` interleaved bins).
501/// Inverse (`ifft_flag != 0`) still runs an `n`-point complex FFT of real+0j;
502/// use [`irfft_q31`] to invert a packed spectrum.
503pub fn rfft_q31(src: &[q31], dst: &mut [q31], n: usize, ifft_flag: u8) {
504    let len = src.len().min(n);
505    if dst.len() < 2 * len {
506        return;
507    }
508
509    if rfft_q_can_pack(len, ifft_flag) {
510        packed_rfft_q31_forward(&src[..len], dst, len);
511        return;
512    }
513
514    let mut c_data = [0; 1024];
515    if 2 * len > c_data.len() {
516        return;
517    }
518
519    for i in 0..len {
520        c_data[2 * i] = src[i];
521        c_data[2 * i + 1] = 0;
522    }
523    cfft_q31(&mut c_data[..2 * len], len, ifft_flag, 1);
524    dst[..2 * len].copy_from_slice(&c_data[..2 * len]);
525}
526
527/// Real FFT for Q15 fixed-point.
528///
529/// Forward packed N/2 algorithm; see [`rfft_q31`]. Scale versus [`rfft_f32`] is
530/// about `1/n`, matching [`cfft_q15`].
531pub fn rfft_q15(src: &[q15], dst: &mut [q15], n: usize, ifft_flag: u8) {
532    let len = src.len().min(n);
533    if dst.len() < 2 * len {
534        return;
535    }
536
537    if rfft_q_can_pack(len, ifft_flag) {
538        packed_rfft_q15_forward(&src[..len], dst, len);
539        return;
540    }
541
542    let mut c_data = [0; 1024];
543    if 2 * len > c_data.len() {
544        return;
545    }
546
547    for i in 0..len {
548        c_data[2 * i] = src[i];
549        c_data[2 * i + 1] = 0;
550    }
551    cfft_q15(&mut c_data[..2 * len], len, ifft_flag, 1);
552    dst[..2 * len].copy_from_slice(&c_data[..2 * len]);
553}
554
555/// Inverse packed real FFT. `src` is `2 * n` interleaved bins from [`rfft_q31`];
556/// `dst` receives `n` real samples. Combined with a forward transform,
557/// `irfft(rfft(x)) ≈ x / n` (same convention as [`cfft_q31`]).
558pub fn irfft_q31(src: &[q31], dst: &mut [q31], n: usize) {
559    if n < 4 || n > TWIDDLE_N || (n & (n - 1)) != 0 || src.len() < 2 * n || dst.len() < n {
560        return;
561    }
562    packed_irfft_q31(&src[..2 * n], dst, n);
563}
564
565/// Inverse packed real FFT. `src` is `2 * n` interleaved bins from [`rfft_q15`];
566/// `dst` receives `n` real samples. Combined with a forward transform,
567/// `irfft(rfft(x)) ≈ x / n` (same convention as [`cfft_q15`]).
568pub fn irfft_q15(src: &[q15], dst: &mut [q15], n: usize) {
569    if n < 4 || n > TWIDDLE_N || (n & (n - 1)) != 0 || src.len() < 2 * n || dst.len() < n {
570        return;
571    }
572    packed_irfft_q15(&src[..2 * n], dst, n);
573}
574
575/// Discrete Cosine Transform Type IV (DCT-IV) for f32.
576pub fn dct4_f32(src: &[f32], dst: &mut [f32], n: usize) {
577    let len = src.len().min(dst.len()).min(n);
578    let pi_over_n = core::f32::consts::PI / (len as f32);
579
580    for k in 0..len {
581        let mut sum = 0.0f32;
582        let k_factor = (k as f32 + 0.5) * pi_over_n;
583        for n_idx in 0..len {
584            let angle = (n_idx as f32 + 0.5) * k_factor;
585            sum += src[n_idx] * angle.cos();
586        }
587        let norm = (2.0 / len as f32).sqrt();
588        dst[k] = sum * norm;
589    }
590}
591
592// --- Fast Walsh-Hadamard Transform (FWHT) ---
593
594/// In-place Fast Walsh-Hadamard Transform (FWHT) for floating point `f32`.
595///
596/// `data.len()` must be a power of 2 (e.g. 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024).
597pub fn fwht_f32(data: &mut [f32]) -> Status {
598    let n = data.len();
599    if n < 2 || (n & (n - 1)) != 0 {
600        return Status::ArgumentError;
601    }
602
603    let mut h = 1;
604    while h < n {
605        let mut i = 0;
606        while i < n {
607            for j in i..(i + h) {
608                let x = data[j];
609                let y = data[j + h];
610                data[j] = x + y;
611                data[j + h] = x - y;
612            }
613            i += h * 2;
614        }
615        h *= 2;
616    }
617
618    Status::Success
619}
620
621/// In-place Inverse Fast Walsh-Hadamard Transform (IFWHT) for floating point `f32` (normalized by $1/N$).
622pub fn ifwht_f32(data: &mut [f32]) -> Status {
623    let status = fwht_f32(data);
624    if status != Status::Success {
625        return status;
626    }
627    let norm = 1.0f32 / (data.len() as f32);
628    for val in data.iter_mut() {
629        *val *= norm;
630    }
631    Status::Success
632}
633
634/// In-place Fast Walsh-Hadamard Transform (FWHT) for 32-bit integers (`i32`).
635pub fn fwht_i32(data: &mut [i32]) -> Status {
636    let n = data.len();
637    if n < 2 || (n & (n - 1)) != 0 {
638        return Status::ArgumentError;
639    }
640
641    let mut h = 1;
642    while h < n {
643        let mut i = 0;
644        while i < n {
645            for j in i..(i + h) {
646                let x = data[j];
647                let y = data[j + h];
648                data[j] = x.wrapping_add(y);
649                data[j + h] = x.wrapping_sub(y);
650            }
651            i += h * 2;
652        }
653        h *= 2;
654    }
655
656    Status::Success
657}
658
659// --- Haar Transform (Jörg Arndt, "Matters Computational", Ch. 24) ---
660
661/// In-place, orthogonal Haar Transform for `f32`: an `O(n)` multiresolution transform using
662/// only additions, subtractions, and a `sqrt(0.5)` scale factor per stage, with no
663/// trigonometric factors at all (unlike the Fourier/Hartley transforms).
664///
665/// `data.len()` must be a power of 2 (e.g. 2, 4, 8, ..., 1024).
666pub fn haar_transform_f32(data: &mut [f32]) -> Status {
667    let n = data.len();
668    if n < 2 || (n & (n - 1)) != 0 {
669        return Status::ArgumentError;
670    }
671
672    let s2 = (0.5f32).sqrt();
673    let mut v = 1.0f32;
674    let mut js = 2;
675    while js <= n {
676        v *= s2;
677        let half = js >> 1;
678        let mut j = 0;
679        while j < n {
680            let t = j + half;
681            let x = data[j];
682            let y = data[t];
683            data[j] = x + y;
684            data[t] = (x - y) * v;
685            j += js;
686        }
687        js <<= 1;
688    }
689    data[0] *= v; // v == 1 / sqrt(n)
690
691    Status::Success
692}
693
694/// In-place Inverse Haar Transform for `f32`, undoing [`haar_transform_f32`].
695pub fn inverse_haar_transform_f32(data: &mut [f32]) -> Status {
696    let n = data.len();
697    if n < 2 || (n & (n - 1)) != 0 {
698        return Status::ArgumentError;
699    }
700
701    let s2 = 2.0f32.sqrt();
702    let mut v = 1.0f32 / (n as f32).sqrt();
703    data[0] *= v;
704
705    let mut js = n;
706    while js >= 2 {
707        let half = js >> 1;
708        let mut j = 0;
709        while j < n {
710            let t = j + half;
711            let x = data[j];
712            let y = data[t] * v;
713            data[j] = x + y;
714            data[t] = x - y;
715            j += js;
716        }
717        v *= s2;
718        js >>= 1;
719    }
720
721    Status::Success
722}
723
724/// In-place, non-normalized Haar Transform for `i32`: a forward-only, integer-exact
725/// decomposition using only wrapping add/subtract (no scaling), analogous to
726/// [`fwht_i32`]. Because the transform is non-normalized, an exact-integer inverse does not
727/// exist in general (undoing it requires dividing by powers of 2 that may not evenly divide
728/// intermediate sums); use [`haar_transform_f32`] / [`inverse_haar_transform_f32`] when an
729/// invertible round trip is required.
730///
731/// `data.len()` must be a power of 2 (e.g. 2, 4, 8, ..., 1024).
732pub fn haar_transform_i32(data: &mut [i32]) -> Status {
733    let n = data.len();
734    if n < 2 || (n & (n - 1)) != 0 {
735        return Status::ArgumentError;
736    }
737
738    let mut js = 2;
739    while js <= n {
740        let half = js >> 1;
741        let mut j = 0;
742        while j < n {
743            let t = j + half;
744            let x = data[j];
745            let y = data[t];
746            data[j] = x.wrapping_add(y);
747            data[t] = x.wrapping_sub(y);
748            j += js;
749        }
750        js <<= 1;
751    }
752
753    Status::Success
754}
755
756// --- Hartley Transform (Jörg Arndt, "Matters Computational", Ch. 25) ---
757
758/// In-place Discrete Hartley Transform for `f32`.
759///
760/// Computed via the identity relating the Hartley and Fourier transforms (Ch. 25):
761/// `H[a] = (Re(F[a]) - Im(F[a])) / sqrt(n)`, built on top of [`cfft_f32`] rather than a
762/// dedicated real-only butterfly network, so it costs a full complex FFT internally
763/// (`n <= 512`) even though its inputs and outputs are purely real.
764///
765/// The Hartley transform is its own inverse (`H[H[a]] = a`): call this function a second time
766/// on its output to invert it, with no separate inverse routine needed.
767///
768/// `data.len()` must be a power of 2 (e.g. 2, 4, 8, ..., 512).
769pub fn hartley_transform_f32(data: &mut [f32]) -> Status {
770    let n = data.len();
771    if n < 2 || (n & (n - 1)) != 0 {
772        return Status::ArgumentError;
773    }
774    if 2 * n > 1024 {
775        return Status::LengthError;
776    }
777
778    let mut c_data = [0.0f32; 1024];
779    for i in 0..n {
780        c_data[2 * i] = data[i];
781        c_data[2 * i + 1] = 0.0;
782    }
783
784    cfft_f32(&mut c_data[..2 * n], n, 0, 1);
785
786    let inv_sqrt_n = 1.0 / (n as f32).sqrt();
787    for i in 0..n {
788        data[i] = (c_data[2 * i] - c_data[2 * i + 1]) * inv_sqrt_n;
789    }
790
791    Status::Success
792}
793
794// --- Generalized Wavelet Transform (Jörg Arndt, "Matters Computational", Ch. 27) ---
795
796/// The Daubechies-4 orthogonal wavelet low-pass filter taps (Ch. 27.1), verified to satisfy
797/// the wavelet conditions `sum(h_j^2) = 1` and `sum(h_j * h_{j+2}) = 0`. Using
798/// `[sqrt(0.5), sqrt(0.5)]` instead recovers the Haar wavelet as a special case.
799pub const DAUBECHIES_4: [f32; 4] = [0.482_962_9, 0.836_516_3, 0.224_143_87, -0.129_409_52];
800
801/// The high-pass filter tap derived from low-pass filter `h` (Ch. 27.1, Eq. 27.1-2):
802/// `g[k] = (-1)^k * h[n - 1 - k]`.
803#[inline(always)]
804fn wavelet_high_pass_tap(h: &[f32], k: usize) -> f32 {
805    let v = h[h.len() - 1 - k];
806    if k % 2 == 0 { v } else { -v }
807}
808
809/// Performs one level of a fast wavelet transform step on the first `m` elements of `data`,
810/// using wavelet filter `h` (low-pass) and its derived high-pass filter. Writes the low-pass
811/// ("scaling") coefficients to `data[0..m/2]` and the high-pass ("wavelet") coefficients to
812/// `data[m/2..m]`; the underlying convolution wraps around cyclically at the block boundary.
813///
814/// `m` must be a power of 2; `h.len()` must be even and `<= m`.
815pub fn wavelet_step_f32(data: &mut [f32], m: usize, h: &[f32]) -> Status {
816    let taps = h.len();
817    if m < 2 || (m & (m - 1)) != 0 || taps == 0 || taps % 2 != 0 || taps > m || data.len() < m {
818        return Status::ArgumentError;
819    }
820    if m > 1024 {
821        return Status::LengthError;
822    }
823
824    let mut scratch = [0.0f32; 1024];
825    let nh = m >> 1;
826    let mut i = 0;
827    while i < m {
828        let mut s = 0.0f32;
829        let mut d = 0.0f32;
830        for k in 0..taps {
831            let idx = (i + k) % m;
832            let x = data[idx];
833            s += h[k] * x;
834            d += wavelet_high_pass_tap(h, k) * x;
835        }
836        let j = i / 2;
837        scratch[j] = s;
838        scratch[nh + j] = d;
839        i += 2;
840    }
841    data[..m].copy_from_slice(&scratch[..m]);
842
843    Status::Success
844}
845
846/// Performs the exact inverse of one [`wavelet_step_f32`] level.
847///
848/// `m` must be a power of 2; `h.len()` must be even and `<= m`.
849pub fn inverse_wavelet_step_f32(data: &mut [f32], m: usize, h: &[f32]) -> Status {
850    let taps = h.len();
851    if m < 2 || (m & (m - 1)) != 0 || taps == 0 || taps % 2 != 0 || taps > m || data.len() < m {
852        return Status::ArgumentError;
853    }
854    if m > 1024 {
855        return Status::LengthError;
856    }
857
858    let mut scratch = [0.0f32; 1024];
859    let nh = m >> 1;
860    for j in 0..nh {
861        let s = data[j];
862        let d = data[nh + j];
863        for k in 0..taps {
864            let idx = (2 * j + k) % m;
865            scratch[idx] += h[k] * s + wavelet_high_pass_tap(h, k) * d;
866        }
867    }
868    data[..m].copy_from_slice(&scratch[..m]);
869
870    Status::Success
871}
872
873/// Performs a full multi-level fast wavelet transform (Ch. 27): repeatedly applies
874/// [`wavelet_step_f32`] to the lower half of the array, halving the active block length each
875/// time, stopping once the block would be smaller than the filter itself (mirroring the Haar
876/// transform's pyramid structure).
877///
878/// `data.len()` must be a power of 2 and `>= h.len()`.
879pub fn wavelet_transform_f32(data: &mut [f32], h: &[f32]) -> Status {
880    let n = data.len();
881    if n < 2 || (n & (n - 1)) != 0 || h.len() > n {
882        return Status::ArgumentError;
883    }
884
885    let mut m = n;
886    while m >= h.len() {
887        let status = wavelet_step_f32(&mut data[..m], m, h);
888        if status != Status::Success {
889            return status;
890        }
891        m >>= 1;
892    }
893
894    Status::Success
895}
896
897/// Performs the exact inverse of [`wavelet_transform_f32`].
898///
899/// `data.len()` must be a power of 2 and `>= h.len()`.
900pub fn inverse_wavelet_transform_f32(data: &mut [f32], h: &[f32]) -> Status {
901    let n = data.len();
902    if n < 2 || (n & (n - 1)) != 0 || h.len() > n {
903        return Status::ArgumentError;
904    }
905
906    let mut smallest = n;
907    while smallest >= h.len() {
908        smallest >>= 1;
909    }
910    smallest <<= 1;
911
912    let mut m = smallest;
913    while m <= n {
914        let status = inverse_wavelet_step_f32(&mut data[..m], m, h);
915        if status != Status::Success {
916            return status;
917        }
918        m <<= 1;
919    }
920
921    Status::Success
922}