Skip to main content

a3s_vec/doc/
vector_api.rs

1//! Native vector conversions and typed document accessors.
2
3use 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    /// Encodes finite f32 coordinates as IEEE 754 half-precision bits.
11    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    /// Converts numeric dense forms to f32 for adapters that require it.
51    ///
52    /// FP64 coordinates may be narrowed. The exact collection executor uses
53    /// [`Self::to_dense_f64`] instead.
54    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    /// Decodes all numeric dense forms without narrowing FP64 coordinates.
68    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    /// Scores a dense query against this vector without materializing a
87    /// converted `Vec<f64>` for every candidate document.
88    ///
89    /// The collection executor intentionally keeps its authoritative score in
90    /// `f64`.  This helper only changes how the stored coordinates are
91    /// traversed; it does not narrow the arithmetic or alter metric semantics.
92    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}