Skip to content

Commit a497351

Browse files
authored
Merge pull request #2189 from fallintoplace/fix/nvexec-ensure-started-destruction
Destroy nvexec ensure_started completion storage before freeing it
2 parents c2d745a + 02729eb commit a497351

2 files changed

Lines changed: 252 additions & 8 deletions

File tree

include/nvexec/stream/ensure_started.cuh

Lines changed: 30 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -78,12 +78,17 @@ namespace nv::execution::_strm
7878

7979
if constexpr (stream_sender<Sender, env_t>)
8080
{
81-
cudaStream_t stream = shared_state_->stream_provider_.own_stream_.value();
82-
shared_state_->index_ = __mapply<__mfind_i<tuple_t>, variant_t>::value;
83-
copy_kernel<Tag, Args&&...>
84-
<<<1, 1, 0, stream>>>(shared_state_->data_, static_cast<Args&&>(args)...);
85-
shared_state_->stream_provider_.status_ = STDEXEC_LOG_CUDA_API(
86-
cudaEventRecord(shared_state_->event_, stream));
81+
if (shared_state_->stream_provider_.status_ == cudaSuccess)
82+
{
83+
cudaStream_t stream = shared_state_->stream_provider_.own_stream_.value();
84+
shared_state_->index_ = __mapply<__mfind_i<tuple_t>, variant_t>::value;
85+
copy_kernel<Tag, Args&&...>
86+
<<<1, 1, 0, stream>>>(shared_state_->data_, static_cast<Args&&>(args)...);
87+
shared_state_->completion_copy_started_ = true;
88+
shared_state_->stream_provider_.status_ = STDEXEC_LOG_CUDA_API(
89+
cudaEventRecord(shared_state_->event_, stream));
90+
shared_state_->event_recorded_ = shared_state_->stream_provider_.status_ == cudaSuccess;
91+
}
8792
}
8893
else
8994
{
@@ -174,6 +179,7 @@ namespace nv::execution::_strm
174179
{
175180
stream_provider_.status_ = STDEXEC_LOG_CUDA_API(
176181
cudaEventCreateWithFlags(&event_, cudaEventDisableTiming));
182+
event_created_ = stream_provider_.status_ == cudaSuccess;
177183
}
178184

179185
STDEXEC::start(opstate2_);
@@ -199,16 +205,29 @@ namespace nv::execution::_strm
199205

200206
~sh_state()
201207
{
202-
if (stream_provider_.status_ == cudaSuccess)
208+
if constexpr (stream_sender<Sender, env_t>)
203209
{
204-
if constexpr (stream_sender<Sender, env_t>)
210+
if (completion_copy_started_)
211+
{
212+
if (event_recorded_)
213+
{
214+
STDEXEC_ASSERT_CUDA_API(cudaEventSynchronize(event_));
215+
}
216+
else
217+
{
218+
STDEXEC_ASSERT_CUDA_API(cudaStreamSynchronize(stream_provider_.own_stream_.value()));
219+
}
220+
}
221+
222+
if (event_created_)
205223
{
206224
STDEXEC_ASSERT_CUDA_API(cudaEventDestroy(event_));
207225
}
208226
}
209227

210228
if (data_)
211229
{
230+
data_->~variant_t();
212231
STDEXEC_ASSERT_CUDA_API(cudaFree(data_));
213232
}
214233
}
@@ -232,6 +251,9 @@ namespace nv::execution::_strm
232251
context ctx_;
233252
stream_provider stream_provider_;
234253
cudaEvent_t event_{};
254+
bool event_created_{false};
255+
bool event_recorded_{false};
256+
bool completion_copy_started_{false};
235257
unsigned int index_{0};
236258
variant_t* data_{nullptr};
237259
task_t* task_{nullptr};

test/nvexec/ensure_started.cpp

Lines changed: 222 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,142 @@ using nvexec::is_on_gpu;
1313

1414
namespace
1515
{
16+
class mapped_flag_t
17+
{
18+
int* h_flag_{};
19+
int* d_flag_{};
20+
21+
public:
22+
mapped_flag_t(mapped_flag_t const &) = delete;
23+
mapped_flag_t(mapped_flag_t&&) = delete;
24+
auto operator=(mapped_flag_t const &) -> mapped_flag_t& = delete;
25+
auto operator=(mapped_flag_t&&) -> mapped_flag_t& = delete;
26+
27+
mapped_flag_t()
28+
{
29+
STDEXEC_TRY_CUDA_API(cudaHostAlloc(&h_flag_, sizeof(int), cudaHostAllocMapped));
30+
STDEXEC_TRY_CUDA_API(cudaHostGetDevicePointer(&d_flag_, h_flag_, 0));
31+
*h_flag_ = 0;
32+
}
33+
34+
~mapped_flag_t()
35+
{
36+
STDEXEC_ASSERT_CUDA_API(cudaFreeHost(h_flag_));
37+
}
38+
39+
class handle_t
40+
{
41+
int* h_flag_{};
42+
int* d_flag_{};
43+
44+
handle_t(int* h_flag, int* d_flag)
45+
: h_flag_(h_flag)
46+
, d_flag_(d_flag)
47+
{}
48+
49+
public:
50+
__host__ __device__ void set() const
51+
{
52+
cuda::atomic_ref<int, cuda::thread_scope_system> flag{*(is_on_gpu() ? d_flag_ : h_flag_)};
53+
flag.store(1, cuda::memory_order_release);
54+
}
55+
56+
friend mapped_flag_t;
57+
};
58+
59+
auto get() -> handle_t
60+
{
61+
return {h_flag_, d_flag_};
62+
}
63+
64+
auto is_set() const -> bool
65+
{
66+
cuda::atomic_ref<int, cuda::thread_scope_system> flag{*h_flag_};
67+
return flag.load(cuda::memory_order_acquire) == 1;
68+
}
69+
};
70+
71+
class lifetime_counter_t
72+
{
73+
int h_counter_storage_{};
74+
int* h_counter_{&h_counter_storage_};
75+
int* d_counter_{};
76+
77+
public:
78+
lifetime_counter_t(lifetime_counter_t const &) = delete;
79+
lifetime_counter_t(lifetime_counter_t&&) = delete;
80+
auto operator=(lifetime_counter_t const &) -> lifetime_counter_t& = delete;
81+
auto operator=(lifetime_counter_t&&) -> lifetime_counter_t& = delete;
82+
83+
lifetime_counter_t()
84+
{
85+
STDEXEC_TRY_CUDA_API(cudaMalloc(&d_counter_, sizeof(int)));
86+
STDEXEC_TRY_CUDA_API(cudaMemset(d_counter_, 0, sizeof(int)));
87+
}
88+
89+
~lifetime_counter_t()
90+
{
91+
STDEXEC_ASSERT_CUDA_API(cudaFree(d_counter_));
92+
}
93+
94+
class handle_t
95+
{
96+
int* h_counter_{};
97+
int* d_counter_{};
98+
99+
handle_t(int* h_counter, int* d_counter)
100+
: h_counter_(h_counter)
101+
, d_counter_(d_counter)
102+
{}
103+
104+
__host__ __device__ void update(int difference) const
105+
{
106+
cuda::std::atomic_ref<int> counter{*(is_on_gpu() ? d_counter_ : h_counter_)};
107+
counter.fetch_add(difference, cuda::std::memory_order_relaxed);
108+
}
109+
110+
friend lifetime_counter_t;
111+
friend class lifetime_tracer_t;
112+
};
113+
114+
auto get() -> handle_t
115+
{
116+
return {h_counter_, d_counter_};
117+
}
118+
119+
auto alive() const -> int
120+
{
121+
int d_counter{};
122+
STDEXEC_TRY_CUDA_API(cudaMemcpy(&d_counter, d_counter_, sizeof(int), cudaMemcpyDeviceToHost));
123+
return *h_counter_ + d_counter;
124+
}
125+
};
126+
127+
class lifetime_tracer_t
128+
{
129+
lifetime_counter_t::handle_t counter_;
130+
131+
public:
132+
lifetime_tracer_t() = delete;
133+
lifetime_tracer_t(lifetime_tracer_t const &) = delete;
134+
135+
__host__ __device__ explicit lifetime_tracer_t(lifetime_counter_t::handle_t counter)
136+
: counter_(counter)
137+
{
138+
counter_.update(1);
139+
}
140+
141+
__host__ __device__ lifetime_tracer_t(lifetime_tracer_t&& other)
142+
: counter_(other.counter_)
143+
{
144+
counter_.update(1);
145+
}
146+
147+
__host__ __device__ ~lifetime_tracer_t()
148+
{
149+
counter_.update(-1);
150+
}
151+
};
16152

17153
TEST_CASE("nvexec ensure_started is eager", "[cuda][stream][adaptors][ensure_started]")
18154
{
@@ -52,6 +188,92 @@ namespace
52188
REQUIRE(v == 1);
53189
}
54190

191+
TEST_CASE("nvexec ensure_started destroys completion storage",
192+
"[cuda][stream][adaptors][ensure_started]")
193+
{
194+
nvexec::stream_context stream_ctx{};
195+
lifetime_counter_t counter{};
196+
auto handle = counter.get();
197+
198+
{
199+
auto snd = exec::ensure_started(
200+
ex::schedule(stream_ctx.get_scheduler())
201+
| a_sender([handle]() -> lifetime_tracer_t { return lifetime_tracer_t{handle}; }));
202+
auto result = STDEXEC::sync_wait(std::move(snd));
203+
REQUIRE(result.has_value());
204+
}
205+
206+
REQUIRE(counter.alive() == 0);
207+
}
208+
209+
TEST_CASE("nvexec ensure_started destroys completion storage when detached",
210+
"[cuda][stream][adaptors][ensure_started]")
211+
{
212+
lifetime_counter_t counter{};
213+
214+
{
215+
nvexec::stream_context stream_ctx{};
216+
auto handle = counter.get();
217+
auto make_tracer = [handle]() -> lifetime_tracer_t
218+
{
219+
NV_IF_TARGET(NV_IS_DEVICE,
220+
(auto const start = clock64(); while (clock64() - start < 10'000'000){}));
221+
return lifetime_tracer_t{handle};
222+
};
223+
auto predecessor = a_sender(ex::schedule(stream_ctx.get_scheduler()), make_tracer);
224+
225+
{
226+
auto snd = exec::ensure_started(std::move(predecessor));
227+
}
228+
229+
STDEXEC_TRY_CUDA_API(cudaDeviceSynchronize());
230+
}
231+
232+
REQUIRE(counter.alive() == 0);
233+
}
234+
235+
TEST_CASE("nvexec ensure_started synchronizes direct stream completion when detached",
236+
"[cuda][stream][adaptors][ensure_started]")
237+
{
238+
int device{};
239+
STDEXEC_TRY_CUDA_API(cudaGetDevice(&device));
240+
241+
int can_map_host_memory{};
242+
STDEXEC_TRY_CUDA_API(
243+
cudaDeviceGetAttribute(&can_map_host_memory, cudaDevAttrCanMapHostMemory, device));
244+
if (!can_map_host_memory)
245+
{
246+
SKIP("device does not support mapped host memory");
247+
}
248+
249+
mapped_flag_t completion{};
250+
251+
{
252+
nvexec::stream_context stream_ctx{};
253+
auto flag = completion.get();
254+
auto predecessor = ex::schedule(stream_ctx.get_scheduler())
255+
| ex::then(
256+
[flag]() -> int
257+
{
258+
NV_IF_TARGET(NV_IS_DEVICE,
259+
(auto const start = clock64();
260+
while (clock64() - start < 10'000'000) {} flag.set();));
261+
return 42;
262+
});
263+
264+
static_assert(
265+
nvexec::_strm::stream_sender<decltype(predecessor), nvexec::_strm::_ensure_started::env_t>);
266+
267+
{
268+
auto snd = exec::ensure_started(std::move(predecessor));
269+
}
270+
271+
bool const completed_before_cleanup = completion.is_set();
272+
STDEXEC_TRY_CUDA_API(cudaDeviceSynchronize());
273+
REQUIRE(completed_before_cleanup);
274+
}
275+
}
276+
55277
TEST_CASE("ensure_started can preceed a sender without values",
56278
"[cuda][stream][adaptors][ensure_started]")
57279
{

0 commit comments

Comments
 (0)