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