use std::{mem::align_of, slice};
use crate::{
error::{Error, Result},
key_composer::{decode_oppv_u64, encode_oppv_u64_fixed},
};
#[derive(Debug, Clone, PartialEq)]
pub struct Sq8Vector {
pub scale: f32,
pub offset: f32,
pub data: Vec<i8>,
}
impl Sq8Vector {
#[inline]
pub fn encode_into(v: &[f64], out_data: &mut Vec<i8>) -> (f32, f32) {
if v.is_empty() {
out_data.clear();
return (1.0, 0.0);
}
let mut min = f64::INFINITY;
let mut max = f64::NEG_INFINITY;
for &x in v {
if x < min {
min = x;
}
if x > max {
max = x;
}
}
let range = (max - min).max(1e-9);
let scale = (range / 255.0) as f32;
let offset = min as f32;
let inv_scale = 255.0 / range;
out_data.clear();
out_data.reserve(v.len());
out_data.extend(v.iter().map(|&x| {
let normalized = (x - min) * inv_scale - 128.0;
normalized.round().clamp(-128.0, 127.0) as i8
}));
(scale, offset)
}
#[inline]
pub fn encode(v: &[f64]) -> Self {
let mut data = Vec::with_capacity(v.len());
let (scale, offset) = Self::encode_into(v, &mut data);
Self {
scale,
offset,
data,
}
}
#[inline]
pub fn decode(&self) -> Vec<f64> {
let mut out = Vec::with_capacity(self.data.len());
self.decode_into(&mut out);
out
}
#[inline]
pub fn decode_into(&self, out: &mut Vec<f64>) {
out.clear();
out.reserve(self.data.len());
let scale = self.scale as f64;
let offset = self.offset as f64;
out.extend(self.data.iter().map(|&q| {
let normalized = (q as f64) + 128.0;
offset + normalized * scale
}));
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum NodePackFormat {
RawF64 = 0x00,
Sq8 = 0x01,
}
#[derive(Debug, Clone)]
pub struct NodePackRef<'a> {
pub format: NodePackFormat,
pub raw_f64_vector: Option<&'a [f64]>,
pub sq8_scale: f32,
pub sq8_offset: f32,
pub sq8_vector: Option<&'a [i8]>,
pub raw_neighbors: &'a [u8],
pub degree: usize,
}
impl<'a> NodePackRef<'a> {
#[inline]
pub fn decode(payload: &'a [u8], dim: usize) -> Result<Self> {
if payload.is_empty() {
return Err(Error::invalid_data("empty node pack payload"));
}
if payload[0] == NodePackFormat::Sq8 as u8 {
let header_len = 1 + 8 + dim + 2;
if payload.len() < header_len {
return Err(Error::invalid_data(
"invalid sq8 node pack payload: too short",
));
}
let scale = f32::from_be_bytes([payload[1], payload[2], payload[3], payload[4]]);
let offset = f32::from_be_bytes([payload[5], payload[6], payload[7], payload[8]]);
let vec_slice = &payload[9..9 + dim];
let sq8_vector: &'a [i8] =
unsafe { slice::from_raw_parts(vec_slice.as_ptr().cast::<i8>(), dim) };
let deg_bytes = &payload[9 + dim..11 + dim];
let degree = u16::from_be_bytes([deg_bytes[0], deg_bytes[1]]) as usize;
let raw_neighbors = &payload[11 + dim..];
return Ok(Self {
format: NodePackFormat::Sq8,
raw_f64_vector: None,
sq8_scale: scale,
sq8_offset: offset,
sq8_vector: Some(sq8_vector),
raw_neighbors,
degree,
});
}
if payload[0] == NodePackFormat::RawF64 as u8 {
let vec_bytes_len = dim * 8;
if payload.len() < 1 + vec_bytes_len + 2 {
return Err(Error::invalid_data(
"invalid raw f64 node pack payload: too short",
));
}
let (vec_bytes, rest) = payload[1..].split_at(vec_bytes_len);
if !(vec_bytes.as_ptr() as usize).is_multiple_of(align_of::<f64>()) {
return Err(Error::invalid_data(
"invalid node pack payload: unaligned f64 vector pointer",
));
}
let vector: &'a [f64] =
unsafe { slice::from_raw_parts(vec_bytes.as_ptr().cast::<f64>(), dim) };
let degree = u16::from_be_bytes([rest[0], rest[1]]) as usize;
let raw_neighbors = &rest[2..];
return Ok(Self {
format: NodePackFormat::RawF64,
raw_f64_vector: Some(vector),
sq8_scale: 1.0,
sq8_offset: 0.0,
sq8_vector: None,
raw_neighbors,
degree,
});
}
let vec_bytes_len = dim * 8;
if payload.len() >= vec_bytes_len + 2 {
let (vec_bytes, rest) = payload.split_at(vec_bytes_len);
if (vec_bytes.as_ptr() as usize).is_multiple_of(align_of::<f64>()) {
let vector: &'a [f64] =
unsafe { slice::from_raw_parts(vec_bytes.as_ptr().cast::<f64>(), dim) };
let degree = u16::from_be_bytes([rest[0], rest[1]]) as usize;
let raw_neighbors = &rest[2..];
return Ok(Self {
format: NodePackFormat::RawF64,
raw_f64_vector: Some(vector),
sq8_scale: 1.0,
sq8_offset: 0.0,
sq8_vector: None,
raw_neighbors,
degree,
});
}
}
Err(Error::invalid_data("unsupported node pack format"))
}
#[inline]
pub fn to_f64_vec(&self) -> Vec<f64> {
if let Some(v) = self.raw_f64_vector {
v.to_vec()
} else if let Some(q) = self.sq8_vector {
let scale = self.sq8_scale as f64;
let offset = self.sq8_offset as f64;
q.iter()
.map(|&val| {
let normalized = (val as f64) + 128.0;
offset + normalized * scale
})
.collect()
} else {
Vec::new()
}
}
#[inline]
pub fn iter_neighbors(&self) -> OppvDeltaNeighborIter<'a> {
OppvDeltaNeighborIter {
rem: self.raw_neighbors,
remaining_count: self.degree,
prev: 0,
is_first: true,
}
}
#[inline]
pub fn to_neighbor_vec(&self) -> Vec<u64> {
self.iter_neighbors().collect()
}
#[inline]
pub fn collect_neighbors_into(&self, out: &mut Vec<u64>) {
out.clear();
out.extend(self.iter_neighbors());
}
#[inline]
pub fn encode_sq8(scale: f32, offset: f32, vector: &[i8], neighbors: &[u64], out: &mut Vec<u8>) {
let dim = vector.len();
out.clear();
out.reserve(1 + 8 + dim + 2 + neighbors.len() * 2);
out.push(NodePackFormat::Sq8 as u8);
out.extend_from_slice(&scale.to_be_bytes());
out.extend_from_slice(&offset.to_be_bytes());
let slice_u8 = unsafe { slice::from_raw_parts(vector.as_ptr().cast::<u8>(), dim) };
out.extend_from_slice(slice_u8);
let deg = neighbors.len().min(u16::MAX as usize) as u16;
out.extend_from_slice(°.to_be_bytes());
if deg == 0 {
return;
}
let valid_neighbors = &neighbors[..deg as usize];
let is_sorted = valid_neighbors.windows(2).all(|w| w[0] <= w[1]);
let mut fixed = [0u8; 9];
let mut prev = 0u64;
let mut sorted_buf;
let slice: &[u64] = if is_sorted {
valid_neighbors
} else {
sorted_buf = valid_neighbors.to_vec();
sorted_buf.sort_unstable();
&sorted_buf
};
for (i, &n) in slice.iter().enumerate() {
let delta = if i == 0 { n } else { n.saturating_sub(prev) };
prev = n;
let len = encode_oppv_u64_fixed(delta, &mut fixed);
out.extend_from_slice(&fixed[..len]);
}
}
#[inline]
pub fn encode(vector: &[f64], neighbors: &[u64], out: &mut Vec<u8>) {
let sq8 = Sq8Vector::encode(vector);
Self::encode_sq8(sq8.scale, sq8.offset, &sq8.data, neighbors, out);
}
}
#[derive(Debug, Clone)]
pub struct OppvDeltaNeighborIter<'a> {
rem: &'a [u8],
remaining_count: usize,
prev: u64,
is_first: bool,
}
impl<'a> Iterator for OppvDeltaNeighborIter<'a> {
type Item = u64;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
if self.remaining_count == 0 || self.rem.is_empty() {
return None;
}
let (delta, len) = match decode_oppv_u64(self.rem) {
Some(res) => res,
None => {
self.remaining_count = 0;
return None;
}
};
self.rem = &self.rem[len..];
let val = if self.is_first {
self.is_first = false;
delta
} else {
self.prev.saturating_add(delta)
};
self.prev = val;
self.remaining_count -= 1;
Some(val)
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
(self.remaining_count, Some(self.remaining_count))
}
}
impl<'a> ExactSizeIterator for OppvDeltaNeighborIter<'a> {}