diff --git a/CHANGELOG.md b/CHANGELOG.md index 4c2baadb..9407e7dd 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 ### Bug Fixes * Fix visited docs tracking that led to assertions on AbstractKnnVectorQuery side [238] (https://github.com/opensearch-project/opensearch-jvector/pull/238) +* Revert of support deletes for incremental insertion [240](https://github.com/opensearch-project/opensearch-jvector/pull/240) ### Infrastructure * Upgrade Gradle to 9.2.0 [220] (https://github.com/opensearch-project/opensearch-jvector/pull/222) * Add support for JDK25 [220] (https://github.com/opensearch-project/opensearch-jvector/pull/222) 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 ce92d203..7fff91e1 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,10 +46,7 @@ public GraphNodeIdToDocMap(IndexInput in) throws IOException { for (int ord = 0; ord < size; ord++) { final int docId = in.readVInt(); graphNodeIdsToDocIds[ord] = docId; - if (docId != -1) { - // ignore deleted documents - docIdsToGraphNodeIds[docId] = ord; - } + docIdsToGraphNodeIds[docId] = ord; } } @@ -71,19 +68,18 @@ 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 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, + "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, 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 34f28d6d..5d25622d 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,7 +36,6 @@ 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 59c37223..da7196d9 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 @@ -173,8 +173,7 @@ public void search(String field, float[] target, KnnCollector knnCollector, Acce if (acceptDocs == null) compatibleBits = ord -> true; else { Bits b = acceptDocs.bits(); - compatibleBits = ord -> b == null - || (jvectorLuceneDocMap.getLuceneDocId(ord) != -1 && b.get(jvectorLuceneDocMap.getLuceneDocId(ord))); + compatibleBits = ord -> b == null || b.get(jvectorLuceneDocMap.getLuceneDocId(ord)); } try (var graphSearcher = new GraphSearcher(index)) { 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 222d837f..492d8a01 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 @@ -15,6 +15,7 @@ import io.github.jbellis.jvector.graph.similarity.BuildScoreProvider; import io.github.jbellis.jvector.quantization.PQVectors; import io.github.jbellis.jvector.quantization.ProductQuantization; +import io.github.jbellis.jvector.util.PhysicalCoreExecutor; import io.github.jbellis.jvector.vector.VectorizationProvider; import io.github.jbellis.jvector.vector.types.VectorFloat; import io.github.jbellis.jvector.vector.types.VectorTypeSupport; @@ -32,7 +33,6 @@ 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; @@ -47,7 +47,8 @@ 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.*; +import static org.opensearch.knn.index.codec.jvector.JVectorFormat.SIMD_POOL_FLUSH; +import static org.opensearch.knn.index.codec.jvector.JVectorFormat.SIMD_POOL_MERGE; /** * JVectorWriter is responsible for writing vector data into index segments using the JVector library. @@ -207,27 +208,23 @@ 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 (remappedRandomAccessVectorValues.size() >= minimumBatchSizeForQuantization) { + if (randomAccessVectorValues.size() >= minimumBatchSizeForQuantization) { log.info("Calculating codebooks and compressed vectors for field {}", fieldInfo.name); - pqVectors = getPQVectors(remappedRandomAccessVectorValues, fieldInfo); + pqVectors = getPQVectors(newToOldOrds, randomAccessVectorValues, 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.", - remappedRandomAccessVectorValues.size(), + randomAccessVectorValues.size(), minimumBatchSizeForQuantization, fieldInfo.name ); pqVectors = null; buildScoreProvider = BuildScoreProvider.randomAccessScoreProvider( - remappedRandomAccessVectorValues, + randomAccessVectorValues, getVectorSimilarityFunction(fieldInfo) ); } @@ -244,30 +241,39 @@ public void flush(int maxDoc, Sorter.DocMap sortMap) throws IOException { OnHeapGraphIndex graph = getGraph( buildScoreProvider, - remappedRandomAccessVectorValues, + randomAccessVectorValues, + newToOldOrds, fieldInfo, segmentWriteState.segmentInfo.name, SIMD_POOL_FLUSH ); - writeField(field.fieldInfo, remappedRandomAccessVectorValues, pqVectors, graphNodeIdToDocMap, graph); + writeField(field.fieldInfo, field.randomAccessVectorValues, pqVectors, newToOldOrds, graphNodeIdToDocMap, graph); } } private void writeField( FieldInfo fieldInfo, - RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, + RandomAccessVectorValues randomAccessVectorValues, PQVectors pqVectors, + int[] newToOldOrds, GraphNodeIdToDocMap graphNodeIdToDocMap, OnHeapGraphIndex graph ) throws IOException { log.info( "Writing field {} with vector count: {}, for segment: {}", fieldInfo.name, - remappedRandomAccessVectorValues.size(), + randomAccessVectorValues.size(), segmentWriteState.segmentInfo.name ); - final var vectorIndexFieldMetadata = writeGraph(graph, remappedRandomAccessVectorValues, fieldInfo, pqVectors, graphNodeIdToDocMap); + final var vectorIndexFieldMetadata = writeGraph( + graph, + randomAccessVectorValues, + fieldInfo, + pqVectors, + newToOldOrds, + graphNodeIdToDocMap + ); meta.writeInt(fieldInfo.number); vectorIndexFieldMetadata.toOutput(meta); @@ -300,16 +306,17 @@ private void writeField( /** * Writes the graph and PQ codebooks and compressed vectors to the vector index file * @param graph graph - * @param remappedRandomAccessVectorValues random access vector values with remapped ordinals + * @param randomAccessVectorValues random access vector values * @param fieldInfo field info * @return Tuple of start offset and length of the graph * @throws IOException IOException */ private VectorIndexFieldMetadata writeGraph( OnHeapGraphIndex graph, - RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, + RandomAccessVectorValues randomAccessVectorValues, FieldInfo fieldInfo, PQVectors pqVectors, + int[] newToOldOrds, GraphNodeIdToDocMap graphNodeIdToDocMap ) throws IOException { // field data file, which contains the graph @@ -334,17 +341,17 @@ private VectorIndexFieldMetadata writeGraph( .fieldNumber(fieldInfo.number) .vectorEncoding(fieldInfo.getVectorEncoding()) .vectorSimilarityFunction(fieldInfo.getVectorSimilarityFunction()) - .vectorDimension(remappedRandomAccessVectorValues.dimension()) + .vectorDimension(randomAccessVectorValues.dimension()) .graphNodeIdToDocMap(graphNodeIdToDocMap); try ( var writer = new OnDiskSequentialGraphIndexWriter.Builder(graph, jVectorIndexWriter).with( - new InlineVectors(remappedRandomAccessVectorValues.dimension()) + new InlineVectors(randomAccessVectorValues.dimension()) ).build() ) { var suppliers = Feature.singleStateFactory( FeatureId.INLINE_VECTORS, - nodeId -> new InlineVectors.State(remappedRandomAccessVectorValues.getVector(nodeId)) + nodeId -> new InlineVectors.State(randomAccessVectorValues.getVector(newToOldOrds[nodeId])) ); writer.write(suppliers); long endGraphOffset = jVectorIndexWriter.position(); @@ -356,7 +363,7 @@ private VectorIndexFieldMetadata writeGraph( log.info( "Writing PQ codebooks and vectors for field {} since the size is {} >= {}", fieldInfo.name, - remappedRandomAccessVectorValues.size(), + randomAccessVectorValues.size(), minimumBatchSizeForQuantization ); resultBuilder.pqCodebooksAndVectorsOffset(endGraphOffset); @@ -374,17 +381,17 @@ private VectorIndexFieldMetadata writeGraph( } } - private PQVectors getPQVectors(RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, FieldInfo fieldInfo) + private PQVectors getPQVectors(int[] newToOldOrds, RandomAccessVectorValues randomAccessVectorValues, FieldInfo fieldInfo) throws IOException { final String fieldName = fieldInfo.name; final VectorSimilarityFunction vectorSimilarityFunction = fieldInfo.getVectorSimilarityFunction(); - log.info("Computing PQ codebooks for field {} for {} vectors", fieldName, remappedRandomAccessVectorValues.size()); + log.info("Computing PQ codebooks for field {} for {} vectors", fieldName, randomAccessVectorValues.size()); final long start = Clock.systemDefaultZone().millis(); - final var M = numberOfSubspacesPerVectorSupplier.apply(remappedRandomAccessVectorValues.dimension()); - final var numberOfClustersPerSubspace = Math.min(256, remappedRandomAccessVectorValues.size()); // number of centroids per + final var M = numberOfSubspacesPerVectorSupplier.apply(randomAccessVectorValues.dimension()); + final var numberOfClustersPerSubspace = Math.min(256, randomAccessVectorValues.size()); // number of centroids per // subspace ProductQuantization pq = ProductQuantization.compute( - remappedRandomAccessVectorValues, + randomAccessVectorValues, M, // number of subspaces numberOfClustersPerSubspace, // number of centroids per subspace vectorSimilarityFunction == VectorSimilarityFunction.EUCLIDEAN, // center the dataset @@ -397,14 +404,9 @@ private PQVectors getPQVectors(RemappedRandomAccessVectorValues remappedRandomAc 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, remappedRandomAccessVectorValues.size()); + log.info("Encoding and building PQ vectors for field {} for {} vectors", fieldName, randomAccessVectorValues.size()); // PQVectors pqVectors = pq.encodeAll(randomAccessVectorValues, SIMD_POOL); - PQVectors pqVectors = PQVectors.encodeAndBuild( - pq, - remappedRandomAccessVectorValues.size(), - remappedRandomAccessVectorValues, - SIMD_POOL_MERGE - ); + PQVectors pqVectors = PQVectors.encodeAndBuild(pq, newToOldOrds.length, newToOldOrds, randomAccessVectorValues, SIMD_POOL_MERGE); log.info( "Encoded and built PQ vectors for field {}, original size: {} bytes, compressed size: {} bytes", fieldName, @@ -609,7 +611,7 @@ class RandomAccessMergedFloatVectorValues implements RandomAccessVectorValues { private final MergeState mergeState; private final GraphNodeIdToDocMap graphNodeIdToDocMap; private final int[] graphNodeIdsToRavvOrds; - private final FixedBitSet[] liveGraphNodesPerReader; + private boolean deletesFound = false; /** * Creates a random access view over merged float vector values. @@ -634,19 +636,9 @@ 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++) { @@ -665,8 +657,8 @@ public RandomAccessMergedFloatVectorValues(FieldInfo fieldInfo, MergeState merge if (liveDocs[i] == null || liveDocs[i].get(it.docID())) { liveVectorCountInReader++; } else { - // This vector is deleted so we need to mark it as deleted in the liveGraphNodesPerReader - liveGraphNodesPerReader[i].clear(it.index()); + deletedOrds[i]++; + deletesFound = true; } } if (liveVectorCountInReader >= vectorsCountInLeadingReader) { @@ -706,10 +698,6 @@ 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]; @@ -718,87 +706,111 @@ 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[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); + final int[] graphNodeIdToDocIds = new int[totalLiveVectorsCount]; + this.graphNodeIdsToRavvOrds = new int[totalLiveVectorsCount]; - int documentsIterated = 0; int graphNodeId = 0; - // 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 - ); - } - - 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(); + 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 + } - for (int docId = it.nextDoc(); docId != DocIdSetIterator.NO_MORE_DOCS; docId = it.nextDoc()) { - if (docMaps[readerIdx].get(docId) == -1) { + documentsIterated++; + } + } + } 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) { log.warn( "Document {} in reader {} is not mapped to a global ordinal from the merge docMaps. Will skip this document for now", docId, - readerIdx + LEADING_READER_IDX ); } 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; + final int ravvLocalOrd = leadingReaderIt.index(); + final int ravvGlobalOrd = ravvLocalOrd + baseOrds[LEADING_READER_IDX]; + graphNodeIdToDocIds[ravvLocalOrd] = newGlobalDocId; + graphNodeIdsToRavvOrds[ravvLocalOrd] = ravvGlobalOrd; graphNodeId++; - ravvOrdToReaderMapping[ravvGlobalOrd][READER_ID] = readerIdx; // Reader index + 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, + 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) { @@ -848,10 +860,6 @@ 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()) { @@ -868,7 +876,7 @@ public void merge() throws IOException { totalVectorsCount, minimumBatchSizeForQuantization ); - pqVectors = getPQVectors(remappedRandomAccessVectorValues, fieldInfo); + pqVectors = getPQVectors(graphNodeIdsToRavvOrds, this, fieldInfo); } else { log.info( "Not enough vectors found for field: {}, totalVectorCount: {}, is below minimumBatchSizeForQuantization: {}", @@ -909,60 +917,41 @@ public void merge() throws IOException { if (pqVectors == null) { buildScoreProvider = BuildScoreProvider.randomAccessScoreProvider( - remappedRandomAccessVectorValues, + this, + graphNodeIdsToRavvOrds, getVectorSimilarityFunction(fieldInfo) ); - 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, - remappedRandomAccessVectorValues.dimension(), - degreeOverflow, - diversityProvider + // 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 ); - ) { - - GraphIndexBuilder builder = new GraphIndexBuilder( + final RandomAccessReader leadingOnHeapGraphReader = leadingReader.getNeighborsScoreCacheForField(fieldName); + final int numBaseVectors = leadingReader.getFloatVectorValues(fieldName).size(); + graph = (OnHeapGraphIndex) buildAndMergeNewNodes( + leadingOnHeapGraphReader, + this, buildScoreProvider, - remappedRandomAccessVectorValues.dimension(), - leadingGraph, + numBaseVectors, + graphNodeIdsToRavvOrds, beamWidth, degreeOverflow, alpha, - true, - SIMD_POOL_MERGE, - PARALLELISM_POOL + 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 ); - - 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"); @@ -971,14 +960,15 @@ public void merge() throws IOException { buildScoreProvider.diversityProviderFor(0); graph = getGraph( buildScoreProvider, - remappedRandomAccessVectorValues, + this, + graphNodeIdsToRavvOrds, fieldInfo, segmentWriteState.segmentInfo.name, SIMD_POOL_MERGE ); } - writeField(fieldInfo, remappedRandomAccessVectorValues, pqVectors, graphNodeIdToDocMap, graph); + writeField(fieldInfo, this, pqVectors, graphNodeIdsToRavvOrds, graphNodeIdToDocMap, graph); } @Override @@ -993,7 +983,7 @@ public int dimension() { @Override public VectorFloat getVector(int ord) { - if (ord < 0 || ord >= ravvOrdToReaderMapping.length) { + if (ord < 0 || ord >= totalDocsCount) { throw new IllegalArgumentException("Ordinal out of bounds: " + ord); } @@ -1023,7 +1013,8 @@ public RandomAccessVectorValues copy() { */ public OnHeapGraphIndex getGraph( BuildScoreProvider buildScoreProvider, - RemappedRandomAccessVectorValues remappedRandomAccessVectorValues, + RandomAccessVectorValues randomAccessVectorValues, + int[] newToOldOrds, FieldInfo fieldInfo, String segmentName, ForkJoinPool SIMD_POOL @@ -1046,12 +1037,12 @@ public OnHeapGraphIndex getGraph( */ final long start = Clock.systemDefaultZone().millis(); final OnHeapGraphIndex graphIndex; - var vv = remappedRandomAccessVectorValues.threadLocalSupplier(); + var vv = randomAccessVectorValues.threadLocalSupplier(); log.info("Building graph from merged float vector"); // parallel graph construction from the merge documents Ids - SIMD_POOL.submit(() -> IntStream.range(0, remappedRandomAccessVectorValues.size()).parallel().forEach(ord -> { - graphIndexBuilder.addGraphNode(ord, vv.get().getVector(ord)); + SIMD_POOL.submit(() -> IntStream.range(0, newToOldOrds.length).parallel().forEach(ord -> { + graphIndexBuilder.addGraphNode(ord, vv.get().getVector(newToOldOrds[ord])); })).join(); graphIndexBuilder.cleanup(); graphIndex = (OnHeapGraphIndex) graphIndexBuilder.getGraph(); @@ -1106,4 +1097,43 @@ public RandomAccessVectorValues copy() { } } + static ImmutableGraphIndex buildAndMergeNewNodes( + RandomAccessReader in, + RandomAccessVectorValues newVectors, + BuildScoreProvider buildScoreProvider, + int startingNodeOffset, + int[] graphToRavvOrdMap, + int beamWidth, + float overflowRatio, + float alpha, + boolean addHierarchy + ) throws IOException { + + var diversityProvider = new VamanaDiversityProvider(buildScoreProvider, alpha); + + try (OnHeapGraphIndex graph = OnHeapGraphIndex.load(in, newVectors.dimension(), overflowRatio, diversityProvider);) { + + GraphIndexBuilder builder = new GraphIndexBuilder( + buildScoreProvider, + newVectors.dimension(), + graph, + beamWidth, + overflowRatio, + alpha, + true, + PhysicalCoreExecutor.pool(), + ForkJoinPool.commonPool() + ); + + var vv = newVectors.threadLocalSupplier(); + + // parallel graph construction from the merge documents Ids + PhysicalCoreExecutor.pool().submit(() -> IntStream.range(startingNodeOffset, newVectors.size()).parallel().forEach(ord -> { + builder.addGraphNode(ord, vv.get().getVector(graphToRavvOrdMap[ord])); + })).join(); + + builder.cleanup(); + return builder.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 index 8c3c5959..ac9947c3 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 @@ -48,25 +48,10 @@ public void testConstructorWithSequentialDocIds() { } } - @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(expected = IllegalStateException.class) + public void testConstructorThrowsWhenMaxDocsLessThanOrdinals() { + int[] invalidMapping = { -1 }; // This would cause issues in Arrays.stream().max() + new GraphNodeIdToDocMap(invalidMapping); } @Test