use crate::celt_band_layout::{CeltFrameSize, CELT_NUM_BANDS};
use crate::celt_e_prob_model::{
e_prob_pair, energy_pred_coef, EProbModelError, EnergyPredictionMode,
};
use crate::celt_laplace::{decay_byte_to_q14, ec_laplace_decode, prob_to_fs};
use crate::range_decoder::RangeDecoder;
pub const E_MEANS_Q4: [i8; 25] = [
103, 100, 92, 85, 81, 77, 72, 70, 78, 75, 73, 71, 78, 74, 69, 72, 70, 74, 76, 71, 60, 60, 60,
60, 60,
];
pub const E_MEANS_Q4_SCALE: f64 = 1.0 / 16.0;
#[inline]
#[must_use]
pub fn e_mean(band: usize) -> Option<f64> {
if band >= CELT_NUM_BANDS {
return None;
}
Some(f64::from(E_MEANS_Q4[band]) * E_MEANS_Q4_SCALE)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CoarseEnergyError {
Model(EProbModelError),
}
impl core::fmt::Display for CoarseEnergyError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
CoarseEnergyError::Model(e) => write!(f, "coarse-energy model lookup failed: {e:?}"),
}
}
}
impl std::error::Error for CoarseEnergyError {}
impl From<EProbModelError> for CoarseEnergyError {
fn from(e: EProbModelError) -> Self {
CoarseEnergyError::Model(e)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct CoarseEnergyFrame {
pub reported_log2: Vec<f64>,
pub residuals: Vec<i32>,
}
#[derive(Debug, Clone)]
pub struct CoarseEnergyState {
history: [f64; CELT_NUM_BANDS],
}
impl Default for CoarseEnergyState {
fn default() -> Self {
Self::new()
}
}
impl CoarseEnergyState {
#[must_use]
pub fn new() -> Self {
Self {
history: [0.0; CELT_NUM_BANDS],
}
}
pub fn reset(&mut self) {
self.history = [0.0; CELT_NUM_BANDS];
}
#[must_use]
pub fn history(&self) -> &[f64; CELT_NUM_BANDS] {
&self.history
}
#[inline]
#[must_use]
fn clamp_history(value: f64) -> f64 {
value
}
pub fn decode_frame(
&mut self,
rd: &mut RangeDecoder<'_>,
frame_size: CeltFrameSize,
intra: bool,
start: usize,
end: usize,
) -> Result<CoarseEnergyFrame, CoarseEnergyError> {
let lm = frame_size.column_index() as u32;
let mode = EnergyPredictionMode::from_intra_flag(intra);
let coef = energy_pred_coef(lm, mode)?;
let alpha = coef.alpha();
let beta = coef.beta();
let end = end.min(CELT_NUM_BANDS);
let start = start.min(end);
let coded = end - start;
let mut reported_log2 = Vec::with_capacity(coded);
let mut residuals = Vec::with_capacity(coded);
let mut pred_freq = 0.0_f64;
#[allow(clippy::needless_range_loop)]
for band in start..end {
let pair = e_prob_pair(lm, mode, band as u32)?;
let fs = prob_to_fs(pair.prob);
let decay = decay_byte_to_q14(pair.decay);
let q = ec_laplace_decode(rd, fs, decay);
let r = f64::from(q);
let d = pred_freq + r;
let prev = self.history[band];
let recon = Self::clamp_history(alpha * prev + d);
let mean = f64::from(E_MEANS_Q4[band]) * E_MEANS_Q4_SCALE;
reported_log2.push(recon + mean);
residuals.push(q);
self.history[band] = recon;
pred_freq += (1.0 - beta) * r;
}
Ok(CoarseEnergyFrame {
reported_log2,
residuals,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn e_means_table_matches_csv() {
assert_eq!(E_MEANS_Q4[0], 103);
assert_eq!(E_MEANS_Q4[7], 70);
assert_eq!(E_MEANS_Q4[20], 60);
assert_eq!(E_MEANS_Q4.len(), 25);
assert!((e_mean(0).unwrap() - 103.0 / 16.0).abs() < 1e-12);
assert!(e_mean(CELT_NUM_BANDS).is_none());
}
#[test]
fn frequency_recurrence_matches_closed_form() {
let beta = 6554.0 / 32768.0; let r = 1.0_f64;
let mut pred_freq = 0.0_f64;
for b in 0..CELT_NUM_BANDS {
let d = pred_freq + r;
let recon = d; let expected = r * (1.0 + (b as f64) * (1.0 - beta));
assert!(
(recon - expected).abs() < 1e-9,
"band {b}: recon {recon} != closed form {expected}"
);
pred_freq += (1.0 - beta) * r;
}
}
#[test]
fn intra_ignores_history() {
let buf = [0x9a, 0x3c, 0x71, 0x55, 0xe2, 0x08, 0xbd, 0x40];
let mut fresh = CoarseEnergyState::new();
let f1 = {
let mut rd = RangeDecoder::new(&buf);
fresh
.decode_frame(&mut rd, CeltFrameSize::Ms20, true, 0, CELT_NUM_BANDS)
.unwrap()
};
let mut primed = CoarseEnergyState::new();
for (i, h) in primed.history.iter_mut().enumerate() {
*h = (i as f64) * 0.37 - 3.0;
}
let f2 = {
let mut rd = RangeDecoder::new(&buf);
primed
.decode_frame(&mut rd, CeltFrameSize::Ms20, true, 0, CELT_NUM_BANDS)
.unwrap()
};
assert_eq!(f1.residuals, f2.residuals);
for (a, b) in f1.reported_log2.iter().zip(f2.reported_log2.iter()) {
assert!((a - b).abs() < 1e-12, "intra energy depends on history");
}
}
#[test]
fn inter_adds_alpha_times_history() {
let buf = [0x12, 0x9f, 0x44, 0xc8, 0x6b, 0x31, 0xaa, 0x05];
let mut zeroed = CoarseEnergyState::new();
let fz = {
let mut rd = RangeDecoder::new(&buf);
zeroed
.decode_frame(&mut rd, CeltFrameSize::Ms10, false, 0, CELT_NUM_BANDS)
.unwrap()
};
let mut primed = CoarseEnergyState::new();
let hist = 2.0_f64;
for h in primed.history.iter_mut() {
*h = hist;
}
let fp = {
let mut rd = RangeDecoder::new(&buf);
primed
.decode_frame(&mut rd, CeltFrameSize::Ms10, false, 0, CELT_NUM_BANDS)
.unwrap()
};
assert_eq!(fz.residuals, fp.residuals);
let alpha = energy_pred_coef(2, EnergyPredictionMode::Inter)
.unwrap()
.alpha();
for (z, p) in fz.reported_log2.iter().zip(fp.reported_log2.iter()) {
assert!(
(p - z - alpha * hist).abs() < 1e-9,
"delta {} != alpha*hist {}",
p - z,
alpha * hist
);
}
}
#[test]
fn hybrid_band_range_decodes_high_bands_only() {
let buf = [0x55, 0xaa, 0x33, 0xcc, 0x0f, 0xf0, 0x5a, 0xa5];
let mut st = CoarseEnergyState::new();
for h in st.history.iter_mut() {
*h = -1.0;
}
let frame = st
.decode_frame(
&mut RangeDecoder::new(&buf),
CeltFrameSize::Ms20,
false,
17,
21,
)
.unwrap();
assert_eq!(frame.reported_log2.len(), 4);
assert_eq!(frame.residuals.len(), 4);
for h in st.history.iter().take(17) {
assert_eq!(*h, -1.0);
}
}
#[test]
fn decode_is_clean_for_all_frame_sizes() {
for fs in [
CeltFrameSize::Ms2_5,
CeltFrameSize::Ms5,
CeltFrameSize::Ms10,
CeltFrameSize::Ms20,
] {
for intra in [false, true] {
let buf = [0x3c, 0x91, 0x7e, 0x22, 0xb5, 0x48, 0xd0, 0x6f, 0x13, 0xee];
let mut rd = RangeDecoder::new(&buf);
let mut st = CoarseEnergyState::new();
let frame = st
.decode_frame(&mut rd, fs, intra, 0, CELT_NUM_BANDS)
.unwrap();
assert_eq!(frame.reported_log2.len(), CELT_NUM_BANDS);
assert!(!rd.has_error(), "fs {fs:?} intra {intra} latched error");
}
}
}
#[test]
fn two_frame_history_threading() {
let buf1 = [0x40, 0x80, 0x10, 0x9f, 0x33, 0x71, 0xc4, 0x05];
let buf2 = [0x88, 0x21, 0x6e, 0xb3, 0x4a, 0xf0, 0x19, 0x5c];
let mut st = CoarseEnergyState::new();
let _f1 = st
.decode_frame(
&mut RangeDecoder::new(&buf1),
CeltFrameSize::Ms20,
false,
0,
CELT_NUM_BANDS,
)
.unwrap();
let hist_after_f1 = *st.history();
let f2 = st
.decode_frame(
&mut RangeDecoder::new(&buf2),
CeltFrameSize::Ms20,
false,
0,
CELT_NUM_BANDS,
)
.unwrap();
let coef = energy_pred_coef(3, EnergyPredictionMode::Inter).unwrap();
let (alpha, beta) = (coef.alpha(), coef.beta());
let mut pred_freq = 0.0;
for (band, &q) in f2.residuals.iter().enumerate() {
let r = f64::from(q);
let recon = alpha * hist_after_f1[band] + pred_freq + r;
let mean = f64::from(E_MEANS_Q4[band]) * E_MEANS_Q4_SCALE;
assert!(
(f2.reported_log2[band] - (recon + mean)).abs() < 1e-9,
"band {band} mismatch"
);
pred_freq += (1.0 - beta) * r;
}
}
}