use crate::{
Result,
codec::{EncodedVector, ResidualCodec, train_quantizer},
kmeans::{assign_points, fit},
};
#[derive(Debug, Clone, Default)]
pub struct InvertedFile {
pub lists: Vec<Vec<u32>>,
}
impl InvertedFile {
pub fn num_centroids(&self) -> usize {
self.lists.len()
}
pub fn docs_for_centroid(&self, centroid_id: usize) -> &[u32] {
self.lists
.get(centroid_id)
.map(Vec::as_slice)
.unwrap_or(&[])
}
pub fn total_doc_postings(&self) -> usize {
self.lists.iter().map(Vec::len).sum()
}
}
#[derive(Debug, Clone)]
pub struct DocumentTokens {
pub doc_id: u64,
pub tokens: Vec<f32>,
pub n_tokens: usize,
}
impl DocumentTokens {
pub fn flat_len(&self) -> usize {
self.tokens.len()
}
}
#[derive(Debug, Clone, Copy)]
pub struct IndexParams {
pub dim: usize,
pub nbits: u32,
pub k_centroids: usize,
pub max_kmeans_iters: usize,
}
#[derive(Debug, Clone)]
pub struct Index {
pub params: IndexParams,
pub codec: ResidualCodec,
pub doc_ids: Vec<u64>,
pub doc_tokens: Vec<Vec<EncodedVector>>,
pub ivf: InvertedFile,
}
impl Index {
pub fn num_documents(&self) -> usize {
self.doc_ids.len()
}
pub fn num_tokens(&self) -> usize {
self.doc_tokens.iter().map(Vec::len).sum()
}
pub fn position_of(&self, doc_id: u64) -> Option<usize> {
self.doc_ids.iter().position(|id| *id == doc_id)
}
}
pub fn build_index(
documents: &[DocumentTokens],
params: IndexParams,
) -> Result<Index> {
assert!(params.dim > 0, "build_index: dim must be positive");
assert!(
params.k_centroids > 0,
"build_index: k_centroids must be positive"
);
assert!(
params.nbits > 0 && params.nbits <= 8,
"build_index: nbits must be in 1..=8, got {}",
params.nbits,
);
for doc in documents {
assert!(
doc.tokens.len() == doc.n_tokens * params.dim,
"build_index: doc {} declared {} tokens but carries {} f32s (dim={})",
doc.doc_id,
doc.n_tokens,
doc.tokens.len(),
params.dim,
);
}
let total_tokens: usize = documents.iter().map(|d| d.n_tokens).sum();
assert!(
total_tokens >= params.k_centroids,
"build_index: need at least {} tokens for {} centroids, got {}",
params.k_centroids,
params.k_centroids,
total_tokens,
);
let mut pool: Vec<f32> = Vec::with_capacity(total_tokens * params.dim);
for doc in documents {
pool.extend_from_slice(&doc.tokens);
}
let centroids = fit(
&pool,
params.k_centroids,
params.dim,
params.max_kmeans_iters.max(1),
)?;
let assignments = assign_points(&pool, ¢roids, params.dim)?;
let mut residual_sample: Vec<f32> = Vec::with_capacity(pool.len());
for (token, &cluster) in pool.chunks_exact(params.dim).zip(&assignments) {
let centroid =
¢roids[cluster * params.dim..(cluster + 1) * params.dim];
for (t, c) in token.iter().zip(centroid) {
residual_sample.push(*t - *c);
}
}
let (bucket_cutoffs, bucket_weights) =
train_quantizer(&residual_sample, params.nbits);
let codec = ResidualCodec {
nbits: params.nbits,
dim: params.dim,
centroids,
bucket_cutoffs,
bucket_weights,
};
codec.validate()?;
let (all_centroid_ids, all_codes) = codec.batch_encode_tokens(&pool)?;
let packed_per_token = codec.packed_bytes();
let mut doc_ids = Vec::with_capacity(documents.len());
let mut doc_tokens = Vec::with_capacity(documents.len());
let mut token_offset = 0usize;
for doc in documents {
doc_ids.push(doc.doc_id);
let n_tok = doc.n_tokens;
let cids = &all_centroid_ids[token_offset..token_offset + n_tok];
let codes_slice = &all_codes[token_offset * packed_per_token
..(token_offset + n_tok) * packed_per_token];
let encoded: Vec<EncodedVector> = (0..n_tok)
.map(|i| EncodedVector {
centroid_id: cids[i],
codes: codes_slice
[i * packed_per_token..(i + 1) * packed_per_token]
.to_vec(),
})
.collect();
doc_tokens.push(encoded);
token_offset += n_tok;
}
let ivf = build_inverted_file(&doc_tokens, params.k_centroids);
Ok(Index {
params,
codec,
doc_ids,
doc_tokens,
ivf,
})
}
pub(crate) fn build_inverted_file(
doc_tokens: &[Vec<EncodedVector>],
k_centroids: usize,
) -> InvertedFile {
let mut lists: Vec<Vec<u32>> = vec![Vec::new(); k_centroids];
for (doc_idx, encoded) in doc_tokens.iter().enumerate() {
let mut touched: Vec<u32> =
encoded.iter().map(|ev| ev.centroid_id).collect();
touched.sort_unstable();
touched.dedup();
for cid in touched {
lists[cid as usize].push(doc_idx as u32);
}
}
InvertedFile { lists }
}
#[cfg(test)]
mod tests {
use super::*;
use crate::distance::squared_l2;
fn small_corpus() -> Vec<DocumentTokens> {
vec![
DocumentTokens {
doc_id: 1,
tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
n_tokens: 3,
},
DocumentTokens {
doc_id: 2,
tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
n_tokens: 3,
},
DocumentTokens {
doc_id: 3,
tokens: vec![0.3, -0.2, 9.7, 10.2],
n_tokens: 2,
},
]
}
fn default_params() -> IndexParams {
IndexParams {
dim: 2,
nbits: 2,
k_centroids: 2,
max_kmeans_iters: 50,
}
}
#[test]
fn build_index_encodes_every_token() {
let docs = small_corpus();
let params = default_params();
let expected_total: usize = docs.iter().map(|d| d.n_tokens).sum();
let index = build_index(&docs, params).unwrap();
assert_eq!(index.num_documents(), docs.len());
assert_eq!(index.num_tokens(), expected_total);
for (encoded, doc) in index.doc_tokens.iter().zip(docs.iter()) {
assert_eq!(encoded.len(), doc.n_tokens);
}
}
#[test]
fn build_index_assigns_tokens_in_each_cluster_to_the_closest_centroid() {
let docs = small_corpus();
let params = default_params();
let index = build_index(&docs, params).unwrap();
let c0 = &index.codec.centroids[0..2];
let c1 = &index.codec.centroids[2..4];
let (near_origin, near_ten) =
if squared_l2(c0, &[0.0, 0.0]) < squared_l2(c0, &[10.0, 10.0]) {
(c0, c1)
} else {
(c1, c0)
};
assert!(squared_l2(near_origin, &[0.0, 0.0]) < 1.0);
assert!(squared_l2(near_ten, &[10.0, 10.0]) < 1.0);
}
#[test]
fn build_index_round_trip_reconstruction_error_is_bounded() {
let docs = small_corpus();
let params = default_params();
let index = build_index(&docs, params).unwrap();
for (doc, encoded_doc) in docs.iter().zip(index.doc_tokens.iter()) {
for (token, encoded) in
doc.tokens.chunks_exact(params.dim).zip(encoded_doc.iter())
{
let decoded = index.codec.decode_vector(encoded).unwrap();
let err = squared_l2(token, &decoded).sqrt();
assert!(
err < 0.6,
"reconstruction error {err} too large for token {token:?}"
);
}
}
}
#[test]
fn build_index_preserves_document_id_order() {
let docs = small_corpus();
let index = build_index(&docs, default_params()).unwrap();
let expected_ids: Vec<u64> = docs.iter().map(|d| d.doc_id).collect();
assert_eq!(index.doc_ids, expected_ids);
assert_eq!(index.position_of(2), Some(1));
assert_eq!(index.position_of(999), None);
}
#[test]
fn build_index_handles_document_with_no_tokens() {
let mut docs = small_corpus();
docs.push(DocumentTokens {
doc_id: 42,
tokens: vec![],
n_tokens: 0,
});
let index = build_index(&docs, default_params()).unwrap();
assert_eq!(index.num_documents(), 4);
assert_eq!(index.doc_tokens[3].len(), 0);
}
#[test]
#[should_panic(expected = "declared")]
fn build_index_panics_on_mismatched_token_count() {
let docs = vec![DocumentTokens {
doc_id: 1,
tokens: vec![0.0, 0.0, 1.0],
n_tokens: 2, }];
let _ = build_index(&docs, default_params()).unwrap();
}
#[test]
fn build_index_ivf_has_one_list_per_centroid() {
let docs = small_corpus();
let index = build_index(&docs, default_params()).unwrap();
assert_eq!(
index.ivf.num_centroids(),
default_params().k_centroids,
"IVF has one list per centroid",
);
}
#[test]
fn build_index_ivf_postings_cover_every_doc_that_touches_each_centroid() {
let docs = small_corpus();
let index = build_index(&docs, default_params()).unwrap();
for (doc_idx, encoded_doc) in index.doc_tokens.iter().enumerate() {
for ev in encoded_doc {
let postings =
index.ivf.docs_for_centroid(ev.centroid_id as usize);
assert!(
postings.contains(&(doc_idx as u32)),
"doc_idx={doc_idx} missing from centroid {} postings",
ev.centroid_id,
);
}
}
}
#[test]
fn build_index_ivf_postings_are_unique_per_centroid() {
let docs = small_corpus();
let index = build_index(&docs, default_params()).unwrap();
for c in 0..index.ivf.num_centroids() {
let postings = index.ivf.docs_for_centroid(c);
let mut unique: Vec<u32> = postings.to_vec();
unique.sort_unstable();
unique.dedup();
assert_eq!(
unique.len(),
postings.len(),
"centroid {c} has duplicate doc entries: {postings:?}",
);
}
}
#[test]
fn build_index_dedupes_repeated_tokens_in_same_centroid() {
let docs = vec![
DocumentTokens {
doc_id: 1,
tokens: vec![0.0, 0.0, 0.05, -0.02, -0.03, 0.01],
n_tokens: 3,
},
DocumentTokens {
doc_id: 2,
tokens: vec![10.0, 10.0, 10.1, 9.9],
n_tokens: 2,
},
];
let index = build_index(&docs, default_params()).unwrap();
for c in 0..index.ivf.num_centroids() {
let postings = index.ivf.docs_for_centroid(c);
let count_of_doc_0 = postings.iter().filter(|&&d| d == 0).count();
assert!(
count_of_doc_0 <= 1,
"doc 0 appears {count_of_doc_0} times in centroid {c}",
);
}
}
#[test]
fn inverted_file_out_of_range_returns_empty_slice() {
let ivf = InvertedFile {
lists: vec![vec![0u32]],
};
assert_eq!(ivf.docs_for_centroid(0).len(), 1);
assert!(ivf.docs_for_centroid(999).is_empty());
}
#[test]
#[should_panic(expected = "at least")]
fn build_index_panics_when_too_few_tokens_for_k() {
let docs = vec![DocumentTokens {
doc_id: 1,
tokens: vec![0.0, 1.0],
n_tokens: 1,
}];
let params = IndexParams {
dim: 2,
nbits: 2,
k_centroids: 4,
max_kmeans_iters: 10,
};
let _ = build_index(&docs, params).unwrap();
}
}