pub(crate) mod walks;
use super::bit_context::BitContext;
use super::Encode;
use walks::half;
#[cfg(test)]
use expect_test::expect;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub struct AtMost<const MAX: usize>(usize);
impl<const MAX: usize> AtMost<MAX> {
#[inline]
pub const fn new(value: usize) -> Self {
if value <= MAX {
AtMost(value)
} else {
panic!("Invalid value in compactly::AtMost")
}
}
}
impl<const MAX: usize> From<AtMost<MAX>> for usize {
#[inline]
fn from(value: AtMost<MAX>) -> Self {
value.0
}
}
impl<const MAX: usize> TryFrom<usize> for AtMost<MAX> {
type Error = ();
#[inline]
fn try_from(value: usize) -> Result<Self, Self::Error> {
if value <= MAX {
Ok(AtMost(value))
} else {
Err(())
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AtMostContext<const MAX: usize> {
pub(crate) bits: [<bool as Encode>::Context; MAX],
}
impl<const MAX: usize> AtMostContext<MAX> {
const SEEDED: [BitContext; MAX] = {
let mut bits = [BitContext::True0False0; MAX];
let mut stack = [(0usize, 0usize); 192];
stack[0] = (0, MAX + 1);
let mut top = 1;
while top > 0 {
top -= 1;
let (start, len) = stack[top];
if len > 1 {
let vc = half(len);
let split = start + vc;
bits[split - 1] = seed_context(vc as u64, (len - vc) as u64);
stack[top] = (start, vc);
stack[top + 1] = (split, len - vc);
top += 2;
}
}
bits
};
}
const fn seed_context(lo: u64, hi: u64) -> BitContext {
let mut best = BitContext::True0False0;
let mut best_err = seed_err(best, lo, hi);
let mut path = 0u32;
while path < 1 << 4 {
let mut state = BitContext::True0False0;
let mut k = 0;
while k < 4 {
state = state.adapt((path >> k) & 1 == 1);
let err = seed_err(state, lo, hi);
if err < best_err {
best_err = err;
best = state;
}
k += 1;
}
path += 1;
}
best
}
const fn seed_err(state: BitContext, lo: u64, hi: u64) -> u64 {
let p = state.probability().prob.get() as u64;
(p * (lo + hi)).abs_diff(256 * lo)
}
impl<const MAX: usize> Default for AtMostContext<MAX> {
#[inline]
fn default() -> Self {
Self { bits: Self::SEEDED }
}
}
impl<const MAX: usize> Encode for AtMost<MAX> {
type Context = AtMostContext<MAX>;
#[inline]
fn encode<E: super::EntropyCoder>(&self, writer: &mut E, ctx: &mut Self::Context) {
writer.encode_atmost(ctx, *self)
}
#[inline]
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<Self, std::io::Error> {
Ok(reader.decode_atmost(ctx))
}
}
#[test]
fn size() {
use super::estimated_bits;
fn test_urange<const MAX: usize>() {
for i in 0..=MAX {
let v = AtMost::<MAX>::new(i);
println!("Testing AtMost::<{MAX}>::new({i})");
let encoded = super::encode(&v);
let decoded = super::decode::<AtMost<MAX>>(&encoded).unwrap();
assert_eq!(decoded, v);
}
}
test_urange::<0>();
test_urange::<1>();
test_urange::<2>();
test_urange::<3>();
test_urange::<4>();
test_urange::<5>();
test_urange::<6>();
test_urange::<7>();
test_urange::<8>();
test_urange::<9>();
test_urange::<254>();
test_urange::<255>();
test_urange::<256>();
fn exact_bits<const MAX: usize>(bits: usize) {
for i in 0..=MAX {
let v = AtMost::<MAX>::new(i);
assert_eq!(
super::Encode::millibits(&v),
super::Millibits::bits(bits),
"AtMost::<{MAX}>::new({i}) should cost exactly {bits} bits"
);
}
}
exact_bits::<1>(1);
exact_bits::<3>(2);
exact_bits::<7>(3);
exact_bits::<15>(4);
exact_bits::<31>(5);
exact_bits::<63>(6);
exact_bits::<127>(7);
exact_bits::<255>(8);
expect!["2"].assert_eq(&estimated_bits!(AtMost::<2>::try_from(0).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<2>::try_from(1).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<2>::try_from(2).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<4>::try_from(0).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<4>::try_from(1).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<4>::try_from(2).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<4>::try_from(3).unwrap()));
expect!["2"].assert_eq(&estimated_bits!(AtMost::<4>::try_from(4).unwrap()));
expect!["3"].assert_eq(&estimated_bits!(AtMost::<5>::try_from(0).unwrap()));
expect!["3"].assert_eq(&estimated_bits!(AtMost::<5>::try_from(1).unwrap()));
expect!["3"].assert_eq(&estimated_bits!(AtMost::<5>::try_from(2).unwrap()));
expect!["3"].assert_eq(&estimated_bits!(AtMost::<5>::try_from(3).unwrap()));
expect!["3"].assert_eq(&estimated_bits!(AtMost::<5>::try_from(4).unwrap()));
expect!["3"].assert_eq(&estimated_bits!(AtMost::<5>::try_from(5).unwrap()));
expect!["7"].assert_eq(&estimated_bits!(AtMost::<127>::try_from(0).unwrap()));
expect!["7"].assert_eq(&estimated_bits!(AtMost::<127>::try_from(1).unwrap()));
expect!["7"].assert_eq(&estimated_bits!(AtMost::<127>::try_from(127).unwrap()));
expect!["8"].assert_eq(&estimated_bits!(AtMost::<255>::try_from(0).unwrap()));
expect!["8"].assert_eq(&estimated_bits!(AtMost::<255>::try_from(1).unwrap()));
expect!["8"].assert_eq(&estimated_bits!(AtMost::<255>::try_from(255).unwrap()));
}
#[test]
fn context_is_const_and_allocation_free() {
const _CTX: AtMostContext<2> = AtMostContext {
bits: [super::bit_context::BitContext::True0False0; 2],
};
assert_eq!(AtMostContext::<2>::default().bits.len(), 2);
assert_eq!(AtMostContext::<255>::default().bits.len(), 255);
assert_eq!(std::mem::size_of::<AtMostContext<0>>(), 0);
}