diff --git a/cpp/include/raft/core/memory_stats_resources.hpp b/cpp/include/raft/core/memory_stats_resources.hpp index 136b48beb6..7e2dd7c166 100644 --- a/cpp/include/raft/core/memory_stats_resources.hpp +++ b/cpp/include/raft/core/memory_stats_resources.hpp @@ -4,6 +4,8 @@ */ #pragma once +#include +#include #include #include #include @@ -20,6 +22,7 @@ #include #include +#include #include #include #include @@ -77,7 +80,8 @@ 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(); } @@ -85,7 +89,17 @@ class memory_stats_resources : public resources { ~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"); + } } memory_stats_resources(memory_stats_resources const&) = delete; @@ -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(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 --- { diff --git a/cpp/include/raft/core/memory_tracking_resources.hpp b/cpp/include/raft/core/memory_tracking_resources.hpp index d06e14ef0e..60121ed9fe 100644 --- a/cpp/include/raft/core/memory_tracking_resources.hpp +++ b/cpp/include/raft/core/memory_tracking_resources.hpp @@ -5,6 +5,8 @@ #pragma once #include +#include +#include #include #include #include @@ -22,6 +24,7 @@ #include #include +#include #include #include #include @@ -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; @@ -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(); } @@ -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(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) --- diff --git a/cpp/tests/core/memory_stats_resources.cpp b/cpp/tests/core/memory_stats_resources.cpp index be615c9041..cd275dace0 100644 --- a/cpp/tests/core/memory_stats_resources.cpp +++ b/cpp/tests/core/memory_stats_resources.cpp @@ -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 #include #include #include +#include #include +#include #include +#include #include #include #include +#include namespace raft { @@ -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(res); + auto wrong_device = device_setter{device1}; + static_cast(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(¤t_mr), nullptr); + } + + { + auto scoped_device = device_setter{device1}; + auto current_mr = rmm::mr::get_current_device_resource_ref(); + EXPECT_NE(cuda::mr::resource_cast(¤t_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(resource::get_device_id(res)); + + auto wrong_device = device_setter{device1}; + static_cast(wrong_device); + auto tracked = std::make_unique(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 diff --git a/cpp/tests/core/monitor_resources.cu b/cpp/tests/core/monitor_resources.cu index c7be4bd285..fb23779d4f 100644 --- a/cpp/tests/core/monitor_resources.cu +++ b/cpp/tests/core/monitor_resources.cu @@ -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 +#include #include #include +#include +#include +#include + +#include + #include #include #include +#include #include #include #include @@ -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(oss); + auto wrong_device = raft::device_setter{device1}; + static_cast(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(¤t_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(¤t_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(raft::resource::get_device_id(res)); + + auto wrong_device = raft::device_setter{device1}; + static_cast(wrong_device); + std::ostringstream oss; + auto tracked = std::make_unique(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 diff --git a/cpp/tests/core/test_memory_resource.hpp b/cpp/tests/core/test_memory_resource.hpp new file mode 100644 index 0000000000..7cb48291cc --- /dev/null +++ b/cpp/tests/core/test_memory_resource.hpp @@ -0,0 +1,68 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include +#include + +#include +#include +#include + +#include + +#include + +#include +#include + +namespace raft::test { + +struct device_resource_restore_guard { + int device_id; + raft::mr::device_resource resource; + + ~device_resource_restore_guard() + { + try { + rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id}, std::move(resource)); + } catch (const std::exception& e) { + ADD_FAILURE() << "Failed to restore device " << device_id << " memory resource: " << e.what(); + } catch (...) { + ADD_FAILURE() << "Failed to restore device " << device_id + << " memory resource: unknown exception"; + } + } +}; + +inline auto install_pool_device_resource(int device_id) -> device_resource_restore_guard +{ + auto scoped_device = raft::device_setter{device_id}; + auto upstream = rmm::mr::get_current_device_resource_ref(); + auto installed_resource = + raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)}; + auto old_resource = rmm::mr::set_current_device_resource(std::move(installed_resource)); + return device_resource_restore_guard{device_id, std::move(old_resource)}; +} + +inline auto install_default_device_resource(int device_id) -> device_resource_restore_guard +{ + auto scoped_device = raft::device_setter{device_id}; + return device_resource_restore_guard{device_id, rmm::mr::reset_current_device_resource()}; +} + +inline auto current_device_uses_pool_resource() -> bool +{ + auto current_mr = rmm::mr::get_current_device_resource_ref(); + return cuda::mr::resource_cast(¤t_mr) != nullptr; +} + +inline auto current_device_uses_default_cuda_resource() -> bool +{ + auto current_mr = rmm::mr::get_current_device_resource_ref(); + return cuda::mr::resource_cast(¤t_mr) != nullptr; +} + +} // namespace raft::test