diff --git a/CHANGELOG.md b/CHANGELOG.md index 3d4bb364..2eaa9c4b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), ### Enhancements * Add script for loading vector data using Parquet [192](https://github.com/opensearch-project/opensearch-jvector/issues/192) * Update script and main docs with minor fixes [196] (https://github.com/opensearch-project/opensearch-jvector/pull/196) +* Support deletes for incremental insertion [203](https://github.com/opensearch-project/opensearch-jvector/pull/203) ### Bug Fixes ### Infrastructure ### Documentation diff --git a/gradle.properties b/gradle.properties index f37a062e..c9adf22e 100644 --- a/gradle.properties +++ b/gradle.properties @@ -5,7 +5,7 @@ version=1.0.0 systemProp.bwc.version=1.3.4 -jvector_version=4.0.0-rc.4 +jvector_version=4.0.0-rc.6 java_release_version=21 # org.gradle.jvmargs=--add-exports jdk.compiler/com.sun.tools.javac.api=ALL-UNNAMED \ diff --git a/src/main/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMap.java b/src/main/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMap.java index 7fff91e1..ce92d203 100644 --- a/src/main/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMap.java +++ b/src/main/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMap.java @@ -46,7 +46,10 @@ public GraphNodeIdToDocMap(IndexInput in) throws IOException { for (int ord = 0; ord < size; ord++) { final int docId = in.readVInt(); graphNodeIdsToDocIds[ord] = docId; - docIdsToGraphNodeIds[docId] = ord; + if (docId != -1) { + // ignore deleted documents + docIdsToGraphNodeIds[docId] = ord; + } } } @@ -68,18 +71,19 @@ public GraphNodeIdToDocMap(int[] graphNodeIdsToDocIds) { // We are going to assume that the number of ordinals is roughly the same as the number of documents in the segment, therefore, // the mapping will not be sparse. if (maxDocs < graphNodeIdsToDocIds.length) { - throw new IllegalStateException("Max docs " + maxDocs + " is less than the number of ordinals " + graphNodeIdsToDocIds.length); - } - if (maxDocId > graphNodeIdsToDocIds.length) { log.warn( - "Max doc id {} is greater than the number of ordinals {}, this implies a lot of deleted documents. Or that some documents are missing vectors. Wasting a lot of memory", - maxDocId, + "Max docs {} is less than the number of ordinals {}, this implies a lot of deleted documents. Or that some documents are missing vectors. Wasting a lot of memory", + maxDocs, graphNodeIdsToDocIds.length ); } this.docIdsToGraphNodeIds = new int[maxDocs]; Arrays.fill(this.docIdsToGraphNodeIds, -1); // -1 means no mapping to ordinal for (int ord = 0; ord < graphNodeIdsToDocIds.length; ord++) { + // -1 means no mapping to docId since document was deleted + if (graphNodeIdsToDocIds[ord] == -1) { + continue; + } this.docIdsToGraphNodeIds[graphNodeIdsToDocIds[ord]] = ord; } } diff --git a/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorFormat.java b/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorFormat.java index 5d25622d..34f28d6d 100644 --- a/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorFormat.java +++ b/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorFormat.java @@ -36,6 +36,7 @@ public class JVectorFormat extends KnnVectorsFormat { // Unfortunately, this can't be managed yet by the OpenSearch ThreadPool because it's not supporting {@link ForkJoinPool} types public static final ForkJoinPool SIMD_POOL_MERGE = getPhysicalCoreExecutor(); public static final ForkJoinPool SIMD_POOL_FLUSH = getPhysicalCoreExecutor(); + public static final ForkJoinPool PARALLELISM_POOL = ForkJoinPool.commonPool(); private final int maxConn; private final int beamWidth; diff --git a/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorReader.java b/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorReader.java index 3c8aa462..7642d6d3 100644 --- a/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorReader.java +++ b/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorReader.java @@ -169,7 +169,7 @@ public void search(String field, float[] target, KnnCollector knnCollector, Bits // Logic works as follows: if acceptDocs is null, we accept all ordinals. Otherwise, we check if the jVector ordinal has a // corresponding Lucene doc ID accepted by acceptDocs filter. io.github.jbellis.jvector.util.Bits compatibleBits = ord -> acceptDocs == null - || acceptDocs.get(jvectorLuceneDocMap.getLuceneDocId(ord)); + || (jvectorLuceneDocMap.getLuceneDocId(ord) != -1 && acceptDocs.get(jvectorLuceneDocMap.getLuceneDocId(ord))); try (var graphSearcher = new GraphSearcher(index)) { final var searchResults = graphSearcher.search( diff --git a/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorWriter.java b/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorWriter.java index 434e08a6..86345611 100644 --- a/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorWriter.java +++ b/src/main/java/org/opensearch/knn/index/codec/jvector/JVectorWriter.java @@ -11,6 +11,7 @@ import io.github.jbellis.jvector.graph.disk.feature.Feature; import io.github.jbellis.jvector.graph.disk.feature.FeatureId; import io.github.jbellis.jvector.graph.disk.feature.InlineVectors; +import io.github.jbellis.jvector.graph.diversity.VamanaDiversityProvider; import io.github.jbellis.jvector.graph.similarity.BuildScoreProvider; import io.github.jbellis.jvector.quantization.PQVectors; import io.github.jbellis.jvector.quantization.ProductQuantization; @@ -31,6 +32,7 @@ import org.apache.lucene.search.DocIdSetIterator; import org.apache.lucene.store.*; import org.apache.lucene.util.Bits; +import io.github.jbellis.jvector.util.FixedBitSet; import org.apache.lucene.util.IOUtils; import org.apache.lucene.util.RamUsageEstimator; import org.opensearch.knn.plugin.stats.KNNCounter; @@ -45,8 +47,7 @@ import static io.github.jbellis.jvector.quantization.KMeansPlusPlusClusterer.UNWEIGHTED; import static org.apache.lucene.codecs.lucene99.Lucene99HnswVectorsReader.readVectorEncoding; -import static org.opensearch.knn.index.codec.jvector.JVectorFormat.SIMD_POOL_FLUSH; -import static org.opensearch.knn.index.codec.jvector.JVectorFormat.SIMD_POOL_MERGE; +import static org.opensearch.knn.index.codec.jvector.JVectorFormat.*; /** * JVectorWriter is responsible for writing vector data into index segments using the JVector library. @@ -205,23 +206,27 @@ public void flush(int maxDoc, Sorter.DocMap sortMap) throws IOException { for (int ord = 0; ord < randomAccessVectorValues.size(); ord++) { newToOldOrds[ord] = ord; } + final RemappedRandomAccessVectorValues remappedRandomAccessVectorValues = new RemappedRandomAccessVectorValues( + randomAccessVectorValues, + newToOldOrds + ); final BuildScoreProvider buildScoreProvider; final PQVectors pqVectors; final FieldInfo fieldInfo = field.fieldInfo; - if (randomAccessVectorValues.size() >= minimumBatchSizeForQuantization) { + if (remappedRandomAccessVectorValues.size() >= minimumBatchSizeForQuantization) { log.info("Calculating codebooks and compressed vectors for field {}", fieldInfo.name); - pqVectors = getPQVectors(newToOldOrds, randomAccessVectorValues, fieldInfo); + pqVectors = getPQVectors(remappedRandomAccessVectorValues, fieldInfo); buildScoreProvider = BuildScoreProvider.pqBuildScoreProvider(getVectorSimilarityFunction(fieldInfo), pqVectors); } else { log.info( "Vector count: {}, less than limit to trigger PQ quantization: {}, for field {}, will use full precision vectors instead.", - randomAccessVectorValues.size(), + remappedRandomAccessVectorValues.size(), minimumBatchSizeForQuantization, fieldInfo.name ); pqVectors = null; buildScoreProvider = BuildScoreProvider.randomAccessScoreProvider( - randomAccessVectorValues, + remappedRandomAccessVectorValues, getVectorSimilarityFunction(fieldInfo) ); } @@ -238,39 +243,30 @@ public void flush(int maxDoc, Sorter.DocMap sortMap) throws IOException { OnHeapGraphIndex graph = getGraph( buildScoreProvider, - randomAccessVectorValues, - newToOldOrds, + remappedRandomAccessVectorValues, fieldInfo, segmentWriteState.segmentInfo.name, SIMD_POOL_FLUSH ); - writeField(field.fieldInfo, field.randomAccessVectorValues, pqVectors, newToOldOrds, graphNodeIdToDocMap, graph); + writeField(field.fieldInfo, remappedRandomAccessVectorValues, pqVectors, graphNodeIdToDocMap, graph); } } private void writeField( FieldInfo fieldInfo, - RandomAccessVectorValues randomAccessVectorValues, + RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, PQVectors pqVectors, - int[] newToOldOrds, GraphNodeIdToDocMap graphNodeIdToDocMap, OnHeapGraphIndex graph ) throws IOException { log.info( "Writing field {} with vector count: {}, for segment: {}", fieldInfo.name, - randomAccessVectorValues.size(), + remappedRandomAccessVectorValues.size(), segmentWriteState.segmentInfo.name ); - final var vectorIndexFieldMetadata = writeGraph( - graph, - randomAccessVectorValues, - fieldInfo, - pqVectors, - newToOldOrds, - graphNodeIdToDocMap - ); + final var vectorIndexFieldMetadata = writeGraph(graph, remappedRandomAccessVectorValues, fieldInfo, pqVectors, graphNodeIdToDocMap); meta.writeInt(fieldInfo.number); vectorIndexFieldMetadata.toOutput(meta); @@ -303,17 +299,16 @@ private void writeField( /** * Writes the graph and PQ codebooks and compressed vectors to the vector index file * @param graph graph - * @param randomAccessVectorValues random access vector values + * @param remappedRandomAccessVectorValues random access vector values with remapped ordinals * @param fieldInfo field info * @return Tuple of start offset and length of the graph * @throws IOException IOException */ private VectorIndexFieldMetadata writeGraph( OnHeapGraphIndex graph, - RandomAccessVectorValues randomAccessVectorValues, + RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, FieldInfo fieldInfo, PQVectors pqVectors, - int[] newToOldOrds, GraphNodeIdToDocMap graphNodeIdToDocMap ) throws IOException { // field data file, which contains the graph @@ -338,17 +333,17 @@ private VectorIndexFieldMetadata writeGraph( .fieldNumber(fieldInfo.number) .vectorEncoding(fieldInfo.getVectorEncoding()) .vectorSimilarityFunction(fieldInfo.getVectorSimilarityFunction()) - .vectorDimension(randomAccessVectorValues.dimension()) + .vectorDimension(remappedRandomAccessVectorValues.dimension()) .graphNodeIdToDocMap(graphNodeIdToDocMap); try ( var writer = new OnDiskSequentialGraphIndexWriter.Builder(graph, jVectorIndexWriter).with( - new InlineVectors(randomAccessVectorValues.dimension()) + new InlineVectors(remappedRandomAccessVectorValues.dimension()) ).build() ) { var suppliers = Feature.singleStateFactory( FeatureId.INLINE_VECTORS, - nodeId -> new InlineVectors.State(randomAccessVectorValues.getVector(newToOldOrds[nodeId])) + nodeId -> new InlineVectors.State(remappedRandomAccessVectorValues.getVector(nodeId)) ); writer.write(suppliers); long endGraphOffset = jVectorIndexWriter.position(); @@ -360,7 +355,7 @@ private VectorIndexFieldMetadata writeGraph( log.info( "Writing PQ codebooks and vectors for field {} since the size is {} >= {}", fieldInfo.name, - randomAccessVectorValues.size(), + remappedRandomAccessVectorValues.size(), minimumBatchSizeForQuantization ); resultBuilder.pqCodebooksAndVectorsOffset(endGraphOffset); @@ -378,17 +373,17 @@ private VectorIndexFieldMetadata writeGraph( } } - private PQVectors getPQVectors(int[] newToOldOrds, RandomAccessVectorValues randomAccessVectorValues, FieldInfo fieldInfo) + private PQVectors getPQVectors(RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, FieldInfo fieldInfo) throws IOException { final String fieldName = fieldInfo.name; final VectorSimilarityFunction vectorSimilarityFunction = fieldInfo.getVectorSimilarityFunction(); - log.info("Computing PQ codebooks for field {} for {} vectors", fieldName, randomAccessVectorValues.size()); + log.info("Computing PQ codebooks for field {} for {} vectors", fieldName, remappedRandomAccessVectorValues.size()); final long start = Clock.systemDefaultZone().millis(); - final var M = numberOfSubspacesPerVectorSupplier.apply(randomAccessVectorValues.dimension()); - final var numberOfClustersPerSubspace = Math.min(256, randomAccessVectorValues.size()); // number of centroids per + final var M = numberOfSubspacesPerVectorSupplier.apply(remappedRandomAccessVectorValues.dimension()); + final var numberOfClustersPerSubspace = Math.min(256, remappedRandomAccessVectorValues.size()); // number of centroids per // subspace ProductQuantization pq = ProductQuantization.compute( - randomAccessVectorValues, + remappedRandomAccessVectorValues, M, // number of subspaces numberOfClustersPerSubspace, // number of centroids per subspace vectorSimilarityFunction == VectorSimilarityFunction.EUCLIDEAN, // center the dataset @@ -401,9 +396,14 @@ private PQVectors getPQVectors(int[] newToOldOrds, RandomAccessVectorValues rand final long trainingTime = end - start; log.info("Computed PQ codebooks for field {}, in {} millis", fieldName, trainingTime); KNNCounter.KNN_QUANTIZATION_TRAINING_TIME.add(trainingTime); - log.info("Encoding and building PQ vectors for field {} for {} vectors", fieldName, randomAccessVectorValues.size()); + log.info("Encoding and building PQ vectors for field {} for {} vectors", fieldName, remappedRandomAccessVectorValues.size()); // PQVectors pqVectors = pq.encodeAll(randomAccessVectorValues, SIMD_POOL); - PQVectors pqVectors = PQVectors.encodeAndBuild(pq, newToOldOrds.length, newToOldOrds, randomAccessVectorValues, SIMD_POOL_MERGE); + PQVectors pqVectors = PQVectors.encodeAndBuild( + pq, + remappedRandomAccessVectorValues.size(), + remappedRandomAccessVectorValues, + SIMD_POOL_MERGE + ); log.info( "Encoded and built PQ vectors for field {}, original size: {} bytes, compressed size: {} bytes", fieldName, @@ -608,7 +608,7 @@ class RandomAccessMergedFloatVectorValues implements RandomAccessVectorValues { private final MergeState mergeState; private final GraphNodeIdToDocMap graphNodeIdToDocMap; private final int[] graphNodeIdsToRavvOrds; - private boolean deletesFound = false; + private final FixedBitSet[] liveGraphNodesPerReader; /** * Creates a random access view over merged float vector values. @@ -633,9 +633,19 @@ public RandomAccessMergedFloatVectorValues(FieldInfo fieldInfo, MergeState merge List allReaders = new ArrayList<>(); final MergeState.DocMap[] docMaps = mergeState.docMaps.clone(); final Bits[] liveDocs = mergeState.liveDocs.clone(); + this.liveGraphNodesPerReader = new FixedBitSet[liveDocs.length]; + // initialize liveGraphNodesPerReader + for (int i = 0; i < liveGraphNodesPerReader.length; i++) { + final int length; + if (liveDocs[i] == null) { + length = mergeState.maxDocs[i]; + } else { + length = liveDocs[i].length(); + } + liveGraphNodesPerReader[i] = new FixedBitSet(length); + liveGraphNodesPerReader[i].set(0, length); + } final int[] baseOrds = new int[mergeState.knnVectorsReaders.length]; - final int[] deletedOrds = new int[mergeState.knnVectorsReaders.length]; // counts the number of deleted documents in each reader - // that previously had a vector // Find the leading reader, count the total number of live vectors, and the base ordinals for each reader for (int i = 0; i < mergeState.knnVectorsReaders.length; i++) { @@ -654,8 +664,8 @@ public RandomAccessMergedFloatVectorValues(FieldInfo fieldInfo, MergeState merge if (liveDocs[i] == null || liveDocs[i].get(it.docID())) { liveVectorCountInReader++; } else { - deletedOrds[i]++; - deletesFound = true; + // This vector is deleted so we need to mark it as deleted in the liveGraphNodesPerReader + liveGraphNodesPerReader[i].clear(it.index()); } } if (liveVectorCountInReader >= vectorsCountInLeadingReader) { @@ -695,6 +705,10 @@ public RandomAccessMergedFloatVectorValues(FieldInfo fieldInfo, MergeState merge final int tempBaseOrd = baseOrds[LEADING_READER_IDX]; baseOrds[LEADING_READER_IDX] = baseOrds[tempLeadingReaderIdx]; baseOrds[tempLeadingReaderIdx] = tempBaseOrd; + // swap liveGraphNodesPerReader + final FixedBitSet tempLiveGraphNodes = liveGraphNodesPerReader[LEADING_READER_IDX]; + liveGraphNodesPerReader[LEADING_READER_IDX] = liveGraphNodesPerReader[tempLeadingReaderIdx]; + liveGraphNodesPerReader[tempLeadingReaderIdx] = tempLiveGraphNodes; } this.perReaderFloatVectorValues = new JVectorFloatVectorValues[readers.length]; @@ -703,111 +717,87 @@ public RandomAccessMergedFloatVectorValues(FieldInfo fieldInfo, MergeState merge // Build mapping from global ordinal to [readerIndex, readerOrd] this.ravvOrdToReaderMapping = new int[totalDocsCount][2]; - int documentsIterated = 0; - // Will be used to build the new graphNodeIdToDocMap with the new graph node id to docId mapping. // This mapping should not be used to access the vectors at any time during construction, but only after the merge is complete // and the new segment is created and used by searchers. - final int[] graphNodeIdToDocIds = new int[totalLiveVectorsCount]; - this.graphNodeIdsToRavvOrds = new int[totalLiveVectorsCount]; + final int[] graphNodeIdToDocIds = new int[totalVectorsCount]; + // The graph node id to ravv ordinal mapping, for the leading reader this mapping is always identity, for the other readers we + // only need to map the live vectors + // Therefore the size of this array is the total number of vectors in the leading reader + the number of live vectors in all + // other readers + final int totalLiveVectorsInLeadingReader = liveGraphNodesPerReader[LEADING_READER_IDX].cardinality(); + final int totalLiveVectorsInOtherReaders = totalLiveVectorsCount - totalLiveVectorsInLeadingReader; + this.graphNodeIdsToRavvOrds = new int[totalLiveVectorsInOtherReaders + readers[LEADING_READER_IDX].getFloatVectorValues( + fieldName + ).size()]; + // initialize the graphNodeIdsToRavvOrds with -1 to indicate that there is no mapping + Arrays.fill(graphNodeIdsToRavvOrds, -1); + int documentsIterated = 0; int graphNodeId = 0; - if (deletesFound) { - // If there are deletes, we need to build a new graph from scratch and compact the graph node ids - // TODO: remove this logic once we support incremental graph building with deletes see - // https://github.com/opensearch-project/opensearch-jvector/issues/171 - for (int readerIdx = 0; readerIdx < readers.length; readerIdx++) { - final JVectorFloatVectorValues values = (JVectorFloatVectorValues) readers[readerIdx].getFloatVectorValues(fieldName); - perReaderFloatVectorValues[readerIdx] = values; - // For each vector in this reader - KnnVectorValues.DocIndexIterator it = values.iterator(); - - for (int docId = it.nextDoc(); docId != DocIdSetIterator.NO_MORE_DOCS; docId = it.nextDoc()) { - if (docMaps[readerIdx].get(docId) == -1) { - log.warn( - "Document {} in reader {} is not mapped to a global ordinal from the merge docMaps. Will skip this document for now", - docId, - readerIdx - ); - } else { - // Mapping from ravv ordinals to [readerIndex, readerOrd] - // Map graph node id to ravv ordinal - // Map graph node id to doc id - final int newGlobalDocId = docMaps[readerIdx].get(docId); - final int ravvLocalOrd = it.index(); - final int ravvGlobalOrd = ravvLocalOrd + baseOrds[readerIdx]; - graphNodeIdToDocIds[graphNodeId] = newGlobalDocId; - graphNodeIdsToRavvOrds[graphNodeId] = ravvGlobalOrd; - graphNodeId++; - ravvOrdToReaderMapping[ravvGlobalOrd][READER_ID] = readerIdx; // Reader index - ravvOrdToReaderMapping[ravvGlobalOrd][READER_ORD] = ravvLocalOrd; // Ordinal in reader - } - - documentsIterated++; - } + // We can reuse the existing graph and simply remap the ravv ordinals to the new global doc ids + // for the leading reader we must preserve the original node Ids and map them to the corresponding ravv vectors originally + // used to build the graph + // This is necessary because we are later going to expand that graph with new vectors from the other readers. + // The leading reader is ALWAYS the first one in the readers array + // Deletes will be applied after the graph is built on heap using the cleanup method + final JVectorFloatVectorValues leadingReaderValues = (JVectorFloatVectorValues) readers[LEADING_READER_IDX] + .getFloatVectorValues(fieldName); + perReaderFloatVectorValues[LEADING_READER_IDX] = leadingReaderValues; + var leadingReaderIt = leadingReaderValues.iterator(); + for (int docId = leadingReaderIt.nextDoc(); docId != DocIdSetIterator.NO_MORE_DOCS; docId = leadingReaderIt.nextDoc()) { + final int newGlobalDocId = docMaps[LEADING_READER_IDX].get(docId); + // we increment anyway even if deleted because we are going to apply cleanup later to remove deleted nodes from the leading + // graph that will be used incremented + graphNodeId++; + if (newGlobalDocId == -1) { + log.debug( + "Document {} in reader {} is not mapped to a global ordinal from the merge docMaps. This means it's deleted, Will skip this document for now", + docId, + LEADING_READER_IDX + ); } - } else { - // If there are no deletes, we can reuse the existing graph and simply remap the ravv ordinals to the new global doc ids - // for the leading reader we must preserve the original node Ids and map them to the corresponding ravv vectors originally - // used to build the graph - // This is necessary because we are later going to expand that graph with new vectors from the other readers. - // The leading reader is ALWAYS the first one in the readers array - final JVectorFloatVectorValues leadingReaderValues = (JVectorFloatVectorValues) readers[LEADING_READER_IDX] - .getFloatVectorValues(fieldName); - perReaderFloatVectorValues[LEADING_READER_IDX] = leadingReaderValues; - var leadingReaderIt = leadingReaderValues.iterator(); - for (int docId = leadingReaderIt.nextDoc(); docId != DocIdSetIterator.NO_MORE_DOCS; docId = leadingReaderIt.nextDoc()) { - final int newGlobalDocId = docMaps[LEADING_READER_IDX].get(docId); - if (newGlobalDocId == -1) { + + final int ravvLocalOrd = leadingReaderIt.index(); + final int ravvGlobalOrd = ravvLocalOrd + baseOrds[LEADING_READER_IDX]; + graphNodeIdToDocIds[ravvLocalOrd] = newGlobalDocId; + graphNodeIdsToRavvOrds[ravvLocalOrd] = ravvGlobalOrd; + ravvOrdToReaderMapping[ravvGlobalOrd][READER_ID] = LEADING_READER_IDX; // Reader index + ravvOrdToReaderMapping[ravvGlobalOrd][READER_ORD] = ravvLocalOrd; // Ordinal in reader + + documentsIterated++; + } + + // For the remaining readers we map the graph node id to the ravv ordinal in the order they appear + for (int readerIdx = 1; readerIdx < readers.length; readerIdx++) { + final JVectorFloatVectorValues values = (JVectorFloatVectorValues) readers[readerIdx].getFloatVectorValues(fieldName); + perReaderFloatVectorValues[readerIdx] = values; + // For each vector in this reader + KnnVectorValues.DocIndexIterator it = values.iterator(); + + for (int docId = it.nextDoc(); docId != DocIdSetIterator.NO_MORE_DOCS; docId = it.nextDoc()) { + if (docMaps[readerIdx].get(docId) == -1) { log.warn( "Document {} in reader {} is not mapped to a global ordinal from the merge docMaps. Will skip this document for now", docId, - LEADING_READER_IDX + readerIdx ); } else { - final int ravvLocalOrd = leadingReaderIt.index(); - final int ravvGlobalOrd = ravvLocalOrd + baseOrds[LEADING_READER_IDX]; - graphNodeIdToDocIds[ravvLocalOrd] = newGlobalDocId; - graphNodeIdsToRavvOrds[ravvLocalOrd] = ravvGlobalOrd; + // Mapping from ravv ordinals to [readerIndex, readerOrd] + // Map graph node id to ravv ordinal + // Map graph node id to doc id + final int newGlobalDocId = docMaps[readerIdx].get(docId); + final int ravvLocalOrd = it.index(); + final int ravvGlobalOrd = ravvLocalOrd + baseOrds[readerIdx]; + graphNodeIdToDocIds[graphNodeId] = newGlobalDocId; + graphNodeIdsToRavvOrds[graphNodeId] = ravvGlobalOrd; graphNodeId++; - ravvOrdToReaderMapping[ravvGlobalOrd][READER_ID] = LEADING_READER_IDX; // Reader index + ravvOrdToReaderMapping[ravvGlobalOrd][READER_ID] = readerIdx; // Reader index ravvOrdToReaderMapping[ravvGlobalOrd][READER_ORD] = ravvLocalOrd; // Ordinal in reader } documentsIterated++; } - - // For the remaining readers we map the graph node id to the ravv ordinal in the order they appear - for (int readerIdx = 1; readerIdx < readers.length; readerIdx++) { - final JVectorFloatVectorValues values = (JVectorFloatVectorValues) readers[readerIdx].getFloatVectorValues(fieldName); - perReaderFloatVectorValues[readerIdx] = values; - // For each vector in this reader - KnnVectorValues.DocIndexIterator it = values.iterator(); - - for (int docId = it.nextDoc(); docId != DocIdSetIterator.NO_MORE_DOCS; docId = it.nextDoc()) { - if (docMaps[readerIdx].get(docId) == -1) { - log.warn( - "Document {} in reader {} is not mapped to a global ordinal from the merge docMaps. Will skip this document for now", - docId, - readerIdx - ); - } else { - // Mapping from ravv ordinals to [readerIndex, readerOrd] - // Map graph node id to ravv ordinal - // Map graph node id to doc id - final int newGlobalDocId = docMaps[readerIdx].get(docId); - final int ravvLocalOrd = it.index(); - final int ravvGlobalOrd = ravvLocalOrd + baseOrds[readerIdx]; - graphNodeIdToDocIds[graphNodeId] = newGlobalDocId; - graphNodeIdsToRavvOrds[graphNodeId] = ravvGlobalOrd; - graphNodeId++; - ravvOrdToReaderMapping[ravvGlobalOrd][READER_ID] = readerIdx; // Reader index - ravvOrdToReaderMapping[ravvGlobalOrd][READER_ORD] = ravvLocalOrd; // Ordinal in reader - } - - documentsIterated++; - } - } } if (documentsIterated < totalVectorsCount) { @@ -857,6 +847,10 @@ public void merge() throws IOException { // Get the leading reader PerFieldKnnVectorsFormat.FieldsReader fieldsReader = (PerFieldKnnVectorsFormat.FieldsReader) readers[LEADING_READER_IDX]; JVectorReader leadingReader = (JVectorReader) fieldsReader.getFieldReader(fieldName); + final RemappedRandomAccessVectorValues remappedRandomAccessVectorValues = new RemappedRandomAccessVectorValues( + this, + graphNodeIdsToRavvOrds + ); final BuildScoreProvider buildScoreProvider; // Check if the leading reader has pre-existing PQ codebooks and if so, refine them with the remaining vectors if (leadingReader.getProductQuantizationForField(fieldInfo.name).isEmpty()) { @@ -873,7 +867,7 @@ public void merge() throws IOException { totalVectorsCount, minimumBatchSizeForQuantization ); - pqVectors = getPQVectors(graphNodeIdsToRavvOrds, this, fieldInfo); + pqVectors = getPQVectors(remappedRandomAccessVectorValues, fieldInfo); } else { log.info( "Not enough vectors found for field: {}, totalVectorCount: {}, is below minimumBatchSizeForQuantization: {}", @@ -914,41 +908,60 @@ public void merge() throws IOException { if (pqVectors == null) { buildScoreProvider = BuildScoreProvider.randomAccessScoreProvider( - this, - graphNodeIdsToRavvOrds, + remappedRandomAccessVectorValues, getVectorSimilarityFunction(fieldInfo) ); - // graph = getGraph(buildScoreProvider, this, newToOldOrds, fieldInfo, segmentWriteState.segmentInfo.name); - if (!deletesFound) { - final String segmentName = segmentWriteState.segmentInfo.name; - log.info( - "No deletes found, and no PQ codebooks found, expanding previous graph with additional vectors for field {} in segment {}", - fieldName, - segmentName - ); - final RandomAccessReader leadingOnHeapGraphReader = leadingReader.getNeighborsScoreCacheForField(fieldName); - final int numBaseVectors = leadingReader.getFloatVectorValues(fieldName).size(); - graph = (OnHeapGraphIndex) GraphIndexBuilder.buildAndMergeNewNodes( + final String segmentName = segmentWriteState.segmentInfo.name; + log.info( + "No pre-existing PQ codebooks found, expanding previous graph with additional vectors for field {} in segment {}", + fieldName, + segmentName + ); + final RandomAccessReader leadingOnHeapGraphReader = leadingReader.getNeighborsScoreCacheForField(fieldName); + final int numBaseVectors = leadingReader.getFloatVectorValues(fieldName).size(); + + var diversityProvider = new VamanaDiversityProvider(buildScoreProvider, alpha); + + try ( + OnHeapGraphIndex leadingGraph = OnHeapGraphIndex.load( leadingOnHeapGraphReader, - this, + remappedRandomAccessVectorValues.dimension(), + degreeOverflow, + diversityProvider + ); + ) { + + GraphIndexBuilder builder = new GraphIndexBuilder( buildScoreProvider, - numBaseVectors, - graphNodeIdsToRavvOrds, + remappedRandomAccessVectorValues.dimension(), + leadingGraph, beamWidth, degreeOverflow, alpha, - hierarchyEnabled - ); - } else { - log.info("Deletes found, and no PQ codebooks found, building new graph from scratch"); - graph = getGraph( - buildScoreProvider, - this, - graphNodeIdsToRavvOrds, - fieldInfo, - segmentWriteState.segmentInfo.name, - SIMD_POOL_MERGE + true, + SIMD_POOL_MERGE, + PARALLELISM_POOL ); + + var vv = remappedRandomAccessVectorValues.threadLocalSupplier(); + + // parallel graph construction from the merge documents Ids + SIMD_POOL_MERGE.submit( + () -> IntStream.range(numBaseVectors, remappedRandomAccessVectorValues.size()).parallel().forEach(ord -> { + builder.addGraphNode(ord, vv.get().getVector(ord)); + }) + ).join(); + + // mark deleted nodes + for (int i = 0; i < numBaseVectors; i++) { + if (!liveGraphNodesPerReader[LEADING_READER_IDX].get(i)) { + builder.markNodeDeleted(i); + } + } + + builder.cleanup(); + + graph = (OnHeapGraphIndex) builder.getGraph(); } } else { log.info("PQ codebooks found, building graph from scratch with PQ vectors"); @@ -957,15 +970,14 @@ public void merge() throws IOException { buildScoreProvider.diversityProviderFor(0); graph = getGraph( buildScoreProvider, - this, - graphNodeIdsToRavvOrds, + remappedRandomAccessVectorValues, fieldInfo, segmentWriteState.segmentInfo.name, SIMD_POOL_MERGE ); } - writeField(fieldInfo, this, pqVectors, graphNodeIdsToRavvOrds, graphNodeIdToDocMap, graph); + writeField(fieldInfo, remappedRandomAccessVectorValues, pqVectors, graphNodeIdToDocMap, graph); } @Override @@ -980,7 +992,7 @@ public int dimension() { @Override public VectorFloat getVector(int ord) { - if (ord < 0 || ord >= totalDocsCount) { + if (ord < 0 || ord >= ravvOrdToReaderMapping.length) { throw new IllegalArgumentException("Ordinal out of bounds: " + ord); } @@ -1010,8 +1022,7 @@ public RandomAccessVectorValues copy() { */ public OnHeapGraphIndex getGraph( BuildScoreProvider buildScoreProvider, - RandomAccessVectorValues randomAccessVectorValues, - int[] newToOldOrds, + RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, FieldInfo fieldInfo, String segmentName, ForkJoinPool SIMD_POOL @@ -1034,12 +1045,12 @@ public OnHeapGraphIndex getGraph( */ final long start = Clock.systemDefaultZone().millis(); final OnHeapGraphIndex graphIndex; - var vv = randomAccessVectorValues.threadLocalSupplier(); + var vv = remappedRandomAccessVectorValues.threadLocalSupplier(); log.info("Building graph from merged float vector"); // parallel graph construction from the merge documents Ids - SIMD_POOL.submit(() -> IntStream.range(0, newToOldOrds.length).parallel().forEach(ord -> { - graphIndexBuilder.addGraphNode(ord, vv.get().getVector(newToOldOrds[ord])); + SIMD_POOL.submit(() -> IntStream.range(0, remappedRandomAccessVectorValues.size()).parallel().forEach(ord -> { + graphIndexBuilder.addGraphNode(ord, vv.get().getVector(ord)); })).join(); graphIndexBuilder.cleanup(); graphIndex = (OnHeapGraphIndex) graphIndexBuilder.getGraph(); diff --git a/src/test/java/org/opensearch/knn/index/codec/jvector/graphNodeIdToDocMapTests.java b/src/test/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMapTests.java similarity index 87% rename from src/test/java/org/opensearch/knn/index/codec/jvector/graphNodeIdToDocMapTests.java rename to src/test/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMapTests.java index 56a29c24..8c3c5959 100644 --- a/src/test/java/org/opensearch/knn/index/codec/jvector/graphNodeIdToDocMapTests.java +++ b/src/test/java/org/opensearch/knn/index/codec/jvector/GraphNodeIdToDocMapTests.java @@ -12,7 +12,7 @@ import java.io.IOException; -public class graphNodeIdToDocMapTests extends LuceneTestCase { +public class GraphNodeIdToDocMapTests extends LuceneTestCase { @Test public void testConstructorWithOrdinalsToDocIds() { @@ -48,10 +48,25 @@ public void testConstructorWithSequentialDocIds() { } } - @Test(expected = IllegalStateException.class) - public void testConstructorThrowsWhenMaxDocsLessThanOrdinals() { - int[] invalidMapping = { -1 }; // This would cause issues in Arrays.stream().max() - new GraphNodeIdToDocMap(invalidMapping); + @Test + public void testConstructorWithDeletedDocuments() { + int[] ordinalsToDocIds = { -1, 2, -1, 6 }; // Mix of deleted (-1) and valid docs + GraphNodeIdToDocMap docMap = new GraphNodeIdToDocMap(ordinalsToDocIds); + + // Test ordinal to doc ID mapping (including deleted docs) + assertEquals(-1, docMap.getLuceneDocId(0)); // deleted + assertEquals(2, docMap.getLuceneDocId(1)); + assertEquals(-1, docMap.getLuceneDocId(2)); // deleted + assertEquals(6, docMap.getLuceneDocId(3)); + + // Test doc ID to ordinal mapping (only valid docs should map back) + assertEquals(1, docMap.getJVectorNodeId(2)); + assertEquals(3, docMap.getJVectorNodeId(6)); + + // Unmapped doc IDs return -1 + assertEquals(-1, docMap.getJVectorNodeId(0)); + assertEquals(-1, docMap.getJVectorNodeId(1)); + assertEquals(-1, docMap.getJVectorNodeId(3)); } @Test