1use 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 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 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#[derive(Clone, Debug)]
145pub struct ProductQuantizationStorage {
146 metadata: ProductQuantizationMetadata,
147 distance_type: DistanceType,
148 batch: RecordBatch,
149
150 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(), 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 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 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 pub async fn write_partition(
404 &self,
405 writer: &mut PreviousFileWriter<ManifestDescribing>,
406 ) -> Result<usize> {
407 let batch_size: usize = 10240; 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 let codebook = match &metadata.codebook {
459 Some(codebook) => codebook.clone(),
460 None => {
461 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 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 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 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 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
836pub 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 debug_assert!(
926 false,
927 "cosine distance should be converted to normalized L2 distance"
928 );
929 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 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}