diff --git a/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java b/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java index 96e31431f8..d67f2165d8 100644 --- a/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java +++ b/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java @@ -8,6 +8,7 @@ import java.io.InputStream; import java.io.OutputStream; import java.nio.file.Path; +import java.util.BitSet; import java.util.Objects; /** @@ -192,6 +193,13 @@ public StandardDataset() {} */ long getGraphDegree(); + /** + * Returns the number of vectors in this index. + * + * @return the number of rows of the indexed dataset + */ + long size(); + /** * A method to persist a CAGRA index using an instance of {@link OutputStream} * for writing index bytes. @@ -310,7 +318,7 @@ static Builder newBuilder(CuVSResources cuvsResources) { * @throws Throwable if an error occurs during the merge operation */ static CagraIndex merge(CagraIndex[] indexes) throws Throwable { - return merge(indexes, null); + return merge(indexes, null, null); } /** @@ -322,6 +330,28 @@ static CagraIndex merge(CagraIndex[] indexes) throws Throwable { * @throws Throwable if an error occurs during the merge operation */ static CagraIndex merge(CagraIndex[] indexes, CagraIndexParams mergeParams) throws Throwable { + return merge(indexes, mergeParams, null); + } + + /** + * Merges multiple CAGRA indexes into a single index, keeping only the rows selected by + * {@code rowFilter}. + * + *

The merge concatenates the input datasets in the order the indexes are given, so bit + * {@code i} of the filter refers to row {@code i} of that concatenation: bits {@code 0} to + * {@code indexes[0].size() - 1} address the first index, the bits that follow address the second, + * and so on. A set bit keeps the row; a clear bit drops it. The rows that survive keep + * their relative order and are packed together, so the merged index has one row per set bit. + * + * @param indexes Array of CAGRA indexes to merge + * @param mergeParams Parameters to control the merge operation, or null to use defaults + * @param rowFilter The rows to keep, or null to keep all of them + * @return A new merged CAGRA index + * @throws IllegalArgumentException if {@code rowFilter} has a bit set beyond the last row + * @throws Throwable if an error occurs during the merge operation + */ + static CagraIndex merge(CagraIndex[] indexes, CagraIndexParams mergeParams, BitSet rowFilter) + throws Throwable { if (indexes == null || indexes.length == 0) { throw new IllegalArgumentException("At least one index must be provided for merging"); } @@ -333,7 +363,7 @@ static CagraIndex merge(CagraIndex[] indexes, CagraIndexParams mergeParams) thro } } - return CuVSProvider.provider().mergeCagraIndexes(indexes, mergeParams); + return CuVSProvider.provider().mergeCagraIndexes(indexes, mergeParams, rowFilter); } /** diff --git a/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/CuVSProvider.java b/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/CuVSProvider.java index 670f44eda6..d6d87720b8 100644 --- a/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/CuVSProvider.java +++ b/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/CuVSProvider.java @@ -9,6 +9,7 @@ import java.lang.invoke.MethodType; import java.nio.file.Path; import java.time.Duration; +import java.util.BitSet; import java.util.List; /** @@ -165,29 +166,6 @@ HnswIndex hnswIndexBuild(CuVSResources resources, HnswIndexParams hnswParams, Cu TieredIndex.Builder newTieredIndexBuilder(CuVSResources cuVSResources) throws UnsupportedOperationException; - /** - * Merges multiple CAGRA indexes into a single index. - * - * @param indexes Array of CAGRA indexes to merge - * @return A new merged CAGRA index - * @throws Throwable if an error occurs during the merge operation - */ - CagraIndex mergeCagraIndexes(CagraIndex[] indexes) throws Throwable; - - /** - * Merges multiple CAGRA indexes into a single index with the specified merge parameters. - * - * @param indexes Array of CAGRA indexes to merge - * @param mergeParams Parameters to control the merge operation, or null to use defaults - * @return A new merged CAGRA index - * @throws Throwable if an error occurs during the merge operation - */ - default CagraIndex mergeCagraIndexes(CagraIndex[] indexes, CagraIndexParams mergeParams) - throws Throwable { - // Default implementation falls back to the method without parameters - return mergeCagraIndexes(indexes); - } - /** * Reports whether the rows of {@code dataset} already sit at the row stride CAGRA requires, which * is the row length in bytes rounded up to a 16 byte boundary. @@ -207,6 +185,20 @@ default boolean isCagraPaddedDataset(CuVSMatrix dataset) { "Padded layout detection is not supported by " + getClass().getName()); } + /** + * Merges multiple CAGRA indexes into a single index, keeping only the rows selected by + * {@code rowFilter}. See {@link CagraIndex#merge(CagraIndex[], CagraIndexParams, BitSet)} for the + * meaning of the filter. + * + * @param indexes Array of CAGRA indexes to merge + * @param mergeParams Parameters to control the merge operation, or null to use defaults + * @param rowFilter The rows to keep, or null to keep all of them + * @return A new merged CAGRA index + * @throws Throwable if an error occurs during the merge operation + */ + CagraIndex mergeCagraIndexes(CagraIndex[] indexes, CagraIndexParams mergeParams, BitSet rowFilter) + throws Throwable; + /** * Creates a device-backed multi-partition filter handle from the pre-packed combined bitset. * Per-partition bit offsets are recomputed inside cuVS from the index sizes. diff --git a/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/UnsupportedProvider.java b/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/UnsupportedProvider.java index 66b0b8cc47..fd1cf7746c 100644 --- a/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/UnsupportedProvider.java +++ b/java/cuvs-java/src/main/java/com/nvidia/cuvs/spi/UnsupportedProvider.java @@ -8,6 +8,7 @@ import java.lang.invoke.MethodHandle; import java.nio.file.Path; import java.time.Duration; +import java.util.BitSet; import java.util.List; import java.util.logging.Level; @@ -81,12 +82,13 @@ public TieredIndex.Builder newTieredIndexBuilder(CuVSResources cuVSResources) { } @Override - public CagraIndex mergeCagraIndexes(CagraIndex[] indexes) { + public boolean isCagraPaddedDataset(CuVSMatrix dataset) { throw new UnsupportedOperationException(reasons); } @Override - public boolean isCagraPaddedDataset(CuVSMatrix dataset) { + public CagraIndex mergeCagraIndexes( + CagraIndex[] indexes, CagraIndexParams mergeParams, BitSet rowFilter) { throw new UnsupportedOperationException(reasons); } diff --git a/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java b/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java index 9506b130cd..bfa2d6a148 100644 --- a/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java +++ b/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java @@ -92,7 +92,7 @@ private CagraIndexImpl( * Used primarily for the merge operation. * * @param indexReference The reference to the existing index - * @param resources The resources instance + * @param resources The resources instance */ private CagraIndexImpl(IndexReference indexReference, CuVSResources resources) { this.resources = resources; @@ -103,10 +103,10 @@ private CagraIndexImpl(IndexReference indexReference, CuVSResources resources) { /** * Constructor for creating an index from a pre-build CAGRA graph * - * @param metric the distance type used - * @param graph a previously built CAGRA graph - * @param dataset the dataset used for indexing - * @param resources an instance of {@link CuVSResources} + * @param metric the distance type used + * @param graph a previously built CAGRA graph + * @param dataset the dataset used for indexing + * @param resources an instance of {@link CuVSResources} */ private CagraIndexImpl( CagraIndexParams.CuvsDistanceType metric, @@ -221,11 +221,14 @@ private IndexReference build(CagraIndexParams indexParameters, CuVSMatrixInterna private static MemorySegment createCagraIndex() { try (var localArena = Arena.ofConfined()) { MemorySegment indexPtrPtr = localArena.allocate(cuvsCagraIndex_t); - // cuvsCagraIndexCreate gets a pointer to a cuvsCagraIndex_t, which is defined as a pointer to + // cuvsCagraIndexCreate gets a pointer to a cuvsCagraIndex_t, which is defined + // as a pointer to // cuvsCagraIndex. - // It's basically an "out" parameter: the C functions will create the index and "return back" + // It's basically an "out" parameter: the C functions will create the index and + // "return back" // a pointer to it: (*index = new cuvsCagraIndex{}; - // The "out parameter" pointer is needed only for the duration of the function invocation (it + // The "out parameter" pointer is needed only for the duration of the function + // invocation (it // could be a stack pointer, in C) so we allocate it from our localArena var returnValue = cuvsCagraIndexCreate(indexPtrPtr); checkCuVSError(returnValue, "cuvsCagraIndexCreate"); @@ -387,7 +390,8 @@ public SearchResults search(CagraQuery query) throws Throwable { prefilter); checkCuVSError(returnValue, "cuvsCagraSearch"); - // TODO: we can avoid/defer this using CuVSDeviceMatrix for neighborsDP and distancesDP + // TODO: we can avoid/defer this using CuVSDeviceMatrix for neighborsDP and + // distancesDP // TODO: also, should we use cuvsMatrixCopy instead? Util.cudaMemcpyAsync( neighborsMemorySegment, @@ -417,7 +421,10 @@ public SearchResults search(CagraQuery query) throws Throwable { } } - /** Returns the underlying {@code cuvsCagraIndex_t} handle for native-side index passing. */ + /** + * Returns the underlying {@code cuvsCagraIndex_t} handle for native-side index + * passing. + */ public MemorySegment getIndexHandle() { return cagraIndexReference.getMemorySegment(); } @@ -604,6 +611,18 @@ public long getGraphDegree() { } } + @Override + public long size() { + checkNotDestroyed(); + try (var localArena = Arena.ofConfined()) { + MemorySegment size = localArena.allocate(int64_t); + checkCuVSError( + cuvsCagraIndexGetSize(cagraIndexReference.getMemorySegment(), size), + "cuvsCagraIndexGetSize"); + return size.get(int64_t, 0); + } + } + private IndexReference fromGraph( CagraIndexParams.CuvsDistanceType metric, CuVSMatrixInternal graph, @@ -789,7 +808,8 @@ CuVSMatrix getDatasetForConversion() { } /** - * Allocates the native CagraIndexParams data structures and fills the configured index parameters in. + * Allocates the native CagraIndexParams data structures and fills the + * configured index parameters in. */ private static CloseableHandle segmentFromIndexParams(CagraIndexParams params) { var handles = new ArrayList(); @@ -923,72 +943,157 @@ public static CagraIndex.Builder newBuilder(CuVSResources cuvsResources) { } /** - * Merges multiple CAGRA indexes into a single index. - * - * @param indexes Array of CAGRA indexes to merge - * @return A new merged CAGRA index - */ - public static CagraIndex merge(CagraIndex[] indexes) { - return merge(indexes, null); - } - - /** - * Merges multiple CAGRA indexes into a single index with specified merge parameters. + * Merges multiple CAGRA indexes into a single index, keeping only the rows selected by + * {@code rowFilter}. See {@link CagraIndex#merge(CagraIndex[], CagraIndexParams, BitSet)} for the + * meaning of the filter. * - * @param indexes Array of CAGRA indexes to merge + * @param indexes Array of CAGRA indexes to merge * @param mergeParams Parameters to control the merge operation, or null to use defaults + * @param rowFilter The rows to keep, or null to keep all of them * @return A new merged CAGRA index */ - public static CagraIndex merge(CagraIndex[] indexes, CagraIndexParams mergeParams) { + public static CagraIndex merge( + CagraIndex[] indexes, CagraIndexParams mergeParams, BitSet rowFilter) { + if (indexes == null || indexes.length == 0) { + throw new IllegalArgumentException("At least one index must be provided for merging"); + } CuVSResources resources = indexes[0].getCuVSResources(); - var mergedIndex = createCagraIndex(); + for (int i = 1; i < indexes.length; i++) { + if (!resources.equals(indexes[i].getCuVSResources())) { + throw new IllegalArgumentException("All indexes must use the same CuVSResources instance"); + } + } try (var localArena = Arena.ofConfined()) { MemorySegment indexesSegment = localArena.allocate(indexes.length * ValueLayout.ADDRESS.byteSize()); + long mergedRowCount = 0; for (int i = 0; i < indexes.length; i++) { CagraIndexImpl indexImpl = (CagraIndexImpl) indexes[i]; indexesSegment.setAtIndex( ValueLayout.ADDRESS, i, indexImpl.cagraIndexReference.getMemorySegment()); + if (rowFilter != null) { + mergedRowCount += indexImpl.size(); + } + } + if (rowFilter != null) { + if (rowFilter.length() > mergedRowCount) { + throw new IllegalArgumentException( + "rowFilter selects row " + + (rowFilter.length() - 1) + + " but the indexes only hold " + + mergedRowCount + + " rows"); + } + if (rowFilter.isEmpty()) { + throw new IllegalArgumentException("rowFilter keeps no rows, there is nothing to merge"); + } } + var mergedIndex = createCagraIndex(); + CagraIndexImpl merged = null; try (var nativeMergeParams = segmentFromIndexParams(mergeParams); var resourcesAccessor = resources.access()) { var cuvsRes = resourcesAccessor.handle(); + // The words the merge filter points at have to outlive the merge call, so the + // allocation is held open around it rather than inside the helper that fills + // the filter in. MemorySegment mergeFilter = cuvsFilter.allocate(localArena); - cuvsFilter.type(mergeFilter, 0); // NO_FILTER - cuvsFilter.addr(mergeFilter, 0); - - MemorySegment mergedDatasetPtr = localArena.allocate(cuvsDataset_t); - checkCuVSError(cuvsDatasetCreate(mergedDatasetPtr), "cuvsDatasetCreate"); - MemorySegment mergedDataset = mergedDatasetPtr.get(cuvsDataset_t, 0); - AutoCloseable datasetOwner = new DatasetCloseDelegate(mergedDataset); - try { - checkCuVSError( - cuvsCagraMerge( - cuvsRes, - nativeMergeParams.handle(), - indexesSegment, - indexes.length, - mergeFilter, - mergedDataset, - mergedIndex), - "cuvsCagraMerge"); - return new CagraIndexImpl(new IndexReference(mergedIndex, null, datasetOwner), resources); - } catch (Throwable e) { + try (@SuppressWarnings("unused") + var filterWords = + allocateRowFilter(cuvsRes, localArena, mergeFilter, rowFilter, mergedRowCount)) { + MemorySegment mergedDatasetPtr = localArena.allocate(cuvsDataset_t); + checkCuVSError(cuvsDatasetCreate(mergedDatasetPtr), "cuvsDatasetCreate"); + MemorySegment mergedDataset = mergedDatasetPtr.get(cuvsDataset_t, 0); + AutoCloseable datasetOwner = new DatasetCloseDelegate(mergedDataset); try { - datasetOwner.close(); - } catch (Exception closeError) { - e.addSuppressed(closeError); + checkCuVSError( + cuvsCagraMerge( + cuvsRes, + nativeMergeParams.handle(), + indexesSegment, + indexes.length, + mergeFilter, + mergedDataset, + mergedIndex), + "cuvsCagraMerge"); + merged = + new CagraIndexImpl(new IndexReference(mergedIndex, null, datasetOwner), resources); + return merged; + } catch (Throwable e) { + try { + datasetOwner.close(); + } catch (Exception closeError) { + e.addSuppressed(closeError); + } + throw e; + } + } + } catch (Throwable t) { + try { + if (merged != null) { + // The merged index owns the dataset by now, so close it rather than only destroying + // the handle. + merged.close(); + } else { + checkCuVSError(cuvsCagraIndexDestroy(mergedIndex), "cuvsCagraIndexDestroy"); } - throw e; + } catch (Throwable cleanupError) { + t.addSuppressed(cleanupError); } + throw t; } } } + /** + * Fills {@code mergeFilter} in and returns the device allocation backing it, + * which the caller has + * to keep open until the merge returns. A null {@code rowFilter} produces a + * NO_FILTER and an + * empty allocation. + * + *

cuvs reads the bitset as a vector of 32 bit words covering + * {@code mergedRowCount} rows, and derives the row count of the merged index from + * the number of bits that are set, so the words have to cover every row rather + * than stop at the last one that survives. + */ + private static CloseableRMMAllocation allocateRowFilter( + long cuvsRes, Arena arena, MemorySegment mergeFilter, BitSet rowFilter, long mergedRowCount) { + if (rowFilter == null) { + cuvsFilter.type(mergeFilter, NO_FILTER()); + cuvsFilter.addr(mergeFilter, 0); + return CloseableRMMAllocation.EMPTY; + } + + long words = (mergedRowCount + 31) / 32; + long bytes = C_INT_BYTE_SIZE * words; + MemorySegment hostWords = + buildMemorySegment(arena, rowFilter.toLongArray(), (mergedRowCount + 63) / 64); + + var deviceWords = allocateRMMSegment(cuvsRes, bytes); + try { + Util.cudaMemcpyAsync( + deviceWords.handle(), hostWords, bytes, HOST_TO_DEVICE, Util.getStream(cuvsRes)); + checkCuVSError(cuvsStreamSync(cuvsRes), "cuvsStreamSync"); + + MemorySegment filterTensor = + prepareTensor(arena, deviceWords.handle(), new long[] {words}, kDLUInt(), 32, kDLCUDA()); + cuvsFilter.type(mergeFilter, BITSET()); + cuvsFilter.addr(mergeFilter, filterTensor.address()); + return deviceWords; + } catch (Throwable t) { + try { + deviceWords.close(); + } catch (Exception closeError) { + t.addSuppressed(closeError); + } + throw t; + } + } + /** * Builder helps configure and create an instance of {@link CagraIndex}. */ @@ -1081,8 +1186,9 @@ public static class IndexReference { * @param indexMemorySegment the MemorySegment instance to use for containing * index reference * @param dataset the dataset used for indexing; the dataset lifetime - * matches the lifetime of the index, we need to keep a reference - * to it so we can close it when the index is closed. + * matches the lifetime of the index, we need to keep + * a reference to it so we can close it when the index + * is closed. * Can be null (e.g. from deserialization or merging) */ private IndexReference(MemorySegment indexMemorySegment, CuVSMatrix dataset) { diff --git a/java/cuvs-java/src/main/java22/com/nvidia/cuvs/spi/JDKProvider.java b/java/cuvs-java/src/main/java22/com/nvidia/cuvs/spi/JDKProvider.java index 9cc4a5499c..1f16a4e904 100644 --- a/java/cuvs-java/src/main/java22/com/nvidia/cuvs/spi/JDKProvider.java +++ b/java/cuvs-java/src/main/java22/com/nvidia/cuvs/spi/JDKProvider.java @@ -26,6 +26,7 @@ import java.nio.file.Files; import java.nio.file.Path; import java.time.Duration; +import java.util.BitSet; import java.util.List; import java.util.Locale; import java.util.Objects; @@ -302,24 +303,14 @@ public TieredIndex.Builder newTieredIndexBuilder(CuVSResources cuVSResources) { } @Override - public CagraIndex mergeCagraIndexes(CagraIndex[] indexes) { - if (indexes == null || indexes.length == 0) { - throw new IllegalArgumentException("At least one index must be provided for merging"); - } - return CagraIndexImpl.merge(indexes); - } - - @Override - public CagraIndex mergeCagraIndexes(CagraIndex[] indexes, CagraIndexParams mergeParams) { - if (indexes == null || indexes.length == 0) { - throw new IllegalArgumentException("At least one index must be provided for merging"); - } - return CagraIndexImpl.merge(indexes, mergeParams); + public boolean isCagraPaddedDataset(CuVSMatrix dataset) { + return CagraIndexImpl.isPaddedDataset(dataset); } @Override - public boolean isCagraPaddedDataset(CuVSMatrix dataset) { - return CagraIndexImpl.isPaddedDataset(dataset); + public CagraIndex mergeCagraIndexes( + CagraIndex[] indexes, CagraIndexParams mergeParams, BitSet rowFilter) { + return CagraIndexImpl.merge(indexes, mergeParams, rowFilter); } @Override diff --git a/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java b/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java index e2287c0a22..22f6fa899b 100644 --- a/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java +++ b/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java @@ -1006,6 +1006,259 @@ public void testMergingIndexes() throws Throwable { } } + /** + * Merges two indexes through a row filter and checks that the merged index holds exactly the rows + * whose bit was set, packed together in the order the inputs were given. + */ + @Test + public void testFilteredMerge() throws Throwable { + float[][] vector1 = { + {0.0f, 0.0f}, + {1.0f, 1.0f}, + {2.0f, 2.0f} + }; + + float[][] vector2 = { + {10.0f, 10.0f}, + {11.0f, 11.0f}, + {12.0f, 12.0f} + }; + + // Bits 0 to 2 address vector1 and bits 3 to 5 address vector2. Drop the middle row of each, + // which leaves four rows that have to end up at positions 0 to 3 of the merged index. + BitSet rowFilter = new BitSet(); + rowFilter.set(0, 6); + rowFilter.clear(1); + rowFilter.clear(4); + + float[][] survivingRows = { + {0.0f, 0.0f}, + {2.0f, 2.0f}, + {10.0f, 10.0f}, + {12.0f, 12.0f} + }; + // A dropped vector is no longer in the index, so its nearest neighbour is two units away. + float[][] droppedRows = { + {1.0f, 1.0f}, + {11.0f, 11.0f} + }; + + try (CuVSResources resources = CheckedCuVSResources.create()) { + CagraIndexParams indexParams = + new CagraIndexParams.Builder() + .withCagraGraphBuildAlgo(CagraGraphBuildAlgo.NN_DESCENT) + .withGraphDegree(1) + .withIntermediateGraphDegree(2) + .withNumWriterThreads(4) + .withMetric(CuvsDistanceType.L2Expanded) + .build(); + + CagraIndex index1 = + CagraIndex.newBuilder(resources) + .withDataset(vector1) + .withIndexParams(indexParams) + .build(); + CagraIndex index2 = + CagraIndex.newBuilder(resources) + .withDataset(vector2) + .withIndexParams(indexParams) + .build(); + + // Host-built indexes are not mergeable. Dim=2 is not 16-byte aligned, so upload to device, + // allocate owning padded copies, and attach them before merge. + try (var device1 = CuVSMatrix.ofArray(vector1).toDevice(resources); + var device2 = CuVSMatrix.ofArray(vector2).toDevice(resources); + var padded1 = index1.makePaddedDataset(device1); + var padded2 = index2.makePaddedDataset(device2)) { + index1.updateDataset(padded1); + index2.updateDataset(padded2); + + assertEquals("Input index sizes", 3, index1.size()); + assertEquals("Input index sizes", 3, index2.size()); + + try (CagraIndex mergedIndex = + CagraIndex.merge(new CagraIndex[] {index1, index2}, null, rowFilter)) { + assertEquals( + "The merged index should hold one row per set bit", + rowFilter.cardinality(), + mergedIndex.size()); + + // Pin SINGLE_CTA; AUTO may pick MULTI_CTA, which drops neighbors on this tiny dataset. + CagraSearchParams searchParams = + new CagraSearchParams.Builder() + .withAlgo(CagraSearchParams.SearchAlgo.SINGLE_CTA) + .build(); + + try (var queryVectors = CuVSMatrix.ofArray(survivingRows)) { + CagraQuery query = + new CagraQuery.Builder(resources) + .withTopK(1) + .withSearchParams(searchParams) + .withQueryVectors(queryVectors) + .withMapping(SearchResults.IDENTITY_MAPPING) + .build(); + + List> results = mergedIndex.search(query).getResults(); + assertEquals(survivingRows.length, results.size()); + for (int row = 0; row < survivingRows.length; row++) { + Map hit = results.get(row); + assertEquals("Expected a single neighbour for row " + row, 1, hit.size()); + int id = hit.keySet().iterator().next(); + assertEquals("Surviving row " + row + " moved", row, id); + assertEquals( + "Surviving row " + row + " is not an exact match", 0.0f, hit.get(id), 1e-5f); + } + } + + try (var queryVectors = CuVSMatrix.ofArray(droppedRows)) { + CagraQuery query = + new CagraQuery.Builder(resources) + .withTopK(1) + .withSearchParams(searchParams) + .withQueryVectors(queryVectors) + .withMapping(SearchResults.IDENTITY_MAPPING) + .build(); + + List> results = mergedIndex.search(query).getResults(); + assertEquals(droppedRows.length, results.size()); + for (int row = 0; row < droppedRows.length; row++) { + Map hit = results.get(row); + assertEquals("Expected a single neighbour for dropped row " + row, 1, hit.size()); + int id = hit.keySet().iterator().next(); + assertEquals( + "Dropped row " + row + " is still in the merged index", 2.0f, hit.get(id), 1e-5f); + } + } + } + index1.close(); + index2.close(); + } + } + } + + /** + * A filter that selects rows beyond the ones the indexes hold is a caller mistake, not something + * to pass on to cuVS. + */ + @Test + public void testFilteredMergeRejectsOversizedFilter() throws Throwable { + float[][] vectors = { + {0.0f, 0.0f}, + {1.0f, 1.0f} + }; + + try (CuVSResources resources = CheckedCuVSResources.create()) { + CagraIndexParams indexParams = + new CagraIndexParams.Builder() + .withCagraGraphBuildAlgo(CagraGraphBuildAlgo.NN_DESCENT) + .withGraphDegree(1) + .withIntermediateGraphDegree(2) + .withMetric(CuvsDistanceType.L2Expanded) + .build(); + + try (CagraIndex index = + CagraIndex.newBuilder(resources) + .withDataset(vectors) + .withIndexParams(indexParams) + .build()) { + // The index holds two rows, so bit 2 is one row past the end of the merge. + BitSet rowFilter = new BitSet(); + rowFilter.set(0, 3); + + assertThrows( + IllegalArgumentException.class, + () -> CagraIndex.merge(new CagraIndex[] {index}, null, rowFilter)); + } + } + } + + /** + * A merge that fails inside cuVS has to release the index it was building, and leave every input + * index untouched and still usable. Host-backed indexes are not mergeable, which fails the native + * call after the output index has been allocated - the one path where the handle used to be + * dropped on the floor. Repeating it is what would surface a release that goes too far, such as + * one freeing an input index or the same handle twice. + */ + @Test + public void testFailedMergeLeavesTheInputsUsable() throws Throwable { + float[][] vector1 = { + {0.0f, 0.0f}, + {1.0f, 1.0f}, + {2.0f, 2.0f} + }; + + float[][] vector2 = { + {10.0f, 10.0f}, + {11.0f, 11.0f}, + {12.0f, 12.0f} + }; + + try (CuVSResources resources = CheckedCuVSResources.create()) { + CagraIndexParams indexParams = + new CagraIndexParams.Builder() + .withCagraGraphBuildAlgo(CagraGraphBuildAlgo.NN_DESCENT) + .withGraphDegree(1) + .withIntermediateGraphDegree(2) + .withMetric(CuvsDistanceType.L2Expanded) + .build(); + + try (CagraIndex index1 = + CagraIndex.newBuilder(resources) + .withDataset(vector1) + .withIndexParams(indexParams) + .build(); + CagraIndex index2 = + CagraIndex.newBuilder(resources) + .withDataset(vector2) + .withIndexParams(indexParams) + .build()) { + + // No device dataset is attached, so cuVS refuses to merge these. + for (int attempt = 0; attempt < 5; attempt++) { + assertThrows( + Throwable.class, () -> CagraIndex.merge(new CagraIndex[] {index1, index2}, null)); + } + + // The failures left the inputs alone: they still report their rows, and a merge that is + // set up correctly still succeeds afterwards. + assertEquals("index1 survived the failed merges", 3, index1.size()); + assertEquals("index2 survived the failed merges", 3, index2.size()); + + try (var device1 = CuVSMatrix.ofArray(vector1).toDevice(resources); + var device2 = CuVSMatrix.ofArray(vector2).toDevice(resources); + var padded1 = index1.makePaddedDataset(device1); + var padded2 = index2.makePaddedDataset(device2)) { + index1.updateDataset(padded1); + index2.updateDataset(padded2); + + try (CagraIndex mergedIndex = CagraIndex.merge(new CagraIndex[] {index1, index2}, null)) { + assertEquals("The merged index holds every row of both inputs", 6, mergedIndex.size()); + + try (var queryVectors = CuVSMatrix.ofArray(new float[][] {{0.0f, 0.0f}})) { + CagraQuery query = + new CagraQuery.Builder(resources) + .withTopK(1) + .withSearchParams( + new CagraSearchParams.Builder() + .withAlgo(CagraSearchParams.SearchAlgo.SINGLE_CTA) + .build()) + .withQueryVectors(queryVectors) + .withMapping(SearchResults.IDENTITY_MAPPING) + .build(); + + List> results = mergedIndex.search(query).getResults(); + assertEquals(1, results.size()); + assertEquals( + "The first row of the merge is the nearest neighbour of the first vector", + 0, + (int) results.getFirst().keySet().iterator().next()); + } + } + } + } + } + } + // Commented out test for Logical merge strategy as it is not yet implemented in C yet @Test public void testMergeStrategies() throws Throwable { diff --git a/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java b/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java index 23fc738a55..906a310ede 100644 --- a/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java +++ b/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsReader.java @@ -297,6 +297,48 @@ private FieldEntry getFieldEntry(String field, VectorEncoding expectedEncoding) return fieldEntry; } + /** + * Gets the FieldEntry for the given field name, or {@code null} when this segment holds no + * cuVS index for it. + * + * @param field name of the field + * @return the meta information for the field, or {@code null} + */ + FieldEntry getFieldEntry(String field) { + final FieldInfo info = fieldInfos.fieldInfo(field); + return info == null ? null : fields.get(info.number); + } + + /** + * Deserializes the CAGRA index of a single field onto the GPU, bypassing the {@link GPUIndex} + * cache. This is how {@link CuVS2510GPUVectorsWriter} obtains the inputs for the cuVS merge API: + * a reader opened with {@link Context#MERGE} has no cached indexes at all, and a cached index + * belongs to the {@link com.nvidia.cuvs.CuVSResources} of whichever thread opened the reader, + * while the merge API requires every input to share one resources instance - the merging + * thread's. + * + *

The returned index is owned by the caller and must be closed by it. + * + * @param field name of the field + * @return a freshly loaded CAGRA index, or {@code null} if this segment has none for the field + * @throws IOException I/O exception + */ + CagraIndex openCagraIndexForMerge(String field) throws IOException { + FieldEntry fieldEntry = getFieldEntry(field); + if (fieldEntry == null || fieldEntry.cagraIndexLength() == 0) { + return null; + } + try (var slice = + cuvsIndexInput.slice( + "cagra index", fieldEntry.cagraIndexOffset(), fieldEntry.cagraIndexLength()); + var in = new IndexInputInputStream(slice)) { + return CagraIndex.newBuilder(getCuVSResourcesInstance()).from(in).build(); + } catch (Throwable t) { + Utils.handleThrowable(t); + throw new AssertionError("unreachable"); + } + } + /** * Invokes loadCuVSIndex for each field and returns the map of {@link GPUIndex}. * diff --git a/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java b/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java index aead683dc4..46735976dc 100644 --- a/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java +++ b/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java @@ -26,6 +26,7 @@ import java.nio.file.Files; import java.nio.file.Path; import java.util.ArrayList; +import java.util.BitSet; import java.util.List; import java.util.Objects; import org.apache.lucene.codecs.CodecUtil; @@ -34,9 +35,9 @@ import org.apache.lucene.codecs.KnnVectorsWriter; import org.apache.lucene.codecs.hnsw.FlatFieldVectorsWriter; import org.apache.lucene.codecs.hnsw.FlatVectorsWriter; +import org.apache.lucene.codecs.perfield.PerFieldKnnVectorsFormat; import org.apache.lucene.index.DocsWithFieldSet; import org.apache.lucene.index.FieldInfo; -import org.apache.lucene.index.FieldInfos; import org.apache.lucene.index.FloatVectorValues; import org.apache.lucene.index.IndexFileNames; import org.apache.lucene.index.KnnVectorValues; @@ -45,8 +46,8 @@ import org.apache.lucene.index.Sorter; import org.apache.lucene.index.Sorter.DocMap; import org.apache.lucene.index.VectorSimilarityFunction; -import org.apache.lucene.internal.hppc.IntObjectHashMap; import org.apache.lucene.store.IndexOutput; +import org.apache.lucene.util.Bits; import org.apache.lucene.util.IOUtils; import org.apache.lucene.util.InfoStream; @@ -391,92 +392,261 @@ static int distFuncToOrd(VectorSimilarityFunction func) { } /** - * Uses the cuVS API to merge CAGRA indexes. + * The inputs the cuVS merge API needs for one field: the CAGRA indexes to concatenate, and the + * rows of that concatenation that survive the merge. * - * This is currently (and intentionally) marked as unused and will be plugged in later. + * @param readers the readers holding a CAGRA index for the field, in merge order + * @param rowFilter the rows of the concatenated data sets to keep, or {@code null} to keep all of + * them + * @param mergedVectorCount the number of vectors the merged index will hold + */ + private record CagraMergeInputs( + List readers, BitSet rowFilter, int mergedVectorCount) {} + + /** + * Collects the CAGRA indexes that can be handed to the cuVS merge API for this field, together + * with the filter that keeps the merged rows lined up with the merged flat vectors. * * @param fieldInfo instance of the FieldInfo * @param mergeState instance of the MergeState + * @return the inputs for the merge, or {@code null} when the cuVS merge API cannot be used for + * this field * @throws IOException I/O Exceptions */ - @SuppressWarnings("unused") - private void mergeCagraIndexes(FieldInfo fieldInfo, MergeState mergeState) throws IOException { - try { - List cagraIndexes = new ArrayList<>(); - // We need this count so that the merged segment's meta information has the vector count. - int totalVectorCount = 0; - for (int i = 0; i < mergeState.knnVectorsReaders.length; i++) { - KnnVectorsReader knnReader = mergeState.knnVectorsReaders[i]; - // Access the CAGRA index for this field from the reader - if (knnReader != null) { - if (knnReader instanceof CuVS2510GPUVectorsReader cvr) { - if (cvr != null) { - totalVectorCount += cvr.getFieldEntries().get(fieldInfo.number).count(); - CagraIndex cagraIndex = getCagraIndexFromReader(cvr, fieldInfo.name); - if (cagraIndex != null) { - cagraIndexes.add(cagraIndex); - } - } - } else { - // This should never happen - throw new RuntimeException( - "Reader is not of CuVSVectorsReader type. Instead it is: " + knnReader.getClass()); - } - } + private CagraMergeInputs cagraMergeInputs(FieldInfo fieldInfo, MergeState mergeState) + throws IOException { + // A brute force index cannot be produced by the merge API; it would have to be rebuilt from the + // vectors on the host anyway, which is what the vector based merge already does. + if (gpuSearchParams.getIndexType() != IndexType.CAGRA) { + reportNotMergeable(fieldInfo, "the index type is " + gpuSearchParams.getIndexType()); + return null; + } + // A sorted merge interleaves the segments' rows instead of concatenating them. + if (mergeState.needsIndexSort) { + reportNotMergeable(fieldInfo, "the merge has to sort the documents"); + return null; + } + List readers = new ArrayList<>(); + BitSet rowFilter = new BitSet(); + // Rows of the concatenation the merge API sees, and the ones of those that survive. + int rowCount = 0; + int survivingCount = 0; + for (int i = 0; i < mergeState.knnVectorsReaders.length; i++) { + if (KnnVectorsWriter.MergedVectorValues.hasVectorValues( + mergeState.fieldInfos[i], fieldInfo.name) + == false) { + continue; } - assert cagraIndexes.size() > 1; - CagraIndex mergedIndex = - CagraIndex.merge(cagraIndexes.toArray(new CagraIndex[cagraIndexes.size()])); - writeMergedCagraIndex(fieldInfo, mergedIndex, totalVectorCount); - info( - infoStream, - COMPONENT, - "Successfully merged " + cagraIndexes.size() + " CAGRA indexes using native merge API"); - } catch (Throwable t) { - Utils.handleThrowable(t); + KnnVectorsReader knnReader = mergeState.knnVectorsReaders[i]; + if (knnReader instanceof PerFieldKnnVectorsFormat.FieldsReader fieldsReader) { + knnReader = fieldsReader.getFieldReader(fieldInfo.name); + } + if (knnReader == null) { + continue; + } + FloatVectorValues values = knnReader.getFloatVectorValues(fieldInfo.name); + if (values == null || values.size() == 0) { + continue; + } + if (!(knnReader instanceof CuVS2510GPUVectorsReader cuvsReader)) { + reportNotMergeable(fieldInfo, "a segment is read by a " + knnReader.getClass().getName()); + return null; + } + CuVS2510GPUVectorsReader.FieldEntry fieldEntry = cuvsReader.getFieldEntry(fieldInfo.name); + // A segment too small for CAGRA, or one whose CAGRA build failed, holds a brute force index + // instead and has nothing to contribute to the merge. + if (fieldEntry == null || fieldEntry.cagraIndexLength() == 0) { + reportNotMergeable( + fieldInfo, "a segment of " + values.size() + " vectors has no CAGRA index"); + return null; + } + if (fieldEntry.count() != values.size()) { + reportNotMergeable( + fieldInfo, + "a segment holds " + + fieldEntry.count() + + " indexed vectors but " + + values.size() + + " flat vectors"); + return null; + } + BitSet liveRows = + liveRows(values, fieldEntry.count(), mergeState.liveDocs[i], mergeState.docMaps[i]); + int liveCount = liveRows.cardinality(); + // A segment whose vectors are all deleted contributes nothing to the merged flat vectors + // either, so leave it out rather than upload an index only to drop every row of it. + if (liveCount == 0) { + continue; + } + for (int ord = liveRows.nextSetBit(0); ord >= 0; ord = liveRows.nextSetBit(ord + 1)) { + rowFilter.set(rowCount + ord); + } + rowCount += fieldEntry.count(); + survivingCount += liveCount; + readers.add(cuvsReader); + } + boolean filtered = survivingCount != rowCount; + // Merging a single index is only worth it when the filter has rows to drop; without one the + // merge would just copy the index it was given. + if (readers.isEmpty() || (readers.size() == 1 && filtered == false)) { + reportNotMergeable(fieldInfo, "there is nothing to merge, " + readers.size() + " segments"); + return null; + } + // The merged index has to be one CAGRA can build, and cuVS refuses a filter that keeps no rows + // at all. + if (survivingCount < MIN_CAGRA_INDEX_SIZE) { + reportNotMergeable( + fieldInfo, "only " + survivingCount + " vectors survive the deletions of the merge"); + return null; + } + // Without deletions every row is kept, and passing no filter lets cuVS skip the gather. + return new CagraMergeInputs(readers, filtered ? rowFilter : null, survivingCount); + } + + /** + * Returns the ordinals of {@code values} whose document survives the merge. + * + *

A cuVS row id is a vector ordinal rather than a document id, so the deletions have to be + * translated through {@link KnnVectorValues#ordToDoc}. The test applied to each document is the + * one {@code DocIDMerger} applies while producing the merged vectors: a document is dropped when + * the merge maps it to {@code -1}. + */ + private static BitSet liveRows( + FloatVectorValues values, int count, Bits liveDocs, MergeState.DocMap docMap) { + BitSet liveRows = new BitSet(count); + if (liveDocs == null) { + liveRows.set(0, count); + return liveRows; } + for (int ord = 0; ord < count; ord++) { + if (docMap.get(values.ordToDoc(ord)) != -1) { + liveRows.set(ord); + } + } + return liveRows; + } + + /** + * Reports why this field cannot go through the cuVS merge API and has to fall back to the vector + * based merge. + */ + private void reportNotMergeable(FieldInfo fieldInfo, String reason) { + info( + infoStream, + COMPONENT, + "Skipping the cuVS merge API for field \"" + fieldInfo.name + "\": " + reason); } /** - * Extracts the CAGRA index for a specific field from a CuVSVectorsReader. + * Uses the cuVS API to merge the segments' CAGRA indexes on the device, which avoids copying + * every vector back to the host and re-uploading it. + * + * @param fieldInfo instance of the FieldInfo + * @param mergeState instance of the MergeState + * @return true if the field was written, false if the caller has to fall back to the vector + * based merge + * @throws IOException I/O Exceptions */ - private CagraIndex getCagraIndexFromReader(CuVS2510GPUVectorsReader reader, String fieldName) { + private boolean mergeCagraIndexes(FieldInfo fieldInfo, MergeState mergeState) throws IOException { + CagraMergeInputs inputs = cagraMergeInputs(fieldInfo, mergeState); + if (inputs == null) { + return false; + } + long cagraIndexOffset = cuvsIndex.getFilePointer(); + long cagraIndexLength; try { - IntObjectHashMap cuvsIndices = reader.getCuvsIndexes(); - FieldInfos fieldInfos = reader.getFieldInfos(); - FieldInfo fieldInfo = fieldInfos.fieldInfo(fieldName); - if (fieldInfo != null) { - GPUIndex cuvsIndex = cuvsIndices.get(fieldInfo.number); - if (cuvsIndex != null) { - return cuvsIndex.getCagraIndex(); - } + writeMergedIndexBytes(fieldInfo, inputs); + cagraIndexLength = cuvsIndex.getFilePointer() - cagraIndexOffset; + } catch (Throwable t) { + if (t instanceof Error error) { + throw error; } - } catch (Exception e) { + // The merge API needs every input data set on the device at once, so it can run out of + // device memory where the vector based merge would not. Nothing is committed until the meta + // entry below is written, and that is what makes the bytes reachable, so falling back here + // only leaves them unreferenced. info( infoStream, COMPONENT, - "Failed to extract CAGRA index for field " + fieldName + ": " + e.getMessage()); - throw e; + "cuVS merge API failed for field \"" + + fieldInfo.name + + "\", falling back to a vector based merge: " + + t); + return false; } - return null; + // The field is committed by this call. Anything that fails from here is a real I/O failure + // rather than something a rebuild could fix, so it propagates: a fallback now would append a + // second meta entry for the same field and leave the segment unreadable. + writeMeta( + fieldInfo, + inputs.mergedVectorCount(), + cagraIndexOffset, + cagraIndexLength, + cuvsIndex.getFilePointer(), + 0L); + info( + infoStream, + COMPONENT, + "Successfully merged " + + inputs.readers().size() + + " CAGRA indexes for field \"" + + fieldInfo.name + + "\" using the cuVS merge API"); + return true; } /** - * Writes a pre-built merged CAGRA index to the output. + * Merges the field's CAGRA indexes and appends the merged index to the cuVS index output, + * releasing every index it opened before returning. Writes no meta entry, so a failure leaves + * nothing but unreferenced bytes behind. + * + * @param fieldInfo instance of the FieldInfo + * @param inputs the indexes to merge and the rows to keep + * @throws Throwable if the merge, the serialization, or releasing an index fails */ - private void writeMergedCagraIndex(FieldInfo fieldInfo, CagraIndex mergedIndex, int vectorCount) - throws IOException { + private void writeMergedIndexBytes(FieldInfo fieldInfo, CagraMergeInputs inputs) + throws Throwable { + List indexes = new ArrayList<>(inputs.readers().size()); try { - long cagraIndexOffset = cuvsIndex.getFilePointer(); - var cagraIndexOutputStream = new IndexOutputOutputStream(cuvsIndex); - Path tmpFile = - Files.createTempFile(getCuVSResourcesInstance().tempDirectory(), "mergedindex", "cag"); - mergedIndex.serialize(cagraIndexOutputStream, tmpFile); - long cagraIndexLength = cuvsIndex.getFilePointer() - cagraIndexOffset; - writeMeta(fieldInfo, vectorCount, cagraIndexOffset, cagraIndexLength, 0L, 0L); - mergedIndex.close(); - } catch (Throwable t) { - Utils.handleThrowable(t); + for (CuVS2510GPUVectorsReader reader : inputs.readers()) { + indexes.add(reader.openCagraIndexForMerge(fieldInfo.name)); + } + // Derive the output parameters the same way a build over the merged data set would, so the + // merged graph degree matches the one a flush of that size produces. + CagraIndexParams mergeParams = + CagraIndexParamsFactory.create( + gpuSearchParams, inputs.mergedVectorCount(), fieldInfo.getVectorDimension()); + try (CagraIndex mergedIndex = + CagraIndex.merge( + indexes.toArray(new CagraIndex[indexes.size()]), mergeParams, inputs.rowFilter())) { + Path tmpFile = + Files.createTempFile(getCuVSResourcesInstance().tempDirectory(), "mergedindex", "cag"); + try { + mergedIndex.serialize(new IndexOutputOutputStream(cuvsIndex), tmpFile); + } finally { + // cuVS removes the file once it has read it back, but not when serializing failed. + Files.deleteIfExists(tmpFile); + } + } + } finally { + // Every index is released even if one of them fails, and the first failure carries the rest + // so that none of them is lost. IOUtils cannot do this here: CagraIndex.close() throws + // Exception rather than IOException, so it is not a Closeable. + Exception failure = null; + for (CagraIndex index : indexes) { + try { + index.close(); + } catch (Exception e) { + if (failure == null) { + failure = e; + } else { + failure.addSuppressed(e); + } + } + } + if (failure != null) { + throw failure; + } } } @@ -517,7 +687,9 @@ private void vectorBasedMerge(FieldInfo fieldInfo, MergeState mergeState) throws @Override public void mergeOneField(FieldInfo fieldInfo, MergeState mergeState) throws IOException { flatVectorsWriter.mergeOneField(fieldInfo, mergeState); - vectorBasedMerge(fieldInfo, mergeState); + if (mergeCagraIndexes(fieldInfo, mergeState) == false) { + vectorBasedMerge(fieldInfo, mergeState); + } } /** diff --git a/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java b/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java index 1c74649843..d22a1e959d 100644 --- a/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java +++ b/java/cuvs-lucene/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java @@ -25,6 +25,7 @@ import com.nvidia.cuvs.spi.CuVSProvider; import java.lang.invoke.MethodHandle; import java.nio.file.Path; +import java.util.BitSet; import java.util.List; import java.util.logging.Level; @@ -80,9 +81,14 @@ public HnswIndex.Builder newHnswIndexBuilder(CuVSResources cuVSResources) return delegate.newHnswIndexBuilder(cuVSResources); } + /** + * Delegates rather than inheriting the default, which refuses what it cannot honour. The two + * narrower overloads route here, so this is the only one that has to be forwarded. + */ @Override - public CagraIndex mergeCagraIndexes(CagraIndex[] arg0) throws Throwable { - return delegate.mergeCagraIndexes(arg0); + public CagraIndex mergeCagraIndexes(CagraIndex[] arg0, CagraIndexParams arg1, BitSet arg2) + throws Throwable { + return delegate.mergeCagraIndexes(arg0, arg1, arg2); } @Override diff --git a/java/cuvs-lucene/src/test/java/com/nvidia/cuvs/lucene/TestMerge.java b/java/cuvs-lucene/src/test/java/com/nvidia/cuvs/lucene/TestMerge.java index a42f55b3fb..c5651481d6 100644 --- a/java/cuvs-lucene/src/test/java/com/nvidia/cuvs/lucene/TestMerge.java +++ b/java/cuvs-lucene/src/test/java/com/nvidia/cuvs/lucene/TestMerge.java @@ -14,8 +14,13 @@ import java.io.IOException; import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.LinkedHashSet; import java.util.List; +import java.util.Map; import java.util.Random; +import java.util.Set; import java.util.logging.Level; import java.util.logging.Logger; import org.apache.lucene.document.Document; @@ -43,6 +48,7 @@ import org.apache.lucene.tests.util.LuceneTestCase; import org.apache.lucene.tests.util.LuceneTestCase.SuppressSysoutChecks; import org.apache.lucene.util.BytesRef; +import org.apache.lucene.util.InfoStream; import org.junit.After; import org.junit.Before; import org.junit.BeforeClass; @@ -1228,6 +1234,308 @@ public void testLargeScaleMerge() throws IOException { } } + /** + * Tests that a plain CAGRA merge goes through the cuVS merge API, and that the rows of the + * merged index still line up with the merged segment's vector ordinals. + * + *

The merge API concatenates the input data sets, so a mismatch between that concatenation + * and the order Lucene writes the flat vectors in would make every search result point at the + * wrong document. The check catches that: for every hit, the score the reader reports has to + * agree with the distance between the query and the vector of the document that was returned. + **/ + @Test + public void testNativeCagraMergeKeepsOrdinalsAligned() throws IOException { + RecordingInfoStream infoStream = new RecordingInfoStream(); + Map vectorsById = new HashMap<>(); + + indexSegments(infoStream, vectorsById, 4 + random().nextInt(3), 0, null); + + try (DirectoryReader reader = DirectoryReader.open(directory)) { + assertEquals("Should have exactly one segment after merge", 1, reader.leaves().size()); + assertTrue( + "The cuVS merge API should have been used, messages: " + infoStream.messages(), + infoStream.nativeMergeUsed()); + assertSearchHitsMatchTheirVectors(reader, vectorsById); + } + } + + /** + * Tests that a segment the merge API produced can be merged again, which exercises the round + * trip of the merged index through serialization and back onto the device. + **/ + @Test + public void testNativeCagraMergeOfAMergedSegment() throws IOException { + Map vectorsById = new HashMap<>(); + + indexSegments(new RecordingInfoStream(), vectorsById, 2 + random().nextInt(2), 0, null); + // Merges the segment the first round produced together with the new ones. The second round + // records on its own InfoStream, so the assertion below is about that merge alone. + RecordingInfoStream infoStream = new RecordingInfoStream(); + indexSegments(infoStream, vectorsById, 2 + random().nextInt(2), 100_000, null); + + try (DirectoryReader reader = DirectoryReader.open(directory)) { + assertEquals("Should have exactly one segment after merge", 1, reader.leaves().size()); + assertTrue( + "The cuVS merge API should have been used, messages: " + infoStream.messages(), + infoStream.nativeMergeUsed()); + assertSearchHitsMatchTheirVectors(reader, vectorsById); + } + } + + /** + * Tests that segments with deletions still go through the cuVS merge API, with the deleted rows + * left out by the row filter, and that the surviving rows stay lined up with the merged segment's + * vector ordinals. + **/ + @Test + public void testNativeCagraMergeWithDeletions() throws IOException { + RecordingInfoStream infoStream = new RecordingInfoStream(); + Map vectorsById = new HashMap<>(); + + // Delete from every segment, so that no segment of the merge is deletion free. + List deletedIds = + indexSegments( + infoStream, vectorsById, 4 + random().nextInt(3), 0, deleteSomeOfEachSegment()); + assertFalse("Test setup should have deleted documents", deletedIds.isEmpty()); + deletedIds.forEach(vectorsById::remove); + + try (DirectoryReader reader = DirectoryReader.open(directory)) { + assertEquals("Should have exactly one segment after merge", 1, reader.leaves().size()); + assertTrue( + "The cuVS merge API should have been used, messages: " + infoStream.messages(), + infoStream.nativeMergeUsed()); + assertNoDeletedDocuments(reader, deletedIds); + assertSearchHitsMatchTheirVectors(reader, vectorsById); + } + } + + /** + * Tests that a segment whose documents are all deleted is left out of the merge entirely. Its + * rows contribute nothing to the merged vectors, so passing its index to the merge API would only + * upload rows the filter drops again. + **/ + @Test + public void testNativeCagraMergeWithAFullyDeletedSegment() throws IOException { + RecordingInfoStream infoStream = new RecordingInfoStream(); + Map vectorsById = new HashMap<>(); + + List deletedIds = + indexSegments( + infoStream, vectorsById, 3 + random().nextInt(3), 0, deleteWholeFirstSegment()); + assertFalse("Test setup should have deleted documents", deletedIds.isEmpty()); + deletedIds.forEach(vectorsById::remove); + + try (DirectoryReader reader = DirectoryReader.open(directory)) { + assertEquals("Should have exactly one segment after merge", 1, reader.leaves().size()); + assertTrue( + "The cuVS merge API should have been used, messages: " + infoStream.messages(), + infoStream.nativeMergeUsed()); + assertNoDeletedDocuments(reader, deletedIds); + assertSearchHitsMatchTheirVectors(reader, vectorsById); + } + } + + /** + * Tests that expunging the deletions of a single segment goes through the merge API too. There is + * no second index to concatenate, but the filter still has rows to drop, which is cheaper than + * rebuilding the index from the vectors on the host. + **/ + @Test + public void testNativeCagraMergeOfASingleSegmentWithDeletions() throws IOException { + RecordingInfoStream infoStream = new RecordingInfoStream(); + Map vectorsById = new HashMap<>(); + + List deletedIds = + indexSegments(infoStream, vectorsById, 1, 0, deleteSomeOfEachSegment()); + assertFalse("Test setup should have deleted documents", deletedIds.isEmpty()); + deletedIds.forEach(vectorsById::remove); + + try (DirectoryReader reader = DirectoryReader.open(directory)) { + assertEquals("Should have exactly one segment after merge", 1, reader.leaves().size()); + assertTrue( + "The cuVS merge API should have been used, messages: " + infoStream.messages(), + infoStream.nativeMergeUsed()); + assertNoDeletedDocuments(reader, deletedIds); + assertSearchHitsMatchTheirVectors(reader, vectorsById); + } + } + + /** Deletes a few of the documents of every segment. */ + private Deletions deleteSomeOfEachSegment() { + return (segment, docIds) -> { + List deleted = new ArrayList<>(); + for (int i = 0; i < 1 + random().nextInt(3); i++) { + deleted.add(docIds.get(random().nextInt(docIds.size()))); + } + return deleted; + }; + } + + /** Deletes every document of the first segment and nothing else. */ + private Deletions deleteWholeFirstSegment() { + return (segment, docIds) -> segment == 0 ? docIds : List.of(); + } + + /** Asserts that none of the given ids is still searchable. */ + private void assertNoDeletedDocuments(DirectoryReader reader, List deletedIds) + throws IOException { + IndexSearcher searcher = new IndexSearcher(reader); + for (int deletedId : deletedIds) { + assertEquals( + "Deleted document " + deletedId + " should be gone", + 0, + searcher.count(new TermQuery(new Term("id", String.valueOf(deletedId))))); + } + } + + /** Chooses which of the documents a segment holds to delete before the merge. */ + @FunctionalInterface + private interface Deletions { + /** + * @param segment the index of the segment, counting from the first one this round indexed + * @param docIds the ids of the documents the segment holds + * @return the ids to delete, which may repeat and may be empty + */ + List forSegment(int segment, List docIds); + } + + /** + * Indexes one segment per commit with a CAGRA only configuration, numbering the documents from + * {@code firstDocId}, then force merges the whole directory into a single segment. Records every + * indexed vector by document id, and, when {@code deletions} is not null, deletes the documents + * it selects from each segment before the merge. + * + * @return the ids of the documents that were deleted, without repetitions + **/ + private List indexSegments( + RecordingInfoStream infoStream, + Map vectorsById, + int segmentCount, + int firstDocId, + Deletions deletions) + throws IOException { + // Enough documents per segment for cuVS to build a CAGRA index rather than fall back to brute + // force, which would disqualify the segment from the merge API. + int docsPerSegment = 64 + random().nextInt(64); + + GPUSearchParams params = + new GPUSearchParams.Builder() + .withCagraGraphBuildAlgo(cagraGraphBuildAlgo) + .withIndexType(IndexType.CAGRA) + .build(); + + IndexWriterConfig config = + new IndexWriterConfig() + .setCodec(alwaysKnnVectorsFormat(new CuVS2510GPUVectorsFormat(params))) + .setInfoStream(infoStream) + // Flush on commit only, so that each commit produces exactly one segment. + .setMaxBufferedDocs(segmentCount * docsPerSegment + 1) + .setRAMBufferSizeMB(IndexWriterConfig.DISABLE_AUTO_FLUSH); + + List> docIdsPerSegment = new ArrayList<>(); + Set deletedIds = new LinkedHashSet<>(); + + try (IndexWriter writer = new IndexWriter(directory, config)) { + for (int segment = 0; segment < segmentCount; segment++) { + List docIds = new ArrayList<>(); + for (int i = 0; i < docsPerSegment; i++) { + int docId = firstDocId + segment * docsPerSegment + i; + float[] vector = generateRandomVector(vectorDimension, random()); + Document doc = new Document(); + doc.add(new StringField("id", String.valueOf(docId), Field.Store.YES)); + doc.add(new KnnFloatVectorField("vector", vector, VectorSimilarityFunction.EUCLIDEAN)); + writer.addDocument(doc); + vectorsById.put(docId, vector); + docIds.add(docId); + } + docIdsPerSegment.add(docIds); + // One commit per segment; the default merge policy leaves this few segments alone until + // the force merge below, so the merge sees every segment at once. + writer.commit(); + } + + if (deletions != null) { + for (int segment = 0; segment < segmentCount; segment++) { + for (int deletedId : deletions.forSegment(segment, docIdsPerSegment.get(segment))) { + writer.deleteDocuments(new Term("id", String.valueOf(deletedId))); + deletedIds.add(deletedId); + } + } + writer.commit(); + } + + writer.forceMerge(1); + } + return List.copyOf(deletedIds); + } + + /** + * Runs a few searches and verifies that the score of every hit matches the distance between the + * query and the vector that was indexed for the document that was returned. + **/ + private void assertSearchHitsMatchTheirVectors( + DirectoryReader reader, Map vectorsById) throws IOException { + // A merged index holding the wrong number of rows would still answer searches, so check the + // count the merge recorded rather than infer it from the hits. + assertEquals( + "Merged segment holds the wrong number of vectors", + vectorsById.size(), + reader.leaves().get(0).reader().getFloatVectorValues("vector").size()); + IndexSearcher searcher = new IndexSearcher(reader); + for (int query = 0; query < 5; query++) { + float[] queryVector = generateRandomVector(vectorDimension, random()); + int k = 1 + random().nextInt(16); + TopDocs topDocs = searcher.search(new KnnFloatVectorQuery("vector", queryVector, k), k); + assertTrue("Should find results after merge", topDocs.scoreDocs.length > 0); + for (ScoreDoc scoreDoc : topDocs.scoreDocs) { + int docId = Integer.parseInt(searcher.storedFields().document(scoreDoc.doc).get("id")); + float[] vector = vectorsById.get(docId); + assertNotNull("Hit on an unknown or deleted document: " + docId, vector); + // The index is built with the default L2Expanded metric, and the reader turns the squared + // euclidean distance cuVS returns into a score of 1 / (1 + distance). + float squaredDistance = 0.0f; + for (int i = 0; i < vector.length; i++) { + float difference = queryVector[i] - vector[i]; + squaredDistance += difference * difference; + } + assertEquals( + "Score of document " + docId + " does not match its vector", + 1.0f / (1.0f + squaredDistance), + scoreDoc.score, + 1e-3f); + } + } + } + + /** An InfoStream that keeps the messages, so that a test can tell which merge path ran. */ + private static class RecordingInfoStream extends InfoStream { + + private final List messages = Collections.synchronizedList(new ArrayList<>()); + + @Override + public void message(String component, String message) { + messages.add(component + ": " + message); + } + + @Override + public boolean isEnabled(String component) { + return true; + } + + @Override + public void close() {} + + List messages() { + synchronized (messages) { + return List.copyOf(messages); + } + } + + boolean nativeMergeUsed() { + return messages().stream().anyMatch(message -> message.contains("using the cuVS merge API")); + } + } + /** Helper method to generate random vectors */ private float[] generateRandomVector(int dimension, Random random) { float[] vector = new float[dimension];