use super::rangecoder::RangeDecoder;
use super::smpl_cc_tables::CcTables;
pub(crate) struct SmplGainResult {
pub(crate) gain_q: [i32; 4],
pub(crate) nrg_res: [i32; 4],
}
pub(crate) fn decode_smpl_gains(
dec: &mut RangeDecoder,
cc: &CcTables,
p3: i32,
subfr_counts: [i32; 4],
) -> SmplGainResult {
let mut res = SmplGainResult {
gain_q: [0; 4],
nrg_res: [0; 4],
};
let gain_main = dec.decode_cdf(cc.nrgres_gain4());
let gain_delta = dec.decode_cdf(cc.nrgres_shape4());
let cfg_sel = 2i32;
let off6 = p3 * gain_delta;
let base7 = gain_main * cc.nrg_step(cfg_sel) - 0x154000;
for sf in 0..(p3 as usize).min(4) {
let cbv = cc.gain_recon(p3 == 4, sf as i32 + off6);
res.gain_q[sf] = base7 + (cbv << 4);
}
for (sf, &cnt) in subfr_counts.iter().enumerate().take((p3 as usize).min(4)) {
if cnt <= 0 {
continue;
}
let bucket = if cnt >= 30 { 3 } else { (cnt & 0xffff) / 10 };
let mut g = (res.gain_q[sf] + 8192) >> 14;
if g < -85 {
g = -85;
}
let neg_part = (g >> 31) & g;
let min_offset = (-neg_part) as usize;
res.nrg_res[sf] =
dec.decode_cdf(cc.fcbg_offset(cfg_sel as usize, bucket as usize, min_offset));
}
log::trace!(
"mlow gains: main={gain_main} delta={gain_delta} gain_q={:?} nrg_res={:?}",
res.gain_q,
res.nrg_res
);
res
}
#[cfg(test)]
mod tests {
use super::super::smpl_cc_tables::load_cc_tables;
use super::super::smpl_decode::{SmplLsfState, decode_smpl_lsf, load_smpl_tables};
use super::super::smpl_pulse::decode_smpl_pulses;
use super::*;
use serde_json::Value;
#[test]
fn gains_match_go() {
let recs: Value = serde_json::from_str(include_str!("testdata/gains_vectors.json"))
.expect("gains_vectors");
let tbl = load_smpl_tables();
let cc = load_cc_tables();
let arr = recs.as_array().unwrap();
assert!(!arr.is_empty());
let as_i32 = |v: &Value| -> Vec<i32> {
v.as_array()
.unwrap()
.iter()
.map(|x| x.as_i64().unwrap() as i32)
.collect()
};
for rec in arr {
let frame = hex::decode(rec["frame"].as_str().unwrap()).unwrap();
let mut st = SmplLsfState::default();
let mut dec = RangeDecoder::new(&frame[1..]);
let lsf = decode_smpl_lsf(&mut dec, tbl, &mut st, 0, 0);
let pulses = decode_smpl_pulses(&mut dec, cc, 320, 4, 1, 0, lsf.stage1);
let g = decode_smpl_gains(&mut dec, cc, 4, pulses.subfr);
assert_eq!(g.gain_q.to_vec(), as_i32(&rec["gain_q"]), "gain_q");
assert_eq!(g.nrg_res.to_vec(), as_i32(&rec["nrg_res"]), "nrg_res");
}
}
}