use std::alloc::Layout;
#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
use std::arch::aarch64::*;
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
use std::borrow::Cow;
use std::iter::repeat_with;
use std::ops::Range;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use crate::common::counter::hardware_counter::HardwareCounterCell;
use crate::common::fs::atomic_save_json;
use crate::common::mmap::MmapFlusher;
use crate::common::typelevel::True;
use crate::common::types::PointOffsetType;
use fs_err as fs;
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use crate::quantization::encoded_storage::{EncodedStorage, EncodedStorageBuilder};
use crate::quantization::encoded_vectors::{EncodedVectors, VectorParameters, validate_vector_parameters};
use crate::quantization::kmeans::kmeans;
use crate::quantization::{ConditionalVariable, EncodingError};
pub const KMEANS_SAMPLE_SIZE: usize = 10_000;
pub const KMEANS_MAX_ITERATIONS: usize = 100;
pub const KMEANS_ACCURACY: f32 = 1e-5;
pub const CENTROIDS_COUNT: usize = 256;
pub struct EncodedVectorsPQ<TStorage: EncodedStorage> {
encoded_vectors: TStorage,
metadata: Metadata,
metadata_path: Option<PathBuf>,
}
pub struct EncodedQueryPQ {
lut: Vec<f32>,
}
#[derive(Serialize, Deserialize)]
pub struct Metadata {
pub centroids: Vec<Vec<f32>>,
pub vector_division: Vec<Range<usize>>,
pub vector_parameters: VectorParameters,
}
impl<TStorage: EncodedStorage> EncodedVectorsPQ<TStorage> {
pub fn storage(&self) -> &TStorage {
&self.encoded_vectors
}
#[allow(clippy::too_many_arguments)]
pub fn encode<'a>(
data: impl Iterator<Item = impl AsRef<[f32]> + 'a> + Clone + Send,
mut storage_builder: impl EncodedStorageBuilder<Storage = TStorage> + Send,
vector_parameters: &VectorParameters,
count: usize,
chunk_size: usize,
max_kmeans_threads: usize,
meta_path: Option<&Path>,
stopped: &AtomicBool,
) -> Result<Self, EncodingError> {
debug_assert!(validate_vector_parameters(data.clone(), vector_parameters).is_ok());
let vector_division = Self::get_vector_division(vector_parameters.dim, chunk_size);
let centroids = Self::find_centroids(
data.clone(),
&vector_division,
vector_parameters,
count,
CENTROIDS_COUNT,
max_kmeans_threads,
stopped,
)?;
Self::encode_storage(
data,
&mut storage_builder,
&vector_division,
¢roids,
max_kmeans_threads,
stopped,
)?;
let encoded_vectors = storage_builder
.build()
.map_err(|e| EncodingError::EncodingError(format!("Failed to build storage: {e}",)))?;
let metadata = Metadata {
centroids,
vector_division,
vector_parameters: *vector_parameters,
};
if let Some(meta_path) = meta_path {
meta_path
.parent()
.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Path must have a parent directory",
)
})
.and_then(fs::create_dir_all)
.map_err(|e| {
EncodingError::EncodingError(format!(
"Failed to create metadata directory: {e}",
))
})?;
atomic_save_json(meta_path, &metadata).map_err(|e| {
EncodingError::EncodingError(format!("Failed to save metadata: {e}",))
})?;
}
if !stopped.load(Ordering::Relaxed) {
Ok(Self {
encoded_vectors,
metadata,
metadata_path: meta_path.map(PathBuf::from),
})
} else {
Err(EncodingError::Stopped)
}
}
pub fn load(encoded_vectors: TStorage, meta_path: &Path) -> std::io::Result<Self> {
let contents = fs::read_to_string(meta_path)?;
let metadata: Metadata = serde_json::from_str(&contents)?;
let result = Self {
encoded_vectors,
metadata,
metadata_path: Some(meta_path.to_path_buf()),
};
Ok(result)
}
fn get_vector_division(dim: usize, chunk_size: usize) -> Vec<Range<usize>> {
(0..dim)
.step_by(chunk_size)
.map(|i| i..std::cmp::min(i + chunk_size, dim))
.collect()
}
fn encode_storage<'a: 'b, 'b>(
data: impl Iterator<Item = impl AsRef<[f32]> + 'a> + Clone + Send + 'b,
storage_builder: &'b mut (impl EncodedStorageBuilder<Storage = TStorage> + Send),
vector_division: &'b [Range<usize>],
centroids: &'b [Vec<f32>],
max_threads: usize,
stopped: &AtomicBool,
) -> Result<(), EncodingError> {
rayon::ThreadPoolBuilder::new()
.thread_name(|idx| format!("pq-encoding-{idx}"))
.num_threads(std::cmp::max(1, max_threads))
.build()
.map_err(|e| {
EncodingError::EncodingError(format!(
"Failed PQ encoding while thread pool init: {e}"
))
})?
.scope(|s| {
Self::encode_storage_rayon(
s,
data,
storage_builder,
vector_division,
centroids,
max_threads,
stopped,
)
})
}
fn encode_storage_rayon<'a: 'b, 'b>(
scope: &rayon::Scope<'b>,
data: impl Iterator<Item = impl AsRef<[f32]> + 'a> + Clone + Send + 'b,
storage_builder: &'b mut (impl EncodedStorageBuilder<Storage = TStorage> + Send),
vector_division: &'b [Range<usize>],
centroids: &'b [Vec<f32>],
max_threads: usize,
stopped: &'b AtomicBool,
) -> Result<(), EncodingError> {
let storage_builder = Arc::new(Mutex::new(storage_builder));
let mut condvars: Vec<ConditionalVariable> =
repeat_with(Default::default).take(max_threads).collect();
condvars[0].notify();
let error = Arc::new(Mutex::new(None));
for thread_index in 0..max_threads {
let data = data.clone().skip(thread_index);
let storage_builder = storage_builder.clone();
let condvar = condvars[thread_index].clone();
let next_condvar = condvars[(thread_index + 1) % max_threads].clone();
let error = error.clone();
scope.spawn(move |_| {
let mut encoded_vector = Vec::with_capacity(vector_division.len());
for vector in data.step_by(max_threads) {
if stopped.load(Ordering::Relaxed) {
return;
}
Self::encode_vector(
vector.as_ref(),
vector_division,
centroids,
&mut encoded_vector,
);
let is_disconnected = condvar.wait();
let insert_result = storage_builder.lock().push_vector_data(&encoded_vector);
if let Err(e) = insert_result {
let mut error = error.lock();
*error = Some(EncodingError::EncodingError(format!(
"Failed to push encoded vector: {e}",
)));
next_condvar.notify();
return;
}
next_condvar.notify();
if is_disconnected {
return;
}
}
});
}
condvars.clear();
if let Some(error) = error.lock().take() {
Err(error)
} else {
Ok(())
}
}
fn encode_vector(
vector_data: &[f32],
vector_division: &[Range<usize>],
centroids: &[Vec<f32>],
encoded_vector: &mut Vec<u8>,
) {
encoded_vector.clear();
for range in vector_division {
let subvector_data = &vector_data[range.clone()];
let mut min_distance = f32::MAX;
let mut min_centroid_index = 0;
for (centroid_index, centroid) in centroids.iter().enumerate() {
let centroid_data = ¢roid[range.clone()];
let distance = subvector_data
.iter()
.zip(centroid_data)
.map(|(a, b)| (a - b).powi(2))
.sum();
if distance < min_distance {
min_distance = distance;
min_centroid_index = centroid_index;
}
}
encoded_vector.push(min_centroid_index as u8);
}
}
fn find_centroids<'a>(
data: impl Iterator<Item = impl AsRef<[f32]> + 'a> + Clone,
vector_division: &[Range<usize>],
vector_parameters: &VectorParameters,
count: usize,
centroids_count: usize,
max_kmeans_threads: usize,
stopped: &AtomicBool,
) -> Result<Vec<Vec<f32>>, EncodingError> {
let sample_size = KMEANS_SAMPLE_SIZE.min(count);
let mut result = vec![vec![]; centroids_count];
if count <= centroids_count {
for (i, vector_data) in data.into_iter().enumerate() {
result[i] = vector_data.as_ref().to_vec();
}
result[count..centroids_count].fill(vec![0.0; vector_parameters.dim]);
return Ok(result);
}
let permutor = permutation_iterator::Permutor::new(count as u64);
let mut selected_vectors: Vec<usize> =
permutor.map(|i| i as usize).take(sample_size).collect();
if stopped.load(Ordering::Relaxed) {
return Err(EncodingError::Stopped);
}
selected_vectors.sort_unstable();
for range in vector_division.iter() {
let mut data_subset = Vec::with_capacity(sample_size * range.len());
let mut selected_index: usize = 0;
for (vector_index, vector_data) in data.clone().enumerate() {
let vector_data = vector_data.as_ref();
if vector_index == selected_vectors[selected_index] {
data_subset.extend_from_slice(&vector_data[range.clone()]);
selected_index += 1;
if selected_index == sample_size {
break;
}
}
}
let centroids = kmeans(
&data_subset,
centroids_count,
range.len(),
KMEANS_MAX_ITERATIONS,
max_kmeans_threads,
KMEANS_ACCURACY,
stopped,
)?;
for (centroid_index, centroid_data) in centroids.chunks_exact(range.len()).enumerate() {
result[centroid_index].extend_from_slice(centroid_data);
}
}
Ok(result)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn score_point_sse(&self, query: &EncodedQueryPQ, centroids: &[u8]) -> f32 {
unsafe {
let len = centroids.len();
let centroids_count = self.metadata.centroids.len();
let mut centroids = centroids.as_ptr();
let mut lut = query.lut.as_ptr();
let mut sum128: __m128 = _mm_setzero_ps();
for _ in 0..len / 4 {
let buffer = [
*lut.add(*centroids as usize),
*lut.add(centroids_count + *centroids.add(1) as usize),
*lut.add(2 * centroids_count + *centroids.add(2) as usize),
*lut.add(3 * centroids_count + *centroids.add(3) as usize),
];
let c = _mm_loadu_ps(buffer.as_ptr());
sum128 = _mm_add_ps(sum128, c);
centroids = centroids.add(4);
lut = lut.add(4 * centroids_count);
}
let sum64: __m128 = _mm_add_ps(sum128, _mm_movehl_ps(sum128, sum128));
let sum32: __m128 = _mm_add_ss(sum64, _mm_shuffle_ps(sum64, sum64, 0x55));
let mut sum = _mm_cvtss_f32(sum32);
for _ in 0..len % 4 {
sum += *lut.add(*centroids as usize);
centroids = centroids.add(1);
lut = lut.add(centroids_count);
}
sum
}
}
#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
unsafe fn score_point_neon(&self, query: &EncodedQueryPQ, centroids: &[u8]) -> f32 {
unsafe {
let len = centroids.len();
let centroids_count = self.metadata.centroids.len();
let mut centroids = centroids.as_ptr();
let mut lut = query.lut.as_ptr();
let mut sum128 = vdupq_n_f32(0.);
for _ in 0..len / 4 {
let buffer = [
*lut.add(*centroids as usize),
*lut.add(centroids_count + *centroids.add(1) as usize),
*lut.add(2 * centroids_count + *centroids.add(2) as usize),
*lut.add(3 * centroids_count + *centroids.add(3) as usize),
];
let c = vld1q_f32(buffer.as_ptr());
sum128 = vaddq_f32(sum128, c);
centroids = centroids.add(4);
lut = lut.add(4 * centroids_count);
}
let mut sum = vaddvq_f32(sum128);
for _ in 0..len % 4 {
sum += *lut.add(*centroids as usize);
centroids = centroids.add(1);
lut = lut.add(centroids_count);
}
sum
}
}
fn score_point_simple(&self, query: &EncodedQueryPQ, centroids: &[u8]) -> f32 {
let len = centroids.len();
let centroids_count = self.metadata.centroids.len();
let mut centroids = centroids.as_ptr();
let mut lut = query.lut.as_ptr();
(0..len)
.map(|_| unsafe {
let value = *lut.add(*centroids as usize);
centroids = centroids.add(1);
lut = lut.add(centroids_count);
value
})
.sum()
}
pub fn get_quantized_vector(&self, i: PointOffsetType) -> Cow<'_, [u8]> {
self.encoded_vectors.get_vector_data(i)
}
pub fn layout(&self) -> Layout {
Layout::from_size_align(self.metadata.vector_division.len(), align_of::<u8>()).unwrap()
}
pub fn get_metadata(&self) -> &Metadata {
&self.metadata
}
}
impl<TStorage: EncodedStorage> EncodedVectors for EncodedVectorsPQ<TStorage> {
type EncodedQuery = EncodedQueryPQ;
fn is_in_ram_or_mmap() -> bool {
TStorage::is_in_ram_or_mmap()
}
fn is_on_disk(&self) -> bool {
self.encoded_vectors.is_on_disk()
}
fn encode_query(&self, query: &[f32]) -> EncodedQueryPQ {
let lut_capacity = self.metadata.vector_division.len() * self.metadata.centroids.len();
let mut lut = Vec::with_capacity(lut_capacity);
for range in &self.metadata.vector_division {
let subquery = &query[range.clone()];
for i in 0..self.metadata.centroids.len() {
let centroid = &self.metadata.centroids[i];
let subcentroid = ¢roid[range.clone()];
let distance = self
.metadata
.vector_parameters
.distance_type
.distance(subquery, subcentroid);
let distance = if self.metadata.vector_parameters.invert {
-distance
} else {
distance
};
lut.push(distance);
}
}
EncodedQueryPQ { lut }
}
fn iter_batch(
&self,
offsets: &[PointOffsetType],
) -> impl Iterator<Item = (usize, Cow<'_, [u8]>)> {
self.encoded_vectors.iter_batch(offsets)
}
fn score(
&self,
query: &Self::EncodedQuery,
encoded_vector: &[u8],
hw_counter: &HardwareCounterCell,
) -> f32 {
self.score_bytes(True, query, encoded_vector, hw_counter)
}
fn score_point(
&self,
query: &EncodedQueryPQ,
i: PointOffsetType,
hw_counter: &HardwareCounterCell,
) -> f32 {
let centroids = self.encoded_vectors.get_vector_data(i);
self.score_bytes(True, query, ¢roids, hw_counter)
}
fn score_internal(
&self,
i: PointOffsetType,
j: PointOffsetType,
hw_counter: &HardwareCounterCell,
) -> f32 {
let centroids_i = self.encoded_vectors.get_vector_data(i);
let centroids_j = self.encoded_vectors.get_vector_data(j);
hw_counter
.vector_io_read()
.incr_delta(self.metadata.vector_division.len() * 2);
hw_counter.cpu_counter().incr_delta(
centroids_i.as_ref().len()
* self
.metadata
.vector_division
.first()
.map(|i| i.len())
.unwrap_or(1),
);
let distance: f32 = centroids_i
.iter()
.zip(centroids_j.as_ref())
.enumerate()
.map(|(range_index, (&c_i, &c_j))| {
let range = &self.metadata.vector_division[range_index];
let data_i = &self.metadata.centroids[c_i as usize][range.clone()];
let data_j = &self.metadata.centroids[c_j as usize][range.clone()];
self.metadata
.vector_parameters
.distance_type
.distance(data_i, data_j)
})
.sum();
if self.metadata.vector_parameters.invert {
-distance
} else {
distance
}
}
fn quantized_vector_size(&self) -> usize {
self.metadata.vector_division.len()
}
fn encode_internal_vector(&self, _id: PointOffsetType) -> Option<EncodedQueryPQ> {
None
}
fn upsert_vector(
&mut self,
_id: PointOffsetType,
_vector: &[f32],
_hw_counter: &HardwareCounterCell,
) -> std::io::Result<()> {
debug_assert!(false, "PQ does not support upsert_vector",);
Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
"PQ does not support upsert_vector",
))
}
fn vectors_count(&self) -> usize {
self.encoded_vectors.vectors_count()
}
fn flusher(&self) -> MmapFlusher {
self.encoded_vectors.flusher()
}
fn files(&self) -> Vec<PathBuf> {
let mut files = self.encoded_vectors.files();
if let Some(meta_path) = &self.metadata_path {
files.push(meta_path.clone());
}
files
}
fn immutable_files(&self) -> Vec<PathBuf> {
let mut files = self.encoded_vectors.immutable_files();
if let Some(meta_path) = &self.metadata_path {
files.push(meta_path.clone());
}
files
}
fn heap_size_bytes(&self) -> usize {
let storage_heap = self.encoded_vectors.heap_size_bytes();
let centroids_heap: usize = self.metadata.centroids.capacity()
* std::mem::size_of::<Vec<f32>>()
+ self
.metadata
.centroids
.iter()
.map(|c| c.capacity() * std::mem::size_of::<f32>())
.sum::<usize>();
let vector_division_heap =
self.metadata.vector_division.capacity() * std::mem::size_of::<Range<usize>>();
storage_heap + centroids_heap + vector_division_heap
}
type SupportsBytes = True;
fn score_bytes(
&self,
_: Self::SupportsBytes,
query: &Self::EncodedQuery,
bytes: &[u8],
hw_counter: &HardwareCounterCell,
) -> f32 {
hw_counter
.cpu_counter()
.incr_delta(self.metadata.vector_division.len());
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
if is_x86_feature_detected!("sse4.1") {
return unsafe { self.score_point_sse(query, bytes) };
}
#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
if std::arch::is_aarch64_feature_detected!("neon") {
return unsafe { self.score_point_neon(query, bytes) };
}
self.score_point_simple(query, bytes)
}
}
pub fn get_quantized_vector_size(vector_parameters: &VectorParameters, chunk_size: usize) -> usize {
(0..vector_parameters.dim).step_by(chunk_size).count()
}