mirror of
https://github.com/milvus-io/milvus.git
synced 2026-07-21 10:15:43 +00:00
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:
co-authored by
MrPresent-Han
parent
88e2ae08b7
commit
af11c336fd
@@ -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
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -73,8 +73,7 @@ const (
|
||||
ReduceSegments = "segments"
|
||||
ReduceShards = "shards"
|
||||
|
||||
BatchReduce = "batch_reduce"
|
||||
StreamReduce = "stream_reduce"
|
||||
BatchReduce = "batch_reduce"
|
||||
|
||||
Pending = "pending"
|
||||
Executing = "executing"
|
||||
|
||||
Reference in New Issue
Block a user