Using Counted Writes with NVSHMEM#

Counted writes transfer data to a remote symmetric destination and then add the number of delivered payload bytes to a remote 64-bit counter. A receiver can wait for an expected byte count and consume the corresponding data after the wait completes.

Motivation#

Counted writes are useful for protocols in which one or more producers write payloads to a receiver and the receiver needs to know how much data is ready. For example, several producer CTAs can write disjoint payload regions while contributing to one receiver-local counter. The receiver waits for the sum of the expected payload sizes instead of managing a separate flag for every producer.

Counted writes differ from ordinary NVSHMEM put-with-signal operations in the following ways:

  • An ordinary put-with-signal applies a caller-selected signal value and signal operation after transferring its payload. A counted write always increments its counter by the number of payload bytes delivered. It does not accept a signal value or signal operation.

  • Contributions from multiple counted writes accumulate on the same counter. Applications therefore use byte-count epochs rather than assigning signal values to individual transfers.

  • Counted writes require the CUDA Compute Fabric Transport (CFT) path. They do not fall back to an ordinary put-with-signal operation or a network transport when CFT counted operations are unavailable.

Requirements and Setup#

Build NVSHMEM with CFT handle support, enable logical endpoints and TMA at run time, and donate shared memory from every CTA that issues counted writes. See Using CFT Handles with NVSHMEM for the complete build, run-time, platform, and shared-memory setup requirements.

Use NVSHMEMX_SMEM_BARRIERS_ONLY when the payload source is already in application shared memory. For a global-memory source, use NVSHMEMX_SMEM_MINIMUM or NVSHMEMX_SMEM_RECOMMENDED so that NVSHMEM can stage the payload through donated shared memory.

Using the API#

NVSHMEM 3.8 provides the following counted-write functions:

__host__ __device__ void
nvshmemx_signal_counted_reset(uint64_t *signal_addr);

__device__ uint64_t
nvshmemx_signal_counted_load(const uint64_t *signal_addr);

__device__ void
nvshmemx_signal_counted_wait_until(const uint64_t *signal_addr,
                                   uint64_t expected);

__device__ int
nvshmemx_putmem_signal_counted_nbi_block(void *dest,
                                         const void *source,
                                         size_t bytes,
                                         uint64_t *signal_addr,
                                         int pe);

See NVSHMEMX_PUTMEM_SIGNAL_COUNTED_NBI_BLOCK and the preceding counted-signal entries in the signal API reference for complete argument, return-value, and usage requirements.

nvshmemx_signal_counted_reset sets the calling PE’s local counter to zero. Use it during initialization or at a protocol boundary after all operations from the previous epoch have completed. The function does not synchronize with producers or waiters.

nvshmemx_signal_counted_load reads the local counter. nvshmemx_signal_counted_wait_until waits until the local counter reaches the expected byte epoch and performs the acquire operation needed before the receiver consumes the delivered payload.

nvshmemx_putmem_signal_counted_nbi_block is a collective operation over the calling CTA. It writes bytes from source to dest on PE pe and increments signal_addr on that PE by bytes after the payload is delivered. All threads in the CTA must call it with the same arguments. The function returns NVSHMEMX_SUCCESS when the operation is accepted. The receiver must still use the counted counter as its data-ready notification.

The following two-PE example has PE 0 send 256 bytes to PE 1. PE 1 waits for the counter before using its destination buffer:

constexpr size_t bytes = 256;

__global__ void counted_put(void *destination, uint64_t *counter,
                            size_t nvshmem_smem_size) {
    extern __shared__ __align__(16) unsigned char smem[];
    unsigned char *payload = smem + nvshmem_smem_size;

    nvshmemx_give_smem(smem, nvshmem_smem_size);
    __syncthreads();

    for (size_t i = threadIdx.x; i < bytes; i += blockDim.x) {
        payload[i] = (unsigned char)i;
    }
    __syncthreads();

    if (nvshmem_my_pe() == 0) {
        int status = nvshmemx_putmem_signal_counted_nbi_block(
            destination, payload, bytes, counter, 1);
        assert(status == NVSHMEMX_SUCCESS);
    } else {
        nvshmemx_signal_counted_wait_until(counter, bytes);
        /* destination now contains the payload from PE 0. */
    }

    __syncthreads();
    nvshmemx_release_smem();
}

void launch_counted_put() {
    assert(nvshmem_n_pes() == 2);

    void *destination = nvshmem_align(16, bytes);
    uint64_t *counter =
        (uint64_t *)nvshmem_align(256, 256);

    nvshmemx_signal_counted_reset(counter);
    nvshmem_barrier_all();

    size_t nvshmem_smem_size = (size_t)nvshmemx_ask_smem(
        NVSHMEMX_SMEM_BARRIERS_ONLY);
    counted_put<<<1, 128, nvshmem_smem_size + bytes>>>(
        destination, counter, nvshmem_smem_size);
    cudaDeviceSynchronize();

    nvshmem_free(counter);
    nvshmem_free(destination);
}