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
21 changes: 18 additions & 3 deletions cpp/include/raft/core/memory_stats_resources.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
*/
#pragma once

#include <raft/core/logger.hpp>
#include <raft/core/resource/device_id.hpp>
#include <raft/core/resource/device_memory_resource.hpp>
#include <raft/core/resource/managed_memory_resource.hpp>
#include <raft/core/resource/pinned_memory_resource.hpp>
Expand All @@ -20,6 +22,7 @@

#include <cstddef>
#include <cstdint>
#include <exception>
#include <memory>
#include <utility>
#include <vector>
Expand Down Expand Up @@ -77,15 +80,26 @@ class memory_stats_resources : public resources {
explicit memory_stats_resources(const resources& existing)
: resources(existing),
old_host_(mr::get_default_host_resource()),
old_device_(rmm::mr::get_current_device_resource_ref())
old_device_(
rmm::mr::get_per_device_resource_ref(rmm::cuda_device_id{resource::get_device_id(*this)}))
{
init();
}

~memory_stats_resources() override
{
mr::set_default_host_resource(old_host_);
rmm::mr::set_current_device_resource(std::move(old_device_));
try {
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
std::move(old_device_));
} catch (const std::exception& e) {
RAFT_LOG_ERROR("memory_stats_resources failed to restore the per-device memory resource: %s",
e.what());
} catch (...) {
RAFT_LOG_ERROR(
"memory_stats_resources failed to restore the per-device memory resource: unknown "
"exception");
}
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

memory_stats_resources(memory_stats_resources const&) = delete;
Expand Down Expand Up @@ -215,7 +229,8 @@ class memory_stats_resources : public resources {
device_stats_adaptor_t sa{rmm::device_async_resource_ref{old_device_}};
device_stats_ = sa.get_stats();
device_adaptor_ = std::make_unique<device_stats_adaptor_t>(std::move(sa));
rmm::mr::set_current_device_resource(*device_adaptor_);
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
*device_adaptor_);
}
// --- Workspace ---
{
Expand Down
21 changes: 18 additions & 3 deletions cpp/include/raft/core/memory_tracking_resources.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
#pragma once

#include <raft/core/detail/macros.hpp>
#include <raft/core/logger.hpp>
#include <raft/core/resource/device_id.hpp>
#include <raft/core/resource/device_memory_resource.hpp>
#include <raft/core/resource/managed_memory_resource.hpp>
#include <raft/core/resource/pinned_memory_resource.hpp>
Expand All @@ -22,6 +24,7 @@
#include <cuda/stream_ref>

#include <chrono>
#include <exception>
#include <fstream>
#include <memory>
#include <ostream>
Expand Down Expand Up @@ -108,7 +111,17 @@ class memory_tracking_resources : public resources {
{
report_.stop();
raft::mr::set_default_host_resource(old_host_);
rmm::mr::set_current_device_resource(old_device_);
try {
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
std::move(old_device_));
} catch (const std::exception& e) {
RAFT_LOG_ERROR(
"memory_tracking_resources failed to restore the per-device memory resource: %s", e.what());
} catch (...) {
RAFT_LOG_ERROR(
"memory_tracking_resources failed to restore the per-device memory resource: unknown "
"exception");
}
}

memory_tracking_resources(memory_tracking_resources const&) = delete;
Expand All @@ -128,7 +141,8 @@ class memory_tracking_resources : public resources {
owned_stream_(std::move(owned_stream)),
report_(out_override ? *out_override : *owned_stream_, sample_interval),
old_host_(raft::mr::get_default_host_resource()),
old_device_(rmm::mr::get_current_device_resource_ref())
old_device_(
rmm::mr::get_per_device_resource_ref(rmm::cuda_device_id{resource::get_device_id(*this)}))
{
init();
}
Expand Down Expand Up @@ -217,7 +231,8 @@ class memory_tracking_resources : public resources {
device_stats_t sa{rmm::device_async_resource_ref{old_device_}};
report_.register_source("device", sa.get_stats());
device_adaptor_ = std::make_unique<device_notify_t>(std::move(sa), report_.get_notifier());
rmm::mr::set_current_device_resource(*device_adaptor_);
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
*device_adaptor_);
}

// --- Workspace (track upstream to preserve limiting_resource_adaptor) ---
Expand Down
84 changes: 83 additions & 1 deletion cpp/tests/core/memory_stats_resources.cpp
Original file line number Diff line number Diff line change
@@ -1,20 +1,27 @@
/*
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION.
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: Apache-2.0
*/

#include "test_memory_resource.hpp"

#include <raft/core/device_setter.hpp>
#include <raft/core/memory_stats_resources.hpp>
#include <raft/core/resource/device_memory_resource.hpp>
#include <raft/core/resources.hpp>

#include <rmm/mr/cuda_memory_resource.hpp>
#include <rmm/mr/per_device_resource.hpp>
#include <rmm/mr/pool_memory_resource.hpp>
#include <rmm/resource_ref.hpp>

#include <cuda/memory_resource>
#include <cuda/stream_ref>

#include <gtest/gtest.h>

#include <cstddef>
#include <memory>

namespace raft {

Expand Down Expand Up @@ -93,4 +100,79 @@ TEST(MemoryStatsResources, IndependentCounting_PoolWorkspace)
dev_mr.deallocate(cuda::stream_ref{cudaStreamLegacy}, dev_ptr, kGlobalSize);
}

TEST(MemoryStatsResources, RestoresDeviceResourceOnConstructionDevice)
{
if (device_setter::get_device_count() < 2) { GTEST_SKIP() << "Requires at least 2 CUDA devices"; }

auto device0 = 0;
auto device1 = 1;

auto device0_guard = test::install_pool_device_resource(device0);
auto device1_guard = test::install_default_device_resource(device1);

{
auto scoped_device = device_setter{device0};
raft::resources res;
auto tracked = std::make_unique<memory_stats_resources>(res);
auto wrong_device = device_setter{device1};
static_cast<void>(wrong_device);
tracked.reset();
}

{
auto scoped_device = device_setter{device0};
auto current_mr = rmm::mr::get_current_device_resource_ref();
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr), nullptr);
}

