use std::sync::Arc;
use arrow_buffer::bit_util::ceil;
use bytes::Bytes;
use futures::future::{BoxFuture, FutureExt};
use log::trace;
use lance_bitpacking::BitPacking;
use lance_core::{Error, Result};
use crate::buffer::LanceBuffer;
use crate::data::BlockInfo;
use crate::data::{DataBlock, FixedWidthDataBlock};
use crate::decoder::{PageScheduler, PrimitivePageDecoder};
use bytemuck::cast_slice;
const LOG_ELEMS_PER_CHUNK: u8 = 10;
const ELEMS_PER_CHUNK: u64 = 1 << LOG_ELEMS_PER_CHUNK;
#[derive(Debug)]
pub struct BitpackedForNonNegScheduler {
compressed_bit_width: u64,
uncompressed_bits_per_value: u64,
buffer_offset: u64,
}
impl BitpackedForNonNegScheduler {
pub fn new(
compressed_bit_width: u64,
uncompressed_bits_per_value: u64,
buffer_offset: u64,
) -> Self {
Self {
compressed_bit_width,
uncompressed_bits_per_value,
buffer_offset,
}
}
fn locate_chunk_start(&self, relative_row_num: u64) -> u64 {
let chunk_size = ELEMS_PER_CHUNK * self.compressed_bit_width / 8;
self.buffer_offset + (relative_row_num / ELEMS_PER_CHUNK * chunk_size)
}
fn locate_chunk_end(&self, relative_row_num: u64) -> u64 {
let chunk_size = ELEMS_PER_CHUNK * self.compressed_bit_width / 8;
self.buffer_offset + (relative_row_num / ELEMS_PER_CHUNK * chunk_size) + chunk_size
}
}
impl PageScheduler for BitpackedForNonNegScheduler {
fn schedule_ranges(
&self,
ranges: &[std::ops::Range<u64>],
scheduler: &Arc<dyn crate::EncodingsIo>,
top_level_row: u64,
) -> BoxFuture<'static, Result<Box<dyn PrimitivePageDecoder>>> {
assert!(!ranges.is_empty());
let mut byte_ranges = vec![];
let mut bytes_idx_to_range_indices = vec![];
let first_byte_range = std::ops::Range {
start: self.locate_chunk_start(ranges[0].start),
end: self.locate_chunk_end(ranges[0].end - 1),
}; byte_ranges.push(first_byte_range);
bytes_idx_to_range_indices.push(vec![ranges[0].clone()]);
for (i, range) in ranges.iter().enumerate().skip(1) {
let this_start = self.locate_chunk_start(range.start);
let this_end = self.locate_chunk_end(range.end - 1);
if this_start == self.locate_chunk_start(ranges[i - 1].end - 1) {
byte_ranges.last_mut().unwrap().end = this_end;
bytes_idx_to_range_indices
.last_mut()
.unwrap()
.push(range.clone());
} else {
byte_ranges.push(this_start..this_end);
bytes_idx_to_range_indices.push(vec![range.clone()]);
}
}
trace!(
"Scheduling I/O for {} ranges spread across byte range {}..{}",
byte_ranges.len(),
byte_ranges[0].start,
byte_ranges.last().unwrap().end
);
let bytes = scheduler.submit_request(byte_ranges.clone(), top_level_row);
let compressed_bit_width = self.compressed_bit_width;
let uncompressed_bits_per_value = self.uncompressed_bits_per_value;
let num_rows = ranges.iter().map(|range| range.end - range.start).sum();
async move {
let bytes = bytes.await?;
let decompressed_output = bitpacked_for_non_neg_decode(
compressed_bit_width,
uncompressed_bits_per_value,
&bytes,
&bytes_idx_to_range_indices,
num_rows,
);
Ok(Box::new(BitpackedForNonNegPageDecoder {
uncompressed_bits_per_value,
decompressed_buf: decompressed_output,
}) as Box<dyn PrimitivePageDecoder>)
}
.boxed()
}
}
#[derive(Debug)]
struct BitpackedForNonNegPageDecoder {
uncompressed_bits_per_value: u64,
decompressed_buf: LanceBuffer,
}
impl PrimitivePageDecoder for BitpackedForNonNegPageDecoder {
fn decode(&self, rows_to_skip: u64, num_rows: u64) -> Result<DataBlock> {
if ![8, 16, 32, 64].contains(&self.uncompressed_bits_per_value) {
return Err(Error::invalid_input_source("BitpackedForNonNegPageDecoder should only has uncompressed_bits_per_value of 8, 16, 32, or 64".into()));
}
let elem_size_in_bytes = self.uncompressed_bits_per_value / 8;
Ok(DataBlock::FixedWidth(FixedWidthDataBlock {
data: self.decompressed_buf.slice_with_length(
(rows_to_skip * elem_size_in_bytes) as usize,
(num_rows * elem_size_in_bytes) as usize,
),
bits_per_value: self.uncompressed_bits_per_value,
num_values: num_rows,
block_info: BlockInfo::new(),
}))
}
}
macro_rules! bitpacked_decode {
($uncompressed_type:ty, $compressed_bit_width:expr, $data:expr, $bytes_idx_to_range_indices:expr, $num_rows:expr) => {{
let mut decompressed: Vec<$uncompressed_type> = Vec::with_capacity($num_rows as usize);
let packed_chunk_size_in_byte: usize = (ELEMS_PER_CHUNK * $compressed_bit_width) as usize / 8;
let mut decompress_chunk_buf = vec![0 as $uncompressed_type; ELEMS_PER_CHUNK as usize];
for (i, bytes) in $data.iter().enumerate() {
let mut ranges_idx = 0;
let mut curr_range_start = $bytes_idx_to_range_indices[i][0].start;
let mut chunk_num = 0;
while chunk_num * packed_chunk_size_in_byte < bytes.len() {
let chunk_in_u8: Vec<u8> = bytes[chunk_num * packed_chunk_size_in_byte..]
[..packed_chunk_size_in_byte]
.to_vec();
chunk_num += 1;
let chunk = cast_slice(&chunk_in_u8);
unsafe {
BitPacking::unchecked_unpack(
$compressed_bit_width as usize,
chunk,
&mut decompress_chunk_buf,
);
}
loop {
let elems_after_curr_range_start_in_this_chunk =
ELEMS_PER_CHUNK - curr_range_start % ELEMS_PER_CHUNK;
if curr_range_start + elems_after_curr_range_start_in_this_chunk
<= $bytes_idx_to_range_indices[i][ranges_idx].end
{
decompressed.extend_from_slice(
&decompress_chunk_buf[(curr_range_start % ELEMS_PER_CHUNK) as usize..],
);
curr_range_start += elems_after_curr_range_start_in_this_chunk;
break;
} else {
let elems_this_range_needed_in_this_chunk =
($bytes_idx_to_range_indices[i][ranges_idx].end - curr_range_start)
.min(ELEMS_PER_CHUNK - curr_range_start % ELEMS_PER_CHUNK);
decompressed.extend_from_slice(
&decompress_chunk_buf[(curr_range_start % ELEMS_PER_CHUNK) as usize..]
[..elems_this_range_needed_in_this_chunk as usize],
);
if curr_range_start + elems_this_range_needed_in_this_chunk
== $bytes_idx_to_range_indices[i][ranges_idx].end
{
ranges_idx += 1;
if ranges_idx == $bytes_idx_to_range_indices[i].len() {
break;
}
curr_range_start = $bytes_idx_to_range_indices[i][ranges_idx].start;
} else {
curr_range_start += elems_this_range_needed_in_this_chunk;
}
}
}
}
}
LanceBuffer::reinterpret_vec(decompressed)
}};
}
fn bitpacked_for_non_neg_decode(
compressed_bit_width: u64,
uncompressed_bits_per_value: u64,
data: &[Bytes],
bytes_idx_to_range_indices: &[Vec<std::ops::Range<u64>>],
num_rows: u64,
) -> LanceBuffer {
match uncompressed_bits_per_value {
8 => bitpacked_decode!(
u8,
compressed_bit_width,
data,
bytes_idx_to_range_indices,
num_rows
),
16 => bitpacked_decode!(
u16,
compressed_bit_width,
data,
bytes_idx_to_range_indices,
num_rows
),
32 => bitpacked_decode!(
u32,
compressed_bit_width,
data,
bytes_idx_to_range_indices,
num_rows
),
64 => bitpacked_decode!(
u64,
compressed_bit_width,
data,
bytes_idx_to_range_indices,
num_rows
),
_ => unreachable!(
"bitpacked_for_non_neg_decode only supports 8, 16, 32, 64 uncompressed_bits_per_value"
),
}
}
#[derive(Debug, Clone, Copy)]
pub struct BitpackedScheduler {
bits_per_value: u64,
uncompressed_bits_per_value: u64,
buffer_offset: u64,
signed: bool,
}
impl BitpackedScheduler {
pub fn new(
bits_per_value: u64,
uncompressed_bits_per_value: u64,
buffer_offset: u64,
signed: bool,
) -> Self {
Self {
bits_per_value,
uncompressed_bits_per_value,
buffer_offset,
signed,
}
}
}
impl PageScheduler for BitpackedScheduler {
fn schedule_ranges(
&self,
ranges: &[std::ops::Range<u64>],
scheduler: &Arc<dyn crate::EncodingsIo>,
top_level_row: u64,
) -> BoxFuture<'static, Result<Box<dyn PrimitivePageDecoder>>> {
let mut min = u64::MAX;
let mut max = 0;
let mut buffer_bit_start_offsets: Vec<u8> = vec![];
let mut buffer_bit_end_offsets: Vec<Option<u8>> = vec![];
let byte_ranges = ranges
.iter()
.map(|range| {
let start_byte_offset = range.start * self.bits_per_value / 8;
let mut end_byte_offset = range.end * self.bits_per_value / 8;
if !(range.end * self.bits_per_value).is_multiple_of(8) {
end_byte_offset += 1;
let end_bit_offset = range.end * self.bits_per_value % 8;
buffer_bit_end_offsets.push(Some(end_bit_offset as u8));
} else {
buffer_bit_end_offsets.push(None);
}
let start_bit_offset = range.start * self.bits_per_value % 8;
buffer_bit_start_offsets.push(start_bit_offset as u8);
let start = self.buffer_offset + start_byte_offset;
let end = self.buffer_offset + end_byte_offset;
min = min.min(start);
max = max.max(end);
start..end
})
.collect::<Vec<_>>();
trace!(
"Scheduling I/O for {} ranges spread across byte range {}..{}",
byte_ranges.len(),
min,
max
);
let bytes = scheduler.submit_request(byte_ranges, top_level_row);
let bits_per_value = self.bits_per_value;
let uncompressed_bits_per_value = self.uncompressed_bits_per_value;
let signed = self.signed;
async move {
let bytes = bytes.await?;
Ok(Box::new(BitpackedPageDecoder {
buffer_bit_start_offsets,
buffer_bit_end_offsets,
bits_per_value,
uncompressed_bits_per_value,
signed,
data: bytes,
}) as Box<dyn PrimitivePageDecoder>)
}
.boxed()
}
}
#[derive(Debug)]
struct BitpackedPageDecoder {
buffer_bit_start_offsets: Vec<u8>,
buffer_bit_end_offsets: Vec<Option<u8>>,
bits_per_value: u64,
uncompressed_bits_per_value: u64,
signed: bool,
data: Vec<Bytes>,
}
impl PrimitivePageDecoder for BitpackedPageDecoder {
fn decode(&self, rows_to_skip: u64, num_rows: u64) -> Result<DataBlock> {
let num_bytes = self.uncompressed_bits_per_value / 8 * num_rows;
let mut dest = vec![0; num_bytes as usize];
debug_assert!(self.bits_per_value <= 64);
let mut rows_to_skip = rows_to_skip;
let mut rows_taken = 0;
let byte_len = self.uncompressed_bits_per_value / 8;
let mut dst_idx = 0;
let mask = u64::MAX >> (64 - self.bits_per_value);
for i in 0..self.data.len() {
let src = &self.data[i];
let (mut src_idx, mut src_offset) = match compute_start_offset(
rows_to_skip,
src.len(),
self.bits_per_value,
self.buffer_bit_start_offsets[i],
self.buffer_bit_end_offsets[i],
) {
StartOffset::SkipFull(rows_to_skip_here) => {
rows_to_skip -= rows_to_skip_here;
continue;
}
StartOffset::SkipSome(buffer_start_offset) => (
buffer_start_offset.index,
buffer_start_offset.bit_offset as u64,
),
};
while src_idx < src.len() && rows_taken < num_rows {
rows_taken += 1;
let mut curr_mask = mask;
let mut curr_src = src[src_idx] & (curr_mask << src_offset) as u8;
let mut src_bits_written = 0;
let mut dst_offset = 0;
let is_negative = is_encoded_item_negative(
src,
src_idx,
src_offset,
self.bits_per_value as usize,
);
while src_bits_written < self.bits_per_value {
dest[dst_idx] += (curr_src >> src_offset) << dst_offset;
let bits_written = (self.bits_per_value - src_bits_written)
.min(8 - src_offset)
.min(8 - dst_offset);
src_bits_written += bits_written;
dst_offset += bits_written;
src_offset += bits_written;
curr_mask >>= bits_written;
if dst_offset == 8 {
dst_idx += 1;
dst_offset = 0;
}
if src_offset == 8 {
src_idx += 1;
src_offset = 0;
if src_idx == src.len() {
break;
}
curr_src = src[src_idx] & curr_mask as u8;
}
}
let mut negative_padded_current_byte = false;
if self.signed && is_negative && dst_offset > 0 {
negative_padded_current_byte = true;
while dst_offset < 8 {
dest[dst_idx] |= 1 << dst_offset;
dst_offset += 1;
}
}
if self.uncompressed_bits_per_value != self.bits_per_value {
let partial_bytes_written = ceil(self.bits_per_value as usize, 8);
let mut to_next_byte = 1;
if self.bits_per_value.is_multiple_of(8) {
to_next_byte = 0;
}
let next_dst_idx =
dst_idx + byte_len as usize - partial_bytes_written + to_next_byte;
if self.signed && is_negative {
if !negative_padded_current_byte {
dest[dst_idx] = 0xFF;
}
for i in dest.iter_mut().take(next_dst_idx).skip(dst_idx + 1) {
*i = 0xFF;
}
}
dst_idx = next_dst_idx;
}
if let Some(buffer_bit_end_offset) = self.buffer_bit_end_offsets[i]
&& src_idx == src.len() - 1
&& src_offset >= buffer_bit_end_offset as u64
{
break;
}
}
}
Ok(DataBlock::FixedWidth(FixedWidthDataBlock {
data: LanceBuffer::from(dest),
bits_per_value: self.uncompressed_bits_per_value,
num_values: num_rows,
block_info: BlockInfo::new(),
}))
}
}
fn is_encoded_item_negative(src: &Bytes, src_idx: usize, src_offset: u64, num_bits: usize) -> bool {
let mut last_byte_idx = src_idx + ((src_offset as usize + num_bits) / 8);
let shift_amount = (src_offset as usize + num_bits) % 8;
let shift_amount = if shift_amount == 0 {
last_byte_idx -= 1;
7
} else {
shift_amount - 1
};
let last_byte = src[last_byte_idx];
let sign_bit_mask = 1 << shift_amount;
let sign_bit = last_byte & sign_bit_mask;
sign_bit > 0
}
#[derive(Debug, PartialEq)]
struct BufferStartOffset {
index: usize,
bit_offset: u8,
}
#[derive(Debug, PartialEq)]
enum StartOffset {
SkipFull(u64),
SkipSome(BufferStartOffset),
}
fn compute_start_offset(
rows_to_skip: u64,
buffer_len: usize,
bits_per_value: u64,
buffer_start_bit_offset: u8,
buffer_end_bit_offset: Option<u8>,
) -> StartOffset {
let rows_in_buffer = rows_in_buffer(
buffer_len,
bits_per_value,
buffer_start_bit_offset,
buffer_end_bit_offset,
);
if rows_to_skip >= rows_in_buffer {
return StartOffset::SkipFull(rows_in_buffer);
}
let start_bit = rows_to_skip * bits_per_value + buffer_start_bit_offset as u64;
let start_byte = start_bit / 8;
StartOffset::SkipSome(BufferStartOffset {
index: start_byte as usize,
bit_offset: (start_bit % 8) as u8,
})
}
fn rows_in_buffer(
buffer_len: usize,
bits_per_value: u64,
buffer_start_bit_offset: u8,
buffer_end_bit_offset: Option<u8>,
) -> u64 {
let mut bits_in_buffer = (buffer_len * 8) as u64 - buffer_start_bit_offset as u64;
if let Some(buffer_end_bit_offset) = buffer_end_bit_offset {
bits_in_buffer -= (8 - buffer_end_bit_offset) as u64;
}
bits_in_buffer / bits_per_value
}
#[cfg(test)]
pub mod test {
use crate::testing::{ArrayGeneratorProvider, TestCases, check_round_trip_encoding_generated};
use super::*;
use std::marker::PhantomData;
use arrow_array::{
ArrowPrimitiveType, PrimitiveArray,
types::{Int16Type, Int32Type, Int64Type, UInt8Type, UInt32Type, UInt64Type},
};
use arrow_schema::{DataType, Field};
use lance_datagen::{ArrayGenerator, array::rand_with_distribution};
use rand::distr::Uniform;
#[test]
fn test_rows_in_buffer() {
let test_cases = vec![
(5usize, 5u64, 0u8, None, 8u64),
(2, 3, 0, Some(5), 4),
(2, 3, 7, Some(6), 2),
];
for (
buffer_len,
bits_per_value,
buffer_start_bit_offset,
buffer_end_bit_offset,
expected,
) in test_cases
{
let result = rows_in_buffer(
buffer_len,
bits_per_value,
buffer_start_bit_offset,
buffer_end_bit_offset,
);
assert_eq!(expected, result);
}
}
#[test]
fn test_compute_start_offset() {
let result = compute_start_offset(0, 5, 5, 0, None);
assert_eq!(
StartOffset::SkipSome(BufferStartOffset {
index: 0,
bit_offset: 0
}),
result
);
let result = compute_start_offset(10, 5, 5, 0, None);
assert_eq!(StartOffset::SkipFull(8), result);
}
struct DistributionArrayGeneratorProvider<
DataType,
Dist: rand::distr::Distribution<DataType::Native> + Clone + Send + Sync + 'static,
>
where
DataType::Native: Copy + 'static,
PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
DataType: ArrowPrimitiveType,
{
phantom: PhantomData<DataType>,
distribution: Dist,
}
impl<DataType, Dist> DistributionArrayGeneratorProvider<DataType, Dist>
where
Dist: rand::distr::Distribution<DataType::Native> + Clone + Send + Sync + 'static,
DataType::Native: Copy + 'static,
PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
DataType: ArrowPrimitiveType,
{
fn new(dist: Dist) -> Self {
Self {
distribution: dist,
phantom: Default::default(),
}
}
}
impl<DataType, Dist> ArrayGeneratorProvider for DistributionArrayGeneratorProvider<DataType, Dist>
where
Dist: rand::distr::Distribution<DataType::Native> + Clone + Send + Sync + 'static,
DataType::Native: Copy + 'static,
PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
DataType: ArrowPrimitiveType,
{
fn provide(&self) -> Box<dyn ArrayGenerator> {
rand_with_distribution::<DataType, Dist>(self.distribution.clone())
}
fn copy(&self) -> Box<dyn ArrayGeneratorProvider> {
Box::new(Self {
phantom: self.phantom,
distribution: self.distribution.clone(),
})
}
}
#[test_log::test(tokio::test)]
async fn test_bitpack_primitive() {
let bitpacked_test_cases: &Vec<(DataType, Box<dyn ArrayGeneratorProvider>)> = &vec![
(
DataType::UInt32,
Box::new(
DistributionArrayGeneratorProvider::<UInt32Type, Uniform<u32>>::new(
Uniform::new(0, 19).unwrap(),
),
),
),
(
DataType::UInt32,
Box::new(
DistributionArrayGeneratorProvider::<UInt32Type, Uniform<u32>>::new(
Uniform::new(5 << 7, 6 << 7).unwrap(),
),
),
),
(
DataType::UInt64,
Box::new(
DistributionArrayGeneratorProvider::<UInt64Type, Uniform<u64>>::new(
Uniform::new(5 << 42, 6 << 42).unwrap(),
),
),
),
(
DataType::UInt8,
Box::new(
DistributionArrayGeneratorProvider::<UInt8Type, Uniform<u8>>::new(
Uniform::new(0, 19).unwrap(),
),
),
),
(
DataType::UInt64,
Box::new(
DistributionArrayGeneratorProvider::<UInt64Type, Uniform<u64>>::new(
Uniform::new(129, 259).unwrap(),
),
),
),
(
DataType::UInt32,
Box::new(
DistributionArrayGeneratorProvider::<UInt32Type, Uniform<u32>>::new(
Uniform::new(200, 250).unwrap(),
),
),
),
(
DataType::UInt64,
Box::new(
DistributionArrayGeneratorProvider::<UInt64Type, Uniform<u64>>::new(
Uniform::new(1, 3).unwrap(), ),
),
),
(
DataType::UInt32,
Box::new(
DistributionArrayGeneratorProvider::<UInt32Type, Uniform<u32>>::new(
Uniform::new(200 << 8, 250 << 8).unwrap(),
),
),
),
(
DataType::UInt64,
Box::new(
DistributionArrayGeneratorProvider::<UInt64Type, Uniform<u64>>::new(
Uniform::new(200 << 16, 250 << 16).unwrap(),
),
),
),
(
DataType::UInt32,
Box::new(
DistributionArrayGeneratorProvider::<UInt32Type, Uniform<u32>>::new(
Uniform::new(0, 1).unwrap(),
),
),
),
(
DataType::Int16,
Box::new(
DistributionArrayGeneratorProvider::<Int16Type, Uniform<i16>>::new(
Uniform::new(-5, 5).unwrap(),
),
),
),
(
DataType::Int64,
Box::new(
DistributionArrayGeneratorProvider::<Int64Type, Uniform<i64>>::new(
Uniform::new(-(5 << 42), 6 << 42).unwrap(),
),
),
),
(
DataType::Int32,
Box::new(
DistributionArrayGeneratorProvider::<Int32Type, Uniform<i32>>::new(
Uniform::new(-(5 << 7), 6 << 7).unwrap(),
),
),
),
(
DataType::Int32,
Box::new(
DistributionArrayGeneratorProvider::<Int32Type, Uniform<i32>>::new(
Uniform::new(-19, 19).unwrap(),
),
),
),
(
DataType::Int32,
Box::new(
DistributionArrayGeneratorProvider::<Int32Type, Uniform<i32>>::new(
Uniform::new(-120, 120).unwrap(),
),
),
),
(
DataType::Int32,
Box::new(
DistributionArrayGeneratorProvider::<Int32Type, Uniform<i32>>::new(
Uniform::new(-120 << 8, 120 << 8).unwrap(),
),
),
),
(
DataType::Int32,
Box::new(
DistributionArrayGeneratorProvider::<Int32Type, Uniform<i32>>::new(
Uniform::new(10, 20).unwrap(),
),
),
),
(
DataType::Int32,
Box::new(
DistributionArrayGeneratorProvider::<Int32Type, Uniform<i32>>::new(
Uniform::new(0, 1).unwrap(),
),
),
),
];
for (data_type, array_gen_provider) in bitpacked_test_cases {
let field = Field::new("", data_type.clone(), false);
let test_cases = TestCases::basic().with_structural_encodings();
check_round_trip_encoding_generated(field, array_gen_provider.copy(), test_cases).await;
}
}
}