use std::collections::HashMap;
use std::sync::Arc;
use crate::error::Result;
use crate::vector::core::distance::DistanceMetric;
use crate::vector::core::distance_quantized::{QuantizedQuery, distance_quantized};
use crate::vector::core::vector::Vector;
use crate::vector::index::quantized_storage::QuantizedVectorPool;
use crate::vector::index::rerank_storage::RerankStoragePool;
#[derive(Debug, Default)]
pub(crate) struct RerankCandidates {
pub doc_ids: Vec<u64>,
pub distances: Vec<f32>,
}
impl RerankCandidates {
pub fn with_capacity(capacity: usize) -> Self {
Self {
doc_ids: Vec::with_capacity(capacity),
distances: Vec::with_capacity(capacity),
}
}
#[inline]
pub fn push(&mut self, doc_id: u64, distance: f32) {
self.doc_ids.push(doc_id);
self.distances.push(distance);
}
#[inline]
pub fn len(&self) -> usize {
self.doc_ids.len()
}
pub fn sort_by_distance(&mut self) {
let mut order: Vec<u32> = (0..self.len() as u32).collect();
order.sort_unstable_by(|&a, &b| {
self.distances[a as usize]
.total_cmp(&self.distances[b as usize])
.then(self.doc_ids[a as usize].cmp(&self.doc_ids[b as usize]))
});
self.doc_ids = order.iter().map(|&i| self.doc_ids[i as usize]).collect();
self.distances = order.iter().map(|&i| self.distances[i as usize]).collect();
}
}
pub(crate) trait RerankStage: Send + Sync {
fn rescore(
&self,
query: &Vector,
candidates: &mut RerankCandidates,
take_n: usize,
) -> Result<bool>;
}
pub(crate) struct F32SidecarStage {
pool: Arc<RerankStoragePool>,
positions: Option<Arc<HashMap<u64, u32>>>,
metric: DistanceMetric,
}
impl F32SidecarStage {
pub fn new(pool: Arc<RerankStoragePool>, field_name: &str, metric: DistanceMetric) -> Self {
let positions = pool.field_position_index(field_name);
Self {
pool,
positions,
metric,
}
}
}
impl RerankStage for F32SidecarStage {
fn rescore(
&self,
query: &Vector,
candidates: &mut RerankCandidates,
take_n: usize,
) -> Result<bool> {
let prepared = self.metric.prepare_query(&query.data);
let take_n = take_n.min(candidates.len());
let mut rescored = RerankCandidates::with_capacity(take_n);
for i in 0..take_n {
let doc_id = candidates.doc_ids[i];
let Some(pos) = self
.positions
.as_ref()
.and_then(|idx| idx.get(&doc_id).copied())
else {
continue;
};
let distance = self
.metric
.distance_with_prepared(&prepared, self.pool.f32_slice_at(pos))?;
rescored.push(doc_id, distance);
}
rescored.sort_by_distance();
*candidates = rescored;
Ok(true)
}
}
pub(crate) struct SqRerankStage {
pool: Arc<QuantizedVectorPool>,
positions: Option<Arc<HashMap<u64, u32>>>,
metric: DistanceMetric,
}
impl SqRerankStage {
pub fn new(pool: Arc<QuantizedVectorPool>, field_name: &str, metric: DistanceMetric) -> Self {
let positions = pool.field_position_index(field_name);
Self {
pool,
positions,
metric,
}
}
}
impl RerankStage for SqRerankStage {
fn rescore(
&self,
query: &Vector,
candidates: &mut RerankCandidates,
take_n: usize,
) -> Result<bool> {
let prepared = QuantizedQuery::prepare(&query.data, &self.pool.params);
let take_n = take_n.min(candidates.len());
let mut rescored = RerankCandidates::with_capacity(take_n);
for i in 0..take_n {
let doc_id = candidates.doc_ids[i];
let Some(pos) = self
.positions
.as_ref()
.and_then(|idx| idx.get(&doc_id).copied())
else {
continue;
};
let (cand, meta) = self.pool.record_at(pos);
let distance = distance_quantized(self.metric, &prepared, cand, meta);
rescored.push(doc_id, distance);
}
rescored.sort_by_distance();
*candidates = rescored;
Ok(true)
}
}
pub(crate) struct RerankPipeline {
stages: Vec<Box<dyn RerankStage>>,
factors: Vec<usize>,
}
impl RerankPipeline {
pub fn new(stages: Vec<Box<dyn RerankStage>>, factors: Vec<usize>) -> Self {
debug_assert_eq!(stages.len(), factors.len());
debug_assert!(!stages.is_empty());
Self { stages, factors }
}
pub fn run(
&self,
query: &Vector,
candidates: &mut RerankCandidates,
top_k: usize,
) -> Result<bool> {
let mut last_applied = false;
for (stage, &factor) in self.stages.iter().zip(&self.factors) {
let take_n = top_k.saturating_mul(factor.max(1));
last_applied = stage.rescore(query, candidates, take_n)?;
}
Ok(last_applied)
}
}
#[cfg(test)]
mod tests {
use super::*;
struct HalveStage;
impl RerankStage for HalveStage {
fn rescore(
&self,
_query: &Vector,
candidates: &mut RerankCandidates,
take_n: usize,
) -> Result<bool> {
let take_n = take_n.min(candidates.len());
let mut out = RerankCandidates::with_capacity(take_n);
for i in 0..take_n {
out.push(candidates.doc_ids[i], candidates.distances[i] / 2.0);
}
out.sort_by_distance();
*candidates = out;
Ok(true)
}
}
struct AbsentStage;
impl RerankStage for AbsentStage {
fn rescore(
&self,
_query: &Vector,
_candidates: &mut RerankCandidates,
_take_n: usize,
) -> Result<bool> {
Ok(false)
}
}
fn buffer(pairs: &[(u64, f32)]) -> RerankCandidates {
let mut c = RerankCandidates::with_capacity(pairs.len());
for &(id, d) in pairs {
c.push(id, d);
}
c
}
#[test]
fn sort_by_distance_breaks_ties_by_doc_id() {
let mut c = buffer(&[(9, 1.0), (2, 0.5), (7, 0.5)]);
c.sort_by_distance();
assert_eq!(c.doc_ids, vec![2, 7, 9]);
assert_eq!(c.distances, vec![0.5, 0.5, 1.0]);
}
#[test]
fn pipeline_narrows_by_per_stage_factor_and_reports_last_stage() {
let pipeline =
RerankPipeline::new(vec![Box::new(HalveStage), Box::new(HalveStage)], vec![3, 1]);
let mut c = buffer(&[(1, 1.0), (2, 2.0), (3, 3.0), (4, 4.0), (5, 5.0), (6, 6.0)]);
let applied = pipeline.run(&Vector::new(vec![0.0]), &mut c, 2).unwrap();
assert!(applied);
assert_eq!(c.doc_ids, vec![1, 2]);
assert_eq!(c.distances, vec![0.25, 0.5]);
}
#[test]
fn absent_final_stage_reports_not_applied_and_preserves_buffer() {
let pipeline = RerankPipeline::new(vec![Box::new(AbsentStage)], vec![4]);
let mut c = buffer(&[(1, 1.0), (2, 2.0)]);
let applied = pipeline.run(&Vector::new(vec![0.0]), &mut c, 1).unwrap();
assert!(!applied);
assert_eq!(c.doc_ids, vec![1, 2]);
assert_eq!(c.distances, vec![1.0, 2.0]);
}
use crate::vector::core::quantization::{QuantizedVectorMeta, ScalarQuantParams};
fn int8_pool_near_and_far() -> Arc<QuantizedVectorPool> {
let params =
ScalarQuantParams::train(&[Vector::new(vec![0.0, 0.0]), Vector::new(vec![10.0, 10.0])])
.unwrap();
let mut records = Vec::new();
for (doc_id, data) in [(1u64, vec![0.0_f32, 0.0]), (2, vec![10.0, 10.0])] {
let q = params.quantize_slice(&data);
let meta = QuantizedVectorMeta::from_quantized(&q, ¶ms);
records.push((doc_id, "f".to_string(), q, meta));
}
Arc::new(QuantizedVectorPool::build(params, 2, records))
}
#[test]
fn sq_stage_rescores_and_reorders_by_int8_distance() {
let stage = SqRerankStage::new(int8_pool_near_and_far(), "f", DistanceMetric::Euclidean);
let mut c = buffer(&[(2, 0.0), (1, 100.0)]);
let applied = stage
.rescore(&Vector::new(vec![0.0, 0.0]), &mut c, 2)
.unwrap();
assert!(applied);
assert_eq!(c.doc_ids, vec![1, 2], "doc 1 (near query) must rank first");
}
#[test]
fn sq_stage_drops_candidates_absent_from_the_pool() {
let stage = SqRerankStage::new(int8_pool_near_and_far(), "f", DistanceMetric::Euclidean);
let mut c = buffer(&[(1, 0.0), (99, 0.0)]); let applied = stage
.rescore(&Vector::new(vec![0.0, 0.0]), &mut c, 2)
.unwrap();
assert!(applied);
assert_eq!(c.doc_ids, vec![1]);
}
}