Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 32 additions & 2 deletions java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;

/**
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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);
}

/**
Expand All @@ -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}.
*
* <p>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 <b>set</b> 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");
}
Expand All @@ -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);
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;

/**
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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);
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -604,6 +604,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,
Expand Down Expand Up @@ -923,69 +935,150 @@ 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;
}
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");
}
} 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.
*
* <p> 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;
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading