Skip to content

Commit 38ca3a2

Browse files
committed
Address memory resource review feedback
1 parent 7c41eaa commit 38ca3a2

5 files changed

Lines changed: 114 additions & 111 deletions

File tree

cpp/include/raft/core/memory_stats_resources.hpp

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

7+
#include <raft/core/logger.hpp>
78
#include <raft/core/resource/device_id.hpp>
89
#include <raft/core/resource/device_memory_resource.hpp>
910
#include <raft/core/resource/managed_memory_resource.hpp>
@@ -21,6 +22,7 @@
2122

2223
#include <cstddef>
2324
#include <cstdint>
25+
#include <exception>
2426
#include <memory>
2527
#include <utility>
2628
#include <vector>
@@ -87,8 +89,17 @@ class memory_stats_resources : public resources {
8789
~memory_stats_resources() override
8890
{
8991
mr::set_default_host_resource(old_host_);
90-
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
91-
std::move(old_device_));
92+
try {
93+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
94+
std::move(old_device_));
95+
} catch (const std::exception& e) {
96+
RAFT_LOG_ERROR("memory_stats_resources failed to restore the per-device memory resource: %s",
97+
e.what());
98+
} catch (...) {
99+
RAFT_LOG_ERROR(
100+
"memory_stats_resources failed to restore the per-device memory resource: unknown "
101+
"exception");
102+
}
92103
}
93104

94105
memory_stats_resources(memory_stats_resources const&) = delete;

cpp/include/raft/core/memory_tracking_resources.hpp

Lines changed: 13 additions & 2 deletions
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/logger.hpp>
89
#include <raft/core/resource/device_id.hpp>
910
#include <raft/core/resource/device_memory_resource.hpp>
1011
#include <raft/core/resource/managed_memory_resource.hpp>
@@ -23,6 +24,7 @@
2324
#include <cuda/stream_ref>
2425

2526
#include <chrono>
27+
#include <exception>
2628
#include <fstream>
2729
#include <memory>
2830
#include <ostream>
@@ -109,8 +111,17 @@ class memory_tracking_resources : public resources {
109111
{
110112
report_.stop();
111113
raft::mr::set_default_host_resource(old_host_);
112-
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
113-
std::move(old_device_));
114+
try {
115+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{resource::get_device_id(*this)},
116+
std::move(old_device_));
117+
} catch (const std::exception& e) {
118+
RAFT_LOG_ERROR(
119+
"memory_tracking_resources failed to restore the per-device memory resource: %s", e.what());
120+
} catch (...) {
121+
RAFT_LOG_ERROR(
122+
"memory_tracking_resources failed to restore the per-device memory resource: unknown "
123+
"exception");
124+
}
114125
}
115126

116127
memory_tracking_resources(memory_tracking_resources const&) = delete;

cpp/tests/core/memory_stats_resources.cpp

Lines changed: 10 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

6+
#include "test_memory_resource.hpp"
7+
68
#include <raft/core/device_setter.hpp>
79
#include <raft/core/memory_stats_resources.hpp>
810
#include <raft/core/resource/device_memory_resource.hpp>
@@ -22,31 +24,6 @@
2224
#include <memory>
2325

2426
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-
auto current_device_uses_pool_resource() -> bool
38-
{
39-
auto current_mr = rmm::mr::get_current_device_resource_ref();
40-
return cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr) != nullptr;
41-
}
42-
43-
auto current_device_uses_default_cuda_resource() -> bool
44-
{
45-
auto current_mr = rmm::mr::get_current_device_resource_ref();
46-
return cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr) != nullptr;
47-
}
48-
49-
} // namespace
5027

5128
TEST(MemoryStatsResources, IndependentCounting_DefaultWorkspace)
5229
{
@@ -130,19 +107,8 @@ TEST(MemoryStatsResources, RestoresDeviceResourceOnConstructionDevice)
130107
auto device0 = 0;
131108
auto device1 = 1;
132109

133-
auto device0_guard = [&]() {
134-
auto scoped_device = device_setter{device0};
135-
auto upstream = rmm::mr::get_current_device_resource_ref();
136-
return device_resource_restore_guard{
137-
device0,
138-
rmm::mr::set_current_device_resource(
139-
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)})};
140-
}();
141-
142-
auto device1_guard = [&]() {
143-
auto scoped_device = device_setter{device1};
144-
return device_resource_restore_guard{device1, rmm::mr::reset_current_device_resource()};
145-
}();
110+
auto device0_guard = test::install_pool_device_resource(device0);
111+
auto device1_guard = test::install_default_device_resource(device1);
146112

