use crate::ics_info::{IcsInfo, PredictorData, WindowSequence, PRED_SFB_MAX};
use crate::swb_offset::long_window_offsets;
use crate::Error;
type Result<T> = core::result::Result<T, Error>;
pub const ALPHA: f32 = 0.90625;
pub const A: f32 = 0.953125;
pub const B: f32 = 0.953125;
pub const NUM_RESET_GROUPS: usize = 30;
pub fn flt_round_inf(pf: f32) -> f32 {
let bits = pf.to_bits();
let flg = bits & 0x0000_8000;
let truncated = bits & 0xffff_0000;
let mut result = f32::from_bits(truncated);
if flg != 0 {
let exp_sign = truncated & 0xff80_0000;
let one_lsb = exp_sign | 0x0001_0000;
result += f32::from_bits(one_lsb);
result -= f32::from_bits(exp_sign);
}
result
}
#[inline]
pub fn flt_trunc(pf: f32) -> f32 {
f32::from_bits(pf.to_bits() & 0xffff_0000)
}
fn flt_round_even(pf: f32) -> f32 {
if pf == 0.0 {
return 0.0;
}
let bits = pf.to_bits();
let biased = ((bits >> 23) & 0xff) as i32;
let exp = biased - 126;
let scale = 2f32.powi(8 - exp);
let tmp = pf * scale;
let mut a = tmp as i64;
if (tmp - a as f32) >= 0.5 {
a += 1;
}
if (tmp - a as f32) == 0.5 {
a &= -2;
}
a as f32 / scale
}
fn mnt_table(i: usize) -> f32 {
let f = f32::from_bits(0x3f80_0000 + ((i as u32) << 16));
flt_round_even(B / f)
}
fn exp_table(i: usize) -> f32 {
let f = f32::from_bits((i as u32) << 23);
if f > 1.0 {
1.0 / f
} else {
0.0
}
}
#[inline]
fn b_over_var(var: f32) -> f32 {
let bits = var.to_bits();
let mant7 = ((bits >> 16) & 0x7f) as usize;
let exp = ((bits >> 23) & 0xff) as usize;
MNT_TABLE[mant7] * EXP_TABLE[exp]
}
static MNT_TABLE: std::sync::LazyLock<[f32; 128]> =
std::sync::LazyLock::new(|| core::array::from_fn(mnt_table));
static EXP_TABLE: std::sync::LazyLock<[f32; 256]> =
std::sync::LazyLock::new(|| core::array::from_fn(exp_table));
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Predictor {
r0: f32,
r1: f32,
cor1: f32,
cor2: f32,
var1: f32,
var2: f32,
}
impl Default for Predictor {
fn default() -> Self {
Self::new()
}
}
impl Predictor {
pub fn new() -> Self {
Self {
r0: 0.0,
r1: 0.0,
cor1: 0.0,
cor2: 0.0,
var1: 1.0,
var2: 1.0,
}
}
pub fn reset(&mut self) {
*self = Self::new();
}
fn coefficients(&self) -> (f32, f32) {
(
self.cor1 * b_over_var(self.var1),
self.cor2 * b_over_var(self.var2),
)
}
pub fn predict(&self) -> f32 {
let (bk1, bk2) = self.coefficients();
let x_est1 = bk1 * self.r0;
let x_est2 = bk2 * self.r1;
flt_round_inf(x_est1 + x_est2)
}
pub fn update(&mut self, x_rec: f32) {
let bk1 = self.cor1 * b_over_var(self.var1);
let e0 = x_rec;
let r0_prev = self.r0;
let r1_prev = self.r1;
let x_est1 = bk1 * r0_prev;
let e1 = e0 - x_est1;
let cor1 = ALPHA * self.cor1 + r0_prev * e0;
let var1 = ALPHA * self.var1 + 0.5 * (r0_prev * r0_prev + e0 * e0);
let cor2 = ALPHA * self.cor2 + r1_prev * e1;
let var2 = ALPHA * self.var2 + 0.5 * (r1_prev * r1_prev + e1 * e1);
let r1_new = A * (r0_prev - bk1 * e0);
let r0_new = A * x_rec;
self.r0 = flt_trunc(r0_new);
self.r1 = flt_trunc(r1_new);
self.cor1 = flt_trunc(cor1);
self.cor2 = flt_trunc(cor2);
self.var1 = flt_trunc(var1);
self.var2 = flt_trunc(var2);
}
}
#[derive(Clone, Debug)]
pub struct PredictorBank {
predictors: Vec<Predictor>,
}
impl PredictorBank {
pub fn new(fs_index: u8) -> Result<Self> {
let offsets = long_window_offsets(fs_index)?;
let pred_sfb_max = PRED_SFB_MAX[fs_index as usize] as usize;
let num_predictors = offsets
.get(pred_sfb_max)
.copied()
.ok_or(Error::PredictorInvalid)? as usize;
Ok(Self {
predictors: vec![Predictor::new(); num_predictors],
})
}
pub fn len(&self) -> usize {
self.predictors.len()
}
pub fn is_empty(&self) -> bool {
self.predictors.is_empty()
}
pub fn reset_all(&mut self) {
for p in &mut self.predictors {
p.reset();
}
}
pub fn reset_group(&mut self, group: u8) -> Result<()> {
if group == 0 || group as usize > NUM_RESET_GROUPS {
return Err(Error::PredictorInvalid);
}
let start = (group - 1) as usize;
let mut idx = start;
while idx < self.predictors.len() {
self.predictors[idx].reset();
idx += NUM_RESET_GROUPS;
}
Ok(())
}
pub fn apply_long(
&mut self,
spec: &mut [f64],
ics_info: &IcsInfo,
pred: Option<&PredictorData>,
fs_index: u8,
) -> Result<bool> {
if ics_info.window_sequence == WindowSequence::EightShort {
self.reset_all();
return Ok(false);
}
if spec.len() < self.predictors.len() {
return Err(Error::PredictorInvalid);
}
let offsets = ics_info.swb_offsets(fs_index)?;
let pred_sfb_max = PRED_SFB_MAX[fs_index as usize] as usize;
let max_sfb = ics_info.max_sfb as usize;
let mut modified = false;
let num_predictors = self.predictors.len();
for sfb in 0..pred_sfb_max {
let fc = offsets[sfb] as usize;
let lc = (offsets[sfb + 1] as usize).min(num_predictors);
if fc >= lc {
continue;
}
let active = sfb < max_sfb
&& pred.is_some_and(|p| p.prediction_used.get(sfb).copied().unwrap_or(false));
for (p, y) in self.predictors[fc..lc]
.iter_mut()
.zip(spec[fc..lc].iter_mut())
{
let x_est = flt_round_inf(p.predict());
let y_rec = *y as f32;
let x_rec = if active {
modified = true;
x_est + y_rec
} else {
y_rec
};
*y = x_rec as f64;
p.update(x_rec);
}
}
if let Some(p) = pred {
if p.reset {
if let Some(group) = p.reset_group_number {
self.reset_group(group)?;
}
}
}
Ok(modified)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ics_info::{WindowSequence, WindowShape};
fn long_ics(max_sfb: u8) -> IcsInfo {
IcsInfo {
family: crate::swb_offset::FrameFamily::Lc1024,
ics_reserved_bit: false,
window_sequence: WindowSequence::OnlyLong,
window_shape: WindowShape::Sine,
max_sfb,
scale_factor_grouping: None,
predictor_data_present: false,
predictor_data: None,
ltp_data_present: false,
ltp_data: None,
ltp_data_present_pair: None,
ltp_data_pair: None,
num_windows: 1,
num_window_groups: 1,
window_group_length: vec![1],
num_swb: 49,
}
}
#[test]
fn flt_round_inf_clears_low_16_bits_when_no_rounding() {
let v = 1.5_f32; assert_eq!(flt_round_inf(v).to_bits() & 0x0000_ffff, 0);
assert_eq!(flt_round_inf(v), v);
}
#[test]
fn flt_round_inf_result_always_has_zero_low_bits() {
for &v in &[0.0_f32, 1.0, -1.0, 3.5_f32, -2.6_f32, 1e-8, 1e8, 0.953125] {
let r = flt_round_inf(v);
assert_eq!(
r.to_bits() & 0x0000_ffff,
0,
"flt_round_inf({v}) left low mantissa bits set"
);
}
}
#[test]
fn flt_round_inf_rounds_toward_infinity() {
let bits = 1.0_f32.to_bits() | 0x0000_8001;
let v = f32::from_bits(bits);
let truncated = f32::from_bits(bits & 0xffff_0000);
let r = flt_round_inf(v);
assert!(r > truncated, "expected round-up: {r} vs trunc {truncated}");
assert_eq!(r.to_bits() & 0x0000_ffff, 0);
}
#[test]
fn fresh_predictor_predicts_zero() {
let p = Predictor::new();
assert_eq!(p.predict(), 0.0);
}
#[test]
fn predictor_initial_state_matches_spec() {
let p = Predictor::new();
assert_eq!(p.r0, 0.0);
assert_eq!(p.r1, 0.0);
assert_eq!(p.cor1, 0.0);
assert_eq!(p.cor2, 0.0);
assert_eq!(p.var1, 1.0);
assert_eq!(p.var2, 1.0);
}
#[test]
fn update_then_reset_returns_to_initial() {
let mut p = Predictor::new();
for _ in 0..16 {
p.update(0.7);
}
assert_ne!(p, Predictor::new());
p.reset();
assert_eq!(p, Predictor::new());
}
#[test]
fn update_advances_lattice_register() {
let mut p = Predictor::new();
let x = 2.0_f32;
p.update(x);
assert_eq!(p.r0, flt_round_inf(A * x));
}
#[test]
fn bank_size_covers_pred_sfb_max() {
let bank = PredictorBank::new(4).unwrap();
assert_eq!(bank.len(), 672);
assert!(!bank.is_empty());
}
#[test]
fn bank_size_24khz() {
let bank = PredictorBank::new(6).unwrap();
assert_eq!(bank.len(), 652);
}
#[test]
fn reset_group_rejects_reserved_numbers() {
let mut bank = PredictorBank::new(4).unwrap();
assert!(matches!(bank.reset_group(0), Err(Error::PredictorInvalid)));
assert!(matches!(bank.reset_group(31), Err(Error::PredictorInvalid)));
assert!(bank.reset_group(1).is_ok());
assert!(bank.reset_group(30).is_ok());
}
#[test]
fn reset_group_only_touches_its_members() {
let mut bank = PredictorBank::new(4).unwrap();
for p in &mut bank.predictors {
p.update(0.5);
}
let before: Vec<Predictor> = bank.predictors.clone();
bank.reset_group(1).unwrap();
for (i, p) in bank.predictors.iter().enumerate() {
if i % NUM_RESET_GROUPS == 0 {
assert_eq!(*p, Predictor::new(), "line {i} should be reset");
} else {
assert_eq!(*p, before[i], "line {i} should be untouched");
}
}
}
#[test]
fn short_block_resets_and_leaves_spectrum_untouched() {
let mut bank = PredictorBank::new(4).unwrap();
for p in &mut bank.predictors {
p.update(0.3);
}
let mut ics = long_ics(40);
ics.window_sequence = WindowSequence::EightShort;
let mut spec = vec![1.0_f64; 1024];
let original = spec.clone();
let modified = bank.apply_long(&mut spec, &ics, None, 4).unwrap();
assert!(!modified);
assert_eq!(spec, original);
for p in &bank.predictors {
assert_eq!(*p, Predictor::new());
}
}
#[test]
fn prediction_off_leaves_spectrum_but_advances_state() {
let mut bank = PredictorBank::new(4).unwrap();
let ics = long_ics(40);
let mut spec = vec![2.0_f64; 1024];
let original = spec.clone();
let modified = bank.apply_long(&mut spec, &ics, None, 4).unwrap();
assert!(!modified);
assert_eq!(spec, original, "prediction-off must not alter the spectrum");
assert_ne!(bank.predictors[0], Predictor::new());
}
#[test]
fn active_band_modifies_spectrum_on_second_frame() {
let mut bank = PredictorBank::new(4).unwrap();
let mut ics = long_ics(40);
ics.predictor_data_present = true;
let pred = PredictorData {
reset: false,
reset_group_number: None,
prediction_used: {
let mut v = vec![false; 40];
v[0] = true;
v
},
};
for _ in 0..6 {
let mut spec = vec![0.0_f64; 1024];
for (c, s) in spec.iter_mut().enumerate().take(8) {
*s = (c as f64) + 1.0;
}
bank.apply_long(&mut spec, &ics, Some(&pred), 4).unwrap();
}
let mut spec2 = vec![1.0_f64; 1024];
let y_rec = spec2.clone();
let modified = bank.apply_long(&mut spec2, &ics, Some(&pred), 4).unwrap();
assert!(modified);
let band0_changed = (0..4).any(|c| spec2[c] != y_rec[c]);
assert!(band0_changed, "active band 0 spectrum did not change");
}
#[test]
fn reset_after_processing_clears_signalled_group() {
let mut bank = PredictorBank::new(4).unwrap();
let mut ics = long_ics(40);
ics.predictor_data_present = true;
let pred = PredictorData {
reset: true,
reset_group_number: Some(1),
prediction_used: vec![true; 40],
};
let mut spec = vec![3.0_f64; 1024];
bank.apply_long(&mut spec, &ics, Some(&pred), 4).unwrap();
assert_eq!(bank.predictors[0], Predictor::new());
assert_eq!(bank.predictors[NUM_RESET_GROUPS], Predictor::new());
assert_ne!(bank.predictors[1], Predictor::new());
}
#[test]
fn spec_shorter_than_bank_is_rejected() {
let mut bank = PredictorBank::new(4).unwrap();
let ics = long_ics(40);
let mut spec = vec![0.0_f64; 100];
assert!(matches!(
bank.apply_long(&mut spec, &ics, None, 4),
Err(Error::PredictorInvalid)
));
}
#[test]
fn bad_fs_index_propagates_error() {
assert!(PredictorBank::new(13).is_err());
}
}