Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions lucene/CHANGES.txt
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,8 @@ Improvements

* GITHUB#15574, GITHUB#15995: Introduce TopGroupsCollectorManager to parallelize search when using TopGroupsCollector. (Binlong Gao)

* GITHUB#15565: Introduce AllGroupHeadsCollectorManager to parallelize search when using AllGroupHeadsCollector. (Binlong Gao)

* GITHUB#15989: DocValuesRangeIterator always tries to use skipper-based block iteration
as its approximation. (Alan Woodward)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -232,6 +232,20 @@ protected void setNextReader(LeafReaderContext ctx) throws IOException {
* @throws IOException If I/O related errors occur
*/
protected abstract void updateDocHead(int doc) throws IOException;

/**
* Returns the sort values for this group head.
*
* @return the sort values, or null if not stored
*/
protected abstract Object[] getSortValues();

/**
* Returns the field comparators used to determine the group head ordering.
*
* @return the comparators, one per sort field
*/
protected abstract FieldComparator[] getComparators();
Comment thread
javanna marked this conversation as resolved.
}

/** General implementation using a {@link FieldComparator} to select the group head */
Expand All @@ -252,6 +266,7 @@ private static class SortingGroupHead<T> extends GroupHead<T> {

final FieldComparator[] comparators;
final LeafFieldComparator[] leafComparators;
final Object[] sortValues;

protected SortingGroupHead(
Sort sort, T groupValue, int doc, LeafReaderContext context, Scorable scorer)
Expand All @@ -260,12 +275,14 @@ protected SortingGroupHead(
final SortField[] sortFields = sort.getSort();
comparators = new FieldComparator[sortFields.length];
leafComparators = new LeafFieldComparator[sortFields.length];
sortValues = new Object[sortFields.length];
for (int i = 0; i < sortFields.length; i++) {
comparators[i] = sortFields[i].getComparator(1, Pruning.NONE);
leafComparators[i] = comparators[i].getLeafComparator(context);
leafComparators[i].setScorer(scorer);
leafComparators[i].copy(0, doc);
leafComparators[i].setBottom(0);
sortValues[i] = comparators[i].value(0);
}
}

Expand All @@ -291,12 +308,23 @@ public int compare(int compIDX, int doc) throws IOException {

@Override
public void updateDocHead(int doc) throws IOException {
for (LeafFieldComparator comparator : leafComparators) {
comparator.copy(0, doc);
comparator.setBottom(0);
for (int i = 0; i < leafComparators.length; i++) {
leafComparators[i].copy(0, doc);
leafComparators[i].setBottom(0);
sortValues[i] = comparators[i].value(0);
}
this.doc = doc + docBase;
}

@Override
protected Object[] getSortValues() {
return sortValues;
}

@Override
protected FieldComparator[] getComparators() {
return comparators;
}
}

/** Specialized implementation for sorting by score */
Expand All @@ -317,12 +345,14 @@ private static class ScoringGroupHead<T> extends GroupHead<T> {

private Scorable scorer;
private float topScore;
private final Object[] sortValues;

protected ScoringGroupHead(Scorable scorer, T groupValue, int doc, int docBase)
throws IOException {
super(groupValue, doc, docBase);
this.scorer = scorer;
this.topScore = scorer.score();
this.sortValues = new Object[] {topScore};
}

@Override
Expand All @@ -344,6 +374,17 @@ protected int compare(int compIDX, int doc) throws IOException {
@Override
protected void updateDocHead(int doc) throws IOException {
this.doc = doc + docBase;
sortValues[0] = topScore;
}

@Override
protected Object[] getSortValues() {
return sortValues;
}

@Override
protected FieldComparator[] getComparators() {
return null;
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
/*
* Licensed to the Apache Software Foundation (ASF) under one or more
* contributor license agreements. See the NOTICE file distributed with
* this work for additional information regarding copyright ownership.
* The ASF licenses this file to You 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.
*/
package org.apache.lucene.search.grouping;

import java.io.IOException;
import java.util.Collection;
import java.util.HashMap;
import java.util.Map;
import java.util.function.Supplier;
import org.apache.lucene.index.IndexReader;
import org.apache.lucene.search.CollectorManager;
import org.apache.lucene.search.FieldComparator;
import org.apache.lucene.search.Sort;
import org.apache.lucene.search.SortField;
import org.apache.lucene.util.Bits;
import org.apache.lucene.util.FixedBitSet;

/**
* A {@link CollectorManager} implementation for {@link AllGroupHeadsCollector} that collects the
* most relevant document (group head) for each group across multiple segments and merges the
* per-segment results into a single {@link GroupHeadsResult}.
*
* <p>Example usage:
*
* <pre class="prettyprint">
* IndexSearcher searcher = ...; // your IndexSearcher
* AllGroupHeadsCollectorManager&lt;BytesRef&gt; manager =
* new AllGroupHeadsCollectorManager&lt;&gt;(
* () -&gt; new TermGroupSelector("category"), Sort.RELEVANCE);
* GroupHeadsResult result = searcher.search(new MatchAllDocsQuery(), manager);
* Bits groupHeadsBits = result.retrieveGroupHeads(searcher.getIndexReader().maxDoc());
* </pre>
*
* @param <T> the type of the group value
* @lucene.experimental
*/
public class AllGroupHeadsCollectorManager<T>
implements CollectorManager<
AllGroupHeadsCollector<T>, AllGroupHeadsCollectorManager.GroupHeadsResult> {

/** Holds the merged group heads and provides access as an {@code int[]} or {@link Bits}. */
public static class GroupHeadsResult {
private final int[] groupHeads;

private GroupHeadsResult(int[] groupHeads) {
this.groupHeads = groupHeads;
}

/** Returns the group head document IDs as an array. */
public int[] retrieveGroupHeads() {
return groupHeads;
}

/**
* Returns the group head document IDs as a {@link Bits} set of size {@code maxDoc}, suitable
* for use as a filter.
*
* @param maxDoc The maxDoc of the top level {@link IndexReader}.
*/
public Bits retrieveGroupHeads(int maxDoc) {
FixedBitSet result = new FixedBitSet(maxDoc);
for (int docId : groupHeads) {
result.set(docId);
}
return result;
}
}

private static final class GroupHeadWithValues {
int doc;
final Object[] sortValues;

GroupHeadWithValues(int doc, Object[] sortValues) {
this.doc = doc;
this.sortValues = sortValues;
}
}

private final Supplier<GroupSelector<T>> groupSelectorFactory;
private final Sort sortWithinGroup;

/**
* Creates a new AllGroupHeadsCollectorManager.
*
* @param groupSelectorFactory factory to create group selectors for each collector
* @param sortWithinGroup the sort to use within each group to determine the group head
*/
public AllGroupHeadsCollectorManager(
Supplier<GroupSelector<T>> groupSelectorFactory, Sort sortWithinGroup) {
this.groupSelectorFactory = groupSelectorFactory;
this.sortWithinGroup = sortWithinGroup;
}

@Override
public AllGroupHeadsCollector<T> newCollector() throws IOException {
return AllGroupHeadsCollector.newCollector(groupSelectorFactory.get(), sortWithinGroup);
}

@Override
public GroupHeadsResult reduce(Collection<AllGroupHeadsCollector<T>> collectors) {
Map<T, GroupHeadWithValues> mergedHeads = new HashMap<>();
SortField[] sortFields = sortWithinGroup.getSort();

for (AllGroupHeadsCollector<T> collector : collectors) {
mergeCollectorHeads(collector, mergedHeads, sortFields);
}

return new GroupHeadsResult(mergedHeads.values().stream().mapToInt(h -> h.doc).toArray());
}

private void mergeCollectorHeads(
AllGroupHeadsCollector<T> collector,
Map<T, GroupHeadWithValues> mergedHeads,
SortField[] sortFields) {
for (AllGroupHeadsCollector.GroupHead<T> head : collector.getCollectedGroupHeads()) {
Object[] sortValues = head.getSortValues();
GroupHeadWithValues existing = mergedHeads.get(head.groupValue);
if (existing == null || isCompetitive(head, sortValues, existing, sortFields)) {
mergedHeads.put(head.groupValue, new GroupHeadWithValues(head.doc, sortValues));
}
}
}

@SuppressWarnings({"rawtypes"})
private boolean isCompetitive(
AllGroupHeadsCollector.GroupHead<T> head,
Object[] sortValues,
GroupHeadWithValues existing,
SortField[] sortFields) {
FieldComparator[] comparators = head.getComparators();
int cmp;
if (sortWithinGroup.equals(Sort.RELEVANCE)) {
cmp = Float.compare((float) sortValues[0], (float) existing.sortValues[0]);
return cmp > 0 || (cmp == 0 && head.doc < existing.doc);
} else {
cmp = 0;
for (int i = 0; i < sortFields.length; i++) {
@SuppressWarnings({"unchecked"})
int c = comparators[i].compareValues(sortValues[i], existing.sortValues[i]);
c = sortFields[i].getReverse() ? -c : c;
if (c != 0) {
cmp = c;
break;
}
}
return cmp < 0 || (cmp == 0 && head.doc < existing.doc);
}
}
}
Loading
Loading