diff --git a/ompi/mca/osc/sm/osc_sm_comm.c b/ompi/mca/osc/sm/osc_sm_comm.c index f20c00c2517..f553aa528d2 100644 --- a/ompi/mca/osc/sm/osc_sm_comm.c +++ b/ompi/mca/osc/sm/osc_sm_comm.c @@ -223,7 +223,7 @@ ompi_osc_sm_win_set_num_notify(struct ompi_win_t *win, int rank = ompi_comm_rank(module->comm); unsigned long requested = (unsigned long) num_notifications; unsigned long *new_caps; - bool grow = false; + bool grow = false, bad; int ret, i; /* "mpi_assert_same_num_notifications" would let us skip the allgather below @@ -232,7 +232,18 @@ ompi_osc_sm_win_set_num_notify(struct ompi_win_t *win, * synchronizing and collective. */ (void) info; - if (num_notifications < 0) { + /* num_notifications is a local argument -- MPI-5.1 12.6.1 allows it to + * differ between MPI processes -- but this is a synchronizing collective. + * A rank that rejected its own value and returned here would leave every + * other rank blocked in the allgather below, turning an erroneous argument + * into a hang. So the validity rides through the collective as a sentinel + * and all ranks fail together. A multi-process window defers the decision; + * a single-process one has nobody to agree with and can answer now. */ + bad = (num_notifications < 0) + || (0 != module->notify_max_assert && + requested > (unsigned long) module->notify_max_assert); + + if (bad && 1 == comm_size) { return MPI_ERR_ARG; } @@ -267,6 +278,7 @@ ompi_osc_sm_win_set_num_notify(struct ompi_win_t *win, return OMPI_SUCCESS; } +agree: new_caps = malloc(sizeof(*new_caps) * comm_size); if (NULL == new_caps) { return OMPI_ERR_TEMP_OUT_OF_RESOURCE; @@ -281,6 +293,16 @@ ompi_osc_sm_win_set_num_notify(struct ompi_win_t *win, return ret; } + for (i = 0 ; i < comm_size ; ++i) { + if (ULONG_MAX == new_caps[i]) { + /* Some rank supplied an invalid count. Every rank sees the same + * gathered array, so they all report the same error and none of + * them reconfigures. */ + free(new_caps); + return MPI_ERR_ARG; + } + } + for (i = 0 ; i < comm_size ; ++i) { if (new_caps[i] > module->node_states[i].notify_counter_capacity) { grow = true; diff --git a/ompi/mca/osc/ucx/osc_ucx.h b/ompi/mca/osc/ucx/osc_ucx.h index bc3dc8a91b3..f79db5d1070 100644 --- a/ompi/mca/osc/ucx/osc_ucx.h +++ b/ompi/mca/osc/ucx/osc_ucx.h @@ -27,6 +27,11 @@ #define OMPI_OSC_UCX_POST_PEER_MAX 32 #define OMPI_OSC_UCX_ATTACH_MAX 48 #define OMPI_OSC_UCX_MEM_ADDR_MAX_LEN 1024 +/* Default number of RMA notification counters reserved per MPI process in each + * window's registered memory region. Overridden per job by the + * osc_ucx_num_notify_counters MCA parameter and per window by the + * "mpi_assert_max_num_notify" info key. */ +#define OMPI_OSC_UCX_DEFAULT_NOTIFY_COUNTERS 16 typedef struct ompi_osc_ucx_component { @@ -43,6 +48,9 @@ typedef struct ompi_osc_ucx_component { bool no_locks; /* Default value of the no_locks info key for new windows */ bool acc_single_intrinsic; unsigned int priority; + /* Number of notification counters reserved per MPI process in each window, + * unless the window's info gives "mpi_assert_max_num_notify". */ + unsigned int num_notify_counters; /* directory where to place backing files */ char *backing_directory; } ompi_osc_ucx_component_t; @@ -122,6 +130,24 @@ typedef struct ompi_osc_ucx_module { struct ompi_communicator_t *comm; int flavor; size_t size; + int *notify_counts; /* per-rank number of notification counters *attached* at each + * rank (size comm_size), as set by MPI_WIN_SET_NUM_NOTIFY and + * kept consistent across the group by an allgather. Always + * <= notify_capacity. */ + unsigned int notify_capacity; /* notification counters currently reserved per rank. + * Agreed on across the group and uniform. Grown on + * demand by MPI_WIN_SET_NUM_NOTIFY unless + * notify_max_assert caps it. */ + unsigned int notify_max_assert; /* non-zero only if *every* rank passed + * "mpi_assert_max_num_notify" at window creation. + * Then the agreed reservation is a hard cap and the + * counters never grow (MPI-5.1 12.2: the assertion + * lets the implementation optimize the allocation). + * Zero means no rank asserted a bound, so the + * standard's "does not assume any limit" applies. */ + uint64_t *notify_addrs; /* per-rank base address of the notification counters + * (size comm_size) */ + void *notify_base; /* this rank's counters; notify_capacity uint64_t */ size_t *sizes; /* used if not every process has the same size */ uint64_t *addrs; uint64_t *state_addrs; @@ -149,6 +175,11 @@ typedef struct ompi_osc_ucx_module { opal_common_ucx_ctx_t *ctx; opal_common_ucx_wpmem_t *mem; opal_common_ucx_wpmem_t *state_mem; + /* Notification counters get their own registration rather than being + * appended to the window data: the data region for MPI_WIN_FLAVOR_CREATE + * belongs to the user and has no room for them, and a dynamic window has + * no data region at all. */ + opal_common_ucx_wpmem_t *notify_mem; ompi_osc_ucx_mem_ranges_t *epoc_outstanding_ops_mems; bool skip_sync_check; bool noncontig_shared_win; @@ -277,6 +308,75 @@ int ompi_osc_find_attached_region_position(ompi_osc_dynamic_win_info_t *dynamic_ int ompi_osc_ucx_dynamic_lock(ompi_osc_ucx_module_t *module, int target); int ompi_osc_ucx_dynamic_unlock(ompi_osc_ucx_module_t *module, int target); +int ompi_osc_ucx_put_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win); +int ompi_osc_ucx_get_notify(void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win); +int ompi_osc_ucx_rput_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win, + struct ompi_request_t **request); +int ompi_osc_ucx_rget_notify(void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win, + struct ompi_request_t **request); +int ompi_osc_ucx_accumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win); +int ompi_osc_ucx_get_accumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + void *result_addr, size_t result_count, + struct ompi_datatype_t *result_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win); +int ompi_osc_ucx_raccumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win, + struct ompi_request_t **request); +int ompi_osc_ucx_rget_accumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + void *result_addr, size_t result_count, + struct ompi_datatype_t *result_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win, + struct ompi_request_t **request); +int ompi_osc_ucx_win_get_notify_value(struct ompi_win_t *win, int notify, + OMPI_MPI_COUNT_TYPE *value); +int ompi_osc_ucx_win_get_notify_bounds(struct ompi_win_t *win, int *num_sb, int *num_ub, + OMPI_MPI_COUNT_TYPE *value_ub); +int ompi_osc_ucx_win_reset_notify_value(struct ompi_win_t *win, int notify, + OMPI_MPI_COUNT_TYPE *value); +int ompi_osc_ucx_win_set_num_notify(struct ompi_win_t *win, struct opal_info_t *info, + int num_notifications); +/* Collectively re-reserve new_capacity notification counters per rank, replacing + * the current registration. Defined in osc_ucx_component.c because it needs the + * component's address-exchange helper. */ +int ompi_osc_ucx_grow_notify_counters(ompi_osc_ucx_module_t *module, + unsigned int new_capacity); + +int ompi_osc_ucx_win_get_num_notify(struct ompi_win_t *win, int target_rank, + int *num_notifications); + /* returns the size at the peer */ static inline size_t ompi_osc_ucx_get_size(ompi_osc_ucx_module_t *module, int rank) { diff --git a/ompi/mca/osc/ucx/osc_ucx_comm.c b/ompi/mca/osc/ucx/osc_ucx_comm.c index 0354edb71c0..1f400f85fb9 100644 --- a/ompi/mca/osc/ucx/osc_ucx_comm.c +++ b/ompi/mca/osc/ucx/osc_ucx_comm.c @@ -17,9 +17,15 @@ #include "ompi/mca/osc/base/osc_base_obj_convert.h" #include "opal/mca/common/ucx/common_ucx.h" +#include + #include "osc_ucx.h" #include "osc_ucx_request.h" +#include + +#include "ompi/attribute/attribute.h" + #define CHECK_VALID_RKEY(_module, _target, _count) \ if (!((_module)->win_info_array[_target]).rkey_init && ((_count) > 0)) { \ @@ -603,6 +609,589 @@ int ompi_osc_ucx_get(void *origin_addr, size_t origin_count, } } +static int osc_ucx_request_over_flush(ompi_osc_ucx_module_t *module, + struct ompi_win_t *win, int target, + ucp_ep_h *ep, enum req_type req_type, + struct ompi_request_t **request); + +static int accumulate_req(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, struct ompi_win_t *win, + ompi_osc_ucx_accumulate_request_t *ucx_req); + +static int get_accumulate_req(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + void *result_addr, size_t result_count, + struct ompi_datatype_t *result_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, struct ompi_win_t *win, + ompi_osc_ucx_accumulate_request_t *ucx_req); + +/* Returns the remote address of notify counter[notify] for the given target. + * The counters have their own registered region (module->notify_mem), separate + * from the window data, so this is independent of the window's flavor and size. */ +static inline uint64_t +osc_ucx_notify_counter_addr(ompi_osc_ucx_module_t *module, int target, int notify) +{ + return module->notify_addrs[target] + (uint64_t)notify * sizeof(uint64_t); +} + +/* A region of module->notify_capacity notification counters is registered per + * rank at window creation (see osc_ucx_component.c), but only the first + * notify_counts[rank] of them are considered *attached* by + * MPI_WIN_SET_NUM_NOTIFY. Per the MPI Standard it is erroneous to reference a + * counter that is out of range at the target, so validate against the target + * rank's attached count. */ +#define CHECK_NOTIFY_IDX(module, notify, rank) \ + if ((notify) < 0 || (notify) >= (module)->notify_counts[rank]) { \ + return MPI_ERR_RMA_NOTIFICATION; \ + } + +/* Increments the target's notification counter once the preceding data + * operation has been ordered ahead of it. Shared by every notified + * operation; they differ only in which base operation they issue first and + * in whether a fence or a flush is needed to order it. + * + * Note that for the request-based variants the data operation has already been + * issued and *request already handed back by the time this can fail, so an + * error return leaves that request outstanding and the caller still has to + * complete it. A transport failure here is not recoverable in any case. */ +static inline int +osc_ucx_notify_target(ompi_osc_ucx_module_t *module, int target, int notify, + ucp_ep_h *ep) +{ + int ret = opal_common_ucx_wpmem_post(module->notify_mem, + UCP_ATOMIC_POST_OP_ADD, 1, + target, sizeof(uint64_t), + osc_ucx_notify_counter_addr(module, target, notify), + ep); + return (OPAL_SUCCESS == ret) ? OMPI_SUCCESS : OMPI_ERROR; +} + +int ompi_osc_ucx_win_get_notify_value(struct ompi_win_t *win, int notify, + OMPI_MPI_COUNT_TYPE *value) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + int my_rank = ompi_comm_rank(module->comm); + + CHECK_NOTIFY_IDX(module, notify, my_rank); + + /* Origins increment this counter with a UCX atomic, which the transport may + * emulate in software on the local worker rather than offload to the NIC. + * In that case the counter only advances while the worker is progressed, so + * a consumer spinning on MPI_WIN_GET_NOTIFY_VALUE -- the natural way to wait + * for a notification -- would never observe the update. Progress the worker + * here so that such a loop makes forward progress on its own, as every other + * spin-wait in this component does. */ + opal_common_ucx_wpool_progress(mca_osc_ucx_component.wpool); + + volatile uint64_t *counter = + (volatile uint64_t *)osc_ucx_notify_counter_addr(module, my_rank, notify); + *value = (OMPI_MPI_COUNT_TYPE)*counter; + opal_atomic_rmb(); + return OMPI_SUCCESS; +} + +int ompi_osc_ucx_win_reset_notify_value(struct ompi_win_t *win, int notify, + OMPI_MPI_COUNT_TYPE *value) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + int my_rank = ompi_comm_rank(module->comm); + uint64_t result_value = 0; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, my_rank); + + OSC_UCX_GET_DEFAULT_EP(ep, module, my_rank); + + /* The counter is incremented by remote origins through UCX network atomic + * operations, so reset it with a UCX atomic swap (targeting our own rank) + * rather than a CPU atomic. That keeps the read-and-zero atomic with + * respect to those concurrent network atomics — a plain CPU swap is not + * ordered against them. The fetch returns the counter's previous value. */ + ret = opal_common_ucx_wpmem_fetch(module->notify_mem, + UCP_ATOMIC_FETCH_OP_SWAP, 0, + my_rank, &result_value, sizeof(result_value), + osc_ucx_notify_counter_addr(module, my_rank, notify), + ep); + if (OPAL_SUCCESS != ret) { + OSC_UCX_VERBOSE(1, "opal_common_ucx_wpmem_fetch failed: %d", ret); + return OMPI_ERROR; + } + + *value = (OMPI_MPI_COUNT_TYPE)result_value; + return OMPI_SUCCESS; +} + +int ompi_osc_ucx_win_get_num_notify(struct ompi_win_t *win, int target_rank, + int *num_notifications) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + + if (target_rank < 0 || target_rank >= ompi_comm_size(module->comm)) { + return MPI_ERR_RANK; + } + + /* Local query (MPI_WIN_GET_NUM_NOTIFY, §12.6.1): return the number of + * notification counters currently attached at target_rank, as last + * published by MPI_WIN_SET_NUM_NOTIFY. */ + *num_notifications = module->notify_counts[target_rank]; + return OMPI_SUCCESS; +} + +int ompi_osc_ucx_win_set_num_notify(struct ompi_win_t *win, struct opal_info_t *info, + int num_notifications) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + int comm_size = ompi_comm_size(module->comm); + int requested = num_notifications; + int *requested_counts; + unsigned int needed; + int ret, i; + + (void) info; /* "mpi_assert_same_num_notifications" is an optimization hint only */ + + /* When every rank asserted "mpi_assert_max_num_notify" at window creation, + * that value is a hard upper bound and asking for more is erroneous. + * Otherwise MPI-5.1 12.2 says no limit is assumed, so a request above the + * current reservation grows the counters below rather than failing. + * + * This is a synchronizing collective, so a rank with a bad argument must not + * return before the allgather below -- that would leave the rest of the + * group blocked in it. Mark the request instead and let every rank discover + * the error from the gathered values. */ + if (requested < 0 || + (0 != module->notify_max_assert && + (unsigned int) requested > module->notify_max_assert)) { + requested = -1; + } + + requested_counts = malloc(comm_size * sizeof(int)); + if (NULL == requested_counts) { + return OMPI_ERR_TEMP_OUT_OF_RESOURCE; + } + + /* All notification counters (existing and newly attached) are reset to zero + * by this call. Resetting before the allgather is what makes the standard's + * "will not return until ... all processes have adjusted the number of + * notification counters" hold: completing the collective implies every rank + * has already reset, so no rank can return and then have a peer wipe the + * notification it just delivered. It is erroneous to call this while an + * access epoch is open, so no concurrent network atomics touch the counters + * and a plain local reset is sufficient. */ + if (NULL != module->notify_base) { + memset(module->notify_base, 0, module->notify_capacity * sizeof(uint64_t)); + } + opal_atomic_wmb(); + + /* Publish every rank's requested count to the whole group so that origins + * can validate notification indices against the target's count. Gathering + * the requested value directly is what makes MPI_WIN_GET_NUM_NOTIFY return + * the value given here, including when it lowers the count. */ + ret = module->comm->c_coll->coll_allgather(&requested, 1, MPI_INT, + requested_counts, 1, MPI_INT, + module->comm, + module->comm->c_coll->coll_allgather_module); + if (OMPI_SUCCESS != ret) { + free(requested_counts); + return ret; + } + + for (i = 0; i < comm_size; i++) { + if (0 > requested_counts[i]) { + /* Some rank asked for a count outside [0, notify_capacity]. Every + * rank sees the same gathered array and bails identically, so the + * attached counts stay as they were rather than the group ending up + * half-reconfigured. The counters have been zeroed, which is + * harmless for a call that is erroneous anyway. */ + free(requested_counts); + return MPI_ERR_ARG; + } + } + + /* Every rank sees the same gathered array, so they all reach the same + * decision about whether to grow and to what size, without extra + * communication. The reservation is uniform across the window (window + * creation agrees it with an allreduce), so it grows to the largest request + * anyone made. Never shrink: a rank that lowered its count keeps the space + * it already has, so only genuine growth costs a re-registration and + * alternating high/low requests do not thrash the NIC. */ + needed = module->notify_capacity; + for (i = 0; i < comm_size; i++) { + if ((unsigned int) requested_counts[i] > needed) { + needed = (unsigned int) requested_counts[i]; + } + } + + if (needed > module->notify_capacity) { + ret = ompi_osc_ucx_grow_notify_counters(module, needed); + if (OMPI_SUCCESS != ret) { + free(requested_counts); + return ret; + } + + /* MPI_WIN_NOTIFICATION_NUM_SB is the count supported without paying for + * a re-registration, so it has to follow the reservation rather than + * stay at whatever was cached when the window was created. */ + ret = ompi_attr_set_int(WIN_ATTR, win, &win->w_keyhash, + MPI_WIN_NOTIFICATION_NUM_SB, + (int) module->notify_capacity, true); + if (OMPI_SUCCESS != ret) { + free(requested_counts); + return ret; + } + } + + memcpy(module->notify_counts, requested_counts, comm_size * sizeof(int)); + free(requested_counts); + + return OMPI_SUCCESS; +} + +int ompi_osc_ucx_win_get_notify_bounds(struct ompi_win_t *win, int *num_sb, int *num_ub, + OMPI_MPI_COUNT_TYPE *value_ub) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + + /* The current reservation is what is supported without paying for a + * re-registration, so it is the suggested bound. The hard bound is only + * real when every rank asserted "mpi_assert_max_num_notify" at window + * creation; otherwise MPI_WIN_SET_NUM_NOTIFY grows the counters on demand + * and the only limit is what can be allocated. Neither depends on the + * window's flavor. */ + *num_sb = (int) module->notify_capacity; + *num_ub = (0 != module->notify_max_assert) ? (int) module->notify_max_assert + : INT_MAX; + + /* Counters are uint64_t and only ever incremented by one per notified + * operation, but they are returned to the user as a signed MPI_Count, so + * the representable maximum of that type is the real bound. */ + *value_ub = (OMPI_MPI_COUNT_TYPE) INT64_MAX; + + return OMPI_SUCCESS; +} + +int ompi_osc_ucx_put_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = ompi_osc_ucx_put(origin_addr, origin_count, origin_dt, + target, target_disp, target_count, target_dt, win); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Flush to ensure the PUT is visible at the target before the counter + * increment arrives. */ + ret = opal_common_ucx_wpmem_fence(module->mem); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + /* Atomically increment the target's notify counter in-place using the + * same mem handle as the window data. */ + return osc_ucx_notify_target(module, target, notify, ep); +} + +int ompi_osc_ucx_get_notify(void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = ompi_osc_ucx_get(origin_addr, origin_count, origin_dt, + target, target_disp, target_count, target_dt, win); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Flush to ensure the GET data is locally available before issuing the + * counter increment back to the target. */ + ret = opal_common_ucx_ctx_flush(module->ctx, OPAL_COMMON_UCX_SCOPE_EP, target); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + return osc_ucx_notify_target(module, target, notify, ep); +} + +int ompi_osc_ucx_rput_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win, + struct ompi_request_t **request) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = check_sync_state(module, target, true); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Issue the data movement and the notification first, and only then build + * the request over a flush -- that way the flush covers both, so completing + * the request means the counter update has been pushed out and not merely + * queued behind the origin's next MPI call. */ + ret = ompi_osc_ucx_put(origin_addr, origin_count, origin_dt, target, + target_disp, target_count, target_dt, win); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Fence to order the PUT before the counter increment. */ + ret = opal_common_ucx_wpmem_fence(module->mem); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + ret = osc_ucx_notify_target(module, target, notify, ep); + if (OMPI_SUCCESS != ret) { + return ret; + } + + return osc_ucx_request_over_flush(module, win, target, ep, RPUT_REQ, request); +} + +int ompi_osc_ucx_rget_notify(void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + int notify, struct ompi_win_t *win, + struct ompi_request_t **request) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = check_sync_state(module, target, true); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* As in rput_notify, the data movement and the notification both precede + * the request-bearing flush so that the request covers both. */ + ret = ompi_osc_ucx_get(origin_addr, origin_count, origin_dt, target, + target_disp, target_count, target_dt, win); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Blocking flush: the notification tells the target its window has been + * read, so the GET must have completed before the counter is incremented + * (MPI-5.1 12.6.4 requires that order). */ + ret = opal_common_ucx_ctx_flush(module->ctx, OPAL_COMMON_UCX_SCOPE_EP, target); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + ret = osc_ucx_notify_target(module, target, notify, ep); + if (OMPI_SUCCESS != ret) { + return ret; + } + + return osc_ucx_request_over_flush(module, win, target, ep, RGET_REQ, request); +} + +int ompi_osc_ucx_accumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = ompi_osc_ucx_accumulate(origin_addr, origin_count, origin_dt, + target, target_disp, target_count, target_dt, + op, win); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Fence so that the accumulate is applied at the target before the counter + * increment, as §12.6.4 requires. */ + ret = opal_common_ucx_wpmem_fence(module->mem); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + return osc_ucx_notify_target(module, target, notify, ep); +} + +int ompi_osc_ucx_get_accumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + void *result_addr, size_t result_count, + struct ompi_datatype_t *result_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = ompi_osc_ucx_get_accumulate(origin_addr, origin_count, origin_dt, + result_addr, result_count, result_dt, + target, target_disp, target_count, target_dt, + op, win); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Flush rather than fence: the result buffer must be locally valid, and the + * update must have been applied at the target, before the target may + * observe the notification. */ + ret = opal_common_ucx_ctx_flush(module->ctx, OPAL_COMMON_UCX_SCOPE_EP, target); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + return osc_ucx_notify_target(module, target, notify, ep); +} + +int ompi_osc_ucx_raccumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win, + struct ompi_request_t **request) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = check_sync_state(module, target, true); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Passing a NULL request runs accumulate_req as the blocking accumulate: it + * ends with a blocking flush, so the read-modify-write has completed at the + * target before we return here. The notification issued below is therefore + * correctly ordered after the window access (MPI-5.1 12.6.4). Building our + * own request afterwards -- rather than taking the pre-completed one from + * ompi_osc_ucx_raccumulate -- is what lets that request also cover the + * notification. */ + ret = accumulate_req(origin_addr, origin_count, origin_dt, target, target_disp, + target_count, target_dt, op, win, NULL); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Fence to order the accumulate ahead of the counter increment. */ + ret = opal_common_ucx_wpmem_fence(module->mem); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + ret = osc_ucx_notify_target(module, target, notify, ep); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* The accumulate is already complete, so the request only has to represent + * the notification reaching the wire; RPUT_REQ selects the plain + * flush-completion behaviour rather than the accumulate state machine. */ + return osc_ucx_request_over_flush(module, win, target, ep, RPUT_REQ, request); +} + +int ompi_osc_ucx_rget_accumulate_notify(const void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + void *result_addr, size_t result_count, + struct ompi_datatype_t *result_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, + struct ompi_op_t *op, int notify, + struct ompi_win_t *win, + struct ompi_request_t **request) +{ + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t *)win->w_osc_module; + ucp_ep_h *ep; + int ret; + + CHECK_NOTIFY_IDX(module, notify, target); + + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + + ret = check_sync_state(module, target, true); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* As in raccumulate_notify: a NULL request makes this the blocking form, + * which ends with a blocking flush, so both the fetched result and the + * update at the target are complete before the notification is issued. */ + ret = get_accumulate_req(origin_addr, origin_count, origin_dt, + result_addr, result_count, result_dt, + target, target_disp, target_count, target_dt, + op, win, NULL); + if (OMPI_SUCCESS != ret) { + return ret; + } + + /* Fence to order the accumulate ahead of the counter increment. */ + ret = opal_common_ucx_wpmem_fence(module->mem); + if (OPAL_SUCCESS != ret) { + return OMPI_ERROR; + } + + ret = osc_ucx_notify_target(module, target, notify, ep); + if (OMPI_SUCCESS != ret) { + return ret; + } + + return osc_ucx_request_over_flush(module, win, target, ep, RGET_REQ, request); +} + static inline bool ompi_osc_need_acc_lock(ompi_osc_ucx_module_t *module, int target) { ompi_osc_ucx_lock_t *lock = NULL; @@ -1490,31 +2079,30 @@ int ompi_osc_ucx_get_accumulate_nb(const void *origin_addr, size_t origin_count, target_count, target_dt, op, win, GET_ACCUMULATE); } -int ompi_osc_ucx_rput(const void *origin_addr, size_t origin_count, - struct ompi_datatype_t *origin_dt, - int target, ptrdiff_t target_disp, size_t target_count, - struct ompi_datatype_t *target_dt, - struct ompi_win_t *win, struct ompi_request_t **request) { - ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t*) win->w_osc_module; - ucp_ep_h *ep; - OSC_UCX_GET_DEFAULT_EP(ep, module, target); +/* Attach an MPI request to a nonblocking worker flush. + * + * The request completes when ucp_worker_flush_nb completes, and that flush + * covers every operation already issued on this worker -- whichever memory + * registration they used. So the request waits for exactly what the caller + * issued *before* getting here. The notified variants rely on this: they issue + * their counter atomic first, so completing the request implies the + * notification has been pushed to the target rather than left queued locally + * (MPI-5.1 12.6.4, advice to implementors). + * + * Allocating the request last also keeps the failure contract clean: anything + * that can fail has already run, so callers return their errors with *request + * untouched, as MPI expects of a call that reports an error. */ +static int osc_ucx_request_over_flush(ompi_osc_ucx_module_t *module, + struct ompi_win_t *win, int target, + ucp_ep_h *ep, enum req_type req_type, + struct ompi_request_t **request) +{ opal_common_ucx_wpmem_t *mem = module->mem; uint64_t remote_addr = (module->state_addrs[target]) + OSC_UCX_STATE_REQ_FLAG_OFFSET; ompi_osc_ucx_generic_request_t *ucx_req = NULL; int ret = OMPI_SUCCESS; - ret = check_sync_state(module, target, true); - if (ret != OMPI_SUCCESS) { - return ret; - } - - ret = ompi_osc_ucx_put(origin_addr, origin_count, origin_dt, target, target_disp, - target_count, target_dt, win); - if (ret != OMPI_SUCCESS) { - return ret; - } - - OMPI_OSC_UCX_GENERIC_REQUEST_ALLOC(win, ucx_req, RPUT_REQ); + OMPI_OSC_UCX_GENERIC_REQUEST_ALLOC(win, ucx_req, req_type); ucx_req->super.module = module; OSC_UCX_INCREMENT_OUTSTANDING_NB_OPS(module); @@ -1544,17 +2132,14 @@ int ompi_osc_ucx_rput(const void *origin_addr, size_t origin_count, return ret; } -int ompi_osc_ucx_rget(void *origin_addr, size_t origin_count, +int ompi_osc_ucx_rput(const void *origin_addr, size_t origin_count, struct ompi_datatype_t *origin_dt, int target, ptrdiff_t target_disp, size_t target_count, - struct ompi_datatype_t *target_dt, struct ompi_win_t *win, - struct ompi_request_t **request) { + struct ompi_datatype_t *target_dt, + struct ompi_win_t *win, struct ompi_request_t **request) { ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t*) win->w_osc_module; ucp_ep_h *ep; OSC_UCX_GET_DEFAULT_EP(ep, module, target); - opal_common_ucx_wpmem_t *mem = module->mem; - uint64_t remote_addr = (module->state_addrs[target]) + OSC_UCX_STATE_REQ_FLAG_OFFSET; - ompi_osc_ucx_generic_request_t *ucx_req = NULL; int ret = OMPI_SUCCESS; ret = check_sync_state(module, target, true); @@ -1562,40 +2147,37 @@ int ompi_osc_ucx_rget(void *origin_addr, size_t origin_count, return ret; } - ret = ompi_osc_ucx_get(origin_addr, origin_count, origin_dt, target, target_disp, + ret = ompi_osc_ucx_put(origin_addr, origin_count, origin_dt, target, target_disp, target_count, target_dt, win); if (ret != OMPI_SUCCESS) { return ret; } - OMPI_OSC_UCX_GENERIC_REQUEST_ALLOC(win, ucx_req, RGET_REQ); - ucx_req->super.module = module; + return osc_ucx_request_over_flush(module, win, target, ep, RPUT_REQ, request); +} - OSC_UCX_INCREMENT_OUTSTANDING_NB_OPS(module); - ret = opal_common_ucx_wpmem_flush_ep_nb(mem, target, ompi_osc_ucx_req_completion, ucx_req, ep); +int ompi_osc_ucx_rget(void *origin_addr, size_t origin_count, + struct ompi_datatype_t *origin_dt, + int target, ptrdiff_t target_disp, size_t target_count, + struct ompi_datatype_t *target_dt, struct ompi_win_t *win, + struct ompi_request_t **request) { + ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t*) win->w_osc_module; + ucp_ep_h *ep; + OSC_UCX_GET_DEFAULT_EP(ep, module, target); + int ret = OMPI_SUCCESS; + ret = check_sync_state(module, target, true); if (ret != OMPI_SUCCESS) { - /* fallback to using an atomic op to acquire a request handle */ - ret = opal_common_ucx_wpmem_fence(mem); - if (ret != OMPI_SUCCESS) { - OSC_UCX_VERBOSE(1, "opal_common_ucx_mem_fence failed: %d", ret); - OMPI_OSC_UCX_REQUEST_RETURN(ucx_req); - return OMPI_ERROR; - } - - ret = opal_common_ucx_wpmem_fetch_nb(mem, UCP_ATOMIC_FETCH_OP_FADD, - 0, target, &(module->req_result), - sizeof(uint64_t), remote_addr & (~0x7), - ompi_osc_ucx_req_completion, ucx_req, ep); - if (ret != OMPI_SUCCESS) { - OMPI_OSC_UCX_REQUEST_RETURN(ucx_req); - return ret; - } + return ret; } - *request = &ucx_req->super.super; + ret = ompi_osc_ucx_get(origin_addr, origin_count, origin_dt, target, target_disp, + target_count, target_dt, win); + if (ret != OMPI_SUCCESS) { + return ret; + } - return ret; + return osc_ucx_request_over_flush(module, win, target, ep, RGET_REQ, request); } int ompi_osc_ucx_raccumulate(const void *origin_addr, size_t origin_count, diff --git a/ompi/mca/osc/ucx/osc_ucx_component.c b/ompi/mca/osc/ucx/osc_ucx_component.c index 635a53a3e0f..4c33b366543 100644 --- a/ompi/mca/osc/ucx/osc_ucx_component.c +++ b/ompi/mca/osc/ucx/osc_ucx_component.c @@ -25,6 +25,10 @@ #include "osc_ucx.h" #include "osc_ucx_request.h" #include "opal/util/sys_limits.h" +#include "opal/util/info.h" +#include "opal/class/opal_cstring.h" + +#include #define memcpy_off(_dst, _src, _len, _off) \ memcpy(((char*)(_dst)) + (_off), _src, _len); \ @@ -102,6 +106,20 @@ ompi_osc_ucx_module_t ompi_osc_ucx_module_template = { .osc_fetch_and_op = ompi_osc_ucx_fetch_and_op, .osc_get_accumulate = ompi_osc_ucx_get_accumulate, + .osc_put_notify = ompi_osc_ucx_put_notify, + .osc_get_notify = ompi_osc_ucx_get_notify, + .osc_rput_notify = ompi_osc_ucx_rput_notify, + .osc_rget_notify = ompi_osc_ucx_rget_notify, + .osc_accumulate_notify = ompi_osc_ucx_accumulate_notify, + .osc_get_accumulate_notify = ompi_osc_ucx_get_accumulate_notify, + .osc_raccumulate_notify = ompi_osc_ucx_raccumulate_notify, + .osc_rget_accumulate_notify = ompi_osc_ucx_rget_accumulate_notify, + .osc_win_get_notify_value = ompi_osc_ucx_win_get_notify_value, + .osc_win_reset_notify_value = ompi_osc_ucx_win_reset_notify_value, + .osc_win_set_num_notify = ompi_osc_ucx_win_set_num_notify, + .osc_win_get_num_notify = ompi_osc_ucx_win_get_num_notify, + .osc_win_get_notify_bounds = ompi_osc_ucx_win_get_notify_bounds, + .osc_rput = ompi_osc_ucx_rput, .osc_rget = ompi_osc_ucx_rget, .osc_raccumulate = ompi_osc_ucx_raccumulate, @@ -150,6 +168,53 @@ static bool check_config_value_bool (char *key, opal_info_t *info) return flag_value[0]; } +/* Read the mpi_assert_max_num_notify info key (MPI-5.1 section 12.2) to decide + * how many notification counters to reserve per MPI process initially, and + * report in *asserted whether the key was actually given. + * + * The key is an assertion by the caller that it will not ask + * MPI_WIN_SET_NUM_NOTIFY for more than this, which lets us size the + * registration once and treat it as a hard upper bound. Without it the + * standard is explicit that "the implementation does not assume any limit", so + * the reservation is only a starting size and the counters grow on demand. */ +static int osc_ucx_reserved_notify_counters(opal_info_t *info, unsigned int *reserved, + bool *asserted) +{ + opal_cstring_t *value_string; + int flag = 0, value = 0; + + *reserved = mca_osc_ucx_component.num_notify_counters; + *asserted = false; + + if (NULL == info) { + return OMPI_SUCCESS; + } + + if (OMPI_SUCCESS != opal_info_get(info, "mpi_assert_max_num_notify", + &value_string, &flag) || !flag) { + return OMPI_SUCCESS; + } + + if (OPAL_SUCCESS != opal_cstring_to_int(value_string, &value)) { + OBJ_RELEASE(value_string); + return MPI_ERR_INFO; + } + OBJ_RELEASE(value_string); + + /* A negative value is a malformed key rather than "no assertion"; only 0 + * carries the "assume nothing" meaning. */ + if (value < 0) { + return MPI_ERR_INFO; + } + + if (0 != value) { + *reserved = (unsigned int) value; + *asserted = true; + } + + return OMPI_SUCCESS; +} + static int component_open(void) { opal_common_ucx_mca_register(); @@ -220,6 +285,23 @@ static int component_register(void) { MCA_BASE_VAR_SCOPE_GROUP, &ompi_osc_ucx_outstanding_ops_flush_threshold); free(description_str); + mca_osc_ucx_component.num_notify_counters = OMPI_OSC_UCX_DEFAULT_NOTIFY_COUNTERS; + + opal_asprintf(&description_str, + "Number of RMA notification counters reserved per MPI process " + "in the registered memory region of each window. Windows whose " + "info gives an mpi_assert_max_num_notify value use that instead. " + "This is a hard upper bound: the counters share the window's " + "registration, so MPI_Win_set_num_notify cannot exceed it " + "(default: %u)", + mca_osc_ucx_component.num_notify_counters); + (void) mca_base_component_var_register(&mca_osc_ucx_component.super.osc_version, + "num_notify_counters", description_str, + MCA_BASE_VAR_TYPE_UNSIGNED_INT, NULL, 0, 0, + OPAL_INFO_LVL_3, MCA_BASE_VAR_SCOPE_GROUP, + &mca_osc_ucx_component.num_notify_counters); + free(description_str); + opal_common_ucx_mca_var_register(&mca_osc_ucx_component.super.osc_version); if (0 == access ("/dev/shm", W_OK)) { @@ -559,7 +641,8 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt opal_common_ucx_mem_type_t mem_type; char *my_mem_addr; int my_mem_addr_size; - uint64_t my_info[3] = {0}; + uint64_t my_info[4] = {0}; + void *notify_base = NULL; char *recv_buf = NULL; void *dynamic_base = NULL; unsigned long adjusted_size = size; @@ -672,6 +755,48 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt module->flavor = flavor; module->size = size; + + /* How many notification counters to reserve per MPI process. Read before + * the counter region is allocated below. A malformed info value must not + * make this rank skip the allreduce that follows -- that would leave the + * rest of the group blocked in window creation -- so the failure is carried + * through the collective as a negative reservation instead. */ + unsigned int notify_reserved = 0; + bool notify_asserted = false; + int notify_values[3]; + /* The reservation is exchanged as an int and is used to size an allocation, + * so a value that does not fit is a bad configuration rather than a request + * to honor. Both failures ride the flag, not the value: MPI_MAX would hide + * a sentinel value behind some other rank's larger reservation. */ + bool notify_bad = (OMPI_SUCCESS != osc_ucx_reserved_notify_counters(info, ¬ify_reserved, + ¬ify_asserted)) + || notify_reserved > (unsigned int) INT_MAX; + notify_values[0] = notify_bad ? 1 : 0; + notify_values[1] = notify_bad ? 0 : (int) notify_reserved; + /* Carried as "some rank did NOT assert" so that it combines under MPI_MAX + * along with the other two values. */ + notify_values[2] = notify_asserted ? 0 : 1; + + /* info is allowed to differ between MPI processes, so agree on one + * reservation for the whole window. Taking the maximum keeps every rank's + * own assertion satisfiable, and propagates any rank's failure flag. */ + ret = module->comm->c_coll->coll_allreduce(MPI_IN_PLACE, notify_values, 3, + MPI_INT, MPI_MAX, module->comm, + module->comm->c_coll->coll_allreduce_module); + if (OMPI_SUCCESS != ret) { + goto error; + } + if (0 != notify_values[0]) { + ret = MPI_ERR_INFO; + goto error; + } + module->notify_capacity = (unsigned int) notify_values[1]; + /* Only treat the reservation as a hard cap when *every* rank asserted a + * bound. A rank that gave no key made no promise, so the window has to + * stay growable for it -- MPI-5.1 12.2 says an absent (zero) key means the + * implementation assumes no limit. */ + module->notify_max_assert = notify_values[2] ? 0 : (unsigned int) notify_values[1]; + module->no_locks = check_config_value_bool ("no_locks", info); module->acc_single_intrinsic = check_config_value_bool ("acc_single_intrinsic", info); module->skip_sync_check = false; @@ -849,7 +974,6 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt goto error; } - for (i = 0, total = 0; i < comm_size ; ++i) { size_t peer_size = ompi_osc_ucx_get_size(module, i); if (peer_size || !module->noncontig_shared_win) { @@ -884,7 +1008,8 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt ret = OMPI_ERR_BAD_PARAM; goto error; } - ret = opal_common_ucx_wpmem_create(module->ctx, mem_base, module->size, + ret = opal_common_ucx_wpmem_create(module->ctx, mem_base, + module->size, mem_type, &exchange_len_info, OPAL_COMMON_UCX_WPMEM_ADDR_EXCHANGE_FULL, (void *)module->comm, @@ -899,6 +1024,37 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt ucp_rkey_buffer_release(my_mem_addr); } + /* Notification counters live in their own registered region, like the + * window state does. Appending them to the window data would mean writing + * past the end of the user's buffer for MPI_WIN_FLAVOR_CREATE, and would + * leave dynamic windows -- which have no data region -- with nowhere to put + * them. */ + if (0 != module->notify_capacity) { + module->notify_base = calloc(module->notify_capacity, sizeof(uint64_t)); + if (NULL == module->notify_base) { + ret = OMPI_ERR_TEMP_OUT_OF_RESOURCE; + goto error; + } + + notify_base = module->notify_base; + ret = opal_common_ucx_wpmem_create(module->ctx, ¬ify_base, + module->notify_capacity * sizeof(uint64_t), + OPAL_COMMON_UCX_MEM_MAP, + &exchange_len_info, + OPAL_COMMON_UCX_WPMEM_ADDR_EXCHANGE_FULL, + (void *)module->comm, + &my_mem_addr, &my_mem_addr_size, + &module->notify_mem); + if (ret != OMPI_SUCCESS) { + goto error; + } + + if (my_mem_addr_size != 0) { + /* rkey object is already distributed among comm processes */ + ucp_rkey_buffer_release(my_mem_addr); + } + } + state_base = (void *)&(module->state); ret = opal_common_ucx_wpmem_create(module->ctx, &state_base, sizeof(ompi_osc_ucx_state_t), @@ -926,6 +1082,7 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt } my_info[1] = (uint64_t)state_base; my_info[2] = ompi_comm_rank(&ompi_mpi_comm_world.comm); + my_info[3] = (uint64_t)module->notify_base; recv_buf = (char *)calloc(comm_size, sizeof(my_info)); ret = comm->c_coll->coll_allgather((void *)my_info, sizeof(my_info), @@ -941,10 +1098,17 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt module->addrs = calloc(comm_size, sizeof(uint64_t)); module->state_addrs = calloc(comm_size, sizeof(uint64_t)); module->comm_world_ranks = calloc(comm_size, sizeof(uint64_t)); + /* Number of notification counters attached at each rank; starts at zero + * everywhere (consistent without communication) and is updated by + * MPI_WIN_SET_NUM_NOTIFY. Counters must be attached before use. */ + module->notify_counts = calloc(comm_size, sizeof(int)); + module->notify_addrs = calloc(comm_size, sizeof(uint64_t)); for (i = 0; i < comm_size; i++) { - memcpy(&(module->addrs[i]), recv_buf + i * 3 * sizeof(uint64_t), sizeof(uint64_t)); - memcpy(&(module->state_addrs[i]), recv_buf + i * 3 * sizeof(uint64_t) + sizeof(uint64_t), sizeof(uint64_t)); - memcpy(&(module->comm_world_ranks[i]), recv_buf + i * 3 * sizeof(uint64_t) + 2 * sizeof(uint64_t), sizeof(uint64_t)); + const char *entry = recv_buf + i * sizeof(my_info); + memcpy(&(module->addrs[i]), entry, sizeof(uint64_t)); + memcpy(&(module->state_addrs[i]), entry + sizeof(uint64_t), sizeof(uint64_t)); + memcpy(&(module->comm_world_ranks[i]), entry + 2 * sizeof(uint64_t), sizeof(uint64_t)); + memcpy(&(module->notify_addrs[i]), entry + 3 * sizeof(uint64_t), sizeof(uint64_t)); } free(recv_buf); @@ -957,6 +1121,7 @@ static int component_select(struct ompi_win_t *win, void **base, size_t size, pt module->state.acc_lock = TARGET_LOCK_UNLOCKED; module->state.dynamic_lock = TARGET_LOCK_UNLOCKED; module->state.dynamic_win_count = 0; + for (i = 0; i < OMPI_OSC_UCX_ATTACH_MAX; i++) { module->local_dynamic_win_info[i].refcnt = 0; } @@ -1091,6 +1256,92 @@ int ompi_osc_ucx_dynamic_unlock(ompi_osc_ucx_module_t *module, int target) { return OMPI_SUCCESS; } +/* Collectively replace the notification-counter registration with a larger one. + * + * Called from MPI_WIN_SET_NUM_NOTIFY, which the standard defines as a blocking, + * synchronizing collective procedure -- that is what makes this safe. Every + * rank has to take part even if its own request fits, because registering the + * memory exchanges rkeys with the whole group. + * + * Two properties of MPI_WIN_SET_NUM_NOTIFY keep this simple. It resets every + * counter to zero, so a freshly calloc'd region is already the required + * contents and no value has to be carried across. And it is erroneous to call + * it while an access epoch is open or with an active notification-threshold + * request, so no remote atomic can be in flight against the old region while it + * is being replaced. + * + * The allgather of the new base addresses doubles as the barrier that lets the + * old region be released: once it completes, every rank has published its new + * address and no rank can issue a notified operation until it returns from the + * enclosing collective. */ +int ompi_osc_ucx_grow_notify_counters(ompi_osc_ucx_module_t *module, + unsigned int new_capacity) +{ + int comm_size = ompi_comm_size(module->comm); + opal_common_ucx_wpmem_t *new_mem = NULL; + void *new_base = NULL, *reg_base; + char *my_mem_addr = NULL; + uint64_t my_addr, *new_addrs = NULL; + int my_mem_addr_size = 0; + int ret; + + new_base = calloc(new_capacity, sizeof(uint64_t)); + if (NULL == new_base) { + return MPI_ERR_NO_MEM; + } + + new_addrs = calloc(comm_size, sizeof(uint64_t)); + if (NULL == new_addrs) { + free(new_base); + return MPI_ERR_NO_MEM; + } + + reg_base = new_base; + ret = opal_common_ucx_wpmem_create(module->ctx, ®_base, + new_capacity * sizeof(uint64_t), + OPAL_COMMON_UCX_MEM_MAP, + &exchange_len_info, + OPAL_COMMON_UCX_WPMEM_ADDR_EXCHANGE_FULL, + (void *)module->comm, + &my_mem_addr, &my_mem_addr_size, + &new_mem); + if (OMPI_SUCCESS != ret) { + free(new_addrs); + free(new_base); + return ret; + } + + if (0 != my_mem_addr_size) { + /* rkey object is already distributed among comm processes */ + ucp_rkey_buffer_release(my_mem_addr); + } + + my_addr = (uint64_t) new_base; + ret = module->comm->c_coll->coll_allgather(&my_addr, sizeof(uint64_t), MPI_BYTE, + new_addrs, sizeof(uint64_t), MPI_BYTE, + module->comm, + module->comm->c_coll->coll_allgather_module); + if (OMPI_SUCCESS != ret) { + opal_common_ucx_wpmem_free(new_mem); + free(new_addrs); + free(new_base); + return ret; + } + + if (NULL != module->notify_mem) { + opal_common_ucx_wpmem_free(module->notify_mem); + } + free(module->notify_base); + free(module->notify_addrs); + + module->notify_mem = new_mem; + module->notify_base = new_base; + module->notify_addrs = new_addrs; + module->notify_capacity = new_capacity; + + return OMPI_SUCCESS; +} + int ompi_osc_ucx_win_attach(struct ompi_win_t *win, void *base, size_t len) { ompi_osc_ucx_module_t *module = (ompi_osc_ucx_module_t*) win->w_osc_module; int insert_index = -1, contain_index; @@ -1245,11 +1496,17 @@ int ompi_osc_ucx_free(struct ompi_win_t *win) { free(module->addrs); free(module->state_addrs); free(module->comm_world_ranks); + free(module->notify_counts); + free(module->notify_addrs); opal_common_ucx_wpmem_free(module->state_mem); if (NULL != module->mem) { opal_common_ucx_wpmem_free(module->mem); } + if (NULL != module->notify_mem) { + opal_common_ucx_wpmem_free(module->notify_mem); + } + free(module->notify_base); opal_common_ucx_wpctx_release(module->ctx); diff --git a/ompi/mpi/c/win_set_num_notify.c.in b/ompi/mpi/c/win_set_num_notify.c.in index cc1d39a9e77..f85bf1c8fd8 100644 --- a/ompi/mpi/c/win_set_num_notify.c.in +++ b/ompi/mpi/c/win_set_num_notify.c.in @@ -30,10 +30,17 @@ PROTOTYPE ERROR_CLASS win_set_num_notify(WIN win, INFO info, INT num_notificatio return OMPI_ERRHANDLER_NOHANDLE_INVOKE(MPI_ERR_WIN, FUNC_NAME); } else if (NULL != info && MPI_INFO_NULL != info && ompi_info_is_freed(info)) { rc = MPI_ERR_INFO; - } else if (num_notifications < 0) { - rc = MPI_ERR_ARG; } + /* num_notifications is deliberately *not* range-checked here. + MPI_WIN_SET_NUM_NOTIFY is a synchronizing collective whose count + argument is local -- MPI-5.1 12.6.1 allows it to differ between MPI + processes -- so a rank that rejected its own value here would return + while every other rank stayed blocked in the osc module's internal + collective, turning an erroneous argument into a hang. The module + carries each rank's validity through that collective instead, so all + ranks agree to fail together. */ + OMPI_ERRHANDLER_CHECK(rc, win, rc, FUNC_NAME); }