Skip to main content

rusty_opus/
quant_bands.rs

1use crate::modes::CeltMode;
2use crate::range_coder::{BITRES, RangeCoder};
3
4pub const PRED_COEF: [f32; 4] = [
5    29440.0 / 32768.0,
6    26112.0 / 32768.0,
7    21248.0 / 32768.0,
8    16384.0 / 32768.0,
9];
10pub const BETA_COEF: [f32; 4] = [
11    30147.0 / 32768.0,
12    22282.0 / 32768.0,
13    12124.0 / 32768.0,
14    6554.0 / 32768.0,
15];
16pub const BETA_INTRA: f32 = 4915.0 / 32768.0;
17
18pub const E_PROB_MODEL: [[[u8; 42]; 2]; 4] = [
19    [
20        [
21            72, 127, 65, 129, 66, 128, 65, 128, 64, 128, 62, 128, 64, 128, 64, 128, 92, 78, 92, 79,
22            92, 78, 90, 79, 116, 41, 115, 40, 114, 40, 132, 26, 132, 26, 145, 17, 161, 12, 176, 10,
23            177, 11,
24        ],
25        [
26            24, 179, 48, 138, 54, 135, 54, 132, 53, 134, 56, 133, 55, 132, 55, 132, 61, 114, 70,
27            96, 74, 88, 75, 88, 87, 74, 89, 66, 91, 67, 100, 59, 108, 50, 120, 40, 122, 37, 97, 43,
28            78, 50,
29        ],
30    ],
31    [
32        [
33            83, 78, 84, 81, 88, 75, 86, 74, 87, 71, 90, 73, 93, 74, 93, 74, 109, 40, 114, 36, 117,
34            34, 117, 34, 143, 17, 145, 18, 146, 19, 162, 12, 165, 10, 178, 7, 189, 6, 190, 8, 177,
35            9,
36        ],
37        [
38            23, 178, 54, 115, 63, 102, 66, 98, 69, 99, 74, 89, 71, 91, 73, 91, 78, 89, 86, 80, 92,
39            66, 93, 64, 102, 59, 103, 60, 104, 60, 117, 52, 123, 44, 138, 35, 133, 31, 97, 38, 77,
40            45,
41        ],
42    ],
43    [
44        [
45            61, 90, 93, 60, 105, 42, 107, 41, 110, 45, 116, 38, 113, 38, 112, 38, 124, 26, 132, 27,
46            136, 19, 140, 20, 155, 14, 159, 16, 158, 18, 170, 13, 177, 10, 187, 8, 192, 6, 175, 9,
47            159, 10,
48        ],
49        [
50            21, 178, 59, 110, 71, 86, 75, 85, 84, 83, 91, 66, 88, 73, 87, 72, 92, 75, 98, 72, 105,
51            58, 107, 54, 115, 52, 114, 55, 112, 56, 129, 51, 132, 40, 150, 33, 140, 29, 98, 35, 77,
52            42,
53        ],
54    ],
55    [
56        [
57            42, 121, 96, 66, 108, 43, 111, 40, 117, 44, 123, 32, 120, 36, 119, 33, 127, 33, 134,
58            34, 139, 21, 147, 23, 152, 20, 158, 25, 154, 26, 166, 21, 173, 16, 184, 13, 184, 10,
59            150, 13, 139, 15,
60        ],
61        [
62            22, 178, 63, 114, 74, 82, 84, 83, 92, 82, 103, 62, 96, 72, 96, 67, 101, 73, 107, 72,
63            113, 55, 118, 52, 125, 52, 118, 52, 117, 55, 135, 49, 137, 39, 157, 32, 145, 29, 97,
64            33, 77, 40,
65        ],
66    ],
67];
68
69pub const SMALL_ENERGY_ICDF: [u8; 3] = [2, 1, 0];
70
71fn loss_distortion(
72    e_bands: &[f32],
73    old_e_bands: &[f32],
74    start: usize,
75    end: usize,
76    len: usize,
77    channels: usize,
78) -> f32 {
79    let mut dist = 0.0f32;
80    for c in 0..channels {
81        let off = c * len;
82        for i in start..end.min(len) {
83            let d = e_bands[off + i] - old_e_bands[off + i];
84            dist += d * d;
85        }
86    }
87    dist.min(200.0)
88}
89
90#[allow(clippy::too_many_arguments)]
91fn quant_coarse_energy_impl(
92    m: &CeltMode,
93    start: usize,
94    end: usize,
95    e_bands: &[f32],
96    old_e_bands: &mut [f32],
97    budget: u32,
98    tell_start: i32,
99    prob_model: &[u8; 42],
100    error: &mut [f32],
101    enc: &mut RangeCoder,
102    channels: usize,
103    lm: usize,
104    intra: bool,
105    max_decay: f32,
106    lfe: bool,
107) -> i32 {
108    let coef = if intra { 0.0 } else { PRED_COEF[lm] };
109    let beta = if intra { BETA_INTRA } else { BETA_COEF[lm] };
110    let mut prev = [0.0f32; 2];
111    let mut badness = 0i32;
112
113    if tell_start + 3 <= budget as i32 {
114        enc.encode_bit_logp(intra, 3);
115    }
116
117    for i in start..end {
118        for c in 0..channels {
119            let x = e_bands[c * m.nb_ebands + i];
120            let old_e_val = old_e_bands[c * m.nb_ebands + i];
121            let old_e = old_e_val.max(-9.0);
122            let f = x - coef * old_e - prev[c];
123
124            let mut qi = ((f + 0.5).floor() as i32).clamp(-32767, 32767);
125            let qi0 = qi;
126
127            let decay_bound = old_e_val.max(-28.0) - max_decay;
128            if qi < 0 && x < decay_bound {
129                qi = qi.saturating_add(((decay_bound - x) as i32).max(0));
130                if qi > 0 {
131                    qi = 0;
132                }
133            }
134
135            let tell = enc.tell();
136            let bits_left = budget as i32 - tell - 3 * channels as i32 * (end - i) as i32;
137            if i != start && bits_left < 30 {
138                if bits_left < 24 {
139                    qi = qi.min(1);
140                }
141                if bits_left < 16 {
142                    qi = qi.max(-1);
143                }
144            }
145            if lfe && i >= 2 {
146                qi = qi.min(0);
147            }
148
149            if tell + 15 <= budget as i32 {
150                let prob_idx = 2 * i.min(20);
151                let fs = (prob_model[prob_idx] as u32) << 7;
152                let decay = (prob_model[prob_idx + 1] as i32) << 6;
153                enc.laplace_encode(&mut qi, fs, decay);
154            } else if tell + 2 <= budget as i32 {
155                qi = qi.clamp(-1, 1);
156                enc.encode_icdf(
157                    (2 * qi) ^ (if qi < 0 { -1 } else { 0 }),
158                    &SMALL_ENERGY_ICDF,
159                    2,
160                );
161            } else if tell < budget as i32 {
162                qi = qi.min(0);
163                enc.encode_bit_logp(qi != 0, 1);
164            } else {
165                qi = -1;
166            }
167
168            badness = badness.saturating_add(qi0.saturating_sub(qi).saturating_abs());
169
170            let q = qi as f32;
171            error[c * m.nb_ebands + i] = f - q;
172            let tmp = coef * old_e + prev[c] + q;
173            old_e_bands[c * m.nb_ebands + i] = tmp;
174            prev[c] = prev[c] + q - beta * q;
175        }
176    }
177
178    if lfe { 0 } else { badness }
179}
180
181#[allow(clippy::too_many_arguments)]
182pub fn quant_coarse_energy_advanced(
183    m: &CeltMode,
184    start: usize,
185    end: usize,
186    eff_end: usize,
187    e_bands: &[f32],
188    old_e_bands: &mut [f32],
189    budget: u32,
190    error: &mut [f32],
191    enc: &mut RangeCoder,
192    channels: usize,
193    lm: usize,
194    nb_available_bytes: usize,
195    force_intra: bool,
196    delayed_intra: &mut f32,
197    mut two_pass: bool,
198    loss_rate: i32,
199    lfe: bool,
200) {
201    let _prof = crate::prof::scope(crate::prof::Stage::CeltCoarse);
202    let mut intra = force_intra
203        || (!two_pass
204            && *delayed_intra > 2.0 * channels as f32 * (end.saturating_sub(start)) as f32
205            && nb_available_bytes > (end.saturating_sub(start)) * channels);
206
207    let intra_bias = ((budget as f32) * (*delayed_intra) * (loss_rate as f32)
208        / ((channels as f32) * 512.0)) as i32;
209    let new_distortion =
210        loss_distortion(e_bands, old_e_bands, start, eff_end, m.nb_ebands, channels);
211
212    let tell = enc.tell();
213    if tell + 3 > budget as i32 {
214        two_pass = false;
215        intra = false;
216    }
217
218    let mut max_decay = if end - start > 10 {
219        16.0f32.min(0.125 * nb_available_bytes as f32)
220    } else {
221        16.0f32
222    };
223    if lfe {
224        max_decay = 3.0;
225    }
226
227    let enc_start_state = enc.clone();
228    let mut old_e_bands_intra = old_e_bands.to_vec();
229    let mut error_intra = error.to_vec();
230    let mut badness1 = 0i32;
231    let mut tell_intra = 0i32;
232    let intra_prob = &E_PROB_MODEL[lm][1];
233
234    if two_pass || intra {
235        badness1 = quant_coarse_energy_impl(
236            m,
237            start,
238            end,
239            e_bands,
240            &mut old_e_bands_intra,
241            budget,
242            tell,
243            intra_prob,
244            &mut error_intra,
245            enc,
246            channels,
247            lm,
248            true,
249            max_decay,
250            lfe,
251        );
252        tell_intra = crate::tell_frac_inline!(enc);
253    }
254
255    if !intra {
256        let enc_intra_state = enc.clone();
257
258        *enc = enc_start_state.clone();
259        let inter_prob = &E_PROB_MODEL[lm][0];
260        let badness2 = quant_coarse_energy_impl(
261            m,
262            start,
263            end,
264            e_bands,
265            old_e_bands,
266            budget,
267            tell,
268            inter_prob,
269            error,
270            enc,
271            channels,
272            lm,
273            false,
274            max_decay,
275            lfe,
276        );
277
278        if two_pass
279            && (badness1 < badness2
280                || (badness1 == badness2
281                    && crate::tell_frac_inline!(enc) + intra_bias > tell_intra))
282        {
283            *enc = enc_intra_state;
284            old_e_bands.copy_from_slice(&old_e_bands_intra);
285            error.copy_from_slice(&error_intra);
286            intra = true;
287        }
288    } else {
289        old_e_bands.copy_from_slice(&old_e_bands_intra);
290        error.copy_from_slice(&error_intra);
291    }
292
293    if intra {
294        *delayed_intra = new_distortion;
295    } else {
296        let pred2 = PRED_COEF[lm] * PRED_COEF[lm];
297        *delayed_intra = pred2 * *delayed_intra + new_distortion;
298    }
299}
300
301#[allow(clippy::too_many_arguments)]
302pub fn quant_coarse_energy(
303    m: &CeltMode,
304    start: usize,
305    end: usize,
306    e_bands: &[f32],
307    old_e_bands: &mut [f32],
308    budget: u32,
309    error: &mut [f32],
310    enc: &mut RangeCoder,
311    channels: usize,
312    lm: usize,
313    force_intra: bool,
314    nb_available_bytes: usize,
315) {
316    let mut delayed_intra = 0.0f32;
317    quant_coarse_energy_advanced(
318        m,
319        start,
320        end,
321        end,
322        e_bands,
323        old_e_bands,
324        budget,
325        error,
326        enc,
327        channels,
328        lm,
329        nb_available_bytes,
330        force_intra,
331        &mut delayed_intra,
332        false,
333        0,
334        false,
335    );
336}
337
338#[allow(clippy::too_many_arguments)]
339pub fn unquant_coarse_energy(
340    m: &CeltMode,
341    start: usize,
342    end: usize,
343    old_e_bands: &mut [f32],
344    intra: bool,
345    dec: &mut RangeCoder,
346    channels: usize,
347    lm: usize,
348) {
349    let prob_model = &E_PROB_MODEL[lm][if intra { 1 } else { 0 }];
350    let coef = if intra { 0.0 } else { PRED_COEF[lm] };
351    let beta = if intra { BETA_INTRA } else { BETA_COEF[lm] };
352    debug_assert!(channels <= 2);
353    let mut prev = [0.0f32; 2];
354    let budget = (dec.storage * 8) as i32;
355
356    for i in start..end {
357        for c in 0..channels {
358            let qi;
359            let tell = dec.tell();
360            if budget - tell >= 15 {
361                let prob_idx = 2 * i.min(20);
362                let fs = (prob_model[prob_idx] as u32) << 7;
363                let decay = (prob_model[prob_idx + 1] as i32) << 6;
364                qi = dec.laplace_decode(fs, decay);
365            } else if budget - tell >= 2 {
366                let s = dec.decode_icdf(&SMALL_ENERGY_ICDF, 2);
367                qi = (s >> 1) ^ -(s & 1);
368            } else if budget - tell >= 1 {
369                qi = if dec.decode_bit_logp(1) { -1 } else { 0 };
370            } else {
371                qi = -1;
372            }
373
374            // Clamp in-place, matching C: oldEBands[i] = MAXG(-GCONST(9.f), oldEBands[i])
375            old_e_bands[c * m.nb_ebands + i] = old_e_bands[c * m.nb_ebands + i].max(-9.0);
376            let old_e = old_e_bands[c * m.nb_ebands + i];
377
378            let q = qi as f32;
379            let tmp = coef * old_e + prev[c] + q;
380            old_e_bands[c * m.nb_ebands + i] = tmp;
381            prev[c] = prev[c] + q - beta * q;
382        }
383    }
384}
385
386#[allow(clippy::too_many_arguments)]
387pub fn quant_fine_energy(
388    m: &CeltMode,
389    start: usize,
390    end: usize,
391    old_e_bands: &mut [f32],
392    error: &mut [f32],
393    fine_quant: &[i32],
394    enc: &mut RangeCoder,
395    channels: usize,
396) {
397    let _prof = crate::prof::scope(crate::prof::Stage::CeltFine);
398    for i in start..end {
399        for c in 0..channels {
400            let bits = fine_quant[i];
401            if bits <= 0 {
402                continue;
403            }
404            let mut q = ((error[c * m.nb_ebands + i] + 0.5) * (1 << bits) as f32).floor() as i32;
405            q = q.max(0).min((1 << bits) - 1);
406            enc.enc_bits(q as u32, bits as u32);
407            let offset = (q as f32 + 0.5) / (1 << bits) as f32 - 0.5;
408            old_e_bands[c * m.nb_ebands + i] += offset;
409            error[c * m.nb_ebands + i] -= offset;
410        }
411    }
412}
413
414pub fn unquant_fine_energy(
415    m: &CeltMode,
416    start: usize,
417    end: usize,
418    old_e_bands: &mut [f32],
419    fine_quant: &[i32],
420    dec: &mut RangeCoder,
421    channels: usize,
422) {
423    for i in start..end {
424        for c in 0..channels {
425            let bits = fine_quant[i];
426            if bits <= 0 {
427                continue;
428            }
429            let q = dec.dec_bits(bits as u32);
430            let offset = (q as f32 + 0.5) / (1 << bits) as f32 - 0.5;
431            old_e_bands[c * m.nb_ebands + i] += offset;
432        }
433    }
434}
435
436#[allow(clippy::too_many_arguments)]
437pub fn quant_energy_finalise(
438    m: &CeltMode,
439    start: usize,
440    end: usize,
441    old_e_bands: &mut [f32],
442    error: &mut [f32],
443    fine_quant: &[i32],
444    fine_priority: &[i32],
445    bits_left: i32,
446    enc: &mut RangeCoder,
447    channels: usize,
448) {
449    let mut bits_left = bits_left;
450    for priority in 0..2 {
451        let mut i = start;
452        while i < end && bits_left >= channels as i32 {
453            if fine_quant[i] >= 8 || fine_priority[i] != priority {
454                i += 1;
455                continue;
456            }
457            let mut c = 0;
458            while c < channels {
459                let q2 = if error[i + c * m.nb_ebands] < 0.0 {
460                    0
461                } else {
462                    1
463                };
464                enc.enc_bits(q2 as u32, 1);
465                let offset =
466                    (q2 as f32 - 0.5) * (1i32 << (14 - fine_quant[i] - 1)) as f32 * (1.0 / 16384.0);
467                old_e_bands[i + c * m.nb_ebands] += offset;
468                error[i + c * m.nb_ebands] -= offset;
469                bits_left -= 1;
470                c += 1;
471            }
472            i += 1;
473        }
474    }
475}
476
477#[allow(clippy::too_many_arguments)]
478pub fn unquant_energy_finalise(
479    m: &CeltMode,
480    start: usize,
481    end: usize,
482    old_e_bands: &mut [f32],
483    fine_quant: &[i32],
484    fine_priority: &[i32],
485    bits_left: i32,
486    dec: &mut RangeCoder,
487    channels: usize,
488) {
489    let mut bits_left = bits_left;
490    for priority in 0..2 {
491        let mut i = start;
492        while i < end && bits_left >= channels as i32 {
493            if fine_quant[i] >= 8 || fine_priority[i] != priority {
494                i += 1;
495                continue;
496            }
497            let mut c = 0;
498            while c < channels {
499                let q2 = dec.dec_bits(1);
500                let offset =
501                    (q2 as f32 - 0.5) * (1i32 << (14 - fine_quant[i] - 1)) as f32 * (1.0 / 16384.0);
502                old_e_bands[i + c * m.nb_ebands] += offset;
503                bits_left -= 1;
504                c += 1;
505            }
506            i += 1;
507        }
508    }
509}
510
511#[cfg(test)]
512mod tests {
513    use super::*;
514    use crate::range_coder::RangeCoder;
515
516    #[test]
517    fn test_coarse_fine_energy() {
518        let mode = crate::modes::default_mode();
519        let mut e_bands = vec![0.0; mode.nb_ebands];
520        for (i, v) in e_bands.iter_mut().enumerate() {
521            *v = 5.0 + (i as f32 * 0.5).sin() * 2.0;
522        }
523
524        let mut old_e_bands = vec![0.0; mode.nb_ebands];
525        let mut error = vec![0.0; mode.nb_ebands];
526        let mut enc = RangeCoder::new_encoder(1000);
527
528        quant_coarse_energy(
529            mode,
530            0,
531            mode.nb_ebands,
532            &e_bands,
533            &mut old_e_bands,
534            10000,
535            &mut error,
536            &mut enc,
537            1,
538            3,
539            false,
540            80,
541        );
542
543        let mut fine_quant = vec![0; mode.nb_ebands];
544        for (i, v) in fine_quant.iter_mut().enumerate() {
545            *v = (i % 3) as i32;
546        }
547
548        quant_fine_energy(
549            mode,
550            0,
551            mode.nb_ebands,
552            &mut old_e_bands,
553            &mut error,
554            &fine_quant,
555            &mut enc,
556            1,
557        );
558
559        let mut fine_priority = vec![0i32; mode.nb_ebands];
560        for (i, v) in fine_priority.iter_mut().enumerate() {
561            *v = (i % 2) as i32;
562        }
563
564        quant_energy_finalise(
565            mode,
566            0,
567            mode.nb_ebands,
568            &mut old_e_bands,
569            &mut error,
570            &fine_quant,
571            &fine_priority,
572            10,
573            &mut enc,
574            1,
575        );
576
577        enc.done();
578        let _compressed = &enc.buf;
579
580        let mut dec = RangeCoder::new_decoder(&enc.buf);
581
582        let mut decoded_old_e_bands = vec![0.0; mode.nb_ebands];
583        let intra = dec.decode_bit_logp(3);
584        unquant_coarse_energy(
585            mode,
586            0,
587            mode.nb_ebands,
588            &mut decoded_old_e_bands,
589            intra,
590            &mut dec,
591            1,
592            3,
593        );
594
595        unquant_fine_energy(
596            mode,
597            0,
598            mode.nb_ebands,
599            &mut decoded_old_e_bands,
600            &fine_quant,
601            &mut dec,
602            1,
603        );
604
605        unquant_energy_finalise(
606            mode,
607            0,
608            mode.nb_ebands,
609            &mut decoded_old_e_bands,
610            &fine_quant,
611            &fine_priority,
612            10,
613            &mut dec,
614            1,
615        );
616
617        for i in 0..mode.nb_ebands {
618            if (decoded_old_e_bands[i] - old_e_bands[i]).abs() >= 1e-5 {
619                println!(
620                    "Mismatch at band {}: enc={} dec={} diff={}",
621                    i,
622                    old_e_bands[i],
623                    decoded_old_e_bands[i],
624                    (decoded_old_e_bands[i] - old_e_bands[i]).abs()
625                );
626            }
627            assert!((decoded_old_e_bands[i] - old_e_bands[i]).abs() < 1e-5);
628        }
629    }
630
631    /// Regression test: extreme/corrupt energy values must not cause an
632    /// "attempt to add with overflow" panic in `quant_coarse_energy_impl`.
633    /// Previously `(qi0 - qi).abs()` overflowed i32 when the float->int cast
634    /// of `qi` saturated near i32::MAX/MIN.
635    #[test]
636    fn test_coarse_energy_extreme_no_overflow() {
637        let mode = crate::modes::default_mode();
638        let n = mode.nb_ebands;
639
640        for &extreme in &[f32::INFINITY, f32::NEG_INFINITY, f32::NAN, 1.0e30, -1.0e30] {
641            let e_bands = vec![extreme; n];
642            let mut old_e_bands = vec![0.0; n];
643            let mut error = vec![0.0; n];
644            let mut enc = RangeCoder::new_encoder(1000);
645
646            // tiny budget forces the `else { qi = -1 }` path, which combined
647            // with a saturated qi0 triggered the overflow at the badness line.
648            quant_coarse_energy(
649                mode,
650                0,
651                n,
652                &e_bands,
653                &mut old_e_bands,
654                0,
655                &mut error,
656                &mut enc,
657                1,
658                3,
659                false,
660                80,
661            );
662
663            // large budget exercises the laplace-encode path with extreme qi.
664            let mut old_e_bands2 = vec![0.0; n];
665            let mut error2 = vec![0.0; n];
666            let mut enc2 = RangeCoder::new_encoder(1000);
667            quant_coarse_energy(
668                mode,
669                0,
670                n,
671                &e_bands,
672                &mut old_e_bands2,
673                10000,
674                &mut error2,
675                &mut enc2,
676                1,
677                3,
678                false,
679                80,
680            );
681        }
682    }
683}