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
236193int64_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_ );
0 commit comments