use candle_core::{Device, Tensor};
use crate::{
Result,
codec::DecodeTable,
device::default_device,
distance::dot,
index::Index,
};
#[derive(Debug, Clone, Copy)]
pub struct SearchParams {
pub top_k: usize,
pub n_probe: usize,
pub n_candidate_docs: Option<usize>,
pub centroid_score_threshold: Option<f32>,
}
impl SearchParams {
pub const fn paper_defaults(top_k: usize) -> Self {
let (n_probe, t_cs, ndocs) = if top_k <= 10 {
(1, 0.5_f32, 256)
} else if top_k <= 100 {
(2, 0.45_f32, 1024)
} else {
(4, 0.4_f32, 4096)
};
let ndocs = if ndocs < top_k.saturating_mul(4) {
top_k.saturating_mul(4)
} else {
ndocs
};
Self {
top_k,
n_probe,
n_candidate_docs: Some(ndocs),
centroid_score_threshold: Some(t_cs),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct SearchResult {
pub doc_id: u64,
pub score: f32,
}
pub fn search(
index: &Index,
query_tokens: &[f32],
params: SearchParams,
) -> Result<Vec<SearchResult>> {
let dim = index.params.dim;
assert!(params.top_k > 0, "search: top_k must be positive");
assert!(params.n_probe > 0, "search: n_probe must be positive");
assert!(
query_tokens.len().is_multiple_of(dim),
"search: query length {} is not a multiple of dim {}",
query_tokens.len(),
dim,
);
if query_tokens.is_empty() || index.num_documents() == 0 {
return Ok(Vec::new());
}
let n_centroids = index.codec.num_centroids();
let n_probe = params.n_probe.min(n_centroids);
let qc_scores: Option<Vec<f32>> = (params
.centroid_score_threshold
.is_some()
|| params.n_candidate_docs.is_some())
.then(|| {
query_centroid_score_matrix(query_tokens, &index.codec.centroids, dim)
});
let pruned_mask: Option<Vec<bool>> =
match (qc_scores.as_ref(), params.centroid_score_threshold) {
(Some(scores), Some(threshold)) => {
let per_cent = per_centroid_max_scores(
scores,
query_tokens.len() / dim,
n_centroids,
);
Some(per_cent.iter().map(|&s| s < threshold).collect())
}
_ => None,
};
let mut candidate_docs: Vec<bool> = vec![false; index.num_documents()];
for query_token in query_tokens.chunks_exact(dim) {
for centroid_id in
top_n_centroids(query_token, &index.codec.centroids, dim, n_probe)
{
if let Some(mask) = pruned_mask.as_ref()
&& mask[centroid_id]
{
continue;
}
for &doc_idx in index.ivf.docs_for_centroid(centroid_id) {
candidate_docs[doc_idx as usize] = true;
}
}
}
let mut candidate_idxs: Vec<usize> = candidate_docs
.iter()
.enumerate()
.filter_map(|(idx, &is_cand)| {
(is_cand && index.doc_token_count(idx) > 0).then_some(idx)
})
.collect();
if let (Some(n_stage2), Some(qc_scores)) =
(params.n_candidate_docs, qc_scores.as_ref())
{
let n_q = query_tokens.len() / dim;
let n_c = index.codec.num_centroids();
candidate_idxs = shortlist_by_approx_score(
candidate_idxs,
index,
qc_scores,
n_q,
n_c,
pruned_mask.as_deref(),
n_stage2,
);
let n_stage3 = n_stage2.div_ceil(4).max(1);
candidate_idxs = shortlist_by_approx_score(
candidate_idxs,
index,
qc_scores,
n_q,
n_c,
None,
n_stage3,
);
}
let mut scored: Vec<SearchResult> = if candidate_idxs.is_empty() {
Vec::new()
} else {
batch_maxsim(query_tokens, &candidate_idxs, index, dim)?
.into_iter()
.map(|(doc_idx, score)| SearchResult {
doc_id: index.doc_ids[doc_idx],
score,
})
.collect()
};
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.doc_id.cmp(&b.doc_id))
});
scored.truncate(params.top_k);
Ok(scored)
}
pub fn top_n_centroids(
point: &[f32],
centroids: &[f32],
dim: usize,
n: usize,
) -> Vec<usize> {
assert_eq!(
point.len(),
dim,
"top_n_centroids: point length {} does not match dim {}",
point.len(),
dim,
);
assert!(dim > 0, "top_n_centroids: dim must be positive");
assert!(
!centroids.is_empty() && centroids.len().is_multiple_of(dim),
"top_n_centroids: centroids length {} is not a positive multiple of dim {}",
centroids.len(),
dim,
);
let k = centroids.len() / dim;
let mut scored: Vec<(usize, f32)> = centroids
.chunks_exact(dim)
.enumerate()
.map(|(i, c)| (i, dot(point, c)))
.collect();
scored.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
scored.into_iter().take(n.min(k)).map(|(i, _)| i).collect()
}
fn query_centroid_score_matrix(
query_tokens: &[f32],
centroids: &[f32],
dim: usize,
) -> Vec<f32> {
let n_q = query_tokens.len() / dim;
let n_c = centroids.len() / dim;
let mut out = vec![0.0f32; n_q * n_c];
for (qi, q) in query_tokens.chunks_exact(dim).enumerate() {
for (ci, c) in centroids.chunks_exact(dim).enumerate() {
out[qi * n_c + ci] = dot(q, c);
}
}
out
}
fn shortlist_by_approx_score(
candidate_idxs: Vec<usize>,
index: &Index,
qc_scores: &[f32],
n_q: usize,
n_centroids: usize,
mask: Option<&[bool]>,
limit: usize,
) -> Vec<usize> {
if candidate_idxs.len() <= limit {
return candidate_idxs;
}
let mut approx: Vec<(usize, f32)> = candidate_idxs
.into_iter()
.map(|doc_idx| {
let score = approx_centroid_interaction_score(
index.doc_centroid_ids(doc_idx),
qc_scores,
n_q,
n_centroids,
mask,
);
(doc_idx, score)
})
.collect();
approx.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
approx.truncate(limit);
approx.into_iter().map(|(idx, _)| idx).collect()
}
fn per_centroid_max_scores(
qc_scores: &[f32],
n_q: usize,
n_centroids: usize,
) -> Vec<f32> {
let mut out = vec![f32::NEG_INFINITY; n_centroids];
for q in 0..n_q {
let row = &qc_scores[q * n_centroids..(q + 1) * n_centroids];
for (c, &s) in row.iter().enumerate() {
if s > out[c] {
out[c] = s;
}
}
}
out
}
fn approx_centroid_interaction_score(
doc_centroid_ids: &[u32],
qc_scores: &[f32],
n_q: usize,
n_centroids: usize,
pruned_mask: Option<&[bool]>,
) -> f32 {
if doc_centroid_ids.is_empty() {
return 0.0;
}
if let Some(mask) = pruned_mask
&& doc_centroid_ids.iter().all(|cid| mask[*cid as usize])
{
return f32::NEG_INFINITY;
}
let mut total = 0.0f32;
for q in 0..n_q {
let row = &qc_scores[q * n_centroids..(q + 1) * n_centroids];
let mut best = f32::NEG_INFINITY;
for &cid in doc_centroid_ids {
if let Some(mask) = pruned_mask
&& mask[cid as usize]
{
continue;
}
let s = row[cid as usize];
if s > best {
best = s;
}
}
if best.is_finite() {
total += best;
}
}
total
}
fn batch_maxsim(
query_tokens: &[f32],
candidate_idxs: &[usize],
index: &Index,
dim: usize,
) -> Result<Vec<(usize, f32)>> {
batch_maxsim_with_cap(
query_tokens,
candidate_idxs,
index,
dim,
decode_chunk_capacity(dim),
)
}
const DECODE_CHUNK_BUDGET_BYTES: usize = 64 * 1024 * 1024;
fn decode_chunk_capacity(dim: usize) -> usize {
let per_token = dim.saturating_mul(4 * 5).max(1);
(DECODE_CHUNK_BUDGET_BYTES / per_token).max(64)
}
fn batch_maxsim_with_cap(
query_tokens: &[f32],
candidate_idxs: &[usize],
index: &Index,
dim: usize,
max_tokens_per_chunk: usize,
) -> Result<Vec<(usize, f32)>> {
let n_q = query_tokens.len() / dim;
let decode_table = DecodeTable::new(&index.codec);
let total_tokens: usize = candidate_idxs
.iter()
.map(|&i| index.doc_token_count(i))
.sum();
if total_tokens == 0 || n_q == 0 {
return Ok(candidate_idxs.iter().map(|&i| (i, 0.0)).collect());
}
let device = default_device();
let q_t = Tensor::from_slice(query_tokens, (n_q, dim), device)?;
let q_transposed = q_t.t()?.contiguous()?;
let mut out = Vec::with_capacity(candidate_idxs.len());
let mut start = 0usize;
while start < candidate_idxs.len() {
let mut end = start + 1;
let mut chunk_tokens = index.doc_token_count(candidate_idxs[start]);
while end < candidate_idxs.len() {
let next = index.doc_token_count(candidate_idxs[end]);
if chunk_tokens + next > max_tokens_per_chunk {
break;
}
chunk_tokens += next;
end += 1;
}
let chunk_candidates = &candidate_idxs[start..end];
let mut chunk_offsets: Vec<usize> =
Vec::with_capacity(chunk_candidates.len() + 1);
chunk_offsets.push(0);
let chunk_ctx = DecodeCtx {
index,
candidate_idxs: chunk_candidates,
dim,
total_tokens: chunk_tokens,
decode_table: &decode_table,
device,
};
let decoded = if matches!(device, Device::Cpu) {
decode_on_cpu(&chunk_ctx, &mut chunk_offsets)?
} else {
decode_on_device(&chunk_ctx, &mut chunk_offsets)?
};
let scores_flat: Vec<f32> = decoded
.matmul(&q_transposed)?
.flatten_all()?
.to_vec1::<f32>()?;
for (i, &doc_idx) in chunk_candidates.iter().enumerate() {
let range_start = chunk_offsets[i];
let range_end = chunk_offsets[i + 1];
if range_start == range_end {
out.push((doc_idx, 0.0));
continue;
}
let mut total = 0.0f32;
for q in 0..n_q {
let mut best = f32::NEG_INFINITY;
for t in range_start..range_end {
let s = scores_flat[t * n_q + q];
if s > best {
best = s;
}
}
if best.is_finite() {
total += best;
}
}
out.push((doc_idx, total));
}
start = end;
}
Ok(out)
}
struct DecodeCtx<'a> {
index: &'a Index,
candidate_idxs: &'a [usize],
dim: usize,
total_tokens: usize,
decode_table: &'a DecodeTable,
device: &'a Device,
}
fn decode_on_cpu(
ctx: &DecodeCtx<'_>,
offsets: &mut Vec<usize>,
) -> Result<Tensor> {
let codec = &ctx.index.codec;
let packed_bytes = codec.packed_bytes();
let mut packed: Vec<f32> = Vec::with_capacity(ctx.total_tokens * ctx.dim);
let mut ev = crate::codec::EncodedVector {
centroid_id: 0,
codes: vec![0u8; packed_bytes],
};
for &doc_idx in ctx.candidate_idxs {
let cids = ctx.index.doc_centroid_ids(doc_idx);
let bytes = ctx.index.doc_residual_bytes(doc_idx);
for (i, &cid) in cids.iter().enumerate() {
ev.centroid_id = cid;
ev.codes.copy_from_slice(
&bytes[i * packed_bytes..(i + 1) * packed_bytes],
);
let decoded =
codec.decode_vector_with_table(&ev, ctx.decode_table)?;
let norm = decoded
.iter()
.map(|v| v * v)
.sum::<f32>()
.sqrt()
.max(1e-12_f32);
packed.extend(decoded.iter().map(|v| v / norm));
}
offsets.push(packed.len() / ctx.dim);
}
Ok(Tensor::from_vec(
packed,
(ctx.total_tokens, ctx.dim),
ctx.device,
)?)
}
fn decode_on_device(
ctx: &DecodeCtx<'_>,
offsets: &mut Vec<usize>,
) -> Result<Tensor> {
let packed_bytes_per_vec = ctx.index.codec.packed_bytes();
let codes_per_byte = ctx.decode_table.codes_per_byte();
let decoded_row_width = packed_bytes_per_vec * codes_per_byte;
debug_assert!(decoded_row_width >= ctx.dim);
let mut centroid_ids: Vec<u32> = Vec::with_capacity(ctx.total_tokens);
let mut packed_bytes: Vec<u32> =
Vec::with_capacity(ctx.total_tokens * packed_bytes_per_vec);
for &doc_idx in ctx.candidate_idxs {
let doc_cids = ctx.index.doc_centroid_ids(doc_idx);
let doc_bytes = ctx.index.doc_residual_bytes(doc_idx);
centroid_ids.extend_from_slice(doc_cids);
packed_bytes.extend(doc_bytes.iter().map(|&b| b as u32));
offsets.push(centroid_ids.len());
}
let n_centroids = ctx.index.codec.num_centroids();
let centroids_t = Tensor::from_slice(
ctx.index.codec.centroids.as_slice(),
(n_centroids, ctx.dim),
ctx.device,
)?;
let weights_t = Tensor::from_slice(
ctx.decode_table.weights_flat(),
(256, codes_per_byte),
ctx.device,
)?;
let cent_idx_t =
Tensor::from_vec(centroid_ids, (ctx.total_tokens,), ctx.device)?;
let byte_idx_t = Tensor::from_vec(
packed_bytes,
(ctx.total_tokens * packed_bytes_per_vec,),
ctx.device,
)?;
let centroid_emb = centroids_t.index_select(¢_idx_t, 0)?;
let residuals_flat = weights_t.index_select(&byte_idx_t, 0)?;
let residuals_padded =
residuals_flat.reshape((ctx.total_tokens, decoded_row_width))?;
let residuals = if decoded_row_width == ctx.dim {
residuals_padded
} else {
residuals_padded.narrow(1, 0, ctx.dim)?
};
let decoded = (centroid_emb + residuals)?;
l2_normalize_rows(&decoded)
}
fn l2_normalize_rows(decoded: &Tensor) -> Result<Tensor> {
let squared = decoded.sqr()?;
let norm_sq = squared.sum_keepdim(1)?; let norm = norm_sq.sqrt()?.clamp(1e-12f32, f32::INFINITY)?;
Ok(decoded.broadcast_div(&norm)?)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
distance::dot,
index::{DocumentTokens, IndexParams, build_index},
};
fn params() -> IndexParams {
IndexParams {
dim: 2,
nbits: 2,
k_centroids: 2,
max_kmeans_iters: 50,
}
}
fn corpus() -> Vec<DocumentTokens> {
let east_a = [1.0f32, 0.0];
let east_b = normalize([0.98, 0.2]);
let east_c = normalize([0.97, -0.24]);
let north_a = [0.0f32, 1.0];
let north_b = normalize([0.2, 0.98]);
let north_c = normalize([-0.24, 0.97]);
let mut doc1 = Vec::new();
doc1.extend_from_slice(&east_a);
doc1.extend_from_slice(&east_b);
doc1.extend_from_slice(&east_c);
let mut doc2 = Vec::new();
doc2.extend_from_slice(&north_a);
doc2.extend_from_slice(&north_b);
doc2.extend_from_slice(&north_c);
let mut doc3 = Vec::new();
doc3.extend_from_slice(&east_a);
doc3.extend_from_slice(&north_a);
vec![
DocumentTokens {
doc_id: 1,
tokens: doc1,
n_tokens: 3,
},
DocumentTokens {
doc_id: 2,
tokens: doc2,
n_tokens: 3,
},
DocumentTokens {
doc_id: 3,
tokens: doc3,
n_tokens: 2,
},
]
}
fn normalize(v: [f32; 2]) -> [f32; 2] {
let norm = (v[0] * v[0] + v[1] * v[1]).sqrt();
[v[0] / norm, v[1] / norm]
}
#[test]
fn top_n_centroids_ranks_by_descending_dot_product() {
let centroids = [0.0, 0.0, 5.0, 0.0, 12.0, 0.0];
let out = top_n_centroids(&[4.0, 0.0], ¢roids, 2, 3);
assert_eq!(out, vec![2, 1, 0]);
}
#[test]
fn top_n_centroids_breaks_ties_toward_earlier_index() {
let centroids = [1.0, 0.0, 0.0, 1.0, -1.0, 0.0];
let out = top_n_centroids(&[0.0, 0.0], ¢roids, 2, 3);
assert_eq!(out, vec![0, 1, 2]);
}
#[test]
fn top_n_centroids_on_unit_norm_centroids_picks_most_aligned() {
let centroids = [
0.6, 0.8, 0.0, 1.0, -1.0, 0.0, 1.0, 0.0,
];
let out = top_n_centroids(&[1.0, 0.0], ¢roids, 2, 2);
assert_eq!(out, vec![3, 0]);
}
#[test]
fn top_n_centroids_caps_at_num_centroids() {
let centroids = [0.0, 0.0, 5.0, 0.0];
let out = top_n_centroids(&[0.0, 0.0], ¢roids, 2, 10);
assert_eq!(out, vec![0, 1]);
}
#[test]
fn search_returns_empty_for_empty_query() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[],
SearchParams {
top_k: 3,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert!(out.is_empty());
}
#[test]
fn search_ranks_matching_cluster_highest() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 3,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert!(!out.is_empty(), "should surface at least one doc");
assert_eq!(out[0].doc_id, 1, "closest doc should rank first");
}
#[test]
fn search_respects_top_k() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 1,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert_eq!(out.len(), 1);
}
#[test]
fn search_scores_are_non_increasing() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[1.0, 0.0, 0.0, 1.0],
SearchParams {
top_k: 3,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
for pair in out.windows(2) {
assert!(
pair[0].score >= pair[1].score,
"scores must be descending: {pair:?}"
);
}
}
#[test]
fn search_with_single_probe_still_finds_the_right_cluster() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[0.0, 1.0],
SearchParams {
top_k: 1,
n_probe: 1,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert_eq!(out[0].doc_id, 2);
}
#[test]
fn search_skips_documents_with_no_tokens() {
let mut docs = corpus();
docs.push(DocumentTokens {
doc_id: 42,
tokens: vec![],
n_tokens: 0,
});
let index = build_index(&docs, params()).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 10,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert!(
out.iter().all(|r| r.doc_id != 42),
"empty doc should not appear in results"
);
}
#[test]
fn search_score_equals_sum_of_maxsim_over_decoded_tokens() {
let index = build_index(&corpus(), params()).unwrap();
let query = [1.0f32, 0.0];
let out = search(
&index,
&query,
SearchParams {
top_k: 1,
n_probe: index.codec.num_centroids(),
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert!(!out.is_empty());
let top = out[0];
let doc_idx = index.position_of(top.doc_id).expect("doc present");
let decoded: Vec<Vec<f32>> = index
.doc_tokens_vec(doc_idx)
.iter()
.map(|ev| {
let raw = index.codec.decode_vector(ev).unwrap();
let norm =
raw.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
raw.iter().map(|v| v / norm).collect()
})
.collect();
let mut expected = 0.0f32;
for q in query.chunks_exact(2) {
let best = decoded
.iter()
.map(|d| dot(q, d))
.fold(f32::NEG_INFINITY, f32::max);
expected += best;
}
assert!(
(expected - top.score).abs() < 1e-5,
"score {} differs from ground-truth MaxSim {}",
top.score,
expected,
);
}
#[test]
fn search_works_with_four_bit_residuals() {
let params_4bit = IndexParams {
dim: 2,
nbits: 4,
k_centroids: 2,
max_kmeans_iters: 50,
};
let index = build_index(&corpus(), params_4bit).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 1,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
assert_eq!(out[0].doc_id, 1);
}
#[test]
fn search_survives_large_synthetic_corpus_on_both_clusters() {
let mut docs = Vec::new();
for i in 0..20 {
let jitter = i as f32 * 0.003;
let east = normalize([1.0 - jitter, 0.05 + jitter]);
docs.push(DocumentTokens {
doc_id: 100 + i,
tokens: east.to_vec(),
n_tokens: 1,
});
let north = normalize([0.05 + jitter, 1.0 - jitter]);
docs.push(DocumentTokens {
doc_id: 200 + i,
tokens: north.to_vec(),
n_tokens: 1,
});
}
let index = build_index(&docs, params()).unwrap();
let east = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 5,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
for r in &east {
assert!(
(100..120).contains(&r.doc_id),
"east query should only surface east docs: got {}",
r.doc_id,
);
}
let north = search(
&index,
&[0.0, 1.0],
SearchParams {
top_k: 5,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
for r in &north {
assert!(
(200..220).contains(&r.doc_id),
"north query should only surface north docs: got {}",
r.doc_id,
);
}
}
#[test]
#[should_panic(expected = "top_k must be positive")]
fn search_panics_on_zero_top_k() {
let index = build_index(&corpus(), params()).unwrap();
let _ = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 0,
n_probe: 1,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
}
#[test]
#[should_panic(expected = "n_probe must be positive")]
fn search_panics_on_zero_n_probe() {
let index = build_index(&corpus(), params()).unwrap();
let _ = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 1,
n_probe: 0,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
}
#[test]
#[should_panic(expected = "query length")]
fn search_panics_on_ragged_query() {
let index = build_index(&corpus(), params()).unwrap();
let _ = search(
&index,
&[1.0, 0.0, 0.5],
SearchParams {
top_k: 1,
n_probe: 1,
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
}
#[test]
fn centroid_interaction_shortlist_caps_candidates_surviving_to_decode() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 5,
n_probe: index.codec.num_centroids(),
n_candidate_docs: Some(1),
centroid_score_threshold: None,
},
)
.unwrap();
assert_eq!(out.len(), 1, "shortlist of 1 must cap output at 1");
assert_eq!(
out[0].doc_id, 1,
"east query should surface the east doc as the survivor",
);
}
#[test]
fn centroid_interaction_with_large_shortlist_matches_no_shortlist() {
let index = build_index(&corpus(), params()).unwrap();
let query = [1.0, 0.0, 0.0, 1.0];
let legacy = search(
&index,
&query,
SearchParams {
top_k: 3,
n_probe: index.codec.num_centroids(),
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
let shortlisted = search(
&index,
&query,
SearchParams {
top_k: 3,
n_probe: index.codec.num_centroids(),
n_candidate_docs: Some(1_000),
centroid_score_threshold: None,
},
)
.unwrap();
assert_eq!(legacy, shortlisted);
}
#[test]
fn centroid_pruning_with_low_threshold_matches_no_pruning() {
let index = build_index(&corpus(), params()).unwrap();
let query = [1.0, 0.0, 0.0, 1.0];
let unpruned = search(
&index,
&query,
SearchParams {
top_k: 3,
n_probe: index.codec.num_centroids(),
n_candidate_docs: None,
centroid_score_threshold: None,
},
)
.unwrap();
let pruned = search(
&index,
&query,
SearchParams {
top_k: 3,
n_probe: index.codec.num_centroids(),
n_candidate_docs: None,
centroid_score_threshold: Some(-1.0),
},
)
.unwrap();
assert_eq!(unpruned, pruned);
}
#[test]
fn two_stage_interaction_caps_survivors_to_ndocs_div_four() {
let mut docs = Vec::new();
for i in 0..20 {
let jitter = i as f32 * 0.003;
let east = normalize([1.0 - jitter, 0.05 + jitter]);
docs.push(DocumentTokens {
doc_id: 100 + i,
tokens: east.to_vec(),
n_tokens: 1,
});
}
let index = build_index(&docs, params()).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 20,
n_probe: index.codec.num_centroids(),
n_candidate_docs: Some(16),
centroid_score_threshold: None,
},
)
.unwrap();
assert!(
out.len() <= 4,
"Stage 3 must cap the decoded set at ndocs/4 = 4; got {}",
out.len(),
);
assert!(
!out.is_empty(),
"Stage 3 should still return at least one result"
);
}
#[test]
fn approx_score_skips_tokens_whose_centroid_is_pruned() {
let n_q = 2;
let n_c = 3;
let qc = [0.9, 0.0, 0.4, 0.0, 0.9, 0.4];
let doc_cids: [u32; 2] = [0, 2];
let unpruned =
approx_centroid_interaction_score(&doc_cids, &qc, n_q, n_c, None);
assert!((unpruned - 1.3).abs() < 1e-5, "unpruned was {unpruned}");
let pruned_mask = [false, false, true];
let pruned = approx_centroid_interaction_score(
&doc_cids,
&qc,
n_q,
n_c,
Some(&pruned_mask),
);
assert!((pruned - 0.9).abs() < 1e-5, "pruned was {pruned}");
}
#[test]
fn approx_score_on_all_pruned_doc_is_negative_infinity() {
let qc = [0.5, 0.5];
let doc_cids: [u32; 1] = [0];
let mask = [true];
let score = approx_centroid_interaction_score(
&doc_cids,
&qc,
1,
1,
Some(&mask),
);
assert!(
score == f32::NEG_INFINITY,
"all-pruned doc should score -inf, got {score}",
);
}
#[test]
fn centroid_pruning_excludes_docs_whose_only_centroid_is_below_threshold() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[1.0f32, 0.0],
SearchParams {
top_k: 5,
n_probe: index.codec.num_centroids(),
n_candidate_docs: None,
centroid_score_threshold: Some(0.5),
},
)
.unwrap();
let ids: Vec<u64> = out.iter().map(|r| r.doc_id).collect();
assert!(
!ids.contains(&2),
"pure-north doc must be pruned when only the east centroid survives the threshold: got {ids:?}",
);
assert!(
ids.contains(&1),
"east-cluster doc must still rank when its centroid survives: got {ids:?}",
);
}
#[test]
fn centroid_pruning_with_unreachable_threshold_drops_every_candidate() {
let index = build_index(&corpus(), params()).unwrap();
let out = search(
&index,
&[1.0, 0.0],
SearchParams {
top_k: 3,
n_probe: index.codec.num_centroids(),
n_candidate_docs: None,
centroid_score_threshold: Some(10.0),
},
)
.unwrap();
assert!(
out.is_empty(),
"every centroid below threshold should yield empty results",
);
}
#[test]
fn centroid_interaction_preserves_exact_maxsim_on_surviving_candidates() {
let index = build_index(&corpus(), params()).unwrap();
let query = [1.0f32, 0.0];
let out = search(
&index,
&query,
SearchParams {
top_k: 1,
n_probe: index.codec.num_centroids(),
n_candidate_docs: Some(1_000),
centroid_score_threshold: None,
},
)
.unwrap();
let top = out[0];
let doc_idx = index.position_of(top.doc_id).unwrap();
let decoded: Vec<Vec<f32>> = index
.doc_tokens_vec(doc_idx)
.iter()
.map(|ev| {
let raw = index.codec.decode_vector(ev).unwrap();
let norm =
raw.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
raw.iter().map(|v| v / norm).collect()
})
.collect();
let expected: f32 = query
.chunks_exact(2)
.map(|q| {
decoded
.iter()
.map(|d| dot(q, d))
.fold(f32::NEG_INFINITY, f32::max)
})
.sum();
assert!((expected - top.score).abs() < 1e-5);
}
#[test]
fn gpu_decode_path_matches_cpu_decode_element_wise() {
let index = build_index(&corpus(), params()).unwrap();
let candidate_idxs: Vec<usize> = (0..index.num_documents()).collect();
let total_tokens: usize = candidate_idxs
.iter()
.map(|&i| index.doc_token_count(i))
.sum();
let decode_table = DecodeTable::new(&index.codec);
let device = Device::Cpu;
let ctx = DecodeCtx {
index: &index,
candidate_idxs: &candidate_idxs,
dim: index.params.dim,
total_tokens,
decode_table: &decode_table,
device: &device,
};
let mut offsets_gpu = vec![0usize];
let decoded_gpu = decode_on_device(&ctx, &mut offsets_gpu).unwrap();
let gpu_rows: Vec<f32> = decoded_gpu
.to_vec2::<f32>()
.unwrap()
.into_iter()
.flatten()
.collect();
let mut offsets_cpu = vec![0usize];
let decoded_cpu = decode_on_cpu(&ctx, &mut offsets_cpu).unwrap();
let cpu_rows: Vec<f32> = decoded_cpu
.to_vec2::<f32>()
.unwrap()
.into_iter()
.flatten()
.collect();
assert_eq!(offsets_gpu, offsets_cpu);
assert_eq!(gpu_rows.len(), cpu_rows.len());
for (g, c) in gpu_rows.iter().zip(cpu_rows.iter()) {
assert!(
(g - c).abs() < 1e-6,
"GPU decode value {g} != CPU decode value {c}",
);
}
}
#[test]
fn paper_defaults_match_table_2_verbatim() {
let p10 = SearchParams::paper_defaults(10);
assert_eq!(p10.n_probe, 1);
assert_eq!(p10.centroid_score_threshold, Some(0.5));
assert_eq!(p10.n_candidate_docs, Some(256));
let p100 = SearchParams::paper_defaults(100);
assert_eq!(p100.n_probe, 2);
assert_eq!(p100.centroid_score_threshold, Some(0.45));
assert_eq!(p100.n_candidate_docs, Some(1024));
let p1000 = SearchParams::paper_defaults(1000);
assert_eq!(p1000.n_probe, 4);
assert_eq!(p1000.centroid_score_threshold, Some(0.4));
assert_eq!(p1000.n_candidate_docs, Some(4096));
}
#[test]
fn paper_defaults_clamp_ndocs_to_at_least_4x_top_k() {
let p = SearchParams::paper_defaults(2000);
assert!(p.n_candidate_docs.unwrap() >= 4 * 2000);
}
#[test]
fn paper_defaults_always_enable_pruning() {
for &k in &[1_usize, 10, 50, 100, 500, 1000, 10_000] {
assert!(
SearchParams::paper_defaults(k)
.centroid_score_threshold
.is_some(),
"paper_defaults({k}) left pruning disabled"
);
}
}
#[test]
fn batch_maxsim_chunked_matches_single_shot() {
let mut docs = Vec::new();
for i in 0..8u64 {
let jitter = i as f32 * 0.015;
let east = normalize([1.0 - jitter, 0.05 + jitter]);
let north = normalize([0.05 + jitter, 1.0 - jitter]);
let mut tokens = Vec::new();
tokens.extend_from_slice(&east);
tokens.extend_from_slice(&east);
tokens.extend_from_slice(&north);
docs.push(DocumentTokens {
doc_id: 200 + i,
tokens,
n_tokens: 3,
});
}
let index = build_index(&docs, params()).unwrap();
let candidate_idxs: Vec<usize> = (0..docs.len()).collect();
let query = [1.0f32, 0.0, 0.0, 1.0];
let single_shot = batch_maxsim_with_cap(
&query,
&candidate_idxs,
&index,
index.params.dim,
usize::MAX,
)
.unwrap();
for cap in [1usize, 2, 5] {
let chunked = batch_maxsim_with_cap(
&query,
&candidate_idxs,
&index,
index.params.dim,
cap,
)
.unwrap();
assert_eq!(
chunked.len(),
single_shot.len(),
"chunked result count differs at cap={cap}"
);
for (i, (got, want)) in
chunked.iter().zip(single_shot.iter()).enumerate()
{
assert_eq!(
got.0, want.0,
"doc index order differs at cap={cap}, i={i}"
);
assert!(
(got.1 - want.1).abs() < 1e-5,
"score diverges at cap={cap}, i={i}: got {} expected {}",
got.1,
want.1,
);
}
}
}
#[test]
fn decode_chunk_capacity_scales_inversely_with_dim() {
let small = decode_chunk_capacity(128);
let large = decode_chunk_capacity(1536);
assert!(
large < small,
"1536-dim cap {large} should be smaller than 128-dim cap {small}"
);
assert!(
decode_chunk_capacity(1) >= 64,
"floor must hold for pathologically small dim"
);
}
}