Skip to main content

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