#![allow(dead_code)]
const PROB_BITS: u32 = 15;
const PROB_BITS_0: u32 = 10;
const PROB_BITS_1: u32 = 14;
const MASK_0: u16 = ((!(!0u32 << PROB_BITS_0)) << (PROB_BITS - PROB_BITS_0)) as u16; const MASK_1: u16 = ((!(!0u32 << PROB_BITS_1)) << (PROB_BITS - PROB_BITS_1)) as u16; const DWS: u8 = 8;
#[derive(Clone, Copy, Debug)]
pub(crate) struct CtxModel {
state: [u16; 2],
rate: u8,
}
impl CtxModel {
pub(crate) fn init(init_value: u8, qp: u8, log2_window_size: u8) -> Self {
let init_id = init_value as i32;
let qp = (qp as i32).clamp(0, 63);
let slope = (init_id >> 3) - 4;
let offset = ((init_id & 7) * 18) + 1;
let inistate = ((slope * (qp - 16)) >> 1) + offset;
let state_clip = inistate.clamp(1, 127);
let p1 = (state_clip << 8) as u16;
let mut c = CtxModel {
state: [p1 & MASK_0, p1 & MASK_1],
rate: DWS,
};
c.set_log2_window_size(log2_window_size);
c
}
pub(crate) fn set_log2_window_size(&mut self, log2_window_size: u8) {
let lws = log2_window_size as u32;
let rate0 = 2 + ((lws >> 2) & 3);
let rate1 = 3 + rate0 + (lws & 3);
self.rate = (16 * rate0 + rate1) as u8;
}
#[inline]
fn prob8(self) -> u16 {
(self.state[0].wrapping_add(self.state[1])) >> 8
}
#[inline]
pub(crate) fn mps(self) -> u8 {
(self.prob8() >> 7) as u8
}
#[inline]
pub(crate) fn get_lps(self, range: u32) -> u32 {
let mut q = self.prob8();
if q & 0x80 != 0 {
q ^= 0xff; }
(((q as u32 >> 2) * (range >> 5)) >> 1) + 4
}
#[inline]
pub(crate) fn update(&mut self, bin: u8) {
let rate0 = (self.rate >> 4) as u32;
let rate1 = (self.rate & 15) as u32;
let mut s0 = self.state[0] as u32;
let mut s1 = self.state[1] as u32;
s0 -= (s0 >> rate0) & MASK_0 as u32;
s1 -= (s1 >> rate1) & MASK_1 as u32;
if bin != 0 {
s0 += (0x7fffu32 >> rate0) & MASK_0 as u32;
s1 += (0x7fffu32 >> rate1) & MASK_1 as u32;
}
self.state[0] = s0 as u16;
self.state[1] = s1 as u16;
}
#[cfg(test)]
pub(crate) fn from_raw(state: [u16; 2], rate: u8) -> Self {
CtxModel { state, rate }
}
}
#[derive(Clone)]
pub(crate) struct CabacEncoder {
low: u32,
m_range: u32,
bits_outstanding: u32,
first_bit: bool,
bit_buffer: u8,
bit_count: u8,
pub(crate) output: Vec<u8>,
}
impl CabacEncoder {
pub(crate) fn new() -> Self {
CabacEncoder {
low: 0,
m_range: 510,
bits_outstanding: 0,
first_bit: true,
bit_buffer: 0,
bit_count: 0,
output: Vec::new(),
}
}
pub(crate) fn range(&self) -> u32 {
self.m_range
}
#[inline]
fn emit_bit(&mut self, b: u32) {
self.bit_buffer = (self.bit_buffer << 1) | (b as u8 & 1);
self.bit_count += 1;
if self.bit_count == 8 {
self.output.push(self.bit_buffer);
self.bit_buffer = 0;
self.bit_count = 0;
}
}
#[inline]
fn put_bit(&mut self, b: u32) {
if self.first_bit {
self.first_bit = false;
} else {
self.emit_bit(b);
}
while self.bits_outstanding > 0 {
self.emit_bit(1 - b);
self.bits_outstanding -= 1;
}
}
#[inline]
fn renorm(&mut self) {
while self.m_range < 256 {
if self.low < 256 {
self.put_bit(0);
} else if self.low >= 512 {
self.low -= 512;
self.put_bit(1);
} else {
self.low -= 256;
self.bits_outstanding += 1;
}
self.m_range <<= 1;
self.low <<= 1;
}
}
#[inline]
pub(crate) fn encode_bin(&mut self, bin_val: u8, ctx: &mut CtxModel) {
let lps = ctx.get_lps(self.m_range);
self.m_range -= lps;
if (bin_val & 1) != ctx.mps() {
self.low += self.m_range;
self.m_range = lps;
}
ctx.update(bin_val & 1);
self.renorm();
}
#[inline]
pub(crate) fn encode_bypass(&mut self, bin_val: u8) {
self.low <<= 1;
if bin_val != 0 {
self.low += self.m_range;
}
if self.low >= 1024 {
self.put_bit(1);
self.low -= 1024;
} else if self.low < 512 {
self.put_bit(0);
} else {
self.low -= 512;
self.bits_outstanding += 1;
}
}
pub(crate) fn encode_bypass_bits(&mut self, value: u32, n: u32) {
let mut i = n;
while i > 0 {
i -= 1;
self.encode_bypass(((value >> i) & 1) as u8);
}
}
pub(crate) fn encode_terminate(&mut self, flag: u8) {
self.m_range -= 2;
if flag != 0 {
self.low += self.m_range;
self.flush();
} else {
self.renorm();
}
}
fn flush(&mut self) {
self.m_range = 2;
self.renorm();
self.put_bit((self.low >> 9) & 1);
let two = ((self.low >> 7) & 3) | 1;
self.emit_bit((two >> 1) & 1);
self.emit_bit(two & 1);
}
pub(crate) fn finish(mut self) -> Vec<u8> {
if self.bit_count > 0 {
self.bit_buffer <<= 8 - self.bit_count;
self.output.push(self.bit_buffer);
self.bit_buffer = 0;
self.bit_count = 0;
}
self.output
}
pub(crate) fn reset(&mut self) {
self.low = 0;
self.m_range = 510;
self.bits_outstanding = 0;
self.first_bit = true;
self.bit_buffer = 0;
self.bit_count = 0;
self.output.clear();
}
pub(crate) fn flushed_len(&self) -> usize {
self.output.len() + usize::from(self.bit_count > 0)
}
}
pub(crate) use decoder::CabacDecoder;
mod decoder {
use super::CtxModel;
pub(crate) struct CabacDecoder<'a> {
range: u32,
offset: u32,
data: &'a [u8],
bitpos: usize, }
impl<'a> CabacDecoder<'a> {
pub(crate) fn new(data: &'a [u8]) -> Self {
let mut d = CabacDecoder {
range: 510,
offset: 0,
data,
bitpos: 0,
};
d.offset = d.read_bits(9);
d
}
#[inline]
fn next_bit(&mut self) -> u32 {
let byte_idx = self.bitpos >> 3;
let bit = if byte_idx < self.data.len() {
let b = self.data[byte_idx];
((b >> (7 - (self.bitpos & 7))) & 1) as u32
} else {
1 };
self.bitpos += 1;
bit
}
#[inline]
fn read_bits(&mut self, n: u32) -> u32 {
let mut v = 0;
for _ in 0..n {
v = (v << 1) | self.next_bit();
}
v
}
#[inline]
fn renorm(&mut self) {
while self.range < 256 {
self.range <<= 1;
self.offset = (self.offset << 1) | self.next_bit();
}
}
#[inline]
pub(crate) fn decode_bin(&mut self, ctx: &mut CtxModel) -> u8 {
let mps = ctx.mps();
let lps = ctx.get_lps(self.range);
self.range -= lps;
let bin = if self.offset >= self.range {
self.offset -= self.range;
self.range = lps;
1 - mps
} else {
mps
};
ctx.update(bin);
self.renorm();
bin
}
#[inline]
pub(crate) fn decode_bypass(&mut self) -> u8 {
self.offset = (self.offset << 1) | self.next_bit();
if self.offset >= self.range {
self.offset -= self.range;
1
} else {
0
}
}
#[inline]
pub(crate) fn decode_terminate(&mut self) -> u8 {
self.range -= 2;
if self.offset >= self.range {
1
} else {
self.renorm();
0
}
}
}
}
#[cfg(test)]
mod tests {
use super::CabacDecoder;
use super::*;
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn bit(&mut self) -> u8 {
(self.next() & 1) as u8
}
}
#[test]
fn mask_constants_match_spec() {
assert_eq!(MASK_0, 0x7FE0);
assert_eq!(MASK_1, 0x7FFE);
}
#[test]
fn init_clamps_and_is_deterministic() {
for &iv in &[0u8, 1, 17, 95, 128, 200, 255] {
for &qp in &[0u8, 16, 26, 37, 51, 63] {
let c = CtxModel::init(iv, qp, 8);
let s = c.state[0] as u32 + c.state[1] as u32;
assert!(s > 0 && s < (1 << 16));
assert!(c.mps() == 0 || c.mps() == 1);
}
}
}
fn roundtrip_mixed(seed: u64, n: usize) {
let mut rng = Rng(seed);
let init_vals = [25u8, 60, 110, 154, 199];
let qp = 32u8;
let make_ctxs = || -> Vec<CtxModel> {
init_vals
.iter()
.map(|&v| CtxModel::init(v, qp, 8))
.collect()
};
enum Op {
Ctx(usize, u8),
Byp(u8),
}
let mut plan = Vec::with_capacity(n);
for _ in 0..n {
if rng.bit() == 0 {
let ci = (rng.next() as usize) % init_vals.len();
plan.push(Op::Ctx(ci, rng.bit()));
} else {
plan.push(Op::Byp(rng.bit()));
}
}
let mut enc = CabacEncoder::new();
let mut ectx = make_ctxs();
for op in &plan {
match op {
Op::Ctx(ci, b) => enc.encode_bin(*b, &mut ectx[*ci]),
Op::Byp(b) => enc.encode_bypass(*b),
}
}
enc.encode_terminate(1);
let bytes = enc.finish();
let mut dec = CabacDecoder::new(&bytes);
let mut dctx = make_ctxs();
for (i, op) in plan.iter().enumerate() {
let got = match op {
Op::Ctx(ci, _) => dec.decode_bin(&mut dctx[*ci]),
Op::Byp(_) => dec.decode_bypass(),
};
let want = match op {
Op::Ctx(_, b) => *b,
Op::Byp(b) => *b,
};
assert_eq!(got, want, "bin {i} mismatch (seed {seed})");
}
assert_eq!(dec.decode_terminate(), 1, "terminate (seed {seed})");
}
#[test]
fn roundtrip_context_and_bypass() {
for seed in 1..200u64 {
roundtrip_mixed(seed.wrapping_mul(0x9E3779B97F4A7C15), 64);
}
}
#[test]
fn roundtrip_long_sequences() {
for seed in 1..20u64 {
roundtrip_mixed(seed.wrapping_mul(0xD1B54A32D192ED03), 4000);
}
}
#[test]
fn roundtrip_all_context_bins() {
let mut rng = Rng(0xABCDEF);
let qp = 28u8;
for &iv in &[12u8, 77, 140, 222] {
let mut plan = Vec::new();
for _ in 0..2000 {
plan.push(rng.bit());
}
let mut enc = CabacEncoder::new();
let mut ec = CtxModel::init(iv, qp, 8);
for &b in &plan {
enc.encode_bin(b, &mut ec);
}
enc.encode_terminate(1);
let bytes = enc.finish();
let mut dec = CabacDecoder::new(&bytes);
let mut dc = CtxModel::init(iv, qp, 8);
for (i, &b) in plan.iter().enumerate() {
assert_eq!(dec.decode_bin(&mut dc), b, "iv {iv} bin {i}");
}
assert_eq!(dec.decode_terminate(), 1);
}
}
#[test]
fn probability_adapts_toward_input() {
let mut c = CtxModel::init(35, 32, 8);
let prob_before = c.state[0] as u32 + c.state[1] as u32;
let lps_before = c.get_lps(384);
for _ in 0..16 {
c.update(1);
}
let prob_after = c.state[0] as u32 + c.state[1] as u32;
let lps_after = c.get_lps(384);
assert!(
prob_after > prob_before,
"estimate must rise ({prob_after} > {prob_before})"
);
assert!(
lps_after <= lps_before,
"LPS must not grow ({lps_after} <= {lps_before})"
);
}
}