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 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 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; let overlap = mode.overlap; 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; 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; let n = mode.mdct.n >> shift; let n2 = n / 2; let n4 = n / 4; let overlap = mode.overlap; let overlap2 = overlap / 2; 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; 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 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; let n2 = n / 2; let n4 = n / 4; let overlap = mode.overlap; let overlap2 = overlap / 2; let stride = 8;
1092
1093 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; 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 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 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}