enhance: remove stream reduce logic(#47516) (#47544)

related: #47516

Signed-off-by: MrPresent-Han <chun.han@gmail.com>
Co-authored-by: MrPresent-Han <chun.han@gmail.com>
This commit is contained in:
Chun Han
2026-02-09 18:25:52 +08:00
committed by GitHub
co-authored by MrPresent-Han
parent 88e2ae08b7
commit af11c336fd
11 changed files with 3 additions and 1892 deletions
-299
View File
@@ -36,305 +36,6 @@
#include "test_utils/PbHelper.h"
#include "test_utils/c_api_test_utils.h"
TEST(CApiTest, StreamReduce) {
int N = 300;
int topK = 100;
int num_queries = 2;
auto collection = NewCollection(get_default_schema_config().c_str());
//1. set up segments
CSegmentInterface segment_1;
auto status = NewSegment(collection, Growing, -1, &segment_1, false);
ASSERT_EQ(status.error_code, Success);
CSegmentInterface segment_2;
status = NewSegment(collection, Growing, -1, &segment_2, false);
ASSERT_EQ(status.error_code, Success);
//2. insert data into segments
auto schema = ((milvus::segcore::Collection*)collection)->get_schema();
auto dataset_1 = DataGen(schema, N, 55, 0, 1, 10, true);
int64_t offset_1;
PreInsert(segment_1, N, &offset_1);
auto insert_data_1 = serialize(dataset_1.raw_);
auto ins_res_1 = Insert(segment_1,
offset_1,
N,
dataset_1.row_ids_.data(),
dataset_1.timestamps_.data(),
insert_data_1.data(),
insert_data_1.size());
ASSERT_EQ(ins_res_1.error_code, Success);
auto dataset_2 = DataGen(schema, N, 66, 0, 1, 10, true);
int64_t offset_2;
PreInsert(segment_2, N, &offset_2);
auto insert_data_2 = serialize(dataset_2.raw_);
auto ins_res_2 = Insert(segment_2,
offset_2,
N,
dataset_2.row_ids_.data(),
dataset_2.timestamps_.data(),
insert_data_2.data(),
insert_data_2.size());
ASSERT_EQ(ins_res_2.error_code, Success);
//3. search two segments
milvus::segcore::ScopedSchemaHandle schema_handle(*schema);
auto binary_plan =
schema_handle.ParseSearch("", // expression (empty for no filter)
"fakevec", // vector field name
topK, // topK
"L2", // metric_type
R"({"nprobe": 10})"); // search_params
auto blob = generate_query_data(num_queries);
void* plan = nullptr;
status = CreateSearchPlanByExpr(
collection, binary_plan.data(), binary_plan.size(), &plan);
ASSERT_EQ(status.error_code, Success);
void* placeholderGroup = nullptr;
status = ParsePlaceholderGroup(
plan, blob.data(), blob.length(), &placeholderGroup);
ASSERT_EQ(status.error_code, Success);
std::vector<CPlaceholderGroup> placeholderGroups;
placeholderGroups.push_back(placeholderGroup);
dataset_1.timestamps_.clear();
dataset_1.timestamps_.push_back(1);
dataset_2.timestamps_.clear();
dataset_2.timestamps_.push_back(1);
CSearchResult res1;
CSearchResult res2;
auto stats1 = CSearch(
segment_1, plan, placeholderGroup, dataset_1.timestamps_[N - 1], &res1);
ASSERT_EQ(stats1.error_code, Success);
auto stats2 = CSearch(
segment_2, plan, placeholderGroup, dataset_2.timestamps_[N - 1], &res2);
ASSERT_EQ(stats2.error_code, Success);
//4. stream reduce two search results
auto slice_nqs = std::vector<int64_t>{num_queries / 2, num_queries / 2};
if (num_queries == 1) {
slice_nqs = std::vector<int64_t>{num_queries};
}
auto slice_topKs = std::vector<int64_t>{topK, topK};
if (topK == 1) {
slice_topKs = std::vector<int64_t>{topK, topK};
}
//5. set up stream reducer
CSearchStreamReducer c_search_stream_reducer;
NewStreamReducer(plan,
slice_nqs.data(),
slice_topKs.data(),
slice_nqs.size(),
&c_search_stream_reducer);
StreamReduce(c_search_stream_reducer, &res1, 1);
StreamReduce(c_search_stream_reducer, &res2, 1);
CSearchResultDataBlobs c_search_result_data_blobs;
GetStreamReduceResult(c_search_stream_reducer, &c_search_result_data_blobs);
SearchResultDataBlobs* search_result_data_blob =
(SearchResultDataBlobs*)(c_search_result_data_blobs);
//6. check
for (size_t i = 0; i < slice_nqs.size(); i++) {
milvus::proto::schema::SearchResultData search_result_data;
auto suc = search_result_data.ParseFromArray(
search_result_data_blob->blobs[i].data(),
search_result_data_blob->blobs[i].size());
ASSERT_TRUE(suc);
ASSERT_EQ(search_result_data.num_queries(), slice_nqs[i]);
ASSERT_EQ(search_result_data.top_k(), slice_topKs[i]);
ASSERT_EQ(search_result_data.ids().int_id().data_size(),
search_result_data.topks().at(0) * slice_nqs[i]);
ASSERT_EQ(search_result_data.scores().size(),
search_result_data.topks().at(0) * slice_nqs[i]);
ASSERT_EQ(search_result_data.topks().size(), slice_nqs[i]);
for (auto real_topk : search_result_data.topks()) {
ASSERT_LE(real_topk, slice_topKs[i]);
}
}
DeleteSearchResultDataBlobs(c_search_result_data_blobs);
DeleteSearchPlan(plan);
DeletePlaceholderGroup(placeholderGroup);
DeleteSearchResult(res1);
DeleteSearchResult(res2);
DeleteCollection(collection);
DeleteSegment(segment_1);
DeleteSegment(segment_2);
DeleteStreamSearchReducer(c_search_stream_reducer);
DeleteStreamSearchReducer(nullptr);
}
TEST(CApiTest, StreamReduceGroupBY) {
int N = 300;
int topK = 100;
int num_queries = 2;
int dim = 4;
namespace schema = milvus::proto::schema;
void* c_collection;
//1. set up schema and collection
{
schema::CollectionSchema collection_schema;
auto pk_field_schema = collection_schema.add_fields();
pk_field_schema->set_name("pk_field");
pk_field_schema->set_fieldid(100);
pk_field_schema->set_data_type(schema::DataType::Int64);
pk_field_schema->set_is_primary_key(true);
auto i8_field_schema = collection_schema.add_fields();
i8_field_schema->set_name("int8_field");
i8_field_schema->set_fieldid(101);
i8_field_schema->set_data_type(schema::DataType::Int8);
i8_field_schema->set_is_primary_key(false);
auto i16_field_schema = collection_schema.add_fields();
i16_field_schema->set_name("int16_field");
i16_field_schema->set_fieldid(102);
i16_field_schema->set_data_type(schema::DataType::Int16);
i16_field_schema->set_is_primary_key(false);
auto i32_field_schema = collection_schema.add_fields();
i32_field_schema->set_name("int32_field");
i32_field_schema->set_fieldid(103);
i32_field_schema->set_data_type(schema::DataType::Int32);
i32_field_schema->set_is_primary_key(false);
auto str_field_schema = collection_schema.add_fields();
str_field_schema->set_name("str_field");
str_field_schema->set_fieldid(104);
str_field_schema->set_data_type(schema::DataType::VarChar);
auto str_type_params = str_field_schema->add_type_params();
str_type_params->set_key(MAX_LENGTH);
str_type_params->set_value(std::to_string(64));
str_field_schema->set_is_primary_key(false);
auto vec_field_schema = collection_schema.add_fields();
vec_field_schema->set_name("fake_vec");
vec_field_schema->set_fieldid(105);
vec_field_schema->set_data_type(schema::DataType::FloatVector);
auto metric_type_param = vec_field_schema->add_index_params();
metric_type_param->set_key("metric_type");
metric_type_param->set_value(knowhere::metric::L2);
auto dim_param = vec_field_schema->add_type_params();
dim_param->set_key("dim");
dim_param->set_value(std::to_string(dim));
c_collection = NewCollection(&collection_schema, knowhere::metric::L2);
}
CSegmentInterface segment;
auto status = NewSegment(c_collection, Growing, -1, &segment, false);
ASSERT_EQ(status.error_code, Success);
//2. generate data and insert
auto c_schema = ((milvus::segcore::Collection*)c_collection)->get_schema();
auto dataset = DataGen(c_schema, N);
int64_t offset;
PreInsert(segment, N, &offset);
auto insert_data = serialize(dataset.raw_);
auto ins_res = Insert(segment,
offset,
N,
dataset.row_ids_.data(),
dataset.timestamps_.data(),
insert_data.data(),
insert_data.size());
ASSERT_EQ(ins_res.error_code, Success);
//3. search
milvus::segcore::ScopedSchemaHandle schema_handle(*c_schema);
auto binary_plan = schema_handle.ParseGroupBySearch(
"", // expression (empty for no filter)
"fake_vec", // vector field name
topK, // topK
"L2", // metric_type
R"({"nprobe": 10})", // search_params
101, // group_by_field_id (int8_field)
1); // group_size
auto blob = generate_query_data(num_queries);
void* plan = nullptr;
status = CreateSearchPlanByExpr(
c_collection, binary_plan.data(), binary_plan.size(), &plan);
ASSERT_EQ(status.error_code, Success);
void* placeholderGroup = nullptr;
status = ParsePlaceholderGroup(
plan, blob.data(), blob.length(), &placeholderGroup);
ASSERT_EQ(status.error_code, Success);
std::vector<CPlaceholderGroup> placeholderGroups;
placeholderGroups.push_back(placeholderGroup);
dataset.timestamps_.clear();
dataset.timestamps_.push_back(1);
CSearchResult res1;
CSearchResult res2;
auto res = CSearch(
segment, plan, placeholderGroup, dataset.timestamps_[N - 1], &res1);
ASSERT_EQ(res.error_code, Success);
res = CSearch(
segment, plan, placeholderGroup, dataset.timestamps_[N - 1], &res2);
ASSERT_EQ(res.error_code, Success);
//4. set up stream reducer
auto slice_nqs = std::vector<int64_t>{num_queries / 2, num_queries / 2};
if (num_queries == 1) {
slice_nqs = std::vector<int64_t>{num_queries};
}
auto slice_topKs = std::vector<int64_t>{topK, topK};
if (topK == 1) {
slice_topKs = std::vector<int64_t>{topK, topK};
}
CSearchStreamReducer c_search_stream_reducer;
NewStreamReducer(plan,
slice_nqs.data(),
slice_topKs.data(),
slice_nqs.size(),
&c_search_stream_reducer);
//5. stream reduce
StreamReduce(c_search_stream_reducer, &res1, 1);
StreamReduce(c_search_stream_reducer, &res2, 1);
CSearchResultDataBlobs c_search_result_data_blobs;
GetStreamReduceResult(c_search_stream_reducer, &c_search_result_data_blobs);
SearchResultDataBlobs* search_result_data_blob =
(SearchResultDataBlobs*)(c_search_result_data_blobs);
//6. check result
for (size_t i = 0; i < slice_nqs.size(); i++) {
milvus::proto::schema::SearchResultData search_result_data;
auto suc = search_result_data.ParseFromArray(
search_result_data_blob->blobs[i].data(),
search_result_data_blob->blobs[i].size());
ASSERT_TRUE(suc);
ASSERT_EQ(search_result_data.num_queries(), slice_nqs[i]);
ASSERT_EQ(search_result_data.top_k(), slice_topKs[i]);
ASSERT_EQ(search_result_data.ids().int_id().data_size(),
search_result_data.topks().at(0) * slice_nqs[i]);
ASSERT_EQ(search_result_data.scores().size(),
search_result_data.topks().at(0) * slice_nqs[i]);
ASSERT_TRUE(search_result_data.has_group_by_field_value());
// check real topks
ASSERT_EQ(search_result_data.topks().size(), slice_nqs[i]);
for (auto real_topk : search_result_data.topks()) {
ASSERT_LE(real_topk, slice_topKs[i]);
}
}
DeleteSearchResultDataBlobs(c_search_result_data_blobs);
DeleteSearchPlan(plan);
DeletePlaceholderGroup(placeholderGroup);
DeleteSearchResult(res1);
DeleteSearchResult(res2);
DeleteCollection(c_collection);
DeleteSegment(segment);
DeleteStreamSearchReducer(c_search_stream_reducer);
DeleteStreamSearchReducer(nullptr);
}
TEST(CApiTest, ReduceSearchResultsAndFillDataCost) {
int N = 100;
int topK = 10;
@@ -1,758 +0,0 @@
// Copyright (C) 2019-2020 Zilliz. All rights reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software distributed under the License
// is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express
// or implied. See the License for the specific language governing permissions and limitations under the License
#include "StreamReduce.h"
#include "NamedType/named_type_impl.hpp"
#include "common/FieldMeta.h"
#include "common/QueryInfo.h"
#include "common/Schema.h"
#include "common/protobuf_utils.h"
#include "fmt/core.h"
#include "pb/schema.pb.h"
#include "query/PlanImpl.h"
#include "query/PlanNode.h"
#include "segcore/ReduceUtils.h"
#include "segcore/SegmentInterface.h"
#include "segcore/Utils.h"
#include "segcore/pkVisitor.h"
#include "segcore/reduce/Reduce.h"
namespace milvus::segcore {
void
StreamReducerHelper::SetNullableVectorValidDataOffsets(
const std::map<FieldId, std::unique_ptr<milvus::DataArray>>&
output_fields_data,
int64_t ki,
MergeBase& merge_base) {
for (auto field_id : plan_->target_entries_) {
auto& field_meta = plan_->schema_->operator[](field_id);
if (field_meta.is_vector() && field_meta.is_nullable()) {
auto it = output_fields_data.find(field_id);
if (it != output_fields_data.end()) {
auto& field_data = it->second;
if (field_data->valid_data_size() > 0) {
int64_t physical_offset = 0;
for (int64_t j = 0; j < ki; ++j) {
if (field_data->valid_data(j)) {
physical_offset++;
}
}
merge_base.setValidDataOffset(field_id, physical_offset);
}
}
}
}
}
void
StreamReducerHelper::FillEntryData() {
for (auto search_result : search_results_to_merge_) {
auto segment = static_cast<milvus::segcore::SegmentInterface*>(
search_result->segment_);
segment->FillTargetEntry(plan_, *search_result);
}
}
void
StreamReducerHelper::AssembleMergedResult() {
if (search_results_to_merge_.size() > 0) {
std::unique_ptr<MergedSearchResult> new_merged_result =
std::make_unique<MergedSearchResult>();
std::vector<PkType> new_merged_pks;
std::vector<float> new_merged_distances;
std::vector<GroupByValueType> new_merged_groupBy_vals;
std::vector<MergeBase> merge_output_data_bases;
std::vector<int64_t> new_result_offsets;
bool need_handle_groupBy =
plan_->plan_node_->search_info_.group_by_field_id_.has_value();
int valid_size = 0;
std::vector<int> real_topKs(total_nq_);
for (int i = 0; i < num_slice_; i++) {
auto nq_begin = slice_nqs_prefix_sum_[i];
auto nq_end = slice_nqs_prefix_sum_[i + 1];
int64_t result_count = 0;
for (auto search_result : search_results_to_merge_) {
AssertInfo(
search_result->topk_per_nq_prefix_sum_.size() ==
search_result->total_nq_ + 1,
"incorrect topk_per_nq_prefix_sum_ size in search result");
result_count +=
search_result->topk_per_nq_prefix_sum_[nq_end] -
search_result->topk_per_nq_prefix_sum_[nq_begin];
}
if (merged_search_result->has_result_) {
result_count +=
merged_search_result->topk_per_nq_prefix_sum_[nq_end] -
merged_search_result->topk_per_nq_prefix_sum_[nq_begin];
}
int nq_base_offset = valid_size;
valid_size += result_count;
new_merged_pks.resize(valid_size);
new_merged_distances.resize(valid_size);
merge_output_data_bases.resize(valid_size);
new_result_offsets.resize(valid_size);
if (need_handle_groupBy) {
new_merged_groupBy_vals.resize(valid_size);
}
for (auto qi = nq_begin; qi < nq_end; qi++) {
for (auto search_result : search_results_to_merge_) {
AssertInfo(search_result != nullptr,
"null search result when reorganize");
if (search_result->result_offsets_.size() == 0) {
continue;
}
auto topK_start =
search_result->topk_per_nq_prefix_sum_[qi];
auto topK_end =
search_result->topk_per_nq_prefix_sum_[qi + 1];
for (auto ki = topK_start; ki < topK_end; ki++) {
auto loc = search_result->result_offsets_[ki];
AssertInfo(loc < result_count && loc >= 0,
"invalid loc when GetSearchResultDataSlice, "
"loc = " +
std::to_string(loc) +
", result_count = " +
std::to_string(result_count));
new_merged_pks[nq_base_offset + loc] =
search_result->primary_keys_[ki];
new_merged_distances[nq_base_offset + loc] =
search_result->distances_[ki];
if (need_handle_groupBy) {
new_merged_groupBy_vals[nq_base_offset + loc] =
search_result->group_by_values_.value()[ki];
}
merge_output_data_bases[nq_base_offset + loc] = {
&search_result->output_fields_data_, ki};
SetNullableVectorValidDataOffsets(
search_result->output_fields_data_,
ki,
merge_output_data_bases[nq_base_offset + loc]);
new_result_offsets[nq_base_offset + loc] = loc;
real_topKs[qi]++;
}
}
if (merged_search_result->has_result_) {
auto topK_start =
merged_search_result->topk_per_nq_prefix_sum_[qi];
auto topK_end =
merged_search_result->topk_per_nq_prefix_sum_[qi + 1];
for (auto ki = topK_start; ki < topK_end; ki++) {
auto loc = merged_search_result->reduced_offsets_[ki];
AssertInfo(loc < result_count && loc >= 0,
"invalid loc when GetSearchResultDataSlice, "
"loc = " +
std::to_string(loc) +
", result_count = " +
std::to_string(result_count));
new_merged_pks[nq_base_offset + loc] =
merged_search_result->primary_keys_[ki];
new_merged_distances[nq_base_offset + loc] =
merged_search_result->distances_[ki];
if (need_handle_groupBy) {
new_merged_groupBy_vals[nq_base_offset + loc] =
merged_search_result->group_by_values_
.value()[ki];
}
merge_output_data_bases[nq_base_offset + loc] = {
&merged_search_result->output_fields_data_, ki};
SetNullableVectorValidDataOffsets(
merged_search_result->output_fields_data_,
ki,
merge_output_data_bases[nq_base_offset + loc]);
new_result_offsets[nq_base_offset + loc] = loc;
real_topKs[qi]++;
}
}
}
}
new_merged_result->primary_keys_ = std::move(new_merged_pks);
new_merged_result->distances_ = std::move(new_merged_distances);
if (need_handle_groupBy) {
new_merged_result->group_by_values_ =
std::move(new_merged_groupBy_vals);
}
new_merged_result->topk_per_nq_prefix_sum_.resize(total_nq_ + 1);
std::partial_sum(
real_topKs.begin(),
real_topKs.end(),
new_merged_result->topk_per_nq_prefix_sum_.begin() + 1);
new_merged_result->result_offsets_ = std::move(new_result_offsets);
for (auto field_id : plan_->target_entries_) {
auto& field_meta = plan_->schema_->operator[](field_id);
auto field_data =
MergeDataArray(merge_output_data_bases, field_meta);
if (field_meta.get_data_type() == DataType::ARRAY) {
field_data->mutable_scalars()
->mutable_array_data()
->set_element_type(
proto::schema::DataType(field_meta.get_element_type()));
} else if (field_meta.get_data_type() == DataType::VECTOR_ARRAY) {
field_data->mutable_vectors()
->mutable_vector_array()
->set_element_type(
proto::schema::DataType(field_meta.get_element_type()));
}
new_merged_result->output_fields_data_[field_id] =
std::move(field_data);
}
merged_search_result = std::move(new_merged_result);
merged_search_result->has_result_ = true;
}
}
void
StreamReducerHelper::MergeReduce() {
GetTotalStorageCost();
FilterSearchResults();
FillPrimaryKeys();
InitializeReduceRecords();
ReduceResultData();
RefreshSearchResult();
FillEntryData();
AssembleMergedResult();
CleanReduceStatus();
}
void*
StreamReducerHelper::SerializeMergedResult() {
std::unique_ptr<SearchResultDataBlobs> search_result_blobs =
std::make_unique<milvus::segcore::SearchResultDataBlobs>();
AssertInfo(num_slice_ > 0,
"Wrong state for num_slice in streamReducer, num_slice:{}",
num_slice_);
search_result_blobs->blobs.resize(num_slice_);
search_result_blobs->costs.resize(num_slice_);
for (int i = 0; i < num_slice_; i++) {
auto [proto, cost] =
GetSearchResultDataSlice(i, total_search_storage_cost_);
search_result_blobs->blobs[i] = std::move(proto);
search_result_blobs->costs[i] = cost;
}
return search_result_blobs.release();
}
void
StreamReducerHelper::ReduceResultData() {
if (search_results_to_merge_.size() > 0) {
for (int i = 0; i < num_segments_; i++) {
auto search_result = search_results_to_merge_[i];
auto result_count = search_result->get_total_result_count();
AssertInfo(search_result != nullptr,
"search result must not equal to nullptr");
AssertInfo(search_result->distances_.size() == result_count,
"incorrect search result distance size");
AssertInfo(search_result->seg_offsets_.size() == result_count,
"incorrect search result seg offset size");
AssertInfo(search_result->primary_keys_.size() == result_count,
"incorrect search result primary key size");
}
for (int64_t slice_index = 0; slice_index < slice_nqs_.size();
slice_index++) {
auto nq_begin = slice_nqs_prefix_sum_[slice_index];
auto nq_end = slice_nqs_prefix_sum_[slice_index + 1];
int64_t offset = 0;
for (int64_t qi = nq_begin; qi < nq_end; qi++) {
StreamReduceSearchResultForOneNQ(
qi, slice_topKs_[slice_index], offset);
}
}
}
}
void
StreamReducerHelper::FilterSearchResults() {
uint32_t valid_index = 0;
for (auto& search_result : search_results_to_merge_) {
// skip when results num is 0
AssertInfo(search_result != nullptr,
"search_result to merge cannot be nullptr, there must be "
"sth wrong in the code");
if (search_result->unity_topK_ == 0) {
continue;
}
FilterInvalidSearchResult(search_result);
search_results_to_merge_[valid_index++] = search_result;
}
search_results_to_merge_.resize(valid_index);
num_segments_ = search_results_to_merge_.size();
}
void
StreamReducerHelper::InitializeReduceRecords() {
// init final_search_records and final_read_topKs
if (merged_search_result->has_result_) {
final_search_records_.resize(num_segments_ + 1);
} else {
final_search_records_.resize(num_segments_);
}
for (auto& search_record : final_search_records_) {
search_record.resize(total_nq_);
}
}
void
StreamReducerHelper::FillPrimaryKeys() {
for (auto& search_result : search_results_to_merge_) {
auto segment = static_cast<SegmentInterface*>(search_result->segment_);
if (search_result->get_total_result_count() > 0) {
segment->FillPrimaryKeys(plan_, *search_result);
}
}
}
void
StreamReducerHelper::FilterInvalidSearchResult(SearchResult* search_result) {
auto total_nq = search_result->total_nq_;
auto topK = search_result->unity_topK_;
AssertInfo(search_result->seg_offsets_.size() == total_nq * topK,
"wrong seg offsets size, size = " +
std::to_string(search_result->seg_offsets_.size()) +
", expected size = " + std::to_string(total_nq * topK));
AssertInfo(search_result->distances_.size() == total_nq * topK,
"wrong distances size, size = " +
std::to_string(search_result->distances_.size()) +
", expected size = " + std::to_string(total_nq * topK));
std::vector<int64_t> real_topKs(total_nq, 0);
uint32_t valid_index = 0;
auto segment = static_cast<SegmentInterface*>(search_result->segment_);
auto& offsets = search_result->seg_offsets_;
auto& distances = search_result->distances_;
if (search_result->group_by_values_.has_value()) {
AssertInfo(search_result->distances_.size() ==
search_result->group_by_values_.value().size(),
"wrong group_by_values size, size:{}, expected size:{} ",
search_result->group_by_values_.value().size(),
search_result->distances_.size());
}
for (auto i = 0; i < total_nq; ++i) {
for (auto j = 0; j < topK; ++j) {
auto index = i * topK + j;
if (offsets[index] != INVALID_SEG_OFFSET) {
AssertInfo(0 <= offsets[index] &&
offsets[index] < segment->get_row_count(),
fmt::format("invalid offset {}, segment {} with "
"rows num {}, data or index corruption",
offsets[index],
segment->get_segment_id(),
segment->get_row_count()));
real_topKs[i]++;
offsets[valid_index] = offsets[index];
distances[valid_index] = distances[index];
if (search_result->group_by_values_.has_value())
search_result->group_by_values_.value()[valid_index] =
search_result->group_by_values_.value()[index];
valid_index++;
}
}
}
offsets.resize(valid_index);
distances.resize(valid_index);
if (search_result->group_by_values_.has_value())
search_result->group_by_values_.value().resize(valid_index);
search_result->topk_per_nq_prefix_sum_.resize(total_nq + 1);
std::partial_sum(real_topKs.begin(),
real_topKs.end(),
search_result->topk_per_nq_prefix_sum_.begin() + 1);
}
void
StreamReducerHelper::StreamReduceSearchResultForOneNQ(int64_t qi,
int64_t topK,
int64_t& offset) {
//1. clear heap for preceding left elements
while (!heap_.empty()) {
heap_.pop();
}
pk_set_.clear();
group_by_val_set_.clear();
//2. push new search results into sort-heap
for (int i = 0; i < num_segments_; i++) {
auto search_result = search_results_to_merge_[i];
auto offset_beg = search_result->topk_per_nq_prefix_sum_[qi];
auto offset_end = search_result->topk_per_nq_prefix_sum_[qi + 1];
if (offset_beg == offset_end) {
continue;
}
auto primary_key = search_result->primary_keys_[offset_beg];
auto distance = search_result->distances_[offset_beg];
if (search_result->group_by_values_.has_value()) {
AssertInfo(
search_result->group_by_values_.value().size() > offset_beg,
"Wrong size for group_by_values size to "
"ReduceSearchResultForOneNQ:{}, not enough for"
"required offset_beg:{}",
search_result->group_by_values_.value().size(),
offset_beg);
}
auto result_pair = std::make_shared<StreamSearchResultPair>(
primary_key,
distance,
search_result,
nullptr,
i,
offset_beg,
offset_end,
search_result->group_by_values_.has_value() &&
search_result->group_by_values_.value().size() > offset_beg
? std::make_optional(
search_result->group_by_values_.value().at(offset_beg))
: std::nullopt);
heap_.push(result_pair);
}
if (heap_.empty()) {
return;
}
//3. if the merged_search_result has previous data
//push merged search result into the heap
if (merged_search_result->has_result_) {
auto merged_off_begin =
merged_search_result->topk_per_nq_prefix_sum_[qi];
auto merged_off_end =
merged_search_result->topk_per_nq_prefix_sum_[qi + 1];
if (merged_off_end > merged_off_begin) {
auto merged_pk =
merged_search_result->primary_keys_[merged_off_begin];
auto merged_distance =
merged_search_result->distances_[merged_off_begin];
auto merged_result_pair = std::make_shared<StreamSearchResultPair>(
merged_pk,
merged_distance,
nullptr,
merged_search_result.get(),
num_segments_, //use last index as the merged segment idex
merged_off_begin,
merged_off_end,
merged_search_result->group_by_values_.has_value() &&
merged_search_result->group_by_values_.value().size() >
merged_off_begin
? std::make_optional(
merged_search_result->group_by_values_.value().at(
merged_off_begin))
: std::nullopt);
heap_.push(merged_result_pair);
}
}
//3. pop heap to sort
int count = 0;
while (count < topK && !heap_.empty()) {
auto pilot = heap_.top();
heap_.pop();
auto seg_index = pilot->segment_index_;
auto pk = pilot->primary_key_;
if (pk == INVALID_PK) {
break; // valid search result for this nq has been run out, break to next
}
if (pk_set_.count(pk) == 0) {
bool skip_for_group_by = false;
if (pilot->group_by_value_.has_value()) {
if (group_by_val_set_.count(pilot->group_by_value_.value()) >
0) {
skip_for_group_by = true;
}
}
if (!skip_for_group_by) {
final_search_records_[seg_index][qi].push_back(pilot->offset_);
if (pilot->search_result_ != nullptr) {
pilot->search_result_->result_offsets_.push_back(offset++);
} else {
merged_search_result->reduced_offsets_.push_back(offset++);
}
pk_set_.insert(pk);
if (pilot->group_by_value_.has_value()) {
group_by_val_set_.insert(pilot->group_by_value_.value());
}
count++;
}
}
pilot->advance();
if (pilot->primary_key_ != INVALID_PK) {
heap_.push(pilot);
}
}
}
void
StreamReducerHelper::RefreshSearchResult() {
//1. refresh new input results
for (int i = 0; i < num_segments_; i++) {
std::vector<int64_t> real_topKs(total_nq_, 0);
auto search_result = search_results_to_merge_[i];
if (search_result->result_offsets_.size() > 0) {
uint32_t final_size = 0;
for (int j = 0; j < total_nq_; j++) {
final_size += final_search_records_[i][j].size();
}
std::vector<milvus::PkType> reduced_pks(final_size);
std::vector<float> reduced_distances(final_size);
std::vector<int64_t> reduced_seg_offsets(final_size);
std::vector<GroupByValueType> reduced_group_by_values(final_size);
uint32_t final_index = 0;
for (int j = 0; j < total_nq_; j++) {
for (auto offset : final_search_records_[i][j]) {
reduced_pks[final_index] =
search_result->primary_keys_[offset];
reduced_distances[final_index] =
search_result->distances_[offset];
reduced_seg_offsets[final_index] =
search_result->seg_offsets_[offset];
if (search_result->group_by_values_.has_value())
reduced_group_by_values[final_index] =
search_result->group_by_values_.value()[offset];
final_index++;
real_topKs[j]++;
}
}
search_result->primary_keys_.swap(reduced_pks);
search_result->distances_.swap(reduced_distances);
search_result->seg_offsets_.swap(reduced_seg_offsets);
if (search_result->group_by_values_.has_value()) {
search_result->group_by_values_.value().swap(
reduced_group_by_values);
}
}
std::partial_sum(real_topKs.begin(),
real_topKs.end(),
search_result->topk_per_nq_prefix_sum_.begin() + 1);
}
//2. refresh merged search result possibly
if (merged_search_result->has_result_) {
std::vector<int64_t> real_topKs(total_nq_, 0);
if (merged_search_result->reduced_offsets_.size() > 0) {
uint32_t final_size = merged_search_result->reduced_offsets_.size();
std::vector<milvus::PkType> reduced_pks(final_size);
std::vector<float> reduced_distances(final_size);
std::vector<int64_t> reduced_seg_offsets(final_size);
std::vector<GroupByValueType> reduced_group_by_values(final_size);
uint32_t final_index = 0;
for (int j = 0; j < total_nq_; j++) {
for (auto offset : final_search_records_[num_segments_][j]) {
reduced_pks[final_index] =
merged_search_result->primary_keys_[offset];
reduced_distances[final_index] =
merged_search_result->distances_[offset];
if (merged_search_result->group_by_values_.has_value())
reduced_group_by_values[final_index] =
merged_search_result->group_by_values_
.value()[offset];
final_index++;
real_topKs[j]++;
}
}
merged_search_result->primary_keys_.swap(reduced_pks);
merged_search_result->distances_.swap(reduced_distances);
if (merged_search_result->group_by_values_.has_value()) {
merged_search_result->group_by_values_.value().swap(
reduced_group_by_values);
}
}
std::partial_sum(
real_topKs.begin(),
real_topKs.end(),
merged_search_result->topk_per_nq_prefix_sum_.begin() + 1);
}
}
std::pair<std::vector<char>, StorageCost>
StreamReducerHelper::GetSearchResultDataSlice(int slice_index,
const StorageCost& total_cost) {
auto nq_begin = slice_nqs_prefix_sum_[slice_index];
auto nq_end = slice_nqs_prefix_sum_[slice_index + 1];
auto search_result_data =
std::make_unique<milvus::proto::schema::SearchResultData>();
// set unify_topK and total_nq
search_result_data->set_top_k(slice_topKs_[slice_index]);
search_result_data->set_num_queries(nq_end - nq_begin);
search_result_data->mutable_topks()->Resize(nq_end - nq_begin, 0);
int64_t result_count = 0;
if (merged_search_result->has_result_) {
AssertInfo(
nq_begin < merged_search_result->topk_per_nq_prefix_sum_.size(),
"nq_begin is incorrect for reduce, nq_begin:{}, topk_size:{}",
nq_begin,
merged_search_result->topk_per_nq_prefix_sum_.size());
AssertInfo(
nq_end < merged_search_result->topk_per_nq_prefix_sum_.size(),
"nq_end is incorrect for reduce, nq_end:{}, topk_size:{}",
nq_end,
merged_search_result->topk_per_nq_prefix_sum_.size());
result_count = merged_search_result->topk_per_nq_prefix_sum_[nq_end] -
merged_search_result->topk_per_nq_prefix_sum_[nq_begin];
}
StorageCost cost = total_cost * (1.0 * (nq_end - nq_begin) / total_nq_);
// `result_pairs` contains the SearchResult and result_offset info, used for filling output fields
std::vector<MergeBase> result_pairs(result_count);
// reserve space for pks
auto primary_field_id =
plan_->schema_->get_primary_field_id().value_or(milvus::FieldId(-1));
AssertInfo(primary_field_id.get() != INVALID_FIELD_ID, "Primary key is -1");
auto pk_type = plan_->schema_->operator[](primary_field_id).get_data_type();
switch (pk_type) {
case milvus::DataType::INT64: {
auto ids = std::make_unique<milvus::proto::schema::LongArray>();
ids->mutable_data()->Resize(result_count, 0);
search_result_data->mutable_ids()->set_allocated_int_id(
ids.release());
break;
}
case milvus::DataType::VARCHAR: {
auto ids = std::make_unique<milvus::proto::schema::StringArray>();
std::vector<std::string> string_pks(result_count);
// TODO: prevent mem copy
*ids->mutable_data() = {string_pks.begin(), string_pks.end()};
search_result_data->mutable_ids()->set_allocated_str_id(
ids.release());
break;
}
default: {
ThrowInfo(DataTypeInvalid,
fmt::format("unsupported primary key type {}", pk_type));
}
}
// reserve space for distances
search_result_data->mutable_scores()->Resize(result_count, 0);
//reserve space for group_by_values
std::vector<GroupByValueType> group_by_values;
if (plan_->plan_node_->search_info_.group_by_field_id_.has_value()) {
group_by_values.resize(result_count);
}
// fill pks and distances
for (auto qi = nq_begin; qi < nq_end; qi++) {
int64_t topk_count = 0;
AssertInfo(merged_search_result != nullptr,
"null merged search result when reorganize");
if (!merged_search_result->has_result_ ||
merged_search_result->result_offsets_.size() == 0) {
continue;
}
auto topk_start = merged_search_result->topk_per_nq_prefix_sum_[qi];
auto topk_end = merged_search_result->topk_per_nq_prefix_sum_[qi + 1];
topk_count += topk_end - topk_start;
for (auto ki = topk_start; ki < topk_end; ki++) {
auto loc = merged_search_result->result_offsets_[ki];
AssertInfo(loc < result_count && loc >= 0,
"invalid loc when GetSearchResultDataSlice, loc = " +
std::to_string(loc) +
", result_count = " + std::to_string(result_count));
// set result pks
switch (pk_type) {
case milvus::DataType::INT64: {
search_result_data->mutable_ids()
->mutable_int_id()
->mutable_data()
->Set(loc,
std::visit(
Int64PKVisitor{},
merged_search_result->primary_keys_[ki]));
break;
}
case milvus::DataType::VARCHAR: {
*search_result_data->mutable_ids()
->mutable_str_id()
->mutable_data()
->Mutable(loc) =
std::visit(StrPKVisitor{},
merged_search_result->primary_keys_[ki]);
break;
}
default: {
ThrowInfo(DataTypeInvalid,
fmt::format("unsupported primary key type {}",
pk_type));
}
}
search_result_data->mutable_scores()->Set(
loc, merged_search_result->distances_[ki]);
// set group by values
if (merged_search_result->group_by_values_.has_value() &&
ki < merged_search_result->group_by_values_.value().size())
group_by_values[loc] =
merged_search_result->group_by_values_.value()[ki];
// set result offset to fill output fields data
result_pairs[loc] = {&merged_search_result->output_fields_data_,
ki};
}
// update result topKs
search_result_data->mutable_topks()->Set(qi - nq_begin, topk_count);
}
AssembleGroupByValues(search_result_data, group_by_values, plan_);
AssertInfo(search_result_data->scores_size() == result_count,
"wrong scores size, size = " +
std::to_string(search_result_data->scores_size()) +
", expected size = " + std::to_string(result_count));
// set output fields
for (auto field_id : plan_->target_entries_) {
auto& field_meta = plan_->schema_->operator[](field_id);
auto field_data =
milvus::segcore::MergeDataArray(result_pairs, field_meta);
if (field_meta.get_data_type() == DataType::ARRAY) {
field_data->mutable_scalars()
->mutable_array_data()
->set_element_type(
proto::schema::DataType(field_meta.get_element_type()));
} else if (field_meta.get_data_type() == DataType::VECTOR_ARRAY) {
field_data->mutable_vectors()
->mutable_vector_array()
->set_element_type(
proto::schema::DataType(field_meta.get_element_type()));
}
search_result_data->mutable_fields_data()->AddAllocated(
field_data.release());
}
// SearchResultData to blob
auto size = search_result_data->ByteSizeLong();
auto buffer = std::vector<char>(size);
search_result_data->SerializePartialToArray(buffer.data(), size);
return {std::move(buffer), cost};
}
void
StreamReducerHelper::GetTotalStorageCost() {
for (auto search_result : search_results_to_merge_) {
total_search_storage_cost_ += search_result->search_storage_cost_;
}
}
void
StreamReducerHelper::CleanReduceStatus() {
this->final_search_records_.clear();
this->merged_search_result->reduced_offsets_.clear();
}
} // namespace milvus::segcore
@@ -1,236 +0,0 @@
// Copyright (C) 2019-2020 Zilliz. All rights reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software distributed under the License
// is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express
// or implied. See the License for the specific language governing permissions and limitations under the License
#pragma once
#include <queue>
#include <unordered_set>
#include "common/Types.h"
#include "segcore/segment_c.h"
#include "query/PlanImpl.h"
#include "common/QueryResult.h"
#include "segcore/ReduceStructure.h"
#include "segcore/Utils.h"
#include "common/EasyAssert.h"
namespace milvus::segcore {
class MergedSearchResult {
public:
bool has_result_;
std::vector<PkType> primary_keys_;
std::vector<float> distances_;
std::optional<std::vector<GroupByValueType>> group_by_values_;
// set output fields data when filling target entity
std::map<FieldId, std::unique_ptr<milvus::DataArray>> output_fields_data_;
// used for reduce, filter invalid pk, get real topks count
std::vector<size_t> topk_per_nq_prefix_sum_;
// fill data during reducing search result
std::vector<int64_t> result_offsets_;
std::vector<int64_t> reduced_offsets_;
};
struct StreamSearchResultPair {
milvus::PkType primary_key_;
float distance_;
milvus::SearchResult* search_result_;
MergedSearchResult* merged_result_;
int64_t segment_index_;
int64_t offset_;
int64_t offset_rb_;
std::optional<milvus::GroupByValueType> group_by_value_;
StreamSearchResultPair(milvus::PkType primary_key,
float distance,
SearchResult* result,
int64_t index,
int64_t lb,
int64_t rb)
: StreamSearchResultPair(primary_key,
distance,
result,
nullptr,
index,
lb,
rb,
std::nullopt) {
}
StreamSearchResultPair(
milvus::PkType primary_key,
float distance,
SearchResult* result,
MergedSearchResult* merged_result,
int64_t index,
int64_t lb,
int64_t rb,
std::optional<milvus::GroupByValueType> group_by_value)
: primary_key_(std::move(primary_key)),
distance_(distance),
search_result_(result),
merged_result_(merged_result),
segment_index_(index),
offset_(lb),
offset_rb_(rb),
group_by_value_(group_by_value) {
AssertInfo(
search_result_ != nullptr || merged_result_ != nullptr,
"For a valid StreamSearchResult pair, "
"at least one of merged_result_ or search_result_ is not nullptr");
}
bool
operator>(const StreamSearchResultPair& other) const {
if (std::fabs(distance_ - other.distance_) < EPSILON) {
return primary_key_ < other.primary_key_;
}
return distance_ > other.distance_;
}
void
advance() {
offset_++;
if (offset_ < offset_rb_) {
if (search_result_ != nullptr) {
primary_key_ = search_result_->primary_keys_.at(offset_);
distance_ = search_result_->distances_.at(offset_);
if (search_result_->group_by_values_.has_value() &&
offset_ < search_result_->group_by_values_.value().size()) {
group_by_value_ =
search_result_->group_by_values_.value().at(offset_);
}
} else {
primary_key_ = merged_result_->primary_keys_.at(offset_);
distance_ = merged_result_->distances_.at(offset_);
if (merged_result_->group_by_values_.has_value() &&
offset_ < merged_result_->group_by_values_.value().size()) {
group_by_value_ =
merged_result_->group_by_values_.value().at(offset_);
}
}
} else {
primary_key_ = INVALID_PK;
distance_ = std::numeric_limits<float>::min();
}
}
};
struct StreamSearchResultPairComparator {
bool
operator()(const std::shared_ptr<StreamSearchResultPair> lhs,
const std::shared_ptr<StreamSearchResultPair> rhs) const {
return (*rhs.get()) > (*lhs.get());
}
};
class StreamReducerHelper {
public:
explicit StreamReducerHelper(milvus::query::Plan* plan,
int64_t* slice_nqs,
int64_t* slice_topKs,
int64_t slice_num)
: plan_(plan),
slice_nqs_(slice_nqs, slice_nqs + slice_num),
slice_topKs_(slice_topKs, slice_topKs + slice_num) {
AssertInfo(slice_nqs_.size() > 0, "empty_nqs");
AssertInfo(slice_nqs_.size() == slice_topKs_.size(),
"unaligned slice_nqs and slice_topKs");
merged_search_result = std::make_unique<MergedSearchResult>();
merged_search_result->has_result_ = false;
num_slice_ = slice_nqs_.size();
slice_nqs_prefix_sum_.resize(num_slice_ + 1);
std::partial_sum(slice_nqs_.begin(),
slice_nqs_.end(),
slice_nqs_prefix_sum_.begin() + 1);
total_nq_ = slice_nqs_prefix_sum_[num_slice_];
AssertInfo(total_nq_ > 0, "empty nq is not allowed");
}
void
SetSearchResultsToMerge(std::vector<SearchResult*>& search_results) {
search_results_to_merge_ = search_results;
num_segments_ = search_results_to_merge_.size();
AssertInfo(num_segments_ > 0, "empty search result");
}
public:
void
MergeReduce();
void*
SerializeMergedResult();
protected:
void
FilterSearchResults();
void
InitializeReduceRecords();
void
FillPrimaryKeys();
void
FilterInvalidSearchResult(SearchResult* search_result);
void
ReduceResultData();
private:
void
RefreshSearchResult();
void
StreamReduceSearchResultForOneNQ(int64_t qi, int64_t topK, int64_t& offset);
void
FillEntryData();
void
AssembleMergedResult();
std::pair<std::vector<char>, StorageCost>
GetSearchResultDataSlice(const int slice_index,
const StorageCost& total_cost);
void
GetTotalStorageCost();
void
CleanReduceStatus();
void
SetNullableVectorValidDataOffsets(
const std::map<FieldId, std::unique_ptr<milvus::DataArray>>&
output_fields_data,
int64_t ki,
MergeBase& merge_base);
std::unique_ptr<MergedSearchResult> merged_search_result;
milvus::query::Plan* plan_;
std::vector<int64_t> slice_nqs_;
std::vector<int64_t> slice_topKs_;
std::vector<SearchResult*> search_results_to_merge_;
int64_t num_segments_{0};
int64_t num_slice_{0};
std::vector<int64_t> slice_nqs_prefix_sum_;
std::priority_queue<std::shared_ptr<StreamSearchResultPair>,
std::vector<std::shared_ptr<StreamSearchResultPair>>,
StreamSearchResultPairComparator>
heap_;
std::unordered_set<milvus::PkType> pk_set_;
std::unordered_set<milvus::GroupByValueType> group_by_val_set_;
std::vector<std::vector<std::vector<int64_t>>> final_search_records_;
int64_t total_nq_{0};
StorageCost total_search_storage_cost_;
};
} // namespace milvus::segcore
-72
View File
@@ -24,70 +24,10 @@
#include "segcore/ReduceStructure.h"
#include "segcore/reduce/GroupReduce.h"
#include "segcore/reduce/Reduce.h"
#include "segcore/reduce/StreamReduce.h"
#include "segcore/reduce_c.h"
using SearchResult = milvus::SearchResult;
CStatus
NewStreamReducer(CSearchPlan c_plan,
int64_t* slice_nqs,
int64_t* slice_topKs,
int64_t num_slices,
CSearchStreamReducer* stream_reducer) {
SCOPE_CGO_CALL_METRIC();
try {
//convert search results and search plan
auto plan = static_cast<milvus::query::Plan*>(c_plan);
auto stream_reduce_helper =
std::make_unique<milvus::segcore::StreamReducerHelper>(
plan, slice_nqs, slice_topKs, num_slices);
*stream_reducer = stream_reduce_helper.release();
return milvus::SuccessCStatus();
} catch (std::exception& e) {
return milvus::FailureCStatus(&e);
}
}
CStatus
StreamReduce(CSearchStreamReducer c_stream_reducer,
CSearchResult* c_search_results,
int64_t num_segments) {
SCOPE_CGO_CALL_METRIC();
try {
auto stream_reducer =
static_cast<milvus::segcore::StreamReducerHelper*>(
c_stream_reducer);
std::vector<SearchResult*> search_results(num_segments);
for (int i = 0; i < num_segments; i++) {
search_results[i] = static_cast<SearchResult*>(c_search_results[i]);
}
stream_reducer->SetSearchResultsToMerge(search_results);
stream_reducer->MergeReduce();
return milvus::SuccessCStatus();
} catch (std::exception& e) {
return milvus::FailureCStatus(&e);
}
}
CStatus
GetStreamReduceResult(CSearchStreamReducer c_stream_reducer,
CSearchResultDataBlobs* c_search_result_data_blobs) {
SCOPE_CGO_CALL_METRIC();
try {
auto stream_reducer =
static_cast<milvus::segcore::StreamReducerHelper*>(
c_stream_reducer);
*c_search_result_data_blobs = stream_reducer->SerializeMergedResult();
return milvus::SuccessCStatus();
} catch (std::exception& e) {
return milvus::FailureCStatus(&e);
}
}
CStatus
ReduceSearchResultsAndFillData(CTraceContext c_trace,
CSearchResultDataBlobs* cSearchResultDataBlobs,
@@ -184,15 +124,3 @@ DeleteSearchResultDataBlobs(CSearchResultDataBlobs cSearchResultDataBlobs) {
cSearchResultDataBlobs);
delete search_result_data_blobs;
}
void
DeleteStreamSearchReducer(CSearchStreamReducer c_stream_reducer) {
SCOPE_CGO_CALL_METRIC();
if (c_stream_reducer == nullptr) {
return;
}
auto stream_reducer =
static_cast<milvus::segcore::StreamReducerHelper*>(c_stream_reducer);
delete stream_reducer;
}
-20
View File
@@ -21,23 +21,6 @@ extern "C" {
#include "segcore/segment_c.h"
typedef void* CSearchResultDataBlobs;
typedef void* CSearchStreamReducer;
CStatus
NewStreamReducer(CSearchPlan c_plan,
int64_t* slice_nqs,
int64_t* slice_topKs,
int64_t num_slices,
CSearchStreamReducer* stream_reducer);
CStatus
StreamReduce(CSearchStreamReducer c_stream_reducer,
CSearchResult* c_search_results,
int64_t num_segments);
CStatus
GetStreamReduceResult(CSearchStreamReducer c_stream_reducer,
CSearchResultDataBlobs* c_search_result_data_blobs);
CStatus
ReduceSearchResultsAndFillData(CTraceContext c_trace,
@@ -59,9 +42,6 @@ GetSearchResultDataBlob(CProto* searchResultDataBlob,
void
DeleteSearchResultDataBlobs(CSearchResultDataBlobs cSearchResultDataBlobs);
void
DeleteStreamSearchReducer(CSearchStreamReducer c_stream_reducer);
#ifdef __cplusplus
}
#endif
-104
View File
@@ -19,9 +19,7 @@ package segments
import (
"context"
"fmt"
"sync"
"go.uber.org/atomic"
"go.uber.org/zap"
"golang.org/x/sync/errgroup"
@@ -116,90 +114,6 @@ func searchSegments(ctx context.Context, mgr *Manager, segments []Segment, segTy
return searchResults, nil
}
// searchSegmentsStreamly performs search on listed segments in a stream mode instead of a batch mode
// all segment ids are validated before calling this function
func searchSegmentsStreamly(ctx context.Context,
mgr *Manager,
segments []Segment,
searchReq *SearchRequest,
streamReduce func(result *SearchResult) error,
) error {
searchLabel := metrics.SealedSegmentLabel
searchResultsToClear := make([]*SearchResult, 0)
var reduceMutex sync.Mutex
var sumReduceDuration atomic.Duration
searcher := func(ctx context.Context, seg Segment) error {
// record search time
tr := timerecord.NewTimeRecorder("searchOnSegments")
searchResult, searchErr := seg.Search(ctx, searchReq)
searchDuration := tr.RecordSpan().Milliseconds()
if searchErr != nil {
return searchErr
}
reduceMutex.Lock()
searchResultsToClear = append(searchResultsToClear, searchResult)
reducedErr := streamReduce(searchResult)
reduceMutex.Unlock()
reduceDuration := tr.RecordSpan()
if reducedErr != nil {
return reducedErr
}
sumReduceDuration.Add(reduceDuration)
// update metrics
metrics.QueryNodeSQSegmentLatency.WithLabelValues(fmt.Sprint(paramtable.GetNodeID()),
metrics.SearchLabel, searchLabel).Observe(float64(searchDuration))
metrics.QueryNodeSegmentSearchLatencyPerVector.WithLabelValues(fmt.Sprint(paramtable.GetNodeID()),
metrics.SearchLabel, searchLabel).Observe(float64(searchDuration) / float64(searchReq.GetNumOfQuery()))
return nil
}
// calling segment search in goroutines
errGroup, ctx := errgroup.WithContext(ctx)
log := log.Ctx(ctx)
for _, segment := range segments {
seg := segment
errGroup.Go(func() error {
if ctx.Err() != nil {
return ctx.Err()
}
var err error
accessRecord := metricsutil.NewSearchSegmentAccessRecord(getSegmentMetricLabel(seg))
defer func() {
accessRecord.Finish(err)
}()
if seg.IsLazyLoad() {
log.Debug("before doing stream search in DiskCache", zap.Int64("segID", seg.ID()))
ctx, cancel := withLazyLoadTimeoutContext(ctx)
defer cancel()
var missing bool
missing, err = mgr.DiskCache.Do(ctx, seg.ID(), searcher)
if missing {
accessRecord.CacheMissing()
}
if err != nil {
log.Warn("failed to do search for disk cache", zap.Int64("segID", seg.ID()), zap.Error(err))
}
log.Debug("after doing stream search in DiskCache", zap.Int64("segID", seg.ID()), zap.Error(err))
return err
}
return searcher(ctx, seg)
})
}
err := errGroup.Wait()
DeleteSearchResults(searchResultsToClear)
if err != nil {
return err
}
metrics.QueryNodeReduceLatency.WithLabelValues(fmt.Sprint(paramtable.GetNodeID()),
metrics.SearchLabel,
metrics.ReduceSegments,
metrics.StreamReduce).Observe(float64(sumReduceDuration.Load().Milliseconds()))
log.Debug("stream reduce sum duration:", zap.Duration("duration", sumReduceDuration.Load()))
return nil
}
// search will search on the historical segments the target segments in historical.
// if segIDs is not specified, it will search on all the historical segments speficied by partIDs.
// if segIDs is specified, it will only search on the segments specified by the segIDs.
@@ -231,21 +145,3 @@ func SearchStreaming(ctx context.Context, manager *Manager, searchReq *SearchReq
searchResults, err := searchSegments(ctx, manager, segments, SegmentTypeGrowing, searchReq)
return searchResults, segments, err
}
func SearchHistoricalStreamly(ctx context.Context, manager *Manager, searchReq *SearchRequest,
collID int64, partIDs []int64, segIDs []int64, plan *planpb.PlanNode, streamReduce func(result *SearchResult) error,
) ([]Segment, error) {
if ctx.Err() != nil {
return nil, ctx.Err()
}
segments, err := validateOnHistorical(ctx, manager, collID, partIDs, segIDs, plan)
if err != nil {
return segments, err
}
err = searchSegmentsStreamly(ctx, manager, segments, searchReq, streamReduce)
if err != nil {
return segments, err
}
return segments, nil
}
@@ -279,117 +279,6 @@ func (suite *SearchSuite) TestSearchWithFilter() {
}
}
func (suite *SearchSuite) TestSearchStreamWithFilter() {
ctx := context.Background()
// create sealed segments with different pk ranges for testing
loader := NewLoader(ctx, suite.manager, suite.chunkManager)
for i := range 10 {
segID := int64(i + 2000)
msgLen := i + 1
binlogs, statslogs, err := mock_segcore.SaveBinLog(ctx,
suite.collectionID,
suite.partitionID,
segID,
msgLen,
suite.schema,
suite.chunkManager,
)
suite.Require().NoError(err)
loadInfo := &querypb.SegmentLoadInfo{
SegmentID: segID,
CollectionID: suite.collectionID,
PartitionID: suite.partitionID,
NumOfRows: int64(msgLen),
BinlogPaths: binlogs,
Statslogs: statslogs,
InsertChannel: fmt.Sprintf("by-dev-rootcoord-dml_0_%dv0", suite.collectionID),
Level: datapb.SegmentLevel_Legacy,
}
seg, err := NewSegment(ctx,
suite.collection,
suite.manager.Segment,
SegmentTypeSealed,
0,
loadInfo,
)
suite.Require().NoError(err)
bfs, err := loader.loadSingleBloomFilterSet(ctx, suite.collectionID, loadInfo, SegmentTypeSealed)
suite.Require().NoError(err)
seg.SetBloomFilter(bfs)
for _, binlog := range binlogs {
err = seg.(*LocalSegment).LoadFieldData(ctx, binlog.FieldID, int64(msgLen), binlog)
suite.Require().NoError(err)
}
suite.manager.Segment.Put(ctx, SegmentTypeSealed, seg)
}
segIDs := []int64{2000, 2001, 2002, 2003, 2004, 2005, 2006, 2007, 2008, 2009}
suite.Run("SearchHistoricalStreamlyWithFilter", func() {
searchReq, err := mock_segcore.GenSearchPlanAndRequests(suite.collection.GetCCollection(), segIDs, mock_segcore.IndexFaissIDMap, 1)
suite.NoError(err)
// create search plan with filter expression "int64Field == 5"
schemaHelper, _ := typeutil.CreateSchemaHelper(suite.schema)
planNode, err := planparserv2.CreateSearchPlan(schemaHelper, "int64Field == 5", "floatVectorField", &planpb.QueryInfo{
Topk: 10,
MetricType: "L2",
SearchParams: `{"nprobe": 10}`,
RoundDecimal: -1,
}, nil, nil)
suite.NoError(err)
streamReduceFunc := func(result *SearchResult) error {
return nil
}
// search with filter - should only search segments containing pk=5 (seg6-seg10)
segments, err := SearchHistoricalStreamly(ctx, suite.manager, searchReq,
suite.collectionID,
nil,
segIDs,
planNode,
streamReduceFunc,
)
suite.NoError(err)
// with sparse filter enabled, only 5 segments (seg6-seg10) should be searched
suite.Len(segments, 5)
suite.manager.Segment.Unpin(segments)
})
suite.Run("SearchHistoricalStreamlyWithoutFilter", func() {
searchReq, err := mock_segcore.GenSearchPlanAndRequests(suite.collection.GetCCollection(), segIDs, mock_segcore.IndexFaissIDMap, 1)
suite.NoError(err)
streamReduceFunc := func(result *SearchResult) error {
return nil
}
// search without filter - should search all segments
segments, err := SearchHistoricalStreamly(ctx, suite.manager, searchReq,
suite.collectionID,
nil,
segIDs,
nil,
streamReduceFunc,
)
suite.NoError(err)
suite.Len(segments, 10)
suite.manager.Segment.Unpin(segments)
})
// cleanup
for _, segID := range segIDs {
suite.manager.Segment.Remove(ctx, segID, querypb.DataScope_Historical)
}
}
func TestSearch(t *testing.T) {
suite.Run(t, new(SearchSuite))
}
+1 -7
View File
@@ -40,7 +40,6 @@ import (
"github.com/milvus-io/milvus/internal/storagev2"
"github.com/milvus-io/milvus/internal/util/analyzer"
"github.com/milvus-io/milvus/internal/util/fileresource"
"github.com/milvus-io/milvus/internal/util/searchutil/scheduler"
"github.com/milvus-io/milvus/internal/util/streamrpc"
"github.com/milvus-io/milvus/internal/util/textmatch"
"github.com/milvus-io/milvus/pkg/v2/common"
@@ -773,12 +772,7 @@ func (node *QueryNode) SearchSegments(ctx context.Context, req *querypb.SearchRe
node.manager.Collection.Unref(req.GetReq().GetCollectionID(), 1)
}()
var task scheduler.Task
if paramtable.Get().QueryNodeCfg.UseStreamComputing.GetAsBool() {
task = tasks.NewStreamingSearchTask(searchCtx, collection, node.manager, req, node.serverID)
} else {
task = tasks.NewSearchTask(searchCtx, collection, node.manager, req, node.serverID)
}
task := tasks.NewSearchTask(searchCtx, collection, node.manager, req, node.serverID)
if err := node.scheduler.Add(task); err != nil {
log.Warn("failed to search channel", zap.Error(err))
-226
View File
@@ -396,229 +396,3 @@ func (t *SearchTask) combinePlaceHolderGroups() error {
t.placeholderGroup, _ = proto.Marshal(ret)
return nil
}
type StreamingSearchTask struct {
SearchTask
others []*StreamingSearchTask
resultBlobs segcore.SearchResultDataBlobs
streamReducer segcore.StreamSearchReducer
}
func NewStreamingSearchTask(ctx context.Context,
collection *segments.Collection,
manager *segments.Manager,
req *querypb.SearchRequest,
serverID int64,
) *StreamingSearchTask {
ctx, span := otel.Tracer(typeutil.QueryNodeRole).Start(ctx, "schedule")
return &StreamingSearchTask{
SearchTask: SearchTask{
ctx: ctx,
collection: collection,
segmentManager: manager,
req: req,
plan: &planpb.PlanNode{},
merged: false,
groupSize: 1,
topk: req.GetReq().GetTopk(),
nq: req.GetReq().GetNq(),
placeholderGroup: req.GetReq().GetPlaceholderGroup(),
originTopks: []int64{req.GetReq().GetTopk()},
originNqs: []int64{req.GetReq().GetNq()},
notifier: make(chan error, 1),
tr: timerecord.NewTimeRecorderWithTrace(ctx, "searchTask"),
scheduleSpan: span,
serverID: serverID,
},
}
}
func (t *StreamingSearchTask) MergeWith(other scheduler.Task) bool {
return false
}
func (t *StreamingSearchTask) Execute() error {
log := log.Ctx(t.ctx).With(
zap.Int64("collectionID", t.collection.ID()),
zap.String("shard", t.req.GetDmlChannels()[0]),
)
// 0. prepare search req
if t.scheduleSpan != nil {
t.scheduleSpan.End()
}
tr := timerecord.NewTimeRecorderWithTrace(t.ctx, "SearchTask")
req := t.req
t.combinePlaceHolderGroups()
searchReq, err := segcore.NewSearchRequest(t.collection.GetCCollection(), req, t.placeholderGroup)
if err != nil {
return err
}
defer searchReq.Delete()
// 1. search&&reduce or streaming-search&&streaming-reduce
metricType := searchReq.Plan().GetMetricType()
var relatedDataSize int64
if req.GetScope() == querypb.DataScope_Historical {
streamReduceFunc := func(result *segments.SearchResult) error {
reduceErr := t.streamReduce(t.ctx, searchReq.Plan(), result, t.originNqs, t.originTopks)
return reduceErr
}
pinnedSegments, err := segments.SearchHistoricalStreamly(
t.ctx,
t.segmentManager,
searchReq,
req.GetReq().GetCollectionID(),
nil,
req.GetSegmentIDs(),
t.plan,
streamReduceFunc)
defer segcore.DeleteStreamReduceHelper(t.streamReducer)
defer t.segmentManager.Segment.Unpin(pinnedSegments)
if err != nil {
log.Error("Failed to search sealed segments streamly", zap.Error(err))
return err
}
t.resultBlobs, err = segcore.GetStreamReduceResult(t.ctx, t.streamReducer)
defer segcore.DeleteSearchResultDataBlobs(t.resultBlobs)
if err != nil {
log.Error("Failed to get stream-reduced search result")
return err
}
relatedDataSize = lo.Reduce(pinnedSegments, func(acc int64, seg segments.Segment, _ int) int64 {
return acc + segments.GetSegmentRelatedDataSize(seg)
}, 0)
} else if req.GetScope() == querypb.DataScope_Streaming {
results, pinnedSegments, err := segments.SearchStreaming(
t.ctx,
t.segmentManager,
searchReq,
req.GetReq().GetCollectionID(),
req.GetReq().GetPartitionIDs(),
req.GetSegmentIDs(),
t.plan,
)
defer segments.DeleteSearchResults(results)
defer t.segmentManager.Segment.Unpin(pinnedSegments)
if err != nil {
return err
}
if t.maybeReturnForEmptyResults(results, metricType, tr) {
return nil
}
tr.RecordSpan()
t.resultBlobs, err = segcore.ReduceSearchResultsAndFillData(
t.ctx,
searchReq.Plan(),
results,
int64(len(results)),
t.originNqs,
t.originTopks,
)
if err != nil {
log.Warn("failed to reduce search results", zap.Error(err))
return err
}
defer segcore.DeleteSearchResultDataBlobs(t.resultBlobs)
metrics.QueryNodeReduceLatency.WithLabelValues(
fmt.Sprint(t.GetNodeID()),
metrics.SearchLabel,
metrics.ReduceSegments,
metrics.BatchReduce).
Observe(float64(tr.RecordSpan().Milliseconds()))
relatedDataSize = lo.Reduce(pinnedSegments, func(acc int64, seg segments.Segment, _ int) int64 {
return acc + segments.GetSegmentRelatedDataSize(seg)
}, 0)
}
// 2. reorganize blobs to original search request
for i := range t.originNqs {
blob, cost, err := segcore.GetSearchResultDataBlob(t.ctx, t.resultBlobs, i)
if err != nil {
return err
}
var task *StreamingSearchTask
if i == 0 {
task = t
} else {
task = t.others[i-1]
}
// Note: blob is unsafe because get from C
bs := make([]byte, len(blob))
copy(bs, blob)
task.result = &internalpb.SearchResults{
Base: &commonpb.MsgBase{
SourceID: t.GetNodeID(),
},
Status: merr.Success(),
MetricType: metricType,
NumQueries: t.originNqs[i],
TopK: t.originTopks[i],
SlicedBlob: bs,
SlicedOffset: 1,
SlicedNumCount: 1,
CostAggregation: &internalpb.CostAggregation{
ServiceTime: tr.ElapseSpan().Milliseconds(),
TotalRelatedDataSize: relatedDataSize,
},
ScannedRemoteBytes: cost.ScannedRemoteBytes,
ScannedTotalBytes: cost.ScannedTotalBytes,
}
}
return nil
}
func (t *StreamingSearchTask) maybeReturnForEmptyResults(results []*segments.SearchResult,
metricType string, tr *timerecord.TimeRecorder,
) bool {
if len(results) == 0 {
for i := range t.originNqs {
var task *StreamingSearchTask
if i == 0 {
task = t
} else {
task = t.others[i-1]
}
task.result = &internalpb.SearchResults{
Base: &commonpb.MsgBase{
SourceID: t.GetNodeID(),
},
Status: merr.Success(),
MetricType: metricType,
NumQueries: t.originNqs[i],
TopK: t.originTopks[i],
SlicedOffset: 1,
SlicedNumCount: 1,
CostAggregation: &internalpb.CostAggregation{
ServiceTime: tr.ElapseSpan().Milliseconds(),
},
ScannedRemoteBytes: 0,
ScannedTotalBytes: 0,
}
}
return true
}
return false
}
func (t *StreamingSearchTask) streamReduce(ctx context.Context,
plan *segcore.SearchPlan,
newResult *segments.SearchResult,
sliceNQs []int64,
sliceTopKs []int64,
) error {
if t.streamReducer == nil {
var err error
t.streamReducer, err = segcore.NewStreamReducer(ctx, plan, sliceNQs, sliceTopKs)
if err != nil {
log.Error("Fail to init stream reducer, return")
return err
}
}
return segcore.StreamReduceSearchResult(ctx, newResult, t.streamReducer)
}
+1 -57
View File
@@ -37,10 +37,7 @@ type SliceInfo struct {
}
// SearchResultDataBlobs is the CSearchResultsDataBlobs in C++
type (
SearchResultDataBlobs = C.CSearchResultDataBlobs
StreamSearchReducer = C.CSearchStreamReducer
)
type SearchResultDataBlobs = C.CSearchResultDataBlobs
type StorageCost struct {
ScannedRemoteBytes int64
@@ -71,55 +68,6 @@ func ParseSliceInfo(originNQs []int64, originTopKs []int64, nqPerSlice int64) *S
return sInfo
}
func NewStreamReducer(ctx context.Context,
plan *SearchPlan,
sliceNQs []int64,
sliceTopKs []int64,
) (StreamSearchReducer, error) {
if plan.cSearchPlan == nil {
return nil, errors.New("nil search plan")
}
if len(sliceNQs) == 0 {
return nil, errors.New("empty slice nqs is not allowed")
}
if len(sliceNQs) != len(sliceTopKs) {
return nil, fmt.Errorf("unaligned sliceNQs(len=%d) and sliceTopKs(len=%d)", len(sliceNQs), len(sliceTopKs))
}
cSliceNQSPtr := (*C.int64_t)(&sliceNQs[0])
cSliceTopKSPtr := (*C.int64_t)(&sliceTopKs[0])
cNumSlices := C.int64_t(len(sliceNQs))
var streamReducer StreamSearchReducer
status := C.NewStreamReducer(plan.cSearchPlan, cSliceNQSPtr, cSliceTopKSPtr, cNumSlices, &streamReducer)
if err := ConsumeCStatusIntoError(&status); err != nil {
return nil, errors.Wrap(err, "MergeSearchResultsWithOutputFields failed")
}
return streamReducer, nil
}
func StreamReduceSearchResult(ctx context.Context,
newResult *SearchResult, streamReducer StreamSearchReducer,
) error {
cSearchResults := make([]C.CSearchResult, 0)
cSearchResults = append(cSearchResults, newResult.cSearchResult)
cSearchResultPtr := &cSearchResults[0]
status := C.StreamReduce(streamReducer, cSearchResultPtr, 1)
if err := ConsumeCStatusIntoError(&status); err != nil {
return errors.Wrap(err, "StreamReduceSearchResult failed")
}
return nil
}
func GetStreamReduceResult(ctx context.Context, streamReducer StreamSearchReducer) (SearchResultDataBlobs, error) {
var cSearchResultDataBlobs SearchResultDataBlobs
status := C.GetStreamReduceResult(streamReducer, &cSearchResultDataBlobs)
if err := ConsumeCStatusIntoError(&status); err != nil {
return nil, errors.Wrap(err, "ReduceSearchResultsAndFillData failed")
}
return cSearchResultDataBlobs, nil
}
func ReduceSearchResultsAndFillData(ctx context.Context, plan *SearchPlan, searchResults []*SearchResult,
numSegments int64, sliceNQs []int64, sliceTopKs []int64,
) (SearchResultDataBlobs, error) {
@@ -171,7 +119,3 @@ func GetSearchResultDataBlob(ctx context.Context, cSearchResultDataBlobs SearchR
func DeleteSearchResultDataBlobs(cSearchResultDataBlobs SearchResultDataBlobs) {
C.DeleteSearchResultDataBlobs(cSearchResultDataBlobs)
}
func DeleteStreamReduceHelper(cStreamReduceHelper StreamSearchReducer) {
C.DeleteStreamSearchReducer(cStreamReduceHelper)
}
+1 -2
View File
@@ -73,8 +73,7 @@ const (
ReduceSegments = "segments"
ReduceShards = "shards"
BatchReduce = "batch_reduce"
StreamReduce = "stream_reduce"
BatchReduce = "batch_reduce"
Pending = "pending"
Executing = "executing"