use std::ptr::read_unaligned;
use crate::{
constants::{AlpFloat, bit_mask_u64},
error::{Error, Result},
};
#[inline(always)]
pub const fn packed_byte_size(count: usize, bit_width: u8) -> usize {
(count * (bit_width as usize)).div_ceil(8)
}
pub fn bitpack_u64(values: &[u64], bit_width: u8, dst: &mut Vec<u8>) {
if values.is_empty() || bit_width == 0 {
return;
}
let total_bytes = packed_byte_size(values.len(), bit_width);
let old_len = dst.len();
dst.reserve(total_bytes);
if bit_width == 8 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
for &v in values {
*dst_ptr = v as u8;
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 16 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len).cast::<[u8; 2]>();
for &v in values {
*dst_ptr = (v as u16).to_le_bytes();
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 32 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len).cast::<[u8; 4]>();
for &v in values {
*dst_ptr = (v as u32).to_le_bytes();
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 64 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len).cast::<[u8; 8]>();
for &v in values {
*dst_ptr = v.to_le_bytes();
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
}
let mask = bit_mask_u64(bit_width);
let mut acc: u128 = 0;
let mut bits: u32 = 0;
for &val in values {
acc |= ((val & mask) as u128) << bits;
bits += bit_width as u32;
if bits >= 64 {
dst.extend_from_slice(&(acc as u64).to_le_bytes());
acc >>= 64;
bits -= 64;
}
}
while bits > 0 {
dst.push(acc as u8);
acc >>= 8;
bits = bits.saturating_sub(8);
}
}
pub fn bitpack_encoded<F: AlpFloat>(
encoded_ints: &[F::Int],
base: F::Int,
bit_width: u8,
dst: &mut Vec<u8>,
) {
if encoded_ints.is_empty() || bit_width == 0 {
return;
}
let total_bytes = packed_byte_size(encoded_ints.len(), bit_width);
let old_len = dst.len();
dst.reserve(total_bytes);
if bit_width == 1 {
let (chunks, rem) = encoded_ints.as_chunks::<8>();
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
for chunk in chunks {
let o0 = (F::int_diff_to_u64(chunk[0], base) as u8) & 0x01;
let o1 = (F::int_diff_to_u64(chunk[1], base) as u8) & 0x01;
let o2 = (F::int_diff_to_u64(chunk[2], base) as u8) & 0x01;
let o3 = (F::int_diff_to_u64(chunk[3], base) as u8) & 0x01;
let o4 = (F::int_diff_to_u64(chunk[4], base) as u8) & 0x01;
let o5 = (F::int_diff_to_u64(chunk[5], base) as u8) & 0x01;
let o6 = (F::int_diff_to_u64(chunk[6], base) as u8) & 0x01;
let o7 = (F::int_diff_to_u64(chunk[7], base) as u8) & 0x01;
*dst_ptr =
o0 | (o1 << 1) | (o2 << 2) | (o3 << 3) | (o4 << 4) | (o5 << 5) | (o6 << 6) | (o7 << 7);
dst_ptr = dst_ptr.add(1);
}
if !rem.is_empty() {
let mut b = 0u8;
for (i, &val) in rem.iter().enumerate() {
let o = (F::int_diff_to_u64(val, base) as u8) & 0x01;
b |= o << i;
}
*dst_ptr = b;
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 2 {
let (chunks, rem) = encoded_ints.as_chunks::<4>();
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
for chunk in chunks {
let o0 = (F::int_diff_to_u64(chunk[0], base) as u8) & 0x03;
let o1 = (F::int_diff_to_u64(chunk[1], base) as u8) & 0x03;
let o2 = (F::int_diff_to_u64(chunk[2], base) as u8) & 0x03;
let o3 = (F::int_diff_to_u64(chunk[3], base) as u8) & 0x03;
*dst_ptr = o0 | (o1 << 2) | (o2 << 4) | (o3 << 6);
dst_ptr = dst_ptr.add(1);
}
if !rem.is_empty() {
let mut b = 0u8;
for (i, &val) in rem.iter().enumerate() {
let o = (F::int_diff_to_u64(val, base) as u8) & 0x03;
b |= o << (i * 2);
}
*dst_ptr = b;
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 4 {
let (chunks, rem) = encoded_ints.as_chunks::<2>();
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
for chunk in chunks {
let o0 = (F::int_diff_to_u64(chunk[0], base) as u8) & 0x0f;
let o1 = (F::int_diff_to_u64(chunk[1], base) as u8) & 0x0f;
*dst_ptr = o0 | (o1 << 4);
dst_ptr = dst_ptr.add(1);
}
if let Some(&last) = rem.first() {
let o0 = (F::int_diff_to_u64(last, base) as u8) & 0x0f;
*dst_ptr = o0;
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 8 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
for &v in encoded_ints {
*dst_ptr = F::int_diff_to_u64(v, base) as u8;
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 16 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len).cast::<[u8; 2]>();
for &v in encoded_ints {
*dst_ptr = (F::int_diff_to_u64(v, base) as u16).to_le_bytes();
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 32 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len).cast::<[u8; 4]>();
for &v in encoded_ints {
*dst_ptr = (F::int_diff_to_u64(v, base) as u32).to_le_bytes();
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
} else if bit_width == 64 {
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len).cast::<[u8; 8]>();
for &v in encoded_ints {
*dst_ptr = F::int_diff_to_u64(v, base).to_le_bytes();
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + total_bytes);
}
return;
}
let mask = bit_mask_u64(bit_width);
let mut acc: u128 = 0;
let mut bits: u32 = 0;
for &val in encoded_ints {
let offset = F::int_diff_to_u64(val, base) & mask;
acc |= (offset as u128) << bits;
bits += bit_width as u32;
if bits >= 64 {
dst.extend_from_slice(&(acc as u64).to_le_bytes());
acc >>= 64;
bits -= 64;
}
}
while bits > 0 {
dst.push(acc as u8);
acc >>= 8;
bits = bits.saturating_sub(8);
}
}
#[inline]
pub fn bitpack_encoded_f64(encoded_ints: &[i64], base: i64, bit_width: u8, dst: &mut Vec<u8>) {
bitpack_encoded::<f64>(encoded_ints, base, bit_width, dst);
}
#[inline]
pub fn bitpack_encoded_f32(encoded_ints: &[i32], base: i32, bit_width: u8, dst: &mut Vec<u8>) {
bitpack_encoded::<f32>(encoded_ints, base, bit_width, dst);
}
pub fn bitunpack_u64(src: &[u8], count: usize, bit_width: u8, dst: &mut Vec<u64>) -> Result<()> {
if count == 0 {
return Ok(());
}
if bit_width == 0 {
dst.resize(dst.len() + count, 0);
return Ok(());
}
let required_bytes = packed_byte_size(count, bit_width);
if src.len() < required_bytes {
return Err(Error::UnexpectedEof {
needed: required_bytes,
available: src.len(),
});
}
let old_len = dst.len();
dst.reserve(count);
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
if bit_width == 8 {
for i in 0..count {
*dst_ptr.add(i) = *src.get_unchecked(i) as u64;
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 16 {
let src_ptr = src.as_ptr().cast::<[u8; 2]>();
for i in 0..count {
let bytes = read_unaligned(src_ptr.add(i));
*dst_ptr.add(i) = u16::from_le_bytes(bytes) as u64;
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 32 {
let src_ptr = src.as_ptr().cast::<[u8; 4]>();
for i in 0..count {
let bytes = read_unaligned(src_ptr.add(i));
*dst_ptr.add(i) = u32::from_le_bytes(bytes) as u64;
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 64 {
let src_ptr = src.as_ptr().cast::<[u8; 8]>();
for i in 0..count {
let bytes = read_unaligned(src_ptr.add(i));
*dst_ptr.add(i) = u64::from_le_bytes(bytes);
}
dst.set_len(old_len + count);
return Ok(());
}
let mask = bit_mask_u64(bit_width);
let mut acc: u128 = 0;
let mut bits_in_acc: u32 = 0;
let mut src_ptr = src.as_ptr();
let src_end = src.as_ptr().add(src.len());
let mut i = 0;
while i < count && src_ptr.add(8) <= src_end {
if bits_in_acc < bit_width as u32 {
let bytes = read_unaligned(src_ptr.cast::<[u8; 8]>());
let chunk = u64::from_le_bytes(bytes);
acc |= (chunk as u128) << bits_in_acc;
bits_in_acc += 64;
src_ptr = src_ptr.add(8);
}
let val = (acc as u64) & mask;
acc >>= bit_width;
bits_in_acc -= bit_width as u32;
*dst_ptr = val;
dst_ptr = dst_ptr.add(1);
i += 1;
}
while i < count {
while bits_in_acc < bit_width as u32 && src_ptr < src_end {
acc |= (*src_ptr as u128) << bits_in_acc;
bits_in_acc += 8;
src_ptr = src_ptr.add(1);
}
let val = (acc as u64) & mask;
acc >>= bit_width;
bits_in_acc = bits_in_acc.saturating_sub(bit_width as u32);
*dst_ptr = val;
dst_ptr = dst_ptr.add(1);
i += 1;
}
dst.set_len(old_len + count);
}
Ok(())
}
#[inline(always)]
pub fn bitunpack_into<F: AlpFloat>(
src: &[u8],
count: usize,
bit_width: u8,
base: F::Int,
fac_int: i64,
frac_flt: F,
dst: &mut Vec<F>,
) -> Result<()> {
if count == 0 {
return Ok(());
}
let required_bytes = packed_byte_size(count, bit_width);
if src.len() < required_bytes {
return Err(Error::UnexpectedEof {
needed: required_bytes,
available: src.len(),
});
}
if bit_width == 0 {
let val = F::decode_from_offset(0, base, fac_int, frac_flt);
dst.resize(dst.len() + count, val);
return Ok(());
}
let old_len = dst.len();
dst.reserve(count);
unsafe {
let mut dst_ptr = dst.as_mut_ptr().add(old_len);
if bit_width == 1 {
let lut = F::build_lut::<2>(base, fac_int, frac_flt);
let full_bytes = count / 8;
for &b in &src[..full_bytes] {
*dst_ptr.add(0) = *lut.get_unchecked((b & 0x01) as usize);
*dst_ptr.add(1) = *lut.get_unchecked(((b >> 1) & 0x01) as usize);
*dst_ptr.add(2) = *lut.get_unchecked(((b >> 2) & 0x01) as usize);
*dst_ptr.add(3) = *lut.get_unchecked(((b >> 3) & 0x01) as usize);
*dst_ptr.add(4) = *lut.get_unchecked(((b >> 4) & 0x01) as usize);
*dst_ptr.add(5) = *lut.get_unchecked(((b >> 5) & 0x01) as usize);
*dst_ptr.add(6) = *lut.get_unchecked(((b >> 6) & 0x01) as usize);
*dst_ptr.add(7) = *lut.get_unchecked(((b >> 7) & 0x01) as usize);
dst_ptr = dst_ptr.add(8);
}
let rem = count % 8;
if rem > 0 {
let b = *src.get_unchecked(full_bytes);
for shift in 0..rem {
let idx = ((b >> shift) & 0x01) as usize;
*dst_ptr = *lut.get_unchecked(idx);
dst_ptr = dst_ptr.add(1);
}
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 2 {
let lut = F::build_lut::<4>(base, fac_int, frac_flt);
let full_bytes = count / 4;
for &b in &src[..full_bytes] {
*dst_ptr.add(0) = *lut.get_unchecked((b & 0x03) as usize);
*dst_ptr.add(1) = *lut.get_unchecked(((b >> 2) & 0x03) as usize);
*dst_ptr.add(2) = *lut.get_unchecked(((b >> 4) & 0x03) as usize);
*dst_ptr.add(3) = *lut.get_unchecked(((b >> 6) & 0x03) as usize);
dst_ptr = dst_ptr.add(4);
}
let rem = count % 4;
if rem > 0 {
let b = *src.get_unchecked(full_bytes);
for i in 0..rem {
let idx = ((b >> (i * 2)) & 0x03) as usize;
*dst_ptr = *lut.get_unchecked(idx);
dst_ptr = dst_ptr.add(1);
}
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 4 {
let lut = F::build_lut::<16>(base, fac_int, frac_flt);
let full_bytes = count / 2;
for &b in &src[..full_bytes] {
*dst_ptr.add(0) = *lut.get_unchecked((b & 0x0f) as usize);
*dst_ptr.add(1) = *lut.get_unchecked((b >> 4) as usize);
dst_ptr = dst_ptr.add(2);
}
if !count.is_multiple_of(2) {
let b = *src.get_unchecked(full_bytes);
*dst_ptr = *lut.get_unchecked((b & 0x0f) as usize);
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 8 {
let lut = F::build_lut::<256>(base, fac_int, frac_flt);
for &b in &src[..count] {
*dst_ptr = *lut.get_unchecked(b as usize);
dst_ptr = dst_ptr.add(1);
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 16 {
let src_ptr = src.as_ptr().cast::<[u8; 2]>();
for i in 0..count {
let bytes = read_unaligned(src_ptr.add(i));
let off = u16::from_le_bytes(bytes) as u64;
*dst_ptr.add(i) = F::decode_from_offset(off, base, fac_int, frac_flt);
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 32 {
let src_ptr = src.as_ptr().cast::<[u8; 4]>();
for i in 0..count {
let bytes = read_unaligned(src_ptr.add(i));
let off = u32::from_le_bytes(bytes) as u64;
*dst_ptr.add(i) = F::decode_from_offset(off, base, fac_int, frac_flt);
}
dst.set_len(old_len + count);
return Ok(());
} else if bit_width == 64 {
let src_ptr = src.as_ptr().cast::<[u8; 8]>();
for i in 0..count {
let bytes = read_unaligned(src_ptr.add(i));
let off = u64::from_le_bytes(bytes);
*dst_ptr.add(i) = F::decode_from_offset(off, base, fac_int, frac_flt);
}
dst.set_len(old_len + count);
return Ok(());
}
let mask = bit_mask_u64(bit_width);
let mut acc: u128 = 0;
let mut bits_in_acc: u32 = 0;
let mut src_ptr = src.as_ptr();
let src_end = src.as_ptr().add(src.len());
let mut i = 0;
while i < count && src_ptr.add(8) <= src_end {
if bits_in_acc < bit_width as u32 {
let bytes = read_unaligned(src_ptr.cast::<[u8; 8]>());
let chunk = u64::from_le_bytes(bytes);
acc |= (chunk as u128) << bits_in_acc;
bits_in_acc += 64;
src_ptr = src_ptr.add(8);
}
let off = (acc as u64) & mask;
acc >>= bit_width;
bits_in_acc -= bit_width as u32;
*dst_ptr = F::decode_from_offset(off, base, fac_int, frac_flt);
dst_ptr = dst_ptr.add(1);
i += 1;
}
while i < count {
while bits_in_acc < bit_width as u32 && src_ptr < src_end {
acc |= (*src_ptr as u128) << bits_in_acc;
bits_in_acc += 8;
src_ptr = src_ptr.add(1);
}
let off = (acc as u64) & mask;
acc >>= bit_width;
bits_in_acc = bits_in_acc.saturating_sub(bit_width as u32);
*dst_ptr = F::decode_from_offset(off, base, fac_int, frac_flt);
dst_ptr = dst_ptr.add(1);
i += 1;
}
dst.set_len(old_len + count);
}
Ok(())
}
#[inline(always)]
pub fn bitunpack_f64_into(
src: &[u8],
count: usize,
bit_width: u8,
base: i64,
fac_int: i64,
frac_flt: f64,
dst: &mut Vec<f64>,
) -> Result<()> {
bitunpack_into::<f64>(src, count, bit_width, base, fac_int, frac_flt, dst)
}
#[inline(always)]
pub fn bitunpack_f32_into(
src: &[u8],
count: usize,
bit_width: u8,
base: i32,
fac_int: i64,
frac_flt: f32,
dst: &mut Vec<f32>,
) -> Result<()> {
bitunpack_into::<f32>(src, count, bit_width, base, fac_int, frac_flt, dst)
}