Skip to content

Commit fdf09c8

Browse files
committed
Simplify
1 parent d1debd8 commit fdf09c8

16 files changed

Lines changed: 708 additions & 1179 deletions

File tree

internal/core/src/common/Consts.h

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,6 @@ const char VEC_OPT_FIELDS[] = "opt_fields";
5151
const char PAGE_RETAIN_ORDER[] = "page_retain_order";
5252
const char TEXT_LOG_ROOT_PATH[] = "text_log";
5353
const char ITERATIVE_FILTER[] = "iterative_filter";
54-
const char GROUP_BY_REFILL[] = "group_by_refill";
5554
const char HINTS[] = "hints";
5655
// json stats related
5756
const char JSON_KEY_INDEX_LOG_ROOT_PATH[] = "json_key_index_log";

internal/core/src/exec/operator/GroupByNode.cpp

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -107,10 +107,9 @@ PhyGroupByNode::GetOutput() {
107107
search_result.unity_topK_);
108108
search_result.topk_per_nq_prefix_sum_.resize(
109109
search_result.total_nq_ + 1);
110-
std::partial_sum(
111-
topks.begin(),
112-
topks.end(),
113-
search_result.topk_per_nq_prefix_sum_.begin() + 1);
110+
std::partial_sum(topks.begin(),
111+
topks.end(),
112+
search_result.topk_per_nq_prefix_sum_.begin() + 1);
114113
}
115114
}
116115
tracer::AddEvent(

internal/core/src/exec/operator/groupby/SearchGroupByOperator.cpp

Lines changed: 96 additions & 436 deletions
Large diffs are not rendered by default.

internal/core/src/exec/operator/groupby/SearchGroupByOperator.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -361,7 +361,7 @@ struct GroupByMap {
361361
bool strict_group_size = false)
362362
: group_capacity_(group_capacity),
363363
group_size_(group_size),
364-
strict_group_size_(strict_group_size){};
364+
strict_group_size_(strict_group_size) {};
365365
bool
366366
IsGroupResEnough() {
367367
bool enough = false;

internal/core/src/query/PlanProto.cpp

Lines changed: 1 addition & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -91,28 +91,7 @@ ProtoParser::PlanNodeFromProto(const planpb::PlanNode& plan_node_proto) {
9191
}
9292
}
9393

94-
if (search_info.search_params_.contains(GROUP_BY_REFILL)) {
95-
auto& refill_param = search_info.search_params_[GROUP_BY_REFILL];
96-
if (refill_param.is_boolean()) {
97-
search_info.proxy_group_by_refill_ = refill_param.get<bool>();
98-
} else if (refill_param.is_string()) {
99-
auto refill_param_str = refill_param.get<std::string>();
100-
if (refill_param_str == "true" || refill_param_str == "True") {
101-
search_info.proxy_group_by_refill_ = true;
102-
} else if (refill_param_str == "false" ||
103-
refill_param_str == "False") {
104-
search_info.proxy_group_by_refill_ = false;
105-
} else {
106-
ThrowInfo(ConfigInvalid,
107-
"group_by_refill: {} not supported",
108-
refill_param);
109-
}
110-
} else {
111-
ThrowInfo(ConfigInvalid,
112-
"group_by_refill: {} not supported",
113-
refill_param);
114-
}
115-
}
94+
search_info.proxy_group_by_refill_ = query_info_proto.group_by_refill();
11695

11796
if (query_info_proto.bm25_avgdl() > 0) {
11897
search_info.search_params_[knowhere::meta::BM25_AVGDL] =

internal/core/src/segcore/reduce/GroupReduce.cpp

Lines changed: 10 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
// or implied. See the License for the specific language governing permissions and limitations under the License
1111

1212
#include "GroupReduce.h"
13-
#include <numeric>
13+
#include <unordered_map>
1414

1515
#include "common/Consts.h"
1616
#include "fmt/format.h"
@@ -186,51 +186,8 @@ GroupReduceHelper::FilterInvalidSearchResult(SearchResult* search_result) {
186186
search_result->seg_offsets_.size(),
187187
"group filter invalid result");
188188

189-
std::vector<int64_t> real_topks(nq, 0);
190-
uint32_t valid_index = 0;
191-
auto segment = static_cast<SegmentInterface*>(search_result->segment_);
192-
auto& offsets = search_result->seg_offsets_;
193-
auto& distances = search_result->distances_;
194189
auto& group_by_values = search_result->group_by_values_.value();
195-
int segment_row_count = segment->get_row_count();
196-
197-
for (auto i = 0; i < nq; ++i) {
198-
for (auto j = 0; j < topK; ++j) {
199-
auto index = i * topK + j;
200-
if (offsets[index] == INVALID_SEG_OFFSET) {
201-
continue;
202-
}
203-
AssertInfo(
204-
0 <= offsets[index] && offsets[index] < segment_row_count,
205-
fmt::format("invalid offset {}, segment {} with "
206-
"rows num {}, data or index corruption",
207-
offsets[index],
208-
segment->get_segment_id(),
209-
segment_row_count));
210-
if (valid_index != index) {
211-
offsets[valid_index] = offsets[index];
212-
distances[valid_index] = distances[index];
213-
group_by_values[valid_index] =
214-
std::move(group_by_values[index]);
215-
if (search_result->element_level_) {
216-
search_result->element_indices_[valid_index] =
217-
search_result->element_indices_[index];
218-
}
219-
}
220-
valid_index++;
221-
real_topks[i]++;
222-
}
223-
}
224-
offsets.resize(valid_index);
225-
distances.resize(valid_index);
226-
group_by_values.resize(valid_index);
227-
if (search_result->element_level_) {
228-
search_result->element_indices_.resize(valid_index);
229-
}
230-
search_result->topk_per_nq_prefix_sum_.resize(nq + 1);
231-
std::partial_sum(real_topks.begin(),
232-
real_topks.end(),
233-
search_result->topk_per_nq_prefix_sum_.begin() + 1);
190+
CompactSearchResult(search_result, &group_by_values);
234191
}
235192

236193
int64_t
@@ -243,7 +200,6 @@ GroupReduceHelper::ReduceSearchResultForOneNQ(int64_t qi,
243200
heap;
244201
pk_set_.clear();
245202
element_result_set_.clear();
246-
group_by_val_count_.clear();
247203
pairs_.clear();
248204

249205
pairs_.reserve(num_segments_);
@@ -266,27 +222,17 @@ GroupReduceHelper::ReduceSearchResultForOneNQ(int64_t qi,
266222
}
267223

268224
int64_t dup_cnt = 0;
225+
auto start = offset;
269226
auto group_size = int64_t(1);
270227
for (auto search_result : search_results_) {
271228
if (search_result->group_size_.has_value()) {
272229
group_size = search_result->group_size_.value();
273230
break;
274231
}
275232
}
276-
auto selected_groups_full = [&]() {
277-
if (static_cast<int64_t>(group_by_val_count_.size()) < topk) {
278-
return false;
279-
}
280-
return std::all_of(group_by_val_count_.begin(),
281-
group_by_val_count_.end(),
282-
[group_size](const auto& item) {
283-
return item.second >= group_size;
284-
});
285-
};
286-
while (!heap.empty()) {
287-
if (selected_groups_full()) {
288-
break;
289-
}
233+
auto result_limit = topk * group_size;
234+
std::unordered_map<GroupByValueType, int64_t> group_counts;
235+
while (offset - start < result_limit && !heap.empty()) {
290236
auto pilot = heap.top();
291237
heap.pop();
292238

@@ -305,16 +251,15 @@ GroupReduceHelper::ReduceSearchResultForOneNQ(int64_t qi,
305251
group_by_values.size());
306252
auto& group_by_value = group_by_values[pilot->offset_];
307253

308-
auto group_iter = group_by_val_count_.find(group_by_value);
309-
auto is_new_group = group_iter == group_by_val_count_.end();
254+
auto group_iter = group_counts.find(group_by_value);
255+
auto is_new_group = group_iter == group_counts.end();
310256
auto group_is_full = !is_new_group && group_iter->second >= group_size;
311257
auto group_capacity_is_full =
312-
is_new_group &&
313-
static_cast<int64_t>(group_by_val_count_.size()) >= topk;
258+
is_new_group && static_cast<int64_t>(group_counts.size()) >= topk;
314259

315260
if (!group_is_full && !group_capacity_is_full) {
316261
if (TryAcceptSearchResult(*pilot)) {
317-
auto& group_count = group_by_val_count_[group_by_value];
262+
auto& group_count = group_counts[group_by_value];
318263
group_count++;
319264
pilot->search_result_->result_offsets_.push_back(offset++);
320265
final_search_records_[index][qi].push_back(pilot->offset_);

internal/core/src/segcore/reduce/GroupReduce.h

Lines changed: 0 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -9,8 +9,6 @@
99
// is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express
1010
// or implied. See the License for the specific language governing permissions and limitations under the License
1111
#pragma once
12-
#include <unordered_map>
13-
1412
#include "Reduce.h"
1513
#include "common/QueryResult.h"
1614
#include "query/PlanImpl.h"
@@ -52,9 +50,6 @@ class GroupReduceHelper : public ReduceHelper {
5250
int64_t nq_end,
5351
std::unique_ptr<milvus::proto::schema::SearchResultData>&
5452
search_res_data) override;
55-
56-
private:
57-
std::unordered_map<milvus::GroupByValueType, int64_t> group_by_val_count_{};
5853
};
5954

6055
} // namespace milvus::segcore

internal/core/src/segcore/reduce/Reduce.cpp

Lines changed: 28 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -99,6 +99,13 @@ ReduceHelper::CheckElementIndicesSize(const SearchResult* search_result,
9999

100100
void
101101
ReduceHelper::FilterInvalidSearchResult(SearchResult* search_result) {
102+
CompactSearchResult(search_result);
103+
}
104+
105+
void
106+
ReduceHelper::CompactSearchResult(
107+
SearchResult* search_result,
108+
std::vector<GroupByValueType>* companion_values) {
102109
auto nq = search_result->total_nq_;
103110
auto topK = search_result->unity_topK_;
104111
AssertInfo(search_result->seg_offsets_.size() == nq * topK,
@@ -117,6 +124,13 @@ ReduceHelper::FilterInvalidSearchResult(SearchResult* search_result) {
117124
auto& offsets = search_result->seg_offsets_;
118125
auto& distances = search_result->distances_;
119126

127+
if (companion_values != nullptr) {
128+
AssertInfo(companion_values->size() == offsets.size(),
129+
"wrong companion values size, size = {}, expected size = {}",
130+
companion_values->size(),
131+
offsets.size());
132+
}
133+
120134
int segment_row_count = segment->get_row_count();
121135
//1. for sealed segment, segment_row_count will not change as delete records will take effect as bitset
122136
//2. for growing segment, segment_row_count is the minimum position acknowledged, which will only increase after
@@ -133,18 +147,27 @@ ReduceHelper::FilterInvalidSearchResult(SearchResult* search_result) {
133147
segment->get_segment_id(),
134148
segment_row_count));
135149
real_topks[i]++;
136-
offsets[valid_index] = offsets[index];
137-
distances[valid_index] = distances[index];
138-
if (search_result->element_level_) {
139-
search_result->element_indices_[valid_index] =
140-
search_result->element_indices_[index];
150+
if (valid_index != index) {
151+
offsets[valid_index] = offsets[index];
152+
distances[valid_index] = distances[index];
153+
if (companion_values != nullptr) {
154+
(*companion_values)[valid_index] =
155+
std::move((*companion_values)[index]);
156+
}
157+
if (search_result->element_level_) {
158+
search_result->element_indices_[valid_index] =
159+
search_result->element_indices_[index];
160+
}
141161
}
142162
valid_index++;
143163
}
144164
}
145165
}
146166
offsets.resize(valid_index);
147167
distances.resize(valid_index);
168+
if (companion_values != nullptr) {
169+
companion_values->resize(valid_index);
170+
}
148171
if (search_result->element_level_) {
149172
search_result->element_indices_.resize(valid_index);
150173
}

internal/core/src/segcore/reduce/Reduce.h

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -91,6 +91,11 @@ class ReduceHelper {
9191
virtual void
9292
FilterInvalidSearchResult(SearchResult* search_result);
9393

94+
void
95+
CompactSearchResult(
96+
SearchResult* search_result,
97+
std::vector<GroupByValueType>* companion_values = nullptr);
98+
9499
void
95100
RefreshSearchResults();
96101

internal/proxy/search_util.go

Lines changed: 0 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@ package proxy
22

33
import (
44
"context"
5-
"encoding/json"
65
"fmt"
76
"regexp"
87
"strconv"
@@ -577,57 +576,6 @@ func parseGroupByInfo(searchParamsPair []*commonpb.KeyValuePair, schema *schemap
577576
return ret, nil
578577
}
579578

580-
func setGroupByRefillOnQueryInfo(queryInfo *planpb.QueryInfo, enabled bool) error {
581-
params := map[string]interface{}{}
582-
if queryInfo.GetSearchParams() != "" {
583-
if err := json.Unmarshal([]byte(queryInfo.GetSearchParams()), &params); err != nil {
584-
return err
585-
}
586-
}
587-
if queryInfo.GetGroupByFieldId() <= 0 {
588-
if _, ok := params[GroupByRefillKey]; !ok {
589-
return nil
590-
}
591-
delete(params, GroupByRefillKey)
592-
bs, err := json.Marshal(params)
593-
if err != nil {
594-
return err
595-
}
596-
queryInfo.SearchParams = string(bs)
597-
return nil
598-
}
599-
normalized := false
600-
if raw, ok := params[GroupByRefillKey]; ok {
601-
switch value := raw.(type) {
602-
case bool:
603-
case string:
604-
parsed, err := strconv.ParseBool(value)
605-
if err != nil {
606-
return merr.WrapErrParameterInvalid("true or false", value,
607-
"value for group_by_refill is invalid")
608-
}
609-
params[GroupByRefillKey] = parsed
610-
normalized = true
611-
default:
612-
return merr.WrapErrParameterInvalid("true or false", fmt.Sprint(value),
613-
"value for group_by_refill is invalid")
614-
}
615-
}
616-
if enabled {
617-
params[GroupByRefillKey] = true
618-
normalized = true
619-
}
620-
if !normalized {
621-
return nil
622-
}
623-
bs, err := json.Marshal(params)
624-
if err != nil {
625-
return err
626-
}
627-
queryInfo.SearchParams = string(bs)
628-
return nil
629-
}
630-
631579
// parseRankParams get limit and offset from rankParams, both are optional.
632580
func parseRankParams(rankParamsPair []*commonpb.KeyValuePair, schema *schemapb.CollectionSchema, largeTopKEnabled bool) (*rankParams, error) {
633581
var (

0 commit comments

Comments
 (0)