use std::collections::HashMap;
use std::sync::Arc;
use crate::error::{LaurusError, Result};
use crate::vector::core::quantization::PqParams;
use crate::vector::core::vector::Vector;
pub const BLOCK_SIZE: usize = 32;
pub const BYTES_PER_SUB_PER_BLOCK: usize = BLOCK_SIZE / 2;
#[derive(Debug)]
pub struct PqFastScanPool {
pub params: PqParams,
pub dim: usize,
pub codebook: Vec<f32>,
pub packed: Vec<u8>,
pub n_vectors: usize,
pub field_index: HashMap<String, Arc<HashMap<u64, u32>>>,
}
impl PqFastScanPool {
#[inline]
pub fn block_count(&self) -> usize {
n_blocks_for(self.n_vectors)
}
#[inline]
pub fn block_stride(&self) -> usize {
self.params.m as usize * BYTES_PER_SUB_PER_BLOCK
}
pub fn build(
params: PqParams,
codebook: Vec<f32>,
records: impl IntoIterator<Item = (u64, String, Vec<u8>)>,
) -> Result<Self> {
if params.k != 16 {
return Err(LaurusError::InvalidOperation(format!(
"PqFastScanPool requires PqParams::k == 16 (got {})",
params.k
)));
}
if codebook.len() != params.codebook_len() {
return Err(LaurusError::InvalidOperation(format!(
"PqFastScanPool: codebook length {} does not match params.codebook_len() {}",
codebook.len(),
params.codebook_len()
)));
}
let m = params.m as usize;
let mut all_codes: Vec<Vec<u8>> = Vec::new();
let mut by_field: HashMap<String, HashMap<u64, u32>> = HashMap::new();
for (doc_id, field, codes) in records {
if codes.len() != m {
return Err(LaurusError::InvalidOperation(format!(
"PqFastScanPool: code length {} does not match params.m {}",
codes.len(),
m
)));
}
for &c in &codes {
if c >= 16 {
return Err(LaurusError::InvalidOperation(format!(
"PqFastScanPool: code value {c} exceeds 4-bit range [0, 15]"
)));
}
}
let pos = all_codes.len() as u32;
all_codes.push(codes);
by_field.entry(field).or_default().insert(doc_id, pos);
}
let n_vectors = all_codes.len();
let packed = pack_codes_into_blocks(&all_codes, m);
let field_index: HashMap<String, Arc<HashMap<u64, u32>>> = by_field
.into_iter()
.map(|(field, map)| (field, Arc::new(map)))
.collect();
Ok(Self {
params,
dim: params.original_dim(),
codebook,
packed,
n_vectors,
field_index,
})
}
pub fn codes_at(&self, vec_idx: usize) -> Vec<u8> {
let m = self.params.m as usize;
let mut out = vec![0u8; m];
let block_idx = vec_idx / BLOCK_SIZE;
let in_block = vec_idx % BLOCK_SIZE;
let block_base = block_idx * self.block_stride();
let (j, shift) = if in_block < 16 {
(in_block, 0)
} else {
(in_block - 16, 4)
};
for (sub, slot) in out.iter_mut().enumerate().take(m) {
let byte = self.packed[block_base + sub * BYTES_PER_SUB_PER_BLOCK + j];
*slot = (byte >> shift) & 0x0F;
}
out
}
#[inline]
pub fn vector_count(&self) -> usize {
self.n_vectors
}
#[inline]
pub fn field_position_index(&self, field: &str) -> Option<Arc<HashMap<u64, u32>>> {
self.field_index.get(field).cloned()
}
#[inline]
pub fn contains(&self, doc_id: u64, field: &str) -> bool {
self.field_index
.get(field)
.is_some_and(|m| m.contains_key(&doc_id))
}
pub fn keys(&self) -> Vec<(u64, String)> {
let mut keys: Vec<(u64, String)> = self
.field_index
.iter()
.flat_map(|(field, map)| map.keys().map(move |id| (*id, field.clone())))
.collect();
keys.sort_by_key(|(id, _)| *id);
keys
}
pub fn dequantize_to_vector(&self, doc_id: u64, field: &str) -> Option<Vector> {
let pos = *self.field_index.get(field)?.get(&doc_id)?;
let codes = self.codes_at(pos as usize);
let data =
crate::vector::core::quantization::pq_decode(&codes, self.params, &self.codebook);
Some(Vector::new(data))
}
}
#[inline]
pub fn n_blocks_for(n_vectors: usize) -> usize {
n_vectors.div_ceil(BLOCK_SIZE)
}
pub fn pack_codes_into_blocks(codes: &[Vec<u8>], m: usize) -> Vec<u8> {
let n_vectors = codes.len();
let n_blocks = n_blocks_for(n_vectors);
let block_stride = m * BYTES_PER_SUB_PER_BLOCK;
let mut packed = vec![0u8; n_blocks * block_stride];
for (vec_idx, v_codes) in codes.iter().enumerate().take(n_vectors) {
let block_idx = vec_idx / BLOCK_SIZE;
let in_block = vec_idx % BLOCK_SIZE;
let (j, shift) = if in_block < 16 {
(in_block, 0)
} else {
(in_block - 16, 4)
};
let block_base = block_idx * block_stride;
for (sub, &code) in v_codes.iter().enumerate() {
let nibble = code & 0x0F;
packed[block_base + sub * BYTES_PER_SUB_PER_BLOCK + j] |= nibble << shift;
}
}
packed
}
#[cfg(test)]
mod tests {
use super::*;
fn dummy_codebook(m: usize, sub_dim: usize) -> Vec<f32> {
let k = 16usize;
let len = m * k * sub_dim;
(0..len).map(|i| i as f32 * 0.01).collect()
}
#[test]
fn round_trip_packs_and_unpacks_codes_for_n_below_block() {
let m = 4;
let codes: Vec<Vec<u8>> = vec![vec![0, 1, 2, 3], vec![4, 5, 6, 7], vec![15, 0, 8, 9]];
let packed = pack_codes_into_blocks(&codes, m);
let params = PqParams::new(m as u16, 16, 2).unwrap();
let pool = PqFastScanPool::build(
params,
dummy_codebook(m, 2),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
assert_eq!(pool.packed, packed);
assert_eq!(pool.codes_at(0), vec![0, 1, 2, 3]);
assert_eq!(pool.codes_at(1), vec![4, 5, 6, 7]);
assert_eq!(pool.codes_at(2), vec![15, 0, 8, 9]);
}
#[test]
fn round_trip_packs_and_unpacks_codes_at_block_boundaries() {
let m = 8;
let n = 33;
let codes: Vec<Vec<u8>> = (0..n)
.map(|i| (0..m).map(|sub| ((i * 7 + sub * 3) % 16) as u8).collect())
.collect();
let params = PqParams::new(m as u16, 16, 2).unwrap();
let pool = PqFastScanPool::build(
params,
dummy_codebook(m, 2),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
assert_eq!(pool.block_count(), 2, "n=33 spans 2 blocks");
for (i, expected) in codes.iter().enumerate().take(n) {
assert_eq!(&pool.codes_at(i), expected, "vector {i} mismatch");
}
}
#[test]
fn round_trip_works_for_exactly_one_block() {
let m = 4;
let n = BLOCK_SIZE;
let codes: Vec<Vec<u8>> = (0..n)
.map(|i| (0..m).map(|sub| ((i + sub * 5) % 16) as u8).collect())
.collect();
let params = PqParams::new(m as u16, 16, 1).unwrap();
let pool = PqFastScanPool::build(
params,
dummy_codebook(m, 1),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
assert_eq!(pool.block_count(), 1);
for (i, expected) in codes.iter().enumerate().take(n) {
assert_eq!(&pool.codes_at(i), expected, "vector {i} mismatch");
}
}
#[test]
fn round_trip_works_for_two_full_blocks() {
let m = 2;
let n = BLOCK_SIZE * 2;
let codes: Vec<Vec<u8>> = (0..n)
.map(|i| vec![(i % 16) as u8, ((i / 2) % 16) as u8])
.collect();
let params = PqParams::new(m as u16, 16, 2).unwrap();
let pool = PqFastScanPool::build(
params,
dummy_codebook(m, 2),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
assert_eq!(pool.block_count(), 2);
for (i, expected) in codes.iter().enumerate().take(n) {
assert_eq!(&pool.codes_at(i), expected);
}
}
#[test]
fn build_rejects_wrong_k() {
let params = PqParams::new(4, 256, 2).unwrap();
let codebook = vec![0.0f32; params.codebook_len()];
let err = PqFastScanPool::build(
params,
codebook,
std::iter::empty::<(u64, String, Vec<u8>)>(),
)
.unwrap_err();
assert!(err.to_string().contains("k == 16"));
}
#[test]
fn build_rejects_out_of_range_codes() {
let m = 4;
let params = PqParams::new(m as u16, 16, 2).unwrap();
let codebook = dummy_codebook(m, 2);
let err = PqFastScanPool::build(
params,
codebook,
std::iter::once((0u64, "f".to_string(), vec![0u8, 1, 16, 3])),
)
.unwrap_err();
assert!(err.to_string().contains("exceeds 4-bit range"));
}
#[test]
fn build_rejects_code_length_mismatch() {
let m = 4;
let params = PqParams::new(m as u16, 16, 2).unwrap();
let codebook = dummy_codebook(m, 2);
let err = PqFastScanPool::build(
params,
codebook,
std::iter::once((0u64, "f".to_string(), vec![0u8, 1, 2])),
)
.unwrap_err();
assert!(err.to_string().contains("does not match params.m"));
}
#[test]
fn build_zero_pads_trailing_partial_block() {
let m = 2;
let n = 5; let codes: Vec<Vec<u8>> = (0..n)
.map(|i| vec![(i % 16) as u8, ((i + 1) % 16) as u8])
.collect();
let params = PqParams::new(m as u16, 16, 1).unwrap();
let pool = PqFastScanPool::build(
params,
dummy_codebook(m, 1),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
assert_eq!(pool.block_count(), 1);
assert_eq!(
pool.packed.len(),
pool.block_count() * pool.block_stride(),
"packed buffer is sized to whole blocks"
);
for (i, expected) in codes.iter().enumerate().take(n) {
assert_eq!(&pool.codes_at(i), expected);
}
for i in n..BLOCK_SIZE {
assert!(
pool.codes_at(i).iter().all(|&c| c == 0),
"vec {i} in padding should be all-zero codes"
);
}
}
}