147113
{
148114
auto scoped_device = device_setter{device0};
@@ -173,19 +139,8 @@ TEST(MemoryStatsResources, InstallsTrackedResourceOnHandleDevice)
173139
auto device0 = 0;
174140
auto device1 = 1;
175141

176-
auto device0_guard = [&]() {
177-
auto scoped_device = device_setter{device0};
178-
auto upstream = rmm::mr::get_current_device_resource_ref();
179-
return device_resource_restore_guard{
180-
device0,
181-
rmm::mr::set_current_device_resource(
182-
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)})};
183-
}();
184-
185-
auto device1_guard = [&]() {
186-
auto scoped_device = device_setter{device1};
187-
return device_resource_restore_guard{device1, rmm::mr::reset_current_device_resource()};
188-
}();
142+
auto device0_guard = test::install_pool_device_resource(device0);
143+
auto device1_guard = test::install_default_device_resource(device1);
189144

190145
{
191146
auto scoped_device = device_setter{device0};
@@ -198,25 +153,25 @@ TEST(MemoryStatsResources, InstallsTrackedResourceOnHandleDevice)
198153

199154
{
200155
auto verify_device0 = device_setter{device0};
201-
EXPECT_FALSE(current_device_uses_pool_resource());
156+
EXPECT_FALSE(test::current_device_uses_pool_resource());
202157
}
203158

204159
{
205160
auto verify_device1 = device_setter{device1};
206-
EXPECT_TRUE(current_device_uses_default_cuda_resource());
161+
EXPECT_TRUE(test::current_device_uses_default_cuda_resource());
207162
}
208163

209164
tracked.reset();
210165
}
211166

212167
{
213168
auto scoped_device = device_setter{device0};
214-
EXPECT_TRUE(current_device_uses_pool_resource());
169+
EXPECT_TRUE(test::current_device_uses_pool_resource());
215170
}
216171

217172
{
218173
auto scoped_device = device_setter{device1};
219-
EXPECT_TRUE(current_device_uses_default_cuda_resource());
174+
EXPECT_TRUE(test::current_device_uses_default_cuda_resource());
220175
}
221176
}
222177

cpp/tests/core/monitor_resources.cu

Lines changed: 10 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

6+
#include "test_memory_resource.hpp"
7+
68
#include <raft/core/device_mdarray.hpp>
79
#include <raft/core/device_setter.hpp>
810
#include <raft/core/memory_tracking_resources.hpp>
@@ -25,28 +27,6 @@
2527

2628
namespace {
2729

28-
struct device_resource_restore_guard {
29-
int device_id;
30-
raft::mr::device_resource resource;
31-
32-
~device_resource_restore_guard()
33-
{
34-
rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id}, std::move(resource));
35-
}
36-
};
37-
38-
auto current_device_uses_pool_resource() -> bool
39-
{
40-
auto current_mr = rmm::mr::get_current_device_resource_ref();
41-
return cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr) != nullptr;
42-
}
43-
44-
auto current_device_uses_default_cuda_resource() -> bool
45-
{
46-
auto current_mr = rmm::mr::get_current_device_resource_ref();
47-
return cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr) != nullptr;
48-
}
49-
5030
TEST(MemoryTrackingResources, TracksDeviceAllocations)
5131
{
5232
using namespace std::chrono_literals;
@@ -88,19 +68,8 @@ TEST(MemoryTrackingResources, RestoresDeviceResourceOnConstructionDevice)
8868
auto device0 = 0;
8969
auto device1 = 1;
9070

91-
auto device0_guard = [&]() {
92-
auto scoped_device = raft::device_setter{device0};
93-
auto upstream = rmm::mr::get_current_device_resource_ref();
94-
return device_resource_restore_guard{
95-
device0,
96-
rmm::mr::set_current_device_resource(
97-
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)})};
98-
}();
99-
100-
auto device1_guard = [&]() {
101-
auto scoped_device = raft::device_setter{device1};
102-
return device_resource_restore_guard{device1, rmm::mr::reset_current_device_resource()};
103-
}();
71+
auto device0_guard = raft::test::install_pool_device_resource(device0);
72+
auto device1_guard = raft::test::install_default_device_resource(device1);
10473