{
auto scoped_device = device_setter{device1};
auto current_mr = rmm::mr::get_current_device_resource_ref();
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr), nullptr);
}
}

TEST(MemoryStatsResources, InstallsTrackedResourceOnHandleDevice)
{
if (device_setter::get_device_count() < 2) { GTEST_SKIP() << "Requires at least 2 CUDA devices"; }

auto device0 = 0;
auto device1 = 1;

auto device0_guard = test::install_pool_device_resource(device0);
auto device1_guard = test::install_default_device_resource(device1);

{
auto scoped_device = device_setter{device0};
raft::resources res;
static_cast<void>(resource::get_device_id(res));

auto wrong_device = device_setter{device1};
static_cast<void>(wrong_device);
auto tracked = std::make_unique<memory_stats_resources>(res);

{
auto verify_device0 = device_setter{device0};
EXPECT_FALSE(test::current_device_uses_pool_resource());
}

{
auto verify_device1 = device_setter{device1};
EXPECT_TRUE(test::current_device_uses_default_cuda_resource());
}

tracked.reset();
}

{
auto scoped_device = device_setter{device0};
EXPECT_TRUE(test::current_device_uses_pool_resource());
}

{
auto scoped_device = device_setter{device1};
EXPECT_TRUE(test::current_device_uses_default_cuda_resource());
}
}

} // namespace raft
92 changes: 91 additions & 1 deletion cpp/tests/core/monitor_resources.cu
Original file line number Diff line number Diff line change
@@ -1,16 +1,26 @@
/*
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION.
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: Apache-2.0
*/

#include "test_memory_resource.hpp"

#include <raft/core/device_mdarray.hpp>
#include <raft/core/device_setter.hpp>
#include <raft/core/memory_tracking_resources.hpp>
#include <raft/core/resources.hpp>

#include <rmm/mr/cuda_memory_resource.hpp>
#include <rmm/mr/per_device_resource.hpp>
#include <rmm/mr/pool_memory_resource.hpp>

#include <cuda/memory_resource>

#include <gtest/gtest.h>

#include <algorithm>
#include <chrono>
#include <memory>
#include <sstream>
#include <string>
#include <thread>
Expand Down Expand Up @@ -49,4 +59,84 @@ TEST(MemoryTrackingResources, TracksDeviceAllocations)
<< output;
}

TEST(MemoryTrackingResources, RestoresDeviceResourceOnConstructionDevice)
{
if (raft::device_setter::get_device_count() < 2) {
GTEST_SKIP() << "Requires at least 2 CUDA devices";
}

auto device0 = 0;
auto device1 = 1;

auto device0_guard = raft::test::install_pool_device_resource(device0);
auto device1_guard = raft::test::install_default_device_resource(device1);

{
auto scoped_device = raft::device_setter{device0};
std::ostringstream oss;
auto tracked = std::make_unique<raft::memory_tracking_resources>(oss);
auto wrong_device = raft::device_setter{device1};
static_cast<void>(wrong_device);
tracked.reset();
}

{
auto scoped_device = raft::device_setter{device0};
auto current_mr = rmm::mr::get_current_device_resource_ref();
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr), nullptr);
}

{
auto scoped_device = raft::device_setter{device1};
auto current_mr = rmm::mr::get_current_device_resource_ref();
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr), nullptr);
}
}

TEST(MemoryTrackingResources, InstallsTrackedResourceOnHandleDevice)
{
if (raft::device_setter::get_device_count() < 2) {
GTEST_SKIP() << "Requires at least 2 CUDA devices";
}

auto device0 = 0;
auto device1 = 1;

auto device0_guard = raft::test::install_pool_device_resource(device0);
auto device1_guard = raft::test::install_default_device_resource(device1);

{
auto scoped_device = raft::device_setter{device0};
raft::resources res;
static_cast<void>(raft::resource::get_device_id(res));

auto wrong_device = raft::device_setter{device1};
static_cast<void>(wrong_device);
std::ostringstream oss;
auto tracked = std::make_unique<raft::memory_tracking_resources>(res, oss);

{
auto verify_device0 = raft::device_setter{device0};
EXPECT_FALSE(raft::test::current_device_uses_pool_resource());
}

{
auto verify_device1 = raft::device_setter{device1};
EXPECT_TRUE(raft::test::current_device_uses_default_cuda_resource());
}

tracked.reset();
}

{
auto scoped_device = raft::device_setter{device0};
EXPECT_TRUE(raft::test::current_device_uses_pool_resource());
}

{
auto scoped_device = raft::device_setter{device1};
EXPECT_TRUE(raft::test::current_device_uses_default_cuda_resource());
}
}

} // namespace
Loading
Loading