#[cfg(feature = "decoder")]
use crate::decoder_buffer::DecoderBuffer;
#[cfg(feature = "decoder")]
use crate::direct_bit_decoder::DirectBitDecoder;
#[cfg(feature = "encoder")]
use crate::direct_bit_encoder::DirectBitEncoder;
#[cfg(feature = "encoder")]
use crate::encoder_buffer::EncoderBuffer;
#[cfg(feature = "decoder")]
use crate::folded_bit32_coder::FoldedBit32Decoder;
#[cfg(feature = "encoder")]
use crate::folded_bit32_coder::FoldedBit32Encoder;
#[cfg(feature = "decoder")]
use crate::rans_bit_decoder::RAnsBitDecoder;
#[cfg(feature = "encoder")]
use crate::rans_bit_encoder::RAnsBitEncoder;
#[cfg(feature = "decoder")]
use crate::status::DracoError;
fn most_significant_bit(value: u32) -> u32 {
debug_assert!(value > 0);
31 - value.leading_zeros()
}
#[cfg(feature = "decoder")]
fn grow_rows(
base: &mut Vec<u32>,
levels: &mut Vec<u32>,
dimension: usize,
rows: usize,
have: usize,
) -> Result<usize, ()> {
let target = rows.max(have.saturating_mul(2));
let needed = target.checked_mul(dimension).ok_or(())?;
for stack in [&mut *base, &mut *levels] {
if stack.len() < needed {
stack.try_reserve(needed - stack.len()).map_err(|_| ())?;
stack.resize(needed, 0);
}
}
Ok(target)
}
#[cfg(feature = "decoder")]
#[inline]
fn copy_row_to_next(stack: &mut [u32], src: usize, dim: usize) {
#[inline(always)]
fn fixed<const N: usize>(stack: &mut [u32], src: usize) {
debug_assert!(stack.len() >= src + 2 * N);
let Some(window) = stack.get_mut(src..src + 2 * N) else {
return;
};
let (row, next) = window.split_at_mut(N);
next.copy_from_slice(row);
}
match dim {
1 => fixed::<1>(stack, src),
2 => fixed::<2>(stack, src),
3 => fixed::<3>(stack, src),
4 => fixed::<4>(stack, src),
5 => fixed::<5>(stack, src),
6 => fixed::<6>(stack, src),
7 => fixed::<7>(stack, src),
8 => fixed::<8>(stack, src),
9 => fixed::<9>(stack, src),
10 => fixed::<10>(stack, src),
11 => fixed::<11>(stack, src),
12 => fixed::<12>(stack, src),
_ => stack.copy_within(src..src + dim, src + dim),
}
}
fn increment_mod(v: u32, m: u32) -> u32 {
let next = v + 1;
if next >= m {
0
} else {
next
}
}
#[derive(Clone)]
pub struct PointDVector {
data: Vec<u32>,
num_points: usize,
dimension: usize,
}
impl PointDVector {
pub fn new(num_points: usize, dimension: usize) -> Self {
Self {
data: vec![0; num_points * dimension],
num_points,
dimension,
}
}
pub fn num_points(&self) -> usize {
self.num_points
}
pub fn dimension(&self) -> usize {
self.dimension
}
pub fn point(&self, index: usize) -> &[u32] {
let start = index * self.dimension;
&self.data[start..start + self.dimension]
}
pub fn point_mut(&mut self, index: usize) -> &mut [u32] {
let start = index * self.dimension;
&mut self.data[start..start + self.dimension]
}
pub fn as_slice(&self) -> &[u32] {
&self.data
}
pub fn as_mut_slice(&mut self) -> &mut [u32] {
&mut self.data
}
pub fn swap_points(&mut self, a: usize, b: usize) {
if a == b {
return;
}
let dim = self.dimension;
let (lo, hi) = if a < b { (a, b) } else { (b, a) };
let (head, tail) = self.data.split_at_mut(hi * dim);
head[lo * dim..][..dim].swap_with_slice(&mut tail[..dim]);
}
pub fn partition(&mut self, begin: usize, end: usize, axis: usize, value: u32) -> usize {
let stride = self.dimension;
let mut first = begin;
let mut last = end;
loop {
loop {
if first == last {
return first;
}
if self.data[first * stride + axis] >= value {
break;
}
first += 1;
}
loop {
last -= 1;
if first == last {
return first;
}
if self.data[last * stride + axis] < value {
break;
}
}
self.swap_points(first, last);
first += 1;
}
}
}
#[cfg(feature = "encoder")]
enum NumbersEncoder {
Direct(DirectBitEncoder),
RAns(RAnsBitEncoder),
Folded(FoldedBit32Encoder),
}
#[cfg(feature = "encoder")]
impl NumbersEncoder {
fn start_encoding(&mut self) {
match self {
NumbersEncoder::Direct(e) => e.start_encoding(),
NumbersEncoder::RAns(e) => e.start_encoding(),
NumbersEncoder::Folded(e) => e.start_encoding(),
}
}
fn encode_least_significant_bits32(&mut self, nbits: u32, value: u32) {
match self {
NumbersEncoder::Direct(e) => e.encode_least_significant_bits32(nbits, value),
NumbersEncoder::RAns(e) => e.encode_least_significant_bits32(nbits, value),
NumbersEncoder::Folded(e) => e.encode_least_significant_bits32(nbits, value),
}
}
fn end_encoding(&mut self, target_buffer: &mut EncoderBuffer) {
match self {
NumbersEncoder::Direct(e) => e.end_encoding(target_buffer),
NumbersEncoder::RAns(e) => e.end_encoding(target_buffer),
NumbersEncoder::Folded(e) => e.end_encoding(target_buffer),
}
}
}
#[cfg(feature = "encoder")]
pub struct DynamicIntegerPointsKdTreeEncoder {
compression_level: u8,
bit_length: u32,
dimension: u32,
deviations: Vec<u32>,
num_remaining_bits: Vec<u32>,
axes: Vec<u32>,
base_stack: Vec<u32>,
levels_stack: Vec<u32>,
numbers_encoder: NumbersEncoder,
remaining_bits_encoder: DirectBitEncoder,
axis_encoder: DirectBitEncoder,
half_encoder: DirectBitEncoder,
}
#[cfg(feature = "encoder")]
impl DynamicIntegerPointsKdTreeEncoder {
pub fn new(compression_level: u8, dimension: u32) -> Self {
assert!(compression_level <= 6);
let stack_len = (32 * dimension + 1) as usize;
let numbers_encoder = match compression_level {
0 | 1 => NumbersEncoder::Direct(DirectBitEncoder::new()),
2 | 3 => NumbersEncoder::RAns(RAnsBitEncoder::new()),
4..=6 => NumbersEncoder::Folded(FoldedBit32Encoder::new()),
_ => unreachable!(),
};
Self {
compression_level,
bit_length: 0,
dimension,
deviations: vec![0; dimension as usize],
num_remaining_bits: vec![0; dimension as usize],
axes: vec![0; dimension as usize],
base_stack: vec![0; stack_len * dimension as usize],
levels_stack: vec![0; stack_len * dimension as usize],
numbers_encoder,
remaining_bits_encoder: DirectBitEncoder::new(),
axis_encoder: DirectBitEncoder::new(),
half_encoder: DirectBitEncoder::new(),
}
}
pub fn encode_points(
&mut self,
points: &mut PointDVector,
bit_length: u32,
buffer: &mut EncoderBuffer,
) {
self.bit_length = bit_length;
buffer.encode_u32(self.bit_length);
buffer.encode_u32(points.num_points() as u32);
if points.num_points() == 0 {
return;
}
self.numbers_encoder.start_encoding();
self.remaining_bits_encoder.start_encoding();
self.axis_encoder.start_encoding();
self.half_encoder.start_encoding();
self.encode_internal(points);
self.numbers_encoder.end_encoding(buffer);
self.remaining_bits_encoder.end_encoding(buffer);
self.axis_encoder.end_encoding(buffer);
self.half_encoder.end_encoding(buffer);
}
fn get_and_encode_axis(
&mut self,
points: &PointDVector,
begin: usize,
end: usize,
old_base: &[u32],
levels: &[u32],
last_axis: u32,
) -> u32 {
if self.compression_level != 6 {
return increment_mod(last_axis, self.dimension);
}
let size = (end - begin) as u32;
debug_assert!(size != 0);
let mut best_axis = 0u32;
if size < 64 {
for axis in 1..self.dimension {
if levels[best_axis as usize] > levels[axis as usize] {
best_axis = axis;
}
}
} else {
for i in 0..self.dimension as usize {
self.deviations[i] = 0;
self.num_remaining_bits[i] = self.bit_length - levels[i];
if self.num_remaining_bits[i] > 0 {
let split = old_base[i] + (1u32 << (self.num_remaining_bits[i] - 1));
let mut cnt = 0u32;
for p in begin..end {
if points.point(p)[i] < split {
cnt += 1;
}
}
let other = size - cnt;
self.deviations[i] = if other > cnt { other } else { cnt };
}
}
let mut max_value = 0u32;
best_axis = 0;
for i in 0..self.dimension as usize {
if self.num_remaining_bits[i] != 0 && self.deviations[i] > max_value {
max_value = self.deviations[i];
best_axis = i as u32;
}
}
self.axis_encoder
.encode_least_significant_bits32(4, best_axis);
}
best_axis
}
fn encode_number(&mut self, nbits: u32, value: u32) {
self.numbers_encoder
.encode_least_significant_bits32(nbits, value);
}
fn encode_internal(&mut self, points: &mut PointDVector) {
#[derive(Clone, Copy)]
struct Status {
begin: usize,
end: usize,
last_axis: u32,
stack_pos: usize,
}
let dimension = self.dimension as usize;
self.base_stack[0..dimension].fill(0);
self.levels_stack[0..dimension].fill(0);
let mut old_base = vec![0; dimension];
let mut levels = vec![0; dimension];
let mut stack: Vec<Status> = Vec::new();
stack.push(Status {
begin: 0,
end: points.num_points(),
last_axis: 0,
stack_pos: 0,
});
while let Some(status) = stack.pop() {
let begin = status.begin;
let end = status.end;
let last_axis = status.last_axis;
let stack_pos = status.stack_pos;
let row_start = stack_pos * dimension;
old_base.copy_from_slice(&self.base_stack[row_start..row_start + dimension]);
levels.copy_from_slice(&self.levels_stack[row_start..row_start + dimension]);
let axis = self.get_and_encode_axis(points, begin, end, &old_base, &levels, last_axis);
let level = levels[axis as usize];
let num_remaining_points = (end - begin) as u32;
if (self.bit_length - level) == 0 {
continue;
}
if num_remaining_points <= 2 {
self.axes[0] = axis;
for i in 1..self.dimension as usize {
self.axes[i] = increment_mod(self.axes[i - 1], self.dimension);
}
for p in begin..end {
let point = points.point(p);
for j in 0..self.dimension as usize {
let num_bits = self.bit_length - levels[self.axes[j] as usize];
if num_bits != 0 {
self.remaining_bits_encoder.encode_least_significant_bits32(
num_bits,
point[self.axes[j] as usize],
);
}
}
}
continue;
}
let num_remaining_bits = self.bit_length - level;
let modifier = 1u32 << (num_remaining_bits - 1);
let child_start = (stack_pos + 1) * dimension;
self.base_stack[child_start..child_start + dimension].copy_from_slice(&old_base);
self.base_stack[child_start + axis as usize] += modifier;
let new_base_axis_value = self.base_stack[child_start + axis as usize];
let split = points.partition(begin, end, axis as usize, new_base_axis_value);
let required_bits = most_significant_bit(num_remaining_points);
let first_half = (split - begin) as u32;
let second_half = (end - split) as u32;
let left = first_half < second_half;
if first_half != second_half {
self.half_encoder.encode_bit(left);
}
if left {
self.encode_number(required_bits, num_remaining_points / 2 - first_half);
} else {
self.encode_number(required_bits, num_remaining_points / 2 - second_half);
}
levels[axis as usize] += 1;
self.levels_stack[row_start..row_start + dimension].copy_from_slice(&levels);
self.levels_stack[child_start..child_start + dimension].copy_from_slice(&levels);
if split != begin {
stack.push(Status {
begin,
end: split,
last_axis: axis,
stack_pos,
});
}
if split != end {
stack.push(Status {
begin: split,
end,
last_axis: axis,
stack_pos: stack_pos + 1,
});
}
}
}
}
#[cfg(feature = "decoder")]
enum NumbersDecoder<'a> {
Direct(DirectBitDecoder),
RAns(RAnsBitDecoder<'a>),
Folded(FoldedBit32Decoder<'a>),
}
#[cfg(feature = "decoder")]
impl<'a> NumbersDecoder<'a> {
fn start_decoding(&mut self, buffer: &mut DecoderBuffer<'a>) -> bool {
match self {
NumbersDecoder::Direct(d) => d.start_decoding(buffer),
NumbersDecoder::RAns(d) => d.start_decoding(buffer),
NumbersDecoder::Folded(d) => d.start_decoding(buffer),
}
}
fn decode_least_significant_bits32(&mut self, nbits: u32, value: &mut u32) -> bool {
match self {
NumbersDecoder::Direct(d) => d.decode_least_significant_bits32(nbits, value),
NumbersDecoder::RAns(d) => d.decode_least_significant_bits32(nbits as i32, value),
NumbersDecoder::Folded(d) => d.decode_least_significant_bits32(nbits, value),
}
}
fn end_decoding(&mut self) {
match self {
NumbersDecoder::Direct(d) => d.end_decoding(),
NumbersDecoder::RAns(d) => d.end_decoding(),
NumbersDecoder::Folded(d) => d.end_decoding(),
}
}
}
#[cfg(feature = "decoder")]
pub struct DynamicIntegerPointsKdTreeDecoder<'a> {
compression_level: u8,
bit_length: u32,
num_points: u32,
num_decoded_points: u32,
dimension: u32,
base_stack: Vec<u32>,
levels_stack: Vec<u32>,
numbers_decoder: NumbersDecoder<'a>,
remaining_bits_decoder: DirectBitDecoder,
axis_decoder: DirectBitDecoder,
half_decoder: DirectBitDecoder,
}
#[cfg(feature = "decoder")]
impl<'a> DynamicIntegerPointsKdTreeDecoder<'a> {
pub fn new(compression_level: u8, dimension: u32) -> Self {
assert!(compression_level <= 6);
let numbers_decoder = match compression_level {
0 | 1 => NumbersDecoder::Direct(DirectBitDecoder::new()),
2 | 3 => NumbersDecoder::RAns(RAnsBitDecoder::new()),
4..=6 => NumbersDecoder::Folded(FoldedBit32Decoder::new()),
_ => unreachable!(),
};
Self {
compression_level,
bit_length: 0,
num_points: 0,
num_decoded_points: 0,
dimension,
base_stack: Vec::new(),
levels_stack: Vec::new(),
numbers_decoder,
remaining_bits_decoder: DirectBitDecoder::new(),
axis_decoder: DirectBitDecoder::new(),
half_decoder: DirectBitDecoder::new(),
}
}
pub fn num_decoded_points(&self) -> u32 {
self.num_decoded_points
}
pub fn decode_points(
&mut self,
buffer: &mut DecoderBuffer<'a>,
oit_max_points: u32,
) -> Result<Vec<u32>, DracoError> {
self.bit_length = buffer
.decode_u32()
.map_err(|_| DracoError::buffer("Buffer ran out reading the KD-tree bit length"))?;
if self.bit_length > 32 {
return Err(DracoError::general(format!(
"KD-tree bit length {} above the 32 a u32 coordinate holds",
self.bit_length
)));
}
self.num_points = buffer
.decode_u32()
.map_err(|_| DracoError::buffer("Buffer ran out reading the KD-tree point count"))?;
if self.num_points == 0 {
self.num_decoded_points = 0;
return Ok(Vec::new());
}
if self.num_points > oit_max_points {
return Err(DracoError::general(format!(
"KD-tree declares {} points against the {oit_max_points} the header allows",
self.num_points
)));
}
self.num_decoded_points = 0;
for (name, started) in [
("numbers", self.numbers_decoder.start_decoding(buffer)),
(
"remaining bits",
self.remaining_bits_decoder.start_decoding(buffer),
),
("axis", self.axis_decoder.start_decoding(buffer)),
("half", self.half_decoder.start_decoding(buffer)),
] {
if !started {
return Err(DracoError::general(format!(
"Failed to start the KD-tree's {name} decoder"
)));
}
}
let out_len = (self.num_points as usize)
.checked_mul(self.dimension as usize)
.ok_or_else(|| {
DracoError::general("KD-tree point count times dimension overflows a usize")
})?;
let mut out: Vec<u32> = Vec::new();
let reserve = out_len.min(buffer.remaining_size().saturating_mul(8));
out.try_reserve(reserve)
.map_err(|_| DracoError::allocation_exceeds_input(reserve * 4, buffer.size()))?;
if !self.decode_internal(self.num_points, &mut out) {
return Err(DracoError::general(format!(
"KD-tree traversal failed after {} of {} points",
self.num_decoded_points, self.num_points
)));
}
self.numbers_decoder.end_decoding();
self.remaining_bits_decoder.end_decoding();
self.axis_decoder.end_decoding();
self.half_decoder.end_decoding();
Ok(out)
}
fn get_axis(
&mut self,
num_remaining_points: u32,
levels: &[u32],
last_axis: u32,
) -> Option<u32> {
if self.compression_level != 6 {
return Some(increment_mod(last_axis, self.dimension));
}
let best_axis = if num_remaining_points < 64 {
levels
.iter()
.enumerate()
.min_by_key(|&(_, level)| *level)
.map_or(0, |(axis, _)| axis as u32)
} else {
let mut v = 0u32;
if !self.axis_decoder.decode_least_significant_bits32(4, &mut v) {
return None;
}
v
};
Some(best_axis)
}
fn decode_number(&mut self, nbits: u32, value: &mut u32) -> bool {
self.numbers_decoder
.decode_least_significant_bits32(nbits, value)
}
fn decode_internal(&mut self, num_points: u32, out: &mut Vec<u32>) -> bool {
let mut base_stack = std::mem::take(&mut self.base_stack);
let mut levels_stack = std::mem::take(&mut self.levels_stack);
base_stack.clear();
levels_stack.clear();
let ok = self.decode_walk(num_points, out, &mut base_stack, &mut levels_stack);
self.base_stack = base_stack;
self.levels_stack = levels_stack;
ok
}
fn decode_walk(
&mut self,
num_points: u32,
out: &mut Vec<u32>,
base_stack: &mut Vec<u32>,
levels_stack: &mut Vec<u32>,
) -> bool {
#[derive(Clone, Copy)]
struct Status {
num_remaining_points: u32,
last_axis: u32,
stack_pos: usize,
}
let dimension = self.dimension as usize;
let Ok(mut rows) = grow_rows(base_stack, levels_stack, dimension, 8, 0) else {
return false;
};
base_stack[0..dimension].fill(0);
levels_stack[0..dimension].fill(0);
let mut stack: Vec<Status> = Vec::new();
stack.push(Status {
num_remaining_points: num_points,
last_axis: 0,
stack_pos: 0,
});
while let Some(status) = stack.pop() {
let num_remaining_points = status.num_remaining_points;
let last_axis = status.last_axis;
let stack_pos = status.stack_pos;
let row_start = stack_pos * dimension;
let row_end = row_start + dimension;
let child_start = row_end;
if base_stack.len() < row_end || levels_stack.len() < row_end {
return false;
}
if num_remaining_points > num_points {
return false;
}
let Some(axis) = self.get_axis(
num_remaining_points,
&levels_stack[row_start..row_end],
last_axis,
) else {
return false;
};
if axis >= self.dimension {
return false;
}
let axis = axis as usize;
let level = levels_stack[row_start + axis];
if (self.bit_length - level) == 0 {
for _ in 0..num_remaining_points {
out.extend_from_slice(&base_stack[row_start..row_end]);
self.num_decoded_points += 1;
}
continue;
}
if num_remaining_points <= 2 {
let old_base = &base_stack[row_start..row_end];
let levels = &levels_stack[row_start..row_end];
for _ in 0..num_remaining_points {
let start = out.len();
out.resize(start + dimension, 0);
let p = &mut out[start..];
let mut axis_j = axis;
for _ in 0..dimension {
let num_bits = self.bit_length - levels[axis_j];
let mut value = 0u32;
if num_bits != 0 {
let ok = self
.remaining_bits_decoder
.decode_least_significant_bits32(num_bits, &mut value);
if !ok {
return false;
}
}
p[axis_j] = value | old_base[axis_j];
axis_j = increment_mod(axis_j as u32, self.dimension) as usize;
}
self.num_decoded_points += 1;
}
continue;
}
if self.num_decoded_points > self.num_points {
return false;
}
if stack_pos + 2 > rows {
let Ok(grown) = grow_rows(base_stack, levels_stack, dimension, stack_pos + 2, rows)
else {
return false;
};
rows = grown;
}
let num_remaining_bits = self.bit_length - level;
let modifier = 1u32 << (num_remaining_bits - 1);
copy_row_to_next(base_stack, row_start, dimension);
base_stack[child_start + axis] += modifier;
let incoming_bits = most_significant_bit(num_remaining_points);
let mut number = 0u32;
if !self.decode_number(incoming_bits, &mut number) {
return false;
}
let mut first_half = num_remaining_points / 2;
if first_half < number {
return false;
}
first_half -= number;
let mut second_half = num_remaining_points - first_half;
if first_half != second_half {
let Some(keep_order) = self.half_decoder.decode_next_bit() else {
return false;
};
if !keep_order {
std::mem::swap(&mut first_half, &mut second_half);
}
}
levels_stack[row_start + axis] += 1;
copy_row_to_next(levels_stack, row_start, dimension);
if first_half != 0 {
stack.push(Status {
num_remaining_points: first_half,
last_axis: axis as u32,
stack_pos,
});
}
if second_half != 0 {
stack.push(Status {
num_remaining_points: second_half,
last_axis: axis as u32,
stack_pos: stack_pos + 1,
});
}
}
true
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn get_axis_rejects_truncated_axis_stream() {
let mut decoder = DynamicIntegerPointsKdTreeDecoder::new(6, 3);
let levels = [0, 0, 0];
assert_eq!(decoder.get_axis(64, &levels, 0), None);
}
}