1use super::vector_codec::{encode_fp16, f64_to_f32, fp16_to_f32, validate_vector};
4use super::{type_error, Doc, VectorValue};
5use crate::error::Result;
6use crate::types::{DataType, MetricType};
7use std::collections::BTreeMap;
8
9impl VectorValue {
10 pub fn encode_fp16(values: &[f32]) -> Result<Self> {
12 Ok(Self::Fp16(encode_fp16(values)?))
13 }
14
15 pub fn data_type(&self) -> DataType {
16 match self {
17 Self::Binary32(_) => DataType::VectorBinary32,
18 Self::Binary64(_) => DataType::VectorBinary64,
19 Self::Fp16(_) => DataType::VectorFp16,
20 Self::Fp32(_) => DataType::VectorFp32,
21 Self::Fp64(_) => DataType::VectorFp64,
22 Self::Int4(_) => DataType::VectorInt4,
23 Self::Int8(_) => DataType::VectorInt8,
24 Self::Int16(_) => DataType::VectorInt16,
25 Self::SparseFp16 { .. } => DataType::SparseVectorFp16,
26 Self::SparseFp32 { .. } => DataType::SparseVectorFp32,
27 }
28 }
29
30 #[allow(clippy::match_same_arms)]
31 pub fn dimension(&self) -> usize {
32 match self {
33 Self::Binary32(v) | Self::Binary64(v) => v.len().saturating_mul(8),
34 Self::Fp16(v) => v.len(),
35 Self::Fp32(v) => v.len(),
36 Self::Fp64(v) => v.len(),
37 Self::Int4(v) | Self::Int8(v) => v.len(),
38 Self::Int16(v) => v.len(),
39 Self::SparseFp16 { indices, .. } | Self::SparseFp32 { indices, .. } => indices
40 .iter()
41 .max()
42 .map_or(0, |v| (*v as usize).saturating_add(1)),
43 }
44 }
45
46 pub fn is_sparse(&self) -> bool {
47 matches!(self, Self::SparseFp16 { .. } | Self::SparseFp32 { .. })
48 }
49
50 pub fn to_dense_f32(&self) -> Option<Vec<f32>> {
55 match self {
56 Self::Fp16(values) => Some(values.iter().map(|v| fp16_to_f32(*v)).collect()),
57 Self::Fp32(values) => Some(values.clone()),
58 Self::Fp64(values) => values.iter().copied().map(f64_to_f32).collect(),
59 Self::Int4(values) | Self::Int8(values) => {
60 Some(values.iter().map(|v| f32::from(*v)).collect())
61 }
62 Self::Int16(values) => Some(values.iter().map(|v| f32::from(*v)).collect()),
63 _ => None,
64 }
65 }
66
67 pub fn to_dense_f64(&self) -> Option<Vec<f64>> {
69 match self {
70 Self::Fp16(values) => Some(
71 values
72 .iter()
73 .map(|value| f64::from(fp16_to_f32(*value)))
74 .collect(),
75 ),
76 Self::Fp32(values) => Some(values.iter().map(|value| f64::from(*value)).collect()),
77 Self::Fp64(values) => Some(values.clone()),
78 Self::Int4(values) | Self::Int8(values) => {
79 Some(values.iter().map(|value| f64::from(*value)).collect())
80 }
81 Self::Int16(values) => Some(values.iter().map(|value| f64::from(*value)).collect()),
82 _ => None,
83 }
84 }
85
86 pub(crate) fn dense_score(
93 &self,
94 query: &[f64],
95 query_norm: f64,
96 metric: MetricType,
97 ) -> Option<f64> {
98 let dimension = match self {
99 Self::Fp16(values) => values.len(),
100 Self::Fp32(values) => values.len(),
101 Self::Fp64(values) => values.len(),
102 Self::Int4(values) | Self::Int8(values) => values.len(),
103 Self::Int16(values) => values.len(),
104 Self::Binary32(_)
105 | Self::Binary64(_)
106 | Self::SparseFp16 { .. }
107 | Self::SparseFp32 { .. } => return None,
108 };
109 if query.len() != dimension {
110 return None;
111 }
112 match self {
113 Self::Fp16(values) => Some(score_dense_iter(
114 query,
115 query_norm,
116 values.len(),
117 values.iter().map(|value| f64::from(fp16_to_f32(*value))),
118 metric,
119 )),
120 Self::Fp32(values) => Some(crate::score_f64::score_f64_f32(
121 query, values, metric, query_norm,
122 )),
123 Self::Fp64(values) => Some(score_dense_iter(
124 query,
125 query_norm,
126 values.len(),
127 values.iter().copied(),
128 metric,
129 )),
130 Self::Int4(values) | Self::Int8(values) => Some(score_dense_iter(
131 query,
132 query_norm,
133 values.len(),
134 values.iter().map(|value| f64::from(*value)),
135 metric,
136 )),
137 Self::Int16(values) => Some(score_dense_iter(
138 query,
139 query_norm,
140 values.len(),
141 values.iter().map(|value| f64::from(*value)),
142 metric,
143 )),
144 Self::Binary32(_)
145 | Self::Binary64(_)
146 | Self::SparseFp16 { .. }
147 | Self::SparseFp32 { .. } => None,
148 }
149 }
150
151 pub fn to_sparse_f64(&self) -> Option<BTreeMap<u32, f64>> {
152 match self {
153 Self::SparseFp16 { indices, values } => {
154 if indices.len() != values.len() {
155 return None;
156 }
157 Some(
158 indices
159 .iter()
160 .copied()
161 .zip(values.iter().map(|value| f64::from(fp16_to_f32(*value))))
162 .collect(),
163 )
164 }
165 Self::SparseFp32 { indices, values } => {
166 if indices.len() != values.len() {
167 return None;
168 }
169 Some(
170 indices
171 .iter()
172 .copied()
173 .zip(values.iter().map(|value| f64::from(*value)))
174 .collect(),
175 )
176 }
177 _ => None,
178 }
179 }
180
181 pub(crate) fn validate(&self) -> Result<()> {
182 validate_vector(self)
183 }
184}
185
186fn score_dense_iter(
187 query: &[f64],
188 query_norm: f64,
189 dimension: usize,
190 values: impl Iterator<Item = f64>,
191 metric: MetricType,
192) -> f64 {
193 debug_assert_eq!(query.len(), dimension);
194 match metric {
195 MetricType::L2 => -query
196 .iter()
197 .copied()
198 .zip(values)
199 .map(|(left, right)| {
200 let difference = left - right;
201 difference * difference
202 })
203 .sum::<f64>(),
204 MetricType::Cosine => {
205 let (dot, value_norm) = query
206 .iter()
207 .copied()
208 .zip(values)
209 .fold((0.0, 0.0), |(dot, value_norm), (left, right)| {
210 (dot + left * right, value_norm + right * right)
211 });
212 if query_norm == 0.0 || value_norm == 0.0 {
213 0.0
214 } else {
215 dot / (query_norm * value_norm.sqrt())
216 }
217 }
218 MetricType::MipsL2 | MetricType::Ip | MetricType::Undefined => query
219 .iter()
220 .copied()
221 .zip(values)
222 .map(|(left, right)| left * right)
223 .sum::<f64>(),
224 }
225}
226
227impl Doc {
228 pub fn add_vector_f32(&mut self, name: &str, vector: &[f32]) -> Result<()> {
229 self.set_vector_value(name, VectorValue::Fp32(vector.to_vec()))
230 }
231
232 pub fn add_vector_f64(&mut self, name: &str, vector: &[f64]) -> Result<()> {
233 self.set_vector_value(name, VectorValue::Fp64(vector.to_vec()))
234 }
235
236 pub fn add_vector_i8(&mut self, name: &str, vector: &[i8]) -> Result<()> {
237 self.set_vector_value(name, VectorValue::Int8(vector.to_vec()))
238 }
239
240 pub fn add_vector_i16(&mut self, name: &str, vector: &[i16]) -> Result<()> {
241 self.set_vector_value(name, VectorValue::Int16(vector.to_vec()))
242 }
243
244 pub fn add_vector_fp16(&mut self, name: &str, vector: &[u16]) -> Result<()> {
245 self.set_vector_value(name, VectorValue::Fp16(vector.to_vec()))
246 }
247
248 pub fn add_vector_fp16_f32(&mut self, name: &str, vector: &[f32]) -> Result<()> {
249 self.set_vector_value(name, VectorValue::encode_fp16(vector)?)
250 }
251
252 pub fn add_vector_i4(&mut self, name: &str, vector: &[i8]) -> Result<()> {
253 self.set_vector_value(name, VectorValue::Int4(vector.to_vec()))
254 }
255
256 pub fn add_vector_binary32(&mut self, name: &str, vector: &[u8]) -> Result<()> {
257 self.set_vector_value(name, VectorValue::Binary32(vector.to_vec()))
258 }
259
260 pub fn add_vector_binary64(&mut self, name: &str, vector: &[u8]) -> Result<()> {
261 self.set_vector_value(name, VectorValue::Binary64(vector.to_vec()))
262 }
263
264 pub fn add_sparse_vector(&mut self, name: &str, indices: &[u32], values: &[f32]) -> Result<()> {
265 self.set_vector_value(
266 name,
267 VectorValue::SparseFp32 {
268 indices: indices.to_vec(),
269 values: values.to_vec(),
270 },
271 )
272 }
273
274 pub fn add_sparse_vector_f32(
275 &mut self,
276 name: &str,
277 indices: &[u32],
278 values: &[f32],
279 ) -> Result<()> {
280 self.add_sparse_vector(name, indices, values)
281 }
282
283 pub fn add_sparse_vector_fp16(
284 &mut self,
285 name: &str,
286 indices: &[u32],
287 values: &[u16],
288 ) -> Result<()> {
289 self.set_vector_value(
290 name,
291 VectorValue::SparseFp16 {
292 indices: indices.to_vec(),
293 values: values.to_vec(),
294 },
295 )
296 }
297
298 pub fn add_sparse_vector_fp16_f32(
299 &mut self,
300 name: &str,
301 indices: &[u32],
302 values: &[f32],
303 ) -> Result<()> {
304 self.set_vector_value(
305 name,
306 VectorValue::SparseFp16 {
307 indices: indices.to_vec(),
308 values: encode_fp16(values)?,
309 },
310 )
311 }
312
313 pub fn get_vector_f32(&self, name: &str) -> Result<Option<Vec<f32>>> {
314 match self.vectors.get(name) {
315 None => Ok(None),
316 Some(VectorValue::Fp32(values)) => Ok(Some(values.clone())),
317 Some(_) => Err(type_error(name, DataType::VectorFp32)),
318 }
319 }
320
321 pub fn get_vector_f64(&self, name: &str) -> Result<Option<Vec<f64>>> {
322 match self.vectors.get(name) {
323 None => Ok(None),
324 Some(VectorValue::Fp64(values)) => Ok(Some(values.clone())),
325 Some(_) => Err(type_error(name, DataType::VectorFp64)),
326 }
327 }
328
329 pub fn get_vector_fp16(&self, name: &str) -> Result<Option<Vec<u16>>> {
330 match self.vectors.get(name) {
331 None => Ok(None),
332 Some(VectorValue::Fp16(values)) => Ok(Some(values.clone())),
333 Some(_) => Err(type_error(name, DataType::VectorFp16)),
334 }
335 }
336
337 pub fn get_vector_i4(&self, name: &str) -> Result<Option<Vec<i8>>> {
338 match self.vectors.get(name) {
339 None => Ok(None),
340 Some(VectorValue::Int4(values)) => Ok(Some(values.clone())),
341 Some(_) => Err(type_error(name, DataType::VectorInt4)),
342 }
343 }
344
345 pub fn get_vector_i8(&self, name: &str) -> Result<Option<Vec<i8>>> {
346 match self.vectors.get(name) {
347 None => Ok(None),
348 Some(VectorValue::Int8(values)) => Ok(Some(values.clone())),
349 Some(_) => Err(type_error(name, DataType::VectorInt8)),
350 }
351 }
352
353 pub fn get_vector_i16(&self, name: &str) -> Result<Option<Vec<i16>>> {
354 match self.vectors.get(name) {
355 None => Ok(None),
356 Some(VectorValue::Int16(values)) => Ok(Some(values.clone())),
357 Some(_) => Err(type_error(name, DataType::VectorInt16)),
358 }
359 }
360
361 pub fn get_vector_binary32(&self, name: &str) -> Result<Option<Vec<u8>>> {
362 match self.vectors.get(name) {
363 None => Ok(None),
364 Some(VectorValue::Binary32(values)) => Ok(Some(values.clone())),
365 Some(_) => Err(type_error(name, DataType::VectorBinary32)),
366 }
367 }
368
369 pub fn get_vector_binary64(&self, name: &str) -> Result<Option<Vec<u8>>> {
370 match self.vectors.get(name) {
371 None => Ok(None),
372 Some(VectorValue::Binary64(values)) => Ok(Some(values.clone())),
373 Some(_) => Err(type_error(name, DataType::VectorBinary64)),
374 }
375 }
376
377 pub fn get_sparse_vector_f32(&self, name: &str) -> Result<Option<(Vec<u32>, Vec<f32>)>> {
378 match self.vectors.get(name) {
379 None => Ok(None),
380 Some(VectorValue::SparseFp32 { indices, values }) => {
381 Ok(Some((indices.clone(), values.clone())))
382 }
383 Some(_) => Err(type_error(name, DataType::SparseVectorFp32)),
384 }
385 }
386
387 pub fn get_sparse_vector_fp16(&self, name: &str) -> Result<Option<(Vec<u32>, Vec<u16>)>> {
388 match self.vectors.get(name) {
389 None => Ok(None),
390 Some(VectorValue::SparseFp16 { indices, values }) => {
391 Ok(Some((indices.clone(), values.clone())))
392 }
393 Some(_) => Err(type_error(name, DataType::SparseVectorFp16)),
394 }
395 }
396}
397
398#[cfg(test)]
399mod tests {
400 use super::VectorValue;
401 use crate::types::MetricType;
402
403 fn reference_score(query: &[f64], values: &[f64], metric: MetricType) -> f64 {
404 match metric {
405 MetricType::L2 => -query
406 .iter()
407 .zip(values)
408 .map(|(left, right)| {
409 let difference = *left - *right;
410 difference * difference
411 })
412 .sum::<f64>(),
413 MetricType::Cosine => {
414 let dot = query
415 .iter()
416 .zip(values)
417 .map(|(left, right)| *left * *right)
418 .sum::<f64>();
419 let query_norm = query.iter().map(|value| value * value).sum::<f64>().sqrt();
420 let value_norm = values.iter().map(|value| value * value).sum::<f64>().sqrt();
421 if query_norm == 0.0 || value_norm == 0.0 {
422 0.0
423 } else {
424 dot / (query_norm * value_norm)
425 }
426 }
427 MetricType::MipsL2 | MetricType::Ip | MetricType::Undefined => query
428 .iter()
429 .zip(values)
430 .map(|(left, right)| *left * *right)
431 .sum(),
432 }
433 }
434
435 #[test]
436 fn borrowed_dense_scoring_matches_materialized_reference() {
437 let query = [0.25_f64, -0.5, 0.75, 0.125, -1.0];
438 let values = [-0.75_f64, -0.25, 0.5, 1.0, 0.125];
439 let query_norm = query.iter().map(|value| value * value).sum::<f64>().sqrt();
440 let fp16 = VectorValue::encode_fp16(&[-0.75_f32, -0.25, 0.5, 1.0, 0.125])
441 .expect("FP16 vector must encode");
442 let vectors = [
443 fp16,
444 VectorValue::Fp32(vec![-0.75_f32, -0.25, 0.5, 1.0, 0.125]),
445 VectorValue::Fp64(values.to_vec()),
446 VectorValue::Int4(vec![-1, 0, 1, 2, 3]),
447 VectorValue::Int8(vec![-7, -2, 4, 8, 1]),
448 VectorValue::Int16(vec![-7, -2, 4, 8, 1]),
449 ];
450 for vector in vectors {
451 let materialized = vector.to_dense_f64().expect("vector must be dense");
452 for metric in [
453 MetricType::L2,
454 MetricType::Cosine,
455 MetricType::Ip,
456 MetricType::MipsL2,
457 ] {
458 let actual = vector
459 .dense_score(&query, query_norm, metric)
460 .expect("dense vector must be scoreable");
461 let expected = reference_score(&query, &materialized, metric);
462 assert!(
463 (actual - expected).abs() <= f64::EPSILON,
464 "metric={metric:?} vector={vector:?} actual={actual} expected={expected}"
465 );
466 }
467 }
468 }
469
470 #[test]
471 fn borrowed_dense_scoring_rejects_non_dense_and_dimension_mismatch() {
472 let query = [1.0_f64, 2.0];
473 let norm = 5.0_f64.sqrt();
474 assert!(VectorValue::Binary32(vec![0, 0, 0, 0])
475 .dense_score(&query, norm, MetricType::Ip)
476 .is_none());
477 assert!(VectorValue::Fp32(vec![1.0])
478 .dense_score(&query, norm, MetricType::Ip)
479 .is_none());
480 }
481}