Skip to content

Commit b0ddd70

Browse files
committed
Restore tracked device resources on their original GPU
1 parent c515501 commit b0ddd70

4 files changed

Lines changed: 133 additions & 2 deletions

File tree

cpp/include/raft/core/memory_stats_resources.hpp

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
*/
55
#pragma once
66

7+
#include <raft/core/device_setter.hpp>
78
#include <raft/core/resource/device_memory_resource.hpp>
89
#include <raft/core/resource/managed_memory_resource.hpp>
910
#include <raft/core/resource/pinned_memory_resource.hpp>
@@ -76,6 +77,7 @@ class memory_stats_resources : public resources {
7677
public:
7778
explicit memory_stats_resources(const resources& existing)
7879
: resources(existing),
80+
device_id_(device_setter::get_current_device()),
7981
old_host_(mr::get_default_host_resource()),
8082
old_device_(rmm::mr::get_current_device_resource_ref())
8183
{
@@ -85,7 +87,7 @@ class memory_stats_resources : public resources {
8587
~memory_stats_resources() override
8688
{
8789
mr::set_default_host_resource(old_host_);
88-
rmm::mr::set_current_device_resource(std::move(old_device_));
90+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id_}, std::move(old_device_));
8991
}
9092

9193
memory_stats_resources(memory_stats_resources const&) = delete;
@@ -143,6 +145,7 @@ class memory_stats_resources : public resources {
143145

144146
std::vector<std::shared_ptr<resource::resource_cell>> snapshot_;
145147

148+
int device_id_;
146149
raft::mr::host_resource old_host_;
147150
raft::mr::device_resource old_device_;
148151

cpp/include/raft/core/memory_tracking_resources.hpp

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
#pragma once
66

77
#include <raft/core/detail/macros.hpp>
8+
#include <raft/core/device_setter.hpp>
89
#include <raft/core/resource/device_memory_resource.hpp>
910
#include <raft/core/resource/managed_memory_resource.hpp>
1011
#include <raft/core/resource/pinned_memory_resource.hpp>
@@ -108,7 +109,7 @@ class memory_tracking_resources : public resources {
108109
{
109110
report_.stop();
110111
raft::mr::set_default_host_resource(old_host_);
111-
rmm::mr::set_current_device_resource(old_device_);
112+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id_}, std::move(old_device_));
112113
}
113114

114115
memory_tracking_resources(memory_tracking_resources const&) = delete;
@@ -127,6 +128,7 @@ class memory_tracking_resources : public resources {
127128
: resources(existing ? *existing : resources{}),
128129
owned_stream_(std::move(owned_stream)),
129130
report_(out_override ? *out_override : *owned_stream_, sample_interval),
131+
device_id_(device_setter::get_current_device()),
130132
old_host_(raft::mr::get_default_host_resource()),
131133
old_device_(rmm::mr::get_current_device_resource_ref())
132134
{
@@ -141,6 +143,7 @@ class memory_tracking_resources : public resources {
141143
std::unique_ptr<std::ofstream> owned_stream_;
142144
raft::mr::resource_monitor report_;
143145

146+
int device_id_;
144147
raft::mr::host_resource old_host_;
145148
raft::mr::device_resource old_device_;
146149

cpp/tests/core/memory_stats_resources.cpp

Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,19 +4,37 @@
44
*/
55

66
#include <raft/core/memory_stats_resources.hpp>
7+
#include <raft/core/device_setter.hpp>
78
#include <raft/core/resource/device_memory_resource.hpp>
89
#include <raft/core/resources.hpp>
910

11+
#include <rmm/mr/cuda_memory_resource.hpp>
1012
#include <rmm/mr/per_device_resource.hpp>
13+
#include <rmm/mr/pool_memory_resource.hpp>
1114
#include <rmm/resource_ref.hpp>
1215

16+
#include <cuda/memory_resource>
1317
#include <cuda/stream_ref>
1418

1519
#include <gtest/gtest.h>
1620

1721
#include <cstddef>
22+
#include <memory>
1823

1924
namespace raft {
25+
namespace {
26+
27+
struct device_resource_restore_guard {
28+
int device_id;
29+
raft::mr::device_resource resource;
30+
31+
~device_resource_restore_guard()
32+
{
33+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id}, std::move(resource));
34+
}
35+
};
36+
37+
} // namespace
2038

2139
TEST(MemoryStatsResources, IndependentCounting_DefaultWorkspace)
2240
{
@@ -93,4 +111,49 @@ TEST(MemoryStatsResources, IndependentCounting_PoolWorkspace)
93111
dev_mr.deallocate(cuda::stream_ref{cudaStreamLegacy}, dev_ptr, kGlobalSize);
94112
}
95113

114+
TEST(MemoryStatsResources, RestoresDeviceResourceOnConstructionDevice)
115+
{
116+
if (device_setter::get_device_count() < 2) {
117+
GTEST_SKIP() << "Requires at least 2 CUDA devices";
118+
}
119+
120+
auto device0 = 0;
121+
auto device1 = 1;
122+
123+
auto device0_guard = [&]() {
124+
auto scoped_device = device_setter{device0};
125+
auto upstream = rmm::mr::get_current_device_resource_ref();
126+
return device_resource_restore_guard{
127+
device0,
128+
rmm::mr::set_current_device_resource(
129+
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)})};
130+
}();
131+
132+
auto device1_guard = [&]() {
133+
auto scoped_device = device_setter{device1};
134+
return device_resource_restore_guard{device1, rmm::mr::reset_current_device_resource()};
135+
}();
136+
137+
{
138+
auto scoped_device = device_setter{device0};
139+
raft::resources res;
140+
auto tracked = std::make_unique<memory_stats_resources>(res);
141+
auto wrong_device = device_setter{device1};
142+
static_cast<void>(wrong_device);
143+
tracked.reset();
144+
}
145+
146+
{
147+
auto scoped_device = device_setter{device0};
148+
auto current_mr = rmm::mr::get_current_device_resource_ref();
149+
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr), nullptr);
150+
}
151+
152+
{
153+
auto scoped_device = device_setter{device1};
154+
auto current_mr = rmm::mr::get_current_device_resource_ref();
155+
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr), nullptr);
156+
}
157+
}
158+
96159
} // namespace raft

