Skip to main content

hermes_core/query/vector/
dense.rs

1//! Dense vector query for similarity search (ANN)
2
3use crate::dsl::Field;
4use crate::segment::SegmentReader;
5use std::sync::Arc;
6
7use super::VectorResultScorer;
8use super::combiner::MultiValueCombiner;
9use crate::query::traits::{CountFuture, Query, Scorer, ScorerFuture};
10
11/// Maximum number of IVF clusters a single dense query may probe.
12///
13/// This guard bounds explicit overrides independently of the trained global
14/// leaf count. Automatic billion-scale fields normally use far fewer probes.
15pub const MAX_DENSE_NPROBE: usize = 65_536;
16
17/// Maximum exact-rerank candidate multiplier accepted by dense search.
18pub const MAX_DENSE_RERANK_FACTOR: f32 = crate::query::MAX_CANDIDATE_OVERSUBSCRIPTION as f32;
19
20/// Default exact-rerank candidate multiplier for dense search.
21pub const DEFAULT_DENSE_RERANK_FACTOR: f32 = MAX_DENSE_RERANK_FACTOR;
22
23/// Dense vector query for similarity search
24#[derive(Debug, Clone)]
25pub struct DenseVectorQuery {
26    /// Field containing the dense vectors
27    pub field: Field,
28    /// Query vector
29    pub vector: Vec<f32>,
30    /// Number of clusters to probe (for IVF indexes)
31    pub nprobe: usize,
32    /// Re-ranking factor multiplied by k for candidate selection (1x to 2x)
33    pub rerank_factor: f32,
34    /// How to combine scores for multi-valued documents
35    pub combiner: MultiValueCombiner,
36    /// Query-global dense plans (IVF-PQ probe route / TQ LUTs), shared by all
37    /// segment scorers spawned for this query. The caches are versioned, so a
38    /// query reused after an index generation change recomputes safely.
39    plan_cache: Arc<crate::segment::DensePlanCache>,
40    /// Shared copy of `vector` for the async per-segment scorer futures,
41    /// which cannot borrow `self`. Built once; revalidated against `vector`
42    /// (which is `pub`) before reuse.
43    shared_vector: std::sync::OnceLock<Arc<[f32]>>,
44}
45
46impl std::fmt::Display for DenseVectorQuery {
47    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48        write!(
49            f,
50            "Dense({}, dim={}, nprobe={}, rerank={})",
51            self.field.0,
52            self.vector.len(),
53            self.nprobe,
54            self.rerank_factor
55        )
56    }
57}
58
59impl DenseVectorQuery {
60    /// Create a new dense vector query
61    pub fn new(field: Field, vector: Vec<f32>) -> Self {
62        Self {
63            field,
64            vector,
65            nprobe: 64,
66            rerank_factor: DEFAULT_DENSE_RERANK_FACTOR,
67            combiner: MultiValueCombiner::Max,
68            plan_cache: Arc::new(Default::default()),
69            shared_vector: std::sync::OnceLock::new(),
70        }
71    }
72
73    /// The query vector as a shared slice for scorer futures. One allocation
74    /// per query instead of one per segment; falls back to a fresh copy if
75    /// the public `vector` was edited after the first scorer was built.
76    fn shared_vector(&self) -> Arc<[f32]> {
77        let shared = self
78            .shared_vector
79            .get_or_init(|| Arc::from(self.vector.as_slice()));
80        if shared.as_ref() == self.vector.as_slice() {
81            Arc::clone(shared)
82        } else {
83            Arc::from(self.vector.as_slice())
84        }
85    }
86
87    /// Set the number of clusters to probe (for IVF indexes)
88    ///
89    /// Values are validated when the query is executed. See
90    /// [`MAX_DENSE_NPROBE`].
91    pub fn with_nprobe(mut self, nprobe: usize) -> Self {
92        self.nprobe = nprobe;
93        self
94    }
95
96    /// Set the re-ranking factor (e.g. 2.0 = fetch 2x candidates for reranking)
97    ///
98    /// Values are validated when the query is executed. See
99    /// [`MAX_DENSE_RERANK_FACTOR`].
100    pub fn with_rerank_factor(mut self, factor: f32) -> Self {
101        self.rerank_factor = factor;
102        self
103    }
104
105    /// Set the multi-value score combiner
106    pub fn with_combiner(mut self, combiner: MultiValueCombiner) -> Self {
107        self.combiner = combiner;
108        self
109    }
110}
111
112impl Query for DenseVectorQuery {
113    fn scorer<'a>(&self, reader: &'a SegmentReader, limit: usize) -> ScorerFuture<'a> {
114        let field = self.field;
115        let vector = self.shared_vector();
116        let nprobe = self.nprobe;
117        let rerank_factor = self.rerank_factor;
118        let combiner = self.combiner;
119        let plan_cache = Arc::clone(&self.plan_cache);
120        Box::pin(async move {
121            let results = reader
122                .search_dense_vector_with_probe_cache(
123                    field,
124                    &vector,
125                    limit,
126                    nprobe,
127                    rerank_factor,
128                    combiner,
129                    &plan_cache,
130                )
131                .await?;
132
133            Ok(Box::new(VectorResultScorer::new(results, field.0)) as Box<dyn Scorer>)
134        })
135    }
136
137    #[cfg(feature = "sync")]
138    fn scorer_sync<'a>(
139        &self,
140        reader: &'a SegmentReader,
141        limit: usize,
142    ) -> crate::Result<Box<dyn Scorer + 'a>> {
143        let results = reader.search_dense_vector_sync_with_probe_cache(
144            self.field,
145            &self.vector,
146            limit,
147            self.nprobe,
148            self.rerank_factor,
149            self.combiner,
150            &self.plan_cache,
151        )?;
152        Ok(Box::new(VectorResultScorer::new(results, self.field.0)) as Box<dyn Scorer>)
153    }
154
155    fn count_estimate<'a>(&self, _reader: &'a SegmentReader) -> CountFuture<'a> {
156        Box::pin(async move { Ok(u32::MAX) })
157    }
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163
164    #[test]
165    fn test_dense_vector_query_builder() {
166        let query = DenseVectorQuery::new(Field(0), vec![1.0, 2.0, 3.0]).with_nprobe(64);
167
168        assert_eq!(query.field, Field(0));
169        assert_eq!(query.vector.len(), 3);
170        assert_eq!(query.nprobe, 64);
171        assert_eq!(query.rerank_factor, 2.0);
172    }
173}