diff --git a/examples/pom.xml b/examples/pom.xml
index 768a9a87..a2af7e60 100644
--- a/examples/pom.xml
+++ b/examples/pom.xml
@@ -30,18 +30,18 @@
org.apache.lucene
lucene-core
- 10.2.0
+ 10.4.0
org.apache.lucene
lucene-codecs
- 10.2.0
+ 10.4.0
test
org.apache.lucene
lucene-backward-codecs
- 10.2.0
+ 10.4.0
commons-io
diff --git a/pom.xml b/pom.xml
index 00114972..bc66c14e 100644
--- a/pom.xml
+++ b/pom.xml
@@ -47,23 +47,23 @@
org.apache.lucene
lucene-core
- 10.2.0
+ 10.4.0
org.apache.lucene
lucene-codecs
- 10.2.0
+ 10.4.0
test
org.apache.lucene
lucene-backward-codecs
- 10.2.0
+ 10.4.0
org.apache.lucene
lucene-test-framework
- 10.2.0
+ 10.4.0
test
diff --git a/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWParams.java b/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWParams.java
index 9e663d2c..ab387a6a 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWParams.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWParams.java
@@ -10,6 +10,7 @@
import java.util.Objects;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
+import java.util.function.Supplier;
public class AcceleratedHNSWParams {
@@ -17,33 +18,40 @@ public class AcceleratedHNSWParams {
* TODO: Update boundaries for all parameters when a consensus is reached.
* Issue: https://github.com/rapidsai/cuvs-lucene/issues/99
*/
- private static final int MIN_WRITER_THREADS = 1;
- private static final int MAX_WRITER_THREADS = 512;
- private static final int MIN_INT_GRAPH_DEG = 2;
- private static final int MAX_INT_GRAPH_DEG = 512;
- private static final int MIN_GRAPH_DEG = 1;
- private static final int MAX_GRAPH_DEG = 512;
- private static final int MIN_HNSW_LAYERS = 1;
- private static final int MAX_HNSW_LAYERS = 3;
- private static final int MIN_MAX_CONN = 1;
- private static final int MAX_MAX_CONN = 512;
- private static final int MIN_BEAM_WIDTH = 1;
- private static final int MAX_BEAM_WIDTH = 512;
- private static final int MIN_NUM_MERGE_WORKERS = 1;
- private static final int MAX_NUM_MERGE_WORKERS = 512;
-
- private static final int DEFAULT_WRITER_THREADS = 1;
- private static final int DEFAULT_INT_GRAPH_DEGREE = 128;
- private static final int DEFAULT_GRAPH_DEGREE = 64;
- private static final int DEFAULT_HNSW_LAYERS = 1;
- private static final int DEFAULT_MAX_CONN = 32;
- private static final int DEFAULT_BEAM_WIDTH = 32;
- private static final CagraGraphBuildAlgo DEFAULT_CAGRA_GRAPH_BUILD_ALGO =
+ public static final int MIN_WRITER_THREADS = 1;
+ public static final int MAX_WRITER_THREADS = 512;
+ public static final int MIN_INT_GRAPH_DEG = 2;
+ public static final int MAX_INT_GRAPH_DEG = 512;
+ public static final int MIN_GRAPH_DEG = 1;
+ public static final int MAX_GRAPH_DEG = 512;
+ public static final int MIN_HNSW_LAYERS = 1;
+ public static final int MAX_HNSW_LAYERS = 3;
+ public static final int MIN_MAX_CONN = 1;
+ public static final int MAX_MAX_CONN = 512;
+ public static final int MIN_BEAM_WIDTH = 1;
+ public static final int MAX_BEAM_WIDTH = 512;
+ public static final int MIN_NUM_MERGE_WORKERS = 1;
+ public static final int MAX_NUM_MERGE_WORKERS = 512;
+
+ public static final int DEFAULT_WRITER_THREADS = 1;
+ public static final int DEFAULT_INT_GRAPH_DEGREE = 128;
+ public static final int DEFAULT_GRAPH_DEGREE = 64;
+ public static final int DEFAULT_HNSW_LAYERS = 1;
+ public static final int DEFAULT_MAX_CONN = 32;
+ public static final int DEFAULT_BEAM_WIDTH = 32;
+ public static final CagraGraphBuildAlgo DEFAULT_CAGRA_GRAPH_BUILD_ALGO =
CagraGraphBuildAlgo.NN_DESCENT;
- private static final CuVSIvfPqParams DEFAULT_IVF_PQ_PARAMS =
- new CuVSIvfPqParams.Builder().build();
- private static final int DEFAULT_NUM_MERGE_WORKERS = 1;
- private static final ExecutorService DEFAULT_MERGE_EXE_SRVC = Executors.newFixedThreadPool(1);
+ public static final int DEFAULT_NUM_MERGE_WORKERS = 1;
+
+ public static final Supplier DEFAULT_IVF_PQ_PARAMS =
+ () -> {
+ return new CuVSIvfPqParams.Builder().build();
+ };
+
+ public static final Supplier DEFAULT_MERGE_EXE_SRVC =
+ () -> {
+ return Executors.newFixedThreadPool(DEFAULT_NUM_MERGE_WORKERS);
+ };
private final int writerThreads;
private final int intermediateGraphDegree;
@@ -57,7 +65,7 @@ public class AcceleratedHNSWParams {
private final ExecutorService mergeExec;
/**
- * Constructs an instance of {@link GPUSearchParams} with specific parameter values.
+ * Constructs an instance of {@link AcceleratedHNSWParams} with specific parameter values.
*
* @param writerThreads Number of cuVS writer threads to use.
* @param intermediateGraphDegree The intermediate graph degree while building the CAGRA index.
@@ -67,7 +75,7 @@ public class AcceleratedHNSWParams {
* @param maxConn The max connection parameter used when building HNSW index with the fallback mechanism.
* @param beamWidth The beam width parameter used when building HNSW index with the fallback mechanism.
* @param cagraGraphBuildAlgo The CAGRA graph build algorithm to use [NN_DESCENT, IVF_PQ].
- * @param cagraGraphBuildAlgo An instance of CuVSIvfPqParams containing IVF_PQ specific parameters.
+ * @param cuVSIvfPqParams An instance of CuVSIvfPqParams containing IVF_PQ specific parameters.
* @param numMergeWorkers The number of merge workers to use with the fallback mechanism.
* @param mergeExec The instance of {@link ExecutorService} to use with the fallback mechanism.
*/
@@ -177,7 +185,7 @@ public int getNumMergeWorkers() {
}
/**
- * Get the instance of the {@link ExecutorService} to be used in the fallback mechanism *
+ * Get the instance of the {@link ExecutorService} to be used in the fallback mechanism
*
* @return the instance of the {@link ExecutorService}
*/
@@ -220,14 +228,14 @@ public static class Builder {
private int maxConn = DEFAULT_MAX_CONN;
private int beamWidth = DEFAULT_BEAM_WIDTH;
private CagraGraphBuildAlgo cagraGraphBuildAlgo = DEFAULT_CAGRA_GRAPH_BUILD_ALGO;
- private CuVSIvfPqParams cuVSIvfPqParams = DEFAULT_IVF_PQ_PARAMS;
private int numMergeWorkers = DEFAULT_NUM_MERGE_WORKERS;
- private ExecutorService mergeExec = DEFAULT_MERGE_EXE_SRVC;
+ private CuVSIvfPqParams cuVSIvfPqParams = null;
+ private ExecutorService mergeExec = null;
/**
* Set the number of cuVS writer threads while building the index
* Valid range - Minimum: {@value MIN_WRITER_THREADS}, Maximum: {@value MAX_WRITER_THREADS}
- * Default value - 64
+ * Default value - {@value DEFAULT_WRITER_THREADS}
*
* @param writerThreads
* @return instance of {@link Builder}
@@ -240,7 +248,7 @@ public Builder withWriterThreads(int writerThreads) {
/**
* Set the intermediate graph degree to use while building CAGRA index
* Valid range - Minimum: {@value MIN_INT_GRAPH_DEG}, Maximum: {@value MAX_INT_GRAPH_DEG}
- * Default value - 128
+ * Default value - {@value DEFAULT_INT_GRAPH_DEGREE}
*
* @param intermediateGraphDegree
* @return instance of {@link Builder}
@@ -253,7 +261,7 @@ public Builder withIntermediateGraphDegree(int intermediateGraphDegree) {
/**
* Set the graph degree to use while building CAGRA index
* Valid range - Minimum: {@value MIN_GRAPH_DEG}, Maximum: {@value MAX_GRAPH_DEG}
- * Default value - 64
+ * Default value - {@value DEFAULT_GRAPH_DEGREE}
*
* @param graphDegree
* @return instance of {@link Builder}
@@ -266,7 +274,7 @@ public Builder withGraphDegree(int graphDegree) {
/**
* Set the number of HNSW layers to construct while building the HNSW index
* Valid range - Minimum: {@value MIN_HNSW_LAYERS}, Maximum: {@value MAX_HNSW_LAYERS}
- * Default value - 2
+ * Default value - {@value DEFAULT_HNSW_LAYERS}
*
* @param hnswLayers the number of HNSW layers
* @return instance of {@link Builder}
@@ -279,7 +287,7 @@ public Builder withHNSWLayer(int hnswLayers) {
/**
* Set the max connections parameter while building HNSW index with fallback mechanism
* Valid range - Minimum: {@value MIN_MAX_CONN}, Maximum: {@value MAX_MAX_CONN}
- * Default value - 8
+ * Default value - {@value DEFAULT_MAX_CONN}
*
* @param maxConn the max connections parameter
* @return instance of {@link Builder}
@@ -292,7 +300,7 @@ public Builder withMaxConn(int maxConn) {
/**
* Set the beam width parameter while building HNSW index with fallback mechanism
* Valid range - Minimum: {@value MIN_BEAM_WIDTH}, Maximum: {@value MAX_BEAM_WIDTH}
- * Default value - 16
+ * Default value - {@value DEFAULT_BEAM_WIDTH}
*
* @param beamWidth the beam width parameter
* @return instance of {@link Builder}
@@ -304,7 +312,7 @@ public Builder withBeamWidth(int beamWidth) {
/**
* Set the CAGRA graph build algorithm to use
- * Default NN_DESCENT
+ * Default value - NN_DESCENT
*
* @param cagraGraphBuildAlgo
* @return instance of {@link Builder}
@@ -327,7 +335,7 @@ public Builder withCuVSIvfPqParams(CuVSIvfPqParams cuVSIvfPqParams) {
/**
* Set the number of merge workers to be used with the fallback mechanism
- * Default value - 1
+ * Default value - {@value DEFAULT_NUM_MERGE_WORKERS}
*
* @param numMergeWorkers number of merge workers to set
* @return instance of {@link Builder}
@@ -407,9 +415,6 @@ private void validate() throws IllegalArgumentException {
if (Objects.isNull(cagraGraphBuildAlgo)) {
throw new IllegalArgumentException("cagraGraphBuildAlgo cannot be null.");
}
- if (Objects.isNull(cuVSIvfPqParams)) {
- throw new IllegalArgumentException("cuVSIvfPqParams cannot be null.");
- }
if (numMergeWorkers < MIN_NUM_MERGE_WORKERS || numMergeWorkers > MAX_NUM_MERGE_WORKERS) {
throw new IllegalArgumentException(
"numMergeWorkers not in valid range. Valid range: ["
@@ -418,9 +423,6 @@ private void validate() throws IllegalArgumentException {
+ MAX_NUM_MERGE_WORKERS
+ "]");
}
- if (Objects.isNull(mergeExec)) {
- throw new IllegalArgumentException("mergeExec cannot be null.");
- }
}
/**
@@ -429,6 +431,12 @@ private void validate() throws IllegalArgumentException {
* @return instance of {@link AcceleratedHNSWParams}
*/
public AcceleratedHNSWParams build() {
+ if (Objects.isNull(cuVSIvfPqParams)) {
+ cuVSIvfPqParams = DEFAULT_IVF_PQ_PARAMS.get();
+ }
+ if (Objects.isNull(mergeExec)) {
+ mergeExec = DEFAULT_MERGE_EXE_SRVC.get();
+ }
validate();
return new AcceleratedHNSWParams(
writerThreads,
diff --git a/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java b/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java
index 457ad22d..c4efdd9d 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java
@@ -34,7 +34,8 @@ public class AcceleratedHNSWUtils {
public enum QuantizationType {
BINARY,
- SCALAR
+ SCALAR,
+ NONE
}
private static final LuceneProvider LUCENE_PROVIDER;
@@ -42,7 +43,7 @@ public enum QuantizationType {
static {
try {
- LUCENE_PROVIDER = LuceneProvider.getInstance("99");
+ LUCENE_PROVIDER = LuceneProvider.getInstance(LuceneProvider.LUCENE_FLOAT_HNSW_LINE);
VECTOR_SIMILARITY_FUNCTIONS = LUCENE_PROVIDER.getSimilarityFunctions();
} catch (Exception e) {
throw new ExceptionInInitializerError(e.getMessage());
@@ -84,7 +85,7 @@ public static GPUBuiltHnswGraph createMultiLayerHnswGraph(
int size,
int dimensions,
CuVSMatrix adjacencyListMatrix,
- List vectors,
+ List> vectors,
int hnswLayers,
int graphDegree,
CagraIndexParams params,
@@ -132,17 +133,32 @@ public static GPUBuiltHnswGraph createMultiLayerHnswGraph(
layerNodes.add(selectedNodes);
- // Extract vectors for selected nodes
- int bytesPerVector = (dimensions + 7) / 8;
- byte[][] selectedVectors = new byte[nextLayerSize][];
- for (int i = 0; i < nextLayerSize; i++) {
- selectedVectors[i] = vectors.get(selectedNodes[i]);
- }
+ if (quantization == QuantizationType.NONE) {
+ // Extract vectors for selected nodes
+ float[][] selectedVectors = new float[nextLayerSize][];
+ for (int i = 0; i < nextLayerSize; i++) {
+ selectedVectors[i] = (float[]) vectors.get(selectedNodes[i]);
+ }
+
+ // Build CAGRA graph for this layer
+ layerAdjacencies.add(
+ buildCagraGraphForSubset(
+ selectedVectors, selectedNodes, 0, params, dimensions, quantization));
+
+ } else {
+
+ // Extract vectors for selected nodes
+ int bytesPerVector = (dimensions + 7) / 8;
+ byte[][] selectedVectors = new byte[nextLayerSize][];
+ for (int i = 0; i < nextLayerSize; i++) {
+ selectedVectors[i] = (byte[]) vectors.get(selectedNodes[i]);
+ }
- // Build CAGRA graph for this layer
- layerAdjacencies.add(
- buildCagraGraphForSubset(
- selectedVectors, selectedNodes, bytesPerVector, params, dimensions, quantization));
+ // Build CAGRA graph for this layer
+ layerAdjacencies.add(
+ buildCagraGraphForSubset(
+ selectedVectors, selectedNodes, bytesPerVector, params, dimensions, quantization));
+ }
// Update for next iteration
currentLayerSize = nextLayerSize;
@@ -160,7 +176,7 @@ public static GPUBuiltHnswGraph createMultiLayerHnswGraph(
* Builds a CAGRA graph for a subset of binary quantized vectors
*/
private static CuVSMatrix buildCagraGraphForSubset(
- byte[][] vectors,
+ Object vectors,
int[] selectedNodes,
int bytesPerVector,
CagraIndexParams params,
@@ -172,9 +188,12 @@ private static CuVSMatrix buildCagraGraphForSubset(
if (quantization == QuantizationType.BINARY) {
subsetDataset =
- createByteMatrixFromArray(vectors, bytesPerVector, getCuVSResourcesInstance());
+ createByteMatrixFromArray((byte[][]) vectors, bytesPerVector, getCuVSResourcesInstance());
+ } else if (quantization == QuantizationType.SCALAR) {
+ subsetDataset =
+ createByteMatrixFromArray((byte[][]) vectors, dimensions, getCuVSResourcesInstance());
} else {
- subsetDataset = createByteMatrixFromArray(vectors, dimensions, getCuVSResourcesInstance());
+ subsetDataset = CuVSMatrix.ofArray((float[][]) vectors);
}
// Build CAGRA index for the subset
@@ -211,8 +230,18 @@ private static CuVSMatrix buildCagraGraphForSubset(
return CuVSMatrix.ofArray(remappedAdjacency);
}
+ private static int[] getSortedNodes(NodesIterator nodesOnLevel) {
+ int[] nodes = new int[nodesOnLevel.size()];
+ int consumed = nodesOnLevel.consume(nodes);
+ assert consumed == nodesOnLevel.size();
+ Arrays.sort(nodes);
+ return nodes;
+ }
+
/**
- * Returns a 2D array of offsets (information written while writing the meta info)
+ * Returns a 2D array of offsets (information written while writing the meta info). Neighbor ids
+ * are written with {@link IndexOutput#writeGroupVInts(int[], int)} to match Apache Lucene 10.4
+ * {@code Lucene99HnswVectorsFormat} (GroupVarInt encoding).
*
* @param graph instance of GPUBuiltHnswGraph
* @param vectorIndex instance of IndexOutput
@@ -226,7 +255,7 @@ public static int[][] writeGraph(GPUBuiltHnswGraph graph, IndexOutput vectorInde
int[][] offsets = new int[graph.numLevels()][];
int[] scratch = new int[graph.maxConn() * 2];
for (int level = 0; level < graph.numLevels(); level++) {
- int[] sortedNodes = NodesIterator.getSortedNodes(graph.getNodesOnLevel(level));
+ int[] sortedNodes = getSortedNodes(graph.getNodesOnLevel(level));
offsets[level] = new int[sortedNodes.length];
int nodeOffsetId = 0;
@@ -259,10 +288,7 @@ public static int[][] writeGraph(GPUBuiltHnswGraph graph, IndexOutput vectorInde
}
// Write the size after duplicates are removed
vectorIndex.writeVInt(actualSize);
- // Write de-duplicated neighbors
- for (int i = 0; i < actualSize; i++) {
- vectorIndex.writeVInt(scratch[i]);
- }
+ vectorIndex.writeGroupVInts(scratch, actualSize);
offsets[level][nodeOffsetId++] =
Math.toIntExact(vectorIndex.getFilePointer() - offsetStart);
}
diff --git a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUSearchCodec.java b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUSearchCodec.java
index 2de761a7..c3befa23 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUSearchCodec.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUSearchCodec.java
@@ -29,7 +29,7 @@ public class CuVS2510GPUSearchCodec extends FilterCodec {
* @throws Exception
*/
public CuVS2510GPUSearchCodec() throws Exception {
- this(NAME, LuceneProvider.getCodec("101"));
+ this(NAME, LuceneProvider.getLucene104Codec());
initializeFormat(new GPUSearchParams.Builder().build());
}
@@ -53,7 +53,7 @@ public CuVS2510GPUSearchCodec(String name, Codec delegate) {
* @throws Exception Exception raised when initializing the codec
*/
public CuVS2510GPUSearchCodec(GPUSearchParams params) throws Exception {
- this(NAME, LuceneProvider.getCodec("101"));
+ this(NAME, LuceneProvider.getLucene104Codec());
initializeFormat(params);
}
diff --git a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsFormat.java b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsFormat.java
index ccc61eae..949f7938 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsFormat.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsFormat.java
@@ -39,7 +39,7 @@ public class CuVS2510GPUVectorsFormat extends KnnVectorsFormat {
static {
try {
- LUCENE_PROVIDER = LuceneProvider.getInstance("99");
+ LUCENE_PROVIDER = LuceneProvider.getInstance(LuceneProvider.LUCENE_FLOAT_HNSW_LINE);
FLAT_VECTORS_FORMAT =
LUCENE_PROVIDER.getLuceneFlatVectorsFormatInstance(DefaultFlatVectorScorer.INSTANCE);
} catch (Exception e) {
diff --git a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java
index c783660d..4357193e 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java
@@ -40,13 +40,14 @@
import org.apache.lucene.index.VectorEncoding;
import org.apache.lucene.index.VectorSimilarityFunction;
import org.apache.lucene.internal.hppc.IntObjectHashMap;
+import org.apache.lucene.search.AcceptDocs;
+import org.apache.lucene.search.DocIdSetIterator;
import org.apache.lucene.search.KnnCollector;
import org.apache.lucene.store.ChecksumIndexInput;
import org.apache.lucene.store.DataInput;
import org.apache.lucene.store.IOContext;
import org.apache.lucene.store.IOContext.Context;
import org.apache.lucene.store.IndexInput;
-import org.apache.lucene.store.ReadAdvice;
import org.apache.lucene.util.Bits;
import org.apache.lucene.util.IOUtils;
import org.apache.lucene.util.hnsw.IntToIntFunction;
@@ -69,7 +70,7 @@ public class CuVS2510GPUVectorsReader extends KnnVectorsReader {
static {
try {
- LUCENE_PROVIDER = LuceneProvider.getInstance("99");
+ LUCENE_PROVIDER = LuceneProvider.getInstance(LuceneProvider.LUCENE_FLOAT_HNSW_LINE);
VECTOR_SIMILARITY_FUNCTIONS = LUCENE_PROVIDER.getSimilarityFunctions();
} catch (Exception e) {
throw new ExceptionInInitializerError(e.getMessage());
@@ -111,8 +112,7 @@ public CuVS2510GPUVectorsReader(SegmentReadState state, FlatVectorsReader flatRe
} finally {
CodecUtil.checkFooter(meta, priorException);
}
- var ioContext = state.context.withReadAdvice(ReadAdvice.SEQUENTIAL);
- cuvsIndexInput = openCuVSInput(state, versionMeta, ioContext);
+ cuvsIndexInput = openCuVSInput(state, versionMeta, state.context);
/*
* Only load indexes on the GPU when this reader is opening for searches.
* Do not load indexes on the GPU when this reader is opening during merge calls.
@@ -392,11 +392,63 @@ private static FloatToFloatFunction getScoreNormalizationFunc(VectorSimilarityFu
return score -> (1f / (1f + score));
}
+ /**
+ * Maps live-document acceptance to vector ordinals. {@link AcceptDocs#bits()} may be null in
+ * Lucene 10.4+ while {@link FloatVectorValues#getAcceptOrds} returns null when passed null bits;
+ * this method fills that gap for GPU prefilter construction.
+ */
+ private static Bits computeAcceptedOrds(FloatVectorValues rawValues, AcceptDocs acceptDocs)
+ throws IOException {
+ if (acceptDocs == null) {
+ return null;
+ }
+ Bits live = acceptDocs.bits();
+ if (live != null) {
+ Bits mapped = rawValues.getAcceptOrds(live);
+ if (mapped != null) {
+ return mapped;
+ }
+ }
+ final Bits docAccept;
+ if (live != null) {
+ docAccept = live;
+ } else {
+ BitSet docBits = new BitSet();
+ DocIdSetIterator disi = acceptDocs.iterator();
+ for (int doc = disi.nextDoc(); doc != DocIdSetIterator.NO_MORE_DOCS; doc = disi.nextDoc()) {
+ docBits.set(doc);
+ }
+ docAccept =
+ new Bits() {
+ @Override
+ public boolean get(int docId) {
+ return docBits.get(docId);
+ }
+
+ @Override
+ public int length() {
+ return docBits.length();
+ }
+ };
+ }
+ return new Bits() {
+ @Override
+ public boolean get(int ord) {
+ return docAccept.get(rawValues.ordToDoc(ord));
+ }
+
+ @Override
+ public int length() {
+ return rawValues.size();
+ }
+ };
+ }
+
/**
* Returns the k nearest neighbor documents using cuVS's CAGRA or brute force algorithm for this field, to the given vector.
*/
@Override
- public void search(String field, float[] target, KnnCollector knnCollector, Bits acceptDocs)
+ public void search(String field, float[] target, KnnCollector knnCollector, AcceptDocs acceptDocs)
throws IOException {
var fieldEntry = getFieldEntry(field, VectorEncoding.FLOAT32);
if (fieldEntry.count() == 0 || knnCollector.k() == 0) {
@@ -410,12 +462,13 @@ public void search(String field, float[] target, KnnCollector knnCollector, Bits
}
final FloatVectorValues rawValues = flatVectorsReader.getFloatVectorValues(field);
- final Bits acceptedOrds = rawValues.getAcceptOrds(acceptDocs);
+ final Bits acceptedOrds = computeAcceptedOrds(rawValues, acceptDocs);
BitSet[] mask = null;
int maskLength = 0;
int topK = knnCollector.k();
if (acceptDocs != null) {
+ assert acceptedOrds != null;
mask = new BitSet[1]; // As there is only one query "target"
mask[0] = new BitSet(acceptedOrds.length());
/*
@@ -528,7 +581,7 @@ public void search(String field, float[] target, KnnCollector knnCollector, Bits
* This is not supported.
*/
@Override
- public void search(String field, byte[] target, KnnCollector knnCollector, Bits acceptDocs)
+ public void search(String field, byte[] target, KnnCollector knnCollector, AcceptDocs acceptDocs)
throws IOException {
throw new UnsupportedOperationException("Byte vectors are not currently supported");
}
diff --git a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java
index 2bb6732c..7bfb4ae2 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java
@@ -73,7 +73,7 @@ public class CuVS2510GPUVectorsWriter extends KnnVectorsWriter {
static {
try {
- LUCENE_PROVIDER = LuceneProvider.getInstance("99");
+ LUCENE_PROVIDER = LuceneProvider.getInstance(LuceneProvider.LUCENE_FLOAT_HNSW_LINE);
VECTOR_SIMILARITY_FUNCTIONS = LUCENE_PROVIDER.getSimilarityFunctions();
} catch (Exception e) {
throw new ExceptionInInitializerError(e.getMessage());
diff --git a/src/main/java/com/nvidia/cuvs/lucene/QuantizedFieldWriter.java b/src/main/java/com/nvidia/cuvs/lucene/FieldWriter.java
similarity index 84%
rename from src/main/java/com/nvidia/cuvs/lucene/QuantizedFieldWriter.java
rename to src/main/java/com/nvidia/cuvs/lucene/FieldWriter.java
index 8a6bbfed..627d931d 100644
--- a/src/main/java/com/nvidia/cuvs/lucene/QuantizedFieldWriter.java
+++ b/src/main/java/com/nvidia/cuvs/lucene/FieldWriter.java
@@ -17,10 +17,10 @@
import org.apache.lucene.index.FieldInfo;
import org.apache.lucene.util.RamUsageEstimator;
-public class QuantizedFieldWriter extends KnnFieldVectorsWriter