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