mirror of
https://github.com/milvus-io/milvus.git
synced 2026-07-21 02:05:41 +00:00
feat: impl StructArray -- support group by for element-level search (#47252)
issue: https://github.com/milvus-io/milvus/issues/42148 design doc: https://github.com/milvus-io/milvus-design-docs/blob/main/design_docs/20260306-struct.md --------- Signed-off-by: SpadeA <tangchenjie1210@gmail.com>
This commit is contained in:
@@ -317,16 +317,6 @@ class QueryContext : public Context {
|
||||
return plan_options_;
|
||||
}
|
||||
|
||||
void
|
||||
set_element_level_query(bool element_level) {
|
||||
element_level_query_ = element_level;
|
||||
}
|
||||
|
||||
bool
|
||||
element_level_query() const {
|
||||
return element_level_query_;
|
||||
}
|
||||
|
||||
void
|
||||
set_struct_name(const std::string& field_name) {
|
||||
struct_name_ = field_name;
|
||||
@@ -417,7 +407,6 @@ class QueryContext : public Context {
|
||||
|
||||
query::PlanOptions plan_options_;
|
||||
|
||||
bool element_level_query_{false};
|
||||
std::string struct_name_;
|
||||
std::shared_ptr<const IArrayOffsets> array_offsets_{nullptr};
|
||||
int64_t active_element_count_{0}; // Total elements in active documents
|
||||
|
||||
@@ -117,7 +117,6 @@ PhyElementFilterBitsNode::GetOutput() {
|
||||
doc_bitset, doc_bitset_valid, array_offsets.get());
|
||||
|
||||
// Step 4: Set query context
|
||||
query_context_->set_element_level_query(true);
|
||||
query_context_->set_struct_name(struct_name_);
|
||||
// Mark that bitset has been converted to element-level
|
||||
query_context_->set_bitset_is_element_level(true);
|
||||
|
||||
@@ -79,6 +79,13 @@ PhySearchGroupByNode::GetOutput() {
|
||||
|
||||
auto op_context = query_context_->get_op_context();
|
||||
auto search_result = query_context_->get_search_result();
|
||||
|
||||
search_info_.array_offsets_ = query_context_->get_array_offsets();
|
||||
if (search_result.element_level_) {
|
||||
AssertInfo(search_info_.array_offsets_ != nullptr,
|
||||
"Array offsets not available");
|
||||
}
|
||||
|
||||
if (search_result.vector_iterators_.has_value()) {
|
||||
AssertInfo(search_result.vector_iterators_.value().size() ==
|
||||
search_result.total_nq_,
|
||||
@@ -105,6 +112,11 @@ PhySearchGroupByNode::GetOutput() {
|
||||
tracer::AddEvent(
|
||||
fmt::format("grouped_results: {}", search_result.seg_offsets_.size()));
|
||||
|
||||
// When group by is used with element-level search, the result is row-level
|
||||
// because group by operates on rows, not elements.
|
||||
search_result.element_level_ = false;
|
||||
search_result.element_indices_.clear();
|
||||
|
||||
query_context_->set_search_result(std::move(search_result));
|
||||
std::chrono::high_resolution_clock::time_point vector_end =
|
||||
std::chrono::high_resolution_clock::now();
|
||||
|
||||
@@ -18,6 +18,7 @@
|
||||
#include <algorithm>
|
||||
#include <cstdint>
|
||||
#include <tuple>
|
||||
#include <unordered_set>
|
||||
#include <variant>
|
||||
|
||||
#include "common/QueryInfo.h"
|
||||
@@ -52,14 +53,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<int8_t>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<int8_t>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -67,14 +65,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<int16_t>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<int16_t>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -82,14 +77,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<int32_t>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<int32_t>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -97,14 +89,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<int64_t>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<int64_t>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -112,14 +101,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<int64_t>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<int64_t>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -127,14 +113,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<bool>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<bool>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -142,14 +125,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
auto dataGetter =
|
||||
GetDataGetter<std::string>(op_ctx, segment, group_by_field_id);
|
||||
GroupIteratorsByType<std::string>(iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
dataGetter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -167,17 +147,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.json_path_,
|
||||
search_info.json_type_,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<bool>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
GroupIteratorsByType<bool>(iterators,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
case DataType::INT8: {
|
||||
@@ -188,17 +164,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.json_path_,
|
||||
search_info.json_type_,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<int8_t>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
GroupIteratorsByType<int8_t>(iterators,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
case DataType::INT16: {
|
||||
@@ -209,17 +181,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.json_path_,
|
||||
search_info.json_type_,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<int16_t>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
GroupIteratorsByType<int16_t>(iterators,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
case DataType::INT32: {
|
||||
@@ -230,17 +198,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.json_path_,
|
||||
search_info.json_type_,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<int32_t>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
GroupIteratorsByType<int32_t>(iterators,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
case DataType::INT64: {
|
||||
@@ -251,17 +215,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.json_path_,
|
||||
search_info.json_type_,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<int64_t>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
GroupIteratorsByType<int64_t>(iterators,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
case DataType::VARCHAR: {
|
||||
@@ -275,14 +235,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<std::string>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
break;
|
||||
}
|
||||
@@ -301,17 +258,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
|
||||
search_info.json_path_,
|
||||
search_info.json_type_,
|
||||
search_info.strict_cast_);
|
||||
GroupIteratorsByType<std::string>(
|
||||
iterators,
|
||||
search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
search_info.metric_type_,
|
||||
topk_per_nq_prefix_sum);
|
||||
GroupIteratorsByType<std::string>(iterators,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
topk_per_nq_prefix_sum);
|
||||
}
|
||||
break;
|
||||
}
|
||||
@@ -328,26 +281,20 @@ template <typename T>
|
||||
void
|
||||
GroupIteratorsByType(
|
||||
const std::vector<std::shared_ptr<VectorIterator>>& iterators,
|
||||
int64_t topK,
|
||||
int64_t group_size,
|
||||
bool strict_group_size,
|
||||
const SearchInfo& search_info,
|
||||
const std::shared_ptr<DataGetter<T>>& data_getter,
|
||||
std::vector<GroupByValueType>& group_by_values,
|
||||
std::vector<int64_t>& seg_offsets,
|
||||
std::vector<float>& distances,
|
||||
const knowhere::MetricType& metrics_type,
|
||||
std::vector<size_t>& topk_per_nq_prefix_sum) {
|
||||
topk_per_nq_prefix_sum.push_back(0);
|
||||
for (auto& iterator : iterators) {
|
||||
GroupIteratorResult<T>(iterator,
|
||||
topK,
|
||||
group_size,
|
||||
strict_group_size,
|
||||
search_info,
|
||||
data_getter,
|
||||
group_by_values,
|
||||
seg_offsets,
|
||||
distances,
|
||||
metrics_type);
|
||||
distances);
|
||||
topk_per_nq_prefix_sum.push_back(seg_offsets.size());
|
||||
}
|
||||
}
|
||||
@@ -355,21 +302,27 @@ GroupIteratorsByType(
|
||||
template <typename T>
|
||||
void
|
||||
GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
|
||||
int64_t topK,
|
||||
int64_t group_size,
|
||||
bool strict_group_size,
|
||||
const SearchInfo& search_info,
|
||||
const std::shared_ptr<DataGetter<T>>& data_getter,
|
||||
std::vector<GroupByValueType>& group_by_values,
|
||||
std::vector<int64_t>& offsets,
|
||||
std::vector<float>& distances,
|
||||
const knowhere::MetricType& metrics_type) {
|
||||
std::vector<float>& distances) {
|
||||
//1.
|
||||
GroupByMap<T> groupMap(topK, group_size, strict_group_size);
|
||||
GroupByMap<T> groupMap(search_info.topk_,
|
||||
search_info.group_size_,
|
||||
search_info.strict_group_size_);
|
||||
|
||||
auto is_element_id = search_info.element_level();
|
||||
auto array_offsets = search_info.array_offsets_;
|
||||
|
||||
//2. do iteration until fill the whole map or run out of all data
|
||||
//note it may enumerate all data inside a segment and can block following
|
||||
//query and search possibly
|
||||
std::vector<std::tuple<int64_t, float, std::optional<T>>> res;
|
||||
// For element-level search, multiple elements from the same row may be
|
||||
// returned by the iterator. We must deduplicate by row_offset so that
|
||||
// the same document does not consume multiple slots in a group.
|
||||
std::unordered_set<int64_t> seen_rows;
|
||||
while (iterator->HasNext() && !groupMap.IsGroupResEnough()) {
|
||||
auto offset_dis_pair = iterator->Next();
|
||||
AssertInfo(
|
||||
@@ -378,16 +331,30 @@ GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
|
||||
"tells hasNext, terminate groupBy operation");
|
||||
auto offset = offset_dis_pair.value().first;
|
||||
auto dis = offset_dis_pair.value().second;
|
||||
std::optional<T> row_data = data_getter->Get(offset);
|
||||
|
||||
// When array_offsets is present, iterator returns element_id,
|
||||
// but DataGetter expects row_id. Convert element_id to row_id.
|
||||
int64_t row_offset = offset;
|
||||
if (is_element_id) {
|
||||
row_offset =
|
||||
array_offsets->ElementIDToRowID(static_cast<int32_t>(offset))
|
||||
.first;
|
||||
// Skip if we already saw this row (closest element wins)
|
||||
if (!seen_rows.insert(row_offset).second) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
std::optional<T> row_data = data_getter->Get(row_offset);
|
||||
if (groupMap.Push(row_data)) {
|
||||
res.emplace_back(offset, dis, row_data);
|
||||
res.emplace_back(row_offset, dis, row_data);
|
||||
}
|
||||
}
|
||||
|
||||
//3. sorted based on distances and metrics
|
||||
auto customComparator = [&](const auto& lhs, const auto& rhs) {
|
||||
return milvus::query::dis_closer(
|
||||
std::get<1>(lhs), std::get<1>(rhs), metrics_type);
|
||||
std::get<1>(lhs), std::get<1>(rhs), search_info.metric_type_);
|
||||
};
|
||||
std::sort(res.begin(), res.end(), customComparator);
|
||||
|
||||
|
||||
@@ -343,14 +343,11 @@ template <typename T>
|
||||
void
|
||||
GroupIteratorsByType(
|
||||
const std::vector<std::shared_ptr<VectorIterator>>& iterators,
|
||||
int64_t topK,
|
||||
int64_t group_size,
|
||||
bool strict_group_size,
|
||||
const SearchInfo& search_info,
|
||||
const std::shared_ptr<DataGetter<T>>& data_getter,
|
||||
std::vector<GroupByValueType>& group_by_values,
|
||||
std::vector<int64_t>& seg_offsets,
|
||||
std::vector<float>& distances,
|
||||
const knowhere::MetricType& metrics_type,
|
||||
std::vector<size_t>& topk_per_nq_prefix_sum);
|
||||
|
||||
template <typename T>
|
||||
@@ -413,13 +410,11 @@ struct GroupByMap {
|
||||
template <typename T>
|
||||
void
|
||||
GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
|
||||
int64_t topK,
|
||||
int64_t group_size,
|
||||
bool strict_group_size,
|
||||
const SearchInfo& search_info,
|
||||
const std::shared_ptr<DataGetter<T>>& data_getter,
|
||||
std::vector<GroupByValueType>& group_by_values,
|
||||
std::vector<int64_t>& offsets,
|
||||
std::vector<float>& distances,
|
||||
const knowhere::MetricType& metrics_type);
|
||||
std::vector<float>& distances);
|
||||
|
||||
} // namespace exec
|
||||
} // namespace milvus
|
||||
|
||||
@@ -13,12 +13,12 @@
|
||||
#include <gtest/gtest.h>
|
||||
#include <stddef.h>
|
||||
#include <cstdint>
|
||||
#include <iostream>
|
||||
#include <map>
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
#include <tuple>
|
||||
#include <unordered_set>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
@@ -321,10 +321,6 @@ TEST_P(ElementFilterSealed, RangeExpr) {
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
int64_t doc_id = search_result->seg_offsets_[i];
|
||||
int32_t elem_idx = search_result->element_indices_[i];
|
||||
float distance = search_result->distances_[i];
|
||||
|
||||
std::cout << "doc_id: " << doc_id << ", element_index: " << elem_idx
|
||||
<< ", distance: " << distance << std::endl;
|
||||
|
||||
// Verify the doc_id satisfies the predicate (id % 2 == 0) only
|
||||
// when predicate is enabled
|
||||
@@ -544,16 +540,6 @@ TEST_P(ElementFilterSealed, UnaryExpr) {
|
||||
ASSERT_LE(search_result->element_indices_.size(),
|
||||
static_cast<size_t>(topK * num_queries));
|
||||
|
||||
std::cout << "Element-level search returned ("
|
||||
<< static_cast<int>(elem_type) << "):" << std::endl;
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
std::cout << "doc_id: " << search_result->seg_offsets_[i]
|
||||
<< ", element_index: "
|
||||
<< search_result->element_indices_[i]
|
||||
<< ", distance: " << search_result->distances_[i]
|
||||
<< std::endl;
|
||||
}
|
||||
|
||||
// Verify distances are sorted
|
||||
for (size_t i = 1; i < search_result->distances_.size(); ++i) {
|
||||
ASSERT_LE(search_result->distances_[i - 1],
|
||||
@@ -1227,16 +1213,9 @@ TEST_P(ElementFilterGrowing, RangeExpr) {
|
||||
static_cast<size_t>(topK * num_queries))
|
||||
<< "Should not exceed topK results";
|
||||
|
||||
std::cout << "Growing segment element-level search ("
|
||||
<< static_cast<int>(elem_type) << "):" << std::endl;
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
int64_t doc_id = search_result->seg_offsets_[i];
|
||||
int32_t elem_idx = search_result->element_indices_[i];
|
||||
float distance = search_result->distances_[i];
|
||||
|
||||
std::cout << " [" << i << "] doc_id=" << doc_id
|
||||
<< ", element_index=" << elem_idx
|
||||
<< ", distance=" << distance << std::endl;
|
||||
|
||||
// Verify the doc_id satisfies the predicate (id % 2 == 0) only
|
||||
// when predicate is enabled
|
||||
@@ -2131,11 +2110,6 @@ TEST_P(ElementFilterNestedIndex, ExecutionMode) {
|
||||
last_distance = search_result->distances_[i];
|
||||
}
|
||||
}
|
||||
|
||||
std::cout << "Test passed: index_type="
|
||||
<< NestedIndexTypeToString(index_type)
|
||||
<< ", offset_mode=" << (offset_mode ? "true" : "false")
|
||||
<< ", valid_results=" << valid_count << std::endl;
|
||||
}
|
||||
|
||||
INSTANTIATE_TEST_SUITE_P(
|
||||
@@ -2154,3 +2128,514 @@ INSTANTIATE_TEST_SUITE_P(
|
||||
name += offset_mode ? "_OffsetMode" : "_FullMode";
|
||||
return name;
|
||||
});
|
||||
|
||||
// Test element-level filter combined with group by on sealed segment with index
|
||||
TEST(ElementFilterGroupBy, SealedWithIndex) {
|
||||
int dim = 4;
|
||||
size_t N = 500;
|
||||
int array_len = 3;
|
||||
|
||||
auto schema = std::make_shared<Schema>();
|
||||
auto vec_fid = schema->AddDebugVectorArrayField("structA[array_vec]",
|
||||
DataType::VECTOR_FLOAT,
|
||||
dim,
|
||||
knowhere::metric::L2);
|
||||
auto int_array_fid = schema->AddDebugArrayField(
|
||||
"structA[price_array]", DataType::INT32, false);
|
||||
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
|
||||
schema->set_primary_field_id(int64_fid);
|
||||
|
||||
// Generate test data
|
||||
auto raw_data = DataGen(schema, N, 42, 0, 1, array_len);
|
||||
|
||||
// Customize int_array data: doc i has elements [i*3+1, i*3+2, i*3+3]
|
||||
for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) {
|
||||
auto* field_data = raw_data.raw_->mutable_fields_data(i);
|
||||
if (field_data->field_id() == int_array_fid.get()) {
|
||||
field_data->mutable_scalars()
|
||||
->mutable_array_data()
|
||||
->mutable_data()
|
||||
->Clear();
|
||||
|
||||
for (int row = 0; row < N; row++) {
|
||||
auto* array_data = field_data->mutable_scalars()
|
||||
->mutable_array_data()
|
||||
->mutable_data()
|
||||
->Add();
|
||||
|
||||
for (int elem = 0; elem < array_len; elem++) {
|
||||
int value = row * array_len + elem + 1;
|
||||
array_data->mutable_int_data()->mutable_data()->Add(value);
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Create sealed segment and load field data
|
||||
auto segment = CreateSealedWithFieldDataLoaded(schema, raw_data);
|
||||
|
||||
// Build and load vector index
|
||||
auto array_vec_values = raw_data.get_col<VectorFieldProto>(vec_fid);
|
||||
std::vector<float> vector_data(dim * N * array_len);
|
||||
for (int i = 0; i < N; i++) {
|
||||
const auto& float_vec = array_vec_values[i].float_vector().data();
|
||||
for (int j = 0; j < array_len * dim; j++) {
|
||||
vector_data[i * array_len * dim + j] = float_vec[j];
|
||||
}
|
||||
}
|
||||
|
||||
auto indexing = GenVecIndexing(N * array_len,
|
||||
dim,
|
||||
vector_data.data(),
|
||||
knowhere::IndexEnum::INDEX_HNSW);
|
||||
LoadIndexInfo load_index_info;
|
||||
load_index_info.field_id = vec_fid.get();
|
||||
load_index_info.index_params = GenIndexParams(indexing.get());
|
||||
load_index_info.cache_index =
|
||||
CreateTestCacheIndex("test", std::move(indexing));
|
||||
load_index_info.index_params["metric_type"] = knowhere::metric::L2;
|
||||
load_index_info.field_type = DataType::VECTOR_ARRAY;
|
||||
load_index_info.element_type = DataType::VECTOR_FLOAT;
|
||||
segment->LoadIndex(load_index_info);
|
||||
|
||||
int topK = 5;
|
||||
int group_size = 4;
|
||||
|
||||
// Execute element-level search with group by primary key
|
||||
ScopedSchemaHandle handle(*schema);
|
||||
std::string expr =
|
||||
"id % 2 == 0 && element_filter(structA, 400 > $[price_array] > 100)";
|
||||
std::string search_params = R"({"ef": 50})";
|
||||
|
||||
auto plan_bytes = handle.ParseGroupBySearch(expr,
|
||||
"structA[array_vec]",
|
||||
topK,
|
||||
knowhere::metric::L2,
|
||||
search_params,
|
||||
int64_fid.get(),
|
||||
group_size);
|
||||
auto plan =
|
||||
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
|
||||
ASSERT_NE(plan, nullptr);
|
||||
|
||||
auto num_queries = 1;
|
||||
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024, true);
|
||||
auto ph_group =
|
||||
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
|
||||
|
||||
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
|
||||
|
||||
ASSERT_NE(search_result, nullptr);
|
||||
ASSERT_TRUE(search_result->group_by_values_.has_value())
|
||||
<< "Group by values should be present";
|
||||
|
||||
// Group by returns row-level results even for element-level search
|
||||
ASSERT_FALSE(search_result->element_level_)
|
||||
<< "Group by should return row-level results";
|
||||
|
||||
auto& group_by_values = search_result->group_by_values_.value();
|
||||
|
||||
ASSERT_LE(search_result->seg_offsets_.size(),
|
||||
static_cast<size_t>(topK * group_size))
|
||||
<< "Should not exceed topK * group_size results";
|
||||
|
||||
std::unordered_map<int64_t, int> group_counts;
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
int64_t doc_id = search_result->seg_offsets_[i];
|
||||
|
||||
if (i < group_by_values.size() && group_by_values[i].has_value()) {
|
||||
if (std::holds_alternative<int64_t>(group_by_values[i].value())) {
|
||||
int64_t group_val =
|
||||
std::get<int64_t>(group_by_values[i].value());
|
||||
|
||||
ASSERT_EQ(group_val, doc_id)
|
||||
<< "Group by primary key: group value should equal doc_id";
|
||||
|
||||
group_counts[group_val]++;
|
||||
ASSERT_LE(group_counts[group_val], group_size)
|
||||
<< "Each group should have at most group_size results";
|
||||
}
|
||||
}
|
||||
|
||||
ASSERT_EQ(doc_id % 2, 0)
|
||||
<< "Result doc_id " << doc_id << " should satisfy (id % 2 == 0)";
|
||||
}
|
||||
|
||||
ASSERT_LE(group_counts.size(), static_cast<size_t>(topK))
|
||||
<< "Should have at most topK distinct groups";
|
||||
|
||||
for (size_t i = 1; i < search_result->distances_.size(); ++i) {
|
||||
ASSERT_LE(search_result->distances_[i - 1],
|
||||
search_result->distances_[i])
|
||||
<< "Distances should be sorted in ascending order";
|
||||
}
|
||||
}
|
||||
|
||||
// Test: element level + group by + growing segment
|
||||
TEST(ElementFilterGroupBy, GrowingSegment) {
|
||||
int dim = 4;
|
||||
auto schema = std::make_shared<Schema>();
|
||||
auto vec_fid = schema->AddDebugVectorArrayField(
|
||||
"structA[array_vec]", DataType::VECTOR_FLOAT, dim, "L2");
|
||||
auto int_array_fid = schema->AddDebugArrayField(
|
||||
"structA[price_array]", DataType::INT32, false);
|
||||
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
|
||||
schema->set_primary_field_id(int64_fid);
|
||||
|
||||
size_t N = 500;
|
||||
int array_len = 3;
|
||||
|
||||
auto raw_data = DataGen(schema, N, 42, 0, 1, array_len);
|
||||
|
||||
// Customize int_array data: doc i has elements [i*3+1, i*3+2, i*3+3]
|
||||
for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) {
|
||||
auto* field_data = raw_data.raw_->mutable_fields_data(i);
|
||||
if (field_data->field_id() == int_array_fid.get()) {
|
||||
field_data->mutable_scalars()
|
||||
->mutable_array_data()
|
||||
->mutable_data()
|
||||
->Clear();
|
||||
|
||||
for (int row = 0; row < N; row++) {
|
||||
auto* array_data = field_data->mutable_scalars()
|
||||
->mutable_array_data()
|
||||
->mutable_data()
|
||||
->Add();
|
||||
|
||||
for (int elem = 0; elem < array_len; elem++) {
|
||||
int value = row * array_len + elem + 1;
|
||||
array_data->mutable_int_data()->mutable_data()->Add(value);
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Create growing segment and insert data
|
||||
auto segment = CreateGrowingSegment(schema, empty_index_meta);
|
||||
segment->PreInsert(N);
|
||||
segment->Insert(0,
|
||||
N,
|
||||
raw_data.row_ids_.data(),
|
||||
raw_data.timestamps_.data(),
|
||||
raw_data.raw_);
|
||||
|
||||
int topK = 5;
|
||||
int group_size = 4;
|
||||
|
||||
// Execute element-level search with group by primary key
|
||||
ScopedSchemaHandle handle(*schema);
|
||||
std::string expr =
|
||||
"id % 2 == 0 && element_filter(structA, 400 > $[price_array] > 100)";
|
||||
std::string search_params = R"({"nprobe": 10})";
|
||||
|
||||
auto plan_bytes = handle.ParseGroupBySearch(expr,
|
||||
"structA[array_vec]",
|
||||
topK,
|
||||
knowhere::metric::L2,
|
||||
search_params,
|
||||
int64_fid.get(),
|
||||
group_size);
|
||||
auto plan =
|
||||
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
|
||||
ASSERT_NE(plan, nullptr);
|
||||
|
||||
auto num_queries = 1;
|
||||
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024, true);
|
||||
auto ph_group =
|
||||
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
|
||||
|
||||
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
|
||||
|
||||
ASSERT_NE(search_result, nullptr);
|
||||
ASSERT_TRUE(search_result->group_by_values_.has_value())
|
||||
<< "Group by values should be present";
|
||||
|
||||
// Verify element_level_ is false (group by returns row-level results)
|
||||
ASSERT_FALSE(search_result->element_level_)
|
||||
<< "Group by should return row-level results";
|
||||
|
||||
auto& group_by_values = search_result->group_by_values_.value();
|
||||
|
||||
ASSERT_FALSE(search_result->seg_offsets_.empty())
|
||||
<< "Should have search results";
|
||||
ASSERT_LE(search_result->seg_offsets_.size(),
|
||||
static_cast<size_t>(topK * group_size))
|
||||
<< "Should not exceed topK * group_size results";
|
||||
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
int64_t doc_id = search_result->seg_offsets_[i];
|
||||
|
||||
// Verify row-level filter
|
||||
ASSERT_EQ(doc_id % 2, 0)
|
||||
<< "Result doc_id " << doc_id << " should satisfy (id % 2 == 0)";
|
||||
|
||||
// Verify element-level filter: 100 < element_value < 400
|
||||
// element_value = doc_id * array_len + elem + 1, elem in [0, array_len)
|
||||
// For doc to have any matching element: doc_id * 3 + 1 < 400 and doc_id * 3 + 3 > 100
|
||||
// So doc_id should be roughly in range [33, 132]
|
||||
ASSERT_GE(doc_id, 33)
|
||||
<< "doc_id " << doc_id << " should be >= 33 (element filter)";
|
||||
ASSERT_LE(doc_id, 133)
|
||||
<< "doc_id " << doc_id << " should be <= 133 (element filter)";
|
||||
|
||||
if (i < group_by_values.size() && group_by_values[i].has_value()) {
|
||||
if (std::holds_alternative<int64_t>(group_by_values[i].value())) {
|
||||
int64_t group_val =
|
||||
std::get<int64_t>(group_by_values[i].value());
|
||||
ASSERT_EQ(group_val, doc_id)
|
||||
<< "Group by primary key: group value should equal doc_id";
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Test: normal group by (without element level)
|
||||
TEST(ElementFilterGroupBy, NormalGroupBy) {
|
||||
int dim = 4;
|
||||
auto schema = std::make_shared<Schema>();
|
||||
auto vec_fid =
|
||||
schema->AddDebugField("vec", DataType::VECTOR_FLOAT, dim, "L2");
|
||||
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
|
||||
auto category_fid = schema->AddDebugField("category", DataType::INT32);
|
||||
schema->set_primary_field_id(int64_fid);
|
||||
|
||||
size_t N = 500;
|
||||
auto raw_data = DataGen(schema, N, 42);
|
||||
|
||||
auto segment = CreateSealedWithFieldDataLoaded(schema, raw_data);
|
||||
|
||||
// Build vector index
|
||||
auto vec_values = raw_data.get_col<float>(vec_fid);
|
||||
auto indexing = GenVecIndexing(
|
||||
N, dim, vec_values.data(), knowhere::IndexEnum::INDEX_HNSW);
|
||||
|
||||
LoadIndexInfo load_index_info;
|
||||
load_index_info.field_id = vec_fid.get();
|
||||
load_index_info.index_params = GenIndexParams(indexing.get());
|
||||
load_index_info.cache_index =
|
||||
CreateTestCacheIndex("test", std::move(indexing));
|
||||
load_index_info.index_params["metric_type"] = knowhere::metric::L2;
|
||||
load_index_info.field_type = DataType::VECTOR_FLOAT;
|
||||
segment->LoadIndex(load_index_info);
|
||||
|
||||
int topK = 5;
|
||||
int group_size = 2;
|
||||
|
||||
// Normal search with group by (no element level)
|
||||
ScopedSchemaHandle handle(*schema);
|
||||
std::string expr = "id % 2 == 0";
|
||||
std::string search_params = R"({"ef": 50})";
|
||||
|
||||
auto plan_bytes = handle.ParseGroupBySearch(expr,
|
||||
"vec",
|
||||
topK,
|
||||
knowhere::metric::L2,
|
||||
search_params,
|
||||
category_fid.get(),
|
||||
group_size);
|
||||
auto plan =
|
||||
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
|
||||
ASSERT_NE(plan, nullptr);
|
||||
|
||||
auto num_queries = 1;
|
||||
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024);
|
||||
auto ph_group =
|
||||
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
|
||||
|
||||
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
|
||||
|
||||
ASSERT_NE(search_result, nullptr);
|
||||
ASSERT_TRUE(search_result->group_by_values_.has_value())
|
||||
<< "Group by values should be present";
|
||||
|
||||
// Verify element_level_ is false
|
||||
ASSERT_FALSE(search_result->element_level_)
|
||||
<< "Normal search should not be element-level";
|
||||
|
||||
auto& group_by_values = search_result->group_by_values_.value();
|
||||
|
||||
ASSERT_FALSE(search_result->seg_offsets_.empty())
|
||||
<< "Should have search results";
|
||||
ASSERT_LE(search_result->seg_offsets_.size(),
|
||||
static_cast<size_t>(topK * group_size))
|
||||
<< "Should not exceed topK * group_size results";
|
||||
|
||||
std::unordered_map<int32_t, int> group_counts;
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
int64_t doc_id = search_result->seg_offsets_[i];
|
||||
|
||||
ASSERT_EQ(doc_id % 2, 0)
|
||||
<< "Result doc_id " << doc_id << " should satisfy (id % 2 == 0)";
|
||||
|
||||
if (i < group_by_values.size() && group_by_values[i].has_value()) {
|
||||
if (std::holds_alternative<int32_t>(group_by_values[i].value())) {
|
||||
int32_t group_val =
|
||||
std::get<int32_t>(group_by_values[i].value());
|
||||
group_counts[group_val]++;
|
||||
ASSERT_LE(group_counts[group_val], group_size)
|
||||
<< "Each group should have at most group_size results";
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ASSERT_LE(group_counts.size(), static_cast<size_t>(topK))
|
||||
<< "Should have at most topK distinct groups";
|
||||
}
|
||||
|
||||
// Test: element-level + group by non-unique field to verify row deduplication.
|
||||
// When multiple elements from the same document are close to the query vector,
|
||||
// each document should only appear once in the results (not consume multiple
|
||||
// group slots).
|
||||
TEST(ElementFilterGroupBy, DeduplicateRowsInGroup) {
|
||||
int dim = 4;
|
||||
size_t N = 500;
|
||||
int array_len = 3;
|
||||
|
||||
auto schema = std::make_shared<Schema>();
|
||||
auto vec_fid = schema->AddDebugVectorArrayField("structA[array_vec]",
|
||||
DataType::VECTOR_FLOAT,
|
||||
dim,
|
||||
knowhere::metric::L2);
|
||||
auto int_array_fid = schema->AddDebugArrayField(
|
||||
"structA[price_array]", DataType::INT32, false);
|
||||
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
|
||||
// category field: non-unique, used for group by
|
||||
// Assign category = id % 10, so ~50 docs per category
|
||||
auto category_fid = schema->AddDebugField("category", DataType::INT32);
|
||||
schema->set_primary_field_id(int64_fid);
|
||||
|
||||
auto raw_data = DataGen(schema, N, 42, 0, 1, array_len);
|
||||
|
||||
// Customize int_array data: doc i has elements [i*3+1, i*3+2, i*3+3]
|
||||
for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) {
|
||||
auto* field_data = raw_data.raw_->mutable_fields_data(i);
|
||||
if (field_data->field_id() == int_array_fid.get()) {
|
||||
field_data->mutable_scalars()
|
||||
->mutable_array_data()
|
||||
->mutable_data()
|
||||
->Clear();
|
||||
|
||||
for (int row = 0; row < N; row++) {
|
||||
auto* array_data = field_data->mutable_scalars()
|
||||
->mutable_array_data()
|
||||
->mutable_data()
|
||||
->Add();
|
||||
for (int elem = 0; elem < array_len; elem++) {
|
||||
int value = row * array_len + elem + 1;
|
||||
array_data->mutable_int_data()->mutable_data()->Add(value);
|
||||
}
|
||||
}
|
||||
}
|
||||
// Set category = id % 10
|
||||
if (field_data->field_id() == category_fid.get()) {
|
||||
field_data->mutable_scalars()
|
||||
->mutable_int_data()
|
||||
->mutable_data()
|
||||
->Clear();
|
||||
for (int row = 0; row < N; row++) {
|
||||
field_data->mutable_scalars()
|
||||
->mutable_int_data()
|
||||
->mutable_data()
|
||||
->Add(row % 10);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto segment = CreateSealedWithFieldDataLoaded(schema, raw_data);
|
||||
|
||||
// Build and load vector index
|
||||
auto array_vec_values = raw_data.get_col<VectorFieldProto>(vec_fid);
|
||||
std::vector<float> vector_data(dim * N * array_len);
|
||||
for (int i = 0; i < N; i++) {
|
||||
const auto& float_vec = array_vec_values[i].float_vector().data();
|
||||
for (int j = 0; j < array_len * dim; j++) {
|
||||
vector_data[i * array_len * dim + j] = float_vec[j];
|
||||
}
|
||||
}
|
||||
|
||||
auto indexing = GenVecIndexing(N * array_len,
|
||||
dim,
|
||||
vector_data.data(),
|
||||
knowhere::IndexEnum::INDEX_HNSW);
|
||||
LoadIndexInfo load_index_info;
|
||||
load_index_info.field_id = vec_fid.get();
|
||||
load_index_info.index_params = GenIndexParams(indexing.get());
|
||||
load_index_info.cache_index =
|
||||
CreateTestCacheIndex("test", std::move(indexing));
|
||||
load_index_info.index_params["metric_type"] = knowhere::metric::L2;
|
||||
load_index_info.field_type = DataType::VECTOR_ARRAY;
|
||||
load_index_info.element_type = DataType::VECTOR_FLOAT;
|
||||
segment->LoadIndex(load_index_info);
|
||||
|
||||
int topK = 5;
|
||||
int group_size = 3;
|
||||
|
||||
ScopedSchemaHandle handle(*schema);
|
||||
std::string expr = "element_filter(structA, 2000 > $[price_array] > 100)";
|
||||
std::string search_params = R"({"ef": 50})";
|
||||
|
||||
auto plan_bytes = handle.ParseGroupBySearch(expr,
|
||||
"structA[array_vec]",
|
||||
topK,
|
||||
knowhere::metric::L2,
|
||||
search_params,
|
||||
category_fid.get(),
|
||||
group_size);
|
||||
auto plan =
|
||||
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
|
||||
ASSERT_NE(plan, nullptr);
|
||||
|
||||
auto num_queries = 1;
|
||||
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024, true);
|
||||
auto ph_group =
|
||||
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
|
||||
|
||||
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
|
||||
|
||||
ASSERT_NE(search_result, nullptr);
|
||||
ASSERT_TRUE(search_result->group_by_values_.has_value())
|
||||
<< "Group by values should be present";
|
||||
ASSERT_FALSE(search_result->element_level_)
|
||||
<< "Group by should return row-level results";
|
||||
|
||||
auto& group_by_values = search_result->group_by_values_.value();
|
||||
|
||||
ASSERT_FALSE(search_result->seg_offsets_.empty())
|
||||
<< "Should have search results";
|
||||
ASSERT_LE(search_result->seg_offsets_.size(),
|
||||
static_cast<size_t>(topK * group_size))
|
||||
<< "Should not exceed topK * group_size results";
|
||||
|
||||
// Verify: no duplicate row_offsets in results.
|
||||
// Without deduplication, the same document could appear multiple times
|
||||
// because different elements from one document may all be close to the
|
||||
// query vector.
|
||||
std::unordered_set<int64_t> seen_docs;
|
||||
std::unordered_map<int32_t, int> group_counts;
|
||||
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
|
||||
int64_t doc_id = search_result->seg_offsets_[i];
|
||||
|
||||
ASSERT_TRUE(seen_docs.insert(doc_id).second)
|
||||
<< "Duplicate doc_id " << doc_id
|
||||
<< " in results: same document should not appear more than once";
|
||||
|
||||
if (i < group_by_values.size() && group_by_values[i].has_value()) {
|
||||
if (std::holds_alternative<int32_t>(group_by_values[i].value())) {
|
||||
int32_t group_val =
|
||||
std::get<int32_t>(group_by_values[i].value());
|
||||
|
||||
ASSERT_EQ(group_val, doc_id % 10)
|
||||
<< "Group value should equal category (doc_id % 10)";
|
||||
|
||||
group_counts[group_val]++;
|
||||
ASSERT_LE(group_counts[group_val], group_size)
|
||||
<< "Each group should have at most group_size results";
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ASSERT_LE(group_counts.size(), static_cast<size_t>(topK))
|
||||
<< "Should have at most topK distinct groups";
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user