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 = 2_000_000;
27const MAX_L2_RERANK_VECTOR_BYTES: usize = 1024 * 1024 * 1024;
28const RERANK_SCORE_BATCH: usize = 4_096;
29const MAX_RERANK_RAW_BATCH_BYTES: usize = 8 * 1024 * 1024;
30const MAX_CONCURRENT_RERANK_SEGMENTS: usize = 8;
31
32#[derive(Clone, Copy, PartialEq, Eq)]
33enum RerankerKind {
34 Dense,
35 Binary,
36}
37
38fn validate_reranker_config<D: crate::directories::Directory + 'static>(
39 searcher: &crate::index::Searcher<D>,
40 config: &RerankerConfig,
41) -> crate::error::Result<RerankerKind> {
42 if !config.rrf_k.is_finite() || config.rrf_k < 0.0 {
43 return Err(crate::Error::Query(format!(
44 "reranker rrf_k must be finite and non-negative, got {}",
45 config.rrf_k
46 )));
47 }
48 config.combiner.validate().map_err(crate::Error::Query)?;
49 if config.vector.is_empty() == config.binary_vector.is_empty() {
50 return Err(crate::Error::Query(
51 "reranker must provide exactly one of vector or binary_vector".to_string(),
52 ));
53 }
54
55 let entry = searcher
56 .schema()
57 .get_field_entry(config.field)
58 .ok_or_else(|| crate::Error::FieldNotFound(config.field.0.to_string()))?;
59 if !config.binary_vector.is_empty() {
60 if entry.field_type != crate::dsl::FieldType::BinaryDenseVector {
61 return Err(crate::Error::InvalidFieldType {
62 expected: "binary_dense_vector".to_string(),
63 got: format!("{:?}", entry.field_type),
64 });
65 }
66 let field_config = entry.binary_dense_vector_config.as_ref().ok_or_else(|| {
67 crate::Error::Schema(format!(
68 "binary dense vector field '{}' has no configuration",
69 entry.name
70 ))
71 })?;
72 if field_config.dim == 0 || !field_config.dim.is_multiple_of(8) {
73 return Err(crate::Error::Schema(format!(
74 "binary dense vector field '{}' has invalid dimension {}",
75 entry.name, field_config.dim
76 )));
77 }
78 if config.binary_vector.len() != field_config.byte_len() {
79 return Err(crate::Error::Query(format!(
80 "reranker binary vector byte length {} does not match field '{}' byte length {}",
81 config.binary_vector.len(),
82 entry.name,
83 field_config.byte_len()
84 )));
85 }
86 if config.matryoshka_dims.is_some() {
87 return Err(crate::Error::Query(
88 "reranker matryoshka_dims is not supported for binary vectors".to_string(),
89 ));
90 }
91 return Ok(RerankerKind::Binary);
92 }
93
94 if entry.field_type != crate::dsl::FieldType::DenseVector {
95 return Err(crate::Error::InvalidFieldType {
96 expected: "dense_vector".to_string(),
97 got: format!("{:?}", entry.field_type),
98 });
99 }
100 let field_config = entry.dense_vector_config.as_ref().ok_or_else(|| {
101 crate::Error::Schema(format!(
102 "dense vector field '{}' has no configuration",
103 entry.name
104 ))
105 })?;
106 if config.vector.len() != field_config.dim {
107 return Err(crate::Error::Query(format!(
108 "reranker vector dimension {} does not match field '{}' dimension {}",
109 config.vector.len(),
110 entry.name,
111 field_config.dim
112 )));
113 }
114 if let Some((index, value)) = config
115 .vector
116 .iter()
117 .enumerate()
118 .find(|(_, value)| !value.is_finite())
119 {
120 return Err(crate::Error::Query(format!(
121 "reranker vector contains non-finite value {value} at index {index}"
122 )));
123 }
124 if config.unit_norm != field_config.unit_norm {
125 return Err(crate::Error::Query(format!(
126 "reranker unit_norm={} does not match field '{}' unit_norm={}",
127 config.unit_norm, entry.name, field_config.unit_norm
128 )));
129 }
130 if let Some(dims) = config.matryoshka_dims
131 && (dims == 0 || dims > field_config.dim)
132 {
133 return Err(crate::Error::Query(format!(
134 "reranker matryoshka_dims must be in 1..={}, got {dims}",
135 field_config.dim
136 )));
137 }
138 Ok(RerankerKind::Dense)
139}
140
141fn reserve_rerank_vectors(
142 vector_budget: &AtomicUsize,
143 byte_budget: &AtomicUsize,
144 count: usize,
145 vector_byte_size: usize,
146) -> crate::error::Result<()> {
147 let bytes = count.checked_mul(vector_byte_size).ok_or_else(|| {
148 crate::Error::Query("reranker stored-vector byte budget overflow".to_string())
149 })?;
150 byte_budget
151 .fetch_update(AtomicOrdering::Relaxed, AtomicOrdering::Relaxed, |used| {
152 used.checked_add(bytes)
153 .filter(|&next| next <= MAX_L2_RERANK_VECTOR_BYTES)
154 })
155 .map_err(|used| {
156 crate::Error::Query(format!(
157 "reranker reads more than {MAX_L2_RERANK_VECTOR_BYTES} stored vector bytes \
158 (already reserved {used}, next document needs {bytes})"
159 ))
160 })?;
161
162 vector_budget
163 .fetch_update(AtomicOrdering::Relaxed, AtomicOrdering::Relaxed, |used| {
164 used.checked_add(count)
165 .filter(|&next| next <= MAX_L2_RERANK_VECTORS)
166 })
167 .map(|_| ())
168 .map_err(|used| {
169 crate::Error::Query(format!(
170 "reranker expands to more than {MAX_L2_RERANK_VECTORS} stored vectors \
171 (already reserved {used}, next document has {count})"
172 ))
173 })
174}
175
176fn plan_flat_read_runs(
182 flat_indexes: impl Iterator<Item = usize>,
183 runs: &mut Vec<(usize, usize, usize)>,
184) {
185 runs.clear();
186 for (buffer_index, flat_index) in flat_indexes.enumerate() {
187 if let Some(run) = runs.last_mut()
188 && run.1.checked_add(run.2) == Some(flat_index)
189 {
190 run.2 += 1;
191 continue;
192 }
193 runs.push((buffer_index, flat_index, 1));
194 }
195}
196
197async fn read_flat_vector_runs(
199 lazy_flat: &crate::segment::LazyFlatVectorData,
200 runs: &[(usize, usize, usize)],
201 raw: &mut [u8],
202) -> crate::error::Result<()> {
203 let vbs = lazy_flat.vector_byte_size();
204 for &(buffer_start, flat_start, count) in runs {
205 let bytes = lazy_flat
206 .read_vectors_batch(flat_start, count)
207 .await
208 .map_err(crate::error::Error::Io)?;
209 let start = buffer_start
210 .checked_mul(vbs)
211 .ok_or_else(|| crate::Error::Query("rerank buffer offset overflow".into()))?;
212 let end = start
213 .checked_add(bytes.len())
214 .ok_or_else(|| crate::Error::Query("rerank buffer range overflow".into()))?;
215 raw.get_mut(start..end)
216 .ok_or_else(|| crate::Error::Corruption("rerank buffer is too short".into()))?
217 .copy_from_slice(bytes.as_slice());
218 }
219 Ok(())
220}
221
222fn report_skipped_candidates<D: crate::directories::Directory + 'static>(
225 searcher: &crate::index::Searcher<D>,
226 kind: &'static str,
227 field_id: u32,
228 skipped: u32,
229 total: usize,
230) {
231 if skipped == 0 {
232 return;
233 }
234 let index_label = searcher.schema().index_label();
235 crate::observe::rerank_candidates_skipped(index_label, kind, u64::from(skipped));
236 log::warn!(
237 "[{kind}_vector_rerank] index={index_label} field {field_id}: {skipped} of {total} \
238 candidates skipped (segment missing or no stored vectors for the field) and dropped \
239 from the reranked result"
240 );
241}
242
243#[inline]
244fn rerank_batch_len(vector_byte_size: usize) -> usize {
245 RERANK_SCORE_BATCH.min((MAX_RERANK_RAW_BATCH_BYTES / vector_byte_size.max(1)).max(1))
246}
247
248struct PrecompQuery<'a> {
250 query: &'a [f32],
251 inv_norm_q: f32,
252 query_f16: &'a [u16],
253}
254
255#[inline]
257#[allow(clippy::too_many_arguments)]
258fn score_batch_precomp(
259 pq: &PrecompQuery<'_>,
260 raw: &[u8],
261 quant: crate::dsl::DenseVectorQuantization,
262 dim: usize,
263 scores: &mut [f32],
264 unit_norm: bool,
265) -> crate::error::Result<()> {
266 let query = pq.query;
267 let inv_norm_q = pq.inv_norm_q;
268 let query_f16 = pq.query_f16;
269 use crate::dsl::DenseVectorQuantization;
270 use crate::structures::simd;
271 let element_size = quant.element_size();
272 let required_bytes = scores
273 .len()
274 .checked_mul(dim)
275 .and_then(|elements| elements.checked_mul(element_size))
276 .ok_or_else(|| {
277 crate::Error::Corruption("dense reranker batch size overflow".to_string())
278 })?;
279 if raw.len() < required_bytes {
280 return Err(crate::Error::Corruption(format!(
281 "dense reranker batch is truncated: need {required_bytes} bytes, got {}",
282 raw.len()
283 )));
284 }
285 if matches!(
286 quant,
287 DenseVectorQuantization::F32 | DenseVectorQuantization::F16
288 ) && required_bytes > 0
289 && !(raw.as_ptr() as usize).is_multiple_of(element_size)
290 {
291 return Err(crate::Error::Corruption(format!(
292 "dense reranker {:?} data is not {}-byte aligned",
293 quant, element_size
294 )));
295 }
296 match (quant, unit_norm) {
297 (DenseVectorQuantization::F32, false) => {
298 let num_floats = scores.len() * dim;
299 let vectors: &[f32] =
303 unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const f32, num_floats) };
304 simd::batch_cosine_scores_precomp(query, vectors, dim, scores, inv_norm_q);
305 }
306 (DenseVectorQuantization::F32, true) => {
307 let num_floats = scores.len() * dim;
308 let vectors: &[f32] =
309 unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const f32, num_floats) };
310 simd::batch_dot_scores_precomp(query, vectors, dim, scores, inv_norm_q);
311 }
312 (DenseVectorQuantization::F16, false) => {
313 simd::batch_cosine_scores_f16_precomp(query_f16, raw, dim, scores, inv_norm_q);
314 }
315 (DenseVectorQuantization::F16, true) => {
316 simd::batch_dot_scores_f16_precomp(query_f16, raw, dim, scores, inv_norm_q);
317 }
318 (DenseVectorQuantization::UInt8, false) => {
319 simd::batch_cosine_scores_u8_precomp(query, raw, dim, scores, inv_norm_q);
320 }
321 (DenseVectorQuantization::UInt8, true) => {
322 simd::batch_dot_scores_u8_precomp(query, raw, dim, scores, inv_norm_q);
323 }
324 (DenseVectorQuantization::Binary, _) => {
325 return Err(crate::Error::InvalidFieldType {
326 expected: "non-binary dense vector".to_string(),
327 got: "binary dense vector".to_string(),
328 });
329 }
330 }
331 Ok(())
332}
333
334#[derive(Debug, Clone)]
336pub struct RerankerConfig {
337 pub field: Field,
339 pub vector: Vec<f32>,
341 pub binary_vector: Vec<u8>,
344 pub combiner: MultiValueCombiner,
346 pub unit_norm: bool,
350 pub matryoshka_dims: Option<usize>,
354 pub rrf_k: f32,
358}
359
360#[cfg(test)]
362use crate::structures::simd::cosine_similarity;
363#[cfg(test)]
364fn score_document(
365 doc: &crate::dsl::Document,
366 config: &RerankerConfig,
367) -> Option<(f32, Vec<ScoredPosition>)> {
368 let query_dim = config.vector.len();
369 let mut values: Vec<(u32, f32)> = doc
370 .get_all(config.field)
371 .filter_map(|fv| fv.as_dense_vector())
372 .enumerate()
373 .filter_map(|(ordinal, vec)| {
374 if vec.len() != query_dim {
375 return None;
376 }
377 let score = cosine_similarity(&config.vector, vec);
378 Some((ordinal as u32, score))
379 })
380 .collect();
381
382 if values.is_empty() {
383 return None;
384 }
385
386 let combined = config.combiner.combine(&values);
387
388 values.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
390 let positions: Vec<ScoredPosition> = values
391 .into_iter()
392 .map(|(ordinal, score)| ScoredPosition::new(ordinal, score))
393 .collect();
394
395 Some((combined, positions))
396}
397
398fn apply_rrf(
408 candidates: &[SearchResult],
409 scored: &mut Vec<SearchResult>,
410 k: f32,
411 final_limit: usize,
412) {
413 let l1_ranks: FxHashMap<(u128, u32), usize> = candidates
415 .iter()
416 .enumerate()
417 .map(|(idx, c)| ((c.segment_id, c.doc_id), idx + 1))
418 .collect();
419
420 for (l2_idx, result) in scored.iter_mut().enumerate() {
422 let l1_rank = l1_ranks
423 .get(&(result.segment_id, result.doc_id))
424 .copied()
425 .unwrap_or(candidates.len() + 1);
426 result.score = super::fusion::rrf_contribution(k, l1_rank)
427 + super::fusion::rrf_contribution(k, l2_idx + 1);
428 }
429
430 scored.sort_unstable_by(compare_search_results_desc);
431 scored.truncate(final_limit);
432}
433
434pub async fn rerank<D: crate::directories::Directory + 'static>(
443 searcher: &crate::index::Searcher<D>,
444 candidates: &[SearchResult],
445 config: &RerankerConfig,
446 final_limit: usize,
447) -> crate::error::Result<Vec<SearchResult>> {
448 let kind = validate_reranker_config(searcher, config)?;
451 if final_limit == 0 || candidates.is_empty() {
452 return Ok(Vec::new());
453 }
454
455 if kind == RerankerKind::Binary {
457 return rerank_binary(searcher, candidates, config, final_limit).await;
458 }
459
460 let t0 = std::time::Instant::now();
461 let field_id = config.field.0;
462 let query = &config.vector;
463 let query_dim = query.len();
464 let segments = searcher.segment_readers();
465 let seg_by_id = searcher.segment_map();
466
467 use crate::structures::simd;
469 let norm_q_sq = simd::dot_product_f32(query, query, query_dim);
470 let inv_norm_q = if norm_q_sq < f32::EPSILON {
471 0.0
472 } else {
473 simd::fast_inv_sqrt(norm_q_sq)
474 };
475 let query_f16: Vec<u16> = query.iter().map(|&v| simd::f32_to_f16(v)).collect();
476 let pq = PrecompQuery {
477 query,
478 inv_norm_q,
479 query_f16: &query_f16,
480 };
481
482 let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
484 let mut skipped = 0u32;
485
486 for (ci, candidate) in candidates.iter().enumerate() {
487 if let Some(&si) = seg_by_id.get(&candidate.segment_id) {
488 segment_groups.entry(si).or_default().push(ci);
489 } else {
490 skipped += 1;
491 }
492 }
493
494 let query_ref = pq.query;
498 let inv_norm_q_val = pq.inv_norm_q;
499 let query_f16_ref = pq.query_f16;
500 let vector_budget = Arc::new(AtomicUsize::new(0));
501 let byte_budget = Arc::new(AtomicUsize::new(0));
502
503 let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
504 |(si, candidate_indices)| {
505 #[allow(clippy::redundant_locals)]
506 let segments = &segments;
507 #[allow(clippy::redundant_locals)]
508 let candidates = candidates;
509 #[allow(clippy::redundant_locals)]
510 let query_ref = query_ref;
511 #[allow(clippy::redundant_locals)]
512 let query_f16_ref = query_f16_ref;
513 #[allow(clippy::redundant_locals)]
514 let config = config;
515 let vector_budget = Arc::clone(&vector_budget);
516 let byte_budget = Arc::clone(&byte_budget);
517 async move {
518 let mut scores: Vec<(usize, u32, f32)> = Vec::new();
519 let mut vectors = 0usize;
520 let mut seg_skipped = 0u32;
521
522 let Some(lazy_flat) = segments[si].flat_vectors().get(&field_id) else {
523 return Ok::<_, crate::error::Error>((
524 scores,
525 vectors,
526 candidate_indices.len() as u32,
527 ));
528 };
529 if lazy_flat.dim != query_dim {
530 return Err(crate::Error::Corruption(format!(
531 "dense reranker field {field_id} stores dimension {}, expected {query_dim}",
532 lazy_flat.dim
533 )));
534 }
535 if lazy_flat.quantization == crate::dsl::DenseVectorQuantization::Binary {
536 return Err(crate::Error::Corruption(format!(
537 "dense reranker field {field_id} unexpectedly uses binary storage"
538 )));
539 }
540
541 let vbs = lazy_flat.vector_byte_size();
542 let quant = lazy_flat.quantization;
543
544 let mut resolved: Vec<(usize, usize, u32)> = Vec::new();
546 for &ci in &candidate_indices {
547 let local_doc_id = candidates[ci].doc_id;
548 let (start, count) = lazy_flat.flat_indexes_for_doc_range(local_doc_id);
549 if count == 0 {
550 seg_skipped += 1;
551 continue;
552 }
553 reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
554 for j in 0..count {
555 let (_, ordinal) = lazy_flat.get_doc_id(start + j);
556 resolved.push((ci, start + j, ordinal as u32));
557 }
558 }
559
560 if resolved.is_empty() {
561 return Ok((scores, vectors, seg_skipped));
562 }
563
564 let n = resolved.len();
565 vectors = n;
566
567 resolved.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
571
572 let batch_len = rerank_batch_len(vbs);
573 let max_batch = batch_len.min(n);
574 let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
575 crate::Error::Query("dense reranker buffer size overflow".into())
576 })?;
577 let mut raw_buf = vec![0u8; max_raw_len];
578 let mut runs: Vec<(usize, usize, usize)> = Vec::new();
579
580 let pq = PrecompQuery {
582 query: query_ref,
583 inv_norm_q: inv_norm_q_val,
584 query_f16: query_f16_ref,
585 };
586
587 if let Some(mdims) = config.matryoshka_dims
589 && mdims < query_dim
590 && n > crate::query::max_candidate_limit(final_limit)
591 {
592 let trunc_dim = mdims;
593 let trunc_pq = PrecompQuery {
594 query: &query_ref[..trunc_dim],
595 inv_norm_q: {
596 let nq = simd::dot_product_f32(
597 &query_ref[..trunc_dim],
598 &query_ref[..trunc_dim],
599 trunc_dim,
600 );
601 if nq < f32::EPSILON {
602 0.0
603 } else {
604 simd::fast_inv_sqrt(nq)
605 }
606 },
607 query_f16: &query_f16_ref[..trunc_dim],
608 };
609 let trunc_vbs = trunc_dim * quant.element_size();
610 let mut scores_buf = vec![0.0f32; n];
611 for (chunk_idx, chunk) in resolved.chunks(batch_len).enumerate() {
612 let raw_len = chunk.len().checked_mul(trunc_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_prefix_raw_into(
622 flat_idx,
623 trunc_vbs,
624 &mut raw[buf_idx * trunc_vbs..(buf_idx + 1) * trunc_vbs],
625 )
626 .await
627 .map_err(crate::error::Error::Io)?;
628 }
629 let score_base = chunk_idx * batch_len;
630 searcher.install_search_cpu(|| {
631 score_batch_precomp(
632 &trunc_pq,
633 raw,
634 quant,
635 trunc_dim,
636 &mut scores_buf[score_base..score_base + chunk.len()],
637 config.unit_norm,
638 )
639 })?;
640 }
641
642 let mut approximate_ordinals: FxHashMap<usize, Vec<(u32, f32)>> =
647 FxHashMap::default();
648 for (resolved_index, &(ci, _, ordinal)) in resolved.iter().enumerate() {
649 approximate_ordinals
650 .entry(ci)
651 .or_default()
652 .push((ordinal, scores_buf[resolved_index]));
653 }
654 let mut ranked: Vec<(usize, f32)> = approximate_ordinals
655 .into_iter()
656 .map(|(ci, ordinals)| (ci, config.combiner.combine(&ordinals)))
657 .collect();
658 searcher.install_search_cpu(|| {
659 ranked.sort_unstable_by(|a, b| {
660 b.1.total_cmp(&a.1)
661 .then_with(|| candidates[a.0].doc_id.cmp(&candidates[b.0].doc_id))
662 });
663 });
664 let approximate_docs = ranked.len();
665 let survivor_doc_limit =
666 crate::query::max_candidate_limit(final_limit).min(approximate_docs);
667 let survivor_docs: FxHashSet<usize> = ranked
668 .into_iter()
669 .take(survivor_doc_limit)
670 .map(|(ci, _)| ci)
671 .collect();
672 let mut survivor_entries: Vec<_> = resolved
675 .iter()
676 .copied()
677 .filter(|(ci, _, _)| survivor_docs.contains(ci))
678 .collect();
679 survivor_entries.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
680 let mut full_scores = vec![0.0f32; max_batch.min(survivor_entries.len())];
681 scores.reserve(survivor_entries.len());
682 for chunk in survivor_entries.chunks(batch_len) {
683 #[cfg(feature = "native")]
684 lazy_flat.prefetch_vectors(
685 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
686 );
687 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
688 crate::Error::Query("dense reranker buffer size overflow".into())
689 })?;
690 let raw = &mut raw_buf[..raw_len];
691 plan_flat_read_runs(
692 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
693 &mut runs,
694 );
695 read_flat_vector_runs(lazy_flat, &runs, raw).await?;
696 searcher.install_search_cpu(|| {
697 score_batch_precomp(
698 &pq,
699 raw,
700 quant,
701 query_dim,
702 &mut full_scores[..chunk.len()],
703 config.unit_norm,
704 )
705 })?;
706 for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
707 scores.push((ci, ordinal, full_scores[buf_idx]));
708 }
709 }
710
711 let survivor_vectors = survivor_entries.len();
712 log::debug!(
713 "[dense_vector_rerank] matryoshka pre-filter: {}/{} dims, {}/{} docs and {}/{} vectors survived",
714 trunc_dim,
715 query_dim,
716 survivor_docs.len(),
717 approximate_docs,
718 survivor_vectors,
719 n,
720 );
721 } else {
722 let mut scores_buf = vec![0.0f32; max_batch];
723 scores.reserve(n);
724 for chunk in resolved.chunks(batch_len) {
725 #[cfg(feature = "native")]
726 lazy_flat.prefetch_vectors(
727 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
728 );
729 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
730 crate::Error::Query("dense reranker buffer size overflow".into())
731 })?;
732 let raw = &mut raw_buf[..raw_len];
733 plan_flat_read_runs(
734 chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
735 &mut runs,
736 );
737 read_flat_vector_runs(lazy_flat, &runs, raw).await?;
738 searcher.install_search_cpu(|| {
739 score_batch_precomp(
740 &pq,
741 raw,
742 quant,
743 query_dim,
744 &mut scores_buf[..chunk.len()],
745 config.unit_norm,
746 )
747 })?;
748 for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
749 scores.push((ci, ordinal, scores_buf[buf_idx]));
750 }
751 }
752 }
753
754 Ok((scores, vectors, seg_skipped))
755 }
756 },
757 ))
758 .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
759 futures::pin_mut!(segment_futs);
760
761 let mut all_scores: Vec<(usize, u32, f32)> = Vec::new();
762 let mut total_vectors = 0usize;
763 while let Some((scores, vectors, seg_skipped)) = segment_futs.try_next().await? {
764 all_scores.extend(scores);
765 total_vectors = total_vectors.saturating_add(vectors);
766 skipped = skipped.saturating_add(seg_skipped);
767 }
768
769 let read_score_elapsed = t0.elapsed();
770 report_skipped_candidates(searcher, "dense", field_id, skipped, candidates.len());
771
772 if total_vectors == 0 {
773 log::debug!(
774 "[dense_vector_rerank] field {}: {} candidates, all skipped (no flat vectors)",
775 field_id,
776 candidates.len()
777 );
778 return Ok(Vec::new());
779 }
780
781 all_scores.sort_unstable_by_key(|&(ci, _, _)| ci);
784
785 let mut scored: Vec<SearchResult> = Vec::with_capacity(
786 candidates
787 .len()
788 .min(crate::query::max_candidate_limit(final_limit)),
789 );
790 let mut ordinal_pairs: Vec<(u32, f32)> = Vec::new();
791 let mut i = 0;
792 while i < all_scores.len() {
793 let ci = all_scores[i].0;
794 let run_start = i;
795 while i < all_scores.len() && all_scores[i].0 == ci {
796 i += 1;
797 }
798 let run = &mut all_scores[run_start..i];
799
800 ordinal_pairs.clear();
802 ordinal_pairs.extend(run.iter().map(|&(_, ord, s)| (ord, s)));
803 let combined = config.combiner.combine(&ordinal_pairs);
804
805 run.sort_unstable_by(|a, b| b.2.total_cmp(&a.2));
807 let positions: Vec<ScoredPosition> = run
808 .iter()
809 .map(|&(_, ord, score)| ScoredPosition::new(ord, score))
810 .collect();
811
812 scored.push(SearchResult {
813 doc_id: candidates[ci].doc_id,
814 score: combined,
815 segment_id: candidates[ci].segment_id,
816 positions: vec![(field_id, positions)],
817 });
818 }
819
820 scored.sort_unstable_by(compare_search_results_desc);
821
822 if config.rrf_k > 0.0 {
823 apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
824 } else {
825 scored.truncate(final_limit);
826 }
827
828 log::debug!(
829 "[dense_vector_rerank] field {}: {} candidates -> {} results (skipped {}, {} vectors, unit_norm={}, rrf_k={}): read+score={:.1}ms total={:.1}ms",
830 field_id,
831 candidates.len(),
832 scored.len(),
833 skipped,
834 total_vectors,
835 config.unit_norm,
836 config.rrf_k,
837 read_score_elapsed.as_secs_f64() * 1000.0,
838 t0.elapsed().as_secs_f64() * 1000.0,
839 );
840
841 Ok(scored)
842}
843
844async fn rerank_binary<D: crate::directories::Directory + 'static>(
846 searcher: &crate::index::Searcher<D>,
847 candidates: &[SearchResult],
848 config: &RerankerConfig,
849 final_limit: usize,
850) -> crate::error::Result<Vec<SearchResult>> {
851 if config.binary_vector.is_empty() || candidates.is_empty() {
852 return Ok(Vec::new());
853 }
854
855 let t0 = std::time::Instant::now();
856 let field_id = config.field.0;
857 let query = &config.binary_vector;
858 let byte_len = query.len();
859 let segments = searcher.segment_readers();
860 let seg_by_id = searcher.segment_map();
861
862 let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
864 let mut skipped = 0u32;
865 for (ci, cand) in candidates.iter().enumerate() {
866 if let Some(&seg_idx) = seg_by_id.get(&cand.segment_id) {
867 let reader = &segments[seg_idx];
868 if reader.flat_vectors().contains_key(&field_id) {
869 segment_groups.entry(seg_idx).or_default().push(ci);
870 continue;
871 }
872 }
873 skipped += 1;
874 }
875
876 let vector_budget = Arc::new(AtomicUsize::new(0));
878 let byte_budget = Arc::new(AtomicUsize::new(0));
879 let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
880 |(seg_idx, cand_indices)| {
881 #[allow(clippy::redundant_locals)]
882 let segments = &segments;
883 #[allow(clippy::redundant_locals)]
884 let candidates = candidates;
885 let vector_budget = Arc::clone(&vector_budget);
886 let byte_budget = Arc::clone(&byte_budget);
887 async move {
888 let mut scores: Vec<(usize, u32, f32)> = Vec::new();
889 let mut seg_skipped = 0u32;
890
891 let Some(lazy_flat) = segments[seg_idx].flat_vectors().get(&field_id) else {
892 return Ok::<_, crate::error::Error>((scores, cand_indices.len() as u32));
893 };
894 if lazy_flat.quantization != crate::dsl::DenseVectorQuantization::Binary
895 || !lazy_flat.dim.is_multiple_of(8)
896 {
897 return Err(crate::Error::Corruption(format!(
898 "binary reranker field {field_id} has invalid flat-vector metadata"
899 )));
900 }
901 let vbs = lazy_flat.vector_byte_size();
902 if vbs != byte_len {
903 return Err(crate::Error::Corruption(format!(
904 "binary reranker field {field_id} stores {vbs} bytes/vector, expected {byte_len}"
905 )));
906 }
907
908 let mut resolved: Vec<(usize, usize)> = Vec::new();
910 for &ci in &cand_indices {
911 let doc_id = candidates[ci].doc_id;
912 let (start, count) = lazy_flat.flat_indexes_for_doc_range(doc_id);
913 if count == 0 {
914 seg_skipped += 1;
915 continue;
916 }
917 reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
918 for j in 0..count {
919 resolved.push((ci, start + j));
920 }
921 }
922 if resolved.is_empty() {
923 return Ok((scores, seg_skipped));
924 }
925
926 resolved.sort_unstable_by_key(|&(_, flat_idx)| flat_idx);
927
928 let n = resolved.len();
929 let batch_len = rerank_batch_len(vbs);
930 let max_batch = batch_len.min(n);
931 let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
932 crate::Error::Query("binary reranker buffer size overflow".into())
933 })?;
934 let mut raw_buf = vec![0u8; max_raw_len];
935 let mut runs: Vec<(usize, usize, usize)> = Vec::new();
936 let mut scores_buf = vec![0f32; max_batch];
937 scores.reserve(n);
938
939 for chunk in resolved.chunks(batch_len) {
940 let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
941 crate::Error::Query("binary reranker buffer size overflow".into())
942 })?;
943 let raw = &mut raw_buf[..raw_len];
944 plan_flat_read_runs(chunk.iter().map(|&(_, flat_idx)| flat_idx), &mut runs);
945 read_flat_vector_runs(lazy_flat, &runs, raw).await?;
946 searcher.install_search_cpu(|| {
947 crate::structures::simd::batch_hamming_scores(
948 query,
949 raw,
950 byte_len,
951 lazy_flat.dim,
952 &mut scores_buf[..chunk.len()],
953 );
954 });
955
956 for (buf_idx, &(ci, flat_idx)) in chunk.iter().enumerate() {
957 let (_, ordinal) = lazy_flat.get_doc_id(flat_idx);
958 scores.push((ci, ordinal as u32, scores_buf[buf_idx]));
959 }
960 }
961
962 Ok((scores, seg_skipped))
963 }
964 },
965 ))
966 .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
967 futures::pin_mut!(segment_futs);
968
969 let mut cand_ordinal_scores: FxHashMap<usize, Vec<(u32, f32)>> = FxHashMap::default();
971 while let Some((scores, seg_skipped)) = segment_futs.try_next().await? {
972 skipped = skipped.saturating_add(seg_skipped);
973 for (ci, ordinal, score) in scores {
974 cand_ordinal_scores
975 .entry(ci)
976 .or_default()
977 .push((ordinal, score));
978 }
979 }
980 report_skipped_candidates(searcher, "binary", field_id, skipped, candidates.len());
981
982 let total_vectors = cand_ordinal_scores.len();
983 let mut scored: Vec<SearchResult> = Vec::with_capacity(total_vectors);
984 for (ci, ordinal_scores) in cand_ordinal_scores {
985 let combined = config.combiner.combine(&ordinal_scores);
986 let positions: Vec<ScoredPosition> = ordinal_scores
987 .iter()
988 .map(|&(ord, s)| ScoredPosition::new(ord, s))
989 .collect();
990 scored.push(SearchResult {
991 doc_id: candidates[ci].doc_id,
992 score: combined,
993 segment_id: candidates[ci].segment_id,
994 positions: vec![(field_id, positions)],
995 });
996 }
997
998 scored.sort_unstable_by(compare_search_results_desc);
999
1000 if config.rrf_k > 0.0 {
1001 apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
1002 } else {
1003 scored.truncate(final_limit);
1004 }
1005
1006 log::debug!(
1007 "[dense_vector_binary_rerank] field {}: {} candidates -> {} results ({} docs scored, bytes_per_vector={}, rrf_k={}): {:.1}ms",
1008 field_id,
1009 candidates.len(),
1010 scored.len(),
1011 total_vectors,
1012 byte_len,
1013 config.rrf_k,
1014 t0.elapsed().as_secs_f64() * 1000.0,
1015 );
1016
1017 Ok(scored)
1018}
1019
1020pub(super) struct CandidateVectorPreparation {
1023 inv_norm_q: f32,
1024 query_f16: Vec<u16>,
1025}
1026
1027pub(super) async fn score_vector_candidates<D: crate::directories::Directory + 'static>(
1030 searcher: &crate::index::Searcher<D>,
1031 flat: &crate::segment::LazyFlatVectorData,
1032 vector: &[f32],
1033 binary_vector: &[u8],
1034 unit_norm: bool,
1035 targets: &[u32],
1036 preparation: &mut Option<CandidateVectorPreparation>,
1037) -> crate::Result<Vec<f32>> {
1038 use crate::structures::simd;
1039 let pq = if binary_vector.is_empty() {
1040 let preparation = preparation.get_or_insert_with(|| {
1041 let norm = simd::dot_product_f32(vector, vector, vector.len());
1042 CandidateVectorPreparation {
1043 inv_norm_q: if norm < f32::EPSILON {
1044 0.0
1045 } else {
1046 simd::fast_inv_sqrt(norm)
1047 },
1048 query_f16: Vec::new(),
1049 }
1050 });
1051 if flat.quantization == crate::dsl::DenseVectorQuantization::F16
1052 && preparation.query_f16.is_empty()
1053 {
1054 preparation
1055 .query_f16
1056 .extend(vector.iter().map(|&v| simd::f32_to_f16(v)));
1057 }
1058 Some(PrecompQuery {
1059 query: vector,
1060 inv_norm_q: preparation.inv_norm_q,
1061 query_f16: &preparation.query_f16,
1062 })
1063 } else {
1064 None
1065 };
1066 let vbs = flat.vector_byte_size();
1067 let batch_len = rerank_batch_len(vbs);
1068 let mut raw = vec![0u8; batch_len.min(targets.len()) * vbs];
1069 let mut runs = Vec::new();
1070 let mut scores = vec![0.0; targets.len()];
1071 for (batch, out) in targets.chunks(batch_len).zip(scores.chunks_mut(batch_len)) {
1072 plan_flat_read_runs(batch.iter().map(|&target| target as usize), &mut runs);
1073 let raw = &mut raw[..batch.len() * vbs];
1074 read_flat_vector_runs(flat, &runs, raw).await?;
1075 if !binary_vector.is_empty() {
1076 searcher.install_search_cpu(|| {
1077 simd::batch_hamming_scores(binary_vector, raw, binary_vector.len(), flat.dim, out)
1078 });
1079 } else {
1080 searcher.install_search_cpu(|| {
1081 score_batch_precomp(
1082 pq.as_ref().expect("dense query"),
1083 raw,
1084 flat.quantization,
1085 flat.dim,
1086 out,
1087 unit_norm,
1088 )
1089 })?;
1090 }
1091 }
1092 Ok(scores)
1093}
1094
1095#[cfg(test)]
1096mod tests {
1097 use super::*;
1098 use crate::dsl::{Document, Field};
1099
1100 #[test]
1104 fn flat_read_runs_coalesce_adjacent_indexes_and_tolerate_duplicates() {
1105 let mut runs = Vec::new();
1106 plan_flat_read_runs([3usize, 4, 5, 9, 10, 20, 20, 21].into_iter(), &mut runs);
1107 assert_eq!(runs, vec![(0, 3, 3), (3, 9, 2), (5, 20, 1), (6, 20, 2)]);
1108
1109 plan_flat_read_runs(std::iter::empty(), &mut runs);
1110 assert!(runs.is_empty());
1111
1112 plan_flat_read_runs([usize::MAX].into_iter(), &mut runs);
1113 assert_eq!(runs, vec![(0, usize::MAX, 1)]);
1114 }
1115
1116 fn make_config(vector: Vec<f32>, combiner: MultiValueCombiner) -> RerankerConfig {
1117 RerankerConfig {
1118 field: Field(0),
1119 vector,
1120 binary_vector: Vec::new(),
1121 combiner,
1122 unit_norm: false,
1123 matryoshka_dims: None,
1124 rrf_k: 0.0,
1125 }
1126 }
1127
1128 #[test]
1129 fn rerank_batches_are_bounded_by_bytes() {
1130 assert_eq!(rerank_batch_len(1), RERANK_SCORE_BATCH);
1131 assert_eq!(
1132 rerank_batch_len(MAX_RERANK_RAW_BATCH_BYTES),
1133 1,
1134 "one very wide vector must still make progress"
1135 );
1136 assert!(
1137 rerank_batch_len(4_096) * 4_096 <= MAX_RERANK_RAW_BATCH_BYTES,
1138 "normal batches must stay within the raw scratch budget"
1139 );
1140 }
1141
1142 #[test]
1143 fn rerank_budget_bounds_count_and_bytes() {
1144 let vectors = AtomicUsize::new(0);
1145 let bytes = AtomicUsize::new(0);
1146 reserve_rerank_vectors(&vectors, &bytes, 2, 32).unwrap();
1147 assert_eq!(vectors.load(AtomicOrdering::Relaxed), 2);
1148 assert_eq!(bytes.load(AtomicOrdering::Relaxed), 64);
1149
1150 let vectors = AtomicUsize::new(0);
1151 let bytes = AtomicUsize::new(0);
1152 assert!(reserve_rerank_vectors(&vectors, &bytes, 2, MAX_L2_RERANK_VECTOR_BYTES).is_err());
1153 }
1154
1155 #[test]
1156 fn test_score_document_single_value() {
1157 let mut doc = Document::new();
1158 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
1159
1160 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1161 let (score, positions) = score_document(&doc, &config).unwrap();
1162 assert!((score - 1.0).abs() < 1e-6);
1164 assert_eq!(positions.len(), 1);
1165 assert_eq!(positions[0].position, 0); }
1167
1168 #[test]
1169 fn test_score_document_orthogonal() {
1170 let mut doc = Document::new();
1171 doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]);
1172
1173 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1174 let (score, _) = score_document(&doc, &config).unwrap();
1175 assert!(score.abs() < 1e-6);
1177 }
1178
1179 #[test]
1180 fn test_score_document_multi_value_max() {
1181 let mut doc = Document::new();
1182 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);
1186 let (score, positions) = score_document(&doc, &config).unwrap();
1187 assert!((score - 1.0).abs() < 1e-6);
1188 assert_eq!(positions.len(), 2);
1190 assert_eq!(positions[0].position, 0); assert!((positions[0].score - 1.0).abs() < 1e-6);
1192 }
1193
1194 #[test]
1195 fn test_score_document_multi_value_avg() {
1196 let mut doc = Document::new();
1197 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);
1201 let (score, _) = score_document(&doc, &config).unwrap();
1202 assert!((score - 0.5).abs() < 1e-6);
1204 }
1205
1206 #[test]
1207 fn test_score_document_missing_field() {
1208 let mut doc = Document::new();
1209 doc.add_dense_vector(Field(1), vec![1.0, 0.0, 0.0]);
1211
1212 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1213 assert!(score_document(&doc, &config).is_none());
1214 }
1215
1216 #[test]
1217 fn test_score_document_wrong_field_type() {
1218 let mut doc = Document::new();
1219 doc.add_text(Field(0), "not a vector");
1220
1221 let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1222 assert!(score_document(&doc, &config).is_none());
1223 }
1224
1225 #[test]
1226 fn test_score_document_dimension_mismatch() {
1227 let mut doc = Document::new();
1228 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());
1232 }
1233
1234 #[test]
1235 fn test_score_document_empty_query_vector() {
1236 let mut doc = Document::new();
1237 doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
1238
1239 let config = make_config(vec![], MultiValueCombiner::Max);
1240 assert!(score_document(&doc, &config).is_none());
1242 }
1243
1244 fn make_result(doc_id: u32, score: f32, segment_id: u128) -> SearchResult {
1245 SearchResult {
1246 doc_id,
1247 score,
1248 segment_id,
1249 positions: Vec::new(),
1250 }
1251 }
1252
1253 #[test]
1254 fn test_rrf_basic_fusion() {
1255 let candidates = vec![
1257 make_result(1, 10.0, 1), make_result(2, 8.0, 1), make_result(3, 5.0, 1), ];
1261
1262 let mut scored = vec![
1264 make_result(3, 0.9, 1), make_result(2, 0.7, 1), make_result(1, 0.3, 1), ];
1268
1269 let k = 60.0;
1270 apply_rrf(&candidates, &mut scored, k, 10);
1271
1272 assert_eq!(scored.len(), 3);
1285 let spread = scored[0].score - scored[2].score;
1287 assert!(
1288 spread < 0.001,
1289 "All docs have near-equal RRF scores, spread={spread}"
1290 );
1291 }
1292
1293 #[test]
1294 fn test_rrf_clear_winner() {
1295 let candidates = vec![
1297 make_result(1, 10.0, 1), make_result(2, 8.0, 1), make_result(3, 5.0, 1), ];
1301
1302 let mut scored = vec![
1304 make_result(1, 0.95, 1), make_result(3, 0.50, 1), make_result(2, 0.30, 1), ];
1308
1309 let k = 60.0;
1310 apply_rrf(&candidates, &mut scored, k, 10);
1311
1312 assert_eq!(scored[0].doc_id, 1, "Doc 1 (rank 1 in both) should be top");
1316 assert!(scored[0].score > scored[1].score);
1317 }
1318
1319 #[test]
1320 fn test_rrf_truncation() {
1321 let candidates = vec![
1322 make_result(1, 10.0, 1),
1323 make_result(2, 8.0, 1),
1324 make_result(3, 5.0, 1),
1325 make_result(4, 3.0, 1),
1326 make_result(5, 1.0, 1),
1327 ];
1328
1329 let mut scored = vec![
1330 make_result(5, 0.9, 1),
1331 make_result(4, 0.8, 1),
1332 make_result(3, 0.7, 1),
1333 make_result(2, 0.6, 1),
1334 make_result(1, 0.5, 1),
1335 ];
1336
1337 apply_rrf(&candidates, &mut scored, 60.0, 3);
1338 assert_eq!(scored.len(), 3, "Should truncate to final_limit=3");
1339 }
1340
1341 #[test]
1342 fn test_rrf_missing_l1_candidate() {
1343 let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1345
1346 let mut scored = vec![
1347 make_result(3, 0.9, 1), make_result(1, 0.5, 1),
1349 ];
1350
1351 apply_rrf(&candidates, &mut scored, 60.0, 10);
1352
1353 assert_eq!(scored[0].doc_id, 1);
1357 }
1358
1359 #[test]
1360 fn test_rrf_small_k() {
1361 let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1363
1364 let mut scored = vec![
1365 make_result(2, 0.9, 1), make_result(1, 0.5, 1), ];
1368
1369 apply_rrf(&candidates, &mut scored, 1.0, 10);
1370
1371 let diff = (scored[0].score - scored[1].score).abs();
1375 assert!(
1376 diff < 1e-6,
1377 "Symmetric ranks should produce equal RRF scores"
1378 );
1379 }
1380}