kaivalnp commented on code in PR #16738: URL: https://github.com/apache/lucene/pull/16738#discussion_r4148856154
########## lucene/join/src/test/org/apache/lucene/search/join/TestParentBlockJoinFloat16KnnVectorQuery.java: ########## @@ -0,0 +1,164 @@ +/* + * 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.join; + +import static org.apache.lucene.index.VectorSimilarityFunction.COSINE; + +import java.io.IOException; +import java.util.ArrayList; +import java.util.List; +import java.util.Random; +import org.apache.lucene.document.Document; +import org.apache.lucene.document.Field; +import org.apache.lucene.document.KnnFloat16VectorField; +import org.apache.lucene.index.DirectoryReader; +import org.apache.lucene.index.IndexReader; +import org.apache.lucene.index.IndexWriter; +import org.apache.lucene.index.IndexWriterConfig; +import org.apache.lucene.index.Term; +import org.apache.lucene.index.VectorSimilarityFunction; +import org.apache.lucene.search.IndexSearcher; +import org.apache.lucene.search.Query; +import org.apache.lucene.search.TermQuery; +import org.apache.lucene.store.Directory; +import org.apache.lucene.tests.util.TestUtil; + +public class TestParentBlockJoinFloat16KnnVectorQuery + extends ParentBlockJoinKnnVectorQueryTestCase { + + @Override + Query getParentJoinKnnQuery( + String fieldName, + float[] queryVector, + Query childFilter, + int k, + BitSetProducer parentBitSet) { + return new DiversifyingChildrenFloat16KnnVectorQuery( + fieldName, fromFloat(queryVector), childFilter, k, parentBitSet); + } + + @Override + Field getKnnVectorField(String name, float[] vector) { + return new KnnFloat16VectorField(name, fromFloat(vector), VectorSimilarityFunction.EUCLIDEAN); + } + + @Override + Field getKnnVectorField( + String name, float[] vector, VectorSimilarityFunction vectorSimilarityFunction) { + return new KnnFloat16VectorField(name, fromFloat(vector), vectorSimilarityFunction); + } + + public void testVectorEncodingMismatch() throws IOException { Review Comment: Should we also add a mismatch test for `BYTE` field x `FLOAT16` query? (same for other test classes, adding a mismatch test for `FLOAT16` fields) ########## lucene/join/src/java/org/apache/lucene/search/join/DiversifyingChildrenFloat16KnnVectorQuery.java: ########## @@ -0,0 +1,176 @@ +/* + * 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.join; + +import static org.apache.lucene.search.knn.KnnSearchStrategy.Hnsw.DEFAULT; + +import java.io.IOException; +import java.util.Arrays; +import java.util.Objects; +import org.apache.lucene.index.Float16VectorValues; +import org.apache.lucene.index.LeafReaderContext; +import org.apache.lucene.index.QueryTimeout; +import org.apache.lucene.search.AcceptDocs; +import org.apache.lucene.search.DocIdSetIterator; +import org.apache.lucene.search.IndexSearcher; +import org.apache.lucene.search.KnnCollector; +import org.apache.lucene.search.KnnFloat16VectorQuery; +import org.apache.lucene.search.Query; +import org.apache.lucene.search.TopDocs; +import org.apache.lucene.search.TopDocsCollector; +import org.apache.lucene.search.VectorScorer; +import org.apache.lucene.search.knn.KnnCollectorManager; +import org.apache.lucene.search.knn.KnnSearchStrategy; +import org.apache.lucene.util.BitSet; + +/** + * kNN float16 vector query that joins matching children vector documents with their parent doc id. + * The top documents returned are the child document ids and the calculated scores. Here is how to + * use this in conjunction with {@link ToParentBlockJoinQuery}. + * + * <pre><code class="language-java"> + * Query knnQuery = new DiversifyingChildrenFloat16KnnVectorQuery(fieldName, queryVector, ...); + * // Rewrite executes kNN search and collects nearest children docIds and their scores + * Query rewrittenKnnQuery = searcher.rewrite(knnQuery); + * // Join the scored children docs with their parents and score the parents + * Query childrenToParents = new ToParentBlockJoinQuery(rewrittenKnnQuery, parentsFilter, ScoreMode.MAX); + * </code></pre> + */ +public class DiversifyingChildrenFloat16KnnVectorQuery extends KnnFloat16VectorQuery { + private static final TopDocs NO_RESULTS = TopDocsCollector.EMPTY_TOPDOCS; + + private final BitSetProducer parentsFilter; + private final Query childFilter; + private final int k; + private final short[] query; + + /** + * Create a DiversifyingChildrenFloat16KnnVectorQuery. + * + * @param field the query field + * @param query the vector query + * @param childFilter the child filter + * @param k how many parent documents to return given the matching children + * @param parentsFilter Filter identifying the parent documents. + */ + public DiversifyingChildrenFloat16KnnVectorQuery( + String field, short[] query, Query childFilter, int k, BitSetProducer parentsFilter) { + this(field, query, childFilter, k, parentsFilter, DEFAULT); + } + + /** + * Create a DiversifyingChildrenFloat16KnnVectorQuery. + * + * @param field the query field + * @param query the vector query + * @param childFilter the child filter + * @param k how many parent documents to return given the matching children + * @param parentsFilter Filter identifying the parent documents. + * @param searchStrategy the search strategy to use. If null, the default strategy will be used. + * The underlying format may not support all strategies and is free to ignore the requested + * strategy. + * @lucene.experimental + */ + public DiversifyingChildrenFloat16KnnVectorQuery( + String field, + short[] query, + Query childFilter, + int k, + BitSetProducer parentsFilter, + KnnSearchStrategy searchStrategy) { + super(field, query, k, childFilter, searchStrategy); + this.childFilter = childFilter; + this.parentsFilter = parentsFilter; + this.k = k; + this.query = query; + } + + @Override + protected TopDocs exactSearch( + LeafReaderContext context, DocIdSetIterator acceptIterator, QueryTimeout queryTimeout) + throws IOException { + Float16VectorValues float16VectorValues = context.reader().getFloat16VectorValues(field); + if (float16VectorValues == null) { + Float16VectorValues.checkField(context.reader(), field); + return NO_RESULTS; + } + + BitSet parentBitSet = parentsFilter.getBitSet(context); + if (parentBitSet == null) { + return NO_RESULTS; + } + VectorScorer float16VectorScorer = float16VectorValues.scorer(query); + if (float16VectorScorer == null) { + return NO_RESULTS; + } + return DiversifyingChildrenVectorScorer.collect( + acceptIterator, parentBitSet, float16VectorScorer, k, queryTimeout); + } + + @Override + protected KnnCollectorManager getKnnCollectorManager(int k, IndexSearcher searcher) { + return new DiversifyingNearestChildrenKnnCollectorManager(k, parentsFilter, searcher); + } + + @Override + protected TopDocs approximateSearch( + LeafReaderContext context, + AcceptDocs acceptDocs, + int visitedLimit, + KnnCollectorManager knnCollectorManager) + throws IOException { + Float16VectorValues.checkField(context.reader(), field); + KnnCollector collector = + knnCollectorManager.newCollector(visitedLimit, searchStrategy, context); + if (collector == null) { + return NO_RESULTS; + } + context.reader().searchNearestVectors(field, query, collector, acceptDocs); + return collector.topDocs(); + } + + @Override + public String toString(String field) { + StringBuilder buffer = new StringBuilder(); + buffer.append(getClass().getSimpleName() + ":"); + buffer.append(this.field + "[" + query[0] + ",...]"); Review Comment: `query[0]` is a `short`, should we convert `fp16` -> `fp32` and print that instead? ########## lucene/CHANGES.txt: ########## @@ -143,6 +143,8 @@ New Features * GITHUB#16473: Add scalar quantization support in Fp16 vector encoding. (Pulkit Gupta) +* GITHUB#16738: Add fp16 query support in diversifying children KNN query. (Pulkit Gupta) Review Comment: Let's add the name of the new query class? ########## lucene/join/src/java/org/apache/lucene/search/join/DiversifyingChildrenVectorScorer.java: ########## @@ -0,0 +1,129 @@ +/* + * 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.join; + +import java.io.IOException; +import org.apache.lucene.index.QueryTimeout; +import org.apache.lucene.search.DocIdSetIterator; +import org.apache.lucene.search.HitQueue; +import org.apache.lucene.search.ScoreDoc; +import org.apache.lucene.search.TopDocs; +import org.apache.lucene.search.TotalHits; +import org.apache.lucene.search.VectorScorer; +import org.apache.lucene.util.BitSet; + +/** + * Iterates the accepted child documents one parent at a time, tracking the best scoring child of + * each parent. Scoring is delegated to the given {@link VectorScorer}. + */ +class DiversifyingChildrenVectorScorer { + private final VectorScorer vectorScorer; + private final DocIdSetIterator vectorIterator; + private final DocIdSetIterator acceptedChildrenIterator; + private final BitSet parentBitSet; + private int currentParent = -1; + private int bestChild = -1; + private float currentScore = Float.NEGATIVE_INFINITY; + + DiversifyingChildrenVectorScorer( + DocIdSetIterator acceptedChildrenIterator, BitSet parentBitSet, VectorScorer vectorScorer) { + this.acceptedChildrenIterator = acceptedChildrenIterator; + this.vectorScorer = vectorScorer; + this.vectorIterator = vectorScorer.iterator(); + this.parentBitSet = parentBitSet; + } + + public int bestChild() { + return bestChild; + } + + public int nextParent() throws IOException { + int nextChild = acceptedChildrenIterator.docID(); + if (nextChild == -1) { + nextChild = acceptedChildrenIterator.nextDoc(); + } + if (nextChild == DocIdSetIterator.NO_MORE_DOCS) { + currentParent = DocIdSetIterator.NO_MORE_DOCS; + return currentParent; + } + currentScore = Float.NEGATIVE_INFINITY; + currentParent = parentBitSet.nextSetBit(nextChild); + do { + vectorIterator.advance(nextChild); + float score = vectorScorer.score(); + if (score > currentScore) { + bestChild = nextChild; + currentScore = score; + } + } while ((nextChild = acceptedChildrenIterator.nextDoc()) != DocIdSetIterator.NO_MORE_DOCS + && nextChild < currentParent); + return currentParent; + } + + public float score() throws IOException { + return currentScore; + } + + /** + * Returns the top {@code k} scoring children, at most one per parent document. The results are + * marked as a lower bound if the given timeout is met before every parent has been visited. + * + * @param acceptedChildrenIterator the child documents to score + * @param parentBitSet the parent documents + * @param scorer scores a child document against the query vector + * @param k how many children to return + * @param queryTimeout the timeout to honour, or null for no timeout + */ + static TopDocs collect( Review Comment: This creates an object on the first line with the parameters as-is, maybe keep this function as non-static? i.e. instead of: ```java DiversifyingChildrenVectorScorer.collect( acceptIterator, parentBitSet, floatVectorScorer, k, queryTimeout) ``` a caller would do: ```java new DiversifyingChildrenVectorScorer(acceptIterator, parentBitSet, floatVectorScorer) .collect(k, queryTimeout) ``` -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
