diff --git a/lib/collection/src/collection_manager/segments_searcher.rs b/lib/collection/src/collection_manager/segments_searcher.rs index f99a2df5b5..e285d872be 100644 --- a/lib/collection/src/collection_manager/segments_searcher.rs +++ b/lib/collection/src/collection_manager/segments_searcher.rs @@ -15,16 +15,16 @@ use segment::data_types::query_context::{FormulaContext, QueryContext, SegmentQu use segment::data_types::segment_record::SegmentRecordRaw; use segment::data_types::vectors::QueryVector; use segment::types::{ - Filter, Indexes, PointIdType, ScoredPoint, SearchParams, SegmentConfig, VectorName, - WithPayload, WithPayloadInterface, WithVector, + Filter, Indexes, PointIdType, ScoredPoint, SegmentConfig, VectorName, WithPayload, WithVector, }; use shard::common::stopping_guard::StoppingGuard; use shard::optimizers::config::DEFAULT_INDEXING_THRESHOLD_KB; use shard::query::query_context::{fill_query_context, init_query_context}; -use shard::query::query_enum::QueryEnum; use shard::retrieve::record_internal::RecordInternal; use shard::retrieve::retrieve_blocking::{retrieve_blocking, retrieve_raw_blocking}; -use shard::search::CoreSearchRequestBatch; +use shard::search::{ + BatchSearchParams, CoreSearchRequestBatch, SearchBatchGroup, group_search_batches, +}; use shard::search_result_aggregator::BatchResultAggregator; use shard::segment_holder::locked::LockedSegmentHolder; use tokio_util::task::AbortOnDropHandle; @@ -570,41 +570,6 @@ impl SegmentsSearcher { } } -#[derive(PartialEq, Default, Debug)] -pub enum SearchType { - #[default] - Nearest, - RecommendBestScore, - RecommendSumScores, - Discover, - Context, - FeedbackNaive, -} - -impl From<&QueryEnum> for SearchType { - fn from(query: &QueryEnum) -> Self { - match query { - QueryEnum::Nearest(_) => Self::Nearest, - QueryEnum::RecommendBestScore(_) => Self::RecommendBestScore, - QueryEnum::RecommendSumScores(_) => Self::RecommendSumScores, - QueryEnum::Discover(_) => Self::Discover, - QueryEnum::Context(_) => Self::Context, - QueryEnum::FeedbackNaive(_) => Self::FeedbackNaive, - } - } -} - -#[derive(PartialEq, Default, Debug)] -struct BatchSearchParams<'a> { - pub search_type: SearchType, - pub vector_name: &'a VectorName, - pub filter: Option<&'a Filter>, - pub with_payload: WithPayload, - pub with_vector: WithVector, - pub top: usize, - pub params: Option<&'a SearchParams>, -} - /// Returns suggested search sampling size for a given number of points and required limit. fn sampling_limit( limit: usize, @@ -667,58 +632,17 @@ fn search_in_segment( let mut result: Vec> = Vec::with_capacity(batch_size); let mut further_results: Vec = Vec::with_capacity(batch_size); // if segment have more points to return - let mut vectors_batch: Vec = vec![]; - let mut prev_params = BatchSearchParams::default(); - for search_query in &request.searches { - let with_payload_interface = search_query - .with_payload - .as_ref() - .unwrap_or(&WithPayloadInterface::Bool(false)); + for group in group_search_batches(&request.searches) { + let SearchBatchGroup { + params, + query_vectors, + } = group; - let params = BatchSearchParams { - search_type: search_query.query.as_ref().into(), - vector_name: search_query.query.get_vector_name(), - filter: search_query.filter.as_ref(), - with_payload: WithPayload::from(with_payload_interface), - with_vector: search_query.with_vector.clone().unwrap_or_default(), - top: search_query.limit + search_query.offset, - params: search_query.params.as_ref(), - }; - - let query = search_query.query.clone().into(); - - // same params enables batching (cmp expensive on large filters) - if params == prev_params { - vectors_batch.push(query); - } else { - // different params means different batches - // execute what has been batched so far - if !vectors_batch.is_empty() { - let (mut res, mut further) = execute_batch_search( - &segment, - &vectors_batch, - &prev_params, - use_sampling, - segment_query_context, - timeout, - )?; - further_results.append(&mut further); - result.append(&mut res); - vectors_batch.clear() - } - // start new batch for current search query - vectors_batch.push(query); - prev_params = params; - } - } - - // run last batch if any - if !vectors_batch.is_empty() { let (mut res, mut further) = execute_batch_search( &segment, - &vectors_batch, - &prev_params, + &query_vectors, + ¶ms, use_sampling, segment_query_context, timeout, diff --git a/lib/edge/ffi/src/ops/query.rs b/lib/edge/ffi/src/ops/query.rs index 5595f95ef1..b601d8b257 100644 --- a/lib/edge/ffi/src/ops/query.rs +++ b/lib/edge/ffi/src/ops/query.rs @@ -50,6 +50,34 @@ impl EdgeShard { Ok(points.into_iter().map(ScoredPoint::from).collect()) }) } + + /// Executes several queries as one planned batch. + /// + /// Returns one result list per request, in the same order as `requests`. + /// Prefer this over repeated [`EdgeShard::query`] calls when issuing + /// several independent queries against the same shard: the batch is + /// planned as a whole, so its searches share one pass over the segments + /// and queries that differ only in their vector are scored together. + /// + /// # Errors + /// + /// Returns [`EdgeError::ShardClosed`](crate::error::EdgeError) if the + /// shard is unloaded, or + /// [`EdgeError::OperationError`](crate::error::EdgeError) if any request + /// is invalid or a required payload index is missing. + pub fn query_batch(&self, requests: Vec) -> Result>> { + self.with_shard(|shard| { + let requests = requests + .into_iter() + .map(edge::QueryRequest::try_from) + .collect::, _>>()?; + let batches = shard.query_batch(requests)?; + Ok(batches + .into_iter() + .map(|points| points.into_iter().map(ScoredPoint::from).collect()) + .collect()) + }) + } } // ── SearchParams ──────────────────────────────────────────────────────────── diff --git a/lib/edge/ffi/tests/integration.rs b/lib/edge/ffi/tests/integration.rs index 23768c6a12..dc9c411e5d 100644 --- a/lib/edge/ffi/tests/integration.rs +++ b/lib/edge/ffi/tests/integration.rs @@ -4091,6 +4091,79 @@ fn search_matrix_relates_samples() { } } +/// `query_batch` returns one result list per request, in request order, and each list matches +/// what the same request returns on its own. +#[test] +fn query_batch_matches_individual_queries() { + use qdrant_edge_ffi::{QueryRequest, ScoringQuery}; + + let dir = tempfile::tempdir().expect("tempdir failed"); + let path = dir.path().to_string_lossy().into_owned(); + let shard: Arc = EdgeShard::load(path, Some(make_config())).expect("load failed"); + upsert_three(&shard); + + let nearest = |limit: u64, values: Vec| QueryRequest { + prefetches: Vec::new(), + query: Some(ScoringQuery::Vector { + query: Query::Nearest { + vector: NamedVector::Dense { values }, + using: Some("vec".to_string()), + }, + }), + limit, + offset: None, + filter: None, + params: None, + with_vector: None, + with_payload: None, + score_threshold: None, + }; + + let requests = vec![ + nearest(1, vec![0.5, 0.5, 0.5, 0.5]), + // Same params as above, so both are pushed down as one batched segment search. + nearest(1, vec![0.9, 0.1, 0.1, 0.1]), + nearest(2, vec![0.5, 0.5, 0.5, 0.5]), + ]; + + let batches = shard + .query_batch(requests.clone()) + .expect("query_batch failed"); + + assert_eq!(batches.len(), 3); + assert_eq!(batches[0].len(), 1); + assert_eq!(batches[1].len(), 1); + assert_eq!(batches[2].len(), 2); + + let num_ids = |points: &[qdrant_edge_ffi::types::ScoredPoint]| -> Vec { + points + .iter() + .map(|point| match &point.id { + PointId::NumId { value } => *value, + PointId::Uuid { value } => panic!("unexpected UUID PointId: {value:?}"), + }) + .collect() + }; + + for (batch, request) in batches.iter().zip(requests) { + let individual = shard.query(request).expect("query failed"); + assert_eq!(num_ids(batch), num_ids(&individual)); + } +} + +/// An empty batch is valid and yields no result lists. +#[test] +fn query_batch_empty_returns_empty() { + let dir = tempfile::tempdir().expect("tempdir failed"); + let path = dir.path().to_string_lossy().into_owned(); + let shard: Arc = EdgeShard::load(path, Some(make_config())).expect("load failed"); + upsert_three(&shard); + + let batches = shard.query_batch(Vec::new()).expect("query_batch failed"); + + assert!(batches.is_empty()); +} + // ── Payload schema in info() ────────────────────────────────────────────────── /// `info()` must report every payload index: a bare-type index shows its diff --git a/lib/edge/python/qdrant_edge.pyi b/lib/edge/python/qdrant_edge.pyi index 1ad0fa6a9d..258764efb5 100644 --- a/lib/edge/python/qdrant_edge.pyi +++ b/lib/edge/python/qdrant_edge.pyi @@ -129,6 +129,22 @@ class EdgeShard: """ ... + def query_batch(self, queries: List["QueryRequest"]) -> List[List["ScoredPoint"]]: + """ + Execute several queries as one planned batch. + + Cheaper than calling `query` once per request: the batch is planned as a + whole, so its searches share one pass over the segments and queries that + differ only in their vector are scored together. + + Args: + queries: The query requests to run together. + + Returns: + One list of scored points per request, in the same order. + """ + ... + def search(self, search: "SearchRequest") -> List["ScoredPoint"]: """ Execute a search against the shard. diff --git a/lib/edge/python/src/lib.rs b/lib/edge/python/src/lib.rs index 0cd33a61a6..029054dafc 100644 --- a/lib/edge/python/src/lib.rs +++ b/lib/edge/python/src/lib.rs @@ -145,6 +145,16 @@ impl PyEdgeShard { Ok(points) } + /// Execute several queries as one planned batch. + /// + /// Cheaper than one `query` per request: the batch shares a single pass over the segments. + /// Returns one result list per request, in the same order as `queries`. + pub fn query_batch(&self, queries: Vec) -> Result>> { + let requests = queries.into_iter().map(Into::into).collect(); + let batches = self.get_shard()?.query_batch(requests)?; + Ok(batches.into_iter().map(PyScoredPoint::wrap_vec).collect()) + } + pub fn search(&self, search: PySearchRequest) -> Result> { let points = self.get_shard()?.search(search.into())?; let points = PyScoredPoint::wrap_vec(points); diff --git a/lib/edge/src/edge_shard/shard_read.rs b/lib/edge/src/edge_shard/shard_read.rs index e998cffdc2..b8a008df13 100644 --- a/lib/edge/src/edge_shard/shard_read.rs +++ b/lib/edge/src/edge_shard/shard_read.rs @@ -52,6 +52,15 @@ impl EdgeShard { EdgeShardRead::query(self, request) } + /// Execute several [`QueryRequest`]s as one planned batch — see + /// [`EdgeShardRead::query_batch`]. + pub fn query_batch( + &self, + requests: Vec, + ) -> OperationResult>> { + EdgeShardRead::query_batch(self, requests) + } + pub fn scroll( &self, request: ScrollRequest, diff --git a/lib/edge/src/read_view/ops/matrix.rs b/lib/edge/src/read_view/ops/matrix.rs index f87625a550..570a2ddf59 100644 --- a/lib/edge/src/read_view/ops/matrix.rs +++ b/lib/edge/src/read_view/ops/matrix.rs @@ -62,38 +62,52 @@ impl EdgeReadView { sample_ids.iter().copied().collect::>(), ))); - let mut nearests = Vec::with_capacity(sampled.len()); - for point in &sampled { - let vector = point - .vector - .as_ref() - .and_then(|v| v.get(&using)) - .map(|v| v.to_owned()) - .ok_or_else(|| { - OperationError::service_error("sampled point is missing its vector") - })?; - let nearest = ShardQueryRequest { - prefetches: vec![], - query: Some(ScoringQuery::Vector(QueryEnum::Nearest(NamedQuery::new( - vector, - using.clone(), - )))), - filter: Some(id_filter.clone()), - score_threshold: None, - limit: limit_per_sample.saturating_add(1), // +1 to drop the point itself afterwards - offset: 0, - params: None, - with_vector: WithVector::Bool(false), - with_payload: WithPayloadInterface::Bool(false), - }; - let mut scores = self.query(nearest)?; - if let Some(pos) = scores.iter().position(|p| p.id == point.id) { - scores.remove(pos); - } else if scores.len() == limit_per_sample.saturating_add(1) { - scores.pop(); - } - nearests.push(scores); - } + // One query per sampled point, but issued as a single batch: they share filter, limit and + // vector name, so every segment scores the whole sample in one batched search. + let nearest_requests = sampled + .iter() + .map(|point| { + let vector = point + .vector + .as_ref() + .and_then(|v| v.get(&using)) + .map(|v| v.to_owned()) + .ok_or_else(|| { + OperationError::service_error("sampled point is missing its vector") + })?; + + Ok(ShardQueryRequest { + prefetches: vec![], + query: Some(ScoringQuery::Vector(QueryEnum::Nearest(NamedQuery::new( + vector, + using.clone(), + )))), + filter: Some(id_filter.clone()), + score_threshold: None, + limit: limit_per_sample.saturating_add(1), // +1 to drop the point itself afterwards + offset: 0, + params: None, + with_vector: WithVector::Bool(false), + with_payload: WithPayloadInterface::Bool(false), + }) + }) + .collect::>>()?; + + let nearests = self + .query_batch(nearest_requests)? + .into_iter() + .zip(&sampled) + .map(|(mut scores, point)| { + // Every point matches its own query best, so drop it; if it is missing (e.g. it + // was deleted concurrently), drop the extra result we asked for instead. + if let Some(pos) = scores.iter().position(|p| p.id == point.id) { + scores.remove(pos); + } else if scores.len() == limit_per_sample.saturating_add(1) { + scores.pop(); + } + scores + }) + .collect(); Ok(SearchMatrixResponse { sample_ids, diff --git a/lib/edge/src/read_view/ops/query.rs b/lib/edge/src/read_view/ops/query.rs index 43ca58cefc..cfd598f32d 100644 --- a/lib/edge/src/read_view/ops/query.rs +++ b/lib/edge/src/read_view/ops/query.rs @@ -27,7 +27,33 @@ use crate::read_view::{EdgeReadView, ReadSegmentHandle}; impl EdgeReadView { pub(crate) fn query(&self, request: ShardQueryRequest) -> OperationResult> { - let planned_query = PlannedQuery::try_from(vec![request])?; + let [points] = + self.query_batch(vec![request])? + .try_into() + .map_err(|unconverted: Vec<_>| { + OperationError::service_error(format!( + "unexpected query batch size: expected 1, received {}", + unconverted.len(), + )) + })?; + + Ok(points) + } + + /// Execute several queries as one planned batch. + /// + /// Planning the whole batch at once puts every request's leaf searches into a single + /// [`search_batch`](Self::search_batch): the segments are visited once for the batch, and + /// leaves that differ only in their query vector are pushed down to each segment as one + /// multi-vector search. Only the plan resolution on top of those leaves — fusion, rescoring, + /// payload fetching — stays per request. + /// + /// Returns one result list per request, in request order. + pub(crate) fn query_batch( + &self, + requests: Vec, + ) -> OperationResult>> { + let planned_query = PlannedQuery::try_from(requests)?; let PlannedQuery { root_plans, @@ -35,17 +61,14 @@ impl EdgeReadView { scrolls, } = planned_query; - let mut search_results = Vec::new(); - for search in &searches { - search_results.push(self.search(search.clone())?); - } + let mut search_results = self.search_batch(&searches)?; - let mut scroll_results = Vec::new(); + let mut scroll_results = Vec::with_capacity(scrolls.len()); for scroll in &scrolls { scroll_results.push(self.query_scroll(scroll)?); } - let mut scored_points_batch = Vec::new(); + let mut scored_points_batch = Vec::with_capacity(root_plans.len()); for root_plan in root_plans { let scored_points = self.resolve_plan( root_plan, @@ -57,16 +80,7 @@ impl EdgeReadView { scored_points_batch.push(scored_points) } - let [scored_points] = scored_points_batch - .try_into() - .map_err(|unconverted: Vec<_>| { - OperationError::service_error(format!( - "unexpected scored points batch size: expected 1, received {}", - unconverted.len(), - )) - })?; - - Ok(scored_points) + Ok(scored_points_batch) } fn resolve_plan( @@ -451,3 +465,222 @@ fn filter_by_point_ids(points: &[Vec]) -> Filter { point_ids, ))) } + +#[cfg(test)] +mod tests { + use segment::data_types::vectors::{NamedQuery, VectorInternal}; + use segment::types::{Condition, WithPayloadInterface}; + use shard::query::query_enum::QueryEnum; + + use super::*; + use crate::test_helpers::{VECTOR_NAME, point, point_with_group, test_config, upsert}; + use crate::{EdgeShard, PrefetchBuilder, QueryRequest, QueryRequestBuilder}; + + fn nearest_query(value: f32) -> ScoringQuery { + ScoringQuery::Vector(QueryEnum::Nearest(NamedQuery::new( + VectorInternal::from(vec![value]), + VECTOR_NAME.to_string(), + ))) + } + + fn nearest(limit: usize) -> QueryRequest { + QueryRequestBuilder::new(limit) + .query(nearest_query(1.0)) + .build() + } + + /// Points 1..=n, dot-product scored against `[1.0]`, so ids rank highest-first. + fn shard_with_points(dir: &tempfile::TempDir, n: u64) -> EdgeShard { + let shard = EdgeShard::new(dir.path(), test_config()).unwrap(); + upsert(&shard, (1..=n).map(point).collect()); + shard + } + + /// The whole point of the batch: it must return exactly what the same requests return one by + /// one, however the requests are grouped when pushed down to the segments. + fn assert_matches_one_by_one(shard: &EdgeShard, requests: Vec) { + let one_by_one: Vec<_> = requests + .iter() + .map(|request| shard.query(request.clone()).unwrap()) + .collect(); + + let batched = shard.query_batch(requests).unwrap(); + + assert_eq!(batched, one_by_one); + } + + #[test] + fn query_batch_returns_one_list_per_request() { + let dir = tempfile::tempdir().unwrap(); + let shard = shard_with_points(&dir, 3); + + let batches = shard + .query_batch(vec![nearest(1), nearest(2), nearest(3)]) + .unwrap(); + + assert_eq!(batches.len(), 3); + assert_eq!(batches[0].len(), 1); + assert_eq!(batches[1].len(), 2); + assert_eq!(batches[2].len(), 3); + // Dot product with [1.0] ranks by vector value, so highest ids first. + assert_eq!(batches[0][0].id, 3.into()); + assert_eq!(batches[1][0].id, 3.into()); + assert_eq!(batches[1][1].id, 2.into()); + } + + #[test] + fn query_batch_empty_returns_empty() { + let dir = tempfile::tempdir().unwrap(); + let shard = EdgeShard::new(dir.path(), test_config()).unwrap(); + + let batches = shard.query_batch(Vec::new()).unwrap(); + + assert!(batches.is_empty()); + } + + /// Requests that agree on everything but the query vector collapse into one segment call; + /// results must still be per request. + #[test] + fn query_batch_groups_identical_params() { + let dir = tempfile::tempdir().unwrap(); + let shard = shard_with_points(&dir, 5); + + let requests: Vec<_> = [1.0, -1.0, 3.0] + .into_iter() + .map(|value| { + QueryRequestBuilder::new(2) + .query(nearest_query(value)) + .build() + }) + .collect(); + + assert_matches_one_by_one(&shard, requests); + } + + /// Requests whose params differ split into several segment calls; the results must still line + /// up with the requests, including for a param that only differs in the middle of the batch. + #[test] + fn query_batch_preserves_order_across_groups() { + let dir = tempfile::tempdir().unwrap(); + let shard = shard_with_points(&dir, 5); + + let only_odd_ids = Filter::new_must(Condition::HasId(HasIdCondition::from( + [1, 3, 5] + .map(Into::into) + .into_iter() + .collect::>(), + ))); + + let requests = vec![ + // Same params as the next one: grouped together. + QueryRequestBuilder::new(3) + .query(nearest_query(1.0)) + .build(), + QueryRequestBuilder::new(3) + .query(nearest_query(2.0)) + .build(), + // Filter, limit and offset are pushed down, so each of these starts a new group. + QueryRequestBuilder::new(3) + .query(nearest_query(1.0)) + .filter(only_odd_ids) + .build(), + QueryRequestBuilder::new(1) + .query(nearest_query(1.0)) + .build(), + QueryRequestBuilder::new(2) + .query(nearest_query(1.0)) + .offset(2) + .build(), + // The threshold is applied to the merged result instead, so this one shares a group + // with the request above only if the pushed-down params match — either way it must + // be cut off for this request alone. + QueryRequestBuilder::new(5) + .query(nearest_query(1.0)) + .score_threshold(3.5) + .build(), + // Back to the params of the first group, but not adjacent to it. + QueryRequestBuilder::new(3) + .query(nearest_query(4.0)) + .build(), + ]; + + assert_matches_one_by_one(&shard, requests); + } + + /// A batch of multi-leaf requests: each request contributes several searches to the same + /// batched pass, and every root plan must pick up its own leaves. + #[test] + fn query_batch_resolves_prefetches_per_request() { + let dir = tempfile::tempdir().unwrap(); + let shard = shard_with_points(&dir, 6); + + // Weighted RRF rather than DBSF: with these points DBSF produces tied fused scores, whose + // relative order is not defined (`score_fusion` sorts the values of an `AHashMap`). + let fusion = |limit: usize, first: f32, second: f32| { + QueryRequestBuilder::new(limit) + .add_prefetch(PrefetchBuilder::new(4).query(nearest_query(first)).build()) + .add_prefetch(PrefetchBuilder::new(4).query(nearest_query(second)).build()) + .query(ScoringQuery::Fusion(FusionInternal::Rrf { + k: 2, + weights: Some(vec![OrderedFloat(1.0), OrderedFloat(0.5)]), + })) + .build() + }; + + let requests = vec![ + fusion(3, 1.0, -1.0), + nearest(3), + fusion(2, 2.0, 5.0), + // A rescore that runs its own search on top of a prefetch. + QueryRequestBuilder::new(2) + .add_prefetch(PrefetchBuilder::new(4).query(nearest_query(1.0)).build()) + .query(nearest_query(-1.0)) + .build(), + ]; + + assert_matches_one_by_one(&shard, requests); + } + + /// Payload and vector fetching happens per root plan, after the shared search pass. + #[test] + fn query_batch_fills_payload_and_vectors_per_request() { + let dir = tempfile::tempdir().unwrap(); + let shard = EdgeShard::new(dir.path(), test_config()).unwrap(); + upsert( + &shard, + vec![point_with_group(1, "a"), point_with_group(2, "b")], + ); + + let requests = vec![ + QueryRequestBuilder::new(2) + .query(nearest_query(1.0)) + .build(), + QueryRequestBuilder::new(2) + .query(nearest_query(1.0)) + .with_payload(WithPayloadInterface::Bool(true)) + .build(), + QueryRequestBuilder::new(2) + .query(nearest_query(1.0)) + .with_vector(WithVector::Bool(true)) + .build(), + ]; + + let batches = shard.query_batch(requests).unwrap(); + + assert!(batches[0].iter().all(|point| point.payload.is_none())); + assert!(batches[1].iter().all(|point| point.payload.is_some())); + assert!(batches[2].iter().all(|point| point.vector.is_some())); + } + + /// An empty shard has no segments to search, so every request gets an empty list — not a + /// short batch. + #[test] + fn query_batch_without_segments_returns_a_list_per_request() { + let dir = tempfile::tempdir().unwrap(); + let shard = EdgeShard::new(dir.path(), test_config()).unwrap(); + + let batches = shard.query_batch(vec![nearest(1), nearest(2)]).unwrap(); + + assert_eq!(batches, vec![vec![], vec![]]); + } +} diff --git a/lib/edge/src/read_view/ops/search.rs b/lib/edge/src/read_view/ops/search.rs index a40875f7fc..d734a9d261 100644 --- a/lib/edge/src/read_view/ops/search.rs +++ b/lib/edge/src/read_view/ops/search.rs @@ -3,15 +3,15 @@ use std::sync::atomic::AtomicBool; use common::counter::hardware_accumulator::HwMeasurementAcc; use common::iterator_ext::IteratorExt; -use segment::common::operation_error::OperationResult; +use segment::common::operation_error::{OperationError, OperationResult}; use segment::data_types::modifier::Modifier; use segment::data_types::query_context::QueryContext; -use segment::data_types::vectors::QueryVector; use segment::entry::ReadSegmentEntry; -use segment::types::{DEFAULT_FULL_SCAN_THRESHOLD, ScoredPoint, WithPayload}; +use segment::types::{DEFAULT_FULL_SCAN_THRESHOLD, Distance, ScoredPoint}; use shard::common::stopping_guard::StoppingGuard; use shard::query::query_context::init_query_context; -use shard::search::CoreSearchRequest; +use shard::query::query_enum::QueryEnum; +use shard::search::{CoreSearchRequest, group_search_batches}; use shard::search_result_aggregator::BatchResultAggregator; use crate::read_view::{EdgeReadView, ReadSegmentHandle}; @@ -19,10 +19,39 @@ use crate::read_view::{EdgeReadView, ReadSegmentHandle}; impl EdgeReadView { /// This method is DEPRECATED and should be replaced with query. pub fn search(&self, search: CoreSearchRequest) -> OperationResult> { + let [points] = + self.search_batch(&[search])? + .try_into() + .map_err(|unconverted: Vec<_>| { + OperationError::service_error(format!( + "unexpected search batch size: expected 1, received {}", + unconverted.len(), + )) + })?; + + Ok(points) + } + + /// Run a whole batch of core searches in a single pass over the segments. + /// + /// Requests that agree on everything but their query vector are handed to each segment as one + /// [`search_batch`](ReadSegmentEntry::search_batch) call, so the work that does not depend on + /// the query vector — evaluating the filter into a candidate set, picking the index to use — is + /// paid once per group instead of once per request. The segments are also visited (and the + /// query context built) once for the whole batch rather than once per request. + /// + /// Returns one result list per request, in request order. + pub(crate) fn search_batch( + &self, + searches: &[CoreSearchRequest], + ) -> OperationResult>> { + if searches.is_empty() { + return Ok(Vec::new()); + } + let is_stopped_guard = StoppingGuard::new(); - let searches = [search]; let query_context = init_query_context( - &searches, + searches, DEFAULT_FULL_SCAN_THRESHOLD, &is_stopped_guard, HwMeasurementAcc::disposable_edge(), @@ -33,7 +62,13 @@ impl EdgeReadView { .is_some_and(|v| v.modifier == Some(Modifier::Idf)) }, )?; - let [search] = searches; + + // Resolved up front so an unknown vector name fails the batch before any search runs. + let distances = searches + .iter() + .map(|search| self.config.get_distance(search.query.get_vector_name())) + .collect::>>()?; + let Some(context) = fill_query_context_over( query_context, &self.segments, @@ -41,108 +76,102 @@ impl EdgeReadView { )? else { // No segments to search - return Ok(vec![]); + return Ok(vec![Vec::new(); searches.len()]); }; - let CoreSearchRequest { - query, - filter, - params, - limit, - offset, - with_payload, - with_vector, - score_threshold, - } = search; - - let vector_name = query.get_vector_name().to_string(); - let query_vector = QueryVector::from(query); - let with_payload = WithPayload::from(with_payload.unwrap_or_default()); - let with_vector = with_vector.unwrap_or_default(); + // Grouping depends on the requests only, so it is computed once and shared by all segments. + let groups = group_search_batches(searches); // Search every segment in parallel on the shard's search pool. Each task derives its own // per-segment query context from the shared `context`. let points_by_segment = self.par_map_segments(|segment| { - let batched_points = segment.read_segment().search_batch( - &vector_name, - &[&query_vector], - &with_payload, - &with_vector, - filter.as_ref(), - offset + limit, - params.as_ref(), - &context.get_segment_query_context(), - )?; + let segment_query_context = context.get_segment_query_context(); + let segment = segment.read_segment(); - debug_assert_eq!(batched_points.len(), 1); + let mut points_by_request = Vec::with_capacity(searches.len()); + for group in &groups { + let query_vectors: Vec<_> = group.query_vectors.iter().collect(); + let batched_points = segment.search_batch( + group.params.vector_name, + &query_vectors, + &group.params.with_payload, + &group.params.with_vector, + group.params.filter, + group.params.top, + group.params.params, + &segment_query_context, + )?; - let [points] = batched_points - .try_into() - .expect("single batched search result"); + debug_assert_eq!(batched_points.len(), group.query_vectors.len()); + points_by_request.extend(batched_points); + } - Ok(points) + Ok(points_by_request) })?; - let mut aggregator = BatchResultAggregator::new([offset + limit]); - aggregator.update_point_versions(points_by_segment.iter().flatten()); + let mut aggregator = + BatchResultAggregator::new(searches.iter().map(|search| search.offset + search.limit)); + aggregator.update_point_versions(points_by_segment.iter().flatten().flatten()); - for points in points_by_segment { - aggregator.update_batch_results(0, points); - } - - let [mut points] = aggregator - .into_topk() - .try_into() - .expect("single batched search result"); - - let distance = { - if let Some(dense) = self.config.vectors.get(&vector_name) { - dense.distance - } else if self.config.sparse_vectors.contains_key(&vector_name) { - segment::types::Distance::Dot - } else { - return Err( - segment::common::operation_error::OperationError::service_error(format!( - "vector config for '{vector_name}' does not exist" - )), - ); - } - }; - - match &query_vector { - QueryVector::Nearest(_) => { - for point in &mut points { - point.score = distance.postprocess_score(point.score); - } - } - QueryVector::RecommendBestScore(_) => (), - QueryVector::RecommendSumScores(_) => (), - QueryVector::Discover(_) => (), - QueryVector::Context(_) => (), - QueryVector::FeedbackNaive(_) => (), - } - - if let Some(score_threshold) = score_threshold { - debug_assert!( - points.is_sorted_by(|left, right| distance.is_ordered(left.score, right.score)), - ); - - let below_threshold = points - .iter() - .enumerate() - .find(|(_, point)| !distance.check_threshold(point.score, score_threshold)); - - if let Some((below_threshold_idx, _)) = below_threshold { - points.truncate(below_threshold_idx); + for points_by_request in points_by_segment { + for (request_idx, points) in points_by_request.into_iter().enumerate() { + aggregator.update_batch_results(request_idx, points); } } - let _ = points.drain(..cmp::min(points.len(), offset)); + // One aggregator was created per request, so the top-k lists line up with `searches`. + let mut points_by_request = aggregator.into_topk(); + debug_assert_eq!(points_by_request.len(), searches.len()); - Ok(points) + for ((points, search), distance) in + points_by_request.iter_mut().zip(searches).zip(distances) + { + postprocess_scores(points, search, distance); + } + + Ok(points_by_request) } } +/// Turn the raw segment scores of a single request into the scores the caller expects: apply the +/// distance's score postprocessing, cut off at the score threshold and skip the requested offset. +fn postprocess_scores( + points: &mut Vec, + search: &CoreSearchRequest, + distance: Distance, +) { + match &search.query { + // Only plain nearest-neighbour scores are raw segment distances; every other query + // already produces a comparable score of its own. + QueryEnum::Nearest(_) => { + for point in points.iter_mut() { + point.score = distance.postprocess_score(point.score); + } + } + QueryEnum::RecommendBestScore(_) => (), + QueryEnum::RecommendSumScores(_) => (), + QueryEnum::Discover(_) => (), + QueryEnum::Context(_) => (), + QueryEnum::FeedbackNaive(_) => (), + } + + if let Some(score_threshold) = search.score_threshold { + debug_assert!( + points.is_sorted_by(|left, right| distance.is_ordered(left.score, right.score)), + ); + + let below_threshold = points + .iter() + .position(|point| !distance.check_threshold(point.score, score_threshold)); + + if let Some(below_threshold_idx) = below_threshold { + points.truncate(below_threshold_idx); + } + } + + let _ = points.drain(..cmp::min(points.len(), search.offset)); +} + /// Fill a [`QueryContext`] from a pre-collected snapshot of read handles. /// /// Read-handle equivalent of [`shard::query::query_context::fill_query_context`], which is hard-typed diff --git a/lib/edge/src/read_view/shard_read.rs b/lib/edge/src/read_view/shard_read.rs index 28bdaeba62..9f68d2fb6b 100644 --- a/lib/edge/src/read_view/shard_read.rs +++ b/lib/edge/src/read_view/shard_read.rs @@ -64,6 +64,15 @@ pub trait EdgeShardRead: sealed::Sealed { fn query(&self, request: QueryRequest) -> OperationResult>; + /// Execute several [`QueryRequest`]s as one planned batch. + /// + /// Cheaper than running the same requests one by one: the batch is planned as a whole, so its + /// leaf searches share a single pass over the segments, and leaves that differ only in their + /// query vector are pushed down to each segment as one multi-vector search. + /// + /// Returns one result list per request, in request order. + fn query_batch(&self, requests: Vec) -> OperationResult>>; + fn scroll( &self, request: ScrollRequest, @@ -99,6 +108,10 @@ impl EdgeShardRead for T { view(self).query(request.into()) } + fn query_batch(&self, requests: Vec) -> OperationResult>> { + view(self).query_batch(requests.into_iter().map(Into::into).collect()) + } + fn scroll( &self, request: ScrollRequest, diff --git a/lib/shard/src/search.rs b/lib/shard/src/search.rs index 46156334ca..ed72c7ce86 100644 --- a/lib/shard/src/search.rs +++ b/lib/shard/src/search.rs @@ -4,7 +4,10 @@ use itertools::Itertools as _; use segment::data_types::load_profile::LoadProfile; #[cfg(feature = "api")] use segment::data_types::vectors::NamedQuery; -use segment::types::{Filter, SearchParams, WithPayloadInterface, WithVector}; +use segment::data_types::vectors::QueryVector; +use segment::types::{ + Filter, SearchParams, VectorName, WithPayload, WithPayloadInterface, WithVector, +}; #[cfg(feature = "api")] use segment::{data_types::vectors::VectorInternal, vector_storage::query::ContextPair}; @@ -248,3 +251,201 @@ impl TryFrom for CoreSearchRequest { pub struct CoreSearchRequestBatch { pub searches: Vec, } + +/// Which scoring query a search runs. Only searches of the same type can share one batched +/// segment call, because a segment scores a whole batch with a single query implementation. +#[derive(PartialEq, Debug)] +pub enum SearchType { + Nearest, + RecommendBestScore, + RecommendSumScores, + Discover, + Context, + FeedbackNaive, +} + +impl From<&QueryEnum> for SearchType { + fn from(query: &QueryEnum) -> Self { + match query { + QueryEnum::Nearest(_) => Self::Nearest, + QueryEnum::RecommendBestScore(_) => Self::RecommendBestScore, + QueryEnum::RecommendSumScores(_) => Self::RecommendSumScores, + QueryEnum::Discover(_) => Self::Discover, + QueryEnum::Context(_) => Self::Context, + QueryEnum::FeedbackNaive(_) => Self::FeedbackNaive, + } + } +} + +/// Everything a segment search takes apart from the query vector itself, i.e. exactly the +/// arguments of [`ReadSegmentEntry::search_batch`] that are shared by a batch. +/// +/// [`ReadSegmentEntry::search_batch`]: segment::entry::ReadSegmentEntry::search_batch +#[derive(PartialEq, Debug)] +pub struct BatchSearchParams<'a> { + pub search_type: SearchType, + pub vector_name: &'a VectorName, + pub filter: Option<&'a Filter>, + pub with_payload: WithPayload, + pub with_vector: WithVector, + pub top: usize, + pub params: Option<&'a SearchParams>, +} + +impl<'a> From<&'a CoreSearchRequest> for BatchSearchParams<'a> { + fn from(request: &'a CoreSearchRequest) -> Self { + let CoreSearchRequest { + query, + filter, + params, + limit, + offset, + with_payload, + with_vector, + score_threshold: _, // applied to the merged result, not by the segment + } = request; + + Self { + search_type: SearchType::from(query), + vector_name: query.get_vector_name(), + filter: filter.as_ref(), + with_payload: WithPayload::from( + with_payload + .as_ref() + .unwrap_or(&WithPayloadInterface::Bool(false)), + ), + with_vector: with_vector.clone().unwrap_or_default(), + top: limit + offset, + params: params.as_ref(), + } + } +} + +/// A run of search requests that a segment can serve with one +/// [`search_batch`](segment::entry::ReadSegmentEntry::search_batch) call: they agree on every +/// parameter and differ only in their query vector. +#[derive(Debug)] +pub struct SearchBatchGroup<'a> { + pub params: BatchSearchParams<'a>, + pub query_vectors: Vec, +} + +/// Split a batch of search requests into groups that can each be pushed down to a segment as a +/// single batched search, so per-query work that does not depend on the query vector — resolving +/// the vector index, building the filtered id context — is paid once per group instead of once +/// per request. +/// +/// Only *consecutive* requests are grouped, so the requests keep their input order: concatenating +/// the per-group results in group order yields one result list per input request, in input order. +/// +/// The grouping depends on the requests alone, so a caller searching several segments computes it +/// once and reuses it for every segment. +pub fn group_search_batches(searches: &[CoreSearchRequest]) -> Vec> { + let mut groups: Vec = Vec::with_capacity(searches.len()); + + for search in searches { + let params = BatchSearchParams::from(search); + let query_vector = QueryVector::from(search.query.clone()); + + // Comparing params is expensive on large filters, but far cheaper than re-running a + // segment search that could have shared one. + match groups.last_mut() { + Some(last) if last.params == params => last.query_vectors.push(query_vector), + Some(_) | None => groups.push(SearchBatchGroup { + params, + query_vectors: vec![query_vector], + }), + } + } + + groups +} + +#[cfg(test)] +mod tests { + use ahash::AHashSet; + use segment::data_types::vectors::{NamedQuery, VectorInternal}; + use segment::types::{Condition, HasIdCondition}; + + use super::*; + + fn nearest(vector: Vec, limit: usize) -> CoreSearchRequest { + CoreSearchRequest { + query: QueryEnum::Nearest(NamedQuery::new( + VectorInternal::from(vector), + "vector".to_string(), + )), + filter: None, + params: None, + limit, + offset: 0, + with_payload: None, + with_vector: None, + score_threshold: None, + } + } + + fn group_sizes(searches: &[CoreSearchRequest]) -> Vec { + group_search_batches(searches) + .iter() + .map(|group| group.query_vectors.len()) + .collect() + } + + #[test] + fn requests_differing_only_by_vector_form_one_group() { + let searches = vec![ + nearest(vec![1.0], 3), + nearest(vec![2.0], 3), + nearest(vec![3.0], 3), + ]; + + assert_eq!(group_sizes(&searches), vec![3]); + } + + #[test] + fn differing_params_split_groups() { + let with_filter = |mut search: CoreSearchRequest| { + search.filter = Some(Filter::new_must(Condition::HasId(HasIdCondition::from( + AHashSet::from_iter([1.into()]), + )))); + search + }; + let with_offset = |mut search: CoreSearchRequest| { + // Offsets are served by raising the segment-level top, so they split too. + search.offset = 1; + search + }; + let with_payload = |mut search: CoreSearchRequest| { + search.with_payload = Some(WithPayloadInterface::Bool(true)); + search + }; + + let searches = vec![ + nearest(vec![1.0], 3), + nearest(vec![1.0], 4), + with_filter(nearest(vec![1.0], 4)), + with_offset(nearest(vec![1.0], 4)), + with_payload(nearest(vec![1.0], 4)), + ]; + + assert_eq!(group_sizes(&searches), vec![1, 1, 1, 1, 1]); + } + + /// Only consecutive requests are grouped, so results can be concatenated back in input order. + #[test] + fn identical_params_are_not_grouped_across_a_different_request() { + let searches = vec![ + nearest(vec![1.0], 3), + nearest(vec![2.0], 5), + nearest(vec![3.0], 3), + ]; + + assert_eq!(group_sizes(&searches), vec![1, 1, 1]); + } + + #[test] + fn empty_batch_has_no_groups() { + assert!(group_search_batches(&[]).is_empty()); + } +}