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 > crate::query::max_candidate_limit(final_limit)
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 =
594 crate::query::max_candidate_limit(final_limit).min(approximate_docs);
595 let survivor_docs: FxHashSet<usize> = ranked
596 .into_iter()
597 .take(survivor_doc_limit)
598 .map(|(ci, _)| ci)
599 .collect();
600 let mut survivor_entries: Vec<_> = resolved
603 .iter()
604 .copied()
605 .filter(|(ci, _, _)| survivor_docs.contains(ci))
606 .collect();
607 survivor_entries.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
608 let mut full_scores = vec![0.0f32; max_batch.min(survivor_entries.len())];
609 scores.reserve(survivor_entries.len());
610 for chunk in survivor_entries.chunks(batch_len) {
611 #[cfg(feature = "native")]
612 lazy_flat.prefetch_vectors(
613 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
614 );
615 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
616 crate::Error::Query("dense reranker buffer size overflow".into())
617 })?;
618 let raw = &mut raw_buf[..raw_len];
619 for (buf_idx, &(_, flat_idx, _)) in chunk.iter().enumerate() {
620 lazy_flat
621 .read_vector_raw_into(
622 flat_idx,
623 &mut raw[buf_idx * vbs..(buf_idx + 1) * vbs],
624 )
625 .await
626 .map_err(crate::error::Error::Io)?;
627 }
628 searcher.install_search_cpu(|| {
629 score_batch_precomp(
630 &pq,
631 raw,
632 quant,
633 query_dim,
634 &mut full_scores[..chunk.len()],
635 config.unit_norm,
636 )
637 })?;
638 for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
639 scores.push((ci, ordinal, full_scores[buf_idx]));
640 }
641 }
642
643 let survivor_vectors = survivor_entries.len();
644 log::debug!(
645 "[dense_vector_rerank] matryoshka pre-filter: {}/{} dims, {}/{} docs and {}/{} vectors survived",
646 trunc_dim,
647 query_dim,
648 survivor_docs.len(),
649 approximate_docs,
650 survivor_vectors,
651 n,
652 );
653 } else {
654 let mut scores_buf = vec![0.0f32; max_batch];
655 scores.reserve(n);
656 for chunk in resolved.chunks(batch_len) {
657 #[cfg(feature = "native")]
658 lazy_flat.prefetch_vectors(
659 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
660 );
661 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
662 crate::Error::Query("dense reranker buffer size overflow".into())
663 })?;
664 let raw = &mut raw_buf[..raw_len];
665 for (buf_idx, &(_, flat_idx, _)) in chunk.iter().enumerate() {
666 lazy_flat
667 .read_vector_raw_into(
668 flat_idx,
669 &mut raw[buf_idx * vbs..(buf_idx + 1) * vbs],
670 )
671 .await
672 .map_err(crate::error::Error::Io)?;
673 }
674 searcher.install_search_cpu(|| {
675 score_batch_precomp(
676 &pq,
677 raw,
678 quant,
679 query_dim,
680 &mut scores_buf[..chunk.len()],
681 config.unit_norm,
682 )
683 })?;
684 for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
685 scores.push((ci, ordinal, scores_buf[buf_idx]));
686 }
687 }
688 }
689
690 Ok((scores, vectors, seg_skipped))
691 }
692 },
693 ))
694 .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
695 futures::pin_mut!(segment_futs);
696
697 let mut all_scores: Vec<(usize, u32, f32)> = Vec::new();
698 let mut total_vectors = 0usize;
699 while let Some((scores, vectors, seg_skipped)) = segment_futs.try_next().await? {
700 all_scores.extend(scores);
701 total_vectors = total_vectors.saturating_add(vectors);
702 skipped = skipped.saturating_add(seg_skipped);
703 }
704
705 let read_score_elapsed = t0.elapsed();
706
707 if total_vectors == 0 {
708 log::debug!(
709 "[dense_vector_rerank] field {}: {} candidates, all skipped (no flat vectors)",
710 field_id,
711 candidates.len()
712 );
713 return Ok(Vec::new());
714 }
715
716 all_scores.sort_unstable_by_key(|&(ci, _, _)| ci);
719
720 let mut scored: Vec<SearchResult> = Vec::with_capacity(
721 candidates
722 .len()
723 .min(crate::query::max_candidate_limit(final_limit)),
724 );
725 let mut ordinal_pairs: Vec<(u32, f32)> = Vec::new();
726 let mut i = 0;
727 while i < all_scores.len() {
728 let ci = all_scores[i].0;
729 let run_start = i;
730 while i < all_scores.len() && all_scores[i].0 == ci {
731 i += 1;
732 }
733 let run = &mut all_scores[run_start..i];
734
735 ordinal_pairs.clear();
737 ordinal_pairs.extend(run.iter().map(|&(_, ord, s)| (ord, s)));
738 let combined = config.combiner.combine(&ordinal_pairs);
739
740 run.sort_unstable_by(|a, b| b.2.total_cmp(&a.2));
742 let positions: Vec<ScoredPosition> = run
743 .iter()
744 .map(|&(_, ord, score)| ScoredPosition::new(ord, score))
745 .collect();
746
747 scored.push(SearchResult {
748 doc_id: candidates[ci].doc_id,
749 score: combined,
750 segment_id: candidates[ci].segment_id,
751 positions: vec![(field_id, positions)],
752 });
753 }
754
755 scored.sort_unstable_by(compare_search_results_desc);
756
757 if config.rrf_k > 0.0 {
758 apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
759 } else {
760 scored.truncate(final_limit);
761 }
762
763 log::debug!(
764 "[dense_vector_rerank] field {}: {} candidates -> {} results (skipped {}, {} vectors, unit_norm={}, rrf_k={}): read+score={:.1}ms total={:.1}ms",
765 field_id,
766 candidates.len(),
767 scored.len(),
768 skipped,
769 total_vectors,
770 config.unit_norm,
771 config.rrf_k,
772 read_score_elapsed.as_secs_f64() * 1000.0,
773 t0.elapsed().as_secs_f64() * 1000.0,
774 );
775
776 Ok(scored)
777}
778
779async fn rerank_binary<D: crate::directories::Directory + 'static>(
781 searcher: &crate::index::Searcher<D>,
782 candidates: &[SearchResult],
783 config: &RerankerConfig,
784 final_limit: usize,
785) -> crate::error::Result<Vec<SearchResult>> {
786 if config.binary_vector.is_empty() || candidates.is_empty() {
787 return Ok(Vec::new());
788 }
789
790 let t0 = std::time::Instant::now();
791 let field_id = config.field.0;
792 let query = &config.binary_vector;
793 let byte_len = query.len();
794 let segments = searcher.segment_readers();
795 let seg_by_id = searcher.segment_map();
796
797 let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
799 for (ci, cand) in candidates.iter().enumerate() {
800 if let Some(&seg_idx) = seg_by_id.get(&cand.segment_id) {
801 let reader = &segments[seg_idx];
802 if reader.flat_vectors().contains_key(&field_id) {
803 segment_groups.entry(seg_idx).or_default().push(ci);
804 }
805 }
806 }
807
808 let vector_budget = Arc::new(AtomicUsize::new(0));
810 let byte_budget = Arc::new(AtomicUsize::new(0));
811 let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
812 |(seg_idx, cand_indices)| {
813 #[allow(clippy::redundant_locals)]
814 let segments = &segments;
815 #[allow(clippy::redundant_locals)]
816 let candidates = candidates;
817 let vector_budget = Arc::clone(&vector_budget);
818 let byte_budget = Arc::clone(&byte_budget);
819 async move {
820 let mut scores: Vec<(usize, u32, f32)> = Vec::new();
821
822 let Some(lazy_flat) = segments[seg_idx].flat_vectors().get(&field_id) else {
823 return Ok::<_, crate::error::Error>(scores);
824 };
825 if lazy_flat.quantization != crate::dsl::DenseVectorQuantization::Binary
826 || !lazy_flat.dim.is_multiple_of(8)
827 {
828 return Err(crate::Error::Corruption(format!(
829 "binary reranker field {field_id} has invalid flat-vector metadata"
830 )));
831 }
832 let vbs = lazy_flat.vector_byte_size();
833 if vbs != byte_len {
834 return Err(crate::Error::Corruption(format!(
835 "binary reranker field {field_id} stores {vbs} bytes/vector, expected {byte_len}"
836 )));
837 }
838
839 let mut resolved: Vec<(usize, usize)> = Vec::new();
841 for &ci in &cand_indices {
842 let doc_id = candidates[ci].doc_id;
843 let (start, count) = lazy_flat.flat_indexes_for_doc_range(doc_id);
844 reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
845 for j in 0..count {
846 resolved.push((ci, start + j));
847 }
848 }
849 if resolved.is_empty() {
850 return Ok(scores);
851 }
852
853 resolved.sort_unstable_by_key(|&(_, flat_idx)| flat_idx);
854
855 let n = resolved.len();
856 let batch_len = rerank_batch_len(vbs);
857 let max_batch = batch_len.min(n);
858 let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
859 crate::Error::Query("binary reranker buffer size overflow".into())
860 })?;
861 let mut raw_buf = vec![0u8; max_raw_len];
862 let mut scores_buf = vec![0f32; max_batch];
863 scores.reserve(n);
864
865 for chunk in resolved.chunks(batch_len) {
866 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
867 crate::Error::Query("binary reranker buffer size overflow".into())
868 })?;
869 let raw = &mut raw_buf[..raw_len];
870 for (buf_idx, &(_, flat_idx)) in chunk.iter().enumerate() {
871 lazy_flat
872 .read_vector_raw_into(
873 flat_idx,
874 &mut raw[buf_idx * vbs..(buf_idx + 1) * vbs],
875 )
876 .await
877 .map_err(crate::error::Error::Io)?;
878 }
879 searcher.install_search_cpu(|| {
880 crate::structures::simd::batch_hamming_scores(
881 query,
882 raw,
883 byte_len,
884 lazy_flat.dim,
885 &mut scores_buf[..chunk.len()],
886 );
887 });
888
889 for (buf_idx, &(ci, flat_idx)) in chunk.iter().enumerate() {
890 let (_, ordinal) = lazy_flat.get_doc_id(flat_idx);
891 scores.push((ci, ordinal as u32, scores_buf[buf_idx]));
892 }
893 }
894
895 Ok(scores)
896 }
897 },
898 ))
899 .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
900 futures::pin_mut!(segment_futs);
901
902 let mut cand_ordinal_scores: FxHashMap<usize, Vec<(u32, f32)>> = FxHashMap::default();
904 while let Some(scores) = segment_futs.try_next().await? {
905 for (ci, ordinal, score) in scores {
906 cand_ordinal_scores
907 .entry(ci)
908 .or_default()
909 .push((ordinal, score));
910 }
911 }
912
913 let total_vectors = cand_ordinal_scores.len();
914 let mut scored: Vec<SearchResult> = Vec::with_capacity(total_vectors);
915 for (ci, ordinal_scores) in cand_ordinal_scores {
916 let combined = config.combiner.combine(&ordinal_scores);
917 let positions: Vec<ScoredPosition> = ordinal_scores
918 .iter()
919 .map(|&(ord, s)| ScoredPosition::new(ord, s))
920 .collect();
921 scored.push(SearchResult {
922 doc_id: candidates[ci].doc_id,
923 score: combined,
924 segment_id: candidates[ci].segment_id,
925 positions: vec![(field_id, positions)],
926 });
927 }
928
929 scored.sort_unstable_by(compare_search_results_desc);
930
931 if config.rrf_k > 0.0 {
932 apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
933 } else {
934 scored.truncate(final_limit);
935 }
936
937 log::debug!(
938 "[dense_vector_binary_rerank] field {}: {} candidates -> {} results ({} docs scored, bytes_per_vector={}, rrf_k={}): {:.1}ms",
939 field_id,
940 candidates.len(),
941 scored.len(),
942 total_vectors,
943 byte_len,
944 config.rrf_k,
945 t0.elapsed().as_secs_f64() * 1000.0,
946 );
947
948 Ok(scored)
949}
950
951#[cfg(test)]
952mod tests {
953 use super::*;
954 use crate::dsl::{Document, Field};
955
956 fn make_config(vector: Vec<f32>, combiner: MultiValueCombiner) -> RerankerConfig {
957 RerankerConfig {
958 field: Field(0),
959 vector,
960 binary_vector: Vec::new(),
961 combiner,
962 unit_norm: false,
963 matryoshka_dims: None,
964 rrf_k: 0.0,
965 }
966 }
967
968 #[test]
969 fn rerank_batches_are_bounded_by_bytes() {
970 assert_eq!(rerank_batch_len(1), RERANK_SCORE_BATCH);
971 assert_eq!(
972 rerank_batch_len(MAX_RERANK_RAW_BATCH_BYTES),
973 1,
974 "one very wide vector must still make progress"
975 );
976 assert!(
977 rerank_batch_len(4_096) * 4_096 <= MAX_RERANK_RAW_BATCH_BYTES,
978 "normal batches must stay within the raw scratch budget"
979 );
980 }
981
982 #[test]
983 fn rerank_budget_bounds_count_and_bytes() {
984 let vectors = AtomicUsize::new(0);
985 let bytes = AtomicUsize::new(0);
986 reserve_rerank_vectors(&vectors, &bytes, 2, 32).unwrap();
987 assert_eq!(vectors.load(AtomicOrdering::Relaxed), 2);
988 assert_eq!(bytes.load(AtomicOrdering::Relaxed), 64);
989
990 let vectors = AtomicUsize::new(0);
991 let bytes = AtomicUsize::new(0);
992 assert!(reserve_rerank_vectors(&vectors, &bytes, 2, MAX_L2_RERANK_VECTOR_BYTES).is_err());
993 }
994
995 #[test]
996 fn test_score_document_single_value() {
997 let mut doc = Document::new();
998 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
999
1000 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1001 let (score, positions) = score_document(&doc, &config).unwrap();
1002 assert!((score - 1.0).abs() < 1e-6);
1004 assert_eq!(positions.len(), 1);
1005 assert_eq!(positions[0].position, 0); }
1007
1008 #[test]
1009 fn test_score_document_orthogonal() {
1010 let mut doc = Document::new();
1011 doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]);
1012
1013 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1014 let (score, _) = score_document(&doc, &config).unwrap();
1015 assert!(score.abs() < 1e-6);
1017 }
1018
1019 #[test]
1020 fn test_score_document_multi_value_max() {
1021 let mut doc = Document::new();
1022 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);
1026 let (score, positions) = score_document(&doc, &config).unwrap();
1027 assert!((score - 1.0).abs() < 1e-6);
1028 assert_eq!(positions.len(), 2);
1030 assert_eq!(positions[0].position, 0); assert!((positions[0].score - 1.0).abs() < 1e-6);
1032 }
1033
1034 #[test]
1035 fn test_score_document_multi_value_avg() {
1036 let mut doc = Document::new();
1037 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);
1041 let (score, _) = score_document(&doc, &config).unwrap();
1042 assert!((score - 0.5).abs() < 1e-6);
1044 }
1045
1046 #[test]
1047 fn test_score_document_missing_field() {
1048 let mut doc = Document::new();
1049 doc.add_dense_vector(Field(1), vec![1.0, 0.0, 0.0]);
1051
1052 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1053 assert!(score_document(&doc, &config).is_none());
1054 }
1055
1056 #[test]
1057 fn test_score_document_wrong_field_type() {
1058 let mut doc = Document::new();
1059 doc.add_text(Field(0), "not a vector");
1060
1061 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1062 assert!(score_document(&doc, &config).is_none());
1063 }
1064
1065 #[test]
1066 fn test_score_document_dimension_mismatch() {
1067 let mut doc = Document::new();
1068 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());
1072 }
1073
1074 #[test]
1075 fn test_score_document_empty_query_vector() {
1076 let mut doc = Document::new();
1077 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
1078
1079 let config = make_config(vec![], MultiValueCombiner::Max);
1080 assert!(score_document(&doc, &config).is_none());
1082 }
1083
1084 fn make_result(doc_id: u32, score: f32, segment_id: u128) -> SearchResult {
1085 SearchResult {
1086 doc_id,
1087 score,
1088 segment_id,
1089 positions: Vec::new(),
1090 }
1091 }
1092
1093 #[test]
1094 fn test_rrf_basic_fusion() {
1095 let candidates = vec![
1097 make_result(1, 10.0, 1), make_result(2, 8.0, 1), make_result(3, 5.0, 1), ];
1101
1102 let mut scored = vec![
1104 make_result(3, 0.9, 1), make_result(2, 0.7, 1), make_result(1, 0.3, 1), ];
1108
1109 let k = 60.0;
1110 apply_rrf(&candidates, &mut scored, k, 10);
1111
1112 assert_eq!(scored.len(), 3);
1125 let spread = scored[0].score - scored[2].score;
1127 assert!(
1128 spread < 0.001,
1129 "All docs have near-equal RRF scores, spread={spread}"
1130 );
1131 }
1132
1133 #[test]
1134 fn test_rrf_clear_winner() {
1135 let candidates = vec![
1137 make_result(1, 10.0, 1), make_result(2, 8.0, 1), make_result(3, 5.0, 1), ];
1141
1142 let mut scored = vec![
1144 make_result(1, 0.95, 1), make_result(3, 0.50, 1), make_result(2, 0.30, 1), ];
1148
1149 let k = 60.0;
1150 apply_rrf(&candidates, &mut scored, k, 10);
1151
1152 assert_eq!(scored[0].doc_id, 1, "Doc 1 (rank 1 in both) should be top");
1156 assert!(scored[0].score > scored[1].score);
1157 }
1158
1159 #[test]
1160 fn test_rrf_truncation() {
1161 let candidates = vec![
1162 make_result(1, 10.0, 1),
1163 make_result(2, 8.0, 1),
1164 make_result(3, 5.0, 1),
1165 make_result(4, 3.0, 1),
1166 make_result(5, 1.0, 1),
1167 ];
1168
1169 let mut scored = vec![
1170 make_result(5, 0.9, 1),
1171 make_result(4, 0.8, 1),
1172 make_result(3, 0.7, 1),
1173 make_result(2, 0.6, 1),
1174 make_result(1, 0.5, 1),
1175 ];
1176
1177 apply_rrf(&candidates, &mut scored, 60.0, 3);
1178 assert_eq!(scored.len(), 3, "Should truncate to final_limit=3");
1179 }
1180
1181 #[test]
1182 fn test_rrf_missing_l1_candidate() {
1183 let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1185
1186 let mut scored = vec![
1187 make_result(3, 0.9, 1), make_result(1, 0.5, 1),
1189 ];
1190
1191 apply_rrf(&candidates, &mut scored, 60.0, 10);
1192
1193 assert_eq!(scored[0].doc_id, 1);
1197 }
1198
1199 #[test]
1200 fn test_rrf_small_k() {
1201 let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1203
1204 let mut scored = vec![
1205 make_result(2, 0.9, 1), make_result(1, 0.5, 1), ];
1208
1209 apply_rrf(&candidates, &mut scored, 1.0, 10);
1210
1211 let diff = (scored[0].score - scored[1].score).abs();
1215 assert!(
1216 diff < 1e-6,
1217 "Symmetric ranks should produce equal RRF scores"
1218 );
1219 }
1220}