docbert_plaid/search.rs
1//! Query-time search over a built [`Index`].
2//!
3//! The flow mirrors PLAID's reference implementation:
4//!
5//! 1. For every query token, find the `n_probe` top dot-product coarse
6//! centroids (matching `S_c,q = C · Q^T`). This is paper Stage 1.
7//! 2. *Optional* centroid pruning: drop probed centroids whose best
8//! score across all query tokens sits below
9//! `centroid_score_threshold`. This is used both here (to shrink the
10//! reachable centroid set before postings are gathered) and below,
11//! at token level, when scoring `D̃`.
12//! 3. Union the doc-level IVF postings for the surviving centroids to
13//! get a set of candidate documents.
14//! 4. *Optional* centroid interaction. When `n_candidate_docs` is set,
15//! the cascade runs in two refinement passes:
16//! - Stage 2 (paper §4.3): approximate MaxSim with the pruned mask
17//! applied at token level; keep top `n_candidate_docs` (paper's
18//! `ndocs`).
19//! - Stage 3 (paper §4.2): approximate MaxSim without pruning on
20//! the Stage 2 survivors; keep top `n_candidate_docs / 4` (paper's
21//! `ndocs/4` empirical heuristic), clamped to at least `top_k`.
22//! 5. Decode the survivors' stored residual codes into approximate
23//! embeddings and compute exact MaxSim against the query.
24//! 6. Sort by exact MaxSim score and return the top-`top_k`.
25//!
26//! The pruning and interaction stages share a precomputed
27//! `[n_query_tokens, n_centroids]` query-centroid score matrix that
28//! step 1 would conceptually compute anyway, and together they
29//! drastically shrink the set of candidates reaching the expensive
30//! decode in step 5. Callers that want the legacy behaviour (every
31//! probed candidate decodes, no pruning) can pass
32//! `n_candidate_docs = None` and `centroid_score_threshold = None`.
33
34use candle_core::{Device, Tensor};
35
36use crate::{
37 Result,
38 codec::DecodeTable,
39 device::default_device,
40 distance::dot,
41 index::Index,
42};
43
44/// Tunable knobs for a single search call.
45#[derive(Debug, Clone, Copy)]
46pub struct SearchParams {
47 /// Number of top-scoring documents to return.
48 pub top_k: usize,
49 /// Number of top dot-product centroids each query token probes.
50 pub n_probe: usize,
51 /// Maximum number of candidates surviving PLAID's
52 /// centroid-interaction stage that proceed to full decode +
53 /// exact MaxSim. `None` disables centroid interaction and sends
54 /// every probed candidate straight to decode — correct but
55 /// slower for large corpora.
56 pub n_candidate_docs: Option<usize>,
57 /// Minimum per-centroid score (max dot product against any query
58 /// token) required for a centroid to contribute to candidate
59 /// generation and centroid interaction. Centroids below this
60 /// threshold are treated as unreachable for this query. `None`
61 /// disables pruning, matching the legacy probe behaviour.
62 pub centroid_score_threshold: Option<f32>,
63}
64
65impl SearchParams {
66 /// Defaults from Table 2 of the PLAID paper (Santhanam et al., 2022),
67 /// bucketed by requested `top_k`.
68 ///
69 /// | top_k bucket | nprobe | t_cs | ndocs |
70 /// |--- |--- |--- |--- |
71 /// | ≤ 10 | 1 | 0.5 | 256 |
72 /// | ≤ 100 | 2 | 0.45 | 1024 |
73 /// | > 100 | 4 | 0.4 | 4096 |
74 ///
75 /// The paper derives these empirically on MS MARCO and LoTTE; they
76 /// trade a small recall hit for large latency wins versus an
77 /// exhaustive probe. `ndocs` is always at least `4 * top_k` so
78 /// Stage 3's `ndocs/4` shortlist can still return enough candidates
79 /// for the caller's requested result count.
80 pub const fn paper_defaults(top_k: usize) -> Self {
81 let (n_probe, t_cs, ndocs) = if top_k <= 10 {
82 (1, 0.5_f32, 256)
83 } else if top_k <= 100 {
84 (2, 0.45_f32, 1024)
85 } else {
86 (4, 0.4_f32, 4096)
87 };
88 let ndocs = if ndocs < top_k.saturating_mul(4) {
89 top_k.saturating_mul(4)
90 } else {
91 ndocs
92 };
93 Self {
94 top_k,
95 n_probe,
96 n_candidate_docs: Some(ndocs),
97 centroid_score_threshold: Some(t_cs),
98 }
99 }
100}
101
102/// One entry in a ranked search result list.
103#[derive(Debug, Clone, Copy, PartialEq)]
104pub struct SearchResult {
105 pub doc_id: u64,
106 pub score: f32,
107}
108
109/// Run a PLAID-style search over `index` and return the top-`top_k`
110/// documents ranked by ColBERT MaxSim against `query_tokens`.
111///
112/// `query_tokens` is a flat row-major `n_query_tokens × dim` buffer.
113/// Results are sorted by score, highest first. Ties are broken by
114/// `doc_id` ascending to keep the output deterministic.
115///
116/// # Errors
117///
118/// Returns [`PlaidError::Tensor`] if the batched MaxSim matmul fails
119/// (e.g. CUDA OOM on the final decode tensor).
120///
121/// # Panics
122///
123/// Panics if `query_tokens.len() % index.params.dim != 0`, if
124/// `params.top_k == 0`, or if `params.n_probe == 0`.
125///
126/// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
127pub fn search(
128 index: &Index,
129 query_tokens: &[f32],
130 params: SearchParams,
131) -> Result<Vec<SearchResult>> {
132 let dim = index.params.dim;
133 assert!(params.top_k > 0, "search: top_k must be positive");
134 assert!(params.n_probe > 0, "search: n_probe must be positive");
135 assert!(
136 query_tokens.len().is_multiple_of(dim),
137 "search: query length {} is not a multiple of dim {}",
138 query_tokens.len(),
139 dim,
140 );
141
142 if query_tokens.is_empty() || index.num_documents() == 0 {
143 return Ok(Vec::new());
144 }
145
146 let n_centroids = index.codec.num_centroids();
147 let n_probe = params.n_probe.min(n_centroids);
148
149 // Precompute the per-centroid max dot product against any query
150 // token when either pruning or centroid interaction is active.
151 // Pruning uses it to build a pruned-centroid bitmask for use in
152 // both candidate generation and centroid interaction; centroid
153 // interaction uses the full qc_scores matrix for approximate
154 // per-doc scoring.
155 let qc_scores: Option<Vec<f32>> = (params
156 .centroid_score_threshold
157 .is_some()
158 || params.n_candidate_docs.is_some())
159 .then(|| {
160 query_centroid_score_matrix(query_tokens, &index.codec.centroids, dim)
161 });
162 let pruned_mask: Option<Vec<bool>> =
163 match (qc_scores.as_ref(), params.centroid_score_threshold) {
164 (Some(scores), Some(threshold)) => {
165 let per_cent = per_centroid_max_scores(
166 scores,
167 query_tokens.len() / dim,
168 n_centroids,
169 );
170 Some(per_cent.iter().map(|&s| s < threshold).collect())
171 }
172 _ => None,
173 };
174
175 // 1-2. Gather the union of candidate doc indices reachable via the
176 // probed centroids. IVF postings are already deduplicated per
177 // doc, so each centroid contributes at most one write per doc.
178 // Centroid pruning drops probed centroids whose best per-query
179 // score sits below the caller's threshold — skipping them
180 // both shrinks the candidate set and saves work later in
181 // centroid interaction.
182 let mut candidate_docs: Vec<bool> = vec![false; index.num_documents()];
183 for query_token in query_tokens.chunks_exact(dim) {
184 for centroid_id in
185 top_n_centroids(query_token, &index.codec.centroids, dim, n_probe)
186 {
187 if let Some(mask) = pruned_mask.as_ref()
188 && mask[centroid_id]
189 {
190 continue;
191 }
192 for &doc_idx in index.ivf.docs_for_centroid(centroid_id) {
193 candidate_docs[doc_idx as usize] = true;
194 }
195 }
196 }
197
198 let mut candidate_idxs: Vec<usize> = candidate_docs
199 .iter()
200 .enumerate()
201 .filter_map(|(idx, &is_cand)| {
202 (is_cand && index.doc_token_count(idx) > 0).then_some(idx)
203 })
204 .collect();
205
206 // 3. Centroid interaction: when the caller set `n_candidate_docs`,
207 // cheaply rank candidates via the precomputed query-centroid
208 // score matrix and keep only the top survivors before paying
209 // for full decode + exact MaxSim.
210 //
211 // The `if let` pattern-matches both fields at once so the
212 // compiler can prove qc_scores is Some under this arm — no
213 // `expect` gymnastics needed.
214 if let (Some(n_stage2), Some(qc_scores)) =
215 (params.n_candidate_docs, qc_scores.as_ref())
216 {
217 let n_q = query_tokens.len() / dim;
218 let n_c = index.codec.num_centroids();
219
220 // Stage 2 — cheap pruned centroid interaction. Keep top
221 // `n_stage2` candidates (the paper's `ndocs`).
222 candidate_idxs = shortlist_by_approx_score(
223 candidate_idxs,
224 index,
225 qc_scores,
226 n_q,
227 n_c,
228 pruned_mask.as_deref(),
229 n_stage2,
230 );
231
232 // Stage 3 — unpruned centroid interaction. Refines the Stage 2
233 // survivors down to `ndocs/4`, matching the paper's empirical
234 // heuristic. Callers that want top-`k` results should set
235 // `n_candidate_docs >= 4 * top_k` so Stage 3 doesn't undershoot
236 // their requested result count. docbert-core's default already
237 // does this (8 * top_k).
238 let n_stage3 = n_stage2.div_ceil(4).max(1);
239 candidate_idxs = shortlist_by_approx_score(
240 candidate_idxs,
241 index,
242 qc_scores,
243 n_q,
244 n_c,
245 None,
246 n_stage3,
247 );
248 }
249
250 // 5. Score every surviving candidate with one batched MaxSim matmul
251 // on decoded tokens.
252 let mut scored: Vec<SearchResult> = if candidate_idxs.is_empty() {
253 Vec::new()
254 } else {
255 batch_maxsim(query_tokens, &candidate_idxs, index, dim)?
256 .into_iter()
257 .map(|(doc_idx, score)| SearchResult {
258 doc_id: index.doc_ids[doc_idx],
259 score,
260 })
261 .collect()
262 };
263
264 // 6. Rank by score (desc), tie-break by doc_id (asc) for determinism.
265 scored.sort_by(|a, b| {
266 b.score
267 .partial_cmp(&a.score)
268 .unwrap_or(std::cmp::Ordering::Equal)
269 .then_with(|| a.doc_id.cmp(&b.doc_id))
270 });
271 scored.truncate(params.top_k);
272 Ok(scored)
273}
274
275/// Indices of the `n` centroids with the highest dot-product against
276/// `point`, most relevant first.
277///
278/// This is the multi-centroid counterpart to
279/// [`crate::kmeans::nearest_centroid`]. It implements PLAID's
280/// query-centroid scoring `S_c,q = C · Q^T` (one row per query token),
281/// which is the ranking the paper uses for candidate generation.
282///
283/// Dot product is used here instead of squared L2 because the paper and
284/// ColBERT's downstream MaxSim both operate on dot-product similarity.
285/// For centroids with varying magnitudes the two metrics disagree, and
286/// dot product is what keeps probe ordering consistent with the scorer.
287/// Ties are broken toward the earlier centroid index for determinism.
288///
289/// The output vector has length `min(n, num_centroids)`.
290///
291/// # Panics
292///
293/// Panics on any shape mismatch between `point`, `centroids`, and `dim`.
294pub fn top_n_centroids(
295 point: &[f32],
296 centroids: &[f32],
297 dim: usize,
298 n: usize,
299) -> Vec<usize> {
300 assert_eq!(
301 point.len(),
302 dim,
303 "top_n_centroids: point length {} does not match dim {}",
304 point.len(),
305 dim,
306 );
307 assert!(dim > 0, "top_n_centroids: dim must be positive");
308 assert!(
309 !centroids.is_empty() && centroids.len().is_multiple_of(dim),
310 "top_n_centroids: centroids length {} is not a positive multiple of dim {}",
311 centroids.len(),
312 dim,
313 );
314
315 let k = centroids.len() / dim;
316 let mut scored: Vec<(usize, f32)> = centroids
317 .chunks_exact(dim)
318 .enumerate()
319 .map(|(i, c)| (i, dot(point, c)))
320 .collect();
321 scored.sort_by(|a, b| {
322 b.1.partial_cmp(&a.1)
323 .unwrap_or(std::cmp::Ordering::Equal)
324 .then_with(|| a.0.cmp(&b.0))
325 });
326 scored.into_iter().take(n.min(k)).map(|(i, _)| i).collect()
327}
328
329/// Compute the row-major `[n_q, n_centroids]` matrix of query-to-centroid
330/// dot products used by PLAID's centroid-interaction stage.
331///
332/// Materialising the whole matrix once amortises the per-centroid work
333/// across every candidate document — the same query-centroid entries
334/// get hit by every doc that touches each centroid.
335fn query_centroid_score_matrix(
336 query_tokens: &[f32],
337 centroids: &[f32],
338 dim: usize,
339) -> Vec<f32> {
340 let n_q = query_tokens.len() / dim;
341 let n_c = centroids.len() / dim;
342 let mut out = vec![0.0f32; n_q * n_c];
343 for (qi, q) in query_tokens.chunks_exact(dim).enumerate() {
344 for (ci, c) in centroids.chunks_exact(dim).enumerate() {
345 out[qi * n_c + ci] = dot(q, c);
346 }
347 }
348 out
349}
350
351/// Rank `candidate_idxs` by approximate centroid-interaction score
352/// and keep the top `limit` survivors.
353///
354/// A no-op when the input is already smaller than `limit`. Used by
355/// both Stage 2 (pruned) and Stage 3 (unpruned) of the paper's
356/// centroid-interaction cascade; the only difference between the two
357/// callers is the `mask` argument.
358fn shortlist_by_approx_score(
359 candidate_idxs: Vec<usize>,
360 index: &Index,
361 qc_scores: &[f32],
362 n_q: usize,
363 n_centroids: usize,
364 mask: Option<&[bool]>,
365 limit: usize,
366) -> Vec<usize> {
367 if candidate_idxs.len() <= limit {
368 return candidate_idxs;
369 }
370 let mut approx: Vec<(usize, f32)> = candidate_idxs
371 .into_iter()
372 .map(|doc_idx| {
373 let score = approx_centroid_interaction_score(
374 index.doc_centroid_ids(doc_idx),
375 qc_scores,
376 n_q,
377 n_centroids,
378 mask,
379 );
380 (doc_idx, score)
381 })
382 .collect();
383 approx.sort_by(|a, b| {
384 b.1.partial_cmp(&a.1)
385 .unwrap_or(std::cmp::Ordering::Equal)
386 .then_with(|| a.0.cmp(&b.0))
387 });
388 approx.truncate(limit);
389 approx.into_iter().map(|(idx, _)| idx).collect()
390}
391
392/// For every centroid, return the max dot-product score achieved by
393/// any query token against that centroid.
394///
395/// PLAID's centroid pruning uses this "best any-query score" to cheaply
396/// rank centroids by query relevance: a centroid whose best score
397/// across all query tokens is still small cannot usefully participate
398/// in MaxSim, so it can be dropped without hurting top-K.
399fn per_centroid_max_scores(
400 qc_scores: &[f32],
401 n_q: usize,
402 n_centroids: usize,
403) -> Vec<f32> {
404 let mut out = vec![f32::NEG_INFINITY; n_centroids];
405 for q in 0..n_q {
406 let row = &qc_scores[q * n_centroids..(q + 1) * n_centroids];
407 for (c, &s) in row.iter().enumerate() {
408 if s > out[c] {
409 out[c] = s;
410 }
411 }
412 }
413 out
414}
415
416/// Approximate MaxSim for one candidate doc using only the centroids
417/// the doc touches. This is PLAID's centroid-interaction score:
418///
419/// `Σ_q max_{c ∈ unique_centroids(doc)} qc_scores[q, c]`
420///
421/// Duplicate centroid ids within the doc are filtered implicitly: a
422/// centroid that appears twice can't beat itself in the max, so the
423/// running-max shortcut is correct without an explicit dedup pass.
424///
425/// When `pruned_mask` is `Some`, tokens whose centroid is marked
426/// pruned are skipped before the max, matching PLAID §4.3's token-level
427/// pruning of `D̃`. If every token in the doc is pruned the result is
428/// `f32::NEG_INFINITY`, which sorts below every non-pruned doc in the
429/// centroid-interaction shortlist.
430fn approx_centroid_interaction_score(
431 doc_centroid_ids: &[u32],
432 qc_scores: &[f32],
433 n_q: usize,
434 n_centroids: usize,
435 pruned_mask: Option<&[bool]>,
436) -> f32 {
437 if doc_centroid_ids.is_empty() {
438 return 0.0;
439 }
440 // If every token's centroid is pruned, the doc has no D̃ rows left;
441 // report −∞ so the caller's ranking drops it instead of a
442 // deceptive 0.0.
443 if let Some(mask) = pruned_mask
444 && doc_centroid_ids.iter().all(|cid| mask[*cid as usize])
445 {
446 return f32::NEG_INFINITY;
447 }
448 let mut total = 0.0f32;
449 for q in 0..n_q {
450 let row = &qc_scores[q * n_centroids..(q + 1) * n_centroids];
451 let mut best = f32::NEG_INFINITY;
452 for &cid in doc_centroid_ids {
453 if let Some(mask) = pruned_mask
454 && mask[cid as usize]
455 {
456 continue;
457 }
458 let s = row[cid as usize];
459 if s > best {
460 best = s;
461 }
462 }
463 if best.is_finite() {
464 total += best;
465 }
466 }
467 total
468}
469
470/// Score a batch of candidate docs against `query_tokens` with a
471/// padding-free packed MaxSim, matching PLAID §4.5.
472///
473/// Two decode strategies share the downstream MaxSim reduction:
474///
475/// - On CPU the fastest path is the hand-rolled LUT walk — a byte
476/// lookup per packed slot beats candle's generic `index_select`
477/// because there's no kernel-launch overhead to amortise.
478/// - On CUDA the right move is to keep the whole flow on-device:
479/// upload the centroid bank and weights LUT once, gather token
480/// residuals via two `index_select` ops, add, and feed straight
481/// into the MaxSim GEMM. This mirrors PLAID's §4.5 CUDA kernel
482/// (one thread per packed byte, centroid-add, matmul) without
483/// re-implementing it by hand.
484///
485/// The per-doc max-then-sum reduction still runs on the host because
486/// candle doesn't expose a native segmented-max primitive and we want
487/// to avoid materialising the padded `[n_docs, max_len, n_q]` tensor
488/// the paper explicitly rejects.
489fn batch_maxsim(
490 query_tokens: &[f32],
491 candidate_idxs: &[usize],
492 index: &Index,
493 dim: usize,
494) -> Result<Vec<(usize, f32)>> {
495 batch_maxsim_with_cap(
496 query_tokens,
497 candidate_idxs,
498 index,
499 dim,
500 decode_chunk_capacity(dim),
501 )
502}
503
504/// Target resident-tensor budget per chunk during on-device decode.
505///
506/// The GPU decode + MaxSim pipeline holds roughly five
507/// `[chunk_tokens, dim]` f32 tensors live at once (centroid
508/// embeddings, residuals-flat, residuals-padded alias, decoded rows
509/// post-add, and the L2-normalise scratch) plus the
510/// `[chunk_tokens, n_q]` scores matmul output. Budgeting around
511/// **64 MiB** of combined working set keeps the transient footprint
512/// small enough that a 12 GiB card still has room left over once
513/// the ColBERT model is resident and a query encode has happened.
514///
515/// 64 MiB is deliberately loose: candle's matmul reserves its own
516/// scratch pool, and we want the chunk loop to compose cleanly with
517/// that reservation without a second-order sizing calculation. For a
518/// 128-dim ColBERT this lands at ~26K tokens per chunk; for a
519/// 1536-dim LateOn it drops to ~2K.
520const DECODE_CHUNK_BUDGET_BYTES: usize = 64 * 1024 * 1024;
521
522/// Rough upper bound on how many tokens we can safely decode in one
523/// on-device chunk at `dim` embedding width.
524///
525/// Returns at least 64 so a pathologically wide model still makes
526/// progress — a single doc's tokens fall into one chunk even if the
527/// strict byte budget would ask for fewer. The floor is also what
528/// keeps tests with `dim=2` corpora from picking an absurdly large
529/// cap that defeats the chunking coverage.
530fn decode_chunk_capacity(dim: usize) -> usize {
531 let per_token = dim.saturating_mul(4 * 5).max(1);
532 (DECODE_CHUNK_BUDGET_BYTES / per_token).max(64)
533}
534
535/// Chunked implementation of [`batch_maxsim`] that caps how many
536/// candidate tokens may be resident on the device at once.
537///
538/// Splits `candidate_idxs` into contiguous runs whose cumulative
539/// token count stays under `max_tokens_per_chunk`, decodes each run
540/// in isolation (one upload, one matmul, one `to_vec2`), and does
541/// the per-doc MaxSim reduction on the host before releasing the
542/// chunk's tensors. Peak VRAM during search is therefore bounded by
543/// the chunk budget rather than by the total candidate set — which
544/// previously blew past available memory for 1536-dim models on
545/// modest cards when 200+ documents survived PLAID's Stage 3
546/// shortlist.
547///
548/// When `max_tokens_per_chunk` is larger than the full candidate
549/// token count, the loop runs exactly once and the behaviour matches
550/// the pre-chunked implementation byte-for-byte.
551///
552/// [`batch_maxsim`]: crate::search::batch_maxsim
553fn batch_maxsim_with_cap(
554 query_tokens: &[f32],
555 candidate_idxs: &[usize],
556 index: &Index,
557 dim: usize,
558 max_tokens_per_chunk: usize,
559) -> Result<Vec<(usize, f32)>> {
560 let n_q = query_tokens.len() / dim;
561
562 let decode_table = DecodeTable::new(&index.codec);
563 let total_tokens: usize = candidate_idxs
564 .iter()
565 .map(|&i| index.doc_token_count(i))
566 .sum();
567
568 if total_tokens == 0 || n_q == 0 {
569 return Ok(candidate_idxs.iter().map(|&i| (i, 0.0)).collect());
570 }
571
572 let device = default_device();
573 let q_t = Tensor::from_slice(query_tokens, (n_q, dim), device)?;
574 let q_transposed = q_t.t()?.contiguous()?;
575
576 let mut out = Vec::with_capacity(candidate_idxs.len());
577 let mut start = 0usize;
578 while start < candidate_idxs.len() {
579 // Grow the chunk until adding the next doc would exceed the
580 // cap, but always include at least one doc — a single doc
581 // larger than the cap still has to decode on its own, and
582 // that's a better failure mode than looping forever.
583 let mut end = start + 1;
584 let mut chunk_tokens = index.doc_token_count(candidate_idxs[start]);
585 while end < candidate_idxs.len() {
586 let next = index.doc_token_count(candidate_idxs[end]);
587 if chunk_tokens + next > max_tokens_per_chunk {
588 break;
589 }
590 chunk_tokens += next;
591 end += 1;
592 }
593
594 let chunk_candidates = &candidate_idxs[start..end];
595 let mut chunk_offsets: Vec<usize> =
596 Vec::with_capacity(chunk_candidates.len() + 1);
597 chunk_offsets.push(0);
598
599 let chunk_ctx = DecodeCtx {
600 index,
601 candidate_idxs: chunk_candidates,
602 dim,
603 total_tokens: chunk_tokens,
604 decode_table: &decode_table,
605 device,
606 };
607 let decoded = if matches!(device, Device::Cpu) {
608 decode_on_cpu(&chunk_ctx, &mut chunk_offsets)?
609 } else {
610 decode_on_device(&chunk_ctx, &mut chunk_offsets)?
611 };
612
613 // `[chunk_tokens, dim] × [dim, n_q]` GEMM, same shape rule as
614 // the pre-chunk implementation but bounded by the chunk cap.
615 // Pull the result back as a single flat host buffer — one
616 // allocation rather than the per-row Vec<Vec<f32>> that
617 // `to_vec2` builds.
618 let scores_flat: Vec<f32> = decoded
619 .matmul(&q_transposed)?
620 .flatten_all()?
621 .to_vec1::<f32>()?;
622
623 for (i, &doc_idx) in chunk_candidates.iter().enumerate() {
624 let range_start = chunk_offsets[i];
625 let range_end = chunk_offsets[i + 1];
626 if range_start == range_end {
627 out.push((doc_idx, 0.0));
628 continue;
629 }
630 let mut total = 0.0f32;
631 for q in 0..n_q {
632 let mut best = f32::NEG_INFINITY;
633 for t in range_start..range_end {
634 let s = scores_flat[t * n_q + q];
635 if s > best {
636 best = s;
637 }
638 }
639 if best.is_finite() {
640 total += best;
641 }
642 }
643 out.push((doc_idx, total));
644 }
645
646 start = end;
647 }
648
649 Ok(out)
650}
651
652/// Inputs the two decode strategies share. Grouping them keeps the
653/// call-site readable and satisfies clippy's `too_many_arguments`
654/// without leaking implementation details into the public API.
655struct DecodeCtx<'a> {
656 index: &'a Index,
657 candidate_idxs: &'a [usize],
658 dim: usize,
659 total_tokens: usize,
660 decode_table: &'a DecodeTable,
661 device: &'a Device,
662}
663
664/// CPU decode path: walk every candidate token, look up each packed
665/// byte in the precomputed LUT, add the centroid dimensions, push
666/// into one big row-major buffer, then upload as a single tensor.
667///
668/// Fast on CPU because the whole loop is contiguous writes plus a
669/// 256-entry table of 4 KiB at `nbits=2` that lives in L1 the entire
670/// call — measurably faster than routing every token through candle
671/// `index_select`, which pays kernel-launch overhead per gather.
672fn decode_on_cpu(
673 ctx: &DecodeCtx<'_>,
674 offsets: &mut Vec<usize>,
675) -> Result<Tensor> {
676 let codec = &ctx.index.codec;
677 let packed_bytes = codec.packed_bytes();
678 let mut packed: Vec<f32> = Vec::with_capacity(ctx.total_tokens * ctx.dim);
679 let mut ev = crate::codec::EncodedVector {
680 centroid_id: 0,
681 codes: vec![0u8; packed_bytes],
682 };
683 for &doc_idx in ctx.candidate_idxs {
684 let cids = ctx.index.doc_centroid_ids(doc_idx);
685 let bytes = ctx.index.doc_residual_bytes(doc_idx);
686 for (i, &cid) in cids.iter().enumerate() {
687 ev.centroid_id = cid;
688 ev.codes.copy_from_slice(
689 &bytes[i * packed_bytes..(i + 1) * packed_bytes],
690 );
691 let decoded =
692 codec.decode_vector_with_table(&ev, ctx.decode_table)?;
693 // Row-wise L2 normalise, mirroring decode_on_device's tail
694 // so CPU and GPU paths return comparable scores. See
695 // `l2_normalize_rows` for why this matters.
696 let norm = decoded
697 .iter()
698 .map(|v| v * v)
699 .sum::<f32>()
700 .sqrt()
701 .max(1e-12_f32);
702 packed.extend(decoded.iter().map(|v| v / norm));
703 }
704 offsets.push(packed.len() / ctx.dim);
705 }
706 Ok(Tensor::from_vec(
707 packed,
708 (ctx.total_tokens, ctx.dim),
709 ctx.device,
710 )?)
711}
712
713/// GPU decode path: upload the centroid bank and the 256-entry LUT,
714/// then recover every candidate token via two `index_select` gathers
715/// and an add — all on-device. This keeps the decode inside the same
716/// kernel stream as the MaxSim matmul, which is the win PLAID §4.5
717/// describes for its CUDA implementation.
718///
719/// When `dim` isn't a multiple of `codes_per_byte`, the packed
720/// residual buffer has `packed_bytes * codes_per_byte - dim` trailing
721/// slack slots. We decode the padded row width and `narrow` back to
722/// `dim` so the output tensor has exactly the expected shape.
723fn decode_on_device(
724 ctx: &DecodeCtx<'_>,
725 offsets: &mut Vec<usize>,
726) -> Result<Tensor> {
727 let packed_bytes_per_vec = ctx.index.codec.packed_bytes();
728 let codes_per_byte = ctx.decode_table.codes_per_byte();
729 let decoded_row_width = packed_bytes_per_vec * codes_per_byte;
730 debug_assert!(decoded_row_width >= ctx.dim);
731
732 // Flat gather buffers — slice directly out of the index's
733 // StridedTensor-style storage. For each candidate we contribute
734 // its `doc_centroid_ids` slice (`[n_tokens]` u32) and its
735 // `doc_residual_bytes` slice (`[n_tokens, packed_bytes_per_vec]`
736 // u8), extended into the flat upload buffers.
737 let mut centroid_ids: Vec<u32> = Vec::with_capacity(ctx.total_tokens);
738 let mut packed_bytes: Vec<u32> =
739 Vec::with_capacity(ctx.total_tokens * packed_bytes_per_vec);
740 for &doc_idx in ctx.candidate_idxs {
741 let doc_cids = ctx.index.doc_centroid_ids(doc_idx);
742 let doc_bytes = ctx.index.doc_residual_bytes(doc_idx);
743 centroid_ids.extend_from_slice(doc_cids);
744 packed_bytes.extend(doc_bytes.iter().map(|&b| b as u32));
745 offsets.push(centroid_ids.len());
746 }
747
748 let n_centroids = ctx.index.codec.num_centroids();
749 let centroids_t = Tensor::from_slice(
750 ctx.index.codec.centroids.as_slice(),
751 (n_centroids, ctx.dim),
752 ctx.device,
753 )?;
754 let weights_t = Tensor::from_slice(
755 ctx.decode_table.weights_flat(),
756 (256, codes_per_byte),
757 ctx.device,
758 )?;
759 let cent_idx_t =
760 Tensor::from_vec(centroid_ids, (ctx.total_tokens,), ctx.device)?;
761 let byte_idx_t = Tensor::from_vec(
762 packed_bytes,
763 (ctx.total_tokens * packed_bytes_per_vec,),
764 ctx.device,
765 )?;
766
767 let centroid_emb = centroids_t.index_select(¢_idx_t, 0)?;
768 let residuals_flat = weights_t.index_select(&byte_idx_t, 0)?;
769 let residuals_padded =
770 residuals_flat.reshape((ctx.total_tokens, decoded_row_width))?;
771 let residuals = if decoded_row_width == ctx.dim {
772 residuals_padded
773 } else {
774 residuals_padded.narrow(1, 0, ctx.dim)?
775 };
776 let decoded = (centroid_emb + residuals)?;
777 l2_normalize_rows(&decoded)
778}
779
780/// Row-wise L2-normalise an `[n, dim]` tensor. Matches fast-plaid's
781/// `decompress_residuals` tail: ColBERT was trained with unit-norm
782/// token embeddings, so MaxSim (a dot product) assumes unit rows
783/// downstream. Quantise → centroid+residual reconstruction drifts
784/// norms a little away from 1; this pulls them back so the matmul
785/// that follows stays a proper cosine similarity.
786fn l2_normalize_rows(decoded: &Tensor) -> Result<Tensor> {
787 let squared = decoded.sqr()?;
788 let norm_sq = squared.sum_keepdim(1)?; // [n, 1]
789 // Clamp away from zero before sqrt — a stray all-zero row (possible
790 // with pathological residuals in tests) would otherwise produce a
791 // division by zero below.
792 let norm = norm_sq.sqrt()?.clamp(1e-12f32, f32::INFINITY)?;
793 Ok(decoded.broadcast_div(&norm)?)
794}
795
796#[cfg(test)]
797mod tests {
798 use super::*;
799 use crate::{
800 distance::dot,
801 index::{DocumentTokens, IndexParams, build_index},
802 };
803
804 fn params() -> IndexParams {
805 IndexParams {
806 dim: 2,
807 nbits: 2,
808 k_centroids: 2,
809 max_kmeans_iters: 50,
810 }
811 }
812
813 /// Two well-separated clusters of **unit-norm** 2-D tokens in three
814 /// documents. This mirrors the real-world ColBERT invariant that
815 /// token embeddings live on the unit sphere, which is what makes
816 /// the dot-product MaxSim meaningful. Doc 1 points east, doc 2
817 /// points north, doc 3 has one token in each cluster.
818 fn corpus() -> Vec<DocumentTokens> {
819 // Unit vectors by construction.
820 let east_a = [1.0f32, 0.0];
821 let east_b = normalize([0.98, 0.2]);
822 let east_c = normalize([0.97, -0.24]);
823 let north_a = [0.0f32, 1.0];
824 let north_b = normalize([0.2, 0.98]);
825 let north_c = normalize([-0.24, 0.97]);
826
827 let mut doc1 = Vec::new();
828 doc1.extend_from_slice(&east_a);
829 doc1.extend_from_slice(&east_b);
830 doc1.extend_from_slice(&east_c);
831
832 let mut doc2 = Vec::new();
833 doc2.extend_from_slice(&north_a);
834 doc2.extend_from_slice(&north_b);
835 doc2.extend_from_slice(&north_c);
836
837 let mut doc3 = Vec::new();
838 doc3.extend_from_slice(&east_a);
839 doc3.extend_from_slice(&north_a);
840
841 vec![
842 DocumentTokens {
843 doc_id: 1,
844 tokens: doc1,
845 n_tokens: 3,
846 },
847 DocumentTokens {
848 doc_id: 2,
849 tokens: doc2,
850 n_tokens: 3,
851 },
852 DocumentTokens {
853 doc_id: 3,
854 tokens: doc3,
855 n_tokens: 2,
856 },
857 ]
858 }
859
860 fn normalize(v: [f32; 2]) -> [f32; 2] {
861 let norm = (v[0] * v[0] + v[1] * v[1]).sqrt();
862 [v[0] / norm, v[1] / norm]
863 }
864
865 #[test]
866 fn top_n_centroids_ranks_by_descending_dot_product() {
867 // Three centroids on the +x axis at magnitudes 0, 5, and 12.
868 // Query at (4, 0). Dot products are 0, 20, 48; PLAID ranks by
869 // descending relevance (`S_c,q = C · Q^T`), so the biggest
870 // dot product wins.
871 let centroids = [0.0, 0.0, 5.0, 0.0, 12.0, 0.0];
872 let out = top_n_centroids(&[4.0, 0.0], ¢roids, 2, 3);
873 assert_eq!(out, vec![2, 1, 0]);
874 }
875
876 #[test]
877 fn top_n_centroids_breaks_ties_toward_earlier_index() {
878 // Zero query makes every dot product 0. The deterministic
879 // tie-break is "earliest index wins".
880 let centroids = [1.0, 0.0, 0.0, 1.0, -1.0, 0.0];
881 let out = top_n_centroids(&[0.0, 0.0], ¢roids, 2, 3);
882 assert_eq!(out, vec![0, 1, 2]);
883 }
884
885 #[test]
886 fn top_n_centroids_on_unit_norm_centroids_picks_most_aligned() {
887 // All centroids and the query are unit-norm. Dot product then
888 // reduces to cosine similarity, so the best-aligned centroid
889 // wins. Here (0.6, 0.8) is the closest direction to (1, 0).
890 let centroids = [
891 0.6, 0.8, //
892 0.0, 1.0, //
893 -1.0, 0.0, //
894 1.0, 0.0,
895 ];
896 let out = top_n_centroids(&[1.0, 0.0], ¢roids, 2, 2);
897 assert_eq!(out, vec![3, 0]);
898 }
899
900 #[test]
901 fn top_n_centroids_caps_at_num_centroids() {
902 let centroids = [0.0, 0.0, 5.0, 0.0];
903 let out = top_n_centroids(&[0.0, 0.0], ¢roids, 2, 10);
904 assert_eq!(out, vec![0, 1]);
905 }
906
907 #[test]
908 fn search_returns_empty_for_empty_query() {
909 let index = build_index(&corpus(), params()).unwrap();
910 let out = search(
911 &index,
912 &[],
913 SearchParams {
914 top_k: 3,
915 n_probe: 2,
916 n_candidate_docs: None,
917 centroid_score_threshold: None,
918 },
919 )
920 .unwrap();
921 assert!(out.is_empty());
922 }
923
924 #[test]
925 fn search_ranks_matching_cluster_highest() {
926 // Unit-norm query pointing east. Doc 1 has three east tokens,
927 // doc 3 has one east and one north token, doc 2 is entirely
928 // north. MaxSim should prefer doc 1.
929 let index = build_index(&corpus(), params()).unwrap();
930 let out = search(
931 &index,
932 &[1.0, 0.0],
933 SearchParams {
934 top_k: 3,
935 n_probe: 2,
936 n_candidate_docs: None,
937 centroid_score_threshold: None,
938 },
939 )
940 .unwrap();
941 assert!(!out.is_empty(), "should surface at least one doc");
942 assert_eq!(out[0].doc_id, 1, "closest doc should rank first");
943 }
944
945 #[test]
946 fn search_respects_top_k() {
947 let index = build_index(&corpus(), params()).unwrap();
948 let out = search(
949 &index,
950 &[1.0, 0.0],
951 SearchParams {
952 top_k: 1,
953 n_probe: 2,
954 n_candidate_docs: None,
955 centroid_score_threshold: None,
956 },
957 )
958 .unwrap();
959 assert_eq!(out.len(), 1);
960 }
961
962 #[test]
963 fn search_scores_are_non_increasing() {
964 let index = build_index(&corpus(), params()).unwrap();
965 // Two query tokens, one east, one north.
966 let out = search(
967 &index,
968 &[1.0, 0.0, 0.0, 1.0],
969 SearchParams {
970 top_k: 3,
971 n_probe: 2,
972 n_candidate_docs: None,
973 centroid_score_threshold: None,
974 },
975 )
976 .unwrap();
977 for pair in out.windows(2) {
978 assert!(
979 pair[0].score >= pair[1].score,
980 "scores must be descending: {pair:?}"
981 );
982 }
983 }
984
985 #[test]
986 fn search_with_single_probe_still_finds_the_right_cluster() {
987 // With n_probe=1, only one centroid is probed per query token.
988 // A query firmly inside one cluster should still surface its
989 // corresponding document first.
990 let index = build_index(&corpus(), params()).unwrap();
991 let out = search(
992 &index,
993 &[0.0, 1.0],
994 SearchParams {
995 top_k: 1,
996 n_probe: 1,
997 n_candidate_docs: None,
998 centroid_score_threshold: None,
999 },
1000 )
1001 .unwrap();
1002 assert_eq!(out[0].doc_id, 2);
1003 }
1004
1005 #[test]
1006 fn search_skips_documents_with_no_tokens() {
1007 let mut docs = corpus();
1008 docs.push(DocumentTokens {
1009 doc_id: 42,
1010 tokens: vec![],
1011 n_tokens: 0,
1012 });
1013 let index = build_index(&docs, params()).unwrap();
1014 let out = search(
1015 &index,
1016 &[1.0, 0.0],
1017 SearchParams {
1018 top_k: 10,
1019 n_probe: 2,
1020 n_candidate_docs: None,
1021 centroid_score_threshold: None,
1022 },
1023 )
1024 .unwrap();
1025 assert!(
1026 out.iter().all(|r| r.doc_id != 42),
1027 "empty doc should not appear in results"
1028 );
1029 }
1030
1031 #[test]
1032 fn search_score_equals_sum_of_maxsim_over_decoded_tokens() {
1033 // Sanity: the score we return should match a ground-truth
1034 // MaxSim computed with the same decoded tokens, confirming our
1035 // per-token similarity is dot product (not -squared_l2).
1036 let index = build_index(&corpus(), params()).unwrap();
1037 let query = [1.0f32, 0.0];
1038
1039 let out = search(
1040 &index,
1041 &query,
1042 SearchParams {
1043 top_k: 1,
1044 n_probe: index.codec.num_centroids(),
1045 n_candidate_docs: None,
1046 centroid_score_threshold: None,
1047 },
1048 )
1049 .unwrap();
1050 assert!(!out.is_empty());
1051
1052 let top = out[0];
1053 let doc_idx = index.position_of(top.doc_id).expect("doc present");
1054 // search() normalises each decoded token to unit norm before
1055 // MaxSim, matching fast-plaid. The reference we compute here
1056 // must mirror that to be byte-comparable.
1057 let decoded: Vec<Vec<f32>> = index
1058 .doc_tokens_vec(doc_idx)
1059 .iter()
1060 .map(|ev| {
1061 let raw = index.codec.decode_vector(ev).unwrap();
1062 let norm =
1063 raw.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
1064 raw.iter().map(|v| v / norm).collect()
1065 })
1066 .collect();
1067
1068 // Standard MaxSim: Σ max_j q_i · d_j.
1069 let mut expected = 0.0f32;
1070 for q in query.as_chunks::<2>().0 {
1071 let best = decoded
1072 .iter()
1073 .map(|d| dot(q, d))
1074 .fold(f32::NEG_INFINITY, f32::max);
1075 expected += best;
1076 }
1077 assert!(
1078 (expected - top.score).abs() < 1e-5,
1079 "score {} differs from ground-truth MaxSim {}",
1080 top.score,
1081 expected,
1082 );
1083 }
1084
1085 #[test]
1086 fn search_works_with_four_bit_residuals() {
1087 // Build an index with 4-bit residual quantization (16 buckets)
1088 // and confirm the MaxSim-ranked top result is still the
1089 // matching-cluster doc.
1090 let params_4bit = IndexParams {
1091 dim: 2,
1092 nbits: 4,
1093 k_centroids: 2,
1094 max_kmeans_iters: 50,
1095 };
1096 let index = build_index(&corpus(), params_4bit).unwrap();
1097 let out = search(
1098 &index,
1099 &[1.0, 0.0],
1100 SearchParams {
1101 top_k: 1,
1102 n_probe: 2,
1103 n_candidate_docs: None,
1104 centroid_score_threshold: None,
1105 },
1106 )
1107 .unwrap();
1108 assert_eq!(out[0].doc_id, 1);
1109 }
1110
1111 #[test]
1112 fn search_survives_large_synthetic_corpus_on_both_clusters() {
1113 // Bigger synthetic corpus with two distinct unit-norm clusters.
1114 // An east query must top-rank an east document; same for north.
1115 let mut docs = Vec::new();
1116 for i in 0..20 {
1117 let jitter = i as f32 * 0.003;
1118 let east = normalize([1.0 - jitter, 0.05 + jitter]);
1119 docs.push(DocumentTokens {
1120 doc_id: 100 + i,
1121 tokens: east.to_vec(),
1122 n_tokens: 1,
1123 });
1124 let north = normalize([0.05 + jitter, 1.0 - jitter]);
1125 docs.push(DocumentTokens {
1126 doc_id: 200 + i,
1127 tokens: north.to_vec(),
1128 n_tokens: 1,
1129 });
1130 }
1131 let index = build_index(&docs, params()).unwrap();
1132
1133 let east = search(
1134 &index,
1135 &[1.0, 0.0],
1136 SearchParams {
1137 top_k: 5,
1138 n_probe: 2,
1139 n_candidate_docs: None,
1140 centroid_score_threshold: None,
1141 },
1142 )
1143 .unwrap();
1144 for r in &east {
1145 assert!(
1146 (100..120).contains(&r.doc_id),
1147 "east query should only surface east docs: got {}",
1148 r.doc_id,
1149 );
1150 }
1151
1152 let north = search(
1153 &index,
1154 &[0.0, 1.0],
1155 SearchParams {
1156 top_k: 5,
1157 n_probe: 2,
1158 n_candidate_docs: None,
1159 centroid_score_threshold: None,
1160 },
1161 )
1162 .unwrap();
1163 for r in &north {
1164 assert!(
1165 (200..220).contains(&r.doc_id),
1166 "north query should only surface north docs: got {}",
1167 r.doc_id,
1168 );
1169 }
1170 }
1171
1172 #[test]
1173 #[should_panic(expected = "top_k must be positive")]
1174 fn search_panics_on_zero_top_k() {
1175 let index = build_index(&corpus(), params()).unwrap();
1176 let _ = search(
1177 &index,
1178 &[1.0, 0.0],
1179 SearchParams {
1180 top_k: 0,
1181 n_probe: 1,
1182 n_candidate_docs: None,
1183 centroid_score_threshold: None,
1184 },
1185 )
1186 .unwrap();
1187 }
1188
1189 #[test]
1190 #[should_panic(expected = "n_probe must be positive")]
1191 fn search_panics_on_zero_n_probe() {
1192 let index = build_index(&corpus(), params()).unwrap();
1193 let _ = search(
1194 &index,
1195 &[1.0, 0.0],
1196 SearchParams {
1197 top_k: 1,
1198 n_probe: 0,
1199 n_candidate_docs: None,
1200 centroid_score_threshold: None,
1201 },
1202 )
1203 .unwrap();
1204 }
1205
1206 #[test]
1207 #[should_panic(expected = "query length")]
1208 fn search_panics_on_ragged_query() {
1209 let index = build_index(&corpus(), params()).unwrap();
1210 let _ = search(
1211 &index,
1212 &[1.0, 0.0, 0.5],
1213 SearchParams {
1214 top_k: 1,
1215 n_probe: 1,
1216 n_candidate_docs: None,
1217 centroid_score_threshold: None,
1218 },
1219 )
1220 .unwrap();
1221 }
1222
1223 #[test]
1224 fn centroid_interaction_shortlist_caps_candidates_surviving_to_decode() {
1225 // With only one candidate allowed past centroid interaction we
1226 // should get exactly one doc in the results, and it must be the
1227 // best-scoring one under the approximate stage.
1228 let index = build_index(&corpus(), params()).unwrap();
1229 let out = search(
1230 &index,
1231 &[1.0, 0.0],
1232 SearchParams {
1233 top_k: 5,
1234 n_probe: index.codec.num_centroids(),
1235 n_candidate_docs: Some(1),
1236 centroid_score_threshold: None,
1237 },
1238 )
1239 .unwrap();
1240 assert_eq!(out.len(), 1, "shortlist of 1 must cap output at 1");
1241 assert_eq!(
1242 out[0].doc_id, 1,
1243 "east query should surface the east doc as the survivor",
1244 );
1245 }
1246
1247 #[test]
1248 fn centroid_interaction_with_large_shortlist_matches_no_shortlist() {
1249 // A shortlist size at least equal to the number of candidates
1250 // is effectively no shortlisting — the centroid-interaction
1251 // stage picks everything through, so results must agree with
1252 // the legacy `None` path byte-for-byte.
1253 let index = build_index(&corpus(), params()).unwrap();
1254 let query = [1.0, 0.0, 0.0, 1.0];
1255
1256 let legacy = search(
1257 &index,
1258 &query,
1259 SearchParams {
1260 top_k: 3,
1261 n_probe: index.codec.num_centroids(),
1262 n_candidate_docs: None,
1263 centroid_score_threshold: None,
1264 },
1265 )
1266 .unwrap();
1267 let shortlisted = search(
1268 &index,
1269 &query,
1270 SearchParams {
1271 top_k: 3,
1272 n_probe: index.codec.num_centroids(),
1273 n_candidate_docs: Some(1_000),
1274 centroid_score_threshold: None,
1275 },
1276 )
1277 .unwrap();
1278 assert_eq!(legacy, shortlisted);
1279 }
1280
1281 #[test]
1282 fn centroid_pruning_with_low_threshold_matches_no_pruning() {
1283 // A threshold below every attainable centroid score is a
1284 // no-op. For unit-norm centroids the dot product ceiling is
1285 // ~1, so using −1.0 guarantees every centroid survives.
1286 let index = build_index(&corpus(), params()).unwrap();
1287 let query = [1.0, 0.0, 0.0, 1.0];
1288 let unpruned = search(
1289 &index,
1290 &query,
1291 SearchParams {
1292 top_k: 3,
1293 n_probe: index.codec.num_centroids(),
1294 n_candidate_docs: None,
1295 centroid_score_threshold: None,
1296 },
1297 )
1298 .unwrap();
1299 let pruned = search(
1300 &index,
1301 &query,
1302 SearchParams {
1303 top_k: 3,
1304 n_probe: index.codec.num_centroids(),
1305 n_candidate_docs: None,
1306 centroid_score_threshold: Some(-1.0),
1307 },
1308 )
1309 .unwrap();
1310 assert_eq!(unpruned, pruned);
1311 }
1312
1313 #[test]
1314 fn two_stage_interaction_caps_survivors_to_ndocs_div_four() {
1315 // Paper §4.2/§4.3: Stage 2 shortlists to `ndocs`, Stage 3
1316 // refines to `ndocs/4`. With 20 synthetic east/north docs, a
1317 // Stage 2 shortlist of 16, and `top_k` well above the Stage 3
1318 // cap, we should see at most `16/4 = 4` docs in the final
1319 // result — the Stage 3 bound, not `top_k` nor Stage 2's 16.
1320 let mut docs = Vec::new();
1321 for i in 0..20 {
1322 let jitter = i as f32 * 0.003;
1323 let east = normalize([1.0 - jitter, 0.05 + jitter]);
1324 docs.push(DocumentTokens {
1325 doc_id: 100 + i,
1326 tokens: east.to_vec(),
1327 n_tokens: 1,
1328 });
1329 }
1330 let index = build_index(&docs, params()).unwrap();
1331 let out = search(
1332 &index,
1333 &[1.0, 0.0],
1334 SearchParams {
1335 top_k: 20,
1336 n_probe: index.codec.num_centroids(),
1337 n_candidate_docs: Some(16),
1338 centroid_score_threshold: None,
1339 },
1340 )
1341 .unwrap();
1342 assert!(
1343 out.len() <= 4,
1344 "Stage 3 must cap the decoded set at ndocs/4 = 4; got {}",
1345 out.len(),
1346 );
1347 assert!(
1348 !out.is_empty(),
1349 "Stage 3 should still return at least one result"
1350 );
1351 }
1352
1353 #[test]
1354 fn approx_score_skips_tokens_whose_centroid_is_pruned() {
1355 // Paper §4.3: D̃ must only be comprised of tokens whose centroid
1356 // passes the t_cs threshold. This directly exercises that
1357 // behaviour on a controlled qc_scores matrix.
1358 //
1359 // 3 centroids, 2 query tokens. Centroid 0 is strong for q0 only,
1360 // centroid 1 is strong for q1 only, centroid 2 is mediocre for
1361 // both. A doc with tokens in centroids [0, 2] would normally
1362 // pull qc_scores[2] into the q1 max (because centroid 0 is weak
1363 // for q1 and centroid 2 at least has 0.4). Pruning centroid 2
1364 // drops that contribution and only centroid 0's qc values
1365 // remain.
1366 let n_q = 2;
1367 let n_c = 3;
1368 // qc_scores layout is row-major [n_q, n_c].
1369 // q0 row: centroid 0 = 0.9, centroid 1 = 0.0, centroid 2 = 0.4.
1370 // q1 row: centroid 0 = 0.0, centroid 1 = 0.9, centroid 2 = 0.4.
1371 let qc = [0.9, 0.0, 0.4, 0.0, 0.9, 0.4];
1372 let doc_cids: [u32; 2] = [0, 2];
1373
1374 let unpruned =
1375 approx_centroid_interaction_score(&doc_cids, &qc, n_q, n_c, None);
1376 // q0: max(0.9, 0.4) = 0.9; q1: max(0.0, 0.4) = 0.4. Total 1.3.
1377 assert!((unpruned - 1.3).abs() < 1e-5, "unpruned was {unpruned}");
1378
1379 // Prune centroid 2 (per-centroid max 0.4 < t_cs=0.5).
1380 let pruned_mask = [false, false, true];
1381 let pruned = approx_centroid_interaction_score(
1382 &doc_cids,
1383 &qc,
1384 n_q,
1385 n_c,
1386 Some(&pruned_mask),
1387 );
1388 // q0: max over {centroid 0} = 0.9; q1: max over {centroid 0} =
1389 // 0.0. Total 0.9.
1390 assert!((pruned - 0.9).abs() < 1e-5, "pruned was {pruned}");
1391 }
1392
1393 #[test]
1394 fn approx_score_on_all_pruned_doc_is_negative_infinity() {
1395 // Every centroid in the doc is masked. The scorer must return a
1396 // score low enough that the doc loses to any non-empty doc in
1397 // the centroid-interaction sort.
1398 let qc = [0.5, 0.5];
1399 let doc_cids: [u32; 1] = [0];
1400 let mask = [true];
1401 let score = approx_centroid_interaction_score(
1402 &doc_cids,
1403 &qc,
1404 1,
1405 1,
1406 Some(&mask),
1407 );
1408 assert!(
1409 score == f32::NEG_INFINITY,
1410 "all-pruned doc should score -inf, got {score}",
1411 );
1412 }
1413
1414 #[test]
1415 fn centroid_pruning_excludes_docs_whose_only_centroid_is_below_threshold() {
1416 // Build an index with two clusters. An east-pointing query
1417 // scores the east centroid near 1.0 and the north centroid
1418 // near 0.0. A threshold of 0.5 prunes the north centroid,
1419 // so doc 2 (pure-north) must be filtered out — but doc 3
1420 // (mixed east + north) still has an east token and survives.
1421 let index = build_index(&corpus(), params()).unwrap();
1422 let out = search(
1423 &index,
1424 &[1.0f32, 0.0],
1425 SearchParams {
1426 top_k: 5,
1427 n_probe: index.codec.num_centroids(),
1428 n_candidate_docs: None,
1429 centroid_score_threshold: Some(0.5),
1430 },
1431 )
1432 .unwrap();
1433 let ids: Vec<u64> = out.iter().map(|r| r.doc_id).collect();
1434 assert!(
1435 !ids.contains(&2),
1436 "pure-north doc must be pruned when only the east centroid survives the threshold: got {ids:?}",
1437 );
1438 assert!(
1439 ids.contains(&1),
1440 "east-cluster doc must still rank when its centroid survives: got {ids:?}",
1441 );
1442 }
1443
1444 #[test]
1445 fn centroid_pruning_with_unreachable_threshold_drops_every_candidate() {
1446 // No centroid can score higher than ~1.0 against a unit-norm
1447 // query on unit-norm centroids. A threshold of 10.0 prunes
1448 // everything, so no docs survive.
1449 let index = build_index(&corpus(), params()).unwrap();
1450 let out = search(
1451 &index,
1452 &[1.0, 0.0],
1453 SearchParams {
1454 top_k: 3,
1455 n_probe: index.codec.num_centroids(),
1456 n_candidate_docs: None,
1457 centroid_score_threshold: Some(10.0),
1458 },
1459 )
1460 .unwrap();
1461 assert!(
1462 out.is_empty(),
1463 "every centroid below threshold should yield empty results",
1464 );
1465 }
1466
1467 #[test]
1468 fn centroid_interaction_preserves_exact_maxsim_on_surviving_candidates() {
1469 // The final ranking stage still runs exact decoded MaxSim.
1470 // With a shortlist that keeps everyone, the top-1 score must
1471 // equal the MaxSim over decoded tokens for that doc.
1472 let index = build_index(&corpus(), params()).unwrap();
1473 let query = [1.0f32, 0.0];
1474 let out = search(
1475 &index,
1476 &query,
1477 SearchParams {
1478 top_k: 1,
1479 n_probe: index.codec.num_centroids(),
1480 n_candidate_docs: Some(1_000),
1481 centroid_score_threshold: None,
1482 },
1483 )
1484 .unwrap();
1485 let top = out[0];
1486 let doc_idx = index.position_of(top.doc_id).unwrap();
1487 // search() L2-normalises each decoded token before MaxSim;
1488 // the reference has to mirror that.
1489 let decoded: Vec<Vec<f32>> = index
1490 .doc_tokens_vec(doc_idx)
1491 .iter()
1492 .map(|ev| {
1493 let raw = index.codec.decode_vector(ev).unwrap();
1494 let norm =
1495 raw.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
1496 raw.iter().map(|v| v / norm).collect()
1497 })
1498 .collect();
1499 let expected: f32 = query
1500 .as_chunks::<2>()
1501 .0
1502 .iter()
1503 .map(|q| {
1504 decoded
1505 .iter()
1506 .map(|d| dot(q, d))
1507 .fold(f32::NEG_INFINITY, f32::max)
1508 })
1509 .sum();
1510 assert!((expected - top.score).abs() < 1e-5);
1511 }
1512
1513 // -- GPU decode parity --
1514
1515 /// When run on CPU-backed candle, `decode_on_device` still lives on
1516 /// the code path the CUDA build takes. This test forces that path
1517 /// and compares every decoded row to the byte-walking CPU decoder
1518 /// at f32 precision. Any drift in the gather/reshape/narrow chain
1519 /// would show up here without needing a real GPU.
1520 #[test]
1521 fn gpu_decode_path_matches_cpu_decode_element_wise() {
1522 let index = build_index(&corpus(), params()).unwrap();
1523 let candidate_idxs: Vec<usize> = (0..index.num_documents()).collect();
1524 let total_tokens: usize = candidate_idxs
1525 .iter()
1526 .map(|&i| index.doc_token_count(i))
1527 .sum();
1528 let decode_table = DecodeTable::new(&index.codec);
1529 let device = Device::Cpu;
1530
1531 let ctx = DecodeCtx {
1532 index: &index,
1533 candidate_idxs: &candidate_idxs,
1534 dim: index.params.dim,
1535 total_tokens,
1536 decode_table: &decode_table,
1537 device: &device,
1538 };
1539
1540 let mut offsets_gpu = vec![0usize];
1541 let decoded_gpu = decode_on_device(&ctx, &mut offsets_gpu).unwrap();
1542 let gpu_rows: Vec<f32> = decoded_gpu
1543 .to_vec2::<f32>()
1544 .unwrap()
1545 .into_iter()
1546 .flatten()
1547 .collect();
1548
1549 let mut offsets_cpu = vec![0usize];
1550 let decoded_cpu = decode_on_cpu(&ctx, &mut offsets_cpu).unwrap();
1551 let cpu_rows: Vec<f32> = decoded_cpu
1552 .to_vec2::<f32>()
1553 .unwrap()
1554 .into_iter()
1555 .flatten()
1556 .collect();
1557
1558 assert_eq!(offsets_gpu, offsets_cpu);
1559 assert_eq!(gpu_rows.len(), cpu_rows.len());
1560 for (g, c) in gpu_rows.iter().zip(cpu_rows.iter()) {
1561 assert!(
1562 (g - c).abs() < 1e-6,
1563 "GPU decode value {g} != CPU decode value {c}",
1564 );
1565 }
1566 }
1567
1568 // -- Paper Table 2 defaults --
1569
1570 #[test]
1571 fn paper_defaults_match_table_2_verbatim() {
1572 // PLAID paper Table 2: (nprobe, t_cs, ndocs)
1573 // k=10 → (1, 0.50, 256)
1574 // k=100 → (2, 0.45, 1024)
1575 // k=1000 → (4, 0.40, 4096)
1576 let p10 = SearchParams::paper_defaults(10);
1577 assert_eq!(p10.n_probe, 1);
1578 assert_eq!(p10.centroid_score_threshold, Some(0.5));
1579 assert_eq!(p10.n_candidate_docs, Some(256));
1580
1581 let p100 = SearchParams::paper_defaults(100);
1582 assert_eq!(p100.n_probe, 2);
1583 assert_eq!(p100.centroid_score_threshold, Some(0.45));
1584 assert_eq!(p100.n_candidate_docs, Some(1024));
1585
1586 let p1000 = SearchParams::paper_defaults(1000);
1587 assert_eq!(p1000.n_probe, 4);
1588 assert_eq!(p1000.centroid_score_threshold, Some(0.4));
1589 assert_eq!(p1000.n_candidate_docs, Some(4096));
1590 }
1591
1592 #[test]
1593 fn paper_defaults_clamp_ndocs_to_at_least_4x_top_k() {
1594 // If a caller asks for top_k = 2000, Table 2's 4096 ndocs would
1595 // leave Stage 3's ndocs/4 shortlist at 1024 — less than the
1596 // caller's requested 2000. The constructor bumps ndocs to
1597 // 4 * top_k so Stage 3 can still return enough candidates.
1598 let p = SearchParams::paper_defaults(2000);
1599 assert!(p.n_candidate_docs.unwrap() >= 4 * 2000);
1600 }
1601
1602 #[test]
1603 fn paper_defaults_always_enable_pruning() {
1604 // Stage 2 centroid pruning is the single biggest algorithmic
1605 // win in Figure 6 of the paper. `paper_defaults` must never
1606 // leave it off by accident.
1607 for &k in &[1_usize, 10, 50, 100, 500, 1000, 10_000] {
1608 assert!(
1609 SearchParams::paper_defaults(k)
1610 .centroid_score_threshold
1611 .is_some(),
1612 "paper_defaults({k}) left pruning disabled"
1613 );
1614 }
1615 }
1616
1617 /// Splitting `batch_maxsim` across multiple decode chunks must
1618 /// produce byte-identical results to a single-chunk run. Keeps
1619 /// the VRAM-bounding loop honest: any divergence between chunk
1620 /// boundaries would mean the per-doc MaxSim reduction is leaking
1621 /// state across calls, which would quietly corrupt search scores
1622 /// on any real corpus big enough to trip the chunk cap.
1623 #[test]
1624 fn batch_maxsim_chunked_matches_single_shot() {
1625 // A slightly bigger corpus than the shared `corpus()` fixture
1626 // so multiple chunk caps actually force different splits.
1627 let mut docs = Vec::new();
1628 for i in 0..8u64 {
1629 let jitter = i as f32 * 0.015;
1630 let east = normalize([1.0 - jitter, 0.05 + jitter]);
1631 let north = normalize([0.05 + jitter, 1.0 - jitter]);
1632 let mut tokens = Vec::new();
1633 tokens.extend_from_slice(&east);
1634 tokens.extend_from_slice(&east);
1635 tokens.extend_from_slice(&north);
1636 docs.push(DocumentTokens {
1637 doc_id: 200 + i,
1638 tokens,
1639 n_tokens: 3,
1640 });
1641 }
1642 let index = build_index(&docs, params()).unwrap();
1643
1644 // Every doc becomes a candidate so the chunk loop has real
1645 // work to do regardless of the query.
1646 let candidate_idxs: Vec<usize> = (0..docs.len()).collect();
1647 let query = [1.0f32, 0.0, 0.0, 1.0];
1648
1649 let single_shot = batch_maxsim_with_cap(
1650 &query,
1651 &candidate_idxs,
1652 &index,
1653 index.params.dim,
1654 usize::MAX,
1655 )
1656 .unwrap();
1657
1658 // A 1-token cap forces one chunk per doc; 2 tokens forces a
1659 // split mid-doc-group; 5 tokens crosses doc boundaries
1660 // asymmetrically. All three must agree with the single-shot
1661 // run byte-for-byte.
1662 for cap in [1usize, 2, 5] {
1663 let chunked = batch_maxsim_with_cap(
1664 &query,
1665 &candidate_idxs,
1666 &index,
1667 index.params.dim,
1668 cap,
1669 )
1670 .unwrap();
1671 assert_eq!(
1672 chunked.len(),
1673 single_shot.len(),
1674 "chunked result count differs at cap={cap}"
1675 );
1676 for (i, (got, want)) in
1677 chunked.iter().zip(single_shot.iter()).enumerate()
1678 {
1679 assert_eq!(
1680 got.0, want.0,
1681 "doc index order differs at cap={cap}, i={i}"
1682 );
1683 assert!(
1684 (got.1 - want.1).abs() < 1e-5,
1685 "score diverges at cap={cap}, i={i}: got {} expected {}",
1686 got.1,
1687 want.1,
1688 );
1689 }
1690 }
1691 }
1692
1693 #[test]
1694 fn decode_chunk_capacity_scales_inversely_with_dim() {
1695 // Bigger models must produce smaller chunks so the resident
1696 // working set stays inside the 64 MiB budget.
1697 let small = decode_chunk_capacity(128);
1698 let large = decode_chunk_capacity(1536);
1699 assert!(
1700 large < small,
1701 "1536-dim cap {large} should be smaller than 128-dim cap {small}"
1702 );
1703 // Floor at 64 tokens keeps tiny-dim corpora from picking a
1704 // uselessly huge cap that defeats the chunking.
1705 assert!(
1706 decode_chunk_capacity(1) >= 64,
1707 "floor must hold for pathologically small dim"
1708 );
1709 }
1710}