1use futures::{StreamExt, TryStreamExt};
11use rustc_hash::{FxHashMap, FxHashSet};
12use std::sync::Arc;
13use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
14
15use crate::dsl::Field;
16
17use super::{MultiValueCombiner, ScoredPosition, SearchResult, compare_search_results_desc};
18
19const MAX_L2_RERANK_VECTORS: usize = 500_000;
23const MAX_L2_RERANK_VECTOR_BYTES: usize = 512 * 1024 * 1024;
24const RERANK_SCORE_BATCH: usize = 4_096;
25const MAX_RERANK_RAW_BATCH_BYTES: usize = 8 * 1024 * 1024;
26const MAX_CONCURRENT_RERANK_SEGMENTS: usize = 8;
27
28#[derive(Clone, Copy, PartialEq, Eq)]
29enum RerankerKind {
30 Dense,
31 Binary,
32}
33
34fn validate_reranker_config<D: crate::directories::Directory + 'static>(
35 searcher: &crate::index::Searcher<D>,
36 config: &RerankerConfig,
37) -> crate::error::Result<RerankerKind> {
38 if !config.rrf_k.is_finite() || config.rrf_k < 0.0 {
39 return Err(crate::Error::Query(format!(
40 "reranker rrf_k must be finite and non-negative, got {}",
41 config.rrf_k
42 )));
43 }
44 config.combiner.validate().map_err(crate::Error::Query)?;
45 if config.vector.is_empty() == config.binary_vector.is_empty() {
46 return Err(crate::Error::Query(
47 "reranker must provide exactly one of vector or binary_vector".to_string(),
48 ));
49 }
50
51 let entry = searcher
52 .schema()
53 .get_field_entry(config.field)
54 .ok_or_else(|| crate::Error::FieldNotFound(config.field.0.to_string()))?;
55 if !config.binary_vector.is_empty() {
56 if entry.field_type != crate::dsl::FieldType::BinaryDenseVector {
57 return Err(crate::Error::InvalidFieldType {
58 expected: "binary_dense_vector".to_string(),
59 got: format!("{:?}", entry.field_type),
60 });
61 }
62 let field_config = entry.binary_dense_vector_config.as_ref().ok_or_else(|| {
63 crate::Error::Schema(format!(
64 "binary dense vector field '{}' has no configuration",
65 entry.name
66 ))
67 })?;
68 if field_config.dim == 0 || !field_config.dim.is_multiple_of(8) {
69 return Err(crate::Error::Schema(format!(
70 "binary dense vector field '{}' has invalid dimension {}",
71 entry.name, field_config.dim
72 )));
73 }
74 if config.binary_vector.len() != field_config.byte_len() {
75 return Err(crate::Error::Query(format!(
76 "reranker binary vector byte length {} does not match field '{}' byte length {}",
77 config.binary_vector.len(),
78 entry.name,
79 field_config.byte_len()
80 )));
81 }
82 if config.matryoshka_dims.is_some() {
83 return Err(crate::Error::Query(
84 "reranker matryoshka_dims is not supported for binary vectors".to_string(),
85 ));
86 }
87 return Ok(RerankerKind::Binary);
88 }
89
90 if entry.field_type != crate::dsl::FieldType::DenseVector {
91 return Err(crate::Error::InvalidFieldType {
92 expected: "dense_vector".to_string(),
93 got: format!("{:?}", entry.field_type),
94 });
95 }
96 let field_config = entry.dense_vector_config.as_ref().ok_or_else(|| {
97 crate::Error::Schema(format!(
98 "dense vector field '{}' has no configuration",
99 entry.name
100 ))
101 })?;
102 if config.vector.len() != field_config.dim {
103 return Err(crate::Error::Query(format!(
104 "reranker vector dimension {} does not match field '{}' dimension {}",
105 config.vector.len(),
106 entry.name,
107 field_config.dim
108 )));
109 }
110 if let Some((index, value)) = config
111 .vector
112 .iter()
113 .enumerate()
114 .find(|(_, value)| !value.is_finite())
115 {
116 return Err(crate::Error::Query(format!(
117 "reranker vector contains non-finite value {value} at index {index}"
118 )));
119 }
120 if config.unit_norm != field_config.unit_norm {
121 return Err(crate::Error::Query(format!(
122 "reranker unit_norm={} does not match field '{}' unit_norm={}",
123 config.unit_norm, entry.name, field_config.unit_norm
124 )));
125 }
126 if let Some(dims) = config.matryoshka_dims
127 && (dims == 0 || dims > field_config.dim)
128 {
129 return Err(crate::Error::Query(format!(
130 "reranker matryoshka_dims must be in 1..={}, got {dims}",
131 field_config.dim
132 )));
133 }
134 Ok(RerankerKind::Dense)
135}
136
137fn reserve_rerank_vectors(
138 vector_budget: &AtomicUsize,
139 byte_budget: &AtomicUsize,
140 count: usize,
141 vector_byte_size: usize,
142) -> crate::error::Result<()> {
143 let bytes = count.checked_mul(vector_byte_size).ok_or_else(|| {
144 crate::Error::Query("reranker stored-vector byte budget overflow".to_string())
145 })?;
146 byte_budget
147 .fetch_update(AtomicOrdering::Relaxed, AtomicOrdering::Relaxed, |used| {
148 used.checked_add(bytes)
149 .filter(|&next| next <= MAX_L2_RERANK_VECTOR_BYTES)
150 })
151 .map_err(|used| {
152 crate::Error::Query(format!(
153 "reranker reads more than {MAX_L2_RERANK_VECTOR_BYTES} stored vector bytes \
154 (already reserved {used}, next document needs {bytes})"
155 ))
156 })?;
157
158 vector_budget
159 .fetch_update(AtomicOrdering::Relaxed, AtomicOrdering::Relaxed, |used| {
160 used.checked_add(count)
161 .filter(|&next| next <= MAX_L2_RERANK_VECTORS)
162 })
163 .map(|_| ())
164 .map_err(|used| {
165 crate::Error::Query(format!(
166 "reranker expands to more than {MAX_L2_RERANK_VECTORS} stored vectors \
167 (already reserved {used}, next document has {count})"
168 ))
169 })
170}
171
172#[inline]
173fn rerank_batch_len(vector_byte_size: usize) -> usize {
174 RERANK_SCORE_BATCH.min((MAX_RERANK_RAW_BATCH_BYTES / vector_byte_size.max(1)).max(1))
175}
176
177struct PrecompQuery<'a> {
179 query: &'a [f32],
180 inv_norm_q: f32,
181 query_f16: &'a [u16],
182}
183
184#[inline]
186#[allow(clippy::too_many_arguments)]
187fn score_batch_precomp(
188 pq: &PrecompQuery<'_>,
189 raw: &[u8],
190 quant: crate::dsl::DenseVectorQuantization,
191 dim: usize,
192 scores: &mut [f32],
193 unit_norm: bool,
194) -> crate::error::Result<()> {
195 let query = pq.query;
196 let inv_norm_q = pq.inv_norm_q;
197 let query_f16 = pq.query_f16;
198 use crate::dsl::DenseVectorQuantization;
199 use crate::structures::simd;
200 let element_size = quant.element_size();
201 let required_bytes = scores
202 .len()
203 .checked_mul(dim)
204 .and_then(|elements| elements.checked_mul(element_size))
205 .ok_or_else(|| {
206 crate::Error::Corruption("dense reranker batch size overflow".to_string())
207 })?;
208 if raw.len() < required_bytes {
209 return Err(crate::Error::Corruption(format!(
210 "dense reranker batch is truncated: need {required_bytes} bytes, got {}",
211 raw.len()
212 )));
213 }
214 if matches!(
215 quant,
216 DenseVectorQuantization::F32 | DenseVectorQuantization::F16
217 ) && required_bytes > 0
218 && !(raw.as_ptr() as usize).is_multiple_of(element_size)
219 {
220 return Err(crate::Error::Corruption(format!(
221 "dense reranker {:?} data is not {}-byte aligned",
222 quant, element_size
223 )));
224 }
225 match (quant, unit_norm) {
226 (DenseVectorQuantization::F32, false) => {
227 let num_floats = scores.len() * dim;
228 let vectors: &[f32] =
232 unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const f32, num_floats) };
233 simd::batch_cosine_scores_precomp(query, vectors, dim, scores, inv_norm_q);
234 }
235 (DenseVectorQuantization::F32, true) => {
236 let num_floats = scores.len() * dim;
237 let vectors: &[f32] =
238 unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const f32, num_floats) };
239 simd::batch_dot_scores_precomp(query, vectors, dim, scores, inv_norm_q);
240 }
241 (DenseVectorQuantization::F16, false) => {
242 simd::batch_cosine_scores_f16_precomp(query_f16, raw, dim, scores, inv_norm_q);
243 }
244 (DenseVectorQuantization::F16, true) => {
245 simd::batch_dot_scores_f16_precomp(query_f16, raw, dim, scores, inv_norm_q);
246 }
247 (DenseVectorQuantization::UInt8, false) => {
248 simd::batch_cosine_scores_u8_precomp(query, raw, dim, scores, inv_norm_q);
249 }
250 (DenseVectorQuantization::UInt8, true) => {
251 simd::batch_dot_scores_u8_precomp(query, raw, dim, scores, inv_norm_q);
252 }
253 (DenseVectorQuantization::Binary, _) => {
254 return Err(crate::Error::InvalidFieldType {
255 expected: "non-binary dense vector".to_string(),
256 got: "binary dense vector".to_string(),
257 });
258 }
259 }
260 Ok(())
261}
262
263#[derive(Debug, Clone)]
265pub struct RerankerConfig {
266 pub field: Field,
268 pub vector: Vec<f32>,
270 pub binary_vector: Vec<u8>,
273 pub combiner: MultiValueCombiner,
275 pub unit_norm: bool,
279 pub matryoshka_dims: Option<usize>,
283 pub rrf_k: f32,
287}
288
289#[cfg(test)]
291use crate::structures::simd::cosine_similarity;
292#[cfg(test)]
293fn score_document(
294 doc: &crate::dsl::Document,
295 config: &RerankerConfig,
296) -> Option<(f32, Vec<ScoredPosition>)> {
297 let query_dim = config.vector.len();
298 let mut values: Vec<(u32, f32)> = doc
299 .get_all(config.field)
300 .filter_map(|fv| fv.as_dense_vector())
301 .enumerate()
302 .filter_map(|(ordinal, vec)| {
303 if vec.len() != query_dim {
304 return None;
305 }
306 let score = cosine_similarity(&config.vector, vec);
307 Some((ordinal as u32, score))
308 })
309 .collect();
310
311 if values.is_empty() {
312 return None;
313 }
314
315 let combined = config.combiner.combine(&values);
316
317 values.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
319 let positions: Vec<ScoredPosition> = values
320 .into_iter()
321 .map(|(ordinal, score)| ScoredPosition::new(ordinal, score))
322 .collect();
323
324 Some((combined, positions))
325}
326
327fn apply_rrf(
337 candidates: &[SearchResult],
338 scored: &mut Vec<SearchResult>,
339 k: f32,
340 final_limit: usize,
341) {
342 let l1_ranks: FxHashMap<(u128, u32), usize> = candidates
344 .iter()
345 .enumerate()
346 .map(|(idx, c)| ((c.segment_id, c.doc_id), idx + 1))
347 .collect();
348
349 for (l2_idx, result) in scored.iter_mut().enumerate() {
351 let l1_rank = l1_ranks
352 .get(&(result.segment_id, result.doc_id))
353 .copied()
354 .unwrap_or(candidates.len() + 1);
355 result.score = super::fusion::rrf_contribution(k, l1_rank)
356 + super::fusion::rrf_contribution(k, l2_idx + 1);
357 }
358
359 scored.sort_unstable_by(compare_search_results_desc);
360 scored.truncate(final_limit);
361}
362
363pub async fn rerank<D: crate::directories::Directory + 'static>(
372 searcher: &crate::index::Searcher<D>,
373 candidates: &[SearchResult],
374 config: &RerankerConfig,
375 final_limit: usize,
376) -> crate::error::Result<Vec<SearchResult>> {
377 let kind = validate_reranker_config(searcher, config)?;
380 if final_limit == 0 || candidates.is_empty() {
381 return Ok(Vec::new());
382 }
383
384 if kind == RerankerKind::Binary {
386 return rerank_binary(searcher, candidates, config, final_limit).await;
387 }
388
389 let t0 = std::time::Instant::now();
390 let field_id = config.field.0;
391 let query = &config.vector;
392 let query_dim = query.len();
393 let segments = searcher.segment_readers();
394 let seg_by_id = searcher.segment_map();
395
396 use crate::structures::simd;
398 let norm_q_sq = simd::dot_product_f32(query, query, query_dim);
399 let inv_norm_q = if norm_q_sq < f32::EPSILON {
400 0.0
401 } else {
402 simd::fast_inv_sqrt(norm_q_sq)
403 };
404 let query_f16: Vec<u16> = query.iter().map(|&v| simd::f32_to_f16(v)).collect();
405 let pq = PrecompQuery {
406 query,
407 inv_norm_q,
408 query_f16: &query_f16,
409 };
410
411 let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
413 let mut skipped = 0u32;
414
415 for (ci, candidate) in candidates.iter().enumerate() {
416 if let Some(&si) = seg_by_id.get(&candidate.segment_id) {
417 segment_groups.entry(si).or_default().push(ci);
418 } else {
419 skipped += 1;
420 }
421 }
422
423 let query_ref = pq.query;
427 let inv_norm_q_val = pq.inv_norm_q;
428 let query_f16_ref = pq.query_f16;
429 let vector_budget = Arc::new(AtomicUsize::new(0));
430 let byte_budget = Arc::new(AtomicUsize::new(0));
431
432 let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
433 |(si, candidate_indices)| {
434 #[allow(clippy::redundant_locals)]
435 let segments = &segments;
436 #[allow(clippy::redundant_locals)]
437 let candidates = candidates;
438 #[allow(clippy::redundant_locals)]
439 let query_ref = query_ref;
440 #[allow(clippy::redundant_locals)]
441 let query_f16_ref = query_f16_ref;
442 #[allow(clippy::redundant_locals)]
443 let config = config;
444 let vector_budget = Arc::clone(&vector_budget);
445 let byte_budget = Arc::clone(&byte_budget);
446 async move {
447 let mut scores: Vec<(usize, u32, f32)> = Vec::new();
448 let mut vectors = 0usize;
449 let mut seg_skipped = 0u32;
450
451 let Some(lazy_flat) = segments[si].flat_vectors().get(&field_id) else {
452 return Ok::<_, crate::error::Error>((
453 scores,
454 vectors,
455 candidate_indices.len() as u32,
456 ));
457 };
458 if lazy_flat.dim != query_dim {
459 return Err(crate::Error::Corruption(format!(
460 "dense reranker field {field_id} stores dimension {}, expected {query_dim}",
461 lazy_flat.dim
462 )));
463 }
464 if lazy_flat.quantization == crate::dsl::DenseVectorQuantization::Binary {
465 return Err(crate::Error::Corruption(format!(
466 "dense reranker field {field_id} unexpectedly uses binary storage"
467 )));
468 }
469
470 let vbs = lazy_flat.vector_byte_size();
471 let quant = lazy_flat.quantization;
472
473 let mut resolved: Vec<(usize, usize, u32)> = Vec::new();
475 for &ci in &candidate_indices {
476 let local_doc_id = candidates[ci].doc_id;
477 let (start, count) = lazy_flat.flat_indexes_for_doc_range(local_doc_id);
478 if count == 0 {
479 seg_skipped += 1;
480 continue;
481 }
482 reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
483 for j in 0..count {
484 let (_, ordinal) = lazy_flat.get_doc_id(start + j);
485 resolved.push((ci, start + j, ordinal as u32));
486 }
487 }
488
489 if resolved.is_empty() {
490 return Ok((scores, vectors, seg_skipped));
491 }
492
493 let n = resolved.len();
494 vectors = n;
495
496 resolved.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
500
501 let batch_len = rerank_batch_len(vbs);
502 let max_batch = batch_len.min(n);
503 let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
504 crate::Error::Query("dense reranker buffer size overflow".into())
505 })?;
506 let mut raw_buf = vec![0u8; max_raw_len];
507
508 let pq = PrecompQuery {
510 query: query_ref,
511 inv_norm_q: inv_norm_q_val,
512 query_f16: query_f16_ref,
513 };
514
515 if let Some(mdims) = config.matryoshka_dims
517 && mdims < query_dim
518 && n > final_limit.saturating_mul(2)
519 {
520 let trunc_dim = mdims;
521 let trunc_pq = PrecompQuery {
522 query: &query_ref[..trunc_dim],
523 inv_norm_q: {
524 let nq = simd::dot_product_f32(
525 &query_ref[..trunc_dim],
526 &query_ref[..trunc_dim],
527 trunc_dim,
528 );
529 if nq < f32::EPSILON {
530 0.0
531 } else {
532 simd::fast_inv_sqrt(nq)
533 }
534 },
535 query_f16: &query_f16_ref[..trunc_dim],
536 };
537 let trunc_vbs = trunc_dim * quant.element_size();
538 let mut scores_buf = vec![0.0f32; n];
539 for (chunk_idx, chunk) in resolved.chunks(batch_len).enumerate() {
540 let raw_len = chunk.len().checked_mul(trunc_vbs).ok_or_else(|| {
544 crate::Error::Query("dense reranker buffer size overflow".into())
545 })?;
546 let raw = &mut raw_buf[..raw_len];
547 for (buf_idx, &(_, flat_idx, _)) in chunk.iter().enumerate() {
548 lazy_flat
549 .read_vector_prefix_raw_into(
550 flat_idx,
551 trunc_vbs,
552 &mut raw[buf_idx * trunc_vbs..(buf_idx + 1) * trunc_vbs],
553 )
554 .await
555 .map_err(crate::error::Error::Io)?;
556 }
557 let score_base = chunk_idx * batch_len;
558 searcher.install_search_cpu(|| {
559 score_batch_precomp(
560 &trunc_pq,
561 raw,
562 quant,
563 trunc_dim,
564 &mut scores_buf[score_base..score_base + chunk.len()],
565 config.unit_norm,
566 )
567 })?;
568 }
569
570 let mut approximate_ordinals: FxHashMap<usize, Vec<(u32, f32)>> =
575 FxHashMap::default();
576 for (resolved_index, &(ci, _, ordinal)) in resolved.iter().enumerate() {
577 approximate_ordinals
578 .entry(ci)
579 .or_default()
580 .push((ordinal, scores_buf[resolved_index]));
581 }
582 let mut ranked: Vec<(usize, f32)> = approximate_ordinals
583 .into_iter()
584 .map(|(ci, ordinals)| (ci, config.combiner.combine(&ordinals)))
585 .collect();
586 searcher.install_search_cpu(|| {
587 ranked.sort_unstable_by(|a, b| {
588 b.1.total_cmp(&a.1)
589 .then_with(|| candidates[a.0].doc_id.cmp(&candidates[b.0].doc_id))
590 });
591 });
592 let approximate_docs = ranked.len();
593 let survivor_doc_limit = final_limit.saturating_mul(2).min(approximate_docs);
594 let survivor_docs: FxHashSet<usize> = ranked
595 .into_iter()
596 .take(survivor_doc_limit)
597 .map(|(ci, _)| ci)
598 .collect();
599 let mut survivor_entries: Vec<_> = resolved
602 .iter()
603 .copied()
604 .filter(|(ci, _, _)| survivor_docs.contains(ci))
605 .collect();
606 survivor_entries.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
607 let mut full_scores = vec![0.0f32; max_batch.min(survivor_entries.len())];
608 scores.reserve(survivor_entries.len());
609 for chunk in survivor_entries.chunks(batch_len) {
610 #[cfg(feature = "native")]
611 lazy_flat.prefetch_vectors(
612 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
613 );
614 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
615 crate::Error::Query("dense reranker buffer size overflow".into())
616 })?;
617 let raw = &mut raw_buf[..raw_len];
618 for (buf_idx, &(_, flat_idx, _)) in chunk.iter().enumerate() {
619 lazy_flat
620 .read_vector_raw_into(
621 flat_idx,
622 &mut raw[buf_idx * vbs..(buf_idx + 1) * vbs],
623 )
624 .await
625 .map_err(crate::error::Error::Io)?;
626 }
627 searcher.install_search_cpu(|| {
628 score_batch_precomp(
629 &pq,
630 raw,
631 quant,
632 query_dim,
633 &mut full_scores[..chunk.len()],
634 config.unit_norm,
635 )
636 })?;
637 for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
638 scores.push((ci, ordinal, full_scores[buf_idx]));
639 }
640 }
641
642 let survivor_vectors = survivor_entries.len();
643 log::debug!(
644 "[reranker] matryoshka pre-filter: {}/{} dims, {}/{} docs and {}/{} vectors survived",
645 trunc_dim,
646 query_dim,
647 survivor_docs.len(),
648 approximate_docs,
649 survivor_vectors,
650 n,
651 );
652 } else {
653 let mut scores_buf = vec![0.0f32; max_batch];
654 scores.reserve(n);
655 for chunk in resolved.chunks(batch_len) {
656 #[cfg(feature = "native")]
657 lazy_flat.prefetch_vectors(
658 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
659 );
660 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
661 crate::Error::Query("dense reranker buffer size overflow".into())
662 })?;
663 let raw = &mut raw_buf[..raw_len];
664 for (buf_idx, &(_, flat_idx, _)) in chunk.iter().enumerate() {
665 lazy_flat
666 .read_vector_raw_into(
667 flat_idx,
668 &mut raw[buf_idx * vbs..(buf_idx + 1) * vbs],
669 )
670 .await
671 .map_err(crate::error::Error::Io)?;
672 }
673 searcher.install_search_cpu(|| {
674 score_batch_precomp(
675 &pq,
676 raw,
677 quant,
678 query_dim,
679 &mut scores_buf[..chunk.len()],
680 config.unit_norm,
681 )
682 })?;
683 for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
684 scores.push((ci, ordinal, scores_buf[buf_idx]));
685 }
686 }
687 }
688
689 Ok((scores, vectors, seg_skipped))
690 }
691 },
692 ))
693 .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
694 futures::pin_mut!(segment_futs);
695
696 let mut all_scores: Vec<(usize, u32, f32)> = Vec::new();
697 let mut total_vectors = 0usize;
698 while let Some((scores, vectors, seg_skipped)) = segment_futs.try_next().await? {
699 all_scores.extend(scores);
700 total_vectors = total_vectors.saturating_add(vectors);
701 skipped = skipped.saturating_add(seg_skipped);
702 }
703
704 let read_score_elapsed = t0.elapsed();
705
706 if total_vectors == 0 {
707 log::debug!(
708 "[reranker] field {}: {} candidates, all skipped (no flat vectors)",
709 field_id,
710 candidates.len()
711 );
712 return Ok(Vec::new());
713 }
714
715 all_scores.sort_unstable_by_key(|&(ci, _, _)| ci);
718
719 let mut scored: Vec<SearchResult> = Vec::with_capacity(candidates.len().min(final_limit * 2));
720 let mut ordinal_pairs: Vec<(u32, f32)> = Vec::new();
721 let mut i = 0;
722 while i < all_scores.len() {
723 let ci = all_scores[i].0;
724 let run_start = i;
725 while i < all_scores.len() && all_scores[i].0 == ci {
726 i += 1;
727 }
728 let run = &mut all_scores[run_start..i];
729
730 ordinal_pairs.clear();
732 ordinal_pairs.extend(run.iter().map(|&(_, ord, s)| (ord, s)));
733 let combined = config.combiner.combine(&ordinal_pairs);
734
735 run.sort_unstable_by(|a, b| b.2.total_cmp(&a.2));
737 let positions: Vec<ScoredPosition> = run
738 .iter()
739 .map(|&(_, ord, score)| ScoredPosition::new(ord, score))
740 .collect();
741
742 scored.push(SearchResult {
743 doc_id: candidates[ci].doc_id,
744 score: combined,
745 segment_id: candidates[ci].segment_id,
746 positions: vec![(field_id, positions)],
747 });
748 }
749
750 scored.sort_unstable_by(compare_search_results_desc);
751
752 if config.rrf_k > 0.0 {
753 apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
754 } else {
755 scored.truncate(final_limit);
756 }
757
758 log::debug!(
759 "[reranker] field {}: {} candidates -> {} results (skipped {}, {} vectors, unit_norm={}, rrf_k={}): read+score={:.1}ms total={:.1}ms",
760 field_id,
761 candidates.len(),
762 scored.len(),
763 skipped,
764 total_vectors,
765 config.unit_norm,
766 config.rrf_k,
767 read_score_elapsed.as_secs_f64() * 1000.0,
768 t0.elapsed().as_secs_f64() * 1000.0,
769 );
770
771 Ok(scored)
772}
773
774async fn rerank_binary<D: crate::directories::Directory + 'static>(
776 searcher: &crate::index::Searcher<D>,
777 candidates: &[SearchResult],
778 config: &RerankerConfig,
779 final_limit: usize,
780) -> crate::error::Result<Vec<SearchResult>> {
781 if config.binary_vector.is_empty() || candidates.is_empty() {
782 return Ok(Vec::new());
783 }
784
785 let t0 = std::time::Instant::now();
786 let field_id = config.field.0;
787 let query = &config.binary_vector;
788 let byte_len = query.len();
789 let segments = searcher.segment_readers();
790 let seg_by_id = searcher.segment_map();
791
792 let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
794 for (ci, cand) in candidates.iter().enumerate() {
795 if let Some(&seg_idx) = seg_by_id.get(&cand.segment_id) {
796 let reader = &segments[seg_idx];
797 if reader.flat_vectors().contains_key(&field_id) {
798 segment_groups.entry(seg_idx).or_default().push(ci);
799 }
800 }
801 }
802
803 let vector_budget = Arc::new(AtomicUsize::new(0));
805 let byte_budget = Arc::new(AtomicUsize::new(0));
806 let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
807 |(seg_idx, cand_indices)| {
808 #[allow(clippy::redundant_locals)]
809 let segments = &segments;
810 #[allow(clippy::redundant_locals)]
811 let candidates = candidates;
812 let vector_budget = Arc::clone(&vector_budget);
813 let byte_budget = Arc::clone(&byte_budget);
814 async move {
815 let mut scores: Vec<(usize, u32, f32)> = Vec::new();
816
817 let Some(lazy_flat) = segments[seg_idx].flat_vectors().get(&field_id) else {
818 return Ok::<_, crate::error::Error>(scores);
819 };
820 if lazy_flat.quantization != crate::dsl::DenseVectorQuantization::Binary
821 || !lazy_flat.dim.is_multiple_of(8)
822 {
823 return Err(crate::Error::Corruption(format!(
824 "binary reranker field {field_id} has invalid flat-vector metadata"
825 )));
826 }
827 let vbs = lazy_flat.vector_byte_size();
828 if vbs != byte_len {
829 return Err(crate::Error::Corruption(format!(
830 "binary reranker field {field_id} stores {vbs} bytes/vector, expected {byte_len}"
831 )));
832 }
833
834 let mut resolved: Vec<(usize, usize)> = Vec::new();
836 for &ci in &cand_indices {
837 let doc_id = candidates[ci].doc_id;
838 let (start, count) = lazy_flat.flat_indexes_for_doc_range(doc_id);
839 reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
840 for j in 0..count {
841 resolved.push((ci, start + j));
842 }
843 }
844 if resolved.is_empty() {
845 return Ok(scores);
846 }
847
848 resolved.sort_unstable_by_key(|&(_, flat_idx)| flat_idx);
849
850 let n = resolved.len();
851 let batch_len = rerank_batch_len(vbs);
852 let max_batch = batch_len.min(n);
853 let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
854 crate::Error::Query("binary reranker buffer size overflow".into())
855 })?;
856 let mut raw_buf = vec![0u8; max_raw_len];
857 let mut scores_buf = vec![0f32; max_batch];
858 scores.reserve(n);
859
860 for chunk in resolved.chunks(batch_len) {
861 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
862 crate::Error::Query("binary reranker buffer size overflow".into())
863 })?;
864 let raw = &mut raw_buf[..raw_len];
865 for (buf_idx, &(_, flat_idx)) in chunk.iter().enumerate() {
866 lazy_flat
867 .read_vector_raw_into(
868 flat_idx,
869 &mut raw[buf_idx * vbs..(buf_idx + 1) * vbs],
870 )
871 .await
872 .map_err(crate::error::Error::Io)?;
873 }
874 searcher.install_search_cpu(|| {
875 crate::structures::simd::batch_hamming_scores(
876 query,
877 raw,
878 byte_len,
879 lazy_flat.dim,
880 &mut scores_buf[..chunk.len()],
881 );
882 });
883
884 for (buf_idx, &(ci, flat_idx)) in chunk.iter().enumerate() {
885 let (_, ordinal) = lazy_flat.get_doc_id(flat_idx);
886 scores.push((ci, ordinal as u32, scores_buf[buf_idx]));
887 }
888 }
889
890 Ok(scores)
891 }
892 },
893 ))
894 .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
895 futures::pin_mut!(segment_futs);
896
897 let mut cand_ordinal_scores: FxHashMap<usize, Vec<(u32, f32)>> = FxHashMap::default();
899 while let Some(scores) = segment_futs.try_next().await? {
900 for (ci, ordinal, score) in scores {
901 cand_ordinal_scores
902 .entry(ci)
903 .or_default()
904 .push((ordinal, score));
905 }
906 }
907
908 let total_vectors = cand_ordinal_scores.len();
909 let mut scored: Vec<SearchResult> = Vec::with_capacity(total_vectors);
910 for (ci, ordinal_scores) in cand_ordinal_scores {
911 let combined = config.combiner.combine(&ordinal_scores);
912 let positions: Vec<ScoredPosition> = ordinal_scores
913 .iter()
914 .map(|&(ord, s)| ScoredPosition::new(ord, s))
915 .collect();
916 scored.push(SearchResult {
917 doc_id: candidates[ci].doc_id,
918 score: combined,
919 segment_id: candidates[ci].segment_id,
920 positions: vec![(field_id, positions)],
921 });
922 }
923
924 scored.sort_unstable_by(compare_search_results_desc);
925
926 if config.rrf_k > 0.0 {
927 apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
928 } else {
929 scored.truncate(final_limit);
930 }
931
932 log::debug!(
933 "[reranker-binary] field {}: {} candidates -> {} results ({} docs scored, {} bytes/vec, rrf_k={}): {:.1}ms",
934 field_id,
935 candidates.len(),
936 scored.len(),
937 total_vectors,
938 byte_len,
939 config.rrf_k,
940 t0.elapsed().as_secs_f64() * 1000.0,
941 );
942
943 Ok(scored)
944}
945
946#[cfg(test)]
947mod tests {
948 use super::*;
949 use crate::dsl::{Document, Field};
950
951 fn make_config(vector: Vec<f32>, combiner: MultiValueCombiner) -> RerankerConfig {
952 RerankerConfig {
953 field: Field(0),
954 vector,
955 binary_vector: Vec::new(),
956 combiner,
957 unit_norm: false,
958 matryoshka_dims: None,
959 rrf_k: 0.0,
960 }
961 }
962
963 #[test]
964 fn rerank_batches_are_bounded_by_bytes() {
965 assert_eq!(rerank_batch_len(1), RERANK_SCORE_BATCH);
966 assert_eq!(
967 rerank_batch_len(MAX_RERANK_RAW_BATCH_BYTES),
968 1,
969 "one very wide vector must still make progress"
970 );
971 assert!(
972 rerank_batch_len(4_096) * 4_096 <= MAX_RERANK_RAW_BATCH_BYTES,
973 "normal batches must stay within the raw scratch budget"
974 );
975 }
976
977 #[test]
978 fn rerank_budget_bounds_count_and_bytes() {
979 let vectors = AtomicUsize::new(0);
980 let bytes = AtomicUsize::new(0);
981 reserve_rerank_vectors(&vectors, &bytes, 2, 32).unwrap();
982 assert_eq!(vectors.load(AtomicOrdering::Relaxed), 2);
983 assert_eq!(bytes.load(AtomicOrdering::Relaxed), 64);
984
985 let vectors = AtomicUsize::new(0);
986 let bytes = AtomicUsize::new(0);
987 assert!(reserve_rerank_vectors(&vectors, &bytes, 2, MAX_L2_RERANK_VECTOR_BYTES).is_err());
988 }
989
990 #[test]
991 fn test_score_document_single_value() {
992 let mut doc = Document::new();
993 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
994
995 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
996 let (score, positions) = score_document(&doc, &config).unwrap();
997 assert!((score - 1.0).abs() < 1e-6);
999 assert_eq!(positions.len(), 1);
1000 assert_eq!(positions[0].position, 0); }
1002
1003 #[test]
1004 fn test_score_document_orthogonal() {
1005 let mut doc = Document::new();
1006 doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]);
1007
1008 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1009 let (score, _) = score_document(&doc, &config).unwrap();
1010 assert!(score.abs() < 1e-6);
1012 }
1013
1014 #[test]
1015 fn test_score_document_multi_value_max() {
1016 let mut doc = Document::new();
1017 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]); doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]); let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1021 let (score, positions) = score_document(&doc, &config).unwrap();
1022 assert!((score - 1.0).abs() < 1e-6);
1023 assert_eq!(positions.len(), 2);
1025 assert_eq!(positions[0].position, 0); assert!((positions[0].score - 1.0).abs() < 1e-6);
1027 }
1028
1029 #[test]
1030 fn test_score_document_multi_value_avg() {
1031 let mut doc = Document::new();
1032 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]); doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]); let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Avg);
1036 let (score, _) = score_document(&doc, &config).unwrap();
1037 assert!((score - 0.5).abs() < 1e-6);
1039 }
1040
1041 #[test]
1042 fn test_score_document_missing_field() {
1043 let mut doc = Document::new();
1044 doc.add_dense_vector(Field(1), vec![1.0, 0.0, 0.0]);
1046
1047 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1048 assert!(score_document(&doc, &config).is_none());
1049 }
1050
1051 #[test]
1052 fn test_score_document_wrong_field_type() {
1053 let mut doc = Document::new();
1054 doc.add_text(Field(0), "not a vector");
1055
1056 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1057 assert!(score_document(&doc, &config).is_none());
1058 }
1059
1060 #[test]
1061 fn test_score_document_dimension_mismatch() {
1062 let mut doc = Document::new();
1063 doc.add_dense_vector(Field(0), vec![1.0, 0.0]); let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max); assert!(score_document(&doc, &config).is_none());
1067 }
1068
1069 #[test]
1070 fn test_score_document_empty_query_vector() {
1071 let mut doc = Document::new();
1072 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
1073
1074 let config = make_config(vec![], MultiValueCombiner::Max);
1075 assert!(score_document(&doc, &config).is_none());
1077 }
1078
1079 fn make_result(doc_id: u32, score: f32, segment_id: u128) -> SearchResult {
1080 SearchResult {
1081 doc_id,
1082 score,
1083 segment_id,
1084 positions: Vec::new(),
1085 }
1086 }
1087
1088 #[test]
1089 fn test_rrf_basic_fusion() {
1090 let candidates = vec![
1092 make_result(1, 10.0, 1), make_result(2, 8.0, 1), make_result(3, 5.0, 1), ];
1096
1097 let mut scored = vec![
1099 make_result(3, 0.9, 1), make_result(2, 0.7, 1), make_result(1, 0.3, 1), ];
1103
1104 let k = 60.0;
1105 apply_rrf(&candidates, &mut scored, k, 10);
1106
1107 assert_eq!(scored.len(), 3);
1120 let spread = scored[0].score - scored[2].score;
1122 assert!(
1123 spread < 0.001,
1124 "All docs have near-equal RRF scores, spread={spread}"
1125 );
1126 }
1127
1128 #[test]
1129 fn test_rrf_clear_winner() {
1130 let candidates = vec![
1132 make_result(1, 10.0, 1), make_result(2, 8.0, 1), make_result(3, 5.0, 1), ];
1136
1137 let mut scored = vec![
1139 make_result(1, 0.95, 1), make_result(3, 0.50, 1), make_result(2, 0.30, 1), ];
1143
1144 let k = 60.0;
1145 apply_rrf(&candidates, &mut scored, k, 10);
1146
1147 assert_eq!(scored[0].doc_id, 1, "Doc 1 (rank 1 in both) should be top");
1151 assert!(scored[0].score > scored[1].score);
1152 }
1153
1154 #[test]
1155 fn test_rrf_truncation() {
1156 let candidates = vec![
1157 make_result(1, 10.0, 1),
1158 make_result(2, 8.0, 1),
1159 make_result(3, 5.0, 1),
1160 make_result(4, 3.0, 1),
1161 make_result(5, 1.0, 1),
1162 ];
1163
1164 let mut scored = vec![
1165 make_result(5, 0.9, 1),
1166 make_result(4, 0.8, 1),
1167 make_result(3, 0.7, 1),
1168 make_result(2, 0.6, 1),
1169 make_result(1, 0.5, 1),
1170 ];
1171
1172 apply_rrf(&candidates, &mut scored, 60.0, 3);
1173 assert_eq!(scored.len(), 3, "Should truncate to final_limit=3");
1174 }
1175
1176 #[test]
1177 fn test_rrf_missing_l1_candidate() {
1178 let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1180
1181 let mut scored = vec![
1182 make_result(3, 0.9, 1), make_result(1, 0.5, 1),
1184 ];
1185
1186 apply_rrf(&candidates, &mut scored, 60.0, 10);
1187
1188 assert_eq!(scored[0].doc_id, 1);
1192 }
1193
1194 #[test]
1195 fn test_rrf_small_k() {
1196 let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1198
1199 let mut scored = vec![
1200 make_result(2, 0.9, 1), make_result(1, 0.5, 1), ];
1203
1204 apply_rrf(&candidates, &mut scored, 1.0, 10);
1205
1206 let diff = (scored[0].score - scored[1].score).abs();
1210 assert!(
1211 diff < 1e-6,
1212 "Symmetric ranks should produce equal RRF scores"
1213 );
1214 }
1215}