use alloc::vec::Vec;
use super::literal_cost::{LiteralCostArena, estimate_bit_costs_for_literals};
use super::nodes::INFINITY;
use crate::shared::command::Command;
use crate::shared::constants::{NUM_COMMAND_SYMBOLS, NUM_LITERAL_SYMBOLS};
use crate::shared::distance::NUM_HISTOGRAM_DISTANCE_SYMBOLS;
use crate::shared::fast_log::fast_log2;
const COMMAND_PRIOR_OFFSET: usize = 11;
const DISTANCE_PRIOR_OFFSET: usize = 20;
fn set_cost(histogram: &[u32], literal_histogram: bool, cost: &mut [f32]) {
let sum: usize = histogram.iter().map(|&count| count as usize).sum();
let log2sum = fast_log2(sum) as f32;
let mut missing_symbol_sum = sum;
if !literal_histogram {
missing_symbol_sum += histogram.iter().filter(|&&count| count == 0).count();
}
let missing_symbol_cost = fast_log2(missing_symbol_sum) as f32 + 2.0;
for (slot, &count) in cost.iter_mut().zip(histogram) {
if count == 0 {
*slot = missing_symbol_cost;
continue;
}
let bits = log2sum - fast_log2(count as usize) as f32;
*slot = if bits < 1.0 { 1.0 } else { bits };
}
}
pub(crate) struct ZopfliCostModel {
cost_cmd: Vec<f32>,
cost_dist: Vec<f32>,
literal_costs: Vec<f32>,
min_cost_cmd: f32,
distance_histogram_size: usize,
histogram_literal: Vec<u32>,
histogram_cmd: Vec<u32>,
histogram_dist: Vec<u32>,
cost_literal: Vec<f32>,
literal_arena: LiteralCostArena,
}
impl ZopfliCostModel {
pub(crate) fn new(alphabet_size: usize) -> Self {
Self {
cost_cmd: vec![0f32; NUM_COMMAND_SYMBOLS],
cost_dist: vec![0f32; alphabet_size.max(1)],
literal_costs: Vec::new(),
min_cost_cmd: 0.0,
distance_histogram_size: alphabet_size,
histogram_literal: vec![0u32; NUM_LITERAL_SYMBOLS],
histogram_cmd: vec![0u32; NUM_COMMAND_SYMBOLS],
histogram_dist: vec![0u32; NUM_HISTOGRAM_DISTANCE_SYMBOLS],
cost_literal: vec![0f32; NUM_LITERAL_SYMBOLS],
literal_arena: LiteralCostArena::default(),
}
}
pub(crate) fn retained_bytes(&self) -> usize {
(self.cost_cmd.capacity()
+ self.cost_dist.capacity()
+ self.literal_costs.capacity()
+ self.cost_literal.capacity())
* size_of::<f32>()
+ (self.histogram_literal.capacity()
+ self.histogram_cmd.capacity()
+ self.histogram_dist.capacity())
* size_of::<u32>()
+ self.literal_arena.retained_bytes()
}
pub(crate) fn reserve(&mut self, num_bytes: usize, alphabet_size: usize) {
if self.literal_costs.len() < num_bytes + 2 {
self.literal_costs.resize(num_bytes + 2, 0.0);
}
if self.cost_dist.len() < alphabet_size {
self.cost_dist.resize(alphabet_size, 0.0);
}
self.distance_histogram_size = alphabet_size;
}
pub(crate) fn set_from_literal_costs(
&mut self,
position: usize,
ringbuffer: &[u8],
mask: usize,
num_bytes: usize,
) {
estimate_bit_costs_for_literals(
position,
num_bytes,
mask,
ringbuffer,
&mut self.literal_arena,
&mut self.literal_costs[1..],
);
accumulate_literal_costs(&mut self.literal_costs, num_bytes);
for (symbol, slot) in self.cost_cmd.iter_mut().enumerate() {
*slot = fast_log2(COMMAND_PRIOR_OFFSET + symbol) as f32;
}
for (symbol, slot) in self
.cost_dist
.iter_mut()
.take(self.distance_histogram_size)
.enumerate()
{
*slot = fast_log2(DISTANCE_PRIOR_OFFSET + symbol) as f32;
}
self.min_cost_cmd = fast_log2(COMMAND_PRIOR_OFFSET) as f32;
}
pub(crate) fn set_from_commands(
&mut self,
position: usize,
ringbuffer: &[u8],
mask: usize,
commands: &[Command],
last_insert_len: usize,
num_bytes: usize,
) {
self.histogram_literal.fill(0);
self.histogram_cmd.fill(0);
self.histogram_dist.fill(0);
let mut pos = position - last_insert_len;
for command in commands {
let inslength = command.insert_len as usize;
let copylength = command.copy_len() as usize;
let distcode = usize::from(command.distance_code());
let cmdcode = usize::from(command.cmd_prefix);
self.histogram_cmd[cmdcode] += 1;
if cmdcode >= 128
&& let Some(slot) = self.histogram_dist.get_mut(distcode)
{
*slot += 1;
}
for offset in 0..inslength {
let literal = ringbuffer.get((pos + offset) & mask).copied().unwrap_or(0);
self.histogram_literal[usize::from(literal)] += 1;
}
pos += inslength + copylength;
}
set_cost(&self.histogram_literal, true, &mut self.cost_literal);
set_cost(&self.histogram_cmd, false, &mut self.cost_cmd);
set_cost(
&self.histogram_dist[..self.distance_histogram_size],
false,
&mut self.cost_dist[..self.distance_histogram_size],
);
self.min_cost_cmd = self.cost_cmd.iter().copied().fold(INFINITY, f32::min);
for index in 0..num_bytes {
let literal = ringbuffer
.get((position + index) & mask)
.copied()
.unwrap_or(0);
self.literal_costs[index + 1] = self.cost_literal[usize::from(literal)];
}
accumulate_literal_costs(&mut self.literal_costs, num_bytes);
}
#[inline(always)]
pub(crate) fn command_cost(&self, cmdcode: u16) -> f32 {
self.cost_cmd
.get(usize::from(cmdcode))
.copied()
.unwrap_or(INFINITY)
}
#[inline(always)]
pub(crate) fn distance_cost(&self, distcode: usize) -> f32 {
self.cost_dist.get(distcode).copied().unwrap_or(INFINITY)
}
#[inline(always)]
pub(crate) fn literal_costs(&self, from: usize, to: usize) -> f32 {
let at = |index: usize| self.literal_costs.get(index).copied().unwrap_or(0.0);
at(to) - at(from)
}
#[inline(always)]
pub(crate) const fn min_cost_cmd(&self) -> f32 {
self.min_cost_cmd
}
}
fn accumulate_literal_costs(costs: &mut [f32], num_bytes: usize) {
let mut literal_carry = 0f32;
costs[0] = 0.0;
for index in 0..num_bytes {
literal_carry += costs[index + 1];
costs[index + 1] = costs[index] + literal_carry;
literal_carry -= costs[index + 1] - costs[index];
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::shared::distance::DistanceParams;
fn alphabet() -> usize {
DistanceParams::default().alphabet_size_limit as usize
}
#[test]
fn a_missing_symbol_costs_more_than_a_rare_one() {
let mut histogram = vec![0u32; 8];
histogram[0] = 1000;
histogram[1] = 1;
let mut cost = vec![0f32; 8];
set_cost(&histogram, false, &mut cost);
assert!(cost[0] < cost[1], "{:?}", cost);
assert!(cost[1] < cost[2], "a missing symbol was not the dearest");
}
#[test]
fn a_literal_histogram_does_not_inflate_its_missing_symbols() {
let mut histogram = vec![0u32; 64];
histogram[0] = 100;
let mut literal = vec![0f32; 64];
let mut other = vec![0f32; 64];
set_cost(&histogram, true, &mut literal);
set_cost(&histogram, false, &mut other);
assert!(literal[1] < other[1]);
}
#[test]
fn no_symbol_costs_less_than_one_bit() {
let mut histogram = vec![0u32; 4];
histogram[0] = 1_000_000;
histogram[1] = 1;
let mut cost = vec![0f32; 4];
set_cost(&histogram, true, &mut cost);
assert_eq!(cost[0], 1.0);
}
#[test]
fn the_literal_prior_prices_every_command_symbol() {
let mut model = ZopfliCostModel::new(alphabet());
model.reserve(4096, alphabet());
let data = vec![b'a'; 4096];
model.set_from_literal_costs(0, &data, usize::MAX, 4096);
assert_eq!(model.command_cost(0), fast_log2(11) as f32);
assert_eq!(model.command_cost(5), fast_log2(16) as f32);
assert_eq!(model.distance_cost(0), fast_log2(20) as f32);
assert_eq!(model.min_cost_cmd(), fast_log2(11) as f32);
}
#[test]
fn cumulative_literal_costs_are_monotone_and_additive() {
let mut model = ZopfliCostModel::new(alphabet());
model.reserve(4096, alphabet());
let data: Vec<u8> = (0..4096u32).map(|i| (i * 31 % 256) as u8).collect();
model.set_from_literal_costs(0, &data, usize::MAX, 4096);
assert_eq!(model.literal_costs(0, 0), 0.0);
let whole = model.literal_costs(0, 4096);
let first = model.literal_costs(0, 2000);
let second = model.literal_costs(2000, 4096);
assert!(whole > 0.0);
assert!((whole - (first + second)).abs() <= whole.abs() * f32::EPSILON * 4.0);
for split in [1usize, 100, 1000, 4095] {
assert!(model.literal_costs(0, split) <= model.literal_costs(0, split + 1));
}
}
#[test]
fn the_carry_beats_a_plain_running_sum() {
let count = 200_000usize;
let addend = 3.7f32;
let mut with_carry = vec![0f32; count + 2];
for slot in with_carry[1..=count].iter_mut() {
*slot = addend;
}
accumulate_literal_costs(&mut with_carry, count);
let mut naive = 0f32;
for _ in 0..count {
naive += addend;
}
let exact = f64::from(addend) * count as f64;
let carried_error = (f64::from(with_carry[count]) - exact).abs();
let naive_error = (f64::from(naive) - exact).abs();
assert!(
naive_error > carried_error * 100.0,
"carry {} (error {carried_error}), naive {naive} (error {naive_error})",
with_carry[count]
);
}
#[test]
fn a_command_model_prices_the_symbols_it_saw() {
let dist = DistanceParams::default();
let data = vec![b'z'; 4096];
let commands: Vec<Command> = (0..20)
.map(|_| Command::new(&dist, 4, 20, 0, 100))
.collect();
let mut model = ZopfliCostModel::new(alphabet());
model.reserve(4096, alphabet());
model.set_from_commands(200, &data, usize::MAX, &commands, 0, 2048);
let used = commands[0].cmd_prefix;
let unused = (0..NUM_COMMAND_SYMBOLS as u16)
.find(|&symbol| symbol != used)
.expect("another symbol exists");
assert!(model.command_cost(used) < model.command_cost(unused));
assert!(model.min_cost_cmd() <= model.command_cost(used));
}
#[test]
fn an_empty_command_list_still_prices_every_symbol() {
let data = vec![b'q'; 1024];
let mut model = ZopfliCostModel::new(alphabet());
model.reserve(1024, alphabet());
model.set_from_commands(0, &data, usize::MAX, &[], 0, 1024);
for symbol in [0u16, 1, 100, (NUM_COMMAND_SYMBOLS - 1) as u16] {
assert!(model.command_cost(symbol).is_finite());
}
assert!(model.literal_costs(0, 1024).is_finite());
}
#[test]
fn rebuilding_the_model_does_not_carry_state_over() {
let dist = DistanceParams::default();
let data = vec![b'k'; 4096];
let commands: Vec<Command> = (0..8).map(|_| Command::new(&dist, 2, 10, 0, 60)).collect();
let mut fresh = ZopfliCostModel::new(alphabet());
fresh.reserve(4096, alphabet());
fresh.set_from_commands(100, &data, usize::MAX, &commands, 0, 1024);
let mut reused = ZopfliCostModel::new(alphabet());
reused.reserve(4096, alphabet());
reused.set_from_literal_costs(0, &data, usize::MAX, 4096);
reused.set_from_commands(100, &data, usize::MAX, &commands, 0, 1024);
for symbol in 0..NUM_COMMAND_SYMBOLS as u16 {
assert_eq!(fresh.command_cost(symbol), reused.command_cost(symbol));
}
assert_eq!(fresh.literal_costs(0, 1024), reused.literal_costs(0, 1024));
}
#[test]
fn reserving_grows_the_model_without_changing_its_prices() {
let data: Vec<u8> = (0..8192u32).map(|i| (i % 251) as u8).collect();
let mut small = ZopfliCostModel::new(alphabet());
small.reserve(1024, alphabet());
small.reserve(8192, alphabet());
small.set_from_literal_costs(0, &data, usize::MAX, 8192);
let mut large = ZopfliCostModel::new(alphabet());
large.reserve(8192, alphabet());
large.set_from_literal_costs(0, &data, usize::MAX, 8192);
assert_eq!(small.literal_costs(0, 8192), large.literal_costs(0, 8192));
}
}