1- #include < cuda_fp16.h>
21#include < cuda_runtime.h>
32#include < nvshmem.h>
43#include < nvshmemx.h>
54#include < stdint.h>
65#include < stdio.h>
7- #include < stdlib.h>
86
97#define CUDA_CHECK (stmt ) \
108 do { \
1614 } \
1715 } while (0 )
1816
19- extern " C" int ag_gemm_workspace_create (int elements_per_rank, void **workspace,
20- uint64_t **ready, int *mype, int *npes,
21- int *mype_in_node, int *npes_in_node) {
22- if (elements_per_rank <= 0 || workspace == nullptr || ready == nullptr ) {
17+ extern " C" int ag_gemm_workspace_create (int elements_per_rank, int element_size,
18+ void **workspace, uint64_t **ready,
19+ int *mype, int *npes, int *mype_in_node,
20+ int *npes_in_node) {
21+ if (elements_per_rank <= 0 || element_size <= 0 || workspace == nullptr ||
22+ ready == nullptr ) {
2323 return -1 ;
2424 }
2525
@@ -29,7 +29,7 @@ extern "C" int ag_gemm_workspace_create(int elements_per_rank, void **workspace,
2929 *npes_in_node = nvshmem_team_n_pes (NVSHMEMX_TEAM_NODE );
3030 CUDA_CHECK (cudaSetDevice (*mype_in_node));
3131
32- size_t workspace_bytes = (size_t )(*npes) * elements_per_rank * sizeof (__half) ;
32+ size_t workspace_bytes = (size_t )(*npes) * elements_per_rank * element_size ;
3333 *workspace = nvshmem_malloc (workspace_bytes);
3434 *ready = (uint64_t *)nvshmem_calloc ((size_t )(*npes), sizeof (uint64_t ));
3535 if (*workspace == nullptr || *ready == nullptr ) {
@@ -61,3 +61,13 @@ extern "C" void *ag_gemm_peer_workspace_ptr(void *workspace, int peer) {
6161extern " C" uint64_t *ag_gemm_peer_ready_ptr (uint64_t *ready, int peer) {
6262 return (uint64_t *)nvshmem_ptr (ready, peer);
6363}
64+
65+ extern " C" void ag_gemm_barrier_all_on_stream (cudaStream_t stream) {
66+ nvshmemx_barrier_all_on_stream (stream);
67+ }
68+
69+ extern " C" void ag_gemm_signal_wait_until_on_stream (uint64_t *signal,
70+ uint64_t value,
71+ cudaStream_t stream) {
72+ nvshmemx_signal_wait_until_on_stream (signal, NVSHMEM_CMP_GE , value, stream);
73+ }
0 commit comments