Skip to main content

rusty_opus/
rate.rs

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; // bits--
43
44    unsafe {
45        for _ in 0..6 {
46            // LOG_MAX_PSEUDO = 6
47            let mid = (lo + hi + 1) >> 1; // round up, matches C
48            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}