use crate::types::{DistanceType, VectorMetadata};
use crate::{RemDbError, Result};
use alloc::vec::Vec;
use core::cmp::Ordering;
#[repr(C)]
pub struct IVFCluster {
pub centroid: Vec<f32>,
pub vector_count: u32,
pub vector_offsets: Vec<usize>,
pub record_ids: Vec<u16>,
}
impl IVFCluster {
pub fn new(dimension: u16) -> Self {
let centroid = vec![0.0; dimension as usize];
IVFCluster {
centroid,
vector_count: 0,
vector_offsets: Vec::new(),
record_ids: Vec::new(),
}
}
pub fn add_vector(&mut self, vector_offset: usize, record_id: u16) -> Result<()> {
self.vector_offsets.push(vector_offset);
self.record_ids.push(record_id);
self.vector_count += 1;
Ok(())
}
}
pub struct IVFIndex {
pub meta: VectorMetadata,
pub vectors: *mut f32,
pub clusters: Vec<IVFCluster>,
pub nlist: u32,
pub nprobe: u32,
pub dimension: u16,
pub vector_count: u32,
pub lock: u32,
}
impl IVFIndex {
pub unsafe fn new(
meta: VectorMetadata,
vectors: *mut f32,
nlist: u32,
nprobe: u32,
) -> Result<Self> {
let dimension = meta.dimension;
let mut clusters = Vec::with_capacity(nlist as usize);
for _ in 0..nlist {
clusters.push(IVFCluster::new(dimension));
}
Ok(IVFIndex {
meta,
vectors,
clusters,
nlist,
nprobe,
dimension,
vector_count: 0,
lock: 0,
})
}
unsafe fn calculate_distance(&self, vec1: *const f32, vec2: *const f32) -> f32 {
match self.meta.distance_type {
DistanceType::L2 => {
let mut sum = 0.0;
for i in 0..self.dimension {
let diff = *vec1.add(i as usize) - *vec2.add(i as usize);
sum += diff * diff;
}
sum.sqrt()
}
DistanceType::InnerProduct => {
let mut sum = 0.0;
for i in 0..self.dimension {
sum += *vec1.add(i as usize) * *vec2.add(i as usize);
}
-sum }
DistanceType::Cosine => {
let mut dot = 0.0;
let mut norm1 = 0.0;
let mut norm2 = 0.0;
for i in 0..self.dimension {
let v1 = *vec1.add(i as usize);
let v2 = *vec2.add(i as usize);
dot += v1 * v2;
norm1 += v1 * v1;
norm2 += v2 * v2;
}
let norm1 = norm1.sqrt();
let norm2 = norm2.sqrt();
if norm1 == 0.0 || norm2 == 0.0 {
-1.0 } else {
-(dot / (norm1 * norm2)) }
}
}
}
#[cfg(feature = "std")]
pub fn save<W: std::io::Write>(&self, writer: &mut W) -> Result<()> {
writer
.write_all(&self.meta.ivf_nlist.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.meta.ivf_nprobe.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.dimension.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.nlist.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.nprobe.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.vector_count.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
for cluster in &self.clusters {
writer
.write_all(&cluster.centroid.len().to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
for &value in &cluster.centroid {
writer
.write_all(&value.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
}
writer
.write_all(&cluster.vector_count.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&cluster.vector_offsets.len().to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
for &offset in &cluster.vector_offsets {
writer
.write_all(&offset.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
}
writer
.write_all(&cluster.record_ids.len().to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
for &record_id in &cluster.record_ids {
writer
.write_all(&record_id.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
}
}
Ok(())
}
#[cfg(feature = "std")]
pub unsafe fn load<R: std::io::Read>(
mut meta: VectorMetadata,
vectors: *mut f32,
reader: &mut R,
) -> Result<Self> {
#[cfg(not(feature = "std"))]
let _ = reader;
let mut nlist_bytes = [0u8; 4];
reader
.read_exact(&mut nlist_bytes)
.map_err(|_| RemDbError::FileIoError)?;
meta.ivf_nlist = u32::from_le_bytes(nlist_bytes);
let mut nprobe_bytes = [0u8; 4];
reader
.read_exact(&mut nprobe_bytes)
.map_err(|_| RemDbError::FileIoError)?;
meta.ivf_nprobe = u32::from_le_bytes(nprobe_bytes);
let mut dimension_bytes = [0u8; 2];
reader
.read_exact(&mut dimension_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let dimension = u16::from_le_bytes(dimension_bytes);
let mut nlist_bytes = [0u8; 4];
reader
.read_exact(&mut nlist_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let nlist = u32::from_le_bytes(nlist_bytes);
let mut nprobe_bytes = [0u8; 4];
reader
.read_exact(&mut nprobe_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let nprobe = u32::from_le_bytes(nprobe_bytes);
let mut vector_count_bytes = [0u8; 8];
reader
.read_exact(&mut vector_count_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let vector_count = u64::from_le_bytes(vector_count_bytes);
let mut clusters = Vec::with_capacity(nlist as usize);
for _ in 0..nlist {
let mut centroid_len_bytes = [0u8; 4];
reader
.read_exact(&mut centroid_len_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let centroid_len = u32::from_le_bytes(centroid_len_bytes) as usize;
let mut centroid = Vec::with_capacity(centroid_len);
for _ in 0..centroid_len {
let mut value_bytes = [0u8; 4];
reader
.read_exact(&mut value_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let value = f32::from_le_bytes(value_bytes);
centroid.push(value);
}
let mut vec_count_bytes = [0u8; 4];
reader
.read_exact(&mut vec_count_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let vec_count = u32::from_le_bytes(vec_count_bytes);
let mut offsets_len_bytes = [0u8; 8];
reader
.read_exact(&mut offsets_len_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let offsets_len = usize::from_le_bytes(offsets_len_bytes);
let mut vector_offsets = Vec::with_capacity(offsets_len);
for _ in 0..offsets_len {
let mut offset_bytes = [0u8; 8];
reader
.read_exact(&mut offset_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let offset = usize::from_le_bytes(offset_bytes);
vector_offsets.push(offset);
}
let mut record_ids_len_bytes = [0u8; 8];
reader
.read_exact(&mut record_ids_len_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let record_ids_len = usize::from_le_bytes(record_ids_len_bytes);
let mut record_ids = Vec::with_capacity(record_ids_len);
for _ in 0..record_ids_len {
let mut record_id_bytes = [0u8; 2];
reader
.read_exact(&mut record_id_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let record_id = u16::from_le_bytes(record_id_bytes);
record_ids.push(record_id);
}
let cluster = IVFCluster {
centroid,
vector_count: vec_count,
vector_offsets,
record_ids,
};
clusters.push(cluster);
}
Ok(IVFIndex {
meta,
vectors,
clusters,
nlist,
nprobe,
dimension,
vector_count: vector_count as u32,
lock: 0,
})
}
unsafe fn find_closest_cluster(&self, vec_ptr: *const f32) -> usize {
let mut closest_cluster = 0;
let mut min_distance = f32::MAX;
for (i, cluster) in self.clusters.iter().enumerate() {
let cluster_ptr = cluster.centroid.as_ptr();
let distance = self.calculate_distance(vec_ptr, cluster_ptr);
if distance < min_distance {
min_distance = distance;
closest_cluster = i;
}
}
closest_cluster
}
pub unsafe fn train(
&mut self,
vectors: &[*const f32],
vector_offsets: &[usize],
record_ids: &[u16],
max_iter: u32,
) -> Result<()> {
let dimension = self.dimension as usize;
let nlist = self.nlist as usize;
if vectors.len() != vector_offsets.len() || vectors.len() != record_ids.len() {
return Err(RemDbError::InternalError);
}
for i in 0..nlist {
if i < vectors.len() {
let vec_ptr = vectors[i];
for j in 0..dimension {
self.clusters[i].centroid[j] = *vec_ptr.add(j);
}
}
}
for _ in 0..max_iter {
for cluster in &mut self.clusters {
cluster.vector_count = 0;
cluster.vector_offsets.clear();
cluster.record_ids.clear();
}
for (i, &vec_ptr) in vectors.iter().enumerate() {
let cluster_idx = self.find_closest_cluster(vec_ptr);
let cluster = &mut self.clusters[cluster_idx];
cluster.vector_offsets.push(vector_offsets[i]);
cluster.record_ids.push(record_ids[i]);
cluster.vector_count += 1;
}
for (_i, cluster) in self.clusters.iter_mut().enumerate() {
if cluster.vector_count > 0 {
let mut new_centroid = vec![0.0; dimension];
for &offset in &cluster.vector_offsets {
let vec_ptr = self.vectors.add(offset);
for j in 0..dimension {
new_centroid[j] += *vec_ptr.add(j);
}
}
let count = cluster.vector_count as f32;
for j in 0..dimension {
new_centroid[j] /= count;
}
let mut converged = true;
for j in 0..dimension {
if (new_centroid[j] - cluster.centroid[j]).abs() > 1e-6 {
converged = false;
break;
}
}
if converged {
break;
}
cluster.centroid = new_centroid;
}
}
}
Ok(())
}
pub unsafe fn search(&self, query_vec: *const f32, k: usize) -> Result<Vec<(f32, u16)>> {
let mut cluster_distances = Vec::new();
for (i, cluster) in self.clusters.iter().enumerate() {
let centroid_ptr = cluster.centroid.as_ptr();
let distance = self.calculate_distance(query_vec, centroid_ptr);
cluster_distances.push((distance, i));
}
cluster_distances.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(Ordering::Equal));
let nprobe = self.nprobe as usize;
let selected_clusters = cluster_distances.iter().take(nprobe).map(|&(_d, i)| i);
let mut results = Vec::new();
for cluster_idx in selected_clusters {
let cluster = &self.clusters[cluster_idx];
for (i, &vector_offset) in cluster.vector_offsets.iter().enumerate() {
let vec_ptr = self.vectors.add(vector_offset);
let distance = self.calculate_distance(query_vec, vec_ptr);
let record_id = cluster.record_ids[i];
results.push((distance, record_id));
}
}
results.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(Ordering::Equal));
let final_results = results.iter().take(k).cloned().collect();
Ok(final_results)
}
pub unsafe fn insert(&mut self, vector_offset: usize, record_id: u16) -> Result<()> {
let vec_ptr = self.vectors.add(vector_offset);
let cluster_idx = self.find_closest_cluster(vec_ptr);
self.clusters[cluster_idx].add_vector(vector_offset, record_id)?;
self.vector_count += 1;
Ok(())
}
pub unsafe fn delete(&mut self, vector_offset: usize) -> Result<()> {
let vec_ptr = self.vectors.add(vector_offset);
let cluster_idx = self.find_closest_cluster(vec_ptr);
let cluster = &mut self.clusters[cluster_idx];
if let Some(pos) = cluster
.vector_offsets
.iter()
.position(|&offset| offset == vector_offset)
{
cluster.vector_offsets.remove(pos);
cluster.record_ids.remove(pos);
cluster.vector_count -= 1;
self.vector_count -= 1;
Ok(())
} else {
Err(RemDbError::RecordNotFound)
}
}
}