Skip to main content

opus_rs/
bands.rs

1use crate::modes::CeltMode;
2use crate::pvq::*;
3use crate::range_coder::RangeCoder;
4use crate::rate::{BITRES, bits2pulses, get_pulses, pulses2bits};
5use crate::tell_frac_inline;
6
7#[cfg(not(feature = "std"))]
8use crate::compat::Math;
9
10const MIN_STEREO_ENERGY: f32 = 1e-10;
11
12pub struct BandCtx<'a> {
13    pub encode: bool,
14    pub m: &'a CeltMode,
15    pub i: usize,
16    pub band_e: &'a [f32],
17    pub rc: &'a mut RangeCoder,
18    pub spread: i32,
19    pub remaining_bits: i32,
20    pub resynth: bool,
21    pub tf_change: i32,
22    pub intensity: usize,
23    pub theta_round: i32,
24    pub avoid_split_noise: bool,
25    pub arch: i32,
26    pub disable_inv: bool,
27    pub seed: u32,
28}
29
30#[inline]
31fn bitexact_cos(x: i16) -> i16 {
32    #[inline(always)]
33    fn frac_mul16(a: i16, b: i16) -> i16 {
34        ((16384i32 + (a as i32) * (b as i32)) >> 15) as i16
35    }
36
37    let tmp = (4096i32 + (x as i32) * (x as i32)) >> 13;
38    let x2 = tmp as i16;
39    let x2 = (32767 - x2 as i32
40        + frac_mul16(x2, -7651 + frac_mul16(x2, 8277 + frac_mul16(-626, x2))) as i32)
41        as i16;
42    1 + x2
43}
44
45#[inline]
46pub fn bitexact_log2tan(isin: i32, icos: i32) -> i32 {
47    let ec_ilog = |x: u32| -> i32 {
48        if x == 0 {
49            0
50        } else {
51            32 - x.leading_zeros() as i32
52        }
53    };
54    let lc = ec_ilog(icos.max(0) as u32);
55    let ls = ec_ilog(isin.max(0) as u32);
56    let icos_shifted = if lc > 0 {
57        icos.max(0) << (15 - lc).max(0)
58    } else {
59        0
60    };
61    let isin_shifted = if ls > 0 {
62        isin.max(0) << (15 - ls).max(0)
63    } else {
64        0
65    };
66    let fract_mul = |a: i32, b: i32| -> i32 { (a * b + 16384) >> 15 };
67    (ls - lc) * (1 << 11) + fract_mul(isin_shifted, fract_mul(isin_shifted, -2597) + 7932)
68        - fract_mul(icos_shifted, fract_mul(icos_shifted, -2597) + 7932)
69}
70
71#[inline(always)]
72fn celt_sudiv(n: i32, d: i32) -> i32 {
73    n / d
74}
75
76#[inline]
77fn isqrt32(mut val: u32) -> u32 {
78    let mut g = 0u32;
79    let mut bshift = ((32 - val.leading_zeros()) as i32 - 1) >> 1;
80    let mut b = 1u32 << bshift;
81    while bshift >= 0 {
82        let t = (((g << 1) + b) as u64) << bshift;
83        if t <= val as u64 {
84            g += b;
85            val -= t as u32;
86        }
87        b >>= 1;
88        bshift -= 1;
89    }
90    g
91}
92
93pub const SPREAD_NONE: i32 = 0;
94pub const SPREAD_LIGHT: i32 = 1;
95pub const SPREAD_NORMAL: i32 = 2;
96pub const SPREAD_AGGRESSIVE: i32 = 3;
97
98#[allow(clippy::too_many_arguments)]
99pub fn spreading_decision(
100    m: &CeltMode,
101    x_buf: &[f32],
102    average: &mut i32,
103    last_decision: i32,
104    hf_average: &mut i32,
105    tapset_decision: &mut i32,
106    update_hf: bool,
107    end: usize,
108    channels: usize,
109    m_val: usize,
110    spread_weight: &[i32],
111) -> i32 {
112    let mut sum = 0;
113    let mut nb_bands = 0;
114    let n0 = m_val * m.short_mdct_size;
115    let mut hf_sum = 0;
116
117    if m_val * (m.e_bands[end] as usize - m.e_bands[end - 1] as usize) <= 8 {
118        return SPREAD_NONE;
119    }
120
121    for c in 0..channels {
122        for (i, &sw) in spread_weight[..end].iter().enumerate() {
123            let n = m_val * (m.e_bands[i + 1] as usize - m.e_bands[i] as usize);
124            if n <= 8 {
125                continue;
126            }
127
128            let mut tcount = [0; 3];
129            let offset = m_val * m.e_bands[i] as usize + c * n0;
130            let x = &x_buf[offset..offset + n];
131
132            for xv in x.iter().copied() {
133                let x2n = xv * xv * (n as f32);
134                if x2n < 0.25 {
135                    tcount[0] += 1;
136                }
137                if x2n < 0.0625 {
138                    tcount[1] += 1;
139                }
140                if x2n < 0.015625 {
141                    tcount[2] += 1;
142                }
143            }
144
145            if i > m.nb_ebands - 4 {
146                hf_sum += 32 * (tcount[1] + tcount[0]) / (n as i32);
147            }
148
149            let tmp = (if 2 * tcount[2] >= (n as i32) { 1 } else { 0 })
150                + (if 2 * tcount[1] >= (n as i32) { 1 } else { 0 })
151                + (if 2 * tcount[0] >= (n as i32) { 1 } else { 0 });
152            sum += tmp * sw;
153            nb_bands += sw;
154        }
155    }
156
157    if update_hf {
158        if hf_sum > 0 {
159            hf_sum /= (channels as i32) * (4 - m.nb_ebands as i32 + end as i32);
160        }
161        *hf_average = (*hf_average + hf_sum) >> 1;
162        hf_sum = *hf_average;
163
164        if *tapset_decision == 2 {
165            hf_sum += 4;
166        } else if *tapset_decision == 0 {
167            hf_sum -= 4;
168        }
169
170        if hf_sum > 22 {
171            *tapset_decision = 2;
172        } else if hf_sum > 18 {
173            *tapset_decision = 1;
174        } else {
175            *tapset_decision = 0;
176        }
177    }
178
179    if nb_bands == 0 {
180        return SPREAD_NORMAL;
181    }
182
183    let mut sum_scaled = (sum << 8) / nb_bands;
184    sum_scaled = (sum_scaled + *average) >> 1;
185    *average = sum_scaled;
186
187    let sum_final = (3 * sum_scaled + (((3 - last_decision) << 7) + 64) + 2) >> 2;
188
189    if sum_final < 80 {
190        SPREAD_AGGRESSIVE
191    } else if sum_final < 256 {
192        SPREAD_NORMAL
193    } else if sum_final < 384 {
194        SPREAD_LIGHT
195    } else {
196        SPREAD_NONE
197    }
198}
199
200pub fn haar1(x: &mut [f32], n0: usize, stride: usize) {
201    #[cfg(target_arch = "aarch64")]
202    {
203        if stride == 1 && n0 >= 64 {
204            haar1_neon(x, n0);
205        } else {
206            haar1_scalar(x, n0, stride);
207        }
208        return;
209    }
210    #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
211    unsafe {
212        if stride == 1 && n0 >= 16 && crate::compat::x86_has_avx() {
213            haar1_avx(x, n0);
214            return;
215        }
216    }
217    #[cfg(not(target_arch = "aarch64"))]
218    haar1_scalar(x, n0, stride);
219}
220
221#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
222#[target_feature(enable = "avx")]
223unsafe fn haar1_avx(x: &mut [f32], n0: usize) {
224    use core::arch::x86_64::*;
225    let n = n0 >> 1;
226    let scale = _mm256_set1_ps(core::f32::consts::FRAC_1_SQRT_2);
227    let mut j = 0;
228    while j + 8 <= n {
229        let ptr = x.as_mut_ptr().add(2 * j);
230        let a = _mm256_loadu_ps(ptr);
231        let b = _mm256_loadu_ps(ptr.add(4));
232
233        let t0 = _mm256_unpacklo_ps(a, b);
234        let t1 = _mm256_unpackhi_ps(a, b);
235
236        let even = _mm256_unpacklo_ps(t0, t1);
237        let odd = _mm256_unpackhi_ps(t0, t1);
238
239        let sum = _mm256_mul_ps(_mm256_add_ps(even, odd), scale);
240        let diff = _mm256_mul_ps(_mm256_sub_ps(even, odd), scale);
241
242        let r0 = _mm256_unpacklo_ps(sum, diff);
243        let r1 = _mm256_unpackhi_ps(sum, diff);
244
245        let out0 = _mm256_permute2f128_ps(r0, r1, 0x20);
246        let out1 = _mm256_permute2f128_ps(r0, r1, 0x31);
247
248        _mm256_storeu_ps(ptr, out0);
249        _mm256_storeu_ps(ptr.add(8), out1);
250        j += 8;
251    }
252
253    let scale = core::f32::consts::FRAC_1_SQRT_2;
254    while j < n {
255        let idx1 = 2 * j;
256        let idx2 = 2 * j + 1;
257        let tmp1 = scale * x[idx1];
258        let tmp2 = scale * x[idx2];
259        x[idx1] = tmp1 + tmp2;
260        x[idx2] = tmp1 - tmp2;
261        j += 1;
262    }
263}
264
265#[cfg_attr(target_arch = "aarch64", allow(dead_code))]
266#[inline]
267fn haar1_scalar(x: &mut [f32], n0: usize, stride: usize) {
268    let n = n0 >> 1;
269    let scale = core::f32::consts::FRAC_1_SQRT_2;
270    for i in 0..stride {
271        for j in 0..n {
272            let idx1 = stride * 2 * j + i;
273            let idx2 = stride * (2 * j + 1) + i;
274            let tmp1 = scale * x[idx1];
275            let tmp2 = scale * x[idx2];
276            x[idx1] = tmp1 + tmp2;
277            x[idx2] = tmp1 - tmp2;
278        }
279    }
280}
281
282#[cfg(target_arch = "aarch64")]
283fn haar1_neon(x: &mut [f32], n0: usize) {
284    use core::arch::aarch64::*;
285
286    let n = n0 >> 1;
287    let scale = core::f32::consts::FRAC_1_SQRT_2;
288
289    unsafe {
290        let vscale = vdupq_n_f32(scale);
291
292        let mut j = 0usize;
293        while j + 4 <= n {
294            let idx = 2 * j;
295            let pairs = vld2q_f32(x.as_ptr().add(idx));
296            let even = vmulq_f32(pairs.0, vscale);
297            let odd = vmulq_f32(pairs.1, vscale);
298
299            let out = float32x4x2_t {
300                0: vaddq_f32(even, odd),
301                1: vsubq_f32(even, odd),
302            };
303            vst2q_f32(x.as_mut_ptr().add(idx), out);
304            j += 4;
305        }
306
307        while j < n {
308            let idx1 = 2 * j;
309            let idx2 = idx1 + 1;
310            let tmp1 = scale * x[idx1];
311            let tmp2 = scale * x[idx2];
312            x[idx1] = tmp1 + tmp2;
313            x[idx2] = tmp1 - tmp2;
314            j += 1;
315        }
316    }
317}
318
319#[inline(always)]
320pub fn compute_qn(n: usize, b: i32, offset: i32, pulse_cap: i32, stereo: bool) -> i32 {
321    static EXP2_TABLE8: [i16; 8] = [16384, 17866, 19483, 21247, 23170, 25267, 27554, 30048];
322    let mut n2 = (2 * n as i32) - 1;
323    if stereo && n == 2 {
324        n2 -= 1;
325    }
326    let mut qb = celt_sudiv(b + n2 * offset, n2);
327    qb = qb.min(b - pulse_cap - (4 << BITRES));
328    qb = qb.min(8 << BITRES);
329    if qb < (1i32 << BITRES >> 1) {
330        1
331    } else {
332        let val = EXP2_TABLE8[(qb & 0x7) as usize] as i32;
333        let shift = 14 - (qb >> BITRES);
334        let raw = if (0..32).contains(&shift) {
335            val >> shift
336        } else {
337            0
338        };
339        let qn = (raw + 1) >> 1 << 1;
340        qn.min(256)
341    }
342}
343
344#[cfg(target_arch = "aarch64")]
345#[inline(always)]
346#[allow(unsafe_op_in_unsafe_fn)]
347unsafe fn stereo_itheta_neon(x: &[f32], y: &[f32], stereo: bool, n: usize) -> i32 {
348    use core::arch::aarch64::*;
349
350    let mut emid = 1e-15f32;
351    let mut eside = 1e-15f32;
352
353    if stereo {
354        let mut sum_mid = vdupq_n_f32(0.0);
355        let mut sum_side = vdupq_n_f32(0.0);
356        let mut i = 0;
357
358        while i + 16 <= n {
359            let x0 = vld1q_f32(x.as_ptr().add(i));
360            let x1 = vld1q_f32(x.as_ptr().add(i + 4));
361            let x2 = vld1q_f32(x.as_ptr().add(i + 8));
362            let x3 = vld1q_f32(x.as_ptr().add(i + 12));
363            let y0 = vld1q_f32(y.as_ptr().add(i));
364            let y1 = vld1q_f32(y.as_ptr().add(i + 4));
365            let y2 = vld1q_f32(y.as_ptr().add(i + 8));
366            let y3 = vld1q_f32(y.as_ptr().add(i + 12));
367
368            let m0 = vaddq_f32(x0, y0);
369            let m1 = vaddq_f32(x1, y1);
370            let m2 = vaddq_f32(x2, y2);
371            let m3 = vaddq_f32(x3, y3);
372            let s0 = vsubq_f32(x0, y0);
373            let s1 = vsubq_f32(x1, y1);
374            let s2 = vsubq_f32(x2, y2);
375            let s3 = vsubq_f32(x3, y3);
376
377            sum_mid = vfmaq_f32(sum_mid, m0, m0);
378            sum_mid = vfmaq_f32(sum_mid, m1, m1);
379            sum_mid = vfmaq_f32(sum_mid, m2, m2);
380            sum_mid = vfmaq_f32(sum_mid, m3, m3);
381            sum_side = vfmaq_f32(sum_side, s0, s0);
382            sum_side = vfmaq_f32(sum_side, s1, s1);
383            sum_side = vfmaq_f32(sum_side, s2, s2);
384            sum_side = vfmaq_f32(sum_side, s3, s3);
385
386            i += 16;
387        }
388
389        while i + 8 <= n {
390            let x0 = vld1q_f32(x.as_ptr().add(i));
391            let x1 = vld1q_f32(x.as_ptr().add(i + 4));
392            let y0 = vld1q_f32(y.as_ptr().add(i));
393            let y1 = vld1q_f32(y.as_ptr().add(i + 4));
394
395            let m0 = vaddq_f32(x0, y0);
396            let m1 = vaddq_f32(x1, y1);
397            let s0 = vsubq_f32(x0, y0);
398            let s1 = vsubq_f32(x1, y1);
399
400            sum_mid = vfmaq_f32(sum_mid, m0, m0);
401            sum_mid = vfmaq_f32(sum_mid, m1, m1);
402            sum_side = vfmaq_f32(sum_side, s0, s0);
403            sum_side = vfmaq_f32(sum_side, s1, s1);
404
405            i += 8;
406        }
407
408        while i + 4 <= n {
409            let x0 = vld1q_f32(x.as_ptr().add(i));
410            let y0 = vld1q_f32(y.as_ptr().add(i));
411            let m0 = vaddq_f32(x0, y0);
412            let s0 = vsubq_f32(x0, y0);
413            sum_mid = vfmaq_f32(sum_mid, m0, m0);
414            sum_side = vfmaq_f32(sum_side, s0, s0);
415            i += 4;
416        }
417
418        emid += vaddvq_f32(sum_mid);
419        eside += vaddvq_f32(sum_side);
420
421        for j in i..n {
422            let m = x[j] + y[j];
423            let s = x[j] - y[j];
424            emid += m * m;
425            eside += s * s;
426        }
427    } else {
428        let mut sum_mid = vdupq_n_f32(0.0);
429        let mut sum_side = vdupq_n_f32(0.0);
430        let mut i = 0;
431
432        while i + 16 <= n {
433            let x0 = vld1q_f32(x.as_ptr().add(i));
434            let x1 = vld1q_f32(x.as_ptr().add(i + 4));
435            let x2 = vld1q_f32(x.as_ptr().add(i + 8));
436            let x3 = vld1q_f32(x.as_ptr().add(i + 12));
437            let y0 = vld1q_f32(y.as_ptr().add(i));
438            let y1 = vld1q_f32(y.as_ptr().add(i + 4));
439            let y2 = vld1q_f32(y.as_ptr().add(i + 8));
440            let y3 = vld1q_f32(y.as_ptr().add(i + 12));
441
442            sum_mid = vfmaq_f32(sum_mid, x0, x0);
443            sum_mid = vfmaq_f32(sum_mid, x1, x1);
444            sum_mid = vfmaq_f32(sum_mid, x2, x2);
445            sum_mid = vfmaq_f32(sum_mid, x3, x3);
446            sum_side = vfmaq_f32(sum_side, y0, y0);
447            sum_side = vfmaq_f32(sum_side, y1, y1);
448            sum_side = vfmaq_f32(sum_side, y2, y2);
449            sum_side = vfmaq_f32(sum_side, y3, y3);
450
451            i += 16;
452        }
453
454        while i + 8 <= n {
455            let x0 = vld1q_f32(x.as_ptr().add(i));
456            let x1 = vld1q_f32(x.as_ptr().add(i + 4));
457            let y0 = vld1q_f32(y.as_ptr().add(i));
458            let y1 = vld1q_f32(y.as_ptr().add(i + 4));
459
460            sum_mid = vfmaq_f32(sum_mid, x0, x0);
461            sum_mid = vfmaq_f32(sum_mid, x1, x1);
462            sum_side = vfmaq_f32(sum_side, y0, y0);
463            sum_side = vfmaq_f32(sum_side, y1, y1);
464
465            i += 8;
466        }
467
468        while i + 4 <= n {
469            let x0 = vld1q_f32(x.as_ptr().add(i));
470            let y0 = vld1q_f32(y.as_ptr().add(i));
471            sum_mid = vfmaq_f32(sum_mid, x0, x0);
472            sum_side = vfmaq_f32(sum_side, y0, y0);
473            i += 4;
474        }
475
476        emid += vaddvq_f32(sum_mid);
477        eside += vaddvq_f32(sum_side);
478
479        for j in i..n {
480            emid += x[j] * x[j];
481            eside += y[j] * y[j];
482        }
483    }
484
485    let mid = emid.sqrt();
486    let side = eside.sqrt();
487    let theta_norm = celt_atan2p_norm(side, mid);
488    (0.5 + 16384.0 * theta_norm) as i32
489}
490
491#[inline(always)]
492#[cfg(target_arch = "aarch64")]
493pub fn stereo_itheta(x: &[f32], y: &[f32], stereo: bool, n: usize) -> i32 {
494    unsafe { stereo_itheta_neon(x, y, stereo, n) }
495}
496
497#[inline(always)]
498#[cfg(not(target_arch = "aarch64"))]
499pub fn stereo_itheta(x: &[f32], y: &[f32], stereo: bool, n: usize) -> i32 {
500    #[cfg(target_arch = "aarch64")]
501    unsafe {
502        return stereo_itheta_neon(x, y, stereo, n);
503    }
504    #[cfg(not(target_arch = "aarch64"))]
505    {
506        let mut emid = 1e-15f32;
507        let mut eside = 1e-15f32;
508        if stereo {
509            for i in 0..n {
510                let m = x[i] + y[i];
511                let s = x[i] - y[i];
512                emid += m * m;
513                eside += s * s;
514            }
515        } else {
516            for i in 0..n {
517                emid += x[i] * x[i];
518                eside += y[i] * y[i];
519            }
520        }
521        let mid = emid.sqrt();
522        let side = eside.sqrt();
523        let theta_norm = celt_atan2p_norm(side, mid);
524        (0.5 + 16384.0 * theta_norm) as i32
525    }
526}
527
528#[inline(always)]
529fn celt_atan2p_norm(y: f32, x: f32) -> f32 {
530    #[inline(always)]
531    fn atan_norm(x: f32) -> f32 {
532        const ATAN2_2_OVER_PI: f32 = core::f32::consts::FRAC_2_PI;
533        const A03: f32 = -3.333_166e-1_f32;
534        const A05: f32 = 1.996_270_4e-1_f32;
535        const A07: f32 = -1.397_658_3e-1_f32;
536        const A09: f32 = 9.794_234_e-2_f32;
537        const A11: f32 = -5.777_359_e-2_f32;
538        const A13: f32 = 2.304_014e-2_f32;
539        const A15: f32 = -4.355_406e-3_f32;
540        let x2 = x * x;
541        ATAN2_2_OVER_PI
542            * x
543            * (1.0
544                + x2 * (A03
545                    + x2 * (A05 + x2 * (A07 + x2 * (A09 + x2 * (A11 + x2 * (A13 + x2 * A15)))))))
546    }
547    if x * x + y * y < 1e-18 {
548        return 0.0;
549    }
550    if y < x {
551        atan_norm(y / x)
552    } else {
553        1.0 - atan_norm(x / y)
554    }
555}
556
557pub struct SplitCtx {
558    pub inv: bool,
559    pub imid: i32,
560    pub iside: i32,
561    pub delta: i32,
562    pub itheta: i32,
563    pub qalloc: i32,
564}
565
566#[allow(clippy::too_many_arguments)]
567#[inline(always)]
568pub fn compute_theta(
569    ctx: &mut BandCtx,
570    sctx: &mut SplitCtx,
571    x: &[f32],
572    y: &[f32],
573    n: usize,
574    b: &mut i32,
575    b_blocks: i32,
576    b0: i32,
577    lm: i32,
578    stereo: bool,
579    fill: &mut u32,
580) {
581    let pulse_cap = ctx.m.log_n[ctx.i] as i32 + (lm << BITRES);
582    let offset = (pulse_cap >> 1) - if stereo && n == 2 { 16 } else { 4 };
583    let mut qn = compute_qn(n, *b, offset, pulse_cap, stereo);
584
585    if stereo && ctx.i >= ctx.intensity {
586        qn = 1;
587    }
588
589    let mut itheta = 0;
590    if ctx.encode {
591        itheta = stereo_itheta(x, y, stereo, n);
592    }
593
594    let tell_start = tell_frac_inline!(ctx.rc);
595
596    if qn != 1 {
597        if ctx.encode {
598            if !stereo || ctx.theta_round == 0 {
599                itheta = (itheta * qn + 8192) >> 14;
600                if !stereo && ctx.avoid_split_noise && itheta > 0 && itheta < qn {
601                    let unquantized = (itheta * 16384) / qn;
602                    let imid = bitexact_cos(unquantized as i16) as i32;
603                    let iside = bitexact_cos((16384 - unquantized) as i16) as i32;
604                    let delta =
605                        (((n as i32 - 1) << 7) * bitexact_log2tan(iside, imid) + 16384) >> 15;
606                    if delta > *b {
607                        itheta = qn;
608                    } else if delta < -*b {
609                        itheta = 0;
610                    }
611                }
612            } else {
613                let bias = if itheta > 8192 {
614                    32767 / qn
615                } else {
616                    -32767 / qn
617                };
618                let down = (itheta * qn + bias) >> 14;
619                let down = down.clamp(0, qn - 1);
620                if ctx.theta_round < 0 {
621                    itheta = down;
622                } else {
623                    itheta = down + 1;
624                }
625            }
626        }
627
628        if stereo && n > 2 {
629            let p0 = 3;
630            let x0 = qn / 2;
631            let ft = p0 * (x0 + 1) + x0;
632            if ctx.encode {
633                let fl = if itheta <= x0 {
634                    p0 * itheta
635                } else {
636                    (itheta - 1 - x0) + (x0 + 1) * p0
637                };
638                let fh = if itheta <= x0 {
639                    p0 * (itheta + 1)
640                } else {
641                    (itheta - x0) + (x0 + 1) * p0
642                };
643                ctx.rc.encode(fl as u32, fh as u32, ft as u32);
644            } else {
645                let fs = ctx.rc.decode(ft as u32);
646                if fs < (x0 + 1) as u32 * p0 as u32 {
647                    itheta = fs as i32 / p0;
648                } else {
649                    itheta = (x0 + 1) + (fs as i32 - (x0 + 1) * p0);
650                }
651                let fl = if itheta <= x0 {
652                    p0 * itheta
653                } else {
654                    (itheta - 1 - x0) + (x0 + 1) * p0
655                };
656                let fh = if itheta <= x0 {
657                    p0 * (itheta + 1)
658                } else {
659                    (itheta - x0) + (x0 + 1) * p0
660                };
661                ctx.rc.update(fl as u32, fh as u32, ft as u32);
662            }
663        } else if b0 > 1 || stereo {
664            if ctx.encode {
665                ctx.rc.enc_uint(itheta as u32, (qn + 1) as u32);
666            } else {
667                itheta = ctx.rc.dec_uint((qn + 1) as u32) as i32;
668            }
669        } else {
670            let ft = ((qn >> 1) + 1) * ((qn >> 1) + 1);
671            if ctx.encode {
672                let fs = if itheta <= (qn >> 1) {
673                    itheta + 1
674                } else {
675                    qn + 1 - itheta
676                };
677                let fl = if itheta <= (qn >> 1) {
678                    (itheta * (itheta + 1)) >> 1
679                } else {
680                    ft - (((qn + 1 - itheta) * (qn + 2 - itheta)) >> 1)
681                };
682                ctx.rc.encode(fl as u32, (fl + fs) as u32, ft as u32);
683            } else {
684                let fm = ctx.rc.decode(ft as u32) as i32;
685                if fm < (((qn >> 1) * ((qn >> 1) + 1)) >> 1) {
686                    itheta = (isqrt32((8 * fm + 1) as u32) as i32 - 1) >> 1;
687                    let fl = (itheta * (itheta + 1)) >> 1;
688                    let fs = itheta + 1;
689                    ctx.rc.update(fl as u32, (fl + fs) as u32, ft as u32);
690                } else {
691                    itheta = (2 * (qn + 1) - isqrt32((8 * (ft - fm - 1) + 1) as u32) as i32) >> 1;
692                    let fs = qn + 1 - itheta;
693                    let fl = ft - (((qn + 1 - itheta) * (qn + 2 - itheta)) >> 1);
694                    ctx.rc.update(fl as u32, (fl + fs) as u32, ft as u32);
695                }
696            }
697        }
698        itheta = (itheta as u32 * 16384 / qn as u32) as i32;
699        if ctx.encode && stereo {
700            let (bx, by) = (x.as_ptr() as *mut f32, y.as_ptr() as *mut f32);
701            let (sx, sy) = unsafe {
702                (
703                    core::slice::from_raw_parts_mut(bx, n),
704                    core::slice::from_raw_parts_mut(by, n),
705                )
706            };
707            if itheta == 0 {
708                intensity_stereo(ctx.m, sx, sy, ctx.band_e, ctx.i, n);
709            } else {
710                stereo_split(sx, sy, n);
711            }
712        }
713    } else if stereo {
714        if ctx.encode {
715            let inv = itheta > 8192 && !ctx.disable_inv;
716            let (bx, by) = (x.as_ptr() as *mut f32, y.as_ptr() as *mut f32);
717            let (sx, sy) = unsafe {
718                (
719                    core::slice::from_raw_parts_mut(bx, n),
720                    core::slice::from_raw_parts_mut(by, n),
721                )
722            };
723            if inv {
724                for yv in sy.iter_mut() {
725                    *yv = -*yv;
726                }
727            }
728            intensity_stereo(ctx.m, sx, sy, ctx.band_e, ctx.i, n);
729            if *b > (2 << BITRES) && ctx.remaining_bits > (2 << BITRES) {
730                ctx.rc.encode_bit_logp(inv, 2);
731            }
732            itheta = 0;
733            sctx.inv = inv;
734        } else {
735            if *b > (2 << BITRES) && ctx.remaining_bits > (2 << BITRES) {
736                sctx.inv = ctx.rc.decode_bit_logp(2);
737            } else {
738                sctx.inv = false;
739            }
740            if ctx.disable_inv {
741                sctx.inv = false;
742            }
743            itheta = 0;
744        }
745    }
746
747    sctx.itheta = itheta;
748
749    sctx.qalloc = tell_frac_inline!(ctx.rc) - tell_start;
750    *b -= sctx.qalloc; // matches C: *b -= qalloc
751
752    if itheta == 0 {
753        sctx.imid = 32767;
754        sctx.iside = 0;
755        sctx.delta = -16384;
756        *fill &= (1 << b_blocks) - 1;
757    } else if itheta == 16384 {
758        sctx.imid = 0;
759        sctx.iside = 32767;
760        sctx.delta = 16384;
761        *fill &= ((1 << b_blocks) - 1) << b_blocks;
762    } else {
763        let imid = bitexact_cos(itheta as i16);
764        sctx.imid = imid as i32;
765        let iside = bitexact_cos((16384 - itheta) as i16);
766        sctx.iside = iside as i32;
767        sctx.delta =
768            (((n as i32 - 1) << 7) * bitexact_log2tan(sctx.iside, sctx.imid) + 16384) >> 15;
769    }
770}
771
772#[inline(always)]
773fn quant_partition_n2_encode(
774    ctx: &mut BandCtx,
775    x: &mut [f32],
776    b: i32,
777    b_blocks: i32,
778    lowband: Option<&mut [f32]>,
779    lm: i32,
780    gain: f32,
781    fill: u32,
782) -> u32 {
783    let mut q = bits2pulses(ctx.m, ctx.i, lm, b);
784    let mut curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
785    ctx.remaining_bits -= curr_bits;
786
787    while ctx.remaining_bits < 0 && q > 0 {
788        ctx.remaining_bits += curr_bits;
789        q -= 1;
790        curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
791        ctx.remaining_bits -= curr_bits;
792    }
793
794    if q != 0 {
795        let k = get_pulses(q);
796        alg_quant(x, 2, k, ctx.spread, b_blocks as usize, ctx.rc, gain, false)
797    } else {
798        let has_lowband = lowband.is_some();
799        if has_lowband {
800            fill
801        } else {
802            (1u32 << b_blocks) - 1
803        }
804    }
805}
806
807#[inline(always)]
808fn quant_partition_n4_encode(
809    ctx: &mut BandCtx,
810    x: &mut [f32],
811    b: i32,
812    b_blocks: i32,
813    lowband: Option<&mut [f32]>,
814    lm: i32,
815    gain: f32,
816    fill: u32,
817) -> u32 {
818    let mut q = bits2pulses(ctx.m, ctx.i, lm, b);
819    let mut curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
820    ctx.remaining_bits -= curr_bits;
821
822    while ctx.remaining_bits < 0 && q > 0 {
823        ctx.remaining_bits += curr_bits;
824        q -= 1;
825        curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
826        ctx.remaining_bits -= curr_bits;
827    }
828
829    if q != 0 {
830        let k = get_pulses(q);
831        alg_quant(x, 4, k, ctx.spread, b_blocks as usize, ctx.rc, gain, false)
832    } else {
833        let has_lowband = lowband.is_some();
834        if has_lowband {
835            fill
836        } else {
837            (1u32 << b_blocks) - 1
838        }
839    }
840}
841
842#[inline(always)]
843fn quant_partition_n8_encode(
844    ctx: &mut BandCtx,
845    x: &mut [f32],
846    b: i32,
847    b_blocks: i32,
848    lowband: Option<&mut [f32]>,
849    lm: i32,
850    gain: f32,
851    fill: u32,
852) -> u32 {
853    let mut q = bits2pulses(ctx.m, ctx.i, lm, b);
854    let mut curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
855    ctx.remaining_bits -= curr_bits;
856
857    while ctx.remaining_bits < 0 && q > 0 {
858        ctx.remaining_bits += curr_bits;
859        q -= 1;
860        curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
861        ctx.remaining_bits -= curr_bits;
862    }
863
864    if q != 0 {
865        let k = get_pulses(q);
866        alg_quant(x, 8, k, ctx.spread, b_blocks as usize, ctx.rc, gain, false)
867    } else {
868        let has_lowband = lowband.is_some();
869        if has_lowband {
870            fill
871        } else {
872            (1u32 << b_blocks) - 1
873        }
874    }
875}
876
877#[inline(always)]
878#[allow(clippy::too_many_arguments)]
879fn quant_partition_direct_encode(
880    ctx: &mut BandCtx,
881    x: &mut [f32],
882    n: usize,
883    b: i32,
884    b_blocks: i32,
885    lowband: Option<&mut [f32]>,
886    lm: i32,
887    gain: f32,
888    fill: u32,
889) -> u32 {
890    let mut q = bits2pulses(ctx.m, ctx.i, lm, b);
891    let mut curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
892    ctx.remaining_bits -= curr_bits;
893
894    while ctx.remaining_bits < 0 && q > 0 {
895        ctx.remaining_bits += curr_bits;
896        q -= 1;
897        curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
898        ctx.remaining_bits -= curr_bits;
899    }
900
901    if q != 0 {
902        let k = get_pulses(q);
903        alg_quant(x, n, k, ctx.spread, b_blocks as usize, ctx.rc, gain, false)
904    } else {
905        let has_lowband = lowband.is_some();
906        if has_lowband {
907            fill
908        } else {
909            (1u32 << b_blocks) - 1
910        }
911    }
912}
913
914#[inline(always)]
915#[allow(clippy::too_many_arguments)]
916fn quant_partition_encode(
917    ctx: &mut BandCtx,
918    x: &mut [f32],
919    n: usize,
920    b: i32,
921    b_blocks: i32,
922    lowband: Option<&mut [f32]>,
923    lm: i32,
924    gain: f32,
925    fill: u32,
926) -> u32 {
927    // N==2 can never split (should_split requires n>2), dispatch immediately
928    if n == 2 {
929        return quant_partition_n2_encode(ctx, x, b, b_blocks, lowband, lm, gain, fill);
930    }
931
932    // Check split condition FIRST (matching C's quant_partition which checks this before dispatch)
933    let should_split = if lm >= 0 && n > 2 {
934        let cache_idx = (lm + 1) as usize * ctx.m.nb_ebands + ctx.i;
935        let cache_base = unsafe { *ctx.m.cache.index.get_unchecked(cache_idx) };
936        if cache_base >= 0 {
937            let cache_base = cache_base as usize;
938            let cache_ptr = ctx.m.cache.bits.as_ptr().wrapping_add(cache_base);
939            let max_q = unsafe { *cache_ptr } as usize;
940            b > (unsafe { *cache_ptr.add(max_q) } as i32) + 12
941        } else {
942            false
943        }
944    } else {
945        false
946    };
947
948    if should_split {
949        let mut sctx = SplitCtx {
950            inv: false,
951            imid: 0,
952            iside: 0,
953            delta: 0,
954            itheta: 0,
955            qalloc: 0,
956        };
957        let mut b_mut = b;
958        let mut fill_mut = fill;
959        let mid = n / 2;
960        let lm = lm - 1;
961        let b0 = b_blocks;
962        if b_blocks == 1 {
963            fill_mut = (fill_mut & 1) | (fill_mut << 1);
964        }
965        let b_blocks = (b_blocks + 1) >> 1;
966        let (x_mid, x_side) = x.split_at_mut(mid);
967
968        compute_theta(
969            ctx,
970            &mut sctx,
971            x_mid,
972            x_side,
973            mid,
974            &mut b_mut,
975            b_blocks,
976            b0,
977            lm,
978            false,
979            &mut fill_mut,
980        );
981
982        ctx.remaining_bits -= sctx.qalloc;
983        let mut delta = sctx.delta;
984        /* Give more bits to low-energy MDCTs than they would otherwise deserve */
985        if b0 > 1 && (sctx.itheta & 0x3fff) != 0 {
986            if sctx.itheta > 8192 {
987                delta -= delta >> (4 - lm);
988            } else {
989                delta = 0.min(delta + ((mid as i32) << BITRES >> (5 - lm)));
990            }
991        }
992        let mbits = (0).max((b_mut - delta) / 2).min(b_mut);
993        let mut sbits = b_mut - mbits;
994        let mut mbits = mbits;
995
996        let mut rebalance = ctx.remaining_bits;
997        let mut cm;
998        let mid_gain = gain * (sctx.imid as f32 / 32768.0);
999        let side_gain = gain * (sctx.iside as f32 / 32768.0);
1000
1001        if mbits >= sbits {
1002            if let Some(lb) = lowband {
1003                let (lb_mid, lb_side) = lb.split_at_mut(mid);
1004                cm = quant_partition_encode(
1005                    ctx,
1006                    x_mid,
1007                    mid,
1008                    mbits,
1009                    b_blocks,
1010                    Some(lb_mid),
1011                    lm,
1012                    mid_gain,
1013                    fill_mut,
1014                );
1015                rebalance = mbits - (rebalance - ctx.remaining_bits);
1016                if rebalance > (3 << 3) && sctx.itheta != 0 {
1017                    sbits += rebalance - (3 << 3);
1018                }
1019                cm |= quant_partition_encode(
1020                    ctx,
1021                    x_side,
1022                    mid,
1023                    sbits,
1024                    b_blocks,
1025                    Some(lb_side),
1026                    lm,
1027                    side_gain,
1028                    fill_mut >> b_blocks,
1029                ) << (b0 >> 1);
1030            } else {
1031                cm = quant_partition_encode(
1032                    ctx, x_mid, mid, mbits, b_blocks, None, lm, mid_gain, fill_mut,
1033                );
1034                rebalance = mbits - (rebalance - ctx.remaining_bits);
1035                if rebalance > (3 << 3) && sctx.itheta != 0 {
1036                    sbits += rebalance - (3 << 3);
1037                }
1038                cm |= quant_partition_encode(
1039                    ctx,
1040                    x_side,
1041                    mid,
1042                    sbits,
1043                    b_blocks,
1044                    None,
1045                    lm,
1046                    side_gain,
1047                    fill_mut >> b_blocks,
1048                ) << (b0 >> 1);
1049            }
1050        } else if let Some(lb) = lowband {
1051            let (lb_mid, lb_side) = lb.split_at_mut(mid);
1052            cm = quant_partition_encode(
1053                ctx,
1054                x_side,
1055                mid,
1056                sbits,
1057                b_blocks,
1058                Some(lb_side),
1059                lm,
1060                side_gain,
1061                fill_mut >> b_blocks,
1062            ) << (b0 >> 1);
1063            rebalance = sbits - (rebalance - ctx.remaining_bits);
1064            if rebalance > (3 << 3) && sctx.itheta != 16384 {
1065                mbits += rebalance - (3 << 3);
1066            }
1067            cm |= quant_partition_encode(
1068                ctx,
1069                x_mid,
1070                mid,
1071                mbits,
1072                b_blocks,
1073                Some(lb_mid),
1074                lm,
1075                mid_gain,
1076                fill_mut,
1077            );
1078        } else {
1079            cm = quant_partition_encode(
1080                ctx,
1081                x_side,
1082                mid,
1083                sbits,
1084                b_blocks,
1085                None,
1086                lm,
1087                side_gain,
1088                fill_mut >> b_blocks,
1089            ) << (b0 >> 1);
1090            rebalance = sbits - (rebalance - ctx.remaining_bits);
1091            if rebalance > (3 << 3) && sctx.itheta != 16384 {
1092                mbits += rebalance - (3 << 3);
1093            }
1094            cm |= quant_partition_encode(
1095                ctx, x_mid, mid, mbits, b_blocks, None, lm, mid_gain, fill_mut,
1096            );
1097        }
1098        cm
1099    } else {
1100        // No split — dispatch to small-N specialized encoders or direct path
1101        if n == 4 {
1102            return quant_partition_n4_encode(ctx, x, b, b_blocks, lowband, lm, gain, fill);
1103        }
1104        if n == 8 {
1105            return quant_partition_n8_encode(ctx, x, b, b_blocks, lowband, lm, gain, fill);
1106        }
1107        if n == 16 {
1108            return quant_partition_direct_encode(ctx, x, n, b, b_blocks, lowband, lm, gain, fill);
1109        }
1110        let mut q = bits2pulses(ctx.m, ctx.i, lm, b);
1111        let mut curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
1112        ctx.remaining_bits -= curr_bits;
1113
1114        while ctx.remaining_bits < 0 && q > 0 {
1115            ctx.remaining_bits += curr_bits;
1116            q -= 1;
1117            curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
1118            ctx.remaining_bits -= curr_bits;
1119        }
1120
1121        if q != 0 {
1122            let k = get_pulses(q);
1123            alg_quant(x, n, k, ctx.spread, b_blocks as usize, ctx.rc, gain, false)
1124        } else if lowband.is_some() {
1125            fill
1126        } else {
1127            (1 << b_blocks) - 1
1128        }
1129    }
1130}
1131
1132#[inline(always)]
1133#[allow(clippy::too_many_arguments)]
1134pub fn quant_partition(
1135    ctx: &mut BandCtx,
1136    x: &mut [f32],
1137    n: usize,
1138    b: i32,
1139    b_blocks: i32,
1140    lowband: Option<&mut [f32]>,
1141    lm: i32,
1142    gain: f32,
1143    fill: u32,
1144) -> u32 {
1145    /* Check split condition FIRST, before dispatching to specialized handlers.
1146    This matches the C code which checks this at the top of quant_partition. */
1147    let should_split = if lm >= 0 && n > 2 {
1148        let cache_idx = (lm + 1) as usize * ctx.m.nb_ebands + ctx.i;
1149        let cache_base = unsafe { *ctx.m.cache.index.get_unchecked(cache_idx) };
1150        if cache_base >= 0 {
1151            let cache_base = cache_base as usize;
1152            let cache_ptr = ctx.m.cache.bits.as_ptr().wrapping_add(cache_base);
1153            let max_q = unsafe { *cache_ptr } as usize;
1154            b > (unsafe { *cache_ptr.add(max_q) } as i32) + 12
1155        } else {
1156            false
1157        }
1158    } else {
1159        false
1160    };
1161    if should_split {
1162        let mut sctx = SplitCtx {
1163            inv: false,
1164            imid: 0,
1165            iside: 0,
1166            delta: 0,
1167            itheta: 0,
1168            qalloc: 0,
1169        };
1170        let mut b_mut = b;
1171        let mut fill_mut = fill;
1172        let mid = n / 2;
1173        let lm = lm - 1;
1174        let b0 = b_blocks; // Save original B0
1175        if b_blocks == 1 {
1176            fill_mut = (fill_mut & 1) | (fill_mut << 1);
1177        }
1178        let b_blocks = (b_blocks + 1) >> 1;
1179        let (x_mid, x_side) = x.split_at_mut(mid);
1180
1181        compute_theta(
1182            ctx,
1183            &mut sctx,
1184            x_mid,
1185            x_side,
1186            mid,
1187            &mut b_mut,
1188            b_blocks,
1189            b0,
1190            lm,
1191            false,
1192            &mut fill_mut,
1193        );
1194
1195        ctx.remaining_bits -= sctx.qalloc;
1196        let mut delta = sctx.delta;
1197        /* Give more bits to low-energy MDCTs than they would otherwise deserve
1198        (matches C quant_partition's B0>1 adjustment) */
1199        if b0 > 1 && (sctx.itheta & 0x3fff) != 0 {
1200            if sctx.itheta > 8192 {
1201                delta -= delta >> (4 - lm);
1202            } else {
1203                delta = 0.min(delta + ((mid as i32) << BITRES >> (5 - lm)));
1204            }
1205        }
1206        let mbits = (0).max((b_mut - delta) / 2).min(b_mut);
1207        let mut sbits = b_mut - mbits;
1208        let mut mbits = mbits;
1209
1210        let mut rebalance = ctx.remaining_bits;
1211        let mut cm;
1212
1213        if mbits >= sbits {
1214            if let Some(lb) = lowband {
1215                let (lb_mid, lb_side) = lb.split_at_mut(mid);
1216                cm = quant_partition(
1217                    ctx,
1218                    x_mid,
1219                    mid,
1220                    mbits,
1221                    b_blocks,
1222                    Some(lb_mid),
1223                    lm,
1224                    gain * (sctx.imid as f32 / 32768.0),
1225                    fill_mut,
1226                );
1227                rebalance = mbits - (rebalance - ctx.remaining_bits);
1228                if rebalance > (3 << 3) && sctx.itheta != 0 {
1229                    sbits += rebalance - (3 << 3);
1230                }
1231                cm |= quant_partition(
1232                    ctx,
1233                    x_side,
1234                    mid,
1235                    sbits,
1236                    b_blocks,
1237                    Some(lb_side),
1238                    lm,
1239                    gain * (sctx.iside as f32 / 32768.0),
1240                    fill_mut >> b_blocks,
1241                ) << (b0 >> 1);
1242            } else {
1243                cm = quant_partition(
1244                    ctx,
1245                    x_mid,
1246                    mid,
1247                    mbits,
1248                    b_blocks,
1249                    None,
1250                    lm,
1251                    gain * (sctx.imid as f32 / 32768.0),
1252                    fill_mut,
1253                );
1254                rebalance = mbits - (rebalance - ctx.remaining_bits);
1255                if rebalance > (3 << 3) && sctx.itheta != 0 {
1256                    sbits += rebalance - (3 << 3);
1257                }
1258                cm |= quant_partition(
1259                    ctx,
1260                    x_side,
1261                    mid,
1262                    sbits,
1263                    b_blocks,
1264                    None,
1265                    lm,
1266                    gain * (sctx.iside as f32 / 32768.0),
1267                    fill_mut >> b_blocks,
1268                ) << (b0 >> 1);
1269            }
1270        } else if let Some(lb) = lowband {
1271            let (lb_mid, lb_side) = lb.split_at_mut(mid);
1272            cm = quant_partition(
1273                ctx,
1274                x_side,
1275                mid,
1276                sbits,
1277                b_blocks,
1278                Some(lb_side),
1279                lm,
1280                gain * (sctx.iside as f32 / 32768.0),
1281                fill_mut >> b_blocks,
1282            ) << (b0 >> 1);
1283            rebalance = sbits - (rebalance - ctx.remaining_bits);
1284            if rebalance > (3 << 3) && sctx.itheta != 16384 {
1285                mbits += rebalance - (3 << 3);
1286            }
1287            cm |= quant_partition(
1288                ctx,
1289                x_mid,
1290                mid,
1291                mbits,
1292                b_blocks,
1293                Some(lb_mid),
1294                lm,
1295                gain * (sctx.imid as f32 / 32768.0),
1296                fill_mut,
1297            );
1298        } else {
1299            cm = quant_partition(
1300                ctx,
1301                x_side,
1302                mid,
1303                sbits,
1304                b_blocks,
1305                None,
1306                lm,
1307                gain * (sctx.iside as f32 / 32768.0),
1308                fill_mut >> b_blocks,
1309            ) << (b0 >> 1);
1310            rebalance = sbits - (rebalance - ctx.remaining_bits);
1311            if rebalance > (3 << 3) && sctx.itheta != 16384 {
1312                mbits += rebalance - (3 << 3);
1313            }
1314            cm |= quant_partition(
1315                ctx,
1316                x_mid,
1317                mid,
1318                mbits,
1319                b_blocks,
1320                None,
1321                lm,
1322                gain * (sctx.imid as f32 / 32768.0),
1323                fill_mut,
1324            );
1325        }
1326        cm
1327    } else {
1328        let mut q = bits2pulses(ctx.m, ctx.i, lm, b);
1329        let mut curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
1330        ctx.remaining_bits -= curr_bits;
1331
1332        while ctx.remaining_bits < 0 && q > 0 {
1333            ctx.remaining_bits += curr_bits;
1334            q -= 1;
1335            curr_bits = pulses2bits(ctx.m, ctx.i, lm, q);
1336            ctx.remaining_bits -= curr_bits;
1337        }
1338
1339        if q != 0 {
1340            let k = get_pulses(q);
1341            if ctx.encode {
1342                alg_quant(
1343                    x,
1344                    n,
1345                    k,
1346                    ctx.spread,
1347                    b_blocks as usize,
1348                    ctx.rc,
1349                    gain,
1350                    ctx.resynth,
1351                )
1352            } else {
1353                alg_unquant(x, n, k, ctx.spread, b_blocks as usize, ctx.rc, gain)
1354            }
1355        } else {
1356            let mut cm = 0u32;
1357            if ctx.resynth {
1358                let cm_mask = (1u32 << b_blocks) - 1;
1359                let fill_masked = fill & cm_mask;
1360                if fill_masked == 0 {
1361                    x[..n].fill(0.0);
1362                } else if let Some(lb) = lowband {
1363                    #[cfg(target_arch = "aarch64")]
1364                    unsafe {
1365                        use core::arch::aarch64::*;
1366                        let n8 = n & !7;
1367                        let mut i = 0;
1368                        while i < n8 {
1369                            let mut vals = [0.0f32; 8];
1370                            for j in 0..8 {
1371                                ctx.seed = celt_lcg_rand(ctx.seed);
1372                                vals[j] = if ctx.seed & 0x8000 != 0 {
1373                                    1.0 / 256.0
1374                                } else {
1375                                    -1.0 / 256.0
1376                                };
1377                            }
1378                            let vnoise = vld1q_f32(vals.as_ptr());
1379                            let vnoise1 = vld1q_f32(vals.as_ptr().add(4));
1380                            let vlb = vld1q_f32(lb.as_ptr().add(i));
1381                            let vlb1 = vld1q_f32(lb.as_ptr().add(i + 4));
1382                            let vres = vaddq_f32(vlb, vnoise);
1383                            let vres1 = vaddq_f32(vlb1, vnoise1);
1384                            vst1q_f32(x.as_mut_ptr().add(i), vres);
1385                            vst1q_f32(x.as_mut_ptr().add(i + 4), vres1);
1386                            i += 8;
1387                        }
1388                        for j in i..n {
1389                            ctx.seed = celt_lcg_rand(ctx.seed);
1390                            x[j] = lb[j]
1391                                + if ctx.seed & 0x8000 != 0 {
1392                                    1.0 / 256.0
1393                                } else {
1394                                    -1.0 / 256.0
1395                                };
1396                        }
1397                    }
1398                    #[cfg(not(target_arch = "aarch64"))]
1399                    {
1400                        for j in 0..n {
1401                            ctx.seed = celt_lcg_rand(ctx.seed);
1402                            x[j] = lb[j]
1403                                + if ctx.seed & 0x8000 != 0 {
1404                                    1.0 / 256.0
1405                                } else {
1406                                    -1.0 / 256.0
1407                                };
1408                        }
1409                    }
1410                    renormalise_vector(x, n, gain);
1411                    cm = fill_masked;
1412                } else {
1413                    for xv in x[..n].iter_mut() {
1414                        ctx.seed = celt_lcg_rand(ctx.seed);
1415                        *xv = ((ctx.seed as i32 >> 20) as f32) / 16384.0;
1416                    }
1417                    renormalise_vector(x, n, gain);
1418                    cm = cm_mask;
1419                }
1420            }
1421            cm
1422        }
1423    }
1424}
1425
1426#[cfg(target_arch = "aarch64")]
1427#[inline(always)]
1428unsafe fn deinterleave_hadamard_neon(x: &mut [f32], n0: usize, stride: usize) {
1429    let n = n0 * stride;
1430    let mut tmp_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
1431    let tmp = core::slice::from_raw_parts_mut(tmp_buf.as_mut_ptr() as *mut f32, n);
1432
1433    for i in 0..stride {
1434        let src_offset = i;
1435        let dst_offset = i * n0;
1436        for j in 0..n0 {
1437            tmp[dst_offset + j] = x[j * stride + src_offset];
1438        }
1439    }
1440
1441    x[..n].copy_from_slice(tmp);
1442}
1443
1444pub fn deinterleave_hadamard(x: &mut [f32], n0: usize, stride: usize, hadamard: bool) {
1445    let n = n0 * stride;
1446
1447    let mut tmp_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
1448
1449    let tmp = unsafe { core::slice::from_raw_parts_mut(tmp_buf.as_mut_ptr() as *mut f32, n) };
1450    if hadamard {
1451        let offset = match stride {
1452            2 => 0,
1453            4 => 2,
1454            8 => 6,
1455            16 => 14,
1456            _ => 0,
1457        };
1458        let ordery = &ORDERY_TABLE[offset..offset + stride];
1459        for i in 0..stride {
1460            for j in 0..n0 {
1461                tmp[ordery[i] as usize * n0 + j] = x[j * stride + i];
1462            }
1463        }
1464    } else {
1465        #[cfg(target_arch = "aarch64")]
1466        unsafe {
1467            if n0 >= 4 {
1468                deinterleave_hadamard_neon(x, n0, stride);
1469                return;
1470            }
1471        }
1472        for i in 0..stride {
1473            for j in 0..n0 {
1474                tmp[i * n0 + j] = x[j * stride + i];
1475            }
1476        }
1477    }
1478    x[..n].copy_from_slice(tmp);
1479}
1480
1481#[cfg(target_arch = "aarch64")]
1482#[inline(always)]
1483unsafe fn interleave_hadamard_neon(x: &mut [f32], n0: usize, stride: usize) {
1484    let n = n0 * stride;
1485    let mut tmp_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
1486    let tmp = core::slice::from_raw_parts_mut(tmp_buf.as_mut_ptr() as *mut f32, n);
1487
1488    for i in 0..stride {
1489        let src_offset = i * n0;
1490        let dst_offset = i;
1491        for j in 0..n0 {
1492            tmp[j * stride + dst_offset] = x[src_offset + j];
1493        }
1494    }
1495
1496    x[..n].copy_from_slice(tmp);
1497}
1498
1499pub fn interleave_hadamard(x: &mut [f32], n0: usize, stride: usize, hadamard: bool) {
1500    let n = n0 * stride;
1501    let mut tmp_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
1502    let tmp = unsafe { core::slice::from_raw_parts_mut(tmp_buf.as_mut_ptr() as *mut f32, n) };
1503    if hadamard {
1504        let offset = match stride {
1505            2 => 0,
1506            4 => 2,
1507            8 => 6,
1508            16 => 14,
1509            _ => 0,
1510        };
1511        let ordery = &ORDERY_TABLE[offset..offset + stride];
1512        for i in 0..stride {
1513            for j in 0..n0 {
1514                tmp[j * stride + i] = x[ordery[i] as usize * n0 + j];
1515            }
1516        }
1517    } else {
1518        #[cfg(target_arch = "aarch64")]
1519        unsafe {
1520            if n0 >= 4 {
1521                interleave_hadamard_neon(x, n0, stride);
1522                return;
1523            }
1524        }
1525        for i in 0..stride {
1526            for j in 0..n0 {
1527                tmp[j * stride + i] = x[i * n0 + j];
1528            }
1529        }
1530    }
1531    x[..n].copy_from_slice(tmp);
1532}
1533
1534const ORDERY_TABLE: [i32; 30] = [
1535    1, 0, 3, 0, 2, 1, 7, 0, 4, 3, 6, 1, 5, 2, 15, 0, 8, 7, 12, 3, 11, 4, 14, 1, 9, 6, 13, 2, 10, 5,
1536];
1537
1538fn quant_band_n1(
1539    ctx: &mut BandCtx,
1540    x: &mut [f32],
1541    y: Option<&mut [f32]>,
1542    lowband_out: Option<&mut [f32]>,
1543) -> u32 {
1544    let mut sign = 0;
1545    if ctx.remaining_bits >= 1 << BITRES {
1546        if ctx.encode {
1547            sign = if x[0] < 0.0 { 1 } else { 0 };
1548            ctx.rc.enc_bits(sign as u32, 1);
1549        } else {
1550            sign = ctx.rc.dec_bits(1) as i32;
1551        }
1552        ctx.remaining_bits -= 1 << BITRES;
1553    }
1554    if ctx.resynth {
1555        x[0] = if sign != 0 { -1.0 } else { 1.0 };
1556    }
1557    if let Some(y_val) = y {
1558        let mut y_sign = 0;
1559        if ctx.remaining_bits >= 1 << BITRES {
1560            if ctx.encode {
1561                y_sign = if y_val[0] < 0.0 { 1 } else { 0 };
1562                ctx.rc.enc_bits(y_sign as u32, 1);
1563            } else {
1564                y_sign = ctx.rc.dec_bits(1) as i32;
1565            }
1566            ctx.remaining_bits -= 1 << BITRES;
1567        }
1568        if ctx.resynth {
1569            y_val[0] = if y_sign != 0 { -1.0 } else { 1.0 };
1570        }
1571    }
1572    if let Some(l_out) = lowband_out {
1573        l_out[0] = x[0] / 16.0;
1574    }
1575    1
1576}
1577
1578#[allow(clippy::too_many_arguments)]
1579#[inline(always)]
1580pub fn quant_band(
1581    ctx: &mut BandCtx,
1582    x: &mut [f32],
1583    n: usize,
1584    b: i32,
1585    b_blocks: i32,
1586    lowband: Option<&mut [f32]>,
1587    lm: i32,
1588    lowband_out: Option<&mut [f32]>,
1589    gain: f32,
1590    fill: u32,
1591) -> u32 {
1592    let n0 = n;
1593    let b0 = b_blocks;
1594    let long_blocks = b0 == 1;
1595
1596    if n == 1 {
1597        return quant_band_n1(ctx, x, None, lowband_out);
1598    }
1599
1600    let mut b_blocks = b_blocks;
1601    let mut n_b = n / b_blocks as usize;
1602    let mut time_divide = 0;
1603    let mut recombine = 0;
1604    let mut tf_change_local = ctx.tf_change;
1605    let mut fill = fill;
1606
1607    if tf_change_local > 0 {
1608        recombine = tf_change_local;
1609    }
1610
1611    let mut lowband_buf = lowband;
1612
1613    static BIT_INTERLEAVE_TABLE: [u8; 16] = [0, 1, 1, 1, 2, 3, 3, 3, 2, 3, 3, 3, 2, 3, 3, 3];
1614
1615    for k in 0..recombine {
1616        if ctx.encode {
1617            haar1(x, n >> k, 1 << k);
1618        }
1619        if let Some(ref mut lb) = lowband_buf {
1620            haar1(lb, n >> k, 1 << k);
1621        }
1622        fill = (BIT_INTERLEAVE_TABLE[(fill & 0xF) as usize] as u32)
1623            | ((BIT_INTERLEAVE_TABLE[(fill >> 4) as usize] as u32) << 2);
1624    }
1625    b_blocks >>= recombine;
1626    n_b <<= recombine;
1627
1628    while n_b & 1 == 0 && tf_change_local < 0 {
1629        if ctx.encode {
1630            haar1(x, n_b, b_blocks as usize);
1631        }
1632        if let Some(ref mut lb) = lowband_buf {
1633            haar1(lb, n_b, b_blocks as usize);
1634        }
1635        fill |= fill << b_blocks;
1636        b_blocks <<= 1;
1637        n_b >>= 1;
1638        time_divide += 1;
1639        tf_change_local += 1;
1640    }
1641
1642    let b0_after = b_blocks;
1643    let n_b0 = n_b;
1644
1645    if b_blocks > 1 {
1646        if ctx.encode {
1647            deinterleave_hadamard(
1648                x,
1649                n_b >> recombine as usize,
1650                (b_blocks << recombine) as usize,
1651                long_blocks,
1652            );
1653        }
1654        if let Some(ref mut lb) = lowband_buf {
1655            deinterleave_hadamard(
1656                lb,
1657                n_b >> recombine as usize,
1658                (b_blocks << recombine) as usize,
1659                long_blocks,
1660            );
1661        }
1662    }
1663
1664    let cm = if ctx.encode {
1665        quant_partition_encode(ctx, x, n, b, b_blocks, lowband_buf, lm, gain, fill)
1666    } else {
1667        quant_partition(ctx, x, n, b, b_blocks, lowband_buf, lm, gain, fill)
1668    };
1669
1670    if ctx.resynth {
1671        let mut cm = cm;
1672
1673        if b_blocks > 1 {
1674            interleave_hadamard(
1675                x,
1676                n_b >> recombine as usize,
1677                (b0_after << recombine) as usize,
1678                long_blocks,
1679            );
1680        }
1681
1682        let mut n_b_undo = n_b0;
1683        let mut b_undo = b0_after;
1684        for _ in 0..time_divide {
1685            b_undo >>= 1;
1686            n_b_undo <<= 1;
1687            cm |= cm >> b_undo;
1688            haar1(x, n_b_undo, b_undo as usize);
1689        }
1690
1691        static BIT_DEINTERLEAVE_TABLE: [u8; 16] = [
1692            0x00, 0x03, 0x0C, 0x0F, 0x30, 0x33, 0x3C, 0x3F, 0xC0, 0xC3, 0xCC, 0xCF, 0xF0, 0xF3,
1693            0xFC, 0xFF,
1694        ];
1695        for k in 0..recombine {
1696            cm = BIT_DEINTERLEAVE_TABLE[cm as usize & 0xF] as u32;
1697            haar1(x, n0 >> k, 1 << k);
1698        }
1699        let mut b_final = b_undo;
1700        b_final <<= recombine;
1701
1702        if let Some(lb_out) = lowband_out {
1703            let scale = (n0 as f32).sqrt();
1704            for j in 0..n0 {
1705                lb_out[j] = scale * x[j];
1706            }
1707        }
1708        cm &= (1u32 << b_final) - 1;
1709        return cm;
1710    }
1711
1712    cm
1713}
1714
1715pub fn stereo_merge(x: &mut [f32], y: &mut [f32], mid: f32, _side: f32, n: usize) {
1716    let mut xp = 0.0f32;
1717    let mut side_e = 0.0f32;
1718    for i in 0..n {
1719        xp += y[i] * x[i];
1720        side_e += y[i] * y[i];
1721    }
1722
1723    xp *= mid;
1724    let el = mid * mid + side_e - 2.0 * xp;
1725    let er = mid * mid + side_e + 2.0 * xp;
1726
1727    if er < 6e-4f32 || el < 6e-4f32 {
1728        y[..n].copy_from_slice(&x[..n]);
1729        return;
1730    }
1731
1732    let lgain = 1.0 / el.sqrt();
1733    let rgain = 1.0 / er.sqrt();
1734
1735    for i in 0..n {
1736        let l = mid * x[i];
1737        let r = y[i];
1738        x[i] = lgain * (l - r);
1739        y[i] = rgain * (l + r);
1740    }
1741}
1742
1743#[inline(always)]
1744fn stereo_split(x: &mut [f32], y: &mut [f32], n: usize) {
1745    let scale = core::f32::consts::FRAC_1_SQRT_2;
1746    for i in 0..n {
1747        let l = scale * x[i];
1748        let r = scale * y[i];
1749        x[i] = l + r;
1750        y[i] = r - l;
1751    }
1752}
1753
1754#[inline(always)]
1755fn intensity_stereo(
1756    m: &CeltMode,
1757    x: &mut [f32],
1758    y: &mut [f32],
1759    band_e: &[f32],
1760    band: usize,
1761    n: usize,
1762) {
1763    let left = band_e[band].max(MIN_STEREO_ENERGY);
1764    let right = band_e[m.nb_ebands + band].max(MIN_STEREO_ENERGY);
1765    let norm = (left * left + right * right).sqrt().max(MIN_STEREO_ENERGY);
1766    let a1 = left / norm;
1767    let a2 = right / norm;
1768    for i in 0..n {
1769        x[i] = a1 * x[i] + a2 * y[i];
1770    }
1771}
1772
1773#[inline(always)]
1774fn special_hybrid_folding(m: &CeltMode, norm: &mut [f32], start: usize, m_val: usize) {
1775    if start + 2 >= m.e_bands.len() {
1776        return;
1777    }
1778    let n1 = m_val * (m.e_bands[start + 1] - m.e_bands[start]) as usize;
1779    let n2 = m_val * (m.e_bands[start + 2] - m.e_bands[start + 1]) as usize;
1780    if n2 <= n1 {
1781        return;
1782    }
1783    let len = n2 - n1;
1784    let src_start = 2 * n1 - n2;
1785    if src_start + len <= norm.len() && n1 + len <= norm.len() {
1786        norm.copy_within(src_start..src_start + len, n1);
1787    }
1788}
1789
1790fn prepare_lowband_views(
1791    norm: &mut [f32],
1792    lowband_scratch_ptr: *mut f32,
1793    allow_lowband_scratch: bool,
1794    effective_lowband: i32,
1795    norm_pos: usize,
1796    n: usize,
1797    want_out: bool,
1798) -> (Option<&mut [f32]>, Option<&mut [f32]>) {
1799    let len = norm.len();
1800    let out_range = if want_out && norm_pos + n <= len {
1801        Some((norm_pos, norm_pos + n))
1802    } else {
1803        None
1804    };
1805
1806    let Some(lb_start) = (if effective_lowband >= 0 {
1807        Some(effective_lowband as usize)
1808    } else {
1809        None
1810    }) else {
1811        let lb_out = out_range.map(|(s, e)| &mut norm[s..e]);
1812        return (None, lb_out);
1813    };
1814    let lb_end = lb_start + n;
1815    if lb_end > len {
1816        let lb_out = out_range.map(|(s, e)| &mut norm[s..e]);
1817        return (None, lb_out);
1818    }
1819
1820    if allow_lowband_scratch {
1821        unsafe {
1822            core::ptr::copy_nonoverlapping(norm.as_ptr().add(lb_start), lowband_scratch_ptr, n)
1823        };
1824        let lb = Some(unsafe { core::slice::from_raw_parts_mut(lowband_scratch_ptr, n) });
1825        let lb_out = out_range.map(|(s, e)| &mut norm[s..e]);
1826        return (lb, lb_out);
1827    }
1828
1829    if let Some((out_start, out_end)) = out_range {
1830        if lb_end <= out_start {
1831            let (left, right) = norm.split_at_mut(out_start);
1832            let lb = Some(&mut left[lb_start..lb_end]);
1833            let lb_out = Some(&mut right[..(out_end - out_start)]);
1834            return (lb, lb_out);
1835        }
1836        if out_end <= lb_start {
1837            let (left, right) = norm.split_at_mut(lb_start);
1838            let lb_out = Some(&mut left[out_start..out_end]);
1839            let lb = Some(&mut right[..n]);
1840            return (lb, lb_out);
1841        }
1842        return (Some(&mut norm[lb_start..lb_end]), None);
1843    }
1844
1845    (Some(&mut norm[lb_start..lb_end]), None)
1846}
1847
1848#[cfg(target_arch = "x86_64")]
1849#[target_feature(enable = "avx2")]
1850#[allow(dead_code)]
1851unsafe fn stereo_merge_avx2(x: &mut [f32], y: &mut [f32], mid: f32, side: f32, n: usize) {
1852    use core::arch::x86_64::*;
1853
1854    let mut i = 0;
1855
1856    let v_mid = _mm256_set1_ps(mid);
1857    let v_side = _mm256_set1_ps(side);
1858
1859    while i + 15 < n {
1860        let x0 = _mm256_loadu_ps(x.as_ptr().add(i));
1861        let x1 = _mm256_loadu_ps(x.as_ptr().add(i + 8));
1862        let y0 = _mm256_loadu_ps(y.as_ptr().add(i));
1863        let y1 = _mm256_loadu_ps(y.as_ptr().add(i + 8));
1864
1865        let x_val0 = _mm256_mul_ps(x0, v_mid);
1866        let x_val1 = _mm256_mul_ps(x1, v_mid);
1867        let y_val0 = _mm256_mul_ps(y0, v_side);
1868        let y_val1 = _mm256_mul_ps(y1, v_side);
1869
1870        let new_x0 = _mm256_sub_ps(x_val0, y_val0);
1871        let new_x1 = _mm256_sub_ps(x_val1, y_val1);
1872        let new_y0 = _mm256_add_ps(x_val0, y_val0);
1873        let new_y1 = _mm256_add_ps(x_val1, y_val1);
1874
1875        _mm256_storeu_ps(x.as_mut_ptr().add(i), new_x0);
1876        _mm256_storeu_ps(x.as_mut_ptr().add(i + 8), new_x1);
1877        _mm256_storeu_ps(y.as_mut_ptr().add(i), new_y0);
1878        _mm256_storeu_ps(y.as_mut_ptr().add(i + 8), new_y1);
1879
1880        i += 16;
1881    }
1882
1883    while i + 7 < n {
1884        let x0 = _mm256_loadu_ps(x.as_ptr().add(i));
1885        let y0 = _mm256_loadu_ps(y.as_ptr().add(i));
1886
1887        let x_val = _mm256_mul_ps(x0, v_mid);
1888        let y_val = _mm256_mul_ps(y0, v_side);
1889
1890        let new_x = _mm256_sub_ps(x_val, y_val);
1891        let new_y = _mm256_add_ps(x_val, y_val);
1892
1893        _mm256_storeu_ps(x.as_mut_ptr().add(i), new_x);
1894        _mm256_storeu_ps(y.as_mut_ptr().add(i), new_y);
1895
1896        i += 8;
1897    }
1898
1899    for j in i..n {
1900        let x_val = x[j] * mid;
1901        let y_val = y[j] * side;
1902        x[j] = x_val - y_val;
1903        y[j] = x_val + y_val;
1904    }
1905}
1906
1907#[allow(dead_code)]
1908#[inline]
1909fn stereo_merge_scalar(x: &mut [f32], y: &mut [f32], mid: f32, side: f32, n: usize) {
1910    for i in 0..n {
1911        let x_val = x[i] * mid;
1912        let y_val = y[i] * side;
1913        x[i] = x_val - y_val;
1914        y[i] = x_val + y_val;
1915    }
1916}
1917
1918#[cfg(target_arch = "aarch64")]
1919#[allow(dead_code)]
1920fn stereo_merge_neon(x: &mut [f32], y: &mut [f32], mid: f32, side: f32, n: usize) {
1921    use core::arch::aarch64::*;
1922
1923    unsafe {
1924        let vmid = vdupq_n_f32(mid);
1925        let vside = vdupq_n_f32(side);
1926
1927        let n16 = n & !15;
1928        for i in (0..n16).step_by(16) {
1929            let x0 = vld1q_f32(x.as_ptr().add(i));
1930            let x1 = vld1q_f32(x.as_ptr().add(i + 4));
1931            let x2 = vld1q_f32(x.as_ptr().add(i + 8));
1932            let x3 = vld1q_f32(x.as_ptr().add(i + 12));
1933
1934            let y0 = vld1q_f32(y.as_ptr().add(i));
1935            let y1 = vld1q_f32(y.as_ptr().add(i + 4));
1936            let y2 = vld1q_f32(y.as_ptr().add(i + 8));
1937            let y3 = vld1q_f32(y.as_ptr().add(i + 12));
1938
1939            let xv0 = vmulq_f32(x0, vmid);
1940            let xv1 = vmulq_f32(x1, vmid);
1941            let xv2 = vmulq_f32(x2, vmid);
1942            let xv3 = vmulq_f32(x3, vmid);
1943
1944            let yv0 = vmulq_f32(y0, vside);
1945            let yv1 = vmulq_f32(y1, vside);
1946            let yv2 = vmulq_f32(y2, vside);
1947            let yv3 = vmulq_f32(y3, vside);
1948
1949            vst1q_f32(x.as_mut_ptr().add(i), vsubq_f32(xv0, yv0));
1950            vst1q_f32(x.as_mut_ptr().add(i + 4), vsubq_f32(xv1, yv1));
1951            vst1q_f32(x.as_mut_ptr().add(i + 8), vsubq_f32(xv2, yv2));
1952            vst1q_f32(x.as_mut_ptr().add(i + 12), vsubq_f32(xv3, yv3));
1953
1954            vst1q_f32(y.as_mut_ptr().add(i), vaddq_f32(xv0, yv0));
1955            vst1q_f32(y.as_mut_ptr().add(i + 4), vaddq_f32(xv1, yv1));
1956            vst1q_f32(y.as_mut_ptr().add(i + 8), vaddq_f32(xv2, yv2));
1957            vst1q_f32(y.as_mut_ptr().add(i + 12), vaddq_f32(xv3, yv3));
1958        }
1959
1960        let n4 = (n & !3) - n16;
1961        for i in (n16..n16 + n4).step_by(4) {
1962            let xv = vld1q_f32(x.as_ptr().add(i));
1963            let yv = vld1q_f32(y.as_ptr().add(i));
1964
1965            let x_val = vmulq_f32(xv, vmid);
1966            let y_val = vmulq_f32(yv, vside);
1967
1968            vst1q_f32(x.as_mut_ptr().add(i), vsubq_f32(x_val, y_val));
1969            vst1q_f32(y.as_mut_ptr().add(i), vaddq_f32(x_val, y_val));
1970        }
1971
1972        for i in (n16 + n4)..n {
1973            let x_val = x[i] * mid;
1974            let y_val = y[i] * side;
1975            x[i] = x_val - y_val;
1976            y[i] = x_val + y_val;
1977        }
1978    }
1979}
1980
1981#[allow(clippy::too_many_arguments)]
1982#[inline(always)]
1983pub fn quant_band_stereo(
1984    ctx: &mut BandCtx,
1985    x: &mut [f32],
1986    y: &mut [f32],
1987    n: usize,
1988    b: i32,
1989    b_blocks: i32,
1990    lowband: Option<&mut [f32]>,
1991    lm: i32,
1992    lowband_out: Option<&mut [f32]>,
1993    _gain: f32,
1994    fill: u32,
1995) -> u32 {
1996    if n == 1 {
1997        return quant_band_n1(ctx, x, Some(y), lowband_out);
1998    }
1999
2000    if ctx.encode
2001        && (ctx.band_e[ctx.i] < MIN_STEREO_ENERGY
2002            || ctx.band_e[ctx.m.nb_ebands + ctx.i] < MIN_STEREO_ENERGY)
2003    {
2004        if ctx.band_e[ctx.i] > ctx.band_e[ctx.m.nb_ebands + ctx.i] {
2005            y.copy_from_slice(x);
2006        } else {
2007            x.copy_from_slice(y);
2008        }
2009    }
2010
2011    let mut sctx = SplitCtx {
2012        inv: false,
2013        imid: 0,
2014        iside: 0,
2015        delta: 0,
2016        itheta: 0,
2017        qalloc: 0,
2018    };
2019    let mut b_mut = b;
2020    let mut fill_mut = fill;
2021    compute_theta(
2022        ctx,
2023        &mut sctx,
2024        x,
2025        y,
2026        n,
2027        &mut b_mut,
2028        b_blocks,
2029        b_blocks,
2030        lm,
2031        true,
2032        &mut fill_mut,
2033    );
2034
2035    let mid_gain = sctx.imid as f32 / 32768.0;
2036    let side_gain = sctx.iside as f32 / 32768.0;
2037
2038    if n == 2 {
2039        let orig_fill = fill;
2040        let mut mbits = b_mut;
2041        let mut sbits = 0;
2042        if sctx.itheta != 0 && sctx.itheta != 16384 {
2043            sbits = 1 << BITRES;
2044        }
2045        mbits -= sbits;
2046        let c = sctx.itheta > 8192;
2047        ctx.remaining_bits -= sctx.qalloc + sbits;
2048
2049        let mut sign = 0;
2050        if sbits != 0 {
2051            if ctx.encode {
2052                sign = if c {
2053                    if (y[0] * x[1] - y[1] * x[0]) < 0.0 {
2054                        1
2055                    } else {
2056                        0
2057                    }
2058                } else if (x[0] * y[1] - x[1] * y[0]) < 0.0 {
2059                    1
2060                } else {
2061                    0
2062                };
2063                ctx.rc.enc_bits(sign as u32, 1);
2064            } else {
2065                sign = ctx.rc.dec_bits(1) as i32;
2066            }
2067        }
2068        let sign_val = (1 - 2 * sign) as f32;
2069        let cm = if c {
2070            let cm = quant_band(
2071                ctx,
2072                y,
2073                n,
2074                mbits,
2075                b_blocks,
2076                lowband,
2077                lm,
2078                lowband_out,
2079                1.0,
2080                orig_fill,
2081            );
2082            x[0] = -sign_val * y[1];
2083            x[1] = sign_val * y[0];
2084            cm
2085        } else {
2086            let cm = quant_band(
2087                ctx,
2088                x,
2089                n,
2090                mbits,
2091                b_blocks,
2092                lowband,
2093                lm,
2094                lowband_out,
2095                1.0,
2096                orig_fill,
2097            );
2098            y[0] = -sign_val * x[1];
2099            y[1] = sign_val * x[0];
2100            cm
2101        };
2102
2103        if ctx.resynth {
2104            let x0 = x[0];
2105            let x1 = x[1];
2106            let y0 = y[0];
2107            let y1 = y[1];
2108            let mx0 = mid_gain * x0;
2109            let mx1 = mid_gain * x1;
2110            let sy0 = side_gain * y0;
2111            let sy1 = side_gain * y1;
2112            x[0] = mx0 - sy0;
2113            x[1] = mx1 - sy1;
2114            y[0] = mx0 + sy0;
2115            y[1] = mx1 + sy1;
2116        }
2117        return cm;
2118    }
2119
2120    ctx.remaining_bits -= sctx.qalloc;
2121    let mut mbits = (0).max((b_mut - sctx.delta) / 2).min(b_mut);
2122    let mut sbits = b_mut - mbits;
2123
2124    let mut rebalance = ctx.remaining_bits;
2125    let mut cm;
2126
2127    if mbits >= sbits {
2128        cm = quant_band(
2129            ctx,
2130            x,
2131            n,
2132            mbits,
2133            b_blocks,
2134            lowband,
2135            lm,
2136            lowband_out,
2137            1.0,
2138            fill_mut,
2139        );
2140        rebalance = mbits - (rebalance - ctx.remaining_bits);
2141        if rebalance > (3 << 3) && sctx.itheta != 0 {
2142            sbits += rebalance - (3 << 3);
2143        }
2144        cm |= quant_band(
2145            ctx,
2146            y,
2147            n,
2148            sbits,
2149            b_blocks,
2150            None,
2151            lm,
2152            None,
2153            side_gain,
2154            fill_mut >> b_blocks,
2155        );
2156    } else {
2157        cm = quant_band(
2158            ctx,
2159            y,
2160            n,
2161            sbits,
2162            b_blocks,
2163            None,
2164            lm,
2165            None,
2166            side_gain,
2167            fill_mut >> b_blocks,
2168        );
2169        rebalance = sbits - (rebalance - ctx.remaining_bits);
2170        if rebalance > (3 << 3) && sctx.itheta != 16384 {
2171            mbits += rebalance - (3 << 3);
2172        }
2173        cm |= quant_band(
2174            ctx,
2175            x,
2176            n,
2177            mbits,
2178            b_blocks,
2179            lowband,
2180            lm,
2181            lowband_out,
2182            1.0,
2183            fill_mut,
2184        );
2185    }
2186
2187    if ctx.resynth {
2188        stereo_merge(x, y, mid_gain, side_gain, n);
2189        if sctx.inv {
2190            for yv in y[..n].iter_mut() {
2191                *yv = -*yv;
2192            }
2193        }
2194    }
2195    cm
2196}
2197
2198#[allow(clippy::too_many_arguments)]
2199pub fn quant_all_bands(
2200    encode: bool,
2201    m: &CeltMode,
2202    start: usize,
2203    end: usize,
2204    x: &mut [f32],
2205    mut y: Option<&mut [f32]>,
2206    collapse_masks: &mut [u32],
2207    band_e: &[f32],
2208    pulses: &[i32],
2209    short_blocks: bool,
2210    spread: i32,
2211    dual_stereo: &mut bool,
2212    intensity: usize,
2213    tf_res: &[i32],
2214    total_bits: i32,
2215    balance: &mut i32,
2216    rc: &mut RangeCoder,
2217    lm: i32,
2218    coded_bands: i32,
2219    resynth: bool,
2220    disable_inv: bool,
2221    seed: &mut u32,
2222) {
2223    let mut balance_val = *balance;
2224    let b_blocks = if short_blocks { 1 << lm } else { 1 };
2225    let c_channels = if y.is_some() { 2 } else { 1 };
2226    let m_val = 1usize << lm as usize;
2227
2228    let norm_offset = m_val * (m.e_bands[start] as usize);
2229    let norm_size = m_val * (m.e_bands[m.nb_ebands - 1] as usize) - norm_offset;
2230
2231    const MAX_NORM_SIZE: usize = 800;
2232    debug_assert!(norm_size <= MAX_NORM_SIZE);
2233
2234    let mut norm_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_NORM_SIZE];
2235    let norm =
2236        unsafe { core::slice::from_raw_parts_mut(norm_buf.as_mut_ptr() as *mut f32, norm_size) };
2237    let mut norm2_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_NORM_SIZE];
2238    let norm2 =
2239        unsafe { core::slice::from_raw_parts_mut(norm2_buf.as_mut_ptr() as *mut f32, norm_size) };
2240
2241    let mut lowband_scratch_buf = [core::mem::MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
2242    let lowband_scratch_ptr = lowband_scratch_buf.as_mut_ptr() as *mut f32;
2243
2244    let mut lowband_offset: usize = 0;
2245    let mut update_lowband = true;
2246    let mut avoid_split_noise = b_blocks > 1;
2247
2248    let e_bands = &m.e_bands;
2249    let mut ctx_seed = *seed;
2250
2251    for i in start..end {
2252        let e_band_i = e_bands[i] as usize;
2253        let e_band_i1 = e_bands[i + 1] as usize;
2254        let offset = m_val * e_band_i;
2255        let n = m_val * (e_band_i1 - e_band_i);
2256        let last = i == end - 1;
2257
2258        let tell = tell_frac_inline!(rc);
2259        if i != start {
2260            balance_val -= tell;
2261        }
2262        let remaining_bits = total_bits - tell - 1;
2263
2264        let mut b = 0i32;
2265        if i < coded_bands as usize {
2266            let curr_balance = celt_sudiv(balance_val, 3i32.min(coded_bands - i as i32));
2267            b = 0i32.max(16383i32.min((remaining_bits + 1).min(pulses[i] + curr_balance)));
2268        }
2269
2270        let norm_pos = m_val * e_band_i - norm_offset;
2271        let tf_change = tf_res[i];
2272
2273        let mut effective_lowband: i32 = -1;
2274        let mut x_cm: u32;
2275        let mut y_cm: u32;
2276
2277        let band_start_abs = m_val * e_band_i;
2278        let start_abs = m_val * (e_bands[start] as usize);
2279        if resynth
2280            && ((band_start_abs as isize - n as isize >= start_abs as isize) || i == start + 1)
2281            && (update_lowband || lowband_offset == 0)
2282        {
2283            lowband_offset = i;
2284        }
2285
2286        if resynth && i == start + 1 {
2287            special_hybrid_folding(m, norm, start, m_val);
2288            if *dual_stereo {
2289                special_hybrid_folding(m, norm2, start, m_val);
2290            }
2291        }
2292
2293        if lowband_offset != 0 && (spread != SPREAD_AGGRESSIVE || b_blocks > 1 || tf_change < 0) {
2294            effective_lowband = 0i32.max(
2295                (m_val * e_bands[lowband_offset] as usize) as i32 - norm_offset as i32 - n as i32,
2296            );
2297            let el_abs = effective_lowband as usize + norm_offset;
2298
2299            let mut fold_start = lowband_offset;
2300            while fold_start > 0 {
2301                fold_start -= 1;
2302                if m_val * (e_bands[fold_start] as usize) <= el_abs {
2303                    break;
2304                }
2305            }
2306
2307            let mut fold_end = lowband_offset.saturating_sub(1);
2308            loop {
2309                fold_end += 1;
2310                if fold_end >= i || m_val * (e_bands[fold_end] as usize) >= el_abs + n {
2311                    break;
2312                }
2313            }
2314
2315            x_cm = 0;
2316            y_cm = 0;
2317            let mut fi = fold_start;
2318            loop {
2319                x_cm |= collapse_masks[fi * c_channels];
2320                y_cm |= collapse_masks[fi * c_channels + c_channels - 1];
2321                fi += 1;
2322                if fi >= fold_end {
2323                    break;
2324                }
2325            }
2326        } else {
2327            x_cm = (1u32 << b_blocks) - 1;
2328            y_cm = (1u32 << b_blocks) - 1;
2329        }
2330
2331        let mut ctx = BandCtx {
2332            encode,
2333            m,
2334            i,
2335            band_e,
2336            rc,
2337            spread,
2338            remaining_bits,
2339            resynth,
2340            tf_change,
2341            intensity,
2342            theta_round: 0,
2343            avoid_split_noise,
2344            arch: 0,
2345            disable_inv,
2346            seed: ctx_seed,
2347        };
2348
2349        let x_slice = &mut x[offset..offset + n];
2350        let band_uses_direct_norm = i >= m.eff_ebands;
2351        let allow_lowband_scratch = !(band_uses_direct_norm || (last && !encode));
2352        if *dual_stereo && i == intensity {
2353            *dual_stereo = false;
2354            if resynth {
2355                for j in 0..norm_pos {
2356                    norm[j] = 0.5 * (norm[j] + norm2[j]);
2357                }
2358            }
2359        }
2360
2361        if *dual_stereo {
2362            let y_slice = &mut y.as_mut().unwrap()[offset..offset + n];
2363
2364            let (lb_x, lb_out_x) = prepare_lowband_views(
2365                norm,
2366                lowband_scratch_ptr,
2367                allow_lowband_scratch,
2368                effective_lowband,
2369                norm_pos,
2370                n,
2371                !last,
2372            );
2373            x_cm = quant_band(
2374                &mut ctx,
2375                x_slice,
2376                n,
2377                b / 2,
2378                b_blocks,
2379                lb_x,
2380                lm,
2381                lb_out_x,
2382                1.0,
2383                x_cm,
2384            );
2385
2386            let (lb_y, lb_out_y) = prepare_lowband_views(
2387                norm2,
2388                lowband_scratch_ptr,
2389                allow_lowband_scratch,
2390                effective_lowband,
2391                norm_pos,
2392                n,
2393                !last,
2394            );
2395            y_cm = quant_band(
2396                &mut ctx,
2397                y_slice,
2398                n,
2399                b / 2,
2400                b_blocks,
2401                lb_y,
2402                lm,
2403                lb_out_y,
2404                1.0,
2405                y_cm,
2406            );
2407        } else if let Some(y_all) = y.as_mut() {
2408            let y_slice = &mut y_all[offset..offset + n];
2409            let (lb, lb_out) = prepare_lowband_views(
2410                norm,
2411                lowband_scratch_ptr,
2412                allow_lowband_scratch,
2413                effective_lowband,
2414                norm_pos,
2415                n,
2416                !last,
2417            );
2418            x_cm = quant_band_stereo(
2419                &mut ctx,
2420                x_slice,
2421                y_slice,
2422                n,
2423                b,
2424                b_blocks,
2425                lb,
2426                lm,
2427                lb_out,
2428                1.0,
2429                x_cm | y_cm,
2430            );
2431            y_cm = x_cm;
2432        } else {
2433            let (lb, lb_out) = prepare_lowband_views(
2434                norm,
2435                lowband_scratch_ptr,
2436                allow_lowband_scratch,
2437                effective_lowband,
2438                norm_pos,
2439                n,
2440                !last,
2441            );
2442            x_cm = quant_band(&mut ctx, x_slice, n, b, b_blocks, lb, lm, lb_out, 1.0, x_cm);
2443            y_cm = x_cm;
2444        }
2445
2446        collapse_masks[i * c_channels] = (x_cm & 0xFF) as u8 as u32;
2447        if c_channels == 2 {
2448            collapse_masks[i * c_channels + 1] = (y_cm & 0xFF) as u8 as u32;
2449        }
2450
2451        balance_val += pulses[i] + tell;
2452        ctx_seed = ctx.seed;
2453        update_lowband = b > ((n as i32) << BITRES);
2454
2455        avoid_split_noise = false;
2456    }
2457    *balance = balance_val;
2458    *seed = ctx_seed;
2459}
2460
2461#[cfg(target_arch = "aarch64")]
2462fn compute_band_energy_neon(band: &[f32]) -> f32 {
2463    use core::arch::aarch64::*;
2464
2465    let n = band.len();
2466    let mut sum = 1e-27f32;
2467
2468    unsafe {
2469        let n16 = n & !15;
2470        if n16 > 0 {
2471            let mut acc0 = vdupq_n_f32(0.0);
2472            let mut acc1 = vdupq_n_f32(0.0);
2473            let mut acc2 = vdupq_n_f32(0.0);
2474            let mut acc3 = vdupq_n_f32(0.0);
2475
2476            for i in (0..n16).step_by(16) {
2477                let v0 = vld1q_f32(band.as_ptr().add(i));
2478                let v1 = vld1q_f32(band.as_ptr().add(i + 4));
2479                let v2 = vld1q_f32(band.as_ptr().add(i + 8));
2480                let v3 = vld1q_f32(band.as_ptr().add(i + 12));
2481
2482                acc0 = vfmaq_f32(acc0, v0, v0);
2483                acc1 = vfmaq_f32(acc1, v1, v1);
2484                acc2 = vfmaq_f32(acc2, v2, v2);
2485                acc3 = vfmaq_f32(acc3, v3, v3);
2486            }
2487
2488            acc0 = vaddq_f32(acc0, acc1);
2489            acc2 = vaddq_f32(acc2, acc3);
2490            acc0 = vaddq_f32(acc0, acc2);
2491            sum += vaddvq_f32(acc0);
2492        }
2493
2494        let n4 = (n & !3) - n16;
2495        if n4 > 0 {
2496            let mut acc = vdupq_n_f32(0.0);
2497            for i in (n16..n16 + n4).step_by(4) {
2498                let v = vld1q_f32(band.as_ptr().add(i));
2499                acc = vfmaq_f32(acc, v, v);
2500            }
2501            sum += vaddvq_f32(acc);
2502        }
2503
2504        for i in (n16 + n4)..n {
2505            let v = band[i];
2506            sum += v * v;
2507        }
2508    }
2509
2510    sum.sqrt()
2511}
2512
2513#[cfg(target_arch = "x86_64")]
2514#[target_feature(enable = "avx2,fma")]
2515unsafe fn compute_band_energy_avx2(band: &[f32]) -> f32 {
2516    use core::arch::x86_64::*;
2517
2518    let n = band.len();
2519    let mut i = 0usize;
2520
2521    let mut acc0 = _mm256_setzero_ps();
2522    let mut acc1 = _mm256_setzero_ps();
2523
2524    while i + 16 <= n {
2525        let v0 = _mm256_loadu_ps(band.as_ptr().add(i));
2526        let v1 = _mm256_loadu_ps(band.as_ptr().add(i + 8));
2527        acc0 = _mm256_fmadd_ps(v0, v0, acc0);
2528        acc1 = _mm256_fmadd_ps(v1, v1, acc1);
2529        i += 16;
2530    }
2531
2532    if i + 8 <= n {
2533        let v0 = _mm256_loadu_ps(band.as_ptr().add(i));
2534        acc0 = _mm256_fmadd_ps(v0, v0, acc0);
2535        i += 8;
2536    }
2537
2538    let acc = _mm256_add_ps(acc0, acc1);
2539    let hi = _mm256_extractf128_ps(acc, 1);
2540    let lo = _mm256_castps256_ps128(acc);
2541    let s4 = _mm_add_ps(lo, hi);
2542    let t1 = _mm_movehl_ps(s4, s4);
2543    let s2 = _mm_add_ps(s4, t1);
2544    let t2 = _mm_shuffle_ps(s2, s2, 0x55);
2545    let mut sum = 1e-27f32 + _mm_cvtss_f32(_mm_add_ss(s2, t2));
2546
2547    for &v in &band[i..] {
2548        sum += v * v;
2549    }
2550
2551    sum.sqrt()
2552}
2553
2554pub fn compute_band_energies(
2555    m: &CeltMode,
2556    x: &[f32],
2557    band_e: &mut [f32],
2558    end: usize,
2559    channels: usize,
2560    lm: usize,
2561) {
2562    let frame_size = m.short_mdct_size << lm;
2563
2564    #[cfg(target_arch = "x86_64")]
2565    let use_avx2 = crate::compat::x86_has_avx2();
2566
2567    for c in 0..channels {
2568        let ch = &x[c * frame_size..(c + 1) * frame_size];
2569        for i in 0..end {
2570            let offset = (m.e_bands[i] as usize) << lm;
2571            let n = ((m.e_bands[i + 1] - m.e_bands[i]) as usize) << lm;
2572            let band = &ch[offset..offset + n];
2573
2574            #[cfg(target_arch = "aarch64")]
2575            {
2576                band_e[c * m.nb_ebands + i] = compute_band_energy_neon(band);
2577            }
2578            #[cfg(target_arch = "x86_64")]
2579            {
2580                if n >= 8 && use_avx2 {
2581                    band_e[c * m.nb_ebands + i] = unsafe { compute_band_energy_avx2(band) };
2582                } else {
2583                    let sum = band.iter().fold(1e-27f32, |acc, &v| acc + v * v);
2584                    band_e[c * m.nb_ebands + i] = sum.sqrt();
2585                }
2586            }
2587            #[cfg(all(not(target_arch = "aarch64"), not(target_arch = "x86_64")))]
2588            {
2589                let sum = band.iter().fold(1e-27f32, |acc, &v| acc + v * v);
2590                band_e[c * m.nb_ebands + i] = sum.sqrt();
2591            }
2592        }
2593    }
2594}
2595
2596pub fn amp2log2(
2597    m: &CeltMode,
2598    start: usize,
2599    end: usize,
2600    band_e: &[f32],
2601    band_log_e: &mut [f32],
2602    channels: usize,
2603) {
2604    for c in 0..channels {
2605        for i in 0..start {
2606            band_log_e[c * m.nb_ebands + i] = -14.0;
2607        }
2608        for i in start..end {
2609            let val = band_e[c * m.nb_ebands + i].max(1e-10);
2610            band_log_e[c * m.nb_ebands + i] = val.log2() - m.e_means[i];
2611        }
2612    }
2613}
2614
2615pub fn log2amp(m: &CeltMode, end: usize, band_e: &mut [f32], band_log_e: &[f32], channels: usize) {
2616    for c in 0..channels {
2617        for i in 0..end {
2618            band_e[c * m.nb_ebands + i] = band_log_e[c * m.nb_ebands + i] + m.e_means[i];
2619        }
2620    }
2621}
2622
2623pub fn normalise_bands(
2624    m: &CeltMode,
2625    freq: &[f32],
2626    x: &mut [f32],
2627    band_e: &[f32],
2628    end: usize,
2629    channels: usize,
2630    m_val: usize,
2631) {
2632    let lm = m_val.trailing_zeros() as usize;
2633    let frame_size = m.short_mdct_size << lm;
2634    #[cfg(target_arch = "x86_64")]
2635    let use_avx2 = crate::compat::x86_has_avx2();
2636    for c in 0..channels {
2637        for i in 0..end {
2638            let base = c * frame_size + ((m.e_bands[i] as usize) << lm);
2639            let n = ((m.e_bands[i + 1] - m.e_bands[i]) as usize) << lm;
2640            let norm = 1.0 / (1e-27 + band_e[c * m.nb_ebands + i]);
2641            let src = &freq[base..base + n];
2642            let dst = &mut x[base..base + n];
2643            #[cfg(target_arch = "x86_64")]
2644            if n >= 8 && use_avx2 {
2645                unsafe { scale_slice_avx2(src, dst, norm, n) };
2646                continue;
2647            }
2648            #[cfg(target_arch = "aarch64")]
2649            if n >= 8 {
2650                unsafe { scale_slice_neon(src, dst, norm, n) };
2651                continue;
2652            }
2653            for (d, &s) in dst.iter_mut().zip(src) {
2654                *d = s * norm;
2655            }
2656        }
2657    }
2658}
2659
2660#[cfg(target_arch = "x86_64")]
2661#[target_feature(enable = "avx2")]
2662unsafe fn scale_slice_avx2(src: &[f32], dst: &mut [f32], scale: f32, n: usize) {
2663    use core::arch::x86_64::*;
2664    let vscale = _mm256_set1_ps(scale);
2665    let mut i = 0;
2666
2667    while i + 16 <= n {
2668        let s0 = _mm256_loadu_ps(src.as_ptr().add(i));
2669        let s1 = _mm256_loadu_ps(src.as_ptr().add(i + 8));
2670        _mm256_storeu_ps(dst.as_mut_ptr().add(i), _mm256_mul_ps(s0, vscale));
2671        _mm256_storeu_ps(dst.as_mut_ptr().add(i + 8), _mm256_mul_ps(s1, vscale));
2672        i += 16;
2673    }
2674    while i + 8 <= n {
2675        let sv = _mm256_loadu_ps(src.as_ptr().add(i));
2676        _mm256_storeu_ps(dst.as_mut_ptr().add(i), _mm256_mul_ps(sv, vscale));
2677        i += 8;
2678    }
2679    for j in i..n {
2680        dst[j] = src[j] * scale;
2681    }
2682}
2683
2684#[cfg(target_arch = "aarch64")]
2685#[inline(always)]
2686#[allow(unsafe_op_in_unsafe_fn)]
2687unsafe fn scale_slice_neon(src: &[f32], dst: &mut [f32], scale: f32, n: usize) {
2688    use core::arch::aarch64::*;
2689    let vscale = vdupq_n_f32(scale);
2690    let mut i = 0;
2691
2692    while i + 16 <= n {
2693        let s0 = vld1q_f32(src.as_ptr().add(i));
2694        let s1 = vld1q_f32(src.as_ptr().add(i + 4));
2695        let s2 = vld1q_f32(src.as_ptr().add(i + 8));
2696        let s3 = vld1q_f32(src.as_ptr().add(i + 12));
2697        vst1q_f32(dst.as_mut_ptr().add(i), vmulq_f32(s0, vscale));
2698        vst1q_f32(dst.as_mut_ptr().add(i + 4), vmulq_f32(s1, vscale));
2699        vst1q_f32(dst.as_mut_ptr().add(i + 8), vmulq_f32(s2, vscale));
2700        vst1q_f32(dst.as_mut_ptr().add(i + 12), vmulq_f32(s3, vscale));
2701        i += 16;
2702    }
2703    while i + 8 <= n {
2704        let s0 = vld1q_f32(src.as_ptr().add(i));
2705        let s1 = vld1q_f32(src.as_ptr().add(i + 4));
2706        vst1q_f32(dst.as_mut_ptr().add(i), vmulq_f32(s0, vscale));
2707        vst1q_f32(dst.as_mut_ptr().add(i + 4), vmulq_f32(s1, vscale));
2708        i += 8;
2709    }
2710    while i + 4 <= n {
2711        let s0 = vld1q_f32(src.as_ptr().add(i));
2712        vst1q_f32(dst.as_mut_ptr().add(i), vmulq_f32(s0, vscale));
2713        i += 4;
2714    }
2715    for j in i..n {
2716        dst[j] = src[j] * scale;
2717    }
2718}
2719
2720#[allow(clippy::too_many_arguments)]
2721pub fn denormalise_bands(
2722    m: &CeltMode,
2723    x: &[f32],
2724    freq: &mut [f32],
2725    band_e: &[f32],
2726    start: usize,
2727    end: usize,
2728    channels: usize,
2729    m_val: usize,
2730) {
2731    let lm = m_val.trailing_zeros() as usize;
2732    let frame_size = m.short_mdct_size << lm;
2733    #[cfg(target_arch = "x86_64")]
2734    let use_avx2 = crate::compat::x86_has_avx2();
2735
2736    for c in 0..channels {
2737        for i in start..end {
2738            let base = c * frame_size + ((m.e_bands[i] as usize) << lm);
2739            let n = ((m.e_bands[i + 1] - m.e_bands[i]) as usize) << lm;
2740            let band_log = band_e[c * m.nb_ebands + i];
2741
2742            // Match C: celt_exp2_db(MIN32(32.f, lg)) — cap gain to prevent overflow
2743            let g = (2.0f32).powf(band_log.min(32.0));
2744            let src = &x[base..base + n];
2745            let dst = &mut freq[base..base + n];
2746            #[cfg(target_arch = "x86_64")]
2747            if n >= 8 && use_avx2 {
2748                unsafe { scale_slice_avx2(src, dst, g, n) };
2749                continue;
2750            }
2751            #[cfg(target_arch = "aarch64")]
2752            if n >= 8 {
2753                unsafe { scale_slice_neon(src, dst, g, n) };
2754                continue;
2755            }
2756            for (d, &s) in dst.iter_mut().zip(src) {
2757                *d = s * g;
2758            }
2759        }
2760    }
2761}
2762
2763pub fn celt_lcg_rand(seed: u32) -> u32 {
2764    seed.wrapping_mul(1664525).wrapping_add(1013904223)
2765}
2766
2767#[cfg(target_arch = "aarch64")]
2768#[inline(always)]
2769#[allow(unsafe_op_in_unsafe_fn)]
2770unsafe fn renormalise_vector_neon(x: &mut [f32], n: usize, gain: f32) {
2771    use core::arch::aarch64::*;
2772
2773    let mut sum_vec = vdupq_n_f32(0.0);
2774    let mut i = 0;
2775
2776    while i + 16 <= n {
2777        let x0 = vld1q_f32(x.as_ptr().add(i));
2778        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2779        let x2 = vld1q_f32(x.as_ptr().add(i + 8));
2780        let x3 = vld1q_f32(x.as_ptr().add(i + 12));
2781        sum_vec = vfmaq_f32(sum_vec, x0, x0);
2782        sum_vec = vfmaq_f32(sum_vec, x1, x1);
2783        sum_vec = vfmaq_f32(sum_vec, x2, x2);
2784        sum_vec = vfmaq_f32(sum_vec, x3, x3);
2785        i += 16;
2786    }
2787
2788    while i + 8 <= n {
2789        let x0 = vld1q_f32(x.as_ptr().add(i));
2790        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2791        sum_vec = vfmaq_f32(sum_vec, x0, x0);
2792        sum_vec = vfmaq_f32(sum_vec, x1, x1);
2793        i += 8;
2794    }
2795
2796    while i + 4 <= n {
2797        let x0 = vld1q_f32(x.as_ptr().add(i));
2798        sum_vec = vfmaq_f32(sum_vec, x0, x0);
2799        i += 4;
2800    }
2801
2802    let mut e = 1e-15f32 + vaddvq_f32(sum_vec);
2803
2804    for j in i..n {
2805        e += x[j] * x[j];
2806    }
2807
2808    let norm = gain / e.sqrt();
2809    let vnorm = vdupq_n_f32(norm);
2810
2811    i = 0;
2812    while i + 16 <= n {
2813        let x0 = vld1q_f32(x.as_ptr().add(i));
2814        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2815        let x2 = vld1q_f32(x.as_ptr().add(i + 8));
2816        let x3 = vld1q_f32(x.as_ptr().add(i + 12));
2817        vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(x0, vnorm));
2818        vst1q_f32(x.as_mut_ptr().add(i + 4), vmulq_f32(x1, vnorm));
2819        vst1q_f32(x.as_mut_ptr().add(i + 8), vmulq_f32(x2, vnorm));
2820        vst1q_f32(x.as_mut_ptr().add(i + 12), vmulq_f32(x3, vnorm));
2821        i += 16;
2822    }
2823
2824    while i + 8 <= n {
2825        let x0 = vld1q_f32(x.as_ptr().add(i));
2826        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2827        vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(x0, vnorm));
2828        vst1q_f32(x.as_mut_ptr().add(i + 4), vmulq_f32(x1, vnorm));
2829        i += 8;
2830    }
2831
2832    while i + 4 <= n {
2833        let x0 = vld1q_f32(x.as_ptr().add(i));
2834        vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(x0, vnorm));
2835        i += 4;
2836    }
2837
2838    for j in i..n {
2839        x[j] *= norm;
2840    }
2841}
2842
2843#[cfg(target_arch = "x86_64")]
2844#[target_feature(enable = "avx2,fma")]
2845unsafe fn renormalise_vector_avx2(x: &mut [f32], n: usize, gain: f32) {
2846    use core::arch::x86_64::*;
2847
2848    let mut i = 0usize;
2849
2850    let mut acc0 = _mm256_setzero_ps();
2851    let mut acc1 = _mm256_setzero_ps();
2852
2853    while i + 16 <= n {
2854        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
2855        let v1 = _mm256_loadu_ps(x.as_ptr().add(i + 8));
2856        acc0 = _mm256_fmadd_ps(v0, v0, acc0);
2857        acc1 = _mm256_fmadd_ps(v1, v1, acc1);
2858        i += 16;
2859    }
2860
2861    if i + 8 <= n {
2862        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
2863        acc0 = _mm256_fmadd_ps(v0, v0, acc0);
2864        i += 8;
2865    }
2866
2867    let acc = _mm256_add_ps(acc0, acc1);
2868    let hi = _mm256_extractf128_ps(acc, 1);
2869    let lo = _mm256_castps256_ps128(acc);
2870    let s4 = _mm_add_ps(lo, hi);
2871    let t1 = _mm_movehl_ps(s4, s4);
2872    let s2 = _mm_add_ps(s4, t1);
2873    let t2 = _mm_shuffle_ps(s2, s2, 0x55);
2874    let mut e = 1e-15f32 + _mm_cvtss_f32(_mm_add_ss(s2, t2));
2875
2876    for &v in &x[i..n] {
2877        e += v * v;
2878    }
2879
2880    let norm = gain / e.sqrt();
2881    let vnorm = _mm256_set1_ps(norm);
2882
2883    i = 0;
2884    while i + 16 <= n {
2885        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
2886        let v1 = _mm256_loadu_ps(x.as_ptr().add(i + 8));
2887        _mm256_storeu_ps(x.as_mut_ptr().add(i), _mm256_mul_ps(v0, vnorm));
2888        _mm256_storeu_ps(x.as_mut_ptr().add(i + 8), _mm256_mul_ps(v1, vnorm));
2889        i += 16;
2890    }
2891    while i + 8 <= n {
2892        let v = _mm256_loadu_ps(x.as_ptr().add(i));
2893        _mm256_storeu_ps(x.as_mut_ptr().add(i), _mm256_mul_ps(v, vnorm));
2894        i += 8;
2895    }
2896    for v in &mut x[i..n] {
2897        *v *= norm;
2898    }
2899}
2900
2901pub fn renormalise_vector(x: &mut [f32], n: usize, gain: f32) {
2902    #[cfg(target_arch = "aarch64")]
2903    unsafe {
2904        renormalise_vector_neon(x, n, gain);
2905    }
2906    #[cfg(target_arch = "x86_64")]
2907    unsafe {
2908        if n >= 16 && crate::compat::x86_has_avx2() {
2909            renormalise_vector_avx2(x, n, gain);
2910            return;
2911        }
2912    }
2913    #[cfg(all(not(target_arch = "aarch64"), not(target_arch = "x86_64")))]
2914    {
2915        let mut e = 1e-15f32;
2916        for &xv in x[..n].iter() {
2917            e += xv * xv;
2918        }
2919        let norm = gain / e.sqrt();
2920        for xv in x[..n].iter_mut() {
2921            *xv *= norm;
2922        }
2923    }
2924    #[cfg(target_arch = "x86_64")]
2925    {
2926        let mut e = 1e-15f32;
2927        for &xv in x[..n].iter() {
2928            e += xv * xv;
2929        }
2930        let norm = gain / e.sqrt();
2931        for xv in x[..n].iter_mut() {
2932            *xv *= norm;
2933        }
2934    }
2935}
2936
2937#[allow(clippy::too_many_arguments)]
2938pub fn anti_collapse(
2939    m: &CeltMode,
2940    x_buf: &mut [f32],
2941    collapse_masks: &[u32],
2942    lm: i32,
2943    channels: usize,
2944    size: usize,
2945    start: usize,
2946    end: usize,
2947    log_e: &[f32],
2948    prev1_log_e: &[f32],
2949    prev2_log_e: &[f32],
2950    pulses: &[i32],
2951    mut seed: u32,
2952) -> u32 {
2953    for i in start..end {
2954        let n0 = (m.e_bands[i + 1] - m.e_bands[i]) as usize;
2955        let depth = if n0 > 0 {
2956            ((1 + pulses[i]) / n0 as i32) >> lm
2957        } else {
2958            0
2959        };
2960
2961        let thresh = 0.5 * (-(0.125 * depth as f32)).exp2();
2962        let sqrt_1 = 1.0 / ((n0 << lm) as f32).sqrt();
2963
2964        for c in 0..channels {
2965            let p1 = prev1_log_e[c * m.nb_ebands + i];
2966            let p2 = prev2_log_e[c * m.nb_ebands + i];
2967
2968            let (p1_adj, p2_adj) = if channels == 1 && prev1_log_e.len() >= 2 * m.nb_ebands {
2969                (
2970                    p1.max(prev1_log_e[m.nb_ebands + i]),
2971                    p2.max(prev2_log_e[m.nb_ebands + i]),
2972                )
2973            } else {
2974                (p1, p2)
2975            };
2976
2977            let e_diff = log_e[c * m.nb_ebands + i] - p1_adj.min(p2_adj);
2978            let e_diff = e_diff.max(0.0);
2979
2980            let mut r = 2.0 * (-e_diff).exp2();
2981            if lm == 3 {
2982                r *= core::f32::consts::SQRT_2;
2983            }
2984            r = r.min(thresh);
2985            r *= sqrt_1;
2986
2987            let x_offset = c * size + ((m.e_bands[i] as usize) << lm);
2988            let mut renormalize = false;
2989            for k in 0..(1 << lm) {
2990                if (collapse_masks[i * channels + c] & (1 << k)) == 0 {
2991                    for j in 0..n0 {
2992                        seed = celt_lcg_rand(seed);
2993                        x_buf[x_offset + (j << lm) + k] = if (seed & 0x8000) != 0 { r } else { -r };
2994                    }
2995                    renormalize = true;
2996                }
2997            }
2998            if renormalize {
2999                renormalise_vector(&mut x_buf[x_offset..x_offset + (n0 << lm)], n0 << lm, 1.0);
3000            }
3001        }
3002    }
3003    seed
3004}
3005
3006#[cfg(all(test, feature = "std"))]
3007mod tests {
3008    use super::*;
3009
3010    #[test]
3011    fn test_bitexact_primitives_reference_values() {
3012        assert_eq!(bitexact_cos(64), 32767);
3013        assert_eq!(bitexact_cos(8192), 23171);
3014        assert_eq!(bitexact_cos(16320), 200);
3015
3016        assert_eq!(bitexact_log2tan(32767, 200), 15059);
3017        assert_eq!(bitexact_log2tan(30274, 12540), 2611);
3018        assert_eq!(bitexact_log2tan(23171, 23171), 0);
3019        assert_eq!(bitexact_log2tan(200, 32767), -15059);
3020    }
3021}