Skip to main content

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(&cent_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], &centroids, 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], &centroids, 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], &centroids, 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], &centroids, 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}