use core::{
mem::{MaybeUninit, size_of_val},
slice::{from_raw_parts, from_raw_parts_mut},
};
use crate::{
bitpack::packed_byte_size,
constants::MAX_EXCEPTIONS,
delta::{delta_range, eval_delta_benefit},
encoder::{
delta::encode_delta,
exception::{Exception, exceptions_byte_size},
kernel::encode_slice,
outlier::try_prune_outliers,
standard::encode_standard,
state::CachedTargetBw,
},
float::AlpFloat,
header::{header_len, raw_header_len, write_header},
params::AlpParams,
sampler::{BestParams, find_best_params, find_identical_base},
};
const HIGH_BW_PRUNE_THRESHOLD: u8 = 16;
const DELTA_EVAL_MIN_BW: u8 = 4;
const LOW_BW_PRUNE_MIN: u8 = 4;
const STACK_BUFFER_CAPACITY: usize = 1024;
const CACHE_VALIDATE_SAMPLE_N: usize = 4;
#[inline(always)]
fn check_roundtrip<F: AlpFloat>(slice: &[F], exp_factor: F, decode: impl Fn(F::Int) -> F) -> bool {
slice.iter().all(|&v| {
let enc = v.fast_round_to_int(exp_factor);
decode(enc).is_exact_same(v)
})
}
#[inline]
pub(crate) fn validate_cached_params<F: AlpFloat>(params: BestParams, sample: &[F]) -> bool {
let exp_factor = F::exp_factor(params.exp, params.fac);
let fac_int = F::fac_int(params.fac);
let frac_exp = F::frac_exp(params.exp);
let check_n = sample.len().min(CACHE_VALIDATE_SAMPLE_N);
let check_slice = &sample[..check_n];
if params.use_div {
check_roundtrip(check_slice, exp_factor, |enc| {
F::decode_from_int_div(enc, exp_factor)
})
} else if fac_int == 1 {
check_roundtrip(check_slice, exp_factor, |enc| {
F::decode_from_int_fac1(enc, frac_exp)
})
} else {
check_roundtrip(check_slice, exp_factor, |enc| {
F::decode_from_int(enc, fac_int, frac_exp)
})
}
}
#[inline]
fn write_raw_fallback<F: AlpFloat>(slice: &[F], count: usize, dst: &mut Vec<u8>) {
let raw_len = size_of_val(slice);
let raw_hdr = raw_header_len(count);
dst.reserve(raw_hdr + raw_len);
write_header(F::TYPE_RAW_BYTE, count, None, dst);
let raw_slice = unsafe { from_raw_parts(slice.as_ptr().cast::<u8>(), raw_len) };
dst.extend_from_slice(raw_slice);
}
#[inline(always)]
unsafe fn encode_pass<F: AlpFloat>(
slice: &[F],
enc_ptr: *mut F::Int,
params: BestParams,
exceptions: &mut Vec<Exception<F::RawBits>>,
) -> (F::Int, F::Int) {
let exp_factor = F::exp_factor(params.exp, params.fac);
let fac_int = F::fac_int(params.fac);
let frac_exp = F::frac_exp(params.exp);
unsafe {
encode_slice(
slice,
enc_ptr,
exp_factor,
fac_int,
frac_exp,
params.use_div,
exceptions,
)
}
}
#[inline(always)]
fn apply_target_bw<F: AlpFloat>(
slice: &[F],
encoded_ints: &mut [F::Int],
base: F::Int,
target_bw: u8,
exceptions: &mut Vec<Exception<F::RawBits>>,
) {
let max_allowed = if target_bw == 0 {
0u64
} else {
(1u64 << target_bw) - 1
};
let had_prev = !exceptions.is_empty();
for (pos, (&v, val_mut)) in slice.iter().zip(encoded_ints.iter_mut()).enumerate() {
let diff = F::int_diff_to_u64(*val_mut, base);
if diff > max_allowed {
exceptions.push(Exception {
pos,
bits: v.to_raw_bits(),
});
*val_mut = base;
}
}
if had_prev && exceptions.len() > 1 {
exceptions.sort_unstable_by_key(|e| e.pos);
exceptions.dedup_by_key(|e| e.pos);
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn compress_into_engine<F: AlpFloat>(
slice: &[F],
dst: &mut Vec<u8>,
force_delta: bool,
cached_params: Option<BestParams>,
cached_target_bw: &mut CachedTargetBw,
cached_use_delta: &mut Option<bool>,
encoded_buf: &mut Vec<F::Int>,
exceptions: &mut Vec<Exception<F::RawBits>>,
) -> Option<BestParams> {
let count = slice.len();
if count == 0 {
let raw_hdr = raw_header_len(0);
dst.reserve(raw_hdr);
write_header(F::TYPE_BYTE, 0, None, dst);
return None;
}
let first = slice[0];
if (count == 1 || slice[1].is_exact_same(first))
&& let Some((exp, base)) = find_identical_base(first)
&& slice[1..].iter().all(|&v| v.is_exact_same(first))
{
let total_needed = header_len(count) + F::BASE_SIZE;
dst.reserve(total_needed);
let params = AlpParams::new(exp, 0, 0, false);
write_header(F::TYPE_BYTE, count, Some(params.pack()), dst);
F::write_base(base, dst);
return Some(BestParams {
exp,
fac: 0,
use_div: false,
});
}
let mut best_params = match cached_params {
Some(p) if validate_cached_params(p, slice) => p,
_ => find_best_params(slice),
};
let mut stack_encoded = MaybeUninit::<[F::Int; STACK_BUFFER_CAPACITY]>::uninit();
exceptions.clear();
encoded_buf.clear();
let use_stack = count <= STACK_BUFFER_CAPACITY && encoded_buf.capacity() < count;
let enc_ptr: *mut F::Int = if use_stack {
stack_encoded.as_mut_ptr().cast::<F::Int>()
} else {
if encoded_buf.capacity() < count {
encoded_buf.reserve(count);
}
encoded_buf.as_mut_ptr()
};
let (mut min_val, mut max_val) = unsafe { encode_pass(slice, enc_ptr, best_params, exceptions) };
if exceptions.len() > MAX_EXCEPTIONS && cached_params.is_some() {
let fresh_params = find_best_params(slice);
if fresh_params != best_params {
exceptions.clear();
let (fresh_min, fresh_max) = unsafe { encode_pass(slice, enc_ptr, fresh_params, exceptions) };
if exceptions.len() <= MAX_EXCEPTIONS {
best_params = fresh_params;
min_val = fresh_min;
max_val = fresh_max;
}
}
}
if exceptions.len() > MAX_EXCEPTIONS {
write_raw_fallback(slice, count, dst);
return None;
}
let encoded_ints: &mut [F::Int] = if use_stack {
unsafe { from_raw_parts_mut(stack_encoded.as_mut_ptr().cast::<F::Int>(), count) }
} else {
unsafe { encoded_buf.set_len(count) };
&mut encoded_buf[..count]
};
let (base, max_offset) = if min_val <= max_val {
(min_val, F::calc_range(min_val, max_val))
} else {
(F::ZERO_INT, 0)
};
let is_large = count > u16::MAX as usize;
let mut for_bit_width = F::bits_needed(max_offset);
let mut did_pre_prune = false;
if for_bit_width >= HIGH_BW_PRUNE_THRESHOLD && exceptions.len() < MAX_EXCEPTIONS {
did_pre_prune = true;
match *cached_target_bw {
CachedTargetBw::Pruned(target_bw) if target_bw < for_bit_width => {
apply_target_bw(slice, encoded_ints, base, target_bw, exceptions);
for_bit_width = target_bw;
}
CachedTargetBw::Disabled | CachedTargetBw::Pruned(_) => {}
CachedTargetBw::Uninit => {
let pruned_bw = try_prune_outliers::<F>(
slice,
encoded_ints,
base,
for_bit_width,
exceptions,
is_large,
);
if pruned_bw < for_bit_width {
*cached_target_bw = CachedTargetBw::Pruned(pruned_bw);
for_bit_width = pruned_bw;
} else {
*cached_target_bw = CachedTargetBw::Disabled;
}
}
}
}
if !exceptions.is_empty() {
for exc in exceptions.iter() {
let patch_val = if exc.pos > 0 {
unsafe { *encoded_ints.get_unchecked(exc.pos - 1) }
} else {
base
};
unsafe {
*encoded_ints.get_unchecked_mut(exc.pos) = patch_val;
}
}
}
let mut for_packed_len = packed_byte_size(count, for_bit_width);
let mut exc_len = exceptions_byte_size::<F>(exceptions.len(), is_large);
let hdr_len = header_len(count);
let for_total = hdr_len + F::BASE_SIZE + for_packed_len + exc_len;
let delta_decision = if count > 1 {
if force_delta {
let first = encoded_ints[0];
let rest = &encoded_ints[1..];
Some(delta_range::<F>(first, rest))
} else {
match *cached_use_delta {
Some(false) => None,
Some(true) => {
let first = encoded_ints[0];
let rest = &encoded_ints[1..];
Some(delta_range::<F>(first, rest))
}
None => {
if for_bit_width >= DELTA_EVAL_MIN_BW {
let first = encoded_ints[0];
let rest = &encoded_ints[1..];
eval_delta_benefit::<F>(first, rest, for_bit_width)
} else {
None
}
}
}
}
} else {
None
};
let (use_delta, min_delta, delta_bit_width, mut total_needed) = match delta_decision {
Some((min_d, delta_bw)) => {
let delta_packed_len = packed_byte_size(count - 1, delta_bw);
let delta_total = hdr_len + F::BASE_SIZE * 2 + delta_packed_len + exc_len;
if delta_total < for_total || force_delta {
(true, min_d, delta_bw, delta_total)
} else {
(false, F::ZERO_INT, 0, for_total)
}
}
None => (false, F::ZERO_INT, 0, for_total),
};
if !force_delta && cached_use_delta.is_none() {
*cached_use_delta = Some(use_delta);
}
if !did_pre_prune
&& !use_delta
&& for_bit_width > LOW_BW_PRUNE_MIN
&& for_bit_width < HIGH_BW_PRUNE_THRESHOLD
{
match *cached_target_bw {
CachedTargetBw::Pruned(target_bw) if target_bw < for_bit_width => {
apply_target_bw(slice, encoded_ints, base, target_bw, exceptions);
for_bit_width = target_bw;
for_packed_len = packed_byte_size(count, for_bit_width);
exc_len = exceptions_byte_size::<F>(exceptions.len(), is_large);
total_needed = hdr_len + F::BASE_SIZE + for_packed_len + exc_len;
}
CachedTargetBw::Disabled | CachedTargetBw::Pruned(_) => {}
CachedTargetBw::Uninit => {
let new_bw = try_prune_outliers::<F>(
slice,
encoded_ints,
base,
for_bit_width,
exceptions,
is_large,
);
if new_bw < for_bit_width {
*cached_target_bw = CachedTargetBw::Pruned(new_bw);
for_bit_width = new_bw;
for_packed_len = packed_byte_size(count, for_bit_width);
exc_len = exceptions_byte_size::<F>(exceptions.len(), is_large);
} else {
*cached_target_bw = CachedTargetBw::Disabled;
}
total_needed = hdr_len + F::BASE_SIZE + for_packed_len + exc_len;
}
}
}
let raw_len = size_of_val(slice);
let raw_hdr = raw_header_len(count);
if total_needed >= raw_len + raw_hdr {
write_raw_fallback(slice, count, dst);
return None;
}
dst.reserve(total_needed);
if use_delta {
let params = AlpParams::from_best_params(best_params, delta_bit_width);
encode_delta::<F>(params, encoded_ints, min_delta, exceptions, dst);
} else {
let params = AlpParams::from_best_params(best_params, for_bit_width);
encode_standard::<F>(params, encoded_ints, base, exceptions, dst);
}
Some(best_params)
}
#[doc(hidden)]
pub fn profile_compress_breakdown<F: AlpFloat>(slice: &[F]) {
use std::time::Instant;
let count = slice.len();
let best_params = find_best_params(slice);
let mut encoded_buf = vec![F::Int::default(); count];
let enc_ptr = encoded_buf.as_mut_ptr();
let mut exceptions = Vec::new();
let iters = 10000;
let start = Instant::now();
let mut min_val = F::MAX_INT;
let mut max_val = F::MIN_INT;
for _ in 0..iters {
exceptions.clear();
let (mn, mx) = unsafe { encode_pass(slice, enc_ptr, best_params, &mut exceptions) };
min_val = mn;
max_val = mx;
}
let t_enc = start.elapsed().as_nanos() as f64 / iters as f64;
let (base, max_offset) = if min_val <= max_val {
(min_val, F::calc_range(min_val, max_val))
} else {
(F::ZERO_INT, 0)
};
let is_large = count > u16::MAX as usize;
let for_bit_width = F::bits_needed(max_offset);
let start = Instant::now();
for _ in 0..iters {
let mut exc_copy = exceptions.clone();
let _ = try_prune_outliers::<F>(
slice,
&mut encoded_buf,
base,
for_bit_width,
&mut exc_copy,
is_large,
);
}
let t_prune = start.elapsed().as_nanos() as f64 / iters as f64;
let start = Instant::now();
let first = encoded_buf[0];
let rest = &encoded_buf[1..];
for _ in 0..iters {
let _ = eval_delta_benefit::<F>(first, rest, for_bit_width);
}
let t_delta = start.elapsed().as_nanos() as f64 / iters as f64;
let mut dst = Vec::with_capacity(count * 8 + 64);
let params = AlpParams::from_best_params(best_params, for_bit_width);
let start = Instant::now();
for _ in 0..iters {
dst.clear();
encode_standard::<F>(params, &encoded_buf, base, &exceptions, &mut dst);
}
let t_pack = start.elapsed().as_nanos() as f64 / iters as f64;
println!(
" Breakdown: enc={:5.1} ns | prune={:5.1} ns | delta={:5.1} ns | pack={:5.1} ns (bw={}, exc={})",
t_enc,
t_prune,
t_delta,
t_pack,
for_bit_width,
exceptions.len()
);
}