10574
{
10675
auto scoped_device = raft::device_setter{device0};
@@ -133,19 +102,8 @@ TEST(MemoryTrackingResources, InstallsTrackedResourceOnHandleDevice)
133102
auto device0 = 0;
134103
auto device1 = 1;
135104

136-
auto device0_guard = [&]() {
137-
auto scoped_device = raft::device_setter{device0};
138-
auto upstream = rmm::mr::get_current_device_resource_ref();
139-
return device_resource_restore_guard{
140-
device0,
141-
rmm::mr::set_current_device_resource(
142-
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)})};
143-
}();
144-
145-
auto device1_guard = [&]() {
146-
auto scoped_device = raft::device_setter{device1};
147-
return device_resource_restore_guard{device1, rmm::mr::reset_current_device_resource()};
148-
}();
105+
auto device0_guard = raft::test::install_pool_device_resource(device0);
106+
auto device1_guard = raft::test::install_default_device_resource(device1);
149107

150108
{
151109
auto scoped_device = raft::device_setter{device0};
@@ -159,25 +117,25 @@ TEST(MemoryTrackingResources, InstallsTrackedResourceOnHandleDevice)
159117

160118
{
161119
auto verify_device0 = raft::device_setter{device0};
162-
EXPECT_FALSE(current_device_uses_pool_resource());
120+
EXPECT_FALSE(raft::test::current_device_uses_pool_resource());
163121
}
164122

165123
{
166124
auto verify_device1 = raft::device_setter{device1};
167-
EXPECT_TRUE(current_device_uses_default_cuda_resource());
125+
EXPECT_TRUE(raft::test::current_device_uses_default_cuda_resource());
168126
}
169127

170128
tracked.reset();
171129
}
172130

173131
{
174132
auto scoped_device = raft::device_setter{device0};
175-
EXPECT_TRUE(current_device_uses_pool_resource());
133+
EXPECT_TRUE(raft::test::current_device_uses_pool_resource());
176134
}
177135

178136
{
179137
auto scoped_device = raft::device_setter{device1};
180-
EXPECT_TRUE(current_device_uses_default_cuda_resource());
138+
EXPECT_TRUE(raft::test::current_device_uses_default_cuda_resource());
181139
}
182140
}
183141

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,68 @@
1+
/*
2+
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
3+
* SPDX-License-Identifier: Apache-2.0
4+
*/
5+
#pragma once
6+
7+
#include <raft/core/device_setter.hpp>
8+
#include <raft/mr/host_device_resource.hpp>
9+
10+
#include <rmm/mr/cuda_memory_resource.hpp>
11+
#include <rmm/mr/per_device_resource.hpp>
12+
#include <rmm/mr/pool_memory_resource.hpp>
13+
14+
#include <cuda/memory_resource>
15+
16+
#include <gtest/gtest.h>
17+
18+
#include <exception>
19+
#include <utility>
20+
21+
namespace raft::test {
22+
23+
struct device_resource_restore_guard {
24+
int device_id;
25+
raft::mr::device_resource resource;
26+
27+
~device_resource_restore_guard()
28+
{
29+
try {
30+
rmm::mr::set_per_device_resource(rmm::cuda_device_id{device_id}, std::move(resource));
31+
} catch (const std::exception& e) {
32+
ADD_FAILURE() << "Failed to restore device " << device_id << " memory resource: " << e.what();
33+
} catch (...) {
34+
ADD_FAILURE() << "Failed to restore device " << device_id
35+
<< " memory resource: unknown exception";
36+
}
37+
}
38+
};
39+
40+
inline auto install_pool_device_resource(int device_id) -> device_resource_restore_guard
41+
{
42+
auto scoped_device = raft::device_setter{device_id};
43+
auto upstream = rmm::mr::get_current_device_resource_ref();
44+
auto installed_resource =
45+
raft::mr::device_resource{rmm::mr::pool_memory_resource(upstream, 1 << 20, 2 << 20)};
46+
auto old_resource = rmm::mr::set_current_device_resource(std::move(installed_resource));
47+
return device_resource_restore_guard{device_id, std::move(old_resource)};
48+
}
49+
50+
inline auto install_default_device_resource(int device_id) -> device_resource_restore_guard
51+
{
52+
auto scoped_device = raft::device_setter{device_id};
53+
return device_resource_restore_guard{device_id, rmm::mr::reset_current_device_resource()};
54+
}
55+
56+
inline auto current_device_uses_pool_resource() -> bool
57+
{
58+
auto current_mr = rmm::mr::get_current_device_resource_ref();
59+
return cuda::mr::resource_cast<rmm::mr::pool_memory_resource>(&current_mr) != nullptr;
60+
}
61+
62+
inline auto current_device_uses_default_cuda_resource() -> bool
63+
{
64+
auto current_mr = rmm::mr::get_current_device_resource_ref();
65+
return cuda::mr::resource_cast<rmm::mr::cuda_memory_resource>(&current_mr) != nullptr;
66+
}
67+
68+
} // namespace raft::test

0 commit comments

Comments
 (0)