Skip to main content

rusty_opus/
mdct.rs

1use crate::kiss_fft::{KissCpx, KissFftState, opus_fft_impl};
2use std::f32::consts::PI;
3use std::mem::MaybeUninit;
4
5const MAX_N2: usize = 960;
6const MAX_N4: usize = 480;
7
8pub struct MdctLookup {
9    pub n: usize,
10    pub max_lm: usize,
11    kfft: Vec<Option<KissFftState>>,
12    trig: Vec<f32>,
13}
14
15impl MdctLookup {
16    pub fn new(n: usize, max_lm: usize) -> Self {
17        let mut kfft = Vec::new();
18        let mut trig = Vec::new();
19        let mut curr_n = n;
20
21        for shift in 0..=max_lm {
22            let n4 = curr_n / 4;
23
24            if shift == 0 {
25                kfft.push(KissFftState::new(n4));
26            } else if let Some(base) = kfft.first().unwrap().as_ref() {
27                kfft.push(KissFftState::new_sub(base, n4));
28            } else {
29                kfft.push(None);
30            }
31
32            let n2 = curr_n / 2;
33            for i in 0..n2 {
34                let angle = 2.0 * PI * (i as f32 + 0.125) / curr_n as f32;
35                trig.push(angle.cos());
36            }
37
38            curr_n >>= 1;
39        }
40
41        Self {
42            n,
43            max_lm,
44            kfft,
45            trig,
46        }
47    }
48
49    fn get_trig(&self, shift: usize) -> (&[f32], usize) {
50        let mut offset = 0;
51        let mut curr_n = self.n;
52        for _ in 0..shift {
53            offset += curr_n / 2;
54            curr_n >>= 1;
55        }
56        (&self.trig[offset..offset + curr_n / 2], curr_n / 4)
57    }
58
59    pub fn get_trig_debug(&self, shift: usize) -> &[f32] {
60        let (trig, _) = self.get_trig(shift);
61        trig
62    }
63
64    #[inline]
65    pub fn forward(
66        &self,
67        input: &[f32],
68        output: &mut [f32],
69        window: &[f32],
70        overlap: usize,
71        shift: usize,
72        stride: usize,
73    ) {
74        let _prof = crate::prof::scope(crate::prof::Stage::CeltMdct);
75        let st = self.kfft[shift]
76            .as_ref()
77            .expect("FFT state not initialized");
78        let n = self.n >> shift;
79        let n2 = n / 2;
80        let n4 = n / 4;
81        let scale = st.scale();
82
83        let (trig, _) = self.get_trig(shift);
84        let overlap2 = overlap / 2;
85
86        let mut f_buf = [MaybeUninit::<f32>::uninit(); MAX_N2];
87        let mut f2_buf = [MaybeUninit::<KissCpx>::uninit(); MAX_N4];
88
89        let f = unsafe { std::slice::from_raw_parts_mut(f_buf.as_mut_ptr() as *mut f32, n2) };
90        let f2 = unsafe { std::slice::from_raw_parts_mut(f2_buf.as_mut_ptr() as *mut KissCpx, n4) };
91
92        assert!(input.len() >= n2 + overlap2);
93        assert!(window.len() >= overlap);
94        assert!(
95            output.len() >= n2,
96            "MDCT forward: output buffer too small (need {}, have {})",
97            n2,
98            output.len()
99        );
100
101        {
102            let mut yp = 0usize;
103            let mut xp1 = overlap2;
104            let mut xp2 = n2 - 1 + overlap2;
105            let mut wp1 = overlap2;
106
107            let mut wp2 = overlap2.saturating_sub(1);
108
109            let limit = overlap.div_ceil(4);
110            let mid = n4.saturating_sub(limit);
111
112            let loop1_iters = limit.min(n4);
113            for _ in 0..loop1_iters {
114                let w1 = window[wp1];
115                let w2 = window[wp2];
116
117                f[yp] = input[xp1 + n2] * w2 + input[xp2] * w1;
118                yp += 1;
119
120                f[yp] = input[xp1] * w1 - input[xp2 - n2] * w2;
121                yp += 1;
122
123                xp1 += 2;
124                xp2 -= 2;
125                wp1 += 2;
126                wp2 = wp2.saturating_sub(2);
127            }
128
129            for _ in limit..mid {
130                f[yp] = input[xp2];
131                yp += 1;
132
133                f[yp] = input[xp1];
134                yp += 1;
135                xp1 += 2;
136                xp2 -= 2;
137            }
138
139            // C: after the middle loop, i == max(limit, N4-limit) and the third
140            // loop runs to N4. The old `if mid > limit {..} else { 0 }` yielded
141            // ZERO iterations when mid <= limit — exactly the short-block case
142            // (N == 2*overlap: n4 = 2*limit), leaving f[2*limit..n2) UNWRITTEN:
143            // uninitialized-stack reads on every transient sub-MDCT (this is
144            // what made the HYB-VBR bitstream hash move across builds).
145            let loop3_iters = n4 - limit.max(mid);
146            let mut wp1_l3 = 0usize;
147            let mut wp2_l3 = overlap.saturating_sub(1);
148            for _ in 0..loop3_iters {
149                let w1 = window[wp1_l3];
150                let w2 = window[wp2_l3];
151
152                f[yp] = -input[xp1 - n2] * w1 + input[xp2] * w2;
153                yp += 1;
154
155                f[yp] = input[xp1] * w2 + input[xp2 + n2] * w1;
156                yp += 1;
157
158                xp1 += 2;
159                xp2 -= 2;
160                wp1_l3 += 2;
161                wp2_l3 -= 2;
162            }
163        }
164
165        #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
166        unsafe {
167            if std::arch::is_x86_feature_detected!("avx") {
168                mdct_pre_rotation_avx(f, f2, trig, &st.bitrev[..n4], n4, scale);
169            } else {
170                for i in 0..n4 {
171                    let re = f[2 * i];
172                    let im = f[2 * i + 1];
173                    let t0 = trig[i];
174                    let t1 = trig[n4 + i];
175
176                    let yr = re * t0 - im * t1;
177                    let yi = im * t0 + re * t1;
178
179                    f2[st.bitrev[i] as usize] = KissCpx::new(yr * scale, yi * scale);
180                }
181            }
182        }
183        #[cfg(all(
184            not(any(target_arch = "x86", target_arch = "x86_64")),
185            target_arch = "aarch64"
186        ))]
187        {
188            mdct_pre_rotation_neon(f, f2, trig, &st.bitrev[..n4], n4, scale);
189        }
190        #[cfg(all(
191            not(any(target_arch = "x86", target_arch = "x86_64")),
192            not(target_arch = "aarch64")
193        ))]
194        for i in 0..n4 {
195            let re = f[2 * i];
196            let im = f[2 * i + 1];
197            let t0 = trig[i];
198            let t1 = trig[n4 + i];
199
200            let yr = re * t0 - im * t1;
201            let yi = im * t0 + re * t1;
202
203            f2[st.bitrev[i] as usize] = KissCpx::new(yr * scale, yi * scale);
204        }
205
206        opus_fft_impl(st, f2);
207
208        #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
209        unsafe {
210            if std::arch::is_x86_feature_detected!("avx") {
211                mdct_post_rotation_avx(f2, trig, output, n4, n2, stride);
212            } else {
213                for i in 0..n4 {
214                    let fp = &f2[i];
215                    let t0 = trig[i];
216                    let t1 = trig[n4 + i];
217
218                    let yr = fp.i * t1 - fp.r * t0;
219                    let yi = fp.r * t1 + fp.i * t0;
220
221                    output[i * 2 * stride] = yr;
222                    output[stride * (n2 - 1 - 2 * i)] = yi;
223                }
224            }
225        }
226        #[cfg(all(
227            not(any(target_arch = "x86", target_arch = "x86_64")),
228            target_arch = "aarch64"
229        ))]
230        {
231            mdct_post_rotation_neon(f2, trig, output, n4, n2, stride);
232        }
233        #[cfg(all(
234            not(any(target_arch = "x86", target_arch = "x86_64")),
235            not(target_arch = "aarch64")
236        ))]
237        for i in 0..n4 {
238            let fp = &f2[i];
239            let t0 = trig[i];
240            let t1 = trig[n4 + i];
241
242            let yr = fp.i * t1 - fp.r * t0;
243            let yi = fp.r * t1 + fp.i * t0;
244
245            output[i * 2 * stride] = yr;
246            output[stride * (n2 - 1 - 2 * i)] = yi;
247        }
248    }
249
250    #[inline]
251    pub fn backward(
252        &self,
253        input: &[f32],
254        output: &mut [f32],
255        window: &[f32],
256        overlap: usize,
257        shift: usize,
258        stride: usize,
259    ) {
260        let _prof = crate::prof::scope(crate::prof::Stage::CeltMdct);
261        // A malformed frame can carry an out-of-range `shift` or a `stride`/size
262        // inconsistent with the input buffer; bail gracefully rather than panic on
263        // the FFT-state index or the (checked) pre/post-rotation reads. Inert for
264        // valid streams (shift in range, buffers correctly sized).
265        let Some(st) = self.kfft.get(shift).and_then(|s| s.as_ref()) else {
266            return;
267        };
268        let n = self.n >> shift;
269        let n2 = n / 2;
270        let n4 = n / 4;
271        let overlap2 = overlap / 2;
272        if n4 == 0 || n4 > MAX_N4 || stride.saturating_mul(n2.saturating_sub(1)) >= input.len() {
273            return;
274        }
275
276        let (trig, _) = self.get_trig(shift);
277
278        let mut f2_buf = [MaybeUninit::<KissCpx>::uninit(); MAX_N4];
279
280        let f2 = unsafe { std::slice::from_raw_parts_mut(f2_buf.as_mut_ptr() as *mut KissCpx, n4) };
281
282        #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
283        unsafe {
284            if std::arch::is_x86_feature_detected!("avx") {
285                mdct_backward_pre_rotation_avx(input, f2, trig, &st.bitrev[..n4], n4, n2, stride);
286            } else {
287                for i in 0..n4 {
288                    let rev = st.bitrev[i] as usize;
289                    let x1 = input[2 * i * stride];
290                    let x2 = input[stride * (n2 - 1 - 2 * i)];
291                    let t0 = trig[i];
292                    let t1 = trig[n4 + i];
293
294                    let yr = x2 * t0 + x1 * t1;
295                    let yi = x1 * t0 - x2 * t1;
296
297                    f2[rev] = KissCpx::new(yi, yr);
298                }
299            }
300        }
301        #[cfg(all(
302            not(any(target_arch = "x86", target_arch = "x86_64")),
303            target_arch = "aarch64"
304        ))]
305        {
306            mdct_backward_pre_rotation_neon(input, f2, trig, &st.bitrev[..n4], n4, n2, stride);
307        }
308        #[cfg(all(
309            not(any(target_arch = "x86", target_arch = "x86_64")),
310            not(target_arch = "aarch64")
311        ))]
312        for i in 0..n4 {
313            let rev = st.bitrev[i] as usize;
314            let x1 = input[2 * i * stride];
315            let x2 = input[stride * (n2 - 1 - 2 * i)];
316            let t0 = trig[i];
317            let t1 = trig[n4 + i];
318
319            let yr = x2 * t0 + x1 * t1;
320            let yi = x1 * t0 - x2 * t1;
321
322            f2[rev] = KissCpx::new(yi, yr);
323        }
324
325        opus_fft_impl(st, f2);
326
327        assert!(output.len() >= overlap2 + n2);
328
329        #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
330        unsafe {
331            if std::arch::is_x86_feature_detected!("avx") {
332                mdct_backward_post_rotation_avx(f2, trig, output, n4, n2, overlap2);
333            } else {
334                for i in 0..((n4 + 1) >> 1) {
335                    let im0 = f2[i].r;
336                    let re0 = f2[i].i;
337                    let t0_0 = trig[i];
338                    let t1_0 = trig[n4 + i];
339
340                    let yr0 = re0 * t0_0 + im0 * t1_0;
341                    let yi0 = re0 * t1_0 - im0 * t0_0;
342
343                    let j = n4 - 1 - i;
344                    let im1 = f2[j].r;
345                    let re1 = f2[j].i;
346                    let t0_1 = trig[j];
347                    let t1_1 = trig[n4 + j];
348
349                    let yr1 = re1 * t0_1 + im1 * t1_1;
350                    let yi1 = re1 * t1_1 - im1 * t0_1;
351
352                    output[overlap2 + 2 * i] = yr0;
353                    output[overlap2 + n2 - 1 - 2 * i] = yi0;
354                    output[overlap2 + n2 - 2 - 2 * i] = yr1;
355                    output[overlap2 + 2 * i + 1] = yi1;
356                }
357            }
358        }
359        #[cfg(all(
360            not(any(target_arch = "x86", target_arch = "x86_64")),
361            target_arch = "aarch64"
362        ))]
363        {
364            mdct_backward_post_rotation_neon(f2, trig, output, n4, n2, overlap2);
365        }
366        #[cfg(all(
367            not(any(target_arch = "x86", target_arch = "x86_64")),
368            not(target_arch = "aarch64")
369        ))]
370        for i in 0..((n4 + 1) >> 1) {
371            let im0 = f2[i].r;
372            let re0 = f2[i].i;
373            let t0_0 = trig[i];
374            let t1_0 = trig[n4 + i];
375
376            let yr0 = re0 * t0_0 + im0 * t1_0;
377            let yi0 = re0 * t1_0 - im0 * t0_0;
378
379            let j = n4 - 1 - i;
380            let im1 = f2[j].r;
381            let re1 = f2[j].i;
382            let t0_1 = trig[j];
383            let t1_1 = trig[n4 + j];
384
385            let yr1 = re1 * t0_1 + im1 * t1_1;
386            let yi1 = re1 * t1_1 - im1 * t0_1;
387
388            output[overlap2 + 2 * i] = yr0;
389            output[overlap2 + n2 - 1 - 2 * i] = yi0;
390            output[overlap2 + n2 - 2 - 2 * i] = yr1;
391            output[overlap2 + 2 * i + 1] = yi1;
392        }
393
394        #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
395        unsafe {
396            if std::arch::is_x86_feature_detected!("avx") {
397                mdct_tdac_avx(output, window, overlap);
398            } else {
399                for i in 0..overlap2 {
400                    let x1 = output[overlap - 1 - i];
401                    let x2 = output[i];
402                    let wp1 = window[i];
403                    let wp2 = window[overlap - 1 - i];
404
405                    output[i] = x2 * wp2 - x1 * wp1;
406                    output[overlap - 1 - i] = x2 * wp1 + x1 * wp2;
407                }
408            }
409        }
410        #[cfg(all(
411            not(any(target_arch = "x86", target_arch = "x86_64")),
412            target_arch = "aarch64"
413        ))]
414        {
415            mdct_tdac_neon(output, window, overlap);
416        }
417        #[cfg(all(
418            not(any(target_arch = "x86", target_arch = "x86_64")),
419            not(target_arch = "aarch64")
420        ))]
421        for i in 0..overlap2 {
422            let x1 = output[overlap - 1 - i];
423            let x2 = output[i];
424            let wp1 = window[i];
425            let wp2 = window[overlap - 1 - i];
426
427            output[i] = x2 * wp2 - x1 * wp1;
428            output[overlap - 1 - i] = x2 * wp1 + x1 * wp2;
429        }
430    }
431}
432
433#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
434#[target_feature(enable = "avx")]
435unsafe fn mdct_pre_rotation_avx(
436    f: &[f32],
437    f2: &mut [KissCpx],
438    trig: &[f32],
439    bitrev: &[i16],
440    n4: usize,
441    scale: f32,
442) {
443    for i in 0..n4 {
444        let re = f[2 * i];
445        let im = f[2 * i + 1];
446        let t0 = trig[i];
447        let t1 = trig[n4 + i];
448
449        let yr = re * t0 - im * t1;
450        let yi = im * t0 + re * t1;
451
452        f2[bitrev[i] as usize] = KissCpx::new(yr * scale, yi * scale);
453    }
454}
455
456#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
457#[target_feature(enable = "avx")]
458unsafe fn mdct_post_rotation_avx(
459    f2: &[KissCpx],
460    trig: &[f32],
461    output: &mut [f32],
462    n4: usize,
463    n2: usize,
464    stride: usize,
465) {
466    for i in 0..n4 {
467        let fp = &f2[i];
468        let t0 = trig[i];
469        let t1 = trig[n4 + i];
470
471        let yr = fp.i * t1 - fp.r * t0;
472        let yi = fp.r * t1 + fp.i * t0;
473
474        output[i * 2 * stride] = yr;
475        output[stride * (n2 - 1 - 2 * i)] = yi;
476    }
477}
478
479#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
480#[target_feature(enable = "avx")]
481unsafe fn mdct_backward_pre_rotation_avx(
482    input: &[f32],
483    f2: &mut [KissCpx],
484    trig: &[f32],
485    bitrev: &[i16],
486    n4: usize,
487    n2: usize,
488    stride: usize,
489) {
490    for i in 0..n4 {
491        let rev = bitrev[i] as usize;
492        let x1 = input[2 * i * stride];
493        let x2 = input[stride * (n2 - 1 - 2 * i)];
494        let t0 = trig[i];
495        let t1 = trig[n4 + i];
496
497        let yr = x2 * t0 + x1 * t1;
498        let yi = x1 * t0 - x2 * t1;
499
500        f2[rev] = KissCpx::new(yi, yr);
501    }
502}
503
504#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
505#[target_feature(enable = "avx")]
506unsafe fn mdct_backward_post_rotation_avx(
507    f2: &[KissCpx],
508    trig: &[f32],
509    output: &mut [f32],
510    n4: usize,
511    n2: usize,
512    overlap2: usize,
513) {
514    for i in 0..((n4 + 1) >> 1) {
515        let im0 = f2[i].r;
516        let re0 = f2[i].i;
517        let t0_0 = trig[i];
518        let t1_0 = trig[n4 + i];
519
520        let yr0 = re0 * t0_0 + im0 * t1_0;
521        let yi0 = re0 * t1_0 - im0 * t0_0;
522
523        let j = n4 - 1 - i;
524        let im1 = f2[j].r;
525        let re1 = f2[j].i;
526        let t0_1 = trig[j];
527        let t1_1 = trig[n4 + j];
528
529        let yr1 = re1 * t0_1 + im1 * t1_1;
530        let yi1 = re1 * t1_1 - im1 * t0_1;
531
532        output[overlap2 + 2 * i] = yr0;
533        output[overlap2 + n2 - 1 - 2 * i] = yi0;
534        output[overlap2 + n2 - 2 - 2 * i] = yr1;
535        output[overlap2 + 2 * i + 1] = yi1;
536    }
537}
538
539#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
540#[target_feature(enable = "avx")]
541unsafe fn mdct_tdac_avx(output: &mut [f32], window: &[f32], overlap: usize) {
542    use std::arch::x86_64::*;
543
544    let overlap2 = overlap / 2;
545    let mut i = 0usize;
546
547    while i + 8 <= overlap2 {
548        let x2 = _mm256_loadu_ps(output.as_ptr().add(i));
549
550        let mut x1_tmp = [0f32; 8];
551        let mut w2_tmp = [0f32; 8];
552        for j in 0..8 {
553            x1_tmp[j] = output[overlap - 1 - (i + j)];
554            w2_tmp[j] = window[overlap - 1 - (i + j)];
555        }
556        let x1 = _mm256_loadu_ps(x1_tmp.as_ptr());
557
558        let w1 = _mm256_loadu_ps(window.as_ptr().add(i));
559        let w2 = _mm256_loadu_ps(w2_tmp.as_ptr());
560
561        let out_fwd = _mm256_sub_ps(_mm256_mul_ps(x2, w2), _mm256_mul_ps(x1, w1));
562        let out_rev = _mm256_add_ps(_mm256_mul_ps(x2, w1), _mm256_mul_ps(x1, w2));
563
564        _mm256_storeu_ps(output.as_mut_ptr().add(i), out_fwd);
565
566        let mut out_rev_tmp = [0f32; 8];
567        _mm256_storeu_ps(out_rev_tmp.as_mut_ptr(), out_rev);
568        for j in 0..8 {
569            output[overlap - 1 - (i + j)] = out_rev_tmp[j];
570        }
571
572        i += 8;
573    }
574
575    for i in i..overlap2 {
576        let x1 = output[overlap - 1 - i];
577        let x2 = output[i];
578        let wp1 = window[i];
579        let wp2 = window[overlap - 1 - i];
580        output[i] = x2 * wp2 - x1 * wp1;
581        output[overlap - 1 - i] = x2 * wp1 + x1 * wp2;
582    }
583}
584
585#[cfg(target_arch = "aarch64")]
586#[inline(always)]
587fn mdct_pre_rotation_neon(
588    f: &[f32],
589    f2: &mut [KissCpx],
590    trig: &[f32],
591    bitrev: &[i16],
592    n4: usize,
593    scale: f32,
594) {
595    use std::arch::aarch64::*;
596
597    unsafe {
598        let vscale = vdupq_n_f32(scale);
599        let f_ptr = f.as_ptr();
600        let trig_ptr = trig.as_ptr();
601        let bitrev_ptr = bitrev.as_ptr();
602        let f2_ptr = f2.as_mut_ptr() as *mut f32;
603
604        let n4_vec = n4 & !3;
605        let mut i = 0;
606
607        while i < n4_vec {
608            let t0 = vld1q_f32(trig_ptr.add(i));
609            let t1 = vld1q_f32(trig_ptr.add(n4 + i));
610
611            let f0 = vld1q_f32(f_ptr.add(2 * i));
612            let f1 = vld1q_f32(f_ptr.add(2 * i + 4));
613
614            let even_odd = vuzpq_f32(f0, f1);
615            let re_v = even_odd.0;
616            let im_v = even_odd.1;
617
618            let yr = vsubq_f32(vmulq_f32(re_v, t0), vmulq_f32(im_v, t1));
619            let yi = vaddq_f32(vmulq_f32(im_v, t0), vmulq_f32(re_v, t1));
620
621            let yr = vmulq_f32(yr, vscale);
622            let yi = vmulq_f32(yi, vscale);
623
624            let yr_arr: [f32; 4] = std::mem::transmute(yr);
625            let yi_arr: [f32; 4] = std::mem::transmute(yi);
626
627            for j in 0..4 {
628                let rev = *bitrev_ptr.add(i + j) as usize;
629                *f2_ptr.add(2 * rev) = yr_arr[j];
630                *f2_ptr.add(2 * rev + 1) = yi_arr[j];
631            }
632
633            i += 4;
634        }
635
636        for i in n4_vec..n4 {
637            let re = *f_ptr.add(2 * i);
638            let im = *f_ptr.add(2 * i + 1);
639            let t0 = *trig_ptr.add(i);
640            let t1 = *trig_ptr.add(n4 + i);
641            let yr = re * t0 - im * t1;
642            let yi = im * t0 + re * t1;
643            let rev = *bitrev_ptr.add(i) as usize;
644            *f2_ptr.add(2 * rev) = yr * scale;
645            *f2_ptr.add(2 * rev + 1) = yi * scale;
646        }
647    }
648}
649
650#[cfg(target_arch = "aarch64")]
651#[inline(always)]
652fn mdct_post_rotation_neon(
653    f2: &[KissCpx],
654    trig: &[f32],
655    output: &mut [f32],
656    n4: usize,
657    n2: usize,
658    stride: usize,
659) {
660    use std::arch::aarch64::*;
661
662    if stride > 1 {
663        for i in 0..n4 {
664            let fp = &f2[i];
665            let t0 = trig[i];
666            let t1 = trig[n4 + i];
667            let yr = fp.i * t1 - fp.r * t0;
668            let yi = fp.r * t1 + fp.i * t0;
669            output[i * 2 * stride] = yr;
670            output[stride * (n2 - 1 - 2 * i)] = yi;
671        }
672        return;
673    }
674
675    unsafe {
676        let f2_ptr = f2.as_ptr() as *const f32;
677        let trig_ptr = trig.as_ptr();
678        let out_ptr = output.as_mut_ptr();
679
680        let n4_vec = n4 & !3;
681        let mut i = 0;
682
683        while i < n4_vec {
684            let c0 = vld1q_f32(f2_ptr.add(2 * i));
685            let c1 = vld1q_f32(f2_ptr.add(2 * i + 4));
686
687            let t0 = vld1q_f32(trig_ptr.add(i));
688            let t1 = vld1q_f32(trig_ptr.add(n4 + i));
689
690            let ri = vuzpq_f32(c0, c1);
691            let r_v = ri.0;
692            let i_v = ri.1;
693
694            let yr = vsubq_f32(vmulq_f32(i_v, t1), vmulq_f32(r_v, t0));
695
696            let yi = vaddq_f32(vmulq_f32(r_v, t1), vmulq_f32(i_v, t0));
697
698            let yr_arr: [f32; 4] = std::mem::transmute(yr);
699            let yi_arr: [f32; 4] = std::mem::transmute(yi);
700
701            for j in 0..4 {
702                *out_ptr.add((i + j) * 2) = yr_arr[j];
703                *out_ptr.add(n2 - 1 - 2 * (i + j)) = yi_arr[j];
704            }
705
706            i += 4;
707        }
708
709        for i in n4_vec..n4 {
710            let fp = &f2[i];
711            let t0 = trig[i];
712            let t1 = trig[n4 + i];
713            let yr = fp.i * t1 - fp.r * t0;
714            let yi = fp.r * t1 + fp.i * t0;
715            output[i * 2] = yr;
716            output[n2 - 1 - 2 * i] = yi;
717        }
718    }
719}
720
721#[cfg(target_arch = "aarch64")]
722#[inline(always)]
723fn mdct_backward_pre_rotation_neon(
724    input: &[f32],
725    f2: &mut [KissCpx],
726    trig: &[f32],
727    bitrev: &[i16],
728    n4: usize,
729    n2: usize,
730    stride: usize,
731) {
732    use std::arch::aarch64::*;
733
734    if stride != 1 {
735        for i in 0..n4 {
736            let rev = bitrev[i] as usize;
737            let x1 = input[2 * i * stride];
738            let x2 = input[stride * (n2 - 1 - 2 * i)];
739            let t0 = trig[i];
740            let t1 = trig[n4 + i];
741            let yr = x2 * t0 + x1 * t1;
742            let yi = x1 * t0 - x2 * t1;
743            f2[rev] = KissCpx::new(yi, yr);
744        }
745        return;
746    }
747
748    unsafe {
749        let in_ptr = input.as_ptr();
750        let trig_ptr = trig.as_ptr();
751        let bitrev_ptr = bitrev.as_ptr();
752        let f2_ptr = f2.as_mut_ptr() as *mut f32;
753
754        let n4_vec = n4 & !3;
755        let mut i = 0;
756
757        while i < n4_vec {
758            let f0 = vld1q_f32(in_ptr.add(2 * i));
759            let f1 = vld1q_f32(in_ptr.add(2 * i + 4));
760            let deint_x1 = vuzpq_f32(f0, f1);
761            let x1_v = deint_x1.0;
762
763            let g0 = vld1q_f32(in_ptr.add(n2 - 7 - 2 * i));
764            let g1 = vld1q_f32(in_ptr.add(n2 - 3 - 2 * i));
765            let deint_x2 = vuzpq_f32(g0, g1);
766
767            let x2_raw = deint_x2.0;
768            let x2_v = vrev64q_f32(x2_raw);
769            let x2_v = vextq_f32(x2_v, x2_v, 2);
770
771            let t0 = vld1q_f32(trig_ptr.add(i));
772            let t1 = vld1q_f32(trig_ptr.add(n4 + i));
773
774            let yr = vaddq_f32(vmulq_f32(x2_v, t0), vmulq_f32(x1_v, t1));
775            let yi = vsubq_f32(vmulq_f32(x1_v, t0), vmulq_f32(x2_v, t1));
776
777            let yr_arr: [f32; 4] = std::mem::transmute(yr);
778            let yi_arr: [f32; 4] = std::mem::transmute(yi);
779
780            for j in 0..4 {
781                let rev = *bitrev_ptr.add(i + j) as usize;
782                *f2_ptr.add(2 * rev) = yi_arr[j];
783                *f2_ptr.add(2 * rev + 1) = yr_arr[j];
784            }
785
786            i += 4;
787        }
788
789        for i in n4_vec..n4 {
790            let rev = *bitrev_ptr.add(i) as usize;
791            let x1 = *in_ptr.add(2 * i);
792            let x2 = *in_ptr.add(n2 - 1 - 2 * i);
793            let t0 = *trig_ptr.add(i);
794            let t1 = *trig_ptr.add(n4 + i);
795            let yr = x2 * t0 + x1 * t1;
796            let yi = x1 * t0 - x2 * t1;
797            *f2_ptr.add(2 * rev) = yi;
798            *f2_ptr.add(2 * rev + 1) = yr;
799        }
800    }
801}
802
803#[cfg(target_arch = "aarch64")]
804#[inline(always)]
805fn mdct_backward_post_rotation_neon(
806    f2: &[KissCpx],
807    trig: &[f32],
808    output: &mut [f32],
809    n4: usize,
810    n2: usize,
811    overlap2: usize,
812) {
813    unsafe {
814        let trig_ptr = trig.as_ptr();
815        let out_base = output.as_mut_ptr().add(overlap2);
816
817        let half = (n4 + 1) >> 1;
818
819        let mut i = 0;
820        while i + 1 < half {
821            let j0 = n4 - 1 - i;
822            let j1 = n4 - 1 - (i + 1);
823
824            let re0 = f2[i].i;
825            let im0 = f2[i].r;
826            let t0_0 = *trig_ptr.add(i);
827            let t1_0 = *trig_ptr.add(n4 + i);
828            let yr0 = re0 * t0_0 + im0 * t1_0;
829            let yi0 = re0 * t1_0 - im0 * t0_0;
830
831            let im1 = f2[j0].r;
832            let re1 = f2[j0].i;
833            let t0_1 = *trig_ptr.add(j0);
834            let t1_1 = *trig_ptr.add(n4 + j0);
835            let yr1 = re1 * t0_1 + im1 * t1_1;
836            let yi1 = re1 * t1_1 - im1 * t0_1;
837
838            *out_base.add(2 * i) = yr0;
839            *out_base.add(n2 - 1 - 2 * i) = yi0;
840            *out_base.add(n2 - 2 - 2 * i) = yr1;
841            *out_base.add(2 * i + 1) = yi1;
842
843            let re0b = f2[i + 1].i;
844            let im0b = f2[i + 1].r;
845            let t0_0b = *trig_ptr.add(i + 1);
846            let t1_0b = *trig_ptr.add(n4 + i + 1);
847            let yr0b = re0b * t0_0b + im0b * t1_0b;
848            let yi0b = re0b * t1_0b - im0b * t0_0b;
849
850            let im1b = f2[j1].r;
851            let re1b = f2[j1].i;
852            let t0_1b = *trig_ptr.add(j1);
853            let t1_1b = *trig_ptr.add(n4 + j1);
854            let yr1b = re1b * t0_1b + im1b * t1_1b;
855            let yi1b = re1b * t1_1b - im1b * t0_1b;
856
857            *out_base.add(2 * (i + 1)) = yr0b;
858            *out_base.add(n2 - 1 - 2 * (i + 1)) = yi0b;
859            *out_base.add(n2 - 2 - 2 * (i + 1)) = yr1b;
860            *out_base.add(2 * (i + 1) + 1) = yi1b;
861
862            i += 2;
863        }
864
865        if i < half {
866            let j = n4 - 1 - i;
867            let im0 = f2[i].r;
868            let re0 = f2[i].i;
869            let t0_0 = *trig_ptr.add(i);
870            let t1_0 = *trig_ptr.add(n4 + i);
871            let yr0 = re0 * t0_0 + im0 * t1_0;
872            let yi0 = re0 * t1_0 - im0 * t0_0;
873
874            let im1 = f2[j].r;
875            let re1 = f2[j].i;
876            let t0_1 = *trig_ptr.add(j);
877            let t1_1 = *trig_ptr.add(n4 + j);
878            let yr1 = re1 * t0_1 + im1 * t1_1;
879            let yi1 = re1 * t1_1 - im1 * t0_1;
880
881            *out_base.add(2 * i) = yr0;
882            *out_base.add(n2 - 1 - 2 * i) = yi0;
883            *out_base.add(n2 - 2 - 2 * i) = yr1;
884            *out_base.add(2 * i + 1) = yi1;
885        }
886    }
887}
888
889#[cfg(target_arch = "aarch64")]
890#[inline(always)]
891fn mdct_tdac_neon(output: &mut [f32], window: &[f32], overlap: usize) {
892    use std::arch::aarch64::*;
893
894    let overlap2 = overlap / 2;
895    if overlap2 < 4 {
896        for i in 0..overlap2 {
897            let x1 = output[overlap - 1 - i];
898            let x2 = output[i];
899            let wp1 = window[i];
900            let wp2 = window[overlap - 1 - i];
901            output[i] = x2 * wp2 - x1 * wp1;
902            output[overlap - 1 - i] = x2 * wp1 + x1 * wp2;
903        }
904        return;
905    }
906
907    unsafe {
908        let out_ptr = output.as_mut_ptr();
909        let win_ptr = window.as_ptr();
910        let n4 = overlap2 & !3;
911        let mut i = 0;
912
913        while i < n4 {
914            let x2_fwd = vld1q_f32(out_ptr.add(i));
915            let x1_rev = vld1q_f32(out_ptr.add(overlap - 4 - i));
916
917            let x1 = vrev64q_f32(x1_rev);
918            let x1 = vextq_f32(x1, x1, 2);
919
920            let wp1_fwd = vld1q_f32(win_ptr.add(i));
921            let wp2_rev = vld1q_f32(win_ptr.add(overlap - 4 - i));
922            let wp2 = vrev64q_f32(wp2_rev);
923            let wp2 = vextq_f32(wp2, wp2, 2);
924            let wp1 = wp1_fwd;
925
926            let out_fwd = vsubq_f32(vmulq_f32(x2_fwd, wp2), vmulq_f32(x1, wp1));
927
928            let out_rev = vaddq_f32(vmulq_f32(x2_fwd, wp1), vmulq_f32(x1, wp2));
929
930            let out_rev = vrev64q_f32(out_rev);
931            let out_rev = vextq_f32(out_rev, out_rev, 2);
932
933            vst1q_f32(out_ptr.add(i), out_fwd);
934            vst1q_f32(out_ptr.add(overlap - 4 - i), out_rev);
935
936            i += 4;
937        }
938
939        for i in n4..overlap2 {
940            let x1 = output[overlap - 1 - i];
941            let x2 = output[i];
942            output[i] = x2 * window[overlap - 1 - i] - x1 * window[i];
943            output[overlap - 1 - i] = x2 * window[i] + x1 * window[overlap - 1 - i];
944        }
945    }
946}
947
948#[cfg(test)]
949mod mdct_tests {
950    #[test]
951    fn test_mdct_backward_transient_no_blowup() {
952        let mode = crate::modes::default_mode();
953        let shift = 3;
954        let n = mode.mdct.n >> shift; // 120
955        let overlap = mode.overlap; // 120
956        let stride = 8;
957
958        let frame_size = 960usize;
959        let mut freq = vec![0.0f32; frame_size];
960        for i in 0..frame_size {
961            freq[i] = ((i as f32) * 0.01).sin() * 10.0;
962        }
963
964        let out_len = n + overlap; // 240
965        let mut output0 = vec![0.0f32; out_len];
966        let mut output1 = vec![0.0f32; out_len];
967
968        mode.mdct.backward(
969            &freq[0..],
970            &mut output0,
971            mode.window,
972            overlap,
973            shift,
974            stride,
975        );
976        mode.mdct.backward(
977            &freq[1..],
978            &mut output1,
979            mode.window,
980            overlap,
981            shift,
982            stride,
983        );
984
985        let max0 = output0.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
986        let max1 = output1.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
987        eprintln!("sub0 max={} sub1 max={}", max0, max1);
988        eprintln!("sub0[60..70]={:?}", &output0[60..70]);
989        eprintln!("sub1[60..70]={:?}", &output1[60..70]);
990
991        assert!(max0.abs() < 500.0, "sub0 blowup: {}", max0);
992        assert!(max1.abs() < 500.0, "sub1 blowup: {}", max1);
993    }
994
995    #[test]
996    fn test_mdct_backward_stride1_neon_matches_scalar() {
997        let mode = crate::modes::default_mode();
998        let shift = 0; // non-transient full-size MDCT
999        let n = mode.mdct.n >> shift; // 1920
1000        let n2 = n / 2; // 960
1001        let n4 = n / 4; // 480
1002        let overlap = mode.overlap; // 120
1003        let overlap2 = overlap / 2; // 60
1004        let stride = 1;
1005
1006        let freq_len = n2;
1007        let mut freq = vec![0.0f32; freq_len + 4];
1008        for i in 0..freq_len {
1009            freq[i] = ((i as f32) * 0.01).sin() * 4577.0;
1010        }
1011
1012        let out_len = overlap2 + n2; // 60 + 960 = 1020
1013        let mut output_hw = vec![0.0f32; out_len + 100];
1014        mode.mdct.backward(
1015            &freq[..],
1016            &mut output_hw,
1017            mode.window,
1018            overlap,
1019            shift,
1020            stride,
1021        );
1022
1023        let st = mode.mdct.kfft[shift].as_ref().unwrap();
1024        let (trig, _) = mode.mdct.get_trig(shift);
1025
1026        use crate::kiss_fft::KissCpx;
1027        let mut f2 = vec![KissCpx::new(0.0, 0.0); n4];
1028        for i in 0..n4 {
1029            let rev = st.bitrev[i] as usize;
1030            let x1 = freq[2 * i * stride];
1031            let x2 = freq[stride * (n2 - 1 - 2 * i)];
1032            let t0 = trig[i];
1033            let t1 = trig[n4 + i];
1034            let yr = x2 * t0 + x1 * t1;
1035            let yi = x1 * t0 - x2 * t1;
1036            f2[rev] = KissCpx::new(yi, yr);
1037        }
1038        crate::kiss_fft::opus_fft_impl(st, &mut f2);
1039
1040        let mut output_scalar = vec![0.0f32; out_len + 100];
1041        for i in 0..((n4 + 1) >> 1) {
1042            let im0 = f2[i].r;
1043            let re0 = f2[i].i;
1044            let t0_0 = trig[i];
1045            let t1_0 = trig[n4 + i];
1046            let yr0 = re0 * t0_0 + im0 * t1_0;
1047            let yi0 = re0 * t1_0 - im0 * t0_0;
1048            let j = n4 - 1 - i;
1049            let im1 = f2[j].r;
1050            let re1 = f2[j].i;
1051            let t0_1 = trig[j];
1052            let t1_1 = trig[n4 + j];
1053            let yr1 = re1 * t0_1 + im1 * t1_1;
1054            let yi1 = re1 * t1_1 - im1 * t0_1;
1055            output_scalar[overlap2 + 2 * i] = yr0;
1056            output_scalar[overlap2 + n2 - 1 - 2 * i] = yi0;
1057            output_scalar[overlap2 + n2 - 2 - 2 * i] = yr1;
1058            output_scalar[overlap2 + 2 * i + 1] = yi1;
1059        }
1060        // TDAC
1061        for i in 0..overlap2 {
1062            let x1 = output_scalar[overlap - 1 - i];
1063            let x2 = output_scalar[i];
1064            let wp1 = mode.window[i];
1065            let wp2 = mode.window[overlap - 1 - i];
1066            output_scalar[i] = x2 * wp2 - x1 * wp1;
1067            output_scalar[overlap - 1 - i] = x2 * wp1 + x1 * wp2;
1068        }
1069
1070        let max_diff = output_hw[..out_len]
1071            .iter()
1072            .zip(output_scalar[..out_len].iter())
1073            .map(|(a, b)| (a - b).abs())
1074            .fold(0.0f32, f32::max);
1075        assert!(
1076            max_diff < 0.5,
1077            "stride=1 NEON vs scalar mismatch: max_diff={}",
1078            max_diff
1079        );
1080    }
1081
1082    #[test]
1083    fn test_mdct_backward_neon_matches_scalar() {
1084        let mode = crate::modes::default_mode();
1085        let shift = 3;
1086        let n = mode.mdct.n >> shift; // 240
1087        let n2 = n / 2; // 120
1088        let n4 = n / 4; // 60
1089        let overlap = mode.overlap; // 120
1090        let overlap2 = overlap / 2; // 60
1091        let stride = 8;
1092
1093        // Build a realistic freq vector (sine wave @ 440Hz)
1094        let frame_size = 960usize;
1095        let mut freq = vec![0.0f32; frame_size];
1096        for i in 0..frame_size {
1097            freq[i] = ((i as f32) * 0.01).sin() * 200.0;
1098        }
1099
1100        let out_len = n + overlap; // 360
1101        let mut output_hw = vec![0.0f32; out_len];
1102        mode.mdct.backward(
1103            &freq[0..],
1104            &mut output_hw,
1105            mode.window,
1106            overlap,
1107            shift,
1108            stride,
1109        );
1110
1111        // Scalar reference
1112        let st = mode.mdct.kfft[shift].as_ref().unwrap();
1113        let (trig, _) = mode.mdct.get_trig(shift);
1114
1115        use crate::kiss_fft::KissCpx;
1116        let mut f2 = vec![KissCpx::new(0.0, 0.0); n4];
1117        for i in 0..n4 {
1118            let rev = st.bitrev[i] as usize;
1119            let x1 = freq[2 * i * stride];
1120            let x2 = freq[stride * (n2 - 1 - 2 * i)];
1121            let t0 = trig[i];
1122            let t1 = trig[n4 + i];
1123            let yr = x2 * t0 + x1 * t1;
1124            let yi = x1 * t0 - x2 * t1;
1125            f2[rev] = KissCpx::new(yi, yr);
1126        }
1127        crate::kiss_fft::opus_fft_impl(st, &mut f2);
1128
1129        let mut output_scalar = vec![0.0f32; out_len];
1130        for i in 0..((n4 + 1) >> 1) {
1131            let im0 = f2[i].r;
1132            let re0 = f2[i].i;
1133            let t0_0 = trig[i];
1134            let t1_0 = trig[n4 + i];
1135            let yr0 = re0 * t0_0 + im0 * t1_0;
1136            let yi0 = re0 * t1_0 - im0 * t0_0;
1137            let j = n4 - 1 - i;
1138            let im1 = f2[j].r;
1139            let re1 = f2[j].i;
1140            let t0_1 = trig[j];
1141            let t1_1 = trig[n4 + j];
1142            let yr1 = re1 * t0_1 + im1 * t1_1;
1143            let yi1 = re1 * t1_1 - im1 * t0_1;
1144            output_scalar[overlap2 + 2 * i] = yr0;
1145            output_scalar[overlap2 + n2 - 1 - 2 * i] = yi0;
1146            output_scalar[overlap2 + n2 - 2 - 2 * i] = yr1;
1147            output_scalar[overlap2 + 2 * i + 1] = yi1;
1148        }
1149        // TDAC
1150        for i in 0..overlap2 {
1151            let x1 = output_scalar[overlap - 1 - i];
1152            let x2 = output_scalar[i];
1153            let wp1 = mode.window[i];
1154            let wp2 = mode.window[overlap - 1 - i];
1155            output_scalar[i] = x2 * wp2 - x1 * wp1;
1156            output_scalar[overlap - 1 - i] = x2 * wp1 + x1 * wp2;
1157        }
1158
1159        for i in 0..out_len {
1160            let diff = (output_hw[i] - output_scalar[i]).abs();
1161            if diff > 1e-3 {
1162                eprintln!(
1163                    "Mismatch at output[{}]: hw={} scalar={} diff={}",
1164                    i, output_hw[i], output_scalar[i], diff
1165                );
1166            }
1167        }
1168        let max_diff = output_hw
1169            .iter()
1170            .zip(output_scalar.iter())
1171            .map(|(a, b)| (a - b).abs())
1172            .fold(0.0f32, f32::max);
1173        assert!(
1174            max_diff < 0.1,
1175            "NEON/HW vs scalar mismatch: max_diff={}",
1176            max_diff
1177        );
1178    }
1179}