pub(crate) const CDF_TOTAL: u32 = 1 << 15;
const TOP: u32 = 1 << 24;
const BOT: u32 = 1 << 16;
#[derive(Clone, Debug)]
pub(crate) struct Cdf {
c: Vec<u16>,
}
impl Cdf {
pub(crate) fn nsyms(&self) -> usize {
self.c.len() - 1
}
pub(crate) fn uniform(n: usize) -> Self {
assert!(n >= 2);
let mut c = vec![0u16; n + 1];
for (i, dst) in c[..n].iter_mut().enumerate() {
*dst = (((i as u64 + 1) * CDF_TOTAL as u64) / n as u64) as u16;
}
c[n - 1] = CDF_TOTAL as u16;
c[n] = 0; Cdf { c }
}
#[allow(unused)]
pub(crate) fn from_cumulative(cum: &[u16]) -> Self {
let n = cum.len();
assert!(n >= 2);
assert_eq!(
cum[n - 1] as u32,
CDF_TOTAL,
"last cumulative must be 32768"
);
let mut c = vec![0u16; n + 1];
c[..n].copy_from_slice(cum);
c[n] = 0;
Cdf { c }
}
#[inline]
fn boundary(&self, k: usize) -> u32 {
if k == 0 { 0 } else { self.c[k - 1] as u32 }
}
#[inline]
fn eff_boundary(&self, k: usize) -> u32 {
self.boundary(k) + k as u32
}
#[inline]
fn eff_total(&self) -> u32 {
CDF_TOTAL + self.nsyms() as u32
}
fn update(&mut self, symbol: usize) {
let n = self.nsyms();
let count = self.c[n] as u32;
let rate = 3
+ (if count > 15 { 1 } else { 0 })
+ (if count > 31 { 1 } else { 0 })
+ floor_log2(n as u32).min(2);
for i in 0..n - 1 {
let cur = self.c[i] as i32;
let tmp = if i >= symbol { CDF_TOTAL as i32 } else { 0 };
let next = if tmp < cur {
cur - ((cur - tmp) >> rate)
} else {
cur + ((tmp - cur) >> rate)
};
self.c[i] = next as u16;
}
if count < 32 {
self.c[n] = (count + 1) as u16;
}
}
}
#[inline]
fn floor_log2(mut x: u32) -> u32 {
let mut r = 0;
while x > 1 {
x >>= 1;
r += 1;
}
r
}
pub(crate) struct RangeEncoder {
low: u32,
range: u32,
out: Vec<u8>,
}
impl Default for RangeEncoder {
fn default() -> Self {
Self::new()
}
}
impl RangeEncoder {
pub(crate) fn new() -> Self {
RangeEncoder {
low: 0,
range: 0xFFFF_FFFF,
out: Vec::new(),
}
}
#[inline]
fn encode_freq(&mut self, cum: u32, freq: u32, tot: u32) {
let r = self.range / tot;
self.low = self.low.wrapping_add(r.wrapping_mul(cum));
self.range = r.wrapping_mul(freq);
loop {
if (self.low ^ self.low.wrapping_add(self.range)) < TOP {
} else if self.range < BOT {
self.range = self.low.wrapping_neg() & (BOT - 1);
} else {
break;
}
self.out.push((self.low >> 24) as u8);
self.low <<= 8;
self.range <<= 8;
}
}
pub(crate) fn encode_symbol(&mut self, symbol: usize, cdf: &mut Cdf) {
debug_assert!(symbol < cdf.nsyms());
let lo = cdf.eff_boundary(symbol);
let hi = cdf.eff_boundary(symbol + 1);
self.encode_freq(lo, hi - lo, cdf.eff_total());
cdf.update(symbol);
}
pub(crate) fn encode_literal(&mut self, value: u32, nbits: u32) {
for i in (0..nbits).rev() {
let bit = (value >> i) & 1;
self.encode_freq(bit * (CDF_TOTAL / 2), CDF_TOTAL / 2, CDF_TOTAL);
}
}
#[allow(unused)]
pub(crate) fn finish(mut self) -> Vec<u8> {
for _ in 0..4 {
self.out.push((self.low >> 24) as u8);
self.low <<= 8;
}
self.out
}
}