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
29 changes: 29 additions & 0 deletions c/include/cuvs/core/c_api.h
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,25 @@ CUVS_EXPORT cuvsError_t cuvsResourcesCreateWithMemoryTracking(cuvsResources_t* r
*/
CUVS_EXPORT cuvsError_t cuvsResourcesDestroy(cuvsResources_t res);

/**
* @brief Set a memory pool on the device used by these resources
*
* @param[in] res cuvsResources_t opaque C handle
* @param[in] percent_of_free_memory Percentage of free device memory to allocate for the pool
* @return cuvsError_t
*/
CUVS_EXPORT cuvsError_t cuvsResourcesSetMemoryPool(cuvsResources_t res,
int percent_of_free_memory);

/**
* @brief Set a CUDA stream pool on these resources
*
* @param[in] res cuvsResources_t opaque C handle
* @param[in] num_streams Number of non-blocking CUDA streams in the pool
* @return cuvsError_t
*/
CUVS_EXPORT cuvsError_t cuvsResourcesSetStreamPool(cuvsResources_t res, size_t num_streams);

/**
* @brief Set cudaStream_t on cuvsResources_t to queue CUDA kernels on APIs
* that accept a cuvsResources_t handle
Expand Down Expand Up @@ -211,6 +230,16 @@ CUVS_EXPORT cuvsError_t cuvsMultiGpuResourcesDestroy(cuvsResources_t res);
* @return cuvsError_t
*/
CUVS_EXPORT cuvsError_t cuvsMultiGpuResourcesSetMemoryPool(cuvsResources_t res, int percent_of_free_memory);

/**
* @brief Set a CUDA stream pool on all devices managed by the multi-GPU resources
*
* @param[in] res cuvsResources_t opaque C handle for multi-GPU resources
* @param[in] num_streams Number of CUDA streams in each device's pool
* @return cuvsError_t
*/
CUVS_EXPORT cuvsError_t cuvsMultiGpuResourcesSetStreamPool(cuvsResources_t res,
Comment thread
tarang-jain marked this conversation as resolved.
size_t num_streams);
/** @} */

/**
Expand Down
84 changes: 83 additions & 1 deletion c/src/core/c_api.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,16 @@
#include <raft/core/device_resources_snmg.hpp>
#include <raft/core/memory_tracking_resources.hpp>
#include <raft/core/resource/cuda_stream.hpp>
#include <raft/core/resource/cuda_stream_pool.hpp>
#include <raft/core/resource/device_id.hpp>
#include <raft/core/resource/device_memory_resource.hpp>
#include <raft/core/resource/multi_gpu.hpp>
#include <raft/core/resource/resource_types.hpp>
#include <raft/core/resources.hpp>
#include <raft/util/cudart_utils.hpp>
#include <rapids_logger/logger.hpp>
#include <rmm/cuda_device.hpp>
#include <rmm/cuda_stream_pool.hpp>
#include <rmm/cuda_stream_view.hpp>
#include <rmm/mr/cuda_async_memory_resource.hpp>
#include <rmm/mr/cuda_memory_resource.hpp>
Expand All @@ -29,18 +33,78 @@
#include <chrono>
#include <cstdint>
#include <memory>
#include <optional>
#include <stdexcept>
#include <string>
#include <thread>

namespace {

class single_gpu_resources : public raft::resources {
public:
~single_gpu_resources() override { reset_memory_pool(); }

void set_memory_pool(int percent_of_free_memory)
{
RAFT_EXPECTS(percent_of_free_memory > 0 && percent_of_free_memory <= 100,
"percent_of_free_memory must be in the range [1, 100]");

reset_memory_pool();
pool_device_id_ = rmm::get_current_cuda_device();
auto pool = rmm::mr::pool_memory_resource{
rmm::mr::get_current_device_resource_ref(),
rmm::percent_of_free_device_memory(percent_of_free_memory)};
previous_memory_resource_.emplace(
rmm::mr::set_per_device_resource(*pool_device_id_, std::move(pool)));
}

private:
void reset_memory_pool()
{
if (!previous_memory_resource_.has_value()) { return; }

rmm::cuda_set_device_raii device_guard{*pool_device_id_};
rmm::mr::set_per_device_resource(*pool_device_id_, std::move(*previous_memory_resource_));
previous_memory_resource_.reset();
pool_device_id_.reset();
}

std::optional<rmm::cuda_device_id> pool_device_id_;
std::optional<raft::mr::device_resource> previous_memory_resource_;
};

} // namespace

extern "C" cuvsError_t cuvsResourcesCreate(cuvsResources_t* res)
{
return cuvs::core::translate_exceptions([=] {
auto res_ptr = new raft::resources{};
auto res_ptr = new single_gpu_resources{};
*res = reinterpret_cast<uintptr_t>(res_ptr);
});
}

extern "C" cuvsError_t cuvsResourcesSetMemoryPool(cuvsResources_t res,
int percent_of_free_memory)
{
return cuvs::core::translate_exceptions([=] {
auto res_ptr = dynamic_cast<single_gpu_resources*>(reinterpret_cast<raft::resources*>(res));
RAFT_EXPECTS(res_ptr != nullptr,
"memory pools are not supported on memory-tracking resources");
res_ptr->set_memory_pool(percent_of_free_memory);
});
}

extern "C" cuvsError_t cuvsResourcesSetStreamPool(cuvsResources_t res, size_t num_streams)
{
return cuvs::core::translate_exceptions([=] {
RAFT_EXPECTS(num_streams > 0, "num_streams must be greater than zero");
auto res_ptr = reinterpret_cast<raft::resources*>(res);
RAFT_EXPECTS(res_ptr != nullptr, "res must not be NULL");
raft::resource::set_cuda_stream_pool(
*res_ptr, std::make_shared<rmm::cuda_stream_pool>(num_streams));
});
}

extern "C" cuvsError_t cuvsResourcesSetWorkspacePool(cuvsResources_t res, size_t initial_size_bytes)
{
return cuvs::core::translate_exceptions([=] {
Expand Down Expand Up @@ -132,6 +196,24 @@ extern "C" cuvsError_t cuvsMultiGpuResourcesSetMemoryPool(cuvsResources_t res,
});
}

extern "C" cuvsError_t cuvsMultiGpuResourcesSetStreamPool(cuvsResources_t res,
size_t num_streams)
{
return cuvs::core::translate_exceptions([=] {
RAFT_EXPECTS(num_streams > 0, "num_streams must be greater than zero");
auto res_ptr = reinterpret_cast<raft::device_resources_snmg*>(res);
RAFT_EXPECTS(res_ptr != nullptr, "res must not be NULL");

auto& device_resources = raft::resource::get_multi_gpu_resource(*res_ptr);
for (auto& device_resource : device_resources) {
rmm::cuda_set_device_raii device_guard{
rmm::cuda_device_id{raft::resource::get_device_id(device_resource)}};
raft::resource::set_cuda_stream_pool(
device_resource, std::make_shared<rmm::cuda_stream_pool>(num_streams));
}
});
}

extern "C" cuvsError_t cuvsStreamSet(cuvsResources_t res, cudaStream_t stream)
{
return cuvs::core::translate_exceptions([=] {
Expand Down
Loading
Loading