Skip to content

Commit 89afaf1

Browse files
authored
Skip two phase rescore for sort query (#1898)
* skip two phase rescore for sort query Signed-off-by: Liyun Xiu <xiliyun@amazon.com> * ad changelog Signed-off-by: Liyun Xiu <xiliyun@amazon.com> --------- Signed-off-by: Liyun Xiu <xiliyun@amazon.com>
1 parent 1778358 commit 89afaf1

3 files changed

Lines changed: 95 additions & 0 deletions

File tree

CHANGELOG.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
1313
* Add `previous_score_field` to `by_field` rerank processor to avoid overwriting existing document fields ([#1880](https://github.com/opensearch-project/neural-search/pull/1880)) (OpenSearch [#21440](https://github.com/opensearch-project/OpenSearch/issues/21440))
1414
* [Hybrid Query] Block hybrid query execution with `search_type=dfs_query_then_fetch` ([#1873](https://github.com/opensearch-project/neural-search/pull/1873))
1515
* [Hybrid Query] Fix `hybrid_score_explanation` returning a single normalization block for hybrid queries on indices that contain a nested field ([#1876](https://github.com/opensearch-project/neural-search/pull/1876))
16+
* [Two-phase] Skip two phase rescore for sort query ([#1898](https://github.com/opensearch-project/neural-search/pull/1898))
1617

1718
### Infrastructure
1819
* [Hybrid Query] Fix flaky `HybridQueryExplainIT` explanation test by asserting per-document explanation trees keyed by `_id` instead of hit position ([#1899](https://github.com/opensearch-project/neural-search/pull/1899))

src/main/java/org/opensearch/neuralsearch/processor/NeuralSparseTwoPhaseProcessor.java

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,11 @@
2727
import org.opensearch.search.pipeline.SearchRequestProcessor;
2828
import org.opensearch.search.rescore.QueryRescorerBuilder;
2929
import org.opensearch.search.rescore.RescorerBuilder;
30+
import org.opensearch.search.sort.ScoreSortBuilder;
31+
import org.opensearch.search.sort.SortBuilder;
32+
import org.opensearch.search.sort.SortOrder;
3033

34+
import java.util.List;
3135
import java.util.Locale;
3236
import java.util.Map;
3337
import java.util.Objects;
@@ -103,6 +107,13 @@ public SearchRequest processRequest(final SearchRequest request) {
103107
if (!enabled || pruneRatio == 0f) {
104108
return request;
105109
}
110+
// Two-phase rescore is incompatible with explicit sort (other than _score DESC).
111+
// OpenSearch rejects sort + rescore in DefaultSearchContext.preProcess(), so when the
112+
// request already has a non-_score-DESC sort, skip the optimization and let the full
113+
// neural_sparse query run as a single phase. Correctness over latency.
114+
if (hasIncompatibleSort(request.source())) {
115+
return request;
116+
}
106117
QueryBuilder queryBuilder = request.source().query();
107118
// Collect the nested NeuralSparseQueryBuilder in the whole query.
108119
Multimap<AbstractNeuralQueryBuilder<?>, Float> queryBuilderMap = collectNeuralQueryBuilderWithSparseEmbedding(
@@ -139,6 +150,28 @@ private QueryBuilder getNestedQueryBuilderFromNeuralSparseQueryBuilderMap(
139150
return boolQueryBuilder;
140151
}
141152

153+
/**
154+
* Returns true when the request specifies a sort that OpenSearch will treat as an explicit
155+
* sort at execution time, in which case adding a rescore would cause
156+
* {@code DefaultSearchContext.preProcess()} to throw
157+
* "Cannot use [sort] option in conjunction with [rescore]".
158+
*
159+
* Mirrors the optimization in {@code SortBuilder.buildSort}: a single {@code _score DESC}
160+
* is collapsed to "no sort" and is therefore compatible. Anything else (a field sort, a
161+
* non-default _score order, or a multi-element sort list) is incompatible.
162+
*/
163+
static boolean hasIncompatibleSort(final SearchSourceBuilder searchSourceBuilder) {
164+
List<SortBuilder<?>> sorts = searchSourceBuilder.sorts();
165+
if (sorts == null || sorts.isEmpty()) {
166+
return false;
167+
}
168+
if (sorts.size() == 1) {
169+
SortBuilder<?> only = sorts.get(0);
170+
return !(only instanceof ScoreSortBuilder) || only.order() != SortOrder.DESC;
171+
}
172+
return true;
173+
}
174+
142175
private float getOriginQueryWeightAfterRescore(final SearchSourceBuilder searchSourceBuilder) {
143176
if (Objects.isNull(searchSourceBuilder.rescores())) {
144177
return 1.0f;

src/test/java/org/opensearch/neuralsearch/processor/NeuralSparseTwoPhaseProcessorTests.java

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,9 @@
1818
import org.opensearch.neuralsearch.util.prune.PruneUtils;
1919
import org.opensearch.search.builder.SearchSourceBuilder;
2020
import org.opensearch.search.rescore.QueryRescorerBuilder;
21+
import org.opensearch.search.sort.ScoreSortBuilder;
22+
import org.opensearch.search.sort.SortBuilders;
23+
import org.opensearch.search.sort.SortOrder;
2124
import org.opensearch.test.OpenSearchTestCase;
2225

2326
import java.util.Collections;
@@ -256,6 +259,64 @@ public void testProcessRequest_whenTwoPhaseEnabledWithNeuralQueryNonSparseEmbedd
256259
assertNull(searchRequest.source().rescores());
257260
}
258261

262+
public void testProcessRequest_whenSortByField_thenSkipRescore() throws Exception {
263+
NeuralSparseTwoPhaseProcessor.Factory factory = new NeuralSparseTwoPhaseProcessor.Factory();
264+
NeuralSparseQueryBuilder neuralQueryBuilder = new NeuralSparseQueryBuilder();
265+
SearchRequest searchRequest = new SearchRequest();
266+
searchRequest.source(
267+
new SearchSourceBuilder().query(neuralQueryBuilder).sort(SortBuilders.fieldSort("entityId").order(SortOrder.ASC))
268+
);
269+
NeuralSparseTwoPhaseProcessor processor = createTestProcessor(factory, 0.5f, true, 4.0f, 10000);
270+
processor.processRequest(searchRequest);
271+
assertNull(searchRequest.source().rescores());
272+
}
273+
274+
public void testProcessRequest_whenSortByScoreDescAndField_thenSkipRescore() throws Exception {
275+
NeuralSparseTwoPhaseProcessor.Factory factory = new NeuralSparseTwoPhaseProcessor.Factory();
276+
NeuralSparseQueryBuilder neuralQueryBuilder = new NeuralSparseQueryBuilder();
277+
SearchRequest searchRequest = new SearchRequest();
278+
searchRequest.source(
279+
new SearchSourceBuilder().query(neuralQueryBuilder)
280+
.sort(new ScoreSortBuilder())
281+
.sort(SortBuilders.fieldSort("entityId").order(SortOrder.ASC))
282+
);
283+
NeuralSparseTwoPhaseProcessor processor = createTestProcessor(factory, 0.5f, true, 4.0f, 10000);
284+
processor.processRequest(searchRequest);
285+
assertNull(searchRequest.source().rescores());
286+
}
287+
288+
public void testProcessRequest_whenSortByScoreDescOnly_thenAddRescore() throws Exception {
289+
NeuralSparseTwoPhaseProcessor.Factory factory = new NeuralSparseTwoPhaseProcessor.Factory();
290+
NeuralSparseQueryBuilder neuralQueryBuilder = new NeuralSparseQueryBuilder();
291+
SearchRequest searchRequest = new SearchRequest();
292+
searchRequest.source(new SearchSourceBuilder().query(neuralQueryBuilder).sort(new ScoreSortBuilder()));
293+
NeuralSparseTwoPhaseProcessor processor = createTestProcessor(factory, 0.5f, true, 4.0f, 10000);
294+
processor.processRequest(searchRequest);
295+
assertNotNull(searchRequest.source().rescores());
296+
}
297+
298+
public void testProcessRequest_whenSortByScoreAsc_thenSkipRescore() throws Exception {
299+
NeuralSparseTwoPhaseProcessor.Factory factory = new NeuralSparseTwoPhaseProcessor.Factory();
300+
NeuralSparseQueryBuilder neuralQueryBuilder = new NeuralSparseQueryBuilder();
301+
SearchRequest searchRequest = new SearchRequest();
302+
searchRequest.source(new SearchSourceBuilder().query(neuralQueryBuilder).sort(new ScoreSortBuilder().order(SortOrder.ASC)));
303+
NeuralSparseTwoPhaseProcessor processor = createTestProcessor(factory, 0.5f, true, 4.0f, 10000);
304+
processor.processRequest(searchRequest);
305+
assertNull(searchRequest.source().rescores());
306+
}
307+
308+
public void testProcessRequest_whenSortByScoreDescWithTrackScores_thenAddRescore() throws Exception {
309+
NeuralSparseTwoPhaseProcessor.Factory factory = new NeuralSparseTwoPhaseProcessor.Factory();
310+
NeuralSparseQueryBuilder neuralQueryBuilder = new NeuralSparseQueryBuilder();
311+
SearchRequest searchRequest = new SearchRequest();
312+
searchRequest.source(
313+
new SearchSourceBuilder().query(neuralQueryBuilder).sort(new ScoreSortBuilder()).trackScores(true)
314+
);
315+
NeuralSparseTwoPhaseProcessor processor = createTestProcessor(factory, 0.5f, true, 4.0f, 10000);
316+
processor.processRequest(searchRequest);
317+
assertNotNull(searchRequest.source().rescores());
318+
}
319+
259320
public void testType() throws Exception {
260321
NeuralSparseTwoPhaseProcessor.Factory factory = new NeuralSparseTwoPhaseProcessor.Factory();
261322
NeuralSparseTwoPhaseProcessor processor = createTestProcessor(factory);

0 commit comments

Comments
 (0)