Skip to main content

lance_index/vector/pq/
storage.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! Product Quantization storage
5//!
6//! Used as storage backend for Graph based algorithms.
7
8use std::{cmp::min, collections::HashMap, sync::Arc};
9
10use arrow::datatypes::{self, UInt8Type};
11use arrow_array::{Array, ArrayRef, ArrowPrimitiveType, PrimitiveArray};
12use arrow_array::{
13    FixedSizeListArray, RecordBatch, UInt8Array, UInt64Array,
14    cast::AsArray,
15    types::{Float32Type, UInt64Type},
16};
17use arrow_schema::{DataType, SchemaRef};
18use async_trait::async_trait;
19use bytes::{Bytes, BytesMut};
20use deepsize::DeepSizeOf;
21use lance_arrow::{FixedSizeListArrayExt, RecordBatchExt};
22use lance_core::{Error, ROW_ID, Result};
23use lance_file::previous::{
24    reader::FileReader as PreviousFileReader, writer::FileWriter as PreviousFileWriter,
25};
26use lance_io::{object_store::ObjectStore, utils::read_message};
27use lance_linalg::distance::{DistanceType, Dot, L2};
28use lance_table::utils::LanceIteratorExtension;
29use lance_table::{format::SelfDescribingFileReader, io::manifest::ManifestDescribing};
30use object_store::path::Path;
31use prost::Message;
32use serde::{Deserialize, Serialize};
33
34use super::ProductQuantizer;
35use super::distance::{build_distance_table_dot, build_distance_table_l2, compute_pq_distance};
36use crate::frag_reuse::FragReuseIndex;
37use crate::vector::graph::{OrderedFloat, OrderedNode};
38use crate::{
39    INDEX_METADATA_SCHEMA_KEY, IndexMetadata, pb,
40    vector::{
41        PQ_CODE_COLUMN,
42        pq::transform::PQTransformer,
43        quantizer::{QuantizerMetadata, QuantizerStorage},
44        storage::{DistCalculator, VectorStore},
45        transform::Transformer,
46    },
47};
48
49pub const PQ_METADATA_KEY: &str = "lance:pq";
50
51#[derive(Debug, Clone, Serialize, Deserialize)]
52pub struct ProductQuantizationMetadata {
53    pub codebook_position: usize,
54    pub nbits: u32,
55    pub num_sub_vectors: usize,
56    pub dimension: usize,
57
58    #[serde(skip)]
59    pub codebook: Option<FixedSizeListArray>,
60
61    // empty for v1 format
62    // used for v3 format
63    // deprecated in later version
64    pub codebook_tensor: Vec<u8>,
65    pub transposed: bool,
66}
67
68impl DeepSizeOf for ProductQuantizationMetadata {
69    fn deep_size_of_children(&self, _context: &mut deepsize::Context) -> usize {
70        self.codebook
71            .as_ref()
72            .map(|codebook| codebook.get_array_memory_size())
73            .unwrap_or(0)
74    }
75}
76
77impl PartialEq for ProductQuantizationMetadata {
78    fn eq(&self, other: &Self) -> bool {
79        self.num_sub_vectors == other.num_sub_vectors
80            && self.nbits == other.nbits
81            && self.dimension == other.dimension
82            && self.codebook == other.codebook
83    }
84}
85
86#[async_trait]
87impl QuantizerMetadata for ProductQuantizationMetadata {
88    fn buffer_index(&self) -> Option<u32> {
89        if self.codebook_position > 0 {
90            // the global buffer index starts from 1
91            Some(self.codebook_position as u32)
92        } else {
93            None
94        }
95    }
96
97    fn set_buffer_index(&mut self, index: u32) {
98        self.codebook_position = index as usize;
99    }
100
101    fn parse_buffer(&mut self, bytes: Bytes) -> Result<()> {
102        debug_assert!(!bytes.is_empty());
103        debug_assert!(self.codebook.is_none());
104        let codebook_tensor: pb::Tensor = pb::Tensor::decode(bytes)?;
105        self.codebook = Some(FixedSizeListArray::try_from(&codebook_tensor)?);
106        Ok(())
107    }
108
109    fn extra_metadata(&self) -> Result<Option<Bytes>> {
110        debug_assert!(self.codebook.is_some());
111        let codebook_tensor: pb::Tensor = pb::Tensor::try_from(self.codebook.as_ref().unwrap())?;
112        let mut bytes = BytesMut::new();
113        codebook_tensor.encode(&mut bytes)?;
114        Ok(Some(bytes.freeze()))
115    }
116
117    async fn load(reader: &PreviousFileReader) -> Result<Self> {
118        let metadata = reader
119            .schema()
120            .metadata
121            .get(PQ_METADATA_KEY)
122            .ok_or(Error::index(format!(
123                "Reading PQ storage: metadata key {} not found",
124                PQ_METADATA_KEY
125            )))?;
126        let mut metadata: Self = serde_json::from_str(metadata)
127            .map_err(|_| Error::index(format!("Failed to parse PQ metadata: {}", metadata)))?;
128
129        debug_assert!(metadata.codebook.is_none());
130        debug_assert!(metadata.codebook_tensor.is_empty());
131
132        let codebook_tensor: pb::Tensor =
133            read_message(reader.object_reader.as_ref(), metadata.codebook_position).await?;
134        metadata.codebook = Some(FixedSizeListArray::try_from(&codebook_tensor)?);
135        Ok(metadata)
136    }
137}
138
139/// Product Quantization Storage
140///
141/// It stores PQ code, as well as the row ID to the original vectors.
142///
143/// It is possible to store additional metadata to accelerate filtering later.
144#[derive(Clone, Debug)]
145pub struct ProductQuantizationStorage {
146    metadata: ProductQuantizationMetadata,
147    distance_type: DistanceType,
148    batch: RecordBatch,
149
150    // For easy access
151    pq_code: Arc<UInt8Array>,
152    row_ids: Arc<UInt64Array>,
153}
154
155impl DeepSizeOf for ProductQuantizationStorage {
156    fn deep_size_of_children(&self, _context: &mut deepsize::Context) -> usize {
157        self.batch.get_array_memory_size()
158            + self
159                .metadata
160                .codebook
161                .as_ref()
162                .map(|codebook| codebook.get_array_memory_size())
163                .unwrap_or(0)
164    }
165}
166
167impl PartialEq for ProductQuantizationStorage {
168    fn eq(&self, other: &Self) -> bool {
169        self.distance_type == other.distance_type
170            && self.metadata.eq(&other.metadata)
171            && self.batch.columns().eq(other.batch.columns())
172    }
173}
174
175impl ProductQuantizationStorage {
176    #[allow(clippy::too_many_arguments)]
177    pub fn new(
178        codebook: FixedSizeListArray,
179        mut batch: RecordBatch,
180        num_bits: u32,
181        num_sub_vectors: usize,
182        dimension: usize,
183        distance_type: DistanceType,
184        transposed: bool,
185        frag_reuse_index: Option<Arc<FragReuseIndex>>,
186    ) -> Result<Self> {
187        if batch.num_columns() != 2 {
188            log::warn!(
189                "PQ storage should have 2 columns, but got {} columns: {}",
190                batch.num_columns(),
191                batch.schema(),
192            );
193            batch = batch.project(&[
194                batch.schema().index_of(ROW_ID)?,
195                batch.schema().index_of(PQ_CODE_COLUMN)?,
196            ])?;
197        }
198
199        let Some(row_ids) = batch.column_by_name(ROW_ID) else {
200            return Err(Error::index(
201                "Row ID column not found from PQ storage".to_string(),
202            ));
203        };
204        let row_ids: Arc<UInt64Array> = row_ids
205            .as_primitive_opt::<UInt64Type>()
206            .ok_or(Error::index(
207                "Row ID column is not of type UInt64".to_string(),
208            ))?
209            .clone()
210            .into();
211
212        if !transposed {
213            let num_sub_vectors_in_byte = if num_bits == 4 {
214                num_sub_vectors / 2
215            } else {
216                num_sub_vectors
217            };
218            let pq_col = batch[PQ_CODE_COLUMN].as_fixed_size_list();
219            let transposed_code = transpose(
220                pq_col.values().as_primitive::<UInt8Type>(),
221                row_ids.len(),
222                num_sub_vectors_in_byte,
223            );
224            let pq_code_fsl = Arc::new(FixedSizeListArray::try_new_from_values(
225                transposed_code,
226                num_sub_vectors_in_byte as i32,
227            )?);
228            batch = batch.replace_column_by_name(PQ_CODE_COLUMN, pq_code_fsl)?;
229        }
230
231        let mut pq_code: Arc<UInt8Array> = batch[PQ_CODE_COLUMN]
232            .as_fixed_size_list()
233            .values()
234            .as_primitive()
235            .clone()
236            .into();
237
238        if let Some(frag_reuse_index_ref) = frag_reuse_index.as_ref() {
239            let transposed_codes = pq_code.values();
240            let mut new_row_ids = Vec::with_capacity(row_ids.len());
241            let mut new_codes = Vec::with_capacity(row_ids.len() * num_sub_vectors);
242
243            let row_ids_values = row_ids.values();
244            for (i, row_id) in row_ids_values.iter().enumerate() {
245                if let Some(mapped_value) = frag_reuse_index_ref.remap_row_id(*row_id) {
246                    new_row_ids.push(mapped_value);
247                    new_codes.extend(get_pq_code(
248                        transposed_codes,
249                        num_bits,
250                        num_sub_vectors,
251                        i as u32,
252                    ));
253                }
254            }
255
256            let new_row_ids = Arc::new(UInt64Array::from(new_row_ids));
257            let new_codes = UInt8Array::from(new_codes);
258            batch = if new_row_ids.is_empty() {
259                RecordBatch::new_empty(batch.schema())
260            } else {
261                let num_bytes_in_code = new_codes.len() / new_row_ids.len();
262                let new_transposed_codes =
263                    transpose(&new_codes, new_row_ids.len(), num_bytes_in_code);
264                let codes_fsl = Arc::new(FixedSizeListArray::try_new_from_values(
265                    new_transposed_codes,
266                    num_bytes_in_code as i32,
267                )?);
268                RecordBatch::try_new(batch.schema(), vec![new_row_ids, codes_fsl])?
269            };
270            pq_code = batch[PQ_CODE_COLUMN]
271                .as_fixed_size_list()
272                .values()
273                .as_primitive::<UInt8Type>()
274                .clone()
275                .into();
276        }
277
278        let distance_type = match distance_type {
279            DistanceType::Cosine => DistanceType::L2,
280            _ => distance_type,
281        };
282        let metadata = ProductQuantizationMetadata {
283            codebook_position: 0,
284            nbits: num_bits,
285            num_sub_vectors,
286            dimension,
287            codebook: Some(codebook),
288            codebook_tensor: Vec::new(), // empty for v1 format
289            transposed: true,
290        };
291        Ok(Self {
292            metadata,
293            distance_type,
294            batch,
295            pq_code,
296            row_ids,
297        })
298    }
299
300    pub fn batch(&self) -> &RecordBatch {
301        &self.batch
302    }
303
304    /// Build a PQ storage from ProductQuantizer and a RecordBatch.
305    ///
306    /// Parameters
307    /// ----------
308    /// quantizer: ProductQuantizer
309    ///    The quantizer used to transform the vectors.
310    /// batch: RecordBatch
311    ///   The batch of vectors to be transformed.
312    /// vector_col: &str
313    ///   The name of the column containing the vectors.
314    pub async fn build(
315        quantizer: ProductQuantizer,
316        batch: &RecordBatch,
317        vector_col: &str,
318        frag_reuse_index: Option<Arc<FragReuseIndex>>,
319    ) -> Result<Self> {
320        let codebook = quantizer.codebook.clone();
321        let num_bits = quantizer.num_bits;
322        let dimension = quantizer.dimension;
323        let num_sub_vectors = quantizer.num_sub_vectors;
324        let metric_type = quantizer.distance_type;
325        let transform = PQTransformer::new(quantizer, vector_col, PQ_CODE_COLUMN);
326        let batch = transform.transform(batch)?;
327        Self::new(
328            codebook,
329            batch,
330            num_bits,
331            num_sub_vectors,
332            dimension,
333            metric_type,
334            false,
335            frag_reuse_index,
336        )
337    }
338
339    pub fn codebook(&self) -> &FixedSizeListArray {
340        self.metadata.codebook.as_ref().unwrap()
341    }
342
343    /// Load full PQ storage from disk.
344    ///
345    /// Parameters
346    /// ----------
347    /// object_store: &ObjectStore
348    ///   The object store to load the storage from.
349    /// path: &Path
350    ///  The path to the storage.
351    ///
352    /// Returns
353    /// --------
354    /// Self
355    ///
356    /// Currently it loads everything in memory.
357    /// TODO: support lazy loading later.
358    pub async fn load(
359        object_store: &ObjectStore,
360        path: &Path,
361        frag_reuse_index: Option<Arc<FragReuseIndex>>,
362    ) -> Result<Self> {
363        let reader = PreviousFileReader::try_new_self_described(object_store, path, None).await?;
364        let schema = reader.schema();
365
366        let metadata_str = schema
367            .metadata
368            .get(INDEX_METADATA_SCHEMA_KEY)
369            .ok_or(Error::index(format!(
370                "Reading PQ storage: index key {} not found",
371                INDEX_METADATA_SCHEMA_KEY
372            )))?;
373        let index_metadata: IndexMetadata = serde_json::from_str(metadata_str).map_err(|_| {
374            Error::index(format!("Failed to parse index metadata: {}", metadata_str))
375        })?;
376        let distance_type: DistanceType =
377            DistanceType::try_from(index_metadata.distance_type.as_str())?;
378
379        let metadata = ProductQuantizationMetadata::load(&reader).await?;
380        Self::load_partition(
381            &reader,
382            0..reader.len(),
383            distance_type,
384            &metadata,
385            frag_reuse_index,
386        )
387        .await
388    }
389
390    pub fn schema(&self) -> SchemaRef {
391        self.batch.schema()
392    }
393
394    pub fn get_row_ids(&self, ids: &[u32]) -> Vec<u64> {
395        ids.iter()
396            .map(|&id| self.row_ids.value(id as usize))
397            .collect()
398    }
399
400    /// Write the PQ storage as a Lance partition to disk,
401    /// and returns the number of rows written.
402    ///
403    pub async fn write_partition(
404        &self,
405        writer: &mut PreviousFileWriter<ManifestDescribing>,
406    ) -> Result<usize> {
407        let batch_size: usize = 10240; // TODO: make it configurable
408        for offset in (0..self.batch.num_rows()).step_by(batch_size) {
409            let length = min(batch_size, self.batch.num_rows() - offset);
410            let slice = self.batch.slice(offset, length);
411            writer.write(&[slice]).await?;
412        }
413        Ok(self.batch.num_rows())
414    }
415}
416
417pub fn transpose<T: ArrowPrimitiveType>(
418    original: &PrimitiveArray<T>,
419    num_rows: usize,
420    num_columns: usize,
421) -> PrimitiveArray<T>
422where
423    PrimitiveArray<T>: From<Vec<T::Native>>,
424{
425    if original.is_empty() {
426        return original.clone();
427    }
428
429    let mut transposed_codes = vec![T::default_value(); original.len()];
430    for (vec_idx, codes) in original.values().chunks_exact(num_columns).enumerate() {
431        for (sub_vec_idx, code) in codes.iter().enumerate() {
432            transposed_codes[sub_vec_idx * num_rows + vec_idx] = *code;
433        }
434    }
435
436    transposed_codes.into()
437}
438
439#[async_trait]
440impl QuantizerStorage for ProductQuantizationStorage {
441    type Metadata = ProductQuantizationMetadata;
442
443    fn try_from_batch(
444        batch: RecordBatch,
445        metadata: &Self::Metadata,
446        distance_type: DistanceType,
447        frag_reuse_index: Option<Arc<FragReuseIndex>>,
448    ) -> Result<Self>
449    where
450        Self: Sized,
451    {
452        let distance_type = match distance_type {
453            DistanceType::Cosine => DistanceType::L2,
454            _ => distance_type,
455        };
456
457        // now it supports only Float32Type
458        let codebook = match &metadata.codebook {
459            Some(codebook) => codebook.clone(),
460            None => {
461                // legacy format would contains codebook tensor but not codebook
462                debug_assert!(!metadata.codebook_tensor.is_empty());
463                let codebook_tensor = pb::Tensor::decode(metadata.codebook_tensor.as_slice())?;
464                FixedSizeListArray::try_from(&codebook_tensor)?
465            }
466        };
467
468        Self::new(
469            codebook,
470            batch,
471            metadata.nbits,
472            metadata.num_sub_vectors,
473            metadata.dimension,
474            distance_type,
475            metadata.transposed,
476            frag_reuse_index,
477        )
478    }
479
480    fn metadata(&self) -> &Self::Metadata {
481        &self.metadata
482    }
483
484    // we can't use the default implementation of remap,
485    // because PQ Storage transposed the PQ codes
486    fn remap(&self, mapping: &HashMap<u64, Option<u64>>) -> Result<Self> {
487        let transposed_codes = self.pq_code.values();
488        let mut new_row_ids = Vec::with_capacity(self.len());
489        let mut new_codes = Vec::with_capacity(self.len() * self.metadata.num_sub_vectors);
490
491        let row_ids = self.row_ids.values();
492        for (i, row_id) in row_ids.iter().enumerate() {
493            match mapping.get(row_id) {
494                Some(Some(new_id)) => {
495                    new_row_ids.push(*new_id);
496                    new_codes.extend(get_pq_code(
497                        transposed_codes,
498                        self.metadata.nbits,
499                        self.metadata.num_sub_vectors,
500                        i as u32,
501                    ));
502                }
503                Some(None) => {}
504                None => {
505                    new_row_ids.push(*row_id);
506                    new_codes.extend(get_pq_code(
507                        transposed_codes,
508                        self.metadata.nbits,
509                        self.metadata.num_sub_vectors,
510                        i as u32,
511                    ));
512                }
513            }
514        }
515
516        let new_row_ids = Arc::new(UInt64Array::from(new_row_ids));
517        let new_codes = UInt8Array::from(new_codes);
518        let batch = if new_row_ids.is_empty() {
519            RecordBatch::new_empty(self.schema())
520        } else {
521            let num_bytes_in_code = new_codes.len() / new_row_ids.len();
522            let new_transposed_codes = transpose(&new_codes, new_row_ids.len(), num_bytes_in_code);
523            let codes_fsl = Arc::new(FixedSizeListArray::try_new_from_values(
524                new_transposed_codes,
525                num_bytes_in_code as i32,
526            )?);
527            RecordBatch::try_new(self.schema(), vec![new_row_ids.clone(), codes_fsl])?
528        };
529        let transposed_codes = batch[PQ_CODE_COLUMN]
530            .as_fixed_size_list()
531            .values()
532            .as_primitive::<UInt8Type>()
533            .clone();
534
535        Ok(Self {
536            metadata: self.metadata.clone(),
537            distance_type: self.distance_type,
538            batch,
539            pq_code: Arc::new(transposed_codes),
540            row_ids: new_row_ids,
541        })
542    }
543
544    /// Load a partition of PQ storage from disk.
545    ///
546    /// Parameters
547    /// ----------
548    /// - *reader: &PreviousFileReader
549    async fn load_partition(
550        reader: &PreviousFileReader,
551        range: std::ops::Range<usize>,
552        distance_type: DistanceType,
553        metadata: &Self::Metadata,
554        frag_reuse_index: Option<Arc<FragReuseIndex>>,
555    ) -> Result<Self> {
556        // Hard coded to float32 for now
557        let codebook = metadata
558            .codebook
559            .as_ref()
560            .ok_or(Error::index(
561                "Codebook not found in PQ metadata".to_string(),
562            ))?
563            .values()
564            .as_primitive::<Float32Type>()
565            .clone();
566
567        let codebook =
568            FixedSizeListArray::try_new_from_values(codebook, metadata.dimension as i32)?;
569
570        let schema = reader.schema();
571        let batch = reader.read_range(range, schema).await?;
572
573        Self::new(
574            codebook,
575            batch,
576            metadata.nbits,
577            metadata.num_sub_vectors,
578            metadata.dimension,
579            distance_type,
580            metadata.transposed,
581            frag_reuse_index,
582        )
583    }
584}
585
586impl VectorStore for ProductQuantizationStorage {
587    type DistanceCalculator<'a> = PQDistCalculator;
588
589    fn to_batches(&self) -> Result<impl Iterator<Item = RecordBatch>> {
590        Ok(std::iter::once(self.batch.clone()))
591    }
592
593    fn append_batch(&self, _batch: RecordBatch, _vector_column: &str) -> Result<Self> {
594        unimplemented!()
595    }
596
597    fn schema(&self) -> &SchemaRef {
598        self.batch.schema_ref()
599    }
600
601    fn as_any(&self) -> &dyn std::any::Any {
602        self
603    }
604
605    fn len(&self) -> usize {
606        self.batch.num_rows()
607    }
608
609    fn distance_type(&self) -> DistanceType {
610        self.distance_type
611    }
612
613    fn row_id(&self, id: u32) -> u64 {
614        self.row_ids.values()[id as usize]
615    }
616
617    fn row_ids(&self) -> impl Iterator<Item = &u64> {
618        self.row_ids.values().iter()
619    }
620
621    fn dist_calculator(&self, query: ArrayRef, _dist_q_c: f32) -> Self::DistanceCalculator<'_> {
622        let codebook = self.metadata.codebook.as_ref().unwrap();
623        match codebook.value_type() {
624            DataType::Float16 => PQDistCalculator::new(
625                codebook
626                    .values()
627                    .as_primitive::<datatypes::Float16Type>()
628                    .values(),
629                self.metadata.nbits,
630                self.metadata.num_sub_vectors,
631                self.pq_code.clone(),
632                query.as_primitive::<datatypes::Float16Type>().values(),
633                self.distance_type,
634            ),
635            DataType::Float32 => PQDistCalculator::new(
636                codebook
637                    .values()
638                    .as_primitive::<datatypes::Float32Type>()
639                    .values(),
640                self.metadata.nbits,
641                self.metadata.num_sub_vectors,
642                self.pq_code.clone(),
643                query.as_primitive::<datatypes::Float32Type>().values(),
644                self.distance_type,
645            ),
646            DataType::Float64 => PQDistCalculator::new(
647                codebook
648                    .values()
649                    .as_primitive::<datatypes::Float64Type>()
650                    .values(),
651                self.metadata.nbits,
652                self.metadata.num_sub_vectors,
653                self.pq_code.clone(),
654                query.as_primitive::<datatypes::Float64Type>().values(),
655                self.distance_type,
656            ),
657            _ => unimplemented!("Unsupported data type: {:?}", codebook.value_type()),
658        }
659    }
660
661    fn dist_calculator_from_id(&self, id: u32) -> Self::DistanceCalculator<'_> {
662        let codes = get_pq_code(
663            self.pq_code.values(),
664            self.metadata.nbits,
665            self.metadata.num_sub_vectors,
666            id,
667        );
668        let codebook = self.metadata.codebook.as_ref().unwrap();
669        match codebook.value_type() {
670            DataType::Float16 => {
671                let codebook = codebook
672                    .values()
673                    .as_primitive::<datatypes::Float16Type>()
674                    .values();
675                let query = get_centroids(
676                    codebook,
677                    self.metadata.nbits,
678                    self.metadata.num_sub_vectors,
679                    self.metadata.dimension,
680                    codes,
681                );
682                PQDistCalculator::new(
683                    codebook,
684                    self.metadata.nbits,
685                    self.metadata.num_sub_vectors,
686                    self.pq_code.clone(),
687                    &query,
688                    self.distance_type,
689                )
690            }
691            DataType::Float32 => {
692                let codebook = codebook
693                    .values()
694                    .as_primitive::<datatypes::Float32Type>()
695                    .values();
696                let query = get_centroids(
697                    codebook,
698                    self.metadata.nbits,
699                    self.metadata.num_sub_vectors,
700                    self.metadata.dimension,
701                    codes,
702                );
703                PQDistCalculator::new(
704                    codebook,
705                    self.metadata.nbits,
706                    self.metadata.num_sub_vectors,
707                    self.pq_code.clone(),
708                    &query,
709                    self.distance_type,
710                )
711            }
712            DataType::Float64 => {
713                let codebook = codebook
714                    .values()
715                    .as_primitive::<datatypes::Float64Type>()
716                    .values();
717                let query = get_centroids(
718                    codebook,
719                    self.metadata.nbits,
720                    self.metadata.num_sub_vectors,
721                    self.metadata.dimension,
722                    codes,
723                );
724                PQDistCalculator::new(
725                    codebook,
726                    self.metadata.nbits,
727                    self.metadata.num_sub_vectors,
728                    self.pq_code.clone(),
729                    &query,
730                    self.distance_type,
731                )
732            }
733            _ => unimplemented!("Unsupported data type: {:?}", codebook.value_type()),
734        }
735    }
736
737    fn dist_between(&self, u: u32, v: u32) -> f32 {
738        // this is a fast way to compute distance between two vectors in the same storage.
739        // it doesn't construct the distance table.
740        let pq_codes = self.pq_code.values();
741        let u_codes = get_pq_code(
742            pq_codes,
743            self.metadata.nbits,
744            self.metadata.num_sub_vectors,
745            u,
746        );
747        let v_codes = get_pq_code(
748            pq_codes,
749            self.metadata.nbits,
750            self.metadata.num_sub_vectors,
751            v,
752        );
753        let codebook = self.metadata.codebook.as_ref().unwrap();
754
755        match codebook.value_type() {
756            DataType::Float16 => {
757                let qu = get_centroids(
758                    codebook
759                        .values()
760                        .as_primitive::<datatypes::Float16Type>()
761                        .values(),
762                    self.metadata.nbits,
763                    self.metadata.num_sub_vectors,
764                    self.metadata.dimension,
765                    u_codes,
766                );
767                let qv = get_centroids(
768                    codebook
769                        .values()
770                        .as_primitive::<datatypes::Float16Type>()
771                        .values(),
772                    self.metadata.nbits,
773                    self.metadata.num_sub_vectors,
774                    self.metadata.dimension,
775                    v_codes,
776                );
777                self.distance_type.func()(&qu, &qv)
778            }
779            DataType::Float32 => {
780                let qu = get_centroids(
781                    codebook
782                        .values()
783                        .as_primitive::<datatypes::Float32Type>()
784                        .values(),
785                    self.metadata.nbits,
786                    self.metadata.num_sub_vectors,
787                    self.metadata.dimension,
788                    u_codes,
789                );
790                let qv = get_centroids(
791                    codebook
792                        .values()
793                        .as_primitive::<datatypes::Float32Type>()
794                        .values(),
795                    self.metadata.nbits,
796                    self.metadata.num_sub_vectors,
797                    self.metadata.dimension,
798                    v_codes,
799                );
800                self.distance_type.func()(&qu, &qv)
801            }
802            DataType::Float64 => {
803                let qu = get_centroids(
804                    codebook
805                        .values()
806                        .as_primitive::<datatypes::Float64Type>()
807                        .values(),
808                    self.metadata.nbits,
809                    self.metadata.num_sub_vectors,
810                    self.metadata.dimension,
811                    u_codes,
812                );
813                let qv = get_centroids(
814                    codebook
815                        .values()
816                        .as_primitive::<datatypes::Float64Type>()
817                        .values(),
818                    self.metadata.nbits,
819                    self.metadata.num_sub_vectors,
820                    self.metadata.dimension,
821                    v_codes,
822                );
823                self.distance_type.func()(&qu, &qv)
824            }
825            _ => unimplemented!("Unsupported data type: {:?}", codebook.value_type()),
826        }
827    }
828
829    fn prefers_candidate(&self, candidate: &OrderedNode, selected: &[OrderedNode]) -> bool {
830        selected
831            .iter()
832            .all(|other| candidate.dist < OrderedFloat(self.dist_between(candidate.id, other.id)))
833    }
834}
835
836/// Distance calculator backed by PQ code.
837pub struct PQDistCalculator {
838    distance_table: Vec<f32>,
839    pq_code: Arc<UInt8Array>,
840    num_sub_vectors: usize,
841    num_bits: u32,
842    distance_type: DistanceType,
843}
844
845impl PQDistCalculator {
846    fn new<T: L2 + Dot>(
847        codebook: &[T],
848        num_bits: u32,
849        num_sub_vectors: usize,
850        pq_code: Arc<UInt8Array>,
851        query: &[T],
852        distance_type: DistanceType,
853    ) -> Self {
854        let distance_table = match distance_type {
855            DistanceType::L2 | DistanceType::Cosine => {
856                build_distance_table_l2(codebook, num_bits, num_sub_vectors, query)
857            }
858            DistanceType::Dot => {
859                build_distance_table_dot(codebook, num_bits, num_sub_vectors, query)
860            }
861            _ => unimplemented!("DistanceType is not supported: {:?}", distance_type),
862        };
863        Self {
864            distance_table,
865            num_sub_vectors,
866            pq_code,
867            num_bits,
868            distance_type,
869        }
870    }
871
872    fn get_pq_code(&self, id: u32) -> impl Iterator<Item = usize> + '_ {
873        get_pq_code(
874            self.pq_code.values(),
875            self.num_bits,
876            self.num_sub_vectors,
877            id,
878        )
879        .map(|v| v as usize)
880    }
881}
882
883impl DistCalculator for PQDistCalculator {
884    fn distance(&self, id: u32) -> f32 {
885        let num_centroids = 2_usize.pow(self.num_bits);
886        let pq_code = self.get_pq_code(id);
887        let diff = self.num_sub_vectors as f32 - 1.0;
888        let dist = if self.num_bits == 4 {
889            pq_code
890                .enumerate()
891                .map(|(i, c)| {
892                    let current_idx = c & 0x0F;
893                    let next_idx = c >> 4;
894
895                    self.distance_table[2 * i * num_centroids + current_idx]
896                        + self.distance_table[(2 * i + 1) * num_centroids + next_idx]
897                })
898                .sum()
899        } else {
900            pq_code
901                .enumerate()
902                .map(|(i, c)| self.distance_table[i * num_centroids + c])
903                .sum()
904        };
905
906        if self.distance_type == DistanceType::Dot {
907            dist - diff
908        } else {
909            dist
910        }
911    }
912
913    fn distance_all(&self, k_hint: usize) -> Vec<f32> {
914        match self.distance_type {
915            DistanceType::L2 => compute_pq_distance(
916                &self.distance_table,
917                self.num_bits,
918                self.num_sub_vectors,
919                self.pq_code.values(),
920                k_hint,
921            ),
922            DistanceType::Cosine => {
923                // it seems we implemented cosine distance at some version,
924                // but from now on, we should use normalized L2 distance.
925                debug_assert!(
926                    false,
927                    "cosine distance should be converted to normalized L2 distance"
928                );
929                // L2 over normalized vectors:  ||x - y|| = x^2 + y^2 - 2 * xy = 1 + 1 - 2 * xy = 2 * (1 - xy)
930                // Cosine distance: 1 - |xy| / (||x|| * ||y||) = 1 - xy / (x^2 * y^2) = 1 - xy / (1 * 1) = 1 - xy
931                // Therefore, Cosine = L2 / 2
932                let l2_dists = compute_pq_distance(
933                    &self.distance_table,
934                    self.num_bits,
935                    self.num_sub_vectors,
936                    self.pq_code.values(),
937                    k_hint,
938                );
939                l2_dists.into_iter().map(|v| v / 2.0).collect()
940            }
941            DistanceType::Dot => {
942                let dot_dists = compute_pq_distance(
943                    &self.distance_table,
944                    self.num_bits,
945                    self.num_sub_vectors,
946                    self.pq_code.values(),
947                    k_hint,
948                );
949                let diff = self.num_sub_vectors as f32 - 1.0;
950                dot_dists.into_iter().map(|v| v - diff).collect()
951            }
952            _ => unimplemented!("distance type is not supported: {:?}", self.distance_type),
953        }
954    }
955}
956
957fn get_pq_code(
958    pq_code: &[u8],
959    num_bits: u32,
960    num_sub_vectors: usize,
961    id: u32,
962) -> impl Iterator<Item = u8> + '_ {
963    let num_bytes = if num_bits == 4 {
964        num_sub_vectors / 2
965    } else {
966        num_sub_vectors
967    };
968
969    let num_vectors = pq_code.len() / num_bytes;
970    pq_code
971        .iter()
972        .skip(id as usize)
973        .step_by(num_vectors)
974        .copied()
975        .exact_size(num_bytes)
976}
977
978fn get_centroids<T: Clone>(
979    codebook: &[T],
980    num_bits: u32,
981    num_sub_vectors: usize,
982    dimension: usize,
983    codes: impl Iterator<Item = u8>,
984) -> Vec<T> {
985    // codebook[i][j] is the j-th centroid of the i-th sub-vector.
986    // the codebook is stored as a flat array, codebook[i * num_centroids + j] = codebook[i][j]
987
988    if num_bits == 4 {
989        return get_centroids_4bit(codebook, num_sub_vectors, dimension, codes);
990    }
991
992    let num_centroids: usize = 2_usize.pow(8);
993    let sub_vector_width = dimension / num_sub_vectors;
994    let mut centroids = Vec::with_capacity(dimension);
995    for (sub_vec_idx, centroid_idx) in codes.enumerate() {
996        let centroid_idx = centroid_idx as usize;
997        let centroid = &codebook[sub_vec_idx * num_centroids * sub_vector_width
998            + centroid_idx * sub_vector_width
999            ..sub_vec_idx * num_centroids * sub_vector_width
1000                + (centroid_idx + 1) * sub_vector_width];
1001        centroids.extend_from_slice(centroid);
1002    }
1003    centroids
1004}
1005
1006fn get_centroids_4bit<T: Clone>(
1007    codebook: &[T],
1008    num_sub_vectors: usize,
1009    dimension: usize,
1010    codes: impl Iterator<Item = u8>,
1011) -> Vec<T> {
1012    let num_centroids: usize = 16;
1013    let sub_vector_width = dimension / num_sub_vectors;
1014    let mut centroids = Vec::with_capacity(dimension);
1015    for (sub_vec_idx, centroid_idx) in codes.into_iter().enumerate() {
1016        let current_idx = (centroid_idx & 0x0F) as usize;
1017        let offset = 2 * sub_vec_idx * num_centroids * sub_vector_width;
1018        let current_centroid = &codebook[offset + current_idx * sub_vector_width
1019            ..offset + (current_idx + 1) * sub_vector_width];
1020        centroids.extend_from_slice(current_centroid);
1021
1022        let next_idx = (centroid_idx >> 4) as usize;
1023        let offset = (2 * sub_vec_idx + 1) * num_centroids * sub_vector_width;
1024        let next_centroid = &codebook
1025            [offset + next_idx * sub_vector_width..offset + (next_idx + 1) * sub_vector_width];
1026        centroids.extend_from_slice(next_centroid);
1027    }
1028    centroids
1029}
1030
1031#[cfg(test)]
1032mod tests {
1033    use crate::vector::storage::StorageBuilder;
1034
1035    use super::*;
1036
1037    use arrow_array::{Float32Array, UInt32Array};
1038    use arrow_schema::{DataType, Field, Schema as ArrowSchema};
1039    use lance_arrow::FixedSizeListArrayExt;
1040    use lance_core::ROW_ID_FIELD;
1041    use rand::Rng;
1042
1043    const DIM: usize = 32;
1044    const TOTAL: usize = 512;
1045    const NUM_SUB_VECTORS: usize = 16;
1046
1047    async fn create_pq_storage() -> ProductQuantizationStorage {
1048        let codebook = Float32Array::from_iter_values((0..256 * DIM).map(|_| rand::random()));
1049        let codebook = FixedSizeListArray::try_new_from_values(codebook, DIM as i32).unwrap();
1050        let pq = ProductQuantizer::new(NUM_SUB_VECTORS, 8, DIM, codebook, DistanceType::Dot);
1051
1052        let schema = ArrowSchema::new(vec![
1053            Field::new(
1054                "vec",
1055                DataType::FixedSizeList(
1056                    Field::new_list_field(DataType::Float32, true).into(),
1057                    DIM as i32,
1058                ),
1059                true,
1060            ),
1061            ROW_ID_FIELD.clone(),
1062        ]);
1063        let vectors = Float32Array::from_iter_values((0..TOTAL * DIM).map(|_| rand::random()));
1064        let row_ids = UInt64Array::from_iter_values((0..TOTAL).map(|v| v as u64));
1065        let fsl = FixedSizeListArray::try_new_from_values(vectors, DIM as i32).unwrap();
1066        let batch =
1067            RecordBatch::try_new(schema.into(), vec![Arc::new(fsl), Arc::new(row_ids)]).unwrap();
1068
1069        StorageBuilder::new("vec".to_owned(), pq.distance_type, pq, None)
1070            .unwrap()
1071            .build(vec![batch])
1072            .unwrap()
1073    }
1074
1075    async fn create_pq_storage_with_extra_column() -> ProductQuantizationStorage {
1076        let codebook = Float32Array::from_iter_values((0..256 * DIM).map(|_| rand::random()));
1077        let codebook = FixedSizeListArray::try_new_from_values(codebook, DIM as i32).unwrap();
1078        let pq = ProductQuantizer::new(NUM_SUB_VECTORS, 8, DIM, codebook, DistanceType::Dot);
1079
1080        let schema = ArrowSchema::new(vec![
1081            Field::new(
1082                "vec",
1083                DataType::FixedSizeList(
1084                    Field::new_list_field(DataType::Float32, true).into(),
1085                    DIM as i32,
1086                ),
1087                true,
1088            ),
1089            ROW_ID_FIELD.clone(),
1090            Field::new("extra", DataType::UInt32, true),
1091        ]);
1092        let vectors = Float32Array::from_iter_values((0..TOTAL * DIM).map(|_| rand::random()));
1093        let row_ids = UInt64Array::from_iter_values((0..TOTAL).map(|v| v as u64));
1094        let extra_column = UInt32Array::from_iter_values((0..TOTAL).map(|v| v as u32));
1095        let fsl = FixedSizeListArray::try_new_from_values(vectors, DIM as i32).unwrap();
1096        let batch = RecordBatch::try_new(
1097            schema.into(),
1098            vec![Arc::new(fsl), Arc::new(row_ids), Arc::new(extra_column)],
1099        )
1100        .unwrap();
1101
1102        StorageBuilder::new("vec".to_owned(), pq.distance_type, pq, None)
1103            .unwrap()
1104            .build(vec![batch])
1105            .unwrap()
1106    }
1107
1108    #[tokio::test]
1109    async fn test_build_pq_storage() {
1110        let storage = create_pq_storage().await;
1111        assert_eq!(storage.len(), TOTAL);
1112        assert_eq!(storage.metadata.num_sub_vectors, NUM_SUB_VECTORS);
1113        assert_eq!(
1114            storage.metadata.codebook.as_ref().unwrap().values().len(),
1115            256 * DIM
1116        );
1117        assert_eq!(storage.pq_code.len(), TOTAL * NUM_SUB_VECTORS);
1118        assert_eq!(storage.row_ids.len(), TOTAL);
1119    }
1120
1121    #[tokio::test]
1122    async fn test_distance_all() {
1123        let storage = create_pq_storage().await;
1124        let query = Arc::new(Float32Array::from_iter_values((0..DIM).map(|v| v as f32)));
1125        let dist_calc = storage.dist_calculator(query, 0.0);
1126        let expected = (0..storage.len())
1127            .map(|id| dist_calc.distance(id as u32))
1128            .collect::<Vec<_>>();
1129        let distances = dist_calc.distance_all(100);
1130        assert_eq!(distances, expected);
1131    }
1132
1133    #[tokio::test]
1134    async fn test_dist_between() {
1135        let mut rng = rand::rng();
1136        let storage = create_pq_storage().await;
1137        let u = rng.random_range(0..storage.len() as u32);
1138        let v = rng.random_range(0..storage.len() as u32);
1139        let dist1 = storage.dist_between(u, v);
1140        let dist2 = storage.dist_between(v, u);
1141        assert_eq!(dist1, dist2);
1142    }
1143
1144    #[tokio::test]
1145    async fn test_remap_with_extra_column() {
1146        let storage = create_pq_storage_with_extra_column().await;
1147        let mut mapping = HashMap::new();
1148        for i in 0..TOTAL / 2 {
1149            mapping.insert(i as u64, Some((TOTAL + i) as u64));
1150        }
1151        for i in TOTAL / 2..TOTAL {
1152            mapping.insert(i as u64, None);
1153        }
1154        let new_storage = storage.remap(&mapping).unwrap();
1155        assert_eq!(new_storage.len(), TOTAL / 2);
1156        assert_eq!(new_storage.row_ids.len(), TOTAL / 2);
1157        for (i, row_id) in new_storage.row_ids().enumerate() {
1158            assert_eq!(*row_id, (TOTAL + i) as u64);
1159        }
1160        assert_eq!(new_storage.batch.num_columns(), 2);
1161        assert!(new_storage.batch.column_by_name(ROW_ID).is_some());
1162        assert!(new_storage.batch.column_by_name(PQ_CODE_COLUMN).is_some());
1163    }
1164}