weavatrix_search_vector/quantized/
scalar.rs1use super::storage::quantize_component;
2use crate::config::{DistanceMetric, IndexConfig};
3use crate::error::SearchError;
4use crate::hit::SearchHit;
5use crate::hnsw::VectorIndex;
6use crate::parallel;
7use crate::vector::{Candidate, inverse_norm};
8use std::collections::BinaryHeap;
9
10#[derive(Debug)]
16pub struct ScalarQuantizedIndex {
17 config: IndexConfig,
18 keys: Vec<u64>,
19 values: Vec<i8>,
20 inverse_norms: Vec<f32>,
21}
22
23impl ScalarQuantizedIndex {
24 pub fn build(config: IndexConfig, vectors: &[(u64, &[f32])]) -> Result<Self, SearchError> {
31 config.validate()?;
32 if config.metric != DistanceMetric::Cosine {
33 return Err(SearchError::InvalidConfig(
34 "scalar-int8 quantization currently requires cosine distance",
35 ));
36 }
37 let elements = config
38 .dimensions
39 .checked_mul(vectors.len())
40 .ok_or(SearchError::CapacityOverflow)?;
41 let mut order = (0..vectors.len()).collect::<Vec<_>>();
42 order.sort_unstable_by_key(|index| vectors[*index].0);
43 for pair in order.windows(2) {
44 if vectors[pair[0]].0 == vectors[pair[1]].0 {
45 return Err(SearchError::DuplicateKey(vectors[pair[0]].0));
46 }
47 }
48 let mut keys = Vec::new();
49 keys.try_reserve_exact(vectors.len())
50 .map_err(|_| SearchError::AllocationFailed)?;
51 let mut values = Vec::new();
52 values
53 .try_reserve_exact(elements)
54 .map_err(|_| SearchError::AllocationFailed)?;
55 let mut inverse_norms = Vec::new();
56 inverse_norms
57 .try_reserve_exact(vectors.len())
58 .map_err(|_| SearchError::AllocationFailed)?;
59 for source in order {
60 let (key, vector) = vectors[source];
61 if vector.len() != config.dimensions {
62 return Err(SearchError::DimensionMismatch {
63 expected: config.dimensions,
64 actual: vector.len(),
65 vector: Some(source),
66 });
67 }
68 let inverse = inverse_norm(vector, Some(source))?;
69 let start = values.len();
70 values.extend(
71 vector
72 .iter()
73 .map(|value| quantize_component(value * inverse)),
74 );
75 let integer_norm = values[start..]
76 .iter()
77 .map(|value| {
78 let value = f32::from(*value);
79 value * value
80 })
81 .sum::<f32>()
82 .sqrt();
83 if integer_norm == 0.0 {
84 return Err(SearchError::ZeroVector {
85 vector: Some(source),
86 });
87 }
88 keys.push(key);
89 inverse_norms.push(integer_norm.recip());
90 }
91 Ok(Self {
92 config,
93 keys,
94 values,
95 inverse_norms,
96 })
97 }
98
99 #[must_use]
100 pub fn len(&self) -> usize {
101 self.keys.len()
102 }
103
104 #[must_use]
105 pub fn is_empty(&self) -> bool {
106 self.keys.is_empty()
107 }
108
109 #[must_use]
110 pub const fn dimensions(&self) -> usize {
111 self.config.dimensions
112 }
113
114 #[must_use]
115 pub const fn config(&self) -> &IndexConfig {
116 &self.config
117 }
118
119 #[must_use]
121 pub fn estimated_memory_bytes(&self) -> usize {
122 self.keys
123 .capacity()
124 .saturating_mul(std::mem::size_of::<u64>())
125 .saturating_add(
126 self.values
127 .capacity()
128 .saturating_mul(std::mem::size_of::<i8>()),
129 )
130 .saturating_add(
131 self.inverse_norms
132 .capacity()
133 .saturating_mul(std::mem::size_of::<f32>()),
134 )
135 }
136
137 #[allow(clippy::cast_precision_loss)]
143 pub fn search(&self, query: &[f32], count: usize) -> Result<Vec<SearchHit>, SearchError> {
144 let (quantized_query, query_inverse_norm) = self.quantize_query(query)?;
145 let limit = count.min(self.len());
146 if limit == 0 {
147 return Ok(Vec::new());
148 }
149 let mut best = BinaryHeap::with_capacity(limit);
150 for index in 0..self.len() {
151 let dot = self
152 .vector(index)
153 .iter()
154 .zip(&quantized_query)
155 .map(|(left, right)| i64::from(*left) * i64::from(*right))
156 .sum::<i64>();
157 let similarity = dot as f32 * self.inverse_norms[index] * query_inverse_norm;
158 let candidate = Candidate::new((1.0 - similarity).clamp(0.0, 2.0), index);
159 if best.len() < limit {
160 best.push(candidate);
161 } else if best
162 .peek()
163 .is_some_and(|worst| candidate.cmp(worst).is_lt())
164 {
165 best.pop();
166 best.push(candidate);
167 }
168 }
169 let mut candidates = best.into_vec();
170 candidates.sort_unstable();
171 Ok(candidates
172 .into_iter()
173 .map(|candidate| SearchHit {
174 key: self.keys[candidate.index()],
175 distance: candidate.distance,
176 })
177 .collect())
178 }
179
180 pub fn search_rerank(
187 &self,
188 oracle: &VectorIndex,
189 query: &[f32],
190 count: usize,
191 candidates: usize,
192 ) -> Result<Vec<SearchHit>, SearchError> {
193 if oracle.dimensions() != self.dimensions() {
194 return Err(SearchError::DimensionMismatch {
195 expected: self.dimensions(),
196 actual: oracle.dimensions(),
197 vector: None,
198 });
199 }
200 let query_inverse = inverse_norm(query, None)?;
201 let candidate_count = candidates.max(count).min(self.len());
202 let mut hits = self.search(query, candidate_count)?;
203 for hit in &mut hits {
204 let vector = oracle
205 .vector(hit.key)
206 .ok_or(SearchError::MissingKey(hit.key))?;
207 let dot = vector
208 .iter()
209 .zip(query)
210 .map(|(left, right)| left * right)
211 .sum::<f32>();
212 hit.distance = (1.0 - dot * query_inverse).clamp(0.0, 2.0);
213 }
214 hits.sort_unstable_by(|left, right| {
215 left.distance
216 .total_cmp(&right.distance)
217 .then_with(|| left.key.cmp(&right.key))
218 });
219 hits.truncate(count.min(hits.len()));
220 Ok(hits)
221 }
222
223 pub fn search_batch(
229 &self,
230 queries: &[&[f32]],
231 count: usize,
232 ) -> Result<Vec<Vec<SearchHit>>, SearchError> {
233 parallel::search_batch(queries, self.config.query_threads, |query| {
234 self.search(query, count)
235 })
236 }
237
238 fn vector(&self, index: usize) -> &[i8] {
239 let start = index * self.dimensions();
240 &self.values[start..start + self.dimensions()]
241 }
242
243 fn quantize_query(&self, query: &[f32]) -> Result<(Vec<i8>, f32), SearchError> {
244 if query.len() != self.dimensions() {
245 return Err(SearchError::DimensionMismatch {
246 expected: self.dimensions(),
247 actual: query.len(),
248 vector: None,
249 });
250 }
251 let inverse = inverse_norm(query, None)?;
252 let mut quantized = Vec::new();
253 quantized
254 .try_reserve_exact(query.len())
255 .map_err(|_| SearchError::AllocationFailed)?;
256 quantized.extend(
257 query
258 .iter()
259 .map(|value| quantize_component(value * inverse)),
260 );
261 let norm = quantized
262 .iter()
263 .map(|value| {
264 let value = f32::from(*value);
265 value * value
266 })
267 .sum::<f32>()
268 .sqrt();
269 if norm == 0.0 {
270 return Err(SearchError::ZeroVector { vector: None });
271 }
272 Ok((quantized, norm.recip()))
273 }
274}