diff --git a/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java b/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java index 2a8ab02f..416da8f2 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java +++ b/src/main/java/com/nvidia/cuvs/lucene/AcceleratedHNSWUtils.java @@ -11,6 +11,7 @@ import com.nvidia.cuvs.CagraIndex; import com.nvidia.cuvs.CagraIndexParams; import com.nvidia.cuvs.CuVSMatrix; +import com.nvidia.cuvs.CuVSResources; import com.nvidia.cuvs.RowView; import java.io.IOException; import java.util.ArrayList; @@ -87,7 +88,8 @@ public static GPUBuiltHnswGraph createMultiLayerHnswGraph( int hnswLayers, int graphDegree, CagraIndexParams params, - QuantizationType quantization) + QuantizationType quantization, + CuVSResources cuVSResources) throws Throwable { // Calculate M as cagraGraphDegree/2 @@ -141,7 +143,13 @@ public static GPUBuiltHnswGraph createMultiLayerHnswGraph( // Build CAGRA graph for this layer layerAdjacencies.add( buildCagraGraphForSubset( - selectedVectors, selectedNodes, 0, params, dimensions, quantization)); + selectedVectors, + selectedNodes, + 0, + params, + dimensions, + quantization, + cuVSResources)); } else { @@ -155,7 +163,13 @@ public static GPUBuiltHnswGraph createMultiLayerHnswGraph( // Build CAGRA graph for this layer layerAdjacencies.add( buildCagraGraphForSubset( - selectedVectors, selectedNodes, bytesPerVector, params, dimensions, quantization)); + selectedVectors, + selectedNodes, + bytesPerVector, + params, + dimensions, + quantization, + cuVSResources)); } // Update for next iteration @@ -179,7 +193,8 @@ private static CuVSMatrix buildCagraGraphForSubset( int bytesPerVector, CagraIndexParams params, int dimensions, - QuantizationType quantization) + QuantizationType quantization, + CuVSResources cuVSResources) throws Throwable { CuVSMatrix subsetDataset; @@ -196,7 +211,7 @@ private static CuVSMatrix buildCagraGraphForSubset( // Build CAGRA index for the subset CagraIndex subsetIndex = - CagraIndex.newBuilder(getCuVSResourcesInstance()) + CagraIndex.newBuilder(cuVSResources) .withDataset(subsetDataset) .withIndexParams(params) .build(); diff --git a/src/main/java/com/nvidia/cuvs/lucene/CagraIndexParamsFactory.java b/src/main/java/com/nvidia/cuvs/lucene/CagraIndexParamsFactory.java index 20fed04b..8995460b 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/CagraIndexParamsFactory.java +++ b/src/main/java/com/nvidia/cuvs/lucene/CagraIndexParamsFactory.java @@ -19,7 +19,7 @@ */ public class CagraIndexParamsFactory { - private static final int ALGO_SWITCH_THRESHOLD = 5_000_000; + private static final int ALGO_SWITCH_THRESHOLD = 1_000_000; /** * Translation of the internal logic found here: diff --git a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java index e8fb304e..a4498d95 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java +++ b/src/main/java/com/nvidia/cuvs/lucene/CuVS2510GPUVectorsWriter.java @@ -11,6 +11,7 @@ import static com.nvidia.cuvs.lucene.CuVS2510GPUVectorsFormat.VERSION_CURRENT; import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.closeCuVSResourcesInstance; import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.getCuVSResourcesInstance; +import static com.nvidia.cuvs.lucene.Utils.Target.HOST; import static com.nvidia.cuvs.lucene.Utils.info; import static org.apache.lucene.index.VectorEncoding.FLOAT32; import static org.apache.lucene.search.DocIdSetIterator.NO_MORE_DOCS; @@ -202,7 +203,7 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) thro try { CuVSMatrix cagraDataset = Utils.createFloatMatrix( - vectors, fieldInfo.getVectorDimension(), getCuVSResourcesInstance()); + vectors, fieldInfo.getVectorDimension(), getCuVSResourcesInstance(), HOST); writeCagraIndex(cagraIndexOutputStream, cagraDataset); } catch (Throwable t) { // Fallback to brute force in a few cases, for now. @@ -216,7 +217,7 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) thro var bruteForceIndexOutputStream = new IndexOutputOutputStream(cuvsIndex); CuVSMatrix bruteforceDataset = Utils.createFloatMatrix( - vectors, fieldInfo.getVectorDimension(), getCuVSResourcesInstance()); + vectors, fieldInfo.getVectorDimension(), getCuVSResourcesInstance(), HOST); writeBruteForceIndex(bruteForceIndexOutputStream, bruteforceDataset); bruteForceIndexLength = cuvsIndex.getFilePointer() - bruteForceIndexOffset; diff --git a/src/main/java/com/nvidia/cuvs/lucene/CuVSResourcesManager.java b/src/main/java/com/nvidia/cuvs/lucene/CuVSResourcesManager.java new file mode 100644 index 00000000..38b8ef06 --- /dev/null +++ b/src/main/java/com/nvidia/cuvs/lucene/CuVSResourcesManager.java @@ -0,0 +1,294 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +package com.nvidia.cuvs.lucene; + +import static com.nvidia.cuvs.CagraIndexParams.CagraGraphBuildAlgo.IVF_PQ; +import static com.nvidia.cuvs.CagraIndexParams.CagraGraphBuildAlgo.NN_DESCENT; + +import com.nvidia.cuvs.CagraIndexParams; +import com.nvidia.cuvs.CagraIndexParams.CagraGraphBuildAlgo; +import com.nvidia.cuvs.CuVSIvfPqIndexParams; +import com.nvidia.cuvs.CuVSIvfPqParams; +import com.nvidia.cuvs.CuVSIvfPqSearchParams; +import com.nvidia.cuvs.CuVSResources; +import com.nvidia.cuvs.GPUInfoProvider; +import com.nvidia.cuvs.spi.CuVSProvider; +import java.util.Arrays; +import java.util.Objects; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.locks.Condition; +import java.util.concurrent.locks.ReentrantLock; +import java.util.logging.Level; +import java.util.logging.Logger; +import java.util.stream.LongStream; + +/** + * Manages a pool of finite {@link ManagedCuVSResources} and allows for the accessing threads + * to lock and acquire available instance and release them back to the pool when finished. + */ +public class CuVSResourcesManager { + + private static final Logger LOG = Logger.getLogger(Utils.class.getName()); + private static final CuVSProvider PROVIDER = CuVSProvider.provider(); + private static final GPUInfoProvider GPU_INFO_PROVIDER = PROVIDER.gpuInfoProvider(); + private static final int MAX_POOL_SIZE = 512; + private static final double ONE_G = Math.pow(1024, 3); + + private ManagedCuVSResources[] pool; + private ReentrantLock lock; + private Condition resourcesAvailable; + private AtomicLong reserveMemory; + private long totalDeviceMemory; + private int capacity; + private CuVSResources cuVSResources; + + public CuVSResourcesManager(int capacity) { + if (capacity > MAX_POOL_SIZE || capacity <= 0) { + throw new IllegalArgumentException( + "Invalid capacity, should be between 1 and " + MAX_POOL_SIZE); + } + this.capacity = capacity; + pool = new ManagedCuVSResources[capacity]; + lock = new ReentrantLock(); + resourcesAvailable = lock.newCondition(); + for (int i = 0; i < capacity; i++) { + pool[i] = new ManagedCuVSResources(getCuVSResourceInstance()); + } + cuVSResources = getCuVSResourceInstance(); + reserveMemory = new AtomicLong(); + totalDeviceMemory = GPU_INFO_PROVIDER.getCurrentInfo(cuVSResources).totalDeviceMemoryInBytes(); + } + + /** + * Acquire an instance of {@link ManagedCuVSResources} when available and enough + * device memory is also available for the request to complete. + * + * @param rows the number of vectors in the dataset + * @param dimension the vector dimension in the dataset + * @param params an instance of {@link CagraIndexParams} + * @return an instance of {@link ManagedCuVSResources} + * @throws InterruptedException + */ + public ManagedCuVSResources acquireResource(long rows, long dimension, CagraIndexParams params) + throws InterruptedException { + try { + lock.lock(); + long neededMemory = getEstimatedMemoryRequirement(rows, dimension, params); + + if (neededMemory > totalDeviceMemory) { + throw new RuntimeException("Not enough GPU device memory available"); + } + + long currentFreeMemory = + GPU_INFO_PROVIDER.getCurrentInfo(cuVSResources).freeDeviceMemoryInBytes(); + + while (getNumberOfUnavailableResources() == capacity + || (totalDeviceMemory - reserveMemory.get()) < neededMemory + || currentFreeMemory < neededMemory) { + resourcesAvailable.await(); + } + reserveMemory.addAndGet(neededMemory); + + ManagedCuVSResources managedCuVSResources = getAvailableResourcesFromPool(); + assert managedCuVSResources != null; + + managedCuVSResources.setNeededMemory(neededMemory); + managedCuVSResources.lock(); + return managedCuVSResources; + } finally { + lock.unlock(); + } + } + + /** + * Releases the acquired instance of {@link ManagedCuVSResources} back to the pool. + * + * @param resource the acquired instance of {@link ManagedCuVSResources} by the thread + */ + public void releaseResource(ManagedCuVSResources resource) { + try { + lock.lock(); + reserveMemory.addAndGet(-resource.getNeededMemory()); + resource.resetNeededMemory(); + resource.unlock(); + resourcesAvailable.signalAll(); + } finally { + lock.unlock(); + } + } + + /** + * Shuts down the instances of wrapped {@link CuVSResources} in the pool. + */ + public void shutdown() { + Arrays.stream(pool) + .forEach( + managedCuVSResources -> { + if (Objects.nonNull(managedCuVSResources) + && Objects.nonNull(managedCuVSResources.getResource())) { + managedCuVSResources.getResource().close(); + } + }); + if (cuVSResources != null) { + cuVSResources.close(); + } + } + + private static CuVSResources getCuVSResourceInstance() { + try { + return CuVSResources.create(); + } catch (UnsupportedOperationException uoe) { + LOG.log( + Level.WARNING, + "cuVS is not supported on this platform or java version: " + uoe.getMessage()); + } catch (Throwable t) { + if (t instanceof ExceptionInInitializerError ex) { + t = ex.getCause(); + } + LOG.log(Level.WARNING, "Exception occurred during creation of cuVS resources. " + t); + } + return null; + } + + private ManagedCuVSResources getAvailableResourcesFromPool() { + return Arrays.stream(pool) + .filter(managedCuVSResources -> !managedCuVSResources.isLocked()) + .findFirst() + .orElse(null); + } + + private long getNumberOfUnavailableResources() { + return Arrays.stream(pool) + .filter(managedCuVSResources -> managedCuVSResources.isLocked()) + .count(); + } + + private long getEstimatedMemoryRequirement(long rows, long dimension, CagraIndexParams params) { + CagraGraphBuildAlgo buildAlgo = params.getCagraGraphBuildAlgo(); + if (buildAlgo.equals(NN_DESCENT)) { + return estimateNNDescentIndexBuildPeakMemory(rows, dimension, params); + } else if (buildAlgo.equals(IVF_PQ)) { + assert params.getCuVSIvfPqParams() != null; + return estimateIVFPQIndexBuildPeakMemory(rows, dimension, params); + } else { + throw new IllegalArgumentException("Unsupported CAGRA build algo"); + } + } + + private long estimateNNDescentIndexBuildPeakMemory( + long rows, long dimension, CagraIndexParams params) { + /* + * N = Number of vectors + * D = Vector dimension + * I = Intermediate Graph Degree + * Sidx = Bytes per graph neighbor ID. This is sizeof(IdxT), usually 4 for int32_t or uint32_t. + * + * NND_device_peak = N × (D × 2 + 276) + * optimize_peak = N × (4 + (Sidx + 1) × I) + * build_peak = dataset_size + max(NND_device_peak, optimize_peak) + */ + + final int sIdx = 4; + long nnDevicePeak = rows * (dimension * 2 + 276); + long optimizePeak = rows * (4 + (sIdx + 1) * params.getIntermediateGraphDegree()); + return (long) ((Math.max(nnDevicePeak, optimizePeak)) + ONE_G); + } + + private long estimateIVFPQIndexBuildPeakMemory( + long rows, long dimension, CagraIndexParams params) { + + /* + * R = IVF-PQ training-set ratio. This is train_set_ratio. + * N = Number of vectors + * D = Vector dimension + * C = Number of IVF-PQ coarse clusters/lists. This is the IVF-PQ n_lists + * Q = Query batch size + * I = Intermediate graph degree + * Sidx = Bytes per graph neighbor ID. This is sizeof(IdxT), usually 4 for int32_t or uint32_t. + * + * IVFPQ_build_peak ​= (R/N × D × 4) + (C × D × 4) + (R/N * sizeof(uint32_t)) + * IVFPQ_search_peak = (Q × D × 4)+ (Q × I × sizeof(uint32_t))+ (Q × I × 4) + * optimize_peak = N × (4 + (Sidx + 1) × I) + * build_peak = dataset_size + max(IVFPQ_build_peak, IVFPQ_search_peak, optimize_peak) + * + * https://github.com/rapidsai/cuvs/blob/main/cpp/src/neighbors/ivf_pq/ivf_pq_build.cuh#L1253 + * trainset_ratio = max(1, n_rows / max(kmeans_trainset_fraction * n_rows, n_lists)) + */ + + CuVSIvfPqParams p = params.getCuVSIvfPqParams(); + CuVSIvfPqIndexParams ip = p.getIndexParams(); + CuVSIvfPqSearchParams sp = p.getSearchParams(); + assert ip != null; + assert sp != null; + + final int sIdx = 4; + + double trainsetRatio = + Math.max(1, rows / Math.max((ip.getKmeansTrainsetFraction() * rows), ip.getnLists())); + + double ivfPQBuildPeak = + (trainsetRatio / rows * dimension * 4) + + (ip.getnLists() * dimension * 4) + + (trainsetRatio / rows * 4); + + long queryBatchSize = 4096; // Not sure about this yet "max_internal_batch_size" + + long ivfPQSearchPeak = + (queryBatchSize * dimension * 4) + + (2 * queryBatchSize * params.getIntermediateGraphDegree() * 4); + + long optimizePeak = rows * (4 + (sIdx + 1) * params.getIntermediateGraphDegree()); + + return (long) + (LongStream.of((long) Math.ceil(ivfPQBuildPeak), ivfPQSearchPeak, optimizePeak) + .max() + .getAsLong() + + (2.5 * ONE_G)); + } + + /** + * Holds reference to CuVSResources with its associated lock, and needed memory. + */ + class ManagedCuVSResources { + + private final CuVSResources cuVSResources; + private final ReentrantLock lock; + private long neededMemory; + + public ManagedCuVSResources(CuVSResources cuVSResources) { + this.cuVSResources = cuVSResources; + lock = new ReentrantLock(); + } + + public CuVSResources getResource() { + return cuVSResources; + } + + public long getNeededMemory() { + return neededMemory; + } + + public void resetNeededMemory() { + setNeededMemory(0); + } + + public void setNeededMemory(long neededMemory) { + this.neededMemory = neededMemory; + } + + public void lock() { + lock.lock(); + } + + public void unlock() { + lock.unlock(); + } + + public boolean isLocked() { + return lock.isLocked(); + } + } +} diff --git a/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java b/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java deleted file mode 100644 index 07e64fb1..00000000 --- a/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSProvider.java +++ /dev/null @@ -1,168 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION. - * SPDX-License-Identifier: Apache-2.0 - */ -package com.nvidia.cuvs.lucene; - -import com.nvidia.cuvs.BruteForceIndex; -import com.nvidia.cuvs.CagraIndex; -import com.nvidia.cuvs.CagraIndexParams; -import com.nvidia.cuvs.CagraIndexParams.CuvsDistanceType; -import com.nvidia.cuvs.CagraIndexParams.HnswHeuristicType; -import com.nvidia.cuvs.CuVSDeviceMatrix; -import com.nvidia.cuvs.CuVSHostMatrix; -import com.nvidia.cuvs.CuVSMatrix; -import com.nvidia.cuvs.CuVSMatrix.Builder; -import com.nvidia.cuvs.CuVSMatrix.DataType; -import com.nvidia.cuvs.CuVSResources; -import com.nvidia.cuvs.GPUInfoProvider; -import com.nvidia.cuvs.HnswIndex; -import com.nvidia.cuvs.HnswIndexParams; -import com.nvidia.cuvs.TieredIndex; -import com.nvidia.cuvs.spi.CuVSProvider; -import java.lang.invoke.MethodHandle; -import java.nio.file.Path; -import java.util.logging.Level; - -class FilterCuVSProvider implements CuVSProvider { - - private final CuVSProvider delegate; - - FilterCuVSProvider(CuVSProvider delegate) { - this.delegate = delegate; - } - - @Override - public Path nativeLibraryPath() { - return CuVSProvider.TMPDIR; - } - - @Override - public CuVSResources newCuVSResources(Path tempPath) throws Throwable { - return delegate.newCuVSResources(tempPath); - } - - @Override - public BruteForceIndex.Builder newBruteForceIndexBuilder(CuVSResources cuVSResources) - throws UnsupportedOperationException { - return delegate.newBruteForceIndexBuilder(cuVSResources); - } - - @Override - public CagraIndex.Builder newCagraIndexBuilder(CuVSResources cuVSResources) - throws UnsupportedOperationException { - return delegate.newCagraIndexBuilder(cuVSResources); - } - - @Override - public HnswIndex.Builder newHnswIndexBuilder(CuVSResources cuVSResources) - throws UnsupportedOperationException { - return delegate.newHnswIndexBuilder(cuVSResources); - } - - @Override - public CagraIndex mergeCagraIndexes(CagraIndex[] arg0) throws Throwable { - return delegate.mergeCagraIndexes(arg0); - } - - @Override - public GPUInfoProvider gpuInfoProvider() { - return delegate.gpuInfoProvider(); - } - - @Override - public Builder newHostMatrixBuilder(long rows, long cols, DataType dataType) { - return delegate.newHostMatrixBuilder(rows, cols, dataType); - } - - @Override - public Builder newHostMatrixBuilder( - long rows, long cols, int maxRows, int maxCols, DataType dataType) { - return delegate.newHostMatrixBuilder(rows, cols, maxRows, maxCols, dataType); - } - - @Override - public Builder newDeviceMatrixBuilder( - CuVSResources resources, long rows, long cols, DataType dataType) { - return delegate.newDeviceMatrixBuilder(resources, rows, cols, dataType); - } - - @Override - public Builder newDeviceMatrixBuilder( - CuVSResources resources, long rows, long cols, int maxRows, int maxCols, DataType dataType) { - return delegate.newDeviceMatrixBuilder(resources, rows, cols, maxRows, maxCols, dataType); - } - - @Override - public MethodHandle newNativeMatrixBuilder() { - return delegate.newNativeMatrixBuilder(); - } - - @Override - public MethodHandle newNativeMatrixBuilderWithStrides() { - return delegate.newNativeMatrixBuilderWithStrides(); - } - - @Override - public CuVSMatrix newMatrixFromArray(float[][] vectors) { - return delegate.newMatrixFromArray(vectors); - } - - @Override - public CuVSMatrix newMatrixFromArray(int[][] vectors) { - return delegate.newMatrixFromArray(vectors); - } - - @Override - public CuVSMatrix newMatrixFromArray(byte[][] vectors) { - return delegate.newMatrixFromArray(vectors); - } - - @Override - public TieredIndex.Builder newTieredIndexBuilder(CuVSResources cuVSResources) - throws UnsupportedOperationException { - return delegate.newTieredIndexBuilder(cuVSResources); - } - - @Override - public CagraIndexParams cagraIndexParamsFromHnswParams( - long arg0, long arg1, int arg2, int arg3, HnswHeuristicType arg4, CuvsDistanceType arg5) { - return delegate.cagraIndexParamsFromHnswParams(arg0, arg1, arg2, arg3, arg4, arg5); - } - - @Override - public Level getLogLevel() { - return delegate.getLogLevel(); - } - - @Override - public void setLogLevel(Level arg0) { - delegate.setLogLevel(arg0); - } - - @Override - public HnswIndex hnswIndexFromCagra(HnswIndexParams arg0, CagraIndex arg1) throws Throwable { - return delegate.hnswIndexFromCagra(arg0, arg1); - } - - @Override - public void enableRMMManagedPooledMemory(int arg0, int arg1) { - delegate.enableRMMManagedPooledMemory(arg0, arg1); - } - - @Override - public void enableRMMPooledMemory(int arg0, int arg1) { - delegate.enableRMMPooledMemory(arg0, arg1); - } - - @Override - public void resetRMMPooledMemory() { - delegate.resetRMMPooledMemory(); - } - - @Override - public HnswIndex hnswIndexBuild(CuVSResources arg0, HnswIndexParams arg1, CuVSMatrix arg2) - throws Throwable { - return delegate.hnswIndexBuild(arg0, arg1, arg2); - } -} diff --git a/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSServiceProvider.java b/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSServiceProvider.java deleted file mode 100644 index 83cebafa..00000000 --- a/src/main/java/com/nvidia/cuvs/lucene/FilterCuVSServiceProvider.java +++ /dev/null @@ -1,24 +0,0 @@ -/* - * SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION. - * SPDX-License-Identifier: Apache-2.0 - */ -package com.nvidia.cuvs.lucene; - -import com.nvidia.cuvs.spi.CuVSProvider; -import com.nvidia.cuvs.spi.CuVSServiceProvider; - -/** - * A provider that creates instances of FilterCuVSProvider. - * - * @since 25.10 - */ -public class FilterCuVSServiceProvider extends CuVSServiceProvider { - - /** - * Initialize and return an CuVSProvider provided by this provider. - */ - @Override - public CuVSProvider get(CuVSProvider builtinProvider) { - return new FilterCuVSProvider(builtinProvider); - } -} diff --git a/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsFormat.java b/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsFormat.java index 3c42707c..ff7dd0c3 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsFormat.java +++ b/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsFormat.java @@ -31,6 +31,7 @@ public class Lucene99AcceleratedHNSWVectorsFormat extends KnnVectorsFormat { private static final FlatVectorsFormat FLAT_VECTORS_FORMAT; private static final int MAX_DIMENSIONS = 4096; private final AcceleratedHNSWParams acceleratedHNSWParams; + private final CuVSResourcesManager cuvsResourcesManager; static final String HNSW_META_CODEC_NAME = "Lucene99HnswVectorsFormatMeta"; static final String HNSW_META_CODEC_EXT = "vem"; @@ -67,6 +68,11 @@ public Lucene99AcceleratedHNSWVectorsFormat() { public Lucene99AcceleratedHNSWVectorsFormat(AcceleratedHNSWParams acceleratedHNSWParams) { super("Lucene99AcceleratedHNSWVectorsFormat"); this.acceleratedHNSWParams = acceleratedHNSWParams; + if (isSupported()) { + cuvsResourcesManager = new CuVSResourcesManager(acceleratedHNSWParams.getWriterThreads()); + } else { + cuvsResourcesManager = null; // Will not be needed in fallback mode. + } } /** @@ -77,7 +83,8 @@ public KnnVectorsWriter fieldsWriter(SegmentWriteState state) throws IOException var flatWriter = FLAT_VECTORS_FORMAT.fieldsWriter(state); if (isSupported()) { log.log(Level.FINE, "cuVS is supported so using the Lucene99AcceleratedHNSWVectorsWriter"); - return new Lucene99AcceleratedHNSWVectorsWriter(state, acceleratedHNSWParams, flatWriter); + return new Lucene99AcceleratedHNSWVectorsWriter( + state, acceleratedHNSWParams, flatWriter, cuvsResourcesManager); } else { log.log( Level.WARNING, diff --git a/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsWriter.java b/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsWriter.java index 35209ce5..0eb68ab9 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsWriter.java +++ b/src/main/java/com/nvidia/cuvs/lucene/Lucene99AcceleratedHNSWVectorsWriter.java @@ -14,8 +14,7 @@ import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_INDEX_EXT; import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_META_CODEC_EXT; import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_META_CODEC_NAME; -import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.closeCuVSResourcesInstance; -import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.getCuVSResourcesInstance; +import static com.nvidia.cuvs.lucene.Utils.Target.HOST; import static com.nvidia.cuvs.lucene.Utils.createListFromMergedVectors; import static org.apache.lucene.index.VectorEncoding.FLOAT32; import static org.apache.lucene.util.RamUsageEstimator.shallowSizeOfInstance; @@ -24,6 +23,7 @@ import com.nvidia.cuvs.CagraIndexParams; import com.nvidia.cuvs.CuVSMatrix; import com.nvidia.cuvs.lucene.AcceleratedHNSWUtils.QuantizationType; +import com.nvidia.cuvs.lucene.CuVSResourcesManager.ManagedCuVSResources; import java.io.IOException; import java.util.ArrayList; import java.util.List; @@ -61,6 +61,7 @@ public class Lucene99AcceleratedHNSWVectorsWriter extends KnnVectorsWriter { private final FlatVectorsWriter flatVectorsWriter; private final List fields = new ArrayList<>(); private final InfoStream infoStream; + private final CuVSResourcesManager cuvsResourcesManager; private IndexOutput hnswMeta = null; private IndexOutput hnswVectorIndex = null; private String vemFileName; @@ -87,12 +88,14 @@ public class Lucene99AcceleratedHNSWVectorsWriter extends KnnVectorsWriter { public Lucene99AcceleratedHNSWVectorsWriter( SegmentWriteState state, AcceleratedHNSWParams acceleratedHNSWParams, - FlatVectorsWriter flatVectorsWriter) + FlatVectorsWriter flatVectorsWriter, + CuVSResourcesManager cuvsResourcesManager) throws IOException { super(); this.flatVectorsWriter = flatVectorsWriter; this.infoStream = state.infoStream; this.acceleratedHNSWParams = acceleratedHNSWParams; + this.cuvsResourcesManager = cuvsResourcesManager; vemFileName = IndexFileNames.segmentFileName( state.segmentInfo.name, state.segmentSuffix, HNSW_META_CODEC_EXT); @@ -144,8 +147,10 @@ public KnnFieldVectorsWriter addField(FieldInfo fieldInfo) throws IOException * @param fieldInfo instance of FieldInfo that has the field description * @param vectors vectors to index * @throws IOException + * @throws InterruptedException */ - private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOException { + private void writeFieldInternal(FieldInfo fieldInfo, List vectors) + throws IOException, InterruptedException { if (vectors.size() == 0) { writeEmpty(fieldInfo, hnswMeta); return; @@ -154,16 +159,17 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) thro writeSingleVectorGraph(fieldInfo, vectors); return; } + CagraIndexParams params = + CagraIndexParamsFactory.create( + acceleratedHNSWParams, vectors.size(), vectors.get(0).length); + ManagedCuVSResources managedCuVSResources = + cuvsResourcesManager.acquireResource(vectors.size(), vectors.get(0).length, params); try { CuVSMatrix dataset = Utils.createFloatMatrix( - vectors, fieldInfo.getVectorDimension(), getCuVSResourcesInstance()); - - CagraIndexParams params = - CagraIndexParamsFactory.create(acceleratedHNSWParams, dataset.size(), dataset.columns()); - + vectors, fieldInfo.getVectorDimension(), managedCuVSResources.getResource(), HOST); CagraIndex cagraIndex = - CagraIndex.newBuilder(getCuVSResourcesInstance()) + CagraIndex.newBuilder(managedCuVSResources.getResource()) .withDataset(dataset) .withIndexParams(params) .build(); @@ -180,7 +186,8 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) thro acceleratedHNSWParams.getHnswLayers(), acceleratedHNSWParams.getGraphdegree(), params, - QuantizationType.NONE); + QuantizationType.NONE, + managedCuVSResources.getResource()); long vectorIndexOffset = hnswVectorIndex.getFilePointer(); int[][] graphLevelNodeOffsets = writeGraph(hnswGraph, hnswVectorIndex); long vectorIndexLength = hnswVectorIndex.getFilePointer() - vectorIndexOffset; @@ -197,6 +204,8 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) thro cagraIndex.close(); } catch (Throwable t) { Utils.handleThrowable(t); + } finally { + cuvsResourcesManager.releaseResource(managedCuVSResources); } } @@ -207,10 +216,14 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) thro public void flush(int maxDoc, DocMap sortMap) throws IOException { flatVectorsWriter.flush(maxDoc, sortMap); for (var field : fields) { - if (sortMap == null) { - writeField(field); - } else { - writeSortingField(field, sortMap); + try { + if (sortMap == null) { + writeField(field); + } else { + writeSortingField(field, sortMap); + } + } catch (Exception e) { + throw new IOException(e.getMessage()); } } } @@ -220,8 +233,9 @@ public void flush(int maxDoc, DocMap sortMap) throws IOException { * * @param fieldData * @throws IOException + * @throws InterruptedException */ - private void writeField(FieldWriter fieldData) throws IOException { + private void writeField(FieldWriter fieldData) throws IOException, InterruptedException { writeFieldInternal(fieldData.fieldInfo(), fieldData.getFloatVectors()); } @@ -231,8 +245,10 @@ private void writeField(FieldWriter fieldData) throws IOException { * @param fieldData instance of GPUFieldWriter * @param sortMap instance of the DocMap * @throws IOException + * @throws InterruptedException */ - private void writeSortingField(FieldWriter fieldData, Sorter.DocMap sortMap) throws IOException { + private void writeSortingField(FieldWriter fieldData, Sorter.DocMap sortMap) + throws IOException, InterruptedException { DocsWithFieldSet oldDocsWithFieldSet = fieldData.getDocsWithFieldSet(); final int[] new2OldOrd = new int[oldDocsWithFieldSet.cardinality()]; mapOldOrdToNewOrd(oldDocsWithFieldSet, sortMap, null, new2OldOrd, null); @@ -325,7 +341,6 @@ public void finish() throws IOException { public void close() throws IOException { printInfoStream(infoStream, COMPONENT, "Closing resources"); IOUtils.close(hnswMeta, hnswVectorIndex, flatVectorsWriter); - closeCuVSResourcesInstance(); } /** diff --git a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsFormat.java b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsFormat.java index 0f8d9602..efc5aebe 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsFormat.java +++ b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsFormat.java @@ -33,6 +33,7 @@ public class LuceneAcceleratedHNSWBinaryQuantizedVectorsFormat extends KnnVector private static final int MAX_DIMENSIONS = 4096; private final AcceleratedHNSWParams acceleratedHNSWParams; + private final CuVSResourcesManager cuvsResourcesManager; static { try { @@ -63,6 +64,11 @@ public LuceneAcceleratedHNSWBinaryQuantizedVectorsFormat( AcceleratedHNSWParams acceleratedHNSWParams) { super("Lucene99AcceleratedHNSWBinaryQuantizedVectorsFormat"); this.acceleratedHNSWParams = acceleratedHNSWParams; + if (isSupported()) { + cuvsResourcesManager = new CuVSResourcesManager(acceleratedHNSWParams.getWriterThreads()); + } else { + cuvsResourcesManager = null; // Will not be needed in fallback mode. + } } /** @@ -76,7 +82,7 @@ public KnnVectorsWriter fieldsWriter(SegmentWriteState state) throws IOException Level.FINE, "cuVS is supported so using the Lucene99AcceleratedHNSWBinaryQuantizedVectorsWriter"); return new LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter( - state, acceleratedHNSWParams, flatWriter); + state, acceleratedHNSWParams, flatWriter, cuvsResourcesManager); } else { try { // Fallback to Lucene's Lucene102HnswBinaryQuantizedVectorsFormat format diff --git a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter.java b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter.java index d8a98cc2..467b5284 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter.java +++ b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter.java @@ -15,8 +15,6 @@ import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_INDEX_EXT; import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_META_CODEC_EXT; import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_META_CODEC_NAME; -import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.closeCuVSResourcesInstance; -import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.getCuVSResourcesInstance; import static org.apache.lucene.index.VectorEncoding.FLOAT32; import static org.apache.lucene.search.DocIdSetIterator.NO_MORE_DOCS; import static org.apache.lucene.util.RamUsageEstimator.shallowSizeOfInstance; @@ -25,6 +23,7 @@ import com.nvidia.cuvs.CagraIndexParams; import com.nvidia.cuvs.CuVSMatrix; import com.nvidia.cuvs.lucene.AcceleratedHNSWUtils.QuantizationType; +import com.nvidia.cuvs.lucene.CuVSResourcesManager.ManagedCuVSResources; import java.io.IOException; import java.util.ArrayList; import java.util.List; @@ -64,6 +63,7 @@ public class LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter extends KnnVector private final List fields = new ArrayList<>(); private final InfoStream infoStream; private final AcceleratedHNSWParams acceleratedHNSWParams; + private final CuVSResourcesManager cuvsResourcesManager; private IndexOutput hnswMeta = null, hnswVectorIndex = null; private boolean finished; private String vemFileName; @@ -80,13 +80,14 @@ public class LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter extends KnnVector public LuceneAcceleratedHNSWBinaryQuantizedVectorsWriter( SegmentWriteState state, AcceleratedHNSWParams acceleratedHNSWParams, - FlatVectorsWriter flatVectorsWriter) + FlatVectorsWriter flatVectorsWriter, + CuVSResourcesManager cuvsResourcesManager) throws IOException { super(); this.acceleratedHNSWParams = acceleratedHNSWParams; this.flatVectorsWriter = flatVectorsWriter; this.infoStream = state.infoStream; - + this.cuvsResourcesManager = cuvsResourcesManager; vemFileName = IndexFileNames.segmentFileName( state.segmentInfo.name, state.segmentSuffix, HNSW_META_CODEC_EXT); @@ -144,30 +145,35 @@ public KnnFieldVectorsWriter addField(FieldInfo fieldInfo) throws IOException * @param fieldInfo instance of FieldInfo that has the field description * @param vectors binary quantized vectors (packed bits as bytes) * @throws IOException + * @throws InterruptedException */ - private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOException { + private void writeFieldInternal(FieldInfo fieldInfo, List vectors) + throws IOException, InterruptedException { if (vectors.size() == 0) { writeEmpty(fieldInfo, hnswMeta); return; } + CagraIndexParams params = + CagraIndexParamsFactory.create( + acceleratedHNSWParams, vectors.size(), vectors.get(0).length); + ManagedCuVSResources managedCuVSResources = + cuvsResourcesManager.acquireResource(vectors.size(), vectors.get(0).length, params); + try { int dimensions = fieldInfo.getVectorDimension(); int bytesPerVector = (dimensions + 7) / 8; CuVSMatrix dataset = - Utils.createByteMatrix(vectors, bytesPerVector, getCuVSResourcesInstance()); + Utils.createByteMatrix(vectors, bytesPerVector, managedCuVSResources.getResource()); if (dataset.size() < 2) { writeSingleVectorGraph(fieldInfo, vectors); return; } - CagraIndexParams params = - CagraIndexParamsFactory.create(acceleratedHNSWParams, dataset.size(), dataset.columns()); - CagraIndex cagraIndex = - CagraIndex.newBuilder(getCuVSResourcesInstance()) + CagraIndex.newBuilder(managedCuVSResources.getResource()) .withDataset(dataset) .withIndexParams(params) .build(); @@ -186,7 +192,8 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throw acceleratedHNSWParams.getHnswLayers(), acceleratedHNSWParams.getGraphdegree(), params, - QuantizationType.BINARY); + QuantizationType.BINARY, + managedCuVSResources.getResource()); long vectorIndexOffset = hnswVectorIndex.getFilePointer(); // Write the graph to the vector index @@ -209,6 +216,8 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throw } catch (Throwable t) { Utils.handleThrowable(t); + } finally { + cuvsResourcesManager.releaseResource(managedCuVSResources); } } @@ -218,12 +227,16 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throw @Override public void flush(int maxDoc, DocMap sortMap) throws IOException { flatVectorsWriter.flush(maxDoc, sortMap); - for (var field : fields) { - if (sortMap == null) { - writeField(field); - } else { - writeSortingField(field, sortMap); + try { + for (var field : fields) { + if (sortMap == null) { + writeField(field); + } else { + writeSortingField(field, sortMap); + } } + } catch (Exception e) { + throw new IOException(e.getMessage()); } } @@ -232,8 +245,9 @@ public void flush(int maxDoc, DocMap sortMap) throws IOException { * * @param fieldData * @throws IOException + * @throws InterruptedException */ - private void writeField(FieldWriter fieldData) throws IOException { + private void writeField(FieldWriter fieldData) throws IOException, InterruptedException { writeFieldInternal(fieldData.fieldInfo(), fieldData.getByteVectors()); } @@ -243,8 +257,10 @@ private void writeField(FieldWriter fieldData) throws IOException { * @param fieldData instance of BinaryQuantizedGPUFieldWriter * @param sortMap instance of the DocMap * @throws IOException + * @throws InterruptedException */ - private void writeSortingField(FieldWriter fieldData, Sorter.DocMap sortMap) throws IOException { + private void writeSortingField(FieldWriter fieldData, Sorter.DocMap sortMap) + throws IOException, InterruptedException { DocsWithFieldSet oldDocsWithFieldSet = fieldData.getDocsWithFieldSet(); final int[] new2OldOrd = new int[oldDocsWithFieldSet.cardinality()]; // new ord to old ord @@ -357,7 +373,6 @@ public void finish() throws IOException { @Override public void close() throws IOException { IOUtils.close(hnswMeta, hnswVectorIndex, flatVectorsWriter); - closeCuVSResourcesInstance(); } /** diff --git a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsFormat.java b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsFormat.java index 534390b8..383504fa 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsFormat.java +++ b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsFormat.java @@ -30,6 +30,7 @@ public class LuceneAcceleratedHNSWScalarQuantizedVectorsFormat extends KnnVector private static final int MAX_DIMENSIONS = 4096; private final AcceleratedHNSWParams acceleratedHNSWParams; + private final CuVSResourcesManager cuvsResourcesManager; static { try { @@ -58,6 +59,11 @@ public LuceneAcceleratedHNSWScalarQuantizedVectorsFormat( AcceleratedHNSWParams acceleratedHNSWParams) { super("Lucene99AcceleratedHNSWScalarQuantizedVectorsFormat"); this.acceleratedHNSWParams = acceleratedHNSWParams; + if (isSupported()) { + cuvsResourcesManager = new CuVSResourcesManager(acceleratedHNSWParams.getWriterThreads()); + } else { + cuvsResourcesManager = null; // Will not be needed in fallback mode. + } } /** @@ -69,7 +75,7 @@ public KnnVectorsWriter fieldsWriter(SegmentWriteState state) throws IOException if (isSupported()) { log.info("cuVS is supported so using the Lucene99AcceleratedHNSWQuantizedVectorsWriter"); return new LuceneAcceleratedHNSWScalarQuantizedVectorsWriter( - state, acceleratedHNSWParams, flatWriter); + state, acceleratedHNSWParams, flatWriter, cuvsResourcesManager); } else { try { // Fallback to Lucene's Lucene99HnswScalarQuantizedVectorsFormat diff --git a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsWriter.java b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsWriter.java index 21b4be3f..39d65ae3 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsWriter.java +++ b/src/main/java/com/nvidia/cuvs/lucene/LuceneAcceleratedHNSWScalarQuantizedVectorsWriter.java @@ -15,8 +15,6 @@ import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_INDEX_EXT; import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_META_CODEC_EXT; import static com.nvidia.cuvs.lucene.Lucene99AcceleratedHNSWVectorsFormat.HNSW_META_CODEC_NAME; -import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.closeCuVSResourcesInstance; -import static com.nvidia.cuvs.lucene.ThreadLocalCuVSResourcesProvider.getCuVSResourcesInstance; import static org.apache.lucene.index.VectorEncoding.FLOAT32; import static org.apache.lucene.search.DocIdSetIterator.NO_MORE_DOCS; import static org.apache.lucene.util.RamUsageEstimator.shallowSizeOfInstance; @@ -25,6 +23,7 @@ import com.nvidia.cuvs.CagraIndexParams; import com.nvidia.cuvs.CuVSMatrix; import com.nvidia.cuvs.lucene.AcceleratedHNSWUtils.QuantizationType; +import com.nvidia.cuvs.lucene.CuVSResourcesManager.ManagedCuVSResources; import java.io.IOException; import java.util.ArrayList; import java.util.List; @@ -65,6 +64,7 @@ public class LuceneAcceleratedHNSWScalarQuantizedVectorsWriter extends KnnVector private final List fields = new ArrayList<>(); private final InfoStream infoStream; private final AcceleratedHNSWParams acceleratedHNSWParams; + private final CuVSResourcesManager cuvsResourcesManager; private IndexOutput hnswMeta = null, hnswVectorIndex = null; private boolean finished; private String vemFileName; @@ -90,12 +90,14 @@ public class LuceneAcceleratedHNSWScalarQuantizedVectorsWriter extends KnnVector public LuceneAcceleratedHNSWScalarQuantizedVectorsWriter( SegmentWriteState state, AcceleratedHNSWParams acceleratedHNSWParams, - FlatVectorsWriter flatVectorsWriter) + FlatVectorsWriter flatVectorsWriter, + CuVSResourcesManager cuvsResourcesManager) throws IOException { super(); this.acceleratedHNSWParams = acceleratedHNSWParams; this.flatVectorsWriter = flatVectorsWriter; this.infoStream = state.infoStream; + this.cuvsResourcesManager = cuvsResourcesManager; vemFileName = IndexFileNames.segmentFileName( @@ -164,13 +166,22 @@ private static byte[] convertSignedToUnsigned(byte[] signedVector) { * @param fieldInfo instance of FieldInfo that has the field description * @param vectors quantized vectors * @throws IOException + * @throws InterruptedException */ - private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOException { + private void writeFieldInternal(FieldInfo fieldInfo, List vectors) + throws IOException, InterruptedException { if (vectors.size() == 0) { writeEmpty(fieldInfo, hnswMeta); return; } + CagraIndexParams params = + CagraIndexParamsFactory.create( + acceleratedHNSWParams, vectors.size(), fieldInfo.getVectorDimension()); + ManagedCuVSResources managedCuVSResources = + cuvsResourcesManager.acquireResource( + vectors.size(), fieldInfo.getVectorDimension(), params); + try { int dimensions = fieldInfo.getVectorDimension(); @@ -182,18 +193,15 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOE // Create CuVSMatrix with BYTE data type (unsigned bytes) CuVSMatrix dataset = - Utils.createByteMatrix(unsignedVectors, dimensions, getCuVSResourcesInstance()); + Utils.createByteMatrix(unsignedVectors, dimensions, managedCuVSResources.getResource()); if (dataset.size() < 2) { writeSingleVectorGraph(fieldInfo, unsignedVectors); return; } - CagraIndexParams params = - CagraIndexParamsFactory.create(acceleratedHNSWParams, dataset.size(), dataset.columns()); - CagraIndex cagraIndex = - CagraIndex.newBuilder(getCuVSResourcesInstance()) + CagraIndex.newBuilder(managedCuVSResources.getResource()) .withDataset(dataset) .withIndexParams(params) .build(); @@ -211,7 +219,8 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOE acceleratedHNSWParams.getHnswLayers(), acceleratedHNSWParams.getGraphdegree(), params, - QuantizationType.SCALAR); + QuantizationType.SCALAR, + managedCuVSResources.getResource()); long vectorIndexOffset = hnswVectorIndex.getFilePointer(); @@ -235,6 +244,8 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOE cagraIndex.close(); } catch (Throwable t) { Utils.handleThrowable(t); + } finally { + cuvsResourcesManager.releaseResource(managedCuVSResources); } } @@ -244,12 +255,16 @@ private void writeFieldInternal(FieldInfo fieldInfo, List vectors) throws IOE @Override public void flush(int maxDoc, DocMap sortMap) throws IOException { flatVectorsWriter.flush(maxDoc, sortMap); - for (var field : fields) { - if (sortMap == null) { - writeField(field); - } else { - writeSortingField(field, sortMap); + try { + for (var field : fields) { + if (sortMap == null) { + writeField(field); + } else { + writeSortingField(field, sortMap); + } } + } catch (Exception e) { + throw new IOException(e.getMessage()); } } @@ -258,8 +273,9 @@ public void flush(int maxDoc, DocMap sortMap) throws IOException { * * @param fieldData * @throws IOException + * @throws InterruptedException */ - private void writeField(FieldWriter fieldData) throws IOException { + private void writeField(FieldWriter fieldData) throws IOException, InterruptedException { writeFieldInternal(fieldData.fieldInfo(), fieldData.getByteVectors()); } @@ -269,8 +285,10 @@ private void writeField(FieldWriter fieldData) throws IOException { * @param fieldData instance of ScalarQuantizedGPUFieldWriter * @param sortMap instance of the DocMap * @throws IOException + * @throws InterruptedException */ - private void writeSortingField(FieldWriter fieldData, Sorter.DocMap sortMap) throws IOException { + private void writeSortingField(FieldWriter fieldData, Sorter.DocMap sortMap) + throws IOException, InterruptedException { DocsWithFieldSet oldDocsWithFieldSet = fieldData.getDocsWithFieldSet(); final int[] new2OldOrd = new int[oldDocsWithFieldSet.cardinality()]; // new ord to old ord @@ -381,7 +399,6 @@ public void finish() throws IOException { @Override public void close() throws IOException { IOUtils.close(hnswMeta, hnswVectorIndex, flatVectorsWriter); - closeCuVSResourcesInstance(); } /** diff --git a/src/main/java/com/nvidia/cuvs/lucene/Utils.java b/src/main/java/com/nvidia/cuvs/lucene/Utils.java index c55aa7c3..28853e44 100644 --- a/src/main/java/com/nvidia/cuvs/lucene/Utils.java +++ b/src/main/java/com/nvidia/cuvs/lucene/Utils.java @@ -27,6 +27,11 @@ public class Utils { static final Logger log = Logger.getLogger(Utils.class.getName()); + public enum Target { + DEVICE, + HOST + } + /** * A utility method that throws specific types of throwable objects based on types. * @@ -51,19 +56,22 @@ static void handleThrowable(Throwable t) throws IOException { * @param data The float vectors * @param dimensions The number float elements in each vector * @param resources The CuVS resources for device matrix creation + * @param target To build the matrix on device or host * @return an instance of CuVSMatrix */ - static CuVSMatrix createFloatMatrix(List data, int dimensions, CuVSResources resources) { - // Use Builder pattern to avoid intermediate float[][] allocation - // and copy directly from List to device memory - CuVSMatrix.Builder builder = - CuVSMatrix.deviceBuilder( - resources, - data.size(), // rows (number of vectors) - dimensions, // columns (vector dimension) - CuVSMatrix.DataType.FLOAT); + static CuVSMatrix createFloatMatrix( + List data, int dimensions, CuVSResources resources, Target target) { + CuVSMatrix.Builder builder = null; + + switch (target) { + case DEVICE: + builder = + CuVSMatrix.deviceBuilder(resources, data.size(), dimensions, CuVSMatrix.DataType.FLOAT); + case HOST: + builder = CuVSMatrix.hostBuilder(data.size(), dimensions, CuVSMatrix.DataType.FLOAT); + } + assert builder != null; - // Add vectors one by one - builder copies directly to device memory for (float[] vector : data) { builder.addVector(vector); }