use crate::encode::entropy::rans;
use crate::encode::entropy::rans::RansSymbolEncoder;
use draco_oxide_core::bit_coder::BitWriter;
use draco_oxide_core::bit_coder::ByteWriter;
use draco_oxide_core::codec::entropy::SymbolEncodingMethod;
use draco_oxide_core::types::{NdVector, Vector};
#[derive(thiserror::Error, Debug, Clone, Copy, PartialEq, Eq)]
pub enum Err {
#[error("RANS encoding error")]
RansEncodingError(#[from] rans::Err),
#[error("Invalid inputs for encode_tagged_symbol(): It must be true that symbol.len()==num_values*num_components, but got symbol.len()={0}, num_values={1}, num_components={2}")]
InvalidInputs(usize, usize, usize),
#[error("Invalid bit length: {0}")]
InvalidBitLength(usize),
}
pub fn encode_symbols<W>(
symbols: Vec<u64>,
num_components: usize,
config: SymbolEncodingMethod,
writer: &mut W,
) -> Result<(), Err>
where
W: ByteWriter,
{
config.write_to(writer);
match config {
SymbolEncodingMethod::LengthCoded => {
let mut bit_lengths = Vec::with_capacity(symbols.len() / num_components);
for i in 0..symbols.len() / num_components {
let mut max_bit_length = 0;
for j in 0..num_components {
let s = symbols[i * num_components + j];
let bit_length = (64 - s.leading_zeros()) as usize;
if bit_length > max_bit_length {
max_bit_length = bit_length;
}
}
bit_lengths.push(max_bit_length as u8);
}
encode_symbols_length_coded(symbols, num_components, bit_lengths, writer)
}
SymbolEncodingMethod::DirectCoded => encode_symbols_direct_coded(symbols, writer),
}
}
fn encode_symbols_length_coded<W>(
symbols: Vec<u64>,
num_components: usize,
bit_lengths: Vec<u8>,
writer: &mut W,
) -> Result<(), Err>
where
W: ByteWriter,
{
let mut freq_counts = Vec::new();
for &bit_length in &bit_lengths {
let bit_length = bit_length as usize;
if freq_counts.len() <= bit_length {
freq_counts.resize(bit_length + 1, 0);
}
freq_counts[bit_length] += 1;
}
let mut encoder = RansSymbolEncoder::new(writer, freq_counts, None, 12)?;
for i in (0..symbols.len() / num_components).rev() {
encoder.write(bit_lengths[i] as usize)?;
}
encoder.flush()?;
let mut writer: BitWriter<_> = BitWriter::spown_from(writer);
for i in 0..symbols.len() / num_components {
let value_bit_length = bit_lengths[i];
for c in 0..num_components {
writer.write_bits((value_bit_length, symbols[i * num_components + c]));
}
}
Ok(())
}
fn encode_symbols_direct_coded<W>(symbols: Vec<u64>, writer: &mut W) -> Result<(), Err>
where
W: ByteWriter,
{
encode_direct_coded_streams(
symbols.iter().map(|&s| s as usize),
symbols.iter().rev().map(|&s| s as usize),
writer,
)
}
pub fn encode_vector_symbols<W, const N: usize>(
values: &[NdVector<N, i32>],
writer: &mut W,
) -> Result<(), Err>
where
W: ByteWriter,
NdVector<N, i32>: Vector<N, Component = i32>,
{
SymbolEncodingMethod::DirectCoded.write_to(writer);
encode_direct_coded_streams(
values
.iter()
.flat_map(|v| (0..N).map(move |i| *v.get(i) as usize)),
values
.iter()
.rev()
.flat_map(|v| (0..N).rev().map(move |i| *v.get(i) as usize)),
writer,
)
}
fn encode_direct_coded_streams<W>(
forward: impl Iterator<Item = usize>,
reversed: impl Iterator<Item = usize>,
writer: &mut W,
) -> Result<(), Err>
where
W: ByteWriter,
{
let mut freq_counts: Vec<usize> = Vec::new();
let mut max_symbol = 0;
for s in forward {
if s >= max_symbol {
max_symbol = s;
freq_counts.resize(max_symbol + 1, 0);
}
freq_counts[s] += 1;
}
let num_unique_symbols = freq_counts.iter().filter(|&&c| c > 0).count();
let bit_length = (usize::BITS - num_unique_symbols.leading_zeros()) as usize;
let bit_length = bit_length.clamp(1, 18);
writer.write_u8(bit_length as u8);
let precision = match bit_length {
1..=8 => 12,
9 => 13,
10 => 15,
11 => 16,
12 => 18,
13 => 19,
_ => 20,
};
let mut encoder = RansSymbolEncoder::new(writer, freq_counts, None, precision)?;
for s in reversed {
encoder.write(s)?;
}
encoder.flush()?;
Ok(())
}