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