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 haar1_scalar(x, n0, stride);
205}
206
207#[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#[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; 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 if n == 2 {
919 return quant_partition_n2_encode(ctx, x, b, b_blocks, lowband, lm, gain, fill);
920 }
921
922 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 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 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 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; 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 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 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 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 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 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 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}