use crate::index::ann::backend::{AnnBackend, AnnBackendCheckpoint, BackendMetric};
use crate::index::ann::product::ProductQuantizer;
use crate::query::AiExecutionContext;
use crate::rowid::RowId;
use crate::schema::ProductQuantizerOptions;
use crate::Result;
use std::collections::BTreeMap;
#[derive(Clone)]
pub(crate) struct PqBackend {
dim: usize,
num_subvectors: usize,
bits: u8,
rerank_factor: usize,
training: ProductQuantizerOptions,
quantizer: Option<ProductQuantizer>,
codes: BTreeMap<RowId, Vec<u8>>,
pending: BTreeMap<RowId, Vec<f32>>,
}
type FrozenProduct = (ProductQuantizer, BTreeMap<RowId, Vec<u8>>);
impl PqBackend {
pub(crate) fn new(
dim: usize,
num_subvectors: usize,
bits: u8,
options: &ProductQuantizerOptions,
) -> Self {
Self {
dim,
num_subvectors,
bits,
rerank_factor: options.rerank_factor,
training: options.clone(),
quantizer: None,
codes: BTreeMap::new(),
pending: BTreeMap::new(),
}
}
fn freeze_active(&self) -> Option<FrozenProduct> {
self.freeze_active_with_checkpoint(&mut || Ok(()))
.expect("infallible product-training checkpoint")
}
fn freeze_active_with_checkpoint(
&self,
checkpoint: &mut dyn FnMut() -> Result<()>,
) -> Result<Option<FrozenProduct>> {
if self.pending.is_empty() {
return Ok(None);
}
let samples: Vec<&[f32]> = self.pending.values().map(|v| v.as_slice()).collect();
let Some(quantizer) = ProductQuantizer::train_with_checkpoint(
self.dim,
self.num_subvectors,
self.bits,
&samples,
&self.training,
checkpoint,
)?
else {
return Ok(None);
};
let mut codes = BTreeMap::new();
for (index, (row_id, vec)) in self.pending.iter().enumerate() {
if index.is_multiple_of(64) {
checkpoint()?;
}
codes.insert(*row_id, quantizer.encode(vec));
}
Ok(Some((quantizer, codes)))
}
pub(crate) fn from_checkpoint(
dim: usize,
num_subvectors: usize,
bits: u8,
options: &ProductQuantizerOptions,
quantizer: ProductQuantizer,
codes: BTreeMap<RowId, Vec<u8>>,
) -> std::result::Result<Self, String> {
if !quantizer.matches_checkpoint(dim, num_subvectors, bits)
|| codes.values().any(|code| code.len() != num_subvectors)
{
return Err("ANN Product checkpoint contains invalid codebook or codes".into());
}
Ok(Self {
dim,
num_subvectors,
bits,
rerank_factor: options.rerank_factor,
training: options.clone(),
quantizer: Some(quantizer),
codes,
pending: BTreeMap::new(),
})
}
fn k(&self) -> usize {
1usize << self.bits
}
}
impl AnnBackend for PqBackend {
fn metric(&self) -> BackendMetric {
BackendMetric::Cosine
}
fn len(&self) -> usize {
self.codes.len() + self.pending.len()
}
fn is_empty(&self) -> bool {
self.codes.is_empty() && self.pending.is_empty()
}
fn insert_validated(
&mut self,
vec: &[f32],
row_id: RowId,
_checkpoint: &mut dyn FnMut() -> Result<()>,
) -> Result<()> {
if self.quantizer.is_some() {
self.quantizer = None;
self.codes.clear();
}
self.pending.insert(row_id, vec.to_vec());
Ok(())
}
fn finalize(&mut self, checkpoint: &mut dyn FnMut() -> Result<()>) -> Result<()> {
if self.quantizer.is_none() {
if let Some((quantizer, codes)) = self.freeze_active_with_checkpoint(checkpoint)? {
self.quantizer = Some(quantizer);
self.codes = codes;
self.pending.clear();
}
}
checkpoint()
}
fn search(
&self,
query: &[f32],
k: usize,
_ef: usize,
context: Option<&AiExecutionContext>,
) -> Result<Vec<(RowId, f64)>> {
let mut scored: Vec<(f32, RowId)> = Vec::new();
if !self.pending.is_empty() {
for (i, (row_id, vec)) in self.pending.iter().enumerate() {
if let Some(context) = context {
if i.is_multiple_of(64) {
let count = (self.pending.len() - i).min(64);
context.consume(crate::query::work_units(
self.dim.saturating_mul(count),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
}
scored.push((squared_l2(query, vec), *row_id));
}
}
if let Some(quantizer) = &self.quantizer {
if !self.codes.is_empty() {
if let Some(context) = context {
context.consume(crate::query::work_units(
self.dim.saturating_mul(self.k()),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
let table = quantizer.adc_table(query);
let k_codes = self.k();
for (i, (row_id, code)) in self.codes.iter().enumerate() {
if let Some(context) = context {
if i.is_multiple_of(64) {
let count = (self.codes.len() - i).min(64);
context.consume(crate::query::work_units(
self.num_subvectors.saturating_mul(count),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
}
let dist =
ProductQuantizer::adc_distance(&table, code, self.num_subvectors, k_codes);
scored.push((dist, *row_id));
}
}
}
if scored.is_empty() {
return Ok(Vec::new());
}
scored.sort_by(|(da, ra), (db, rb)| da.total_cmp(db).then_with(|| ra.cmp(rb)));
let rerank_set = if self.rerank_factor > 0 {
(k.saturating_mul(self.rerank_factor)).min(scored.len())
} else {
k.min(scored.len())
};
if self.rerank_factor > 0 && rerank_set > k {
if let Some(quantizer) = &self.quantizer {
let mut reranked = Vec::with_capacity(rerank_set);
for (index, (_, row_id)) in scored[..rerank_set].iter().enumerate() {
if let Some(context) = context {
if index.is_multiple_of(64) {
let count = (rerank_set - index).min(64);
context.consume(crate::query::work_units(
self.dim.saturating_mul(count),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
}
let scored_row = if let Some(vec) = self.pending.get(row_id) {
(squared_l2(query, vec), *row_id)
} else if let Some(code) = self.codes.get(row_id) {
let recon = quantizer.reconstruct(code);
(squared_l2(query, &recon), *row_id)
} else {
(f32::INFINITY, *row_id)
};
reranked.push(scored_row);
}
reranked.sort_by(|(da, ra), (db, rb)| da.total_cmp(db).then_with(|| ra.cmp(rb)));
return Ok(reranked
.into_iter()
.take(k)
.map(|(dist, row_id)| (row_id, f64::from(dist)))
.collect());
}
}
Ok(scored
.into_iter()
.take(k)
.map(|(dist, row_id)| (row_id, f64::from(dist)))
.collect())
}
fn entries(&self) -> Vec<(Vec<u8>, RowId)> {
let mut out: Vec<(Vec<u8>, RowId)> = Vec::new();
for (row_id, vec) in &self.pending {
let mut bytes = Vec::with_capacity(8 + vec.len() * 4);
bytes.extend_from_slice(&row_id.0.to_le_bytes());
for value in vec {
bytes.extend_from_slice(&value.to_le_bytes());
}
out.push((bytes, *row_id));
}
for (row_id, code) in &self.codes {
if let Some(quantizer) = &self.quantizer {
let recon = quantizer.reconstruct(code);
let mut bytes = Vec::with_capacity(8 + recon.len() * 4);
bytes.extend_from_slice(&row_id.0.to_le_bytes());
for value in &recon {
bytes.extend_from_slice(&value.to_le_bytes());
}
out.push((bytes, *row_id));
}
}
out
}
fn freeze(&self) -> AnnBackendCheckpoint {
let (quantizer, codes) = if let Some(quantizer) = &self.quantizer {
(quantizer.clone(), self.codes.clone())
} else if let Some((quantizer, codes)) = self.freeze_active() {
(quantizer, codes)
} else {
let zero = vec![0.0f32; self.dim];
let quantizer = ProductQuantizer::train(
self.dim,
self.num_subvectors,
self.bits,
&[zero.as_slice()],
&self.training,
)
.expect("non-empty training set");
(quantizer, BTreeMap::new())
};
AnnBackendCheckpoint::Product {
dim: self.dim,
num_subvectors: self.num_subvectors,
bits: self.bits,
rerank_factor: self.rerank_factor,
quantizer,
codes,
}
}
fn empty_active(&self) -> Box<dyn AnnBackend> {
Box::new(Self::new(
self.dim,
self.num_subvectors,
self.bits,
&self.training,
))
}
fn rebuild_from_entries(&self, entries: &[(Vec<u8>, RowId)]) -> Box<dyn AnnBackend> {
let mut pending = BTreeMap::new();
for (bytes, _) in entries {
if bytes.len() < 8 + self.dim * 4 {
continue;
}
let mut rid_bytes = [0u8; 8];
rid_bytes.copy_from_slice(&bytes[..8]);
let row_id = RowId(u64::from_le_bytes(rid_bytes));
let vec = (0..self.dim)
.map(|i| {
let offset = 8 + i * 4;
f32::from_le_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
])
})
.collect();
pending.insert(row_id, vec);
}
let mut rebuilt = Self {
dim: self.dim,
num_subvectors: self.num_subvectors,
bits: self.bits,
rerank_factor: self.rerank_factor,
training: self.training.clone(),
quantizer: None,
codes: BTreeMap::new(),
pending,
};
if let Some((quantizer, codes)) = rebuilt.freeze_active() {
rebuilt.quantizer = Some(quantizer);
rebuilt.codes = codes;
rebuilt.pending.clear();
}
Box::new(rebuilt)
}
fn clone_box(&self) -> Box<dyn AnnBackend> {
Box::new(self.clone())
}
}
fn squared_l2(a: &[f32], b: &[f32]) -> f32 {
let mut sum = 0.0f32;
for (x, y) in a.iter().zip(b.iter()) {
let d = x - y;
sum += d * d;
}
sum
}