1use crate::modes::CeltMode;
2use crate::range_coder::RangeCoder;
3use std::cmp::{max, min};
4
5const MAX_EBANDS: usize = 21;
6pub const BITRES: i32 = 3;
7pub const FINE_OFFSET: i32 = 21;
8pub const QTHETA_OFFSET: i32 = 4;
9pub const QTHETA_OFFSET_TWOPHASE: i32 = 16;
10pub const MAX_FINE_BITS: i32 = 8;
11
12pub const LOG2_FRAC_TABLE: [u8; 24] = [
13 0, 8, 13, 16, 19, 21, 23, 24, 26, 27, 28, 29, 30, 31, 32, 32, 33, 34, 34, 35, 36, 36, 37, 37,
14];
15
16#[inline(always)]
17pub fn get_pulses(i: i32) -> i32 {
18 if i < 8 {
19 i
20 } else {
21 let shift = (i >> 3) - 1;
22 if shift >= 31 {
23 return 0x7FFFFFFF;
24 }
25 (8 + (i & 7)) << shift
26 }
27}
28
29#[inline(always)]
30pub fn bits2pulses(m: &CeltMode, band: usize, mut lm: i32, bits: i32) -> i32 {
31 lm += 1;
32 let idx = lm as usize * m.nb_ebands + band;
33 let cache_index = unsafe { *m.cache.index.get_unchecked(idx) };
34 if cache_index < 0 {
35 return 0;
36 }
37 let cache = &m.cache.bits[cache_index as usize..];
38 let cache_ptr = cache.as_ptr();
39
40 let mut lo = 0i32;
41 let mut hi = unsafe { *cache_ptr } as i32;
42 let bits = bits - 1; unsafe {
45 for _ in 0..6 {
46 let mid = (lo + hi + 1) >> 1; if *cache_ptr.add(mid as usize) as i32 >= bits {
49 hi = mid;
50 } else {
51 lo = mid;
52 }
53 }
54
55 let lo_val = if lo == 0 {
56 -1i32
57 } else {
58 *cache_ptr.add(lo as usize) as i32
59 };
60 let hi_val = *cache_ptr.add(hi as usize) as i32;
61 if bits - lo_val <= hi_val - bits {
62 lo
63 } else {
64 hi
65 }
66 }
67}
68
69#[inline(always)]
70pub fn pulses2bits(m: &CeltMode, band: usize, mut lm: i32, pulses: i32) -> i32 {
71 if pulses == 0 {
72 return 0;
73 }
74 lm += 1;
75 let idx = lm as usize * m.nb_ebands + band;
76 let cache_index = unsafe { *m.cache.index.get_unchecked(idx) };
77 if cache_index < 0 {
78 return 0;
79 }
80 let cache = &m.cache.bits[cache_index as usize..];
81
82 unsafe { (*cache.as_ptr().add(pulses as usize) as i32) + 1 }
83}
84
85#[allow(clippy::too_many_arguments)]
86pub fn clt_compute_allocation(
87 m: &CeltMode,
88 start: usize,
89 end: usize,
90 offsets: &[i32],
91 cap: &[i32],
92 alloc_trim: i32,
93 intensity: &mut i32,
94 dual_stereo: &mut i32,
95 mut total: i32,
96 balance_out: &mut i32,
97 pulses: &mut [i32],
98 ebits: &mut [i32],
99 fine_priority: &mut [i32],
100 c: i32,
101 lm: i32,
102 rc: &mut RangeCoder,
103 encode: bool,
104 prev: i32,
105 signal_bandwidth: i32,
106) -> i32 {
107 let _prof = crate::prof::scope(crate::prof::Stage::CeltAlloc);
108 total = max(total, 0);
109 let nb_ebands = m.nb_ebands;
110 let mut skip_start = start;
111
112 let skip_rsv = if total >= (1 << BITRES) {
113 1 << BITRES
114 } else {
115 0
116 };
117 total -= skip_rsv;
118
119 let mut intensity_rsv = 0;
120 let mut dual_stereo_rsv = 0;
121 if c == 2 {
122 intensity_rsv = LOG2_FRAC_TABLE[end - start] as i32;
123 if intensity_rsv > total {
124 intensity_rsv = 0;
125 } else {
126 total -= intensity_rsv;
127 dual_stereo_rsv = if total >= (1 << BITRES) {
128 1 << BITRES
129 } else {
130 0
131 };
132 total -= dual_stereo_rsv;
133 }
134 }
135
136 let mut thresh_buf = [0i32; MAX_EBANDS];
137 let thresh = &mut thresh_buf[..nb_ebands];
138 let mut trim_offset_buf = [0i32; MAX_EBANDS];
139 let trim_offset = &mut trim_offset_buf[..nb_ebands];
140
141 for j in start..end {
142 thresh[j] = max(
143 c << BITRES,
144 ((3 * (m.e_bands[j + 1] - m.e_bands[j]) as i32) << (lm + BITRES)) >> 4,
145 );
146 trim_offset[j] = (c
147 * (m.e_bands[j + 1] - m.e_bands[j]) as i32
148 * (alloc_trim - 5 - lm)
149 * (end - j - 1) as i32
150 * (1 << (lm + BITRES)))
151 >> 6;
152 if (m.e_bands[j + 1] - m.e_bands[j]) << lm == 1 {
153 trim_offset[j] -= c << BITRES;
154 }
155 }
156
157 let mut lo = 1;
158 let mut hi = m.nb_alloc_vectors as i32 - 1;
159 while lo <= hi {
160 let mut done = false;
161 let mut psum = 0;
162 let mid = (lo + hi) >> 1;
163 for j in (start..end).rev() {
164 let n = (m.e_bands[j + 1] - m.e_bands[j]) as i32;
165 let raw = m.alloc_vectors[mid as usize * m.alloc_stride + j] as i32;
166 let mut bitsj = (c * n * raw) << lm >> 2;
167 if bitsj > 0 {
168 bitsj = max(0, bitsj + trim_offset[j]);
169 }
170 bitsj += offsets[j];
171 if bitsj >= thresh[j] || done {
172 done = true;
173 psum += min(bitsj, cap[j]);
174 } else if bitsj >= (c << BITRES) {
175 psum += c << BITRES;
176 }
177 }
178 if psum > total {
179 hi = mid - 1;
180 } else {
181 lo = mid + 1;
182 }
183 }
184
185 let hi_final = lo as usize;
186 let lo_final = (lo - 1) as usize;
187
188 let mut bits1_buf = [0i32; MAX_EBANDS];
189 let bits1 = &mut bits1_buf[..nb_ebands];
190 let mut bits2_buf = [0i32; MAX_EBANDS];
191 let bits2 = &mut bits2_buf[..nb_ebands];
192
193 for j in start..end {
194 let n = (m.e_bands[j + 1] - m.e_bands[j]) as i32;
195 let mut bits1j = (c * n * m.alloc_vectors[lo_final * m.alloc_stride + j] as i32) << lm >> 2;
196 let mut bits2j = if hi_final >= m.nb_alloc_vectors {
197 cap[j]
198 } else {
199 (c * n * m.alloc_vectors[hi_final * m.alloc_stride + j] as i32) << lm >> 2
200 };
201
202 if bits1j > 0 {
203 bits1j = max(0, bits1j + trim_offset[j]);
204 }
205 if bits2j > 0 {
206 bits2j = max(0, bits2j + trim_offset[j]);
207 }
208 if lo_final > 0 {
209 bits1j += offsets[j];
210 }
211 bits2j += offsets[j];
212 if offsets[j] > 0 {
213 skip_start = j;
214 }
215 bits2j = max(0, bits2j - bits1j);
216 bits1[j] = bits1j;
217 bits2[j] = bits2j;
218 }
219
220 interp_bits2pulses(
221 m,
222 start,
223 end,
224 skip_start,
225 bits1,
226 bits2,
227 thresh,
228 cap,
229 total,
230 balance_out,
231 skip_rsv,
232 intensity,
233 intensity_rsv,
234 dual_stereo,
235 dual_stereo_rsv,
236 pulses,
237 ebits,
238 fine_priority,
239 c,
240 lm,
241 rc,
242 encode,
243 prev,
244 signal_bandwidth,
245 )
246}
247
248#[allow(clippy::too_many_arguments)]
249fn interp_bits2pulses(
250 m: &CeltMode,
251 start: usize,
252 end: usize,
253 skip_start: usize,
254 bits1: &[i32],
255 bits2: &[i32],
256 thresh: &[i32],
257 cap: &[i32],
258 total: i32,
259 balance_out: &mut i32,
260 skip_rsv: i32,
261 intensity: &mut i32,
262 mut intensity_rsv: i32,
263 dual_stereo: &mut i32,
264 dual_stereo_rsv: i32,
265 pulses: &mut [i32],
266 ebits: &mut [i32],
267 fine_priority: &mut [i32],
268 c: i32,
269 lm: i32,
270 rc: &mut RangeCoder,
271 encode: bool,
272 prev: i32,
273 signal_bandwidth: i32,
274) -> i32 {
275 let mut psum: i32;
276 let mut lo = 0;
277 let mut hi = 1 << 6;
278 let alloc_floor = c << BITRES;
279 let stereo = if c > 1 { 1 } else { 0 };
280 let log_m = lm << BITRES;
281
282 let mut bits_buf = [0i32; MAX_EBANDS];
283 let bits = &mut bits_buf[..m.nb_ebands];
284
285 for _ in 0..6 {
286 let mid = (lo + hi) >> 1;
287 psum = 0;
288 let mut done = false;
289 for j in (start..end).rev() {
290 let tmp = bits1[j] + ((mid * bits2[j]) >> 6);
291 if tmp >= thresh[j] || done {
292 done = true;
293 psum += min(tmp, cap[j]);
294 } else if tmp >= alloc_floor {
295 psum += alloc_floor;
296 }
297 }
298 if psum > total {
299 hi = mid;
300 } else {
301 lo = mid;
302 }
303 }
304 psum = 0;
305 let mut done = false;
306 for j in (start..end).rev() {
307 let mut tmp = bits1[j] + ((lo * bits2[j]) >> 6);
308 if tmp < thresh[j] && !done {
309 if tmp >= alloc_floor {
310 tmp = alloc_floor;
311 } else {
312 tmp = 0;
313 }
314 } else {
315 done = true;
316 }
317 tmp = min(tmp, cap[j]);
318 bits[j] = tmp;
319 psum += tmp;
320 }
321
322 let mut coded_bands = end;
323 let mut total_with_rsv = total;
324 loop {
325 if coded_bands <= start {
326 break;
327 }
328 let j = coded_bands - 1;
329 if j <= skip_start {
330 total_with_rsv += skip_rsv;
331 break;
332 }
333
334 let left = total_with_rsv - psum;
335 let nb_samples = (m.e_bands[coded_bands] - m.e_bands[start]) as i32;
336 let percoeff = left / nb_samples;
337 let left_rem = left - nb_samples * percoeff;
338 let rem = max(left_rem - (m.e_bands[j] - m.e_bands[start]) as i32, 0);
339 let band_width = (m.e_bands[coded_bands] - m.e_bands[j]) as i32;
340 let mut band_bits = bits[j] + percoeff * band_width + rem;
341
342 if band_bits >= max(thresh[j], alloc_floor + (1 << BITRES)) {
343 if encode {
344 let depth_threshold = if coded_bands > 17 {
345 if (j as i32) < prev { 7 } else { 9 }
346 } else {
347 0
348 };
349 if coded_bands <= start + 2
350 || (band_bits > ((depth_threshold * band_width) << lm << BITRES) >> 4
351 && (j as i32) <= signal_bandwidth)
352 {
353 rc.encode_bit_logp(true, 1);
354 break;
355 }
356 rc.encode_bit_logp(false, 1);
357 } else {
358 let bit = rc.decode_bit_logp(1);
359 if bit {
360 break;
361 }
362 }
363 psum += 1 << BITRES;
364 band_bits -= 1 << BITRES;
365 }
366 psum -= bits[j] + intensity_rsv;
367 if intensity_rsv > 0 {
368 intensity_rsv = LOG2_FRAC_TABLE[j - start] as i32;
369 }
370 psum += intensity_rsv;
371 if band_bits >= alloc_floor {
372 psum += alloc_floor;
373 bits[j] = alloc_floor;
374 } else {
375 bits[j] = 0;
376 }
377 coded_bands -= 1;
378 }
379
380 if intensity_rsv > 0 {
381 if encode {
382 *intensity = min(*intensity, coded_bands as i32);
383 rc.enc_uint(
384 (*intensity - start as i32) as u32,
385 (coded_bands + 1 - start) as u32,
386 );
387 } else {
388 *intensity = start as i32 + rc.dec_uint((coded_bands + 1 - start) as u32) as i32;
389 }
390 } else {
391 *intensity = 0;
392 }
393
394 let mut dual_stereo_rsv_final = dual_stereo_rsv;
395 if *intensity <= start as i32 {
396 total_with_rsv += dual_stereo_rsv_final;
397 dual_stereo_rsv_final = 0;
398 }
399 if dual_stereo_rsv_final > 0 {
400 if encode {
401 rc.encode_bit_logp(*dual_stereo != 0, 1);
402 } else {
403 *dual_stereo = if rc.decode_bit_logp(1) { 1 } else { 0 };
404 }
405 } else {
406 *dual_stereo = 0;
407 }
408
409 let mut left = total_with_rsv - psum;
410 let nb_samples = (m.e_bands[coded_bands] - m.e_bands[start]) as i32;
411 let percoeff = left / nb_samples;
412 left -= nb_samples * percoeff;
413 for (j, bits_j) in bits[start..coded_bands]
414 .iter_mut()
415 .enumerate()
416 .map(|(i, v)| (i + start, v))
417 {
418 *bits_j += percoeff * (m.e_bands[j + 1] - m.e_bands[j]) as i32;
419 }
420 for (j, bits_j) in bits[start..coded_bands]
421 .iter_mut()
422 .enumerate()
423 .map(|(i, v)| (i + start, v))
424 {
425 let tmp = min(left, (m.e_bands[j + 1] - m.e_bands[j]) as i32);
426 *bits_j += tmp;
427 left -= tmp;
428 }
429
430 let mut balance = 0;
431 for j in start..coded_bands {
432 let n0 = (m.e_bands[j + 1] - m.e_bands[j]) as i32;
433 let n = n0 << lm;
434 let bit = bits[j] + balance;
435
436 let mut excess;
437 if n > 1 {
438 excess = max(bit - cap[j], 0);
439 bits[j] = bit - excess;
440
441 let den = c * n
442 + (if c == 2 && n > 2 && *dual_stereo == 0 && (j as i32) < *intensity {
443 1
444 } else {
445 0
446 });
447 let nc_log_n = den * (m.log_n[j] as i32 + log_m);
448 let mut offset = (nc_log_n >> 1) - den * FINE_OFFSET;
449
450 if n == 2 {
451 offset += den << BITRES >> 2;
452 }
453
454 if bits[j] + offset < (den * 2) << BITRES {
455 offset += nc_log_n >> 2;
456 } else if bits[j] + offset < (den * 3) << BITRES {
457 offset += nc_log_n >> 3;
458 }
459
460 ebits[j] = max(0, bits[j] + offset + (den << (BITRES - 1)));
461
462 let num = ebits[j];
463 if den > 0 {
464 ebits[j] = ((num as u32 / den as u32) >> BITRES) as i32;
465 } else {
466 ebits[j] = 0;
467 }
468
469 if c * ebits[j] > (bits[j] >> BITRES) {
470 ebits[j] = bits[j] >> stereo >> BITRES;
471 }
472 ebits[j] = min(ebits[j], MAX_FINE_BITS);
473 fine_priority[j] = if ebits[j] * (den << BITRES) >= bits[j] + offset {
474 1
475 } else {
476 0
477 };
478 bits[j] -= (c * ebits[j]) << BITRES;
479 } else {
480 excess = max(0, bit - (c << BITRES));
481 bits[j] = bit - excess;
482 ebits[j] = 0;
483 fine_priority[j] = 1;
484 }
485
486 if excess > 0 {
487 let extra_fine = min(excess >> (stereo + BITRES), MAX_FINE_BITS - ebits[j]);
488 ebits[j] += extra_fine;
489 let extra_bits = (extra_fine * c) << BITRES;
490 fine_priority[j] = if extra_bits >= excess - balance { 1 } else { 0 };
491 excess -= extra_bits;
492 }
493 balance = excess;
494 pulses[j] = bits[j];
495 }
496 *balance_out = balance;
497
498 for j in coded_bands..end {
499 ebits[j] = bits[j] >> stereo >> BITRES;
500 bits[j] = 0;
501 fine_priority[j] = if ebits[j] < 1 { 1 } else { 0 };
502 pulses[j] = 0;
503 }
504
505 coded_bands as i32
506}