cpp/tests/core/monitor_resources.cu

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,19 +4,36 @@
44
*/
55

66
#include <raft/core/device_mdarray.hpp>
7+
#include <raft/core/device_setter.hpp>
78
#include <raft/core/memory_tracking_resources.hpp>
89
#include <raft/core/resources.hpp>
910

11+
#include <rmm/mr/cuda_memory_resource.hpp>
12+
#include <rmm/mr/per_device_resource.hpp>
13+
#include <rmm/mr/pool_memory_resource.hpp>
14+
15+
#include <cuda/memory_resource>
1016
#include <gtest/gtest.h>
1117

1218
#include <algorithm>
1319
#include <chrono>
20+
#include <memory>
1421
#include <sstream>
1522
#include <string>
1623
#include <thread>
1724

1825
namespace {
1926

27+
struct device_resource_restore_guard {
28+
int device_id;
29+
raft::mr::device_resource resource;
30+
31+
~device_resource_restore_guard()
32+
{
33+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id}, std::move(resource));
34+
}
35+
};
36+
2037
TEST(MemoryTrackingResources, TracksDeviceAllocations)
2138
{
2239
using namespace std::chrono_literals;
@@ -49,4 +66,49 @@ TEST(MemoryTrackingResources, TracksDeviceAllocations)
4966
<< output;
5067
}
5168

69+
TEST(MemoryTrackingResources, RestoresDeviceResourceOnConstructionDevice)
70+
{
71+
if (raft::device_setter::get_device_count() < 2) {
72+
GTEST_SKIP() << "Requires at least 2 CUDA devices";
73+
}
74+
75+
auto device0 = 0;
76+
auto device1 = 1;
77+
78+
auto device0_guard = [&]() {
79+
auto scoped_device = raft::device_setter{device0};
80+
auto upstream = rmm::mr::get_current_device_resource_ref();
81+
return device_resource_restore_guard{
82+
device0,
83+
rmm::mr::set_current_device_resource(
84+
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)})};
85+
}();
86+
87+
auto device1_guard = [&]() {
88+
auto scoped_device = raft::device_setter{device1};
89+
return device_resource_restore_guard{device1, rmm::mr::reset_current_device_resource()};
90+
}();
91+
92+
{
93+
auto scoped_device = raft::device_setter{device0};
94+
std::ostringstream oss;
95+
auto tracked = std::make_unique<raft::memory_tracking_resources>(oss);
96+
auto wrong_device = raft::device_setter{device1};
97+
static_cast<void>(wrong_device);
98+
tracked.reset();
99+
}
100+
101+
{
102+
auto scoped_device = raft::device_setter{device0};
103+
auto current_mr = rmm::mr::get_current_device_resource_ref();
104+
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr), nullptr);
105+
}
106+
107+
{
108+
auto scoped_device = raft::device_setter{device1};
109+
auto current_mr = rmm::mr::get_current_device_resource_ref();
110+
EXPECT_NE(cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr), nullptr);
111+
}
112+
}
113+
52114
} // namespace

0 commit comments

Comments
 (0)