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