Skip to content

Commit 8c692ca

Browse files
committed
UCP/RKEY: Acquire context lock when calling ucp_rkey_pack_memh
1 parent c4e66dc commit 8c692ca

5 files changed

Lines changed: 82 additions & 25 deletions

File tree

src/ucp/core/ucp_rkey.c

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,7 @@ ucp_rkey_unpack_distance(const ucp_rkey_packed_distance_t *packed_distance,
125125
distance->bandwidth = UCS_FP8_UNPACK(BANDWIDTH, packed_distance->bandwidth);
126126
}
127127

128+
/* context->mt_lock must be held */
128129
UCS_PROFILE_FUNC(ssize_t, ucp_rkey_pack_memh,
129130
(context, md_map, memh, address, length, mem_info, sys_dev_map,
130131
sys_distance, uct_flags, buffer),

src/ucp/proto/proto_common.inl

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -165,12 +165,12 @@ ucp_proto_request_set_stage(ucp_request_t *req, uint8_t proto_stage)
165165
{
166166
const ucp_proto_t *proto = req->send.proto_config->proto;
167167

168-
ucs_assertv(proto_stage < UCP_PROTO_STAGE_LAST, "stage=%"PRIu8,
168+
ucs_assertv(proto_stage < UCP_PROTO_STAGE_LAST, "stage=%" PRIu8,
169169
proto_stage);
170170
ucs_assert(proto->progress[proto_stage] != NULL);
171171

172172
ucp_trace_req(req, "set to stage %u, progress function '%s'", proto_stage,
173-
ucs_debug_get_symbol_name(proto->progress[proto_stage]));
173+
ucs_debug_get_symbol_name((void *)proto->progress[proto_stage]));
174174
req->send.proto_stage = proto_stage;
175175

176176
/* Set pointer to progress function */
@@ -186,7 +186,7 @@ static void ucp_proto_request_set_proto(ucp_request_t *req,
186186
const ucp_proto_config_t *proto_config,
187187
size_t msg_length)
188188
{
189-
ucs_assertv(req->flags & UCP_REQUEST_FLAG_PROTO_SEND, "flags=0x%"PRIx32,
189+
ucs_assertv(req->flags & UCP_REQUEST_FLAG_PROTO_SEND, "flags=0x%" PRIx32,
190190
req->flags);
191191

192192
req->send.proto_config = proto_config;
@@ -346,6 +346,7 @@ ucp_proto_request_pack_rkey(ucp_request_t *req, ucp_md_map_t md_map,
346346
const ucs_sys_dev_distance_t *dev_distance,
347347
void *rkey_buffer)
348348
{
349+
ucp_context_h context = req->send.ep->worker->context;
349350
const ucp_datatype_iter_t *dt_iter = &req->send.state.dt_iter;
350351
ucp_mem_h memh;
351352
ssize_t packed_rkey_size;
@@ -366,17 +367,19 @@ ucp_proto_request_pack_rkey(ucp_request_t *req, ucp_md_map_t md_map,
366367
ucs_unlikely(memh->flags & UCP_MEMH_FLAG_HAS_AUTO_GVA)) {
367368
ucp_memh_disable_gva(memh, md_map);
368369
}
369-
370370
if (!ucs_test_all_flags(memh->md_map, md_map)) {
371-
ucs_trace("dt_iter_md_map=0x%"PRIx64" md_map=0x%"PRIx64, memh->md_map,
372-
md_map);
371+
ucs_trace("dt_iter_md_map=0x%" PRIx64 " md_map=0x%" PRIx64,
372+
memh->md_map, md_map);
373373
}
374374

375+
/* TODO: context lock is not scalable. Consider fine-grained lock per memh,
376+
* immutable memh with rkey cache, RCU/COW */
377+
UCP_THREAD_CS_ENTER(&context->mt_lock);
375378
packed_rkey_size = ucp_rkey_pack_memh(
376-
req->send.ep->worker->context, md_map & memh->md_map, memh,
377-
dt_iter->type.contig.buffer, dt_iter->length, &dt_iter->mem_info,
378-
distance_dev_map, dev_distance,
379+
context, md_map & memh->md_map, memh, dt_iter->type.contig.buffer,
380+
dt_iter->length, &dt_iter->mem_info, distance_dev_map, dev_distance,
379381
ucp_ep_config(req->send.ep)->uct_rkey_pack_flags, rkey_buffer);
382+
UCP_THREAD_CS_EXIT(&context->mt_lock);
380383

381384
if (packed_rkey_size < 0) {
382385
ucs_error("failed to pack remote key: %s",

src/ucp/rndv/rndv.c

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -178,11 +178,15 @@ size_t ucp_rndv_rts_pack(ucp_request_t *sreq, ucp_rndv_rts_hdr_t *rndv_rts_hdr,
178178
rndv_rts_hdr->address = (uintptr_t)sreq->send.buffer;
179179
rkey_buf = UCS_PTR_BYTE_OFFSET(rndv_rts_hdr,
180180
sizeof(*rndv_rts_hdr));
181-
packed_rkey_size = ucp_rkey_pack_memh(
181+
182+
UCP_THREAD_CS_ENTER(&worker->context->mt_lock);
183+
packed_rkey_size = ucp_rkey_pack_memh(
182184
worker->context, sreq->send.rndv.md_map,
183185
sreq->send.state.dt.dt.contig.memh, sreq->send.buffer,
184186
sreq->send.length, &mem_info, 0, NULL,
185187
ucp_ep_config(sreq->send.ep)->uct_rkey_pack_flags, rkey_buf);
188+
UCP_THREAD_CS_EXIT(&worker->context->mt_lock);
189+
186190
if (packed_rkey_size < 0) {
187191
ucs_fatal("failed to pack rendezvous remote key: %s",
188192
ucs_status_string((ucs_status_t)packed_rkey_size));
@@ -205,6 +209,7 @@ static size_t ucp_rndv_rtr_pack(void *dest, void *arg)
205209
ucp_rndv_rtr_hdr_t *rndv_rtr_hdr = dest;
206210
ucp_request_t *rreq = ucp_request_get_super(rndv_req);
207211
ucp_ep_h ep = rndv_req->send.ep;
212+
ucp_context_h context = ep->worker->context;
208213
ucp_memory_info_t mem_info;
209214
ssize_t packed_rkey_size;
210215

@@ -221,12 +226,15 @@ static size_t ucp_rndv_rtr_pack(void *dest, void *arg)
221226
mem_info.type = rreq->recv.dt_iter.mem_info.type;
222227
mem_info.sys_dev = UCS_SYS_DEVICE_ID_UNKNOWN;
223228

229+
UCP_THREAD_CS_ENTER(&context->mt_lock);
224230
packed_rkey_size = ucp_rkey_pack_memh(
225-
ep->worker->context, rndv_req->send.rndv.md_map,
231+
context, rndv_req->send.rndv.md_map,
226232
rreq->recv.dt_iter.type.contig.memh,
227233
rreq->recv.dt_iter.type.contig.buffer, rndv_req->send.length,
228234
&mem_info, 0, NULL, ucp_ep_config(ep)->uct_rkey_pack_flags,
229235
rndv_rtr_hdr + 1);
236+
UCP_THREAD_CS_EXIT(&context->mt_lock);
237+
230238
if (packed_rkey_size < 0) {
231239
return packed_rkey_size;
232240
}

src/ucp/rndv/rndv_rtr.c

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -252,6 +252,7 @@ static size_t ucp_proto_rndv_rtr_mtype_pack(void *dest, void *arg)
252252
{
253253
ucp_rndv_rtr_hdr_t *rtr = dest;
254254
ucp_request_t *req = arg;
255+
ucp_context_h context = req->send.ep->worker->context;
255256
const ucp_proto_rndv_rtr_priv_t *rpriv = req->send.proto_config->priv;
256257
ucp_md_map_t md_map = rpriv->super.md_map;
257258
ucp_mem_desc_t *mdesc = req->send.rndv.mdesc;
@@ -266,10 +267,14 @@ static size_t ucp_proto_rndv_rtr_mtype_pack(void *dest, void *arg)
266267
/* Pack remote key for the fragment */
267268
mem_info.type = mdesc->memh->mem_type;
268269
mem_info.sys_dev = UCS_SYS_DEVICE_ID_UNKNOWN;
269-
packed_rkey_size = ucp_rkey_pack_memh(req->send.ep->worker->context, md_map,
270-
mdesc->memh, mdesc->ptr,
270+
271+
UCP_THREAD_CS_ENTER(&context->mt_lock);
272+
packed_rkey_size = ucp_rkey_pack_memh(context, md_map, mdesc->memh,
273+
mdesc->ptr,
271274
req->send.state.dt_iter.length,
272275
&mem_info, 0, NULL, 0, rtr + 1);
276+
UCP_THREAD_CS_EXIT(&context->mt_lock);
277+
273278
if (packed_rkey_size < 0) {
274279
ucs_error("failed to pack remote key: %s",
275280
ucs_status_string((ucs_status_t)packed_rkey_size));

test/gtest/ucp/test_ucp_rma_mt.cc

Lines changed: 52 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,10 @@
99

1010
#include <common/test_helpers.h>
1111

12+
extern "C" {
13+
#include <ucp/proto/proto_common.inl>
14+
}
15+
1216
#if _OPENMP
1317
#include "omp.h"
1418
#endif
@@ -35,6 +39,22 @@ class test_ucp_rma_mt : public ucp_test {
3539
add_variant(variants, UCP_FEATURE_RMA, MULTI_THREAD_CONTEXT);
3640
add_variant(variants, UCP_FEATURE_RMA, MULTI_THREAD_WORKER);
3741
}
42+
43+
ucp_mem_h mem_map(ucp_context_h context, void *data, size_t size)
44+
{
45+
ucp_mem_map_params_t params;
46+
ucp_mem_h memh;
47+
48+
params.field_mask = UCP_MEM_MAP_PARAM_FIELD_ADDRESS |
49+
UCP_MEM_MAP_PARAM_FIELD_LENGTH |
50+
UCP_MEM_MAP_PARAM_FIELD_FLAGS;
51+
params.address = data;
52+
params.length = size;
53+
params.flags = get_variant_value();
54+
55+
ASSERT_UCS_OK(ucp_mem_map(context, &params, &memh));
56+
return memh;
57+
}
3858
};
3959

4060
UCS_TEST_P(test_ucp_rma_mt, put_get) {
@@ -43,19 +63,9 @@ UCS_TEST_P(test_ucp_rma_mt, put_get) {
4363
uint64_t orig_data[num_threads] GTEST_ATTRIBUTE_UNUSED_;
4464
uint64_t target_data[num_threads] GTEST_ATTRIBUTE_UNUSED_;
4565

46-
ucp_mem_map_params_t params;
47-
ucp_mem_h memh;
4866
void *memheap = target_data;
49-
50-
params.field_mask = UCP_MEM_MAP_PARAM_FIELD_ADDRESS |
51-
UCP_MEM_MAP_PARAM_FIELD_LENGTH |
52-
UCP_MEM_MAP_PARAM_FIELD_FLAGS;
53-
params.address = memheap;
54-
params.length = sizeof(uint64_t) * num_threads;
55-
params.flags = get_variant_value();
56-
57-
st = ucp_mem_map(receiver().ucph(), &params, &memh);
58-
ASSERT_UCS_OK(st);
67+
ucp_mem_h memh = mem_map(receiver().ucph(), memheap,
68+
sizeof(uint64_t) * num_threads);
5969

6070
void *rkey_buffer;
6171
size_t rkey_buffer_size;
@@ -200,4 +210,34 @@ UCS_TEST_P(test_ucp_rma_mt, put_get) {
200210
ASSERT_UCS_OK(st);
201211
}
202212

213+
UCS_TEST_P(test_ucp_rma_mt, rkey_pack) {
214+
uint8_t data[1024] GTEST_ATTRIBUTE_UNUSED_;
215+
ucp_context_h context = sender().ucph();
216+
ucp_mem_h memh = mem_map(context, data, sizeof(data));
217+
218+
#if _OPENMP && ENABLE_MT
219+
#pragma omp parallel for
220+
for (int i = 0; i < mt_num_threads(); i++) {
221+
if (i % 2 == 0) {
222+
void *rkey;
223+
size_t rkey_size;
224+
ASSERT_UCS_OK(ucp_rkey_pack(context, memh, &rkey, &rkey_size));
225+
ucp_rkey_buffer_release(rkey);
226+
} else {
227+
ucs_sys_dev_distance_t sys_dev = {};
228+
ucp_request req = {};
229+
req.send.ep = sender().ep();
230+
req.send.state.dt_iter.type.contig.memh = memh;
231+
req.send.state.dt_iter.type.contig.buffer = data;
232+
req.send.state.dt_iter.length = sizeof(data);
233+
234+
uint8_t rkey[1024];
235+
ucp_proto_request_pack_rkey(&req, memh->md_map, 0, &sys_dev, rkey);
236+
}
237+
}
238+
#endif
239+
240+
ASSERT_UCS_OK(ucp_mem_unmap(context, memh));
241+
}
242+
203243
UCP_INSTANTIATE_TEST_CASE(test_ucp_rma_mt)

0 commit comments

Comments
 (0)