diff --git a/examples/device/ep/README.md b/examples/device/ep/README.md index 21e0c8c83b..79299e8367 100644 --- a/examples/device/ep/README.md +++ b/examples/device/ep/README.md @@ -18,7 +18,7 @@ import nixl_ep # Initialize buffer with dynamic rank support buffer = nixl_ep.Buffer(rank, explicitly_destroy=True) -buffer.update_memory_buffers(num_ranks, num_experts_per_rank, rdma_bytes) +buffer.update_memory_buffers(num_ranks, num_experts_per_rank, 0, rdma_bytes) buffer.connect_ranks(initial_ranks) # Dispatch & Combine calls @@ -39,7 +39,7 @@ buffer.disconnect_ranks(ranks) ## Key APIs - `Buffer(rank_id, nvlink_backend, explicitly_destroy)`: Initialize the NIXL communication buffer -- `update_memory_buffers(num_ranks, num_experts_per_rank, num_rdma_bytes)`: Prepare buffers for up to `num_ranks` ranks and `num_experts_per_rank` experts +- `update_memory_buffers(num_ranks, num_experts_per_rank, num_nvl_bytes, num_rdma_bytes)`: Prepare buffers for up to `num_ranks` ranks and `num_experts_per_rank` experts - `connect_ranks(remote_ranks)`: Establish NIXL connections to new peers (can be called multiple times) - `disconnect_ranks(remote_ranks)`: Clean up connections to departing peers diff --git a/examples/device/ep/csrc/config.hpp b/examples/device/ep/csrc/config.hpp index f17ac113d7..2d7c1b21af 100644 --- a/examples/device/ep/csrc/config.hpp +++ b/examples/device/ep/csrc/config.hpp @@ -37,6 +37,82 @@ dtype_t align_up(dtype_t a, dtype_t b) { return ceil_div(a, b) * b; } +struct Config { + int num_sms; + int num_max_nvl_chunked_send_tokens; + int num_max_nvl_chunked_recv_tokens; + int num_max_rdma_chunked_send_tokens; + int num_max_rdma_chunked_recv_tokens; + + Config(int num_sms, + int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + int num_max_rdma_chunked_send_tokens, int num_max_rdma_chunked_recv_tokens) : + num_sms(num_sms), + num_max_nvl_chunked_send_tokens(num_max_nvl_chunked_send_tokens), + num_max_nvl_chunked_recv_tokens(num_max_nvl_chunked_recv_tokens), + num_max_rdma_chunked_send_tokens(num_max_rdma_chunked_send_tokens), + num_max_rdma_chunked_recv_tokens(num_max_rdma_chunked_recv_tokens) { + EP_HOST_ASSERT(num_sms >= 0); + EP_HOST_ASSERT(num_max_nvl_chunked_send_tokens > 0 and num_max_nvl_chunked_recv_tokens > 0); + EP_HOST_ASSERT(num_max_nvl_chunked_send_tokens < num_max_nvl_chunked_recv_tokens); + EP_HOST_ASSERT(num_max_rdma_chunked_send_tokens > 0 and num_max_rdma_chunked_recv_tokens > 0); + + // Ceil up RDMA buffer size + this->num_max_rdma_chunked_recv_tokens = align_up(num_max_rdma_chunked_recv_tokens, num_max_rdma_chunked_send_tokens); + EP_HOST_ASSERT(num_max_rdma_chunked_send_tokens < num_max_rdma_chunked_recv_tokens); + // NOTES: this assertion is related to RDMA lazy head update, we must ensure senders always have space to push + EP_HOST_ASSERT(num_max_rdma_chunked_send_tokens <= num_max_rdma_chunked_recv_tokens / 2); + } + + size_t get_nvl_buffer_size_hint(size_t hidden_bytes, int num_ranks) const { + // Below are some assumptions + // TODO: add assertions + constexpr int kNumMaxTopK = 128; + constexpr int kNumMaxScales = 128; + EP_HOST_ASSERT(num_ranks < NUM_MAX_NVL_PEERS or num_ranks % NUM_MAX_NVL_PEERS == 0); + EP_HOST_ASSERT(num_ranks <= NUM_MAX_NVL_PEERS or num_sms % 2 == 0); + const auto num_rdma_ranks = std::max(num_ranks / NUM_MAX_NVL_PEERS, 1); + const auto num_nvl_ranks = std::min(num_ranks, NUM_MAX_NVL_PEERS); + const int num_channels = num_sms / 2; + + size_t num_bytes = 0; + num_bytes += num_channels * num_nvl_ranks * (2 * num_rdma_ranks + 3) * sizeof(int); + num_bytes += num_channels * num_nvl_ranks * num_max_nvl_chunked_recv_tokens * hidden_bytes; + num_bytes += num_channels * num_nvl_ranks * num_max_nvl_chunked_recv_tokens * internode::get_source_meta_bytes(); + num_bytes += num_channels * num_nvl_ranks * num_max_nvl_chunked_recv_tokens * kNumMaxTopK * sizeof(int64_t); + num_bytes += num_channels * num_nvl_ranks * num_max_nvl_chunked_recv_tokens * kNumMaxTopK * sizeof(float); + num_bytes += num_channels * num_nvl_ranks * num_max_nvl_chunked_recv_tokens * kNumMaxScales * sizeof(float); + num_bytes = ((num_bytes + 127) / 128) * 128; + return num_bytes; + } + + size_t get_rdma_buffer_size_hint(int64_t hidden_bytes, int num_ranks) const { + // Legacy mode + if (num_ranks <= NUM_MAX_NVL_PEERS) + return 0; + + // Below are some assumptions + // TODO: add assertions + constexpr int kNumMaxTopK = 128; + constexpr int kNumMaxScales = 128; + EP_HOST_ASSERT(num_ranks % NUM_MAX_NVL_PEERS == 0); + EP_HOST_ASSERT(num_sms % 2 == 0); + const int num_rdma_ranks = num_ranks / NUM_MAX_NVL_PEERS; + const int num_channels = num_sms / 2; + + size_t num_bytes = 0; + num_bytes += num_channels * num_rdma_ranks * (NUM_MAX_NVL_PEERS * 2 + 2) * 2 * sizeof(int); + num_bytes += num_channels * num_rdma_ranks * num_max_rdma_chunked_recv_tokens * hidden_bytes * 2; + num_bytes += num_channels * num_rdma_ranks * num_max_rdma_chunked_recv_tokens * internode::get_source_meta_bytes() * 2; + num_bytes += num_channels * num_rdma_ranks * num_max_rdma_chunked_recv_tokens * kNumMaxTopK * sizeof(int64_t) * 2; + num_bytes += num_channels * num_rdma_ranks * num_max_rdma_chunked_recv_tokens * kNumMaxTopK * sizeof(float) * 2; + num_bytes += num_channels * num_rdma_ranks * num_max_rdma_chunked_recv_tokens * kNumMaxScales * sizeof(float) * 2; + num_bytes += num_channels * num_rdma_ranks * num_max_rdma_chunked_recv_tokens * sizeof(int4) * 2; + num_bytes = ((num_bytes + 127) / 128) * 128; + return num_bytes; + } +}; + struct EPBuffer { int num_clean_int = 0; diff --git a/examples/device/ep/csrc/kernels/api.cuh b/examples/device/ep/csrc/kernels/api.cuh index f4da764e1d..c4d539a927 100644 --- a/examples/device/ep/csrc/kernels/api.cuh +++ b/examples/device/ep/csrc/kernels/api.cuh @@ -30,8 +30,12 @@ namespace nixl_ep { -// EP kernels -namespace ep_kernels { +namespace intranode { + +void barrier(int** barrier_signal_ptrs, int rank, int num_nvl_ranks, cudaStream_t stream); + +} // namespace intranode + struct gpu_nixl_ctx { nixlGpuXferReqH *batch_reqs; // [dest_rank] nixlGpuXferReqH *remote_barrier_reqs; // [dest_rank] @@ -40,8 +44,13 @@ struct gpu_nixl_ctx { void **rdma_p2p_ptrs; // [num_ranks] void *rdma_buffer_ptr; int num_ranks; + int num_rdma_ranks; int rank; + // Internode-specific barrier counters (used only by internode kernels) + uint64_t *last_barrier_counter; // Tracks barrier epochs + uint64_t *local_barrier_counter_ptr; // Receives barrier signals from remote ranks + /* Double buffering considerations are handled by the caller */ __device__ inline void *rdma_p2p_ptr_get(uint64_t ptr, int dst_rank) { if (rdma_p2p_ptrs[dst_rank] == nullptr) @@ -63,46 +72,254 @@ struct gpu_nixl_ctx { } }; -void clean_buffer(int* clean_0, int num_clean_int_0, - int* clean_1, int num_clean_int_1, - int rank, int num_ranks, int* mask_buffer, int* sync_buffer, - cudaStream_t stream); -void dispatch(void* packed_recv_x, void* packed_recv_x_scales, - int* packed_recv_src_info, int64_t* packed_recv_layout_range, +// Internode runtime +namespace internode { + +void *alloc(size_t size, size_t alignment); + +void free(void *ptr); + +} // namespace internode + +// Layout kernels +namespace layout { + +void get_dispatch_layout(const topk_idx_t* topk_idx, + int* num_tokens_per_rank, + int* num_tokens_per_rdma_rank, + int* num_tokens_per_expert, + bool* is_token_in_rank, + int num_tokens, + int num_topk, + int num_ranks, + int num_experts, + cudaStream_t stream); + +} // namespace layout + +// Internode kernels +namespace internode { + +int get_source_meta_bytes(); + +void notify_dispatch(const int* num_tokens_per_rank, + int* moe_recv_counter_mapped, + int num_ranks, + const int* num_tokens_per_rdma_rank, + int* moe_recv_rdma_counter_mapped, + const int* num_tokens_per_expert, + int* moe_recv_expert_counter_mapped, + int num_experts, + const bool* is_token_in_rank, + int num_tokens, + int num_channels, + int hidden_int4, + int num_scales, + int num_topk, + int expert_alignment, + int* rdma_channel_prefix_matrix, + int* recv_rdma_rank_prefix_sum, + int* gbl_channel_prefix_matrix, + int* recv_gbl_rank_prefix_sum, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_recv_tokens, + int** barrier_signal_ptrs, + int rank, + cudaStream_t stream, + int64_t num_rdma_bytes, + int64_t num_nvl_bytes, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx); + +void dispatch(void* recv_x, + float* recv_x_scales, + topk_idx_t* recv_topk_idx, + float* recv_topk_weights, + void* recv_src_meta, + const void* x, + const float* x_scales, + const topk_idx_t* topk_idx, + const float* topk_weights, + int* send_rdma_head, + int* send_nvl_head, + int* recv_rdma_channel_prefix_matrix, + int* recv_gbl_channel_prefix_matrix, + const int* rdma_channel_prefix_matrix, + const int* recv_rdma_rank_prefix_sum, + const int* gbl_channel_prefix_matrix, + const int* recv_gbl_rank_prefix_sum, + const bool* is_token_in_rank, + int num_tokens, + int hidden_int4, + int num_scales, + int num_topk, + int num_experts, + int scale_token_stride, + int scale_hidden_stride, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_send_tokens, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_send_tokens, + int num_max_nvl_chunked_recv_tokens, + int rank, + int num_ranks, + bool is_cached_dispatch, + cudaStream_t stream, + int num_channels, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx); + +void cached_notify(int hidden_int4, + int num_scales, + int num_topk_idx, + int num_topk_weights, + int num_ranks, + int num_channels, + int num_combined_tokens, + int* combined_rdma_head, + const int* rdma_channel_prefix_matrix, + const int* rdma_rank_prefix_sum, + int* combined_nvl_head, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_recv_tokens, + int** barrier_signal_ptrs, + int rank, + cudaStream_t stream, + int64_t num_rdma_bytes, + int64_t num_nvl_bytes, + bool is_cached_dispatch, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx); + +void combine(cudaDataType_t type, + void* combined_x, + float* combined_topk_weights, + const bool* is_combined_token_in_rank, + const void* x, + const float* topk_weights, + const void* bias_0, + const void* bias_1, + const int* combined_rdma_head, + const int* combined_nvl_head, + const void* src_meta, + const int* rdma_channel_prefix_matrix, + const int* rdma_rank_prefix_sum, + const int* gbl_channel_prefix_matrix, + int num_tokens, + int num_combined_tokens, + int hidden, + int num_topk, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_send_tokens, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_send_tokens, + int num_max_nvl_chunked_recv_tokens, + int rank, + int num_ranks, + cudaStream_t stream, + int num_channels, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx); + +} // namespace internode + + +// EP kernels +namespace ep_kernels { +void clean_buffer(int* clean_0, + int num_clean_int_0, + int* clean_1, + int num_clean_int_1, + int rank, + int num_ranks, + int* mask_buffer, + int* sync_buffer, + cudaStream_t stream); + +void dispatch(void* packed_recv_x, + void* packed_recv_x_scales, + int* packed_recv_src_info, + int64_t* packed_recv_layout_range, int* packed_recv_count, int* mask_buffer, int* cumulative_local_expert_recv_stats, int64_t* dispatch_wait_recv_cost_stats, - void* rdma_recv_x, uint64_t* rdma_recv_count, void* rdma_x, - const void* x, const topk_idx_t* topk_idx, - uint64_t* next_clean, int num_next_clean_int, - int num_tokens, int hidden, int num_max_dispatch_tokens_per_rank, - int num_topk, int num_experts, int rank, int num_ranks, - bool use_fp8, bool round_scale, bool use_ue8m0, - void* workspace, int num_device_sms, - cudaStream_t stream, int phases, ep_kernels::gpu_nixl_ctx nixl_ctx); + void* rdma_recv_x, + uint64_t* rdma_recv_count, + void* rdma_x, + const void* x, + const topk_idx_t* topk_idx, + uint64_t* next_clean, + int num_next_clean_int, + int num_tokens, + int hidden, + int num_max_dispatch_tokens_per_rank, + int num_topk, + int num_experts, + int rank, + int num_ranks, + bool use_fp8, + bool round_scale, + bool use_ue8m0, + void* workspace, + int num_device_sms, + cudaStream_t stream, + int phases, + nixl_ep::gpu_nixl_ctx nixl_ctx); void combine(void* combined_x, - void* rdma_recv_x, uint64_t* rdma_recv_flag, void* rdma_send_x, - const void* x, const topk_idx_t* topk_idx, const float* topk_weights, - const int* src_info, const int64_t* layout_range, + void* rdma_recv_x, + uint64_t* rdma_recv_flag, + void* rdma_send_x, + const void* x, + const topk_idx_t* topk_idx, + const float* topk_weights, + const int* src_info, + const int64_t* layout_range, int* mask_buffer, int64_t* combine_wait_recv_cost_stats, - uint64_t* next_clean, int num_next_clean_int, - int num_combined_tokens, int hidden, int num_max_dispatch_tokens_per_rank, - int num_topk, int num_experts, int rank, int num_ranks, + uint64_t* next_clean, + int num_next_clean_int, + int num_combined_tokens, + int hidden, + int num_max_dispatch_tokens_per_rank, + int num_topk, + int num_experts, + int rank, + int num_ranks, bool use_logfmt, - void* workspace, int num_device_sms, - cudaStream_t stream, int phases, bool zero_copy, ep_kernels::gpu_nixl_ctx nixl_ctx); + void* workspace, + int num_device_sms, + cudaStream_t stream, + int phases, + bool zero_copy, + nixl_ep::gpu_nixl_ctx nixl_ctx); -void barrier(ep_kernels::gpu_nixl_ctx nixl_ctx, int* mask_buffer_ptr, int* sync_buffer_ptr, cudaStream_t stream); +void barrier(nixl_ep::gpu_nixl_ctx nixl_ctx, + int* mask_buffer_ptr, + int* sync_buffer_ptr, + cudaStream_t stream); -void query_mask_buffer(int* mask_buffer_ptr, int num_ranks, int* output_mask_tensor, cudaStream_t stream); +void query_mask_buffer(int* mask_buffer_ptr, + int num_ranks, + int* output_mask_tensor, + cudaStream_t stream); -void update_mask_buffer(int* mask_buffer_ptr, int rank_to_mask, bool mask, cudaStream_t stream); +void update_mask_buffer(int* mask_buffer_ptr, + int rank_to_mask, + bool mask, + cudaStream_t stream); -void clean_mask_buffer(int* mask_buffer_ptr, int num_ranks, cudaStream_t stream); +void clean_mask_buffer(int* mask_buffer_ptr, + int num_ranks, + cudaStream_t stream); } // namespace ep_kernels diff --git a/examples/device/ep/csrc/kernels/buffer.cuh b/examples/device/ep/csrc/kernels/buffer.cuh new file mode 100644 index 0000000000..98ec442ecb --- /dev/null +++ b/examples/device/ep/csrc/kernels/buffer.cuh @@ -0,0 +1,140 @@ +#pragma once + +#include "configs.cuh" +#include "exception.cuh" + +namespace nixl_ep { + +template +struct Buffer { +private: + uint8_t* ptr; + +public: + int total_bytes; + + __device__ __forceinline__ Buffer() : ptr(nullptr), total_bytes(0) {} + + __device__ __forceinline__ Buffer(void* &gbl_ptr, int num_elems, int offset = 0) { + total_bytes = num_elems * sizeof(dtype_t); + ptr = reinterpret_cast(gbl_ptr) + offset * sizeof(dtype_t); + gbl_ptr = reinterpret_cast(gbl_ptr) + total_bytes; + } + + __device__ __forceinline__ Buffer advance_also(void* &gbl_ptr) { + gbl_ptr = reinterpret_cast(gbl_ptr) + total_bytes; + return *this; + } + + __device__ __forceinline__ dtype_t* buffer() { + return reinterpret_cast(ptr); + } + + __device__ __forceinline__ dtype_t& operator[](int idx) { + return buffer()[idx]; + } +}; + +template +struct AsymBuffer { +private: + uint8_t* ptrs[kNumRanks]; + int num_bytes; + +public: + int total_bytes; + + __device__ __forceinline__ AsymBuffer(void* &gbl_ptr, int num_elems, int num_ranks, + int sm_id = 0, int num_sms = 1, int offset = 0) { + EP_STATIC_ASSERT(kNumRanks == 1, ""); + num_bytes = num_elems * sizeof(dtype_t); + + int per_channel_bytes = num_bytes * num_ranks; + total_bytes = per_channel_bytes * num_sms; + ptrs[0] = reinterpret_cast(gbl_ptr) + per_channel_bytes * sm_id + num_bytes * offset; + gbl_ptr = reinterpret_cast(gbl_ptr) + total_bytes; + } + + __device__ __forceinline__ AsymBuffer(void** gbl_ptrs, int num_elems, int num_ranks, + int sm_id = 0, int num_sms = 1, int offset = 0) { + EP_STATIC_ASSERT(kNumRanks > 1, ""); + num_bytes = num_elems * sizeof(dtype_t); + + int per_channel_bytes = num_bytes * num_ranks; + total_bytes = per_channel_bytes * num_sms; + for (int i = 0; i < kNumRanks; ++ i) { + ptrs[i] = reinterpret_cast(gbl_ptrs[i]) + per_channel_bytes * sm_id + num_bytes * offset; + gbl_ptrs[i] = reinterpret_cast(gbl_ptrs[i]) + total_bytes; + } + } + + __device__ __forceinline__ void advance(int shift) { +#ifdef __CUDACC__ + #pragma unroll +#endif + for (int i = 0; i < kNumRanks; ++ i) + ptrs[i] = ptrs[i] + shift * sizeof(dtype_t); + } + + __device__ __forceinline__ AsymBuffer advance_also(void* &gbl_ptr) { + gbl_ptr = reinterpret_cast(gbl_ptr) + total_bytes; + return *this; + } + + template + __device__ __forceinline__ AsymBuffer advance_also(void** gbl_ptrs) { + for (int i = 0; i < kNumAlsoRanks; ++ i) + gbl_ptrs[i] = reinterpret_cast(gbl_ptrs[i]) + total_bytes; + return *this; + } + + __device__ __forceinline__ dtype_t* buffer(int idx = 0) { + EP_STATIC_ASSERT(kNumRanks == 1, "`buffer` is only available for single rank case"); + return reinterpret_cast(ptrs[0] + num_bytes * idx); + } + + __device__ __forceinline__ dtype_t* buffer_by(int rank_idx, int idx = 0) { + EP_STATIC_ASSERT(kNumRanks > 1, "`buffer` is only available for single rank case"); + return reinterpret_cast(ptrs[rank_idx] + num_bytes * idx); + } +}; + +template +struct SymBuffer { +private: + // NOTES: for non-decoupled case, `recv_ptr` is not used + uint8_t* send_ptr; + uint8_t* recv_ptr; + int num_bytes; + +public: + int total_bytes; + + __device__ __forceinline__ SymBuffer(void* &gbl_ptr, int num_elems, int num_ranks, + int sm_id = 0, int num_sms = 1) { + num_bytes = num_elems * sizeof(dtype_t); + + int per_channel_bytes = num_bytes * num_ranks; + total_bytes = per_channel_bytes * num_sms * (static_cast(kDecoupled) + 1); + send_ptr = reinterpret_cast(gbl_ptr) + per_channel_bytes * sm_id; + recv_ptr = reinterpret_cast(gbl_ptr) + per_channel_bytes * (sm_id + num_sms); + gbl_ptr = reinterpret_cast(gbl_ptr) + total_bytes; + } + + __device__ __forceinline__ dtype_t* send_buffer(int idx = 0) { + EP_STATIC_ASSERT(kDecoupled, "`send_buffer` is only available for non-decoupled case"); + return reinterpret_cast(send_ptr + num_bytes * idx); + } + + __device__ __forceinline__ dtype_t* recv_buffer(int idx = 0) { + EP_STATIC_ASSERT(kDecoupled, "`recv_buffer` is only available for non-decoupled case"); + return reinterpret_cast(recv_ptr + num_bytes * idx); + } + + __device__ __forceinline__ dtype_t* buffer(int idx = 0) { + EP_STATIC_ASSERT(not kDecoupled, "`buffer` is only available for decoupled case"); + return reinterpret_cast(send_ptr + num_bytes * idx); + } +}; + +} // namespace nixl_ep diff --git a/examples/device/ep/csrc/kernels/configs.cuh b/examples/device/ep/csrc/kernels/configs.cuh index cd3ce777d5..73b3c5daf6 100644 --- a/examples/device/ep/csrc/kernels/configs.cuh +++ b/examples/device/ep/csrc/kernels/configs.cuh @@ -22,11 +22,14 @@ #pragma once +#define NUM_MAX_NVL_PEERS 8 +#define NUM_MAX_RDMA_PEERS 20 #define NUM_WORKSPACE_BYTES (32 * 1024 * 1024) #define NUM_MAX_LOCAL_EXPERTS 1024 #define NUM_BUFFER_ALIGNMENT_BYTES 128 #define FINISHED_SUM_TAG 1024 +#define NUM_WAIT_NANOSECONDS 500 #ifndef ENABLE_FAST_DEBUG #define NUM_CPU_TIMEOUT_SECS 100 diff --git a/examples/device/ep/csrc/kernels/internode.cu b/examples/device/ep/csrc/kernels/internode.cu new file mode 100644 index 0000000000..2c56595f22 --- /dev/null +++ b/examples/device/ep/csrc/kernels/internode.cu @@ -0,0 +1,2545 @@ +#include +#include + +#include "buffer.cuh" +#include "configs.cuh" +#include "exception.cuh" +#include "launch.cuh" +#include "utils.cuh" +#include "api.cuh" +#include "nixl_device.cuh" +#include +#include +#include "log.cuh" +#include + +namespace nixl_ep { + +namespace internode { + +struct SourceMeta { + int src_rdma_rank, is_token_in_nvl_rank_bits; + + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS == 8, "Invalid number of maximum NVL peers"); + + __forceinline__ SourceMeta() = default; + + // TODO: faster encoding + __device__ __forceinline__ SourceMeta(int rdma_rank, const bool* is_token_in_nvl_ranks) { + src_rdma_rank = rdma_rank; + is_token_in_nvl_rank_bits = is_token_in_nvl_ranks[0]; + #pragma unroll + for (int i = 1; i < NUM_MAX_NVL_PEERS; ++i) + is_token_in_nvl_rank_bits |= is_token_in_nvl_ranks[i] << i; + } + + __device__ __forceinline__ bool is_token_in_nvl_rank(int nvl_rank) const { return (is_token_in_nvl_rank_bits >> nvl_rank) & 1; } +}; + +EP_STATIC_ASSERT(sizeof(SourceMeta) % sizeof(int) == 0, "Invalid size of `SourceMeta`"); + +int get_source_meta_bytes() { + return sizeof(SourceMeta); +} + +__host__ __device__ __forceinline__ int get_num_bytes_per_token(int hidden_int4, int num_scales, int num_topk_idx, int num_topk_weights) { + return static_cast(align_up(hidden_int4 * sizeof(int4) + sizeof(SourceMeta) + num_scales * sizeof(float) + + num_topk_idx * sizeof(int) + num_topk_weights * sizeof(float), + sizeof(int4))); +} + +__host__ __device__ __forceinline__ std::pair get_rdma_clean_meta(int hidden_int4, + int num_scales, + int num_topk_idx, + int num_topk_weights, + int num_rdma_ranks, + int num_rdma_recv_buffer_tokens, + int num_channels) { + // Return `int32_t` offset and count to clean + return {(get_num_bytes_per_token(hidden_int4, num_scales, num_topk_idx, num_topk_weights) * num_rdma_recv_buffer_tokens * + num_rdma_ranks * 2 * num_channels) / + sizeof(int), + (NUM_MAX_NVL_PEERS * 2 + 4) * num_rdma_ranks * 2 * num_channels}; +} + +__host__ __device__ __forceinline__ std::pair get_nvl_clean_meta(int hidden_int4, + int num_scales, + int num_topk_idx, + int num_topk_weights, + int num_rdma_ranks, + int num_nvl_ranks, + int num_nvl_recv_buffer_tokens, + int num_channels, + bool is_dispatch) { + // Return `int32_t` offset and to clean + EP_STATIC_ASSERT(sizeof(SourceMeta) % sizeof(int) == 0, "Invalid size of `SourceMeta`"); + + return { + (num_nvl_recv_buffer_tokens * get_num_bytes_per_token(hidden_int4, num_scales, num_topk_idx, num_topk_weights) * num_nvl_ranks * + num_channels) / + sizeof(int), + num_nvl_ranks * (2 * num_rdma_ranks + 2) * num_channels, + }; +} + +template +__forceinline__ __device__ int translate_dst_rdma_rank(const int dst_rdma_rank, const int nvl_rank) { + return kLowLatencyMode ? (dst_rdma_rank * NUM_MAX_NVL_PEERS + nvl_rank) : dst_rdma_rank; +} + +__forceinline__ __device__ void nixl_barrier(nixl_ep::gpu_nixl_ctx nixl_ctx, int num_channels) { + int rdma_rank = nixl_ctx.rank / NUM_MAX_NVL_PEERS; + + // Send barrier signals to all other RDMA ranks + for (int j=0; j( + remote_barrier_handle, 0, 1, 0, j, true, &request_status); + EP_DEVICE_ASSERT(status == NIXL_IN_PROG); + while ((status = nixlGpuGetXferStatus(request_status)) == NIXL_IN_PROG); + EP_DEVICE_ASSERT(status == NIXL_SUCCESS); + } + + // Wait for all other RDMA ranks to signal us + uint64_t poll_counter = 0; + uint64_t epoch = ld_acquire_sys_global(nixl_ctx.last_barrier_counter); + uint64_t expected_counter = (epoch + 1) * (nixl_ctx.num_rdma_ranks - 1); + DEVICE_LOG_DEBUG("rank %d nixl_barrier epoch: %lu, expected_counter: %lu", nixl_ctx.rank, epoch, expected_counter); + + while(ld_acquire_sys_global(nixl_ctx.local_barrier_counter_ptr) < expected_counter) + if (++poll_counter % 5000000 == 0) + DEVICE_LOG_DEBUG("rank %d nixl_barrier waiting for epoch: %ld, local_barrier_counter: %ld, poll_counter: %ld", nixl_ctx.rank, epoch, ld_acquire_sys_global(nixl_ctx.local_barrier_counter_ptr), poll_counter); + + st_release_sys_global(nixl_ctx.last_barrier_counter, epoch + 1); + DEVICE_LOG_DEBUG("rank %d nixl_barrier completed epoch: %ld, local_barrier_counter: %ld", nixl_ctx.rank, epoch + 1, ld_acquire_sys_global(nixl_ctx.local_barrier_counter_ptr)); + } +} + +template +__global__ void notify_dispatch(const int* num_tokens_per_rank, + int* moe_recv_counter_mapped, + int num_ranks, + const int* num_tokens_per_rdma_rank, + int* moe_recv_rdma_counter_mapped, + const int* num_tokens_per_expert, + int* moe_recv_expert_counter_mapped, + int num_experts, + const bool* is_token_in_rank, + int num_tokens, + int num_channels, + int expert_alignment, + const int rdma_clean_offset, + const int rdma_num_int_clean, + const int nvl_clean_offset, + const int nvl_num_int_clean, + int* rdma_channel_prefix_matrix, + int* recv_rdma_rank_prefix_sum, + int* gbl_channel_prefix_matrix, + int* recv_gbl_rank_prefix_sum, + void* rdma_buffer_ptr, + void** buffer_ptrs, + int** barrier_signal_ptrs, + int rank, + nixl_ep::gpu_nixl_ctx nixl_ctx) { + auto sm_id = static_cast(blockIdx.x); + auto thread_id = static_cast(threadIdx.x), warp_id = thread_id / 32, lane_id = get_lane_id(); + auto num_threads = static_cast(blockDim.x), num_warps = num_threads / 32; + + auto rdma_rank = rank / NUM_MAX_NVL_PEERS, nvl_rank = rank % NUM_MAX_NVL_PEERS; + auto num_rdma_experts = num_experts / kNumRDMARanks, num_nvl_experts = num_rdma_experts / NUM_MAX_NVL_PEERS; + DEVICE_LOG_DEBUG_LANE_SYNC(0,"notify_dispatch rank: %d, sm_id: %d, num_channels: %d, warp_id: %d, lane_id: %d", rank, sm_id, num_channels, warp_id, lane_id); + if (sm_id == 0) { + // Communication with others + // Global barrier: the first warp does intra-node sync, the second warp does internode sync + EP_DEVICE_ASSERT(num_warps > 1); + EP_DEVICE_ASSERT(kNumRDMARanks <= num_threads); + + __syncthreads(); + + if (thread_id == 32) { + DEVICE_LOG_DEBUG("rank %d notify_dispatch calling nixl_barrier 1", rank); + nixl_barrier(nixl_ctx, num_channels); + } + barrier_block(barrier_signal_ptrs, nvl_rank); + + // Send numbers of tokens per rank/expert to RDMA ranks + auto rdma_buffer_ptr_int = static_cast(rdma_buffer_ptr); + auto rdma_recv_num_tokens_mixed = SymBuffer(rdma_buffer_ptr, NUM_MAX_NVL_PEERS + num_rdma_experts + 1, kNumRDMARanks); + + // Clean up for later data dispatch + EP_DEVICE_ASSERT(rdma_recv_num_tokens_mixed.total_bytes <= rdma_clean_offset * sizeof(int)); + #pragma unroll + for (int i = thread_id; i < rdma_num_int_clean; i += num_threads) + rdma_buffer_ptr_int[rdma_clean_offset + i] = 0; + + // Copy to send buffer + #pragma unroll + for (int i = thread_id; i < num_ranks; i += num_threads) + rdma_recv_num_tokens_mixed.send_buffer(i / NUM_MAX_NVL_PEERS)[i % NUM_MAX_NVL_PEERS] = num_tokens_per_rank[i]; + #pragma unroll + for (int i = thread_id; i < num_experts; i += num_threads) + rdma_recv_num_tokens_mixed.send_buffer(i / num_rdma_experts)[NUM_MAX_NVL_PEERS + i % num_rdma_experts] = + num_tokens_per_expert[i]; + if (thread_id < kNumRDMARanks) + rdma_recv_num_tokens_mixed.send_buffer(thread_id)[NUM_MAX_NVL_PEERS + num_rdma_experts] = num_tokens_per_rdma_rank[thread_id]; + __syncthreads(); + + // Issue send + // TODO: more light fence or barrier or signaling + // TODO: overlap EP barrier and NVL cleaning + for (int i = warp_id; i < kNumRDMARanks; i += num_warps) { + if (i != rdma_rank) { + size_t src_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_recv_num_tokens_mixed.send_buffer(i))); + size_t dst_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_recv_num_tokens_mixed.recv_buffer(rdma_rank))); + size_t msg_size = (NUM_MAX_NVL_PEERS + num_rdma_experts + 1) * sizeof(int); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(i, nvl_rank)); + nixl_status_t status = nixlGpuPostSingleWriteXferReq( + batch_req, 0, src_offset, dst_offset, msg_size, 0, true); + EP_DEVICE_ASSERT(status == NIXL_IN_PROG); + } else { + UNROLLED_WARP_COPY(1, + lane_id, + NUM_MAX_NVL_PEERS + num_rdma_experts + 1, + rdma_recv_num_tokens_mixed.recv_buffer(rdma_rank), + rdma_recv_num_tokens_mixed.send_buffer(i), + ld_volatile_global, + st_na_global); + } + } + __syncthreads(); + + + // Barrier + if (thread_id == 0){ + DEVICE_LOG_DEBUG("rank %d notify_dispatch calling nixl_barrier 2", rank); + nixl_barrier(nixl_ctx, num_channels); + } + __syncthreads(); + + // NVL buffers + auto nvl_send_buffer = thread_id < NUM_MAX_NVL_PEERS ? buffer_ptrs[thread_id] : nullptr; + auto nvl_recv_buffer = buffer_ptrs[nvl_rank]; + auto nvl_reduced_num_tokens_per_expert = Buffer(nvl_recv_buffer, num_rdma_experts).advance_also(nvl_send_buffer); + auto nvl_send_num_tokens_per_rank = AsymBuffer(nvl_send_buffer, kNumRDMARanks, NUM_MAX_NVL_PEERS); + auto nvl_send_num_tokens_per_expert = AsymBuffer(nvl_send_buffer, num_nvl_experts, NUM_MAX_NVL_PEERS); + auto nvl_recv_num_tokens_per_rank = AsymBuffer(nvl_recv_buffer, kNumRDMARanks, NUM_MAX_NVL_PEERS); + auto nvl_recv_num_tokens_per_expert = AsymBuffer(nvl_recv_buffer, num_nvl_experts, NUM_MAX_NVL_PEERS); + + // Clean up for later data dispatch + auto nvl_buffer_ptr_int = static_cast(buffer_ptrs[nvl_rank]); + EP_DEVICE_ASSERT(nvl_reduced_num_tokens_per_expert.total_bytes + nvl_send_num_tokens_per_rank.total_bytes + + nvl_send_num_tokens_per_expert.total_bytes <= + nvl_clean_offset * sizeof(int)); + #pragma unroll + for (int i = thread_id; i < nvl_num_int_clean; i += num_threads) + nvl_buffer_ptr_int[nvl_clean_offset + i] = 0; + + // Reduce number of tokens per expert into the NVL send buffer + // TODO: may use NVSHMEM reduction + EP_DEVICE_ASSERT(num_rdma_experts <= num_threads); + if (thread_id < num_rdma_experts) { + int sum = 0; + #pragma unroll + for (int i = 0; i < kNumRDMARanks; ++i) + sum += rdma_recv_num_tokens_mixed.recv_buffer(i)[NUM_MAX_NVL_PEERS + thread_id]; + nvl_reduced_num_tokens_per_expert[thread_id] = sum; + } + __syncthreads(); + + // Reduce RDMA received tokens + if (thread_id == 0) { + int sum = 0; + #pragma unroll + for (int i = 0; i < kNumRDMARanks; ++i) { + sum += rdma_recv_num_tokens_mixed.recv_buffer(i)[NUM_MAX_NVL_PEERS + num_rdma_experts]; + DEVICE_LOG_DEBUG("rank %d |rdma_recv_num_tokens_mixed.recv_buffer(%d)[%d]address:%p, value:%d", rank, i, NUM_MAX_NVL_PEERS + num_rdma_experts,rdma_recv_num_tokens_mixed.recv_buffer(i),rdma_recv_num_tokens_mixed.recv_buffer(i)[NUM_MAX_NVL_PEERS + num_rdma_experts]); + recv_rdma_rank_prefix_sum[i] = sum; + DEVICE_LOG_DEBUG("rank %d | Receiving RDMA tokens from RDMA rank %d: %d", rank, i, recv_rdma_rank_prefix_sum[i]- (i == 0 ? 0 : recv_rdma_rank_prefix_sum[i - 1])); + } + DEVICE_LOG_DEBUG("rank %d | moe_recv_rdma_counter_mapped: %p, sum: %d", rank, moe_recv_rdma_counter_mapped, sum); + while (ld_volatile_global(moe_recv_rdma_counter_mapped) != -1); + *moe_recv_rdma_counter_mapped = sum; + } + + // Send numbers of tokens per rank/expert to NVL ranks + EP_DEVICE_ASSERT(NUM_MAX_NVL_PEERS <= num_threads); + if (thread_id < NUM_MAX_NVL_PEERS) { + #pragma unroll + for (int i = 0; i < kNumRDMARanks; ++i) + nvl_send_num_tokens_per_rank.buffer(nvl_rank)[i] = rdma_recv_num_tokens_mixed.recv_buffer(i)[thread_id]; + #pragma unroll + for (int i = 0; i < num_nvl_experts; ++i) + nvl_send_num_tokens_per_expert.buffer(nvl_rank)[i] = nvl_reduced_num_tokens_per_expert[thread_id * num_nvl_experts + i]; + } + barrier_block(barrier_signal_ptrs, nvl_rank); + + // Reduce the number of tokens per rank/expert + EP_DEVICE_ASSERT(num_nvl_experts <= num_threads); + if (thread_id == 0) { + int sum = 0; + #pragma unroll + for (int i = 0; i < num_ranks; ++i) { + int src_rdma_rank = i / NUM_MAX_NVL_PEERS, src_nvl_rank = i % NUM_MAX_NVL_PEERS; + sum += nvl_recv_num_tokens_per_rank.buffer(src_nvl_rank)[src_rdma_rank]; + recv_gbl_rank_prefix_sum[i] = sum; + } + while (ld_volatile_global(moe_recv_counter_mapped) != -1); + *moe_recv_counter_mapped = sum; + } + if (thread_id < num_nvl_experts) { + int sum = 0; + #pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) + sum += nvl_recv_num_tokens_per_expert.buffer(i)[thread_id]; + sum = (sum + expert_alignment - 1) / expert_alignment * expert_alignment; + while (ld_volatile_global(moe_recv_expert_counter_mapped + thread_id) != -1); + moe_recv_expert_counter_mapped[thread_id] = sum; + } + + // Finally barrier + if (thread_id == 32) + nixl_barrier(nixl_ctx, num_channels); + barrier_block(barrier_signal_ptrs, nvl_rank); + } else { + // Calculate meta data + int dst_rdma_rank = sm_id - 1; + for (int channel_id = warp_id; channel_id < num_channels; channel_id += num_warps) { + int token_start_idx, token_end_idx; + get_channel_task_range(num_tokens, num_channels, channel_id, token_start_idx, token_end_idx); + + // Iterate over tokens + int total_count = 0, per_nvl_rank_count[NUM_MAX_NVL_PEERS] = {0}; + for (int64_t i = token_start_idx + lane_id; i < token_end_idx; i += 32) { + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * sizeof(bool) == sizeof(uint64_t), "Invalid number of NVL peers"); + auto is_token_in_rank_uint64 = + *reinterpret_cast(is_token_in_rank + i * num_ranks + dst_rdma_rank * NUM_MAX_NVL_PEERS); + auto is_token_in_rank_values = reinterpret_cast(&is_token_in_rank_uint64); + #pragma unroll + for (int j = 0; j < NUM_MAX_NVL_PEERS; ++j) + per_nvl_rank_count[j] += is_token_in_rank_values[j]; + total_count += (is_token_in_rank_uint64 != 0); + } + + // Warp reduce + total_count = warp_reduce_sum(total_count); + #pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) + per_nvl_rank_count[i] = warp_reduce_sum(per_nvl_rank_count[i]); + + // Write into channel matrix + if (elect_one_sync()) { + #pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) + gbl_channel_prefix_matrix[(dst_rdma_rank * NUM_MAX_NVL_PEERS + i) * num_channels + channel_id] = per_nvl_rank_count[i]; + rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id] = total_count; + } + } + + // Calculate prefix sum + __syncthreads(); + if (thread_id == 0) { + auto prefix_row = rdma_channel_prefix_matrix + dst_rdma_rank * num_channels; + #pragma unroll + for (int i = 1; i < num_channels; ++i) + prefix_row[i] += prefix_row[i - 1]; + } + + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS <= 32, "Invalid number of NVL peers"); + if (thread_id < NUM_MAX_NVL_PEERS) { + auto prefix_row = gbl_channel_prefix_matrix + (dst_rdma_rank * NUM_MAX_NVL_PEERS + thread_id) * num_channels; + #pragma unroll + for (int i = 1; i < num_channels; ++i) + prefix_row[i] += prefix_row[i - 1]; + } + } +} + +void notify_dispatch(const int* num_tokens_per_rank, + int* moe_recv_counter_mapped, + int num_ranks, + const int* num_tokens_per_rdma_rank, + int* moe_recv_rdma_counter_mapped, + const int* num_tokens_per_expert, + int* moe_recv_expert_counter_mapped, + int num_experts, + const bool* is_token_in_rank, + int num_tokens, + int num_channels, + int hidden_int4, + int num_scales, + int num_topk, + int expert_alignment, + int* rdma_channel_prefix_matrix, + int* recv_rdma_rank_prefix_sum, + int* gbl_channel_prefix_matrix, + int* recv_gbl_rank_prefix_sum, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_recv_tokens, + int** barrier_signal_ptrs, + int rank, + cudaStream_t stream, + int64_t num_rdma_bytes, + int64_t num_nvl_bytes, + bool low_latency_mode, + nixl_ep::gpu_nixl_ctx nixl_ctx) { +#define NOTIFY_DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ + { \ + auto notify_dispatch_func = low_latency_mode ? notify_dispatch : notify_dispatch; \ + LAUNCH_KERNEL(&cfg, \ + notify_dispatch_func, \ + num_tokens_per_rank, \ + moe_recv_counter_mapped, \ + num_ranks, \ + num_tokens_per_rdma_rank, \ + moe_recv_rdma_counter_mapped, \ + num_tokens_per_expert, \ + moe_recv_expert_counter_mapped, \ + num_experts, \ + is_token_in_rank, \ + num_tokens, \ + num_channels, \ + expert_alignment, \ + rdma_clean_meta.first, \ + rdma_clean_meta.second, \ + nvl_clean_meta.first, \ + nvl_clean_meta.second, \ + rdma_channel_prefix_matrix, \ + recv_rdma_rank_prefix_sum, \ + gbl_channel_prefix_matrix, \ + recv_gbl_rank_prefix_sum, \ + rdma_buffer_ptr, \ + buffer_ptrs, \ + barrier_signal_ptrs, \ + rank, \ + nixl_ctx); \ + } \ + break + + constexpr int kNumThreads = 512; + const auto num_rdma_ranks = num_ranks / NUM_MAX_NVL_PEERS; + + // Get clean meta + auto rdma_clean_meta = + get_rdma_clean_meta(hidden_int4, num_scales, num_topk, num_topk, num_rdma_ranks, num_max_rdma_chunked_recv_tokens, num_channels); + auto nvl_clean_meta = get_nvl_clean_meta(hidden_int4, + num_scales, + num_topk, + num_topk, + num_rdma_ranks, + NUM_MAX_NVL_PEERS, + num_max_nvl_chunked_recv_tokens, + num_channels, + true); + EP_HOST_ASSERT((rdma_clean_meta.first + rdma_clean_meta.second) * sizeof(int) <= num_rdma_bytes); + EP_HOST_ASSERT((nvl_clean_meta.first + nvl_clean_meta.second) * sizeof(int) <= num_nvl_bytes); + EP_HOST_ASSERT(num_rdma_bytes < std::numeric_limits::max()); + EP_HOST_ASSERT(num_nvl_bytes < std::numeric_limits::max()); + + // Launch kernel + SETUP_LAUNCH_CONFIG(1 + num_rdma_ranks, kNumThreads, stream); + SWITCH_RDMA_RANKS(NOTIFY_DISPATCH_LAUNCH_CASE); +#undef NOTIFY_DISPATCH_LAUNCH_CASE +} + +// At most 8 RDMA ranks to be sent +constexpr int get_num_topk_rdma_ranks(int num_rdma_ranks) { + return num_rdma_ranks < 8 ? num_rdma_ranks : 8; +} + +template +__global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NVL_PEERS) * 32), 1) + dispatch(int4* recv_x, + float* recv_x_scales, + topk_idx_t* recv_topk_idx, + float* recv_topk_weights, + SourceMeta* recv_src_meta, + const int4* x, + const float* x_scales, + const topk_idx_t* topk_idx, + const float* topk_weights, + int* send_rdma_head, + int* send_nvl_head, + int* recv_rdma_channel_prefix_matrix, + int* recv_gbl_channel_prefix_matrix, + const int* rdma_channel_prefix_matrix, + const int* recv_rdma_rank_prefix_sum, + const int* gbl_channel_prefix_matrix, + const int* recv_gbl_rank_prefix_sum, + const bool* is_token_in_rank, + int num_tokens, + int hidden_int4, + int num_scales, + int num_topk, + int num_experts, + int scale_token_stride, + int scale_hidden_stride, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_send_tokens, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_send_tokens, + int num_max_nvl_chunked_recv_tokens, + int rank, + int num_ranks, + nixl_ep::gpu_nixl_ctx nixl_ctx) { + enum class WarpRole { kRDMASender, kRDMASenderCoordinator, kRDMAAndNVLForwarder, kForwarderCoordinator, kNVLReceivers }; + + const auto num_sms = static_cast(gridDim.x); + const auto sm_id = static_cast(blockIdx.x); + const auto num_threads = static_cast(blockDim.x), num_warps = num_threads / 32; + const auto thread_id = static_cast(threadIdx.x), warp_id = thread_id / 32, lane_id = get_lane_id(); + const auto num_channels = num_sms / 2, channel_id = sm_id / 2; + const bool is_forwarder = sm_id % 2 == 0; + const auto rdma_rank = rank / NUM_MAX_NVL_PEERS, nvl_rank = rank % NUM_MAX_NVL_PEERS; + + if(thread_id ==0) { + DEVICE_LOG_DEBUG("dispatch rank: %d, sm_id: %d, num_channels: %d, channel_id: %d, warp_id: %d, lane_id: %d, is_forwarder: %d, rdma_rank: %d, nvl_rank: %d, send_nvl_head: %p, send_rdma_head: %p", rank, sm_id, num_channels, channel_id, warp_id, lane_id, is_forwarder, rdma_rank, nvl_rank, send_nvl_head, send_rdma_head); + } + + const auto role_meta = [=]() -> std::pair { + if (is_forwarder) { + if (warp_id < NUM_MAX_NVL_PEERS) { + return {WarpRole::kRDMAAndNVLForwarder, (warp_id + channel_id) % NUM_MAX_NVL_PEERS}; + } else { + return {WarpRole::kForwarderCoordinator, warp_id - NUM_MAX_NVL_PEERS}; + } + } else if (warp_id < kNumDispatchRDMASenderWarps) { + return {WarpRole::kRDMASender, -1}; + } else if (warp_id == kNumDispatchRDMASenderWarps) { + return {WarpRole::kRDMASenderCoordinator, -1}; + } else { + return {WarpRole::kNVLReceivers, (warp_id + channel_id - kNumDispatchRDMASenderWarps) % NUM_MAX_NVL_PEERS}; + } + }(); + auto warp_role = role_meta.first; + auto target_rank = role_meta.second; // Not applicable for RDMA senders + EP_DEVICE_ASSERT(num_warps == kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NVL_PEERS); + + // Data checks + EP_DEVICE_ASSERT(num_topk <= 32); + + // RDMA symmetric layout + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * sizeof(bool) == sizeof(uint64_t), "Invalid number of NVL peers"); + auto hidden_bytes = hidden_int4 * sizeof(int4); + auto scale_bytes = num_scales * sizeof(float); + auto num_bytes_per_token = get_num_bytes_per_token(hidden_int4, num_scales, num_topk, num_topk); + auto rdma_channel_data = SymBuffer( + rdma_buffer_ptr, num_max_rdma_chunked_recv_tokens * num_bytes_per_token, kNumRDMARanks, channel_id, num_channels); + auto rdma_channel_meta = SymBuffer(rdma_buffer_ptr, NUM_MAX_NVL_PEERS * 2 + 2, kNumRDMARanks, channel_id, num_channels); + auto rdma_channel_head = SymBuffer(rdma_buffer_ptr, 1, kNumRDMARanks, channel_id, num_channels); + auto rdma_channel_tail = SymBuffer(rdma_buffer_ptr, 1, kNumRDMARanks, channel_id, num_channels); + + // NVL buffer layouts + // NOTES: `rs_wr_buffer_ptr` means "Read for Senders, Write for Receivers", `ws_rr_buffer_ptr` means "Write for Senders, Read for + // Receivers" + void *rs_wr_buffer_ptr = nullptr, *ws_rr_buffer_ptr = nullptr; + int rs_wr_rank = 0, ws_rr_rank = 0; + if (warp_role == WarpRole::kRDMAAndNVLForwarder) + rs_wr_buffer_ptr = buffer_ptrs[nvl_rank], ws_rr_buffer_ptr = buffer_ptrs[target_rank], rs_wr_rank = nvl_rank, + ws_rr_rank = target_rank; + if (warp_role == WarpRole::kNVLReceivers) + rs_wr_buffer_ptr = buffer_ptrs[target_rank], ws_rr_buffer_ptr = buffer_ptrs[nvl_rank], rs_wr_rank = target_rank, + ws_rr_rank = nvl_rank; + + // Allocate buffers + auto nvl_channel_x = AsymBuffer(ws_rr_buffer_ptr, + num_max_nvl_chunked_recv_tokens * num_bytes_per_token, + NUM_MAX_NVL_PEERS, + channel_id, + num_channels, + rs_wr_rank) + .advance_also(rs_wr_buffer_ptr); + auto nvl_channel_prefix_start = + AsymBuffer(ws_rr_buffer_ptr, kNumRDMARanks, NUM_MAX_NVL_PEERS, channel_id, num_channels, rs_wr_rank) + .advance_also(rs_wr_buffer_ptr); + auto nvl_channel_prefix_end = AsymBuffer(ws_rr_buffer_ptr, kNumRDMARanks, NUM_MAX_NVL_PEERS, channel_id, num_channels, rs_wr_rank) + .advance_also(rs_wr_buffer_ptr); + auto nvl_channel_head = + AsymBuffer(rs_wr_buffer_ptr, 1, NUM_MAX_NVL_PEERS, channel_id, num_channels, ws_rr_rank).advance_also(ws_rr_buffer_ptr); + auto nvl_channel_tail = + AsymBuffer(ws_rr_buffer_ptr, 1, NUM_MAX_NVL_PEERS, channel_id, num_channels, rs_wr_rank).advance_also(rs_wr_buffer_ptr); + + // RDMA sender warp synchronization + // NOTES: `rdma_send_channel_tail` means the latest released tail + // NOTES: `rdma_send_channel_window` means the ongoing 32 transactions' status + __shared__ int rdma_send_channel_lock[kNumRDMARanks]; + __shared__ int rdma_send_channel_tail[kNumRDMARanks]; + __shared__ uint32_t rdma_send_channel_window[kNumRDMARanks]; + auto sync_rdma_sender_smem = []() { asm volatile("barrier.sync 0, %0;" ::"r"((kNumDispatchRDMASenderWarps + 1) * 32)); }; + + // TMA stuffs + extern __shared__ __align__(1024) uint8_t smem_tma_buffer[]; + auto tma_buffer = smem_tma_buffer + target_rank * kNumTMABytesPerWarp; + auto tma_mbarrier = reinterpret_cast(tma_buffer + num_bytes_per_token); + uint32_t tma_phase = 0; + if ((warp_role == WarpRole::kRDMAAndNVLForwarder or warp_role == WarpRole::kNVLReceivers) and elect_one_sync()) { + mbarrier_init(tma_mbarrier, 1); + fence_barrier_init(); + EP_DEVICE_ASSERT(num_bytes_per_token + sizeof(uint64_t) <= kNumTMABytesPerWarp); + } + __syncwarp(); + + // Forward warp synchronization + __shared__ volatile int forward_channel_head[NUM_MAX_NVL_PEERS][kNumRDMARanks]; + __shared__ volatile bool forward_channel_retired[NUM_MAX_NVL_PEERS]; + auto sync_forwarder_smem = []() { asm volatile("barrier.sync 1, %0;" ::"r"((NUM_MAX_NVL_PEERS + 1) * 32)); }; + + if (warp_role == WarpRole::kRDMASender) { + // Get tasks + int token_start_idx, token_end_idx; + get_channel_task_range(num_tokens, num_channels, channel_id, token_start_idx, token_end_idx); + + // Send number of tokens in this channel by `-value - 1` + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * 2 + 2 <= 32, "Invalid number of NVL peers"); + for (int dst_rdma_rank = warp_id; dst_rdma_rank < kNumRDMARanks; dst_rdma_rank += kNumDispatchRDMASenderWarps) { + auto dst_ptr = + dst_rdma_rank == rdma_rank ? rdma_channel_meta.recv_buffer(dst_rdma_rank) : rdma_channel_meta.send_buffer(dst_rdma_rank); + if (lane_id < NUM_MAX_NVL_PEERS) { + dst_ptr[lane_id] = + -(channel_id == 0 + ? 0 + : gbl_channel_prefix_matrix[(dst_rdma_rank * NUM_MAX_NVL_PEERS + lane_id) * num_channels + channel_id - 1]) - + 1; + } else if (lane_id < NUM_MAX_NVL_PEERS * 2) { + dst_ptr[lane_id] = + -gbl_channel_prefix_matrix[(dst_rdma_rank * NUM_MAX_NVL_PEERS + lane_id - NUM_MAX_NVL_PEERS) * num_channels + + channel_id] - + 1; + } else if (lane_id == NUM_MAX_NVL_PEERS * 2) { + dst_ptr[lane_id] = -(channel_id == 0 ? 0 : rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id - 1]) - 1; + } else if (lane_id == NUM_MAX_NVL_PEERS * 2 + 1) { + dst_ptr[lane_id] = -rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id] - 1; +#ifdef ENABLE_DEBUG_LOGS + if (channel_id != 0) { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | dst_rdma_rank: %d, sending %d tokens over channel %d ", rank, warp_id, channel_id, dst_rdma_rank, rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id] - rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id - 1] , channel_id); + } else { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | dst_rdma_rank: %d, sending %d tokens over channel %d ", rank, warp_id, channel_id, dst_rdma_rank, rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id] , channel_id); + } +#endif + } + __syncwarp(); + + // Issue RDMA for non-local ranks + if (dst_rdma_rank != rdma_rank) { + size_t src_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_channel_meta.send_buffer(dst_rdma_rank))); + size_t dst_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_channel_meta.recv_buffer(rdma_rank))); + size_t msg_size = sizeof(int) * (NUM_MAX_NVL_PEERS * 2 + 2); + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d channel %d | dst_rdma_rank: %d| RDMA SENDER ISSUE RDMA", rank, warp_id, channel_id, dst_rdma_rank); + int translated_rank = translate_dst_rdma_rank(dst_rdma_rank, nvl_rank); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translated_rank); + EP_DEVICE_ASSERT(nixlGpuPostSingleWriteXferReq( + batch_req, 0, src_offset, dst_offset, msg_size, channel_id, true) == + NIXL_IN_PROG); + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d channel %d | dst_rdma_rank: %d| RDMA SENDER ISSUE RDMA DONE", rank, warp_id, channel_id, dst_rdma_rank); + } + } + sync_rdma_sender_smem(); + + // Iterate over tokens and copy into buffer + int64_t token_idx; + int cached_rdma_channel_head = 0, global_rdma_tail_idx = 0; + auto send_buffer = lane_id == rdma_rank ? rdma_channel_data.recv_buffer(lane_id) : rdma_channel_data.send_buffer(lane_id); + for (token_idx = token_start_idx; token_idx < token_end_idx; ++token_idx) { + // Read RDMA rank existence + uint64_t is_token_in_rank_uint64 = 0; + if (lane_id < kNumRDMARanks) { + is_token_in_rank_uint64 = + __ldg(reinterpret_cast(is_token_in_rank + token_idx * num_ranks + lane_id * NUM_MAX_NVL_PEERS)); + global_rdma_tail_idx += (is_token_in_rank_uint64 != 0); + } + __syncwarp(); + + // Skip the token which does not belong to this warp + if ((token_idx - token_start_idx) % kNumDispatchRDMASenderWarps != warp_id) + continue; + auto rdma_tail_idx = is_token_in_rank_uint64 == 0 ? -1 : global_rdma_tail_idx - 1; + + // Wait the remote buffer to be released + auto start_time = clock64(); +#ifdef ENABLE_DEBUG_LOGS + int remote_head_poll_counter = 1; + int last_cached_rdma_channel_head = cached_rdma_channel_head; +#endif + while (is_token_in_rank_uint64 != 0 and rdma_tail_idx - cached_rdma_channel_head >= num_max_rdma_chunked_recv_tokens) { + cached_rdma_channel_head = static_cast(ld_volatile_global(rdma_channel_head.buffer(lane_id))); +#ifdef ENABLE_DEBUG_LOGS + if (last_cached_rdma_channel_head != cached_rdma_channel_head) { + DEVICE_LOG_DEBUG_LANE(lane_id, "[DISPATCH DEBUG] rank %d warp %d channel %d | RDMA SENDER HEAD UPDATED | lane_id: %d | cached_rdma_channel_head: %d | last_cached_rdma_channel_head: %d | head pointer: %p", rank, warp_id, channel_id, lane_id, cached_rdma_channel_head, last_cached_rdma_channel_head, (void*)rdma_channel_head.buffer(lane_id)); + last_cached_rdma_channel_head = cached_rdma_channel_head; + } + + if (remote_head_poll_counter % 50000000 == 0) { + DEVICE_LOG_DEBUG_LANE(lane_id, "[DISPATCH DEBUG] rank %d warp %d channel %d | RDMA SENDER - polling head update| polling counter: %d | lane_id: %d | cached_rdma_channel_head: %d | rdma_tail_idx: %d | head pointer: %p", rank, warp_id, channel_id, remote_head_poll_counter, lane_id, cached_rdma_channel_head, rdma_tail_idx, (void*)rdma_channel_head.buffer(lane_id)); + } + remote_head_poll_counter++; +#endif + + // Timeout check + if (clock64() - start_time >= NUM_TIMEOUT_CYCLES) { + printf("NixlEP dispatch RDMA sender timeout, channel: %d, RDMA: %d, nvl: %d, dst RDMA lane: %d, head: %d, tail: %d\n", + channel_id, + rdma_rank, + nvl_rank, + lane_id, + cached_rdma_channel_head, + rdma_tail_idx); + trap(); + } + } + __syncwarp(); + + // Store RDMA head for combine + if (lane_id < kNumRDMARanks and not kCachedMode) + send_rdma_head[token_idx * kNumRDMARanks + lane_id] = rdma_tail_idx; + + // Broadcast tails + SourceMeta src_meta; + int num_topk_ranks = 0, topk_ranks[kNumTopkRDMARanks]; + void* dst_send_buffers[kNumTopkRDMARanks]; + #pragma unroll + for (int i = 0, slot_idx; i < kNumRDMARanks; ++i) + if ((slot_idx = __shfl_sync(0xffffffff, rdma_tail_idx, i)) >= 0) { + slot_idx = slot_idx % num_max_rdma_chunked_recv_tokens; + topk_ranks[num_topk_ranks] = i; + auto recv_is_token_in_rank_uint64 = broadcast(is_token_in_rank_uint64, i); + auto recv_is_token_in_rank_values = reinterpret_cast(&recv_is_token_in_rank_uint64); + if (lane_id == num_topk_ranks) + src_meta = SourceMeta(rdma_rank, recv_is_token_in_rank_values); + dst_send_buffers[num_topk_ranks++] = + reinterpret_cast(broadcast(send_buffer, i)) + slot_idx * num_bytes_per_token; + } + EP_DEVICE_ASSERT(num_topk_ranks <= kNumTopkRDMARanks); + + // Copy `x` into symmetric send buffer + auto st_broadcast = [=](const int key, const int4& value) { + #pragma unroll + for (int j = 0; j < num_topk_ranks; ++j) + st_na_global(reinterpret_cast(dst_send_buffers[j]) + key, value); + }; + UNROLLED_WARP_COPY(5, lane_id, hidden_int4, 0, x + token_idx * hidden_int4, ld_nc_global, st_broadcast); + #pragma unroll + for (int i = 0; i < num_topk_ranks; ++i) + dst_send_buffers[i] = reinterpret_cast(dst_send_buffers[i]) + hidden_int4; + + // Copy `x_scales` into symmetric send buffer + #pragma unroll + for (int i = lane_id; i < num_scales; i += 32) { + auto offset = token_idx * scale_token_stride + i * scale_hidden_stride; + auto value = ld_nc_global(x_scales + offset); + #pragma unroll + for (int j = 0; j < num_topk_ranks; ++j) + st_na_global(reinterpret_cast(dst_send_buffers[j]) + i, value); + } + #pragma unroll + for (int i = 0; i < num_topk_ranks; ++i) + dst_send_buffers[i] = reinterpret_cast(dst_send_buffers[i]) + num_scales; + + // Copy source metadata into symmetric send buffer + if (lane_id < num_topk_ranks) + st_na_global(reinterpret_cast(dst_send_buffers[lane_id]), src_meta); + #pragma unroll + for (int i = 0; i < num_topk_ranks; ++i) + dst_send_buffers[i] = reinterpret_cast(dst_send_buffers[i]) + 1; + + // Copy `topk_idx` and `topk_weights` into symmetric send buffer + #pragma unroll + for (int i = lane_id; i < num_topk * num_topk_ranks; i += 32) { + auto rank_idx = i / num_topk, copy_idx = i % num_topk; + auto idx_value = static_cast(ld_nc_global(topk_idx + token_idx * num_topk + copy_idx)); + auto weight_value = ld_nc_global(topk_weights + token_idx * num_topk + copy_idx); + st_na_global(reinterpret_cast(dst_send_buffers[rank_idx]) + copy_idx, idx_value); + st_na_global(reinterpret_cast(dst_send_buffers[rank_idx]) + num_topk + copy_idx, weight_value); + } + __syncwarp(); + + // Release the transaction in the window + if (is_token_in_rank_uint64 != 0) { + // Acquire lock first + acquire_lock(rdma_send_channel_lock + lane_id); + auto latest_tail = rdma_send_channel_tail[lane_id]; + auto offset = rdma_tail_idx - latest_tail; + while (offset >= 32) { + release_lock(rdma_send_channel_lock + lane_id); + acquire_lock(rdma_send_channel_lock + lane_id); + latest_tail = rdma_send_channel_tail[lane_id]; + offset = rdma_tail_idx - latest_tail; + } + + // Release the transaction slot + // Add the bit and move the ones if possible + auto window = rdma_send_channel_window[lane_id] | (1u << offset); + if (offset == 0) { + auto num_empty_slots = (~window) == 0 ? 32 : __ffs(~window) - 1; + st_release_cta(rdma_send_channel_tail + lane_id, latest_tail + num_empty_slots); + window >>= num_empty_slots; + } + rdma_send_channel_window[lane_id] = window; + + // Release lock + release_lock(rdma_send_channel_lock + lane_id); + } + __syncwarp(); + } + DEVICE_LOG_DEBUG_LANE(0,"rank %d warp %d | channel %d | RDMA SENDER DONE |token %ld| global_rdma_tail_idx: %d", rank, warp_id, channel_id, token_idx, global_rdma_tail_idx); + } else if (warp_role == WarpRole::kRDMASenderCoordinator) { + // NOTES: in case of splitting, the issued put at the end of the buffer + EP_DEVICE_ASSERT(num_max_rdma_chunked_recv_tokens % num_max_rdma_chunked_send_tokens == 0); + + // Clean shared memory + EP_STATIC_ASSERT(kNumRDMARanks <= 32, "Invalid number of RDMA ranks"); + (lane_id < kNumRDMARanks) ? (rdma_send_channel_lock[lane_id] = 0) : 0; + (lane_id < kNumRDMARanks) ? (rdma_send_channel_tail[lane_id] = 0) : 0; + (lane_id < kNumRDMARanks) ? (rdma_send_channel_window[lane_id] = 0) : 0; + + // Synchronize shared memory + sync_rdma_sender_smem(); + + // Get number of tokens to send for each RDMA rank + int num_tokens_to_send = 0; + if (lane_id < kNumRDMARanks) { + num_tokens_to_send = rdma_channel_prefix_matrix[lane_id * num_channels + channel_id]; + if (channel_id > 0) + num_tokens_to_send -= rdma_channel_prefix_matrix[lane_id * num_channels + channel_id - 1]; + } + + // Iterate all RDMA ranks + int last_issued_tail = 0; + auto start_time = clock64(); +#ifdef ENABLE_DEBUG_LOGS + bool approaching_timeout_synced_tokens_printed=false; + bool approaching_timeout_num_tokens_processed_printed=false; +#endif + while (__any_sync(0xffffffff, num_tokens_to_send > 0)) { + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES and lane_id < kNumRDMARanks) { + printf("NixlEP RDMA sender coordinator timeout, channel: %d, IB: %d, nvl %d, dst IB: %d, tail: %d, remaining: %d\n", + channel_id, + rdma_rank, + nvl_rank, + lane_id, + last_issued_tail, + num_tokens_to_send); + trap(); + } + + // TODO: try thread-level `put_nbi`? + for (int i = 0, synced_num_tokens_to_send; i < kNumRDMARanks; ++i) { + // To mitigate incast congestion, shuffle the starting index of target rank for different ranks and channels + int dst_rdma_rank = (i + channel_id + rdma_rank) % kNumRDMARanks; + synced_num_tokens_to_send = __shfl_sync(0xffffffff, num_tokens_to_send, dst_rdma_rank); +#ifdef ENABLE_DEBUG_LOGS + if (clock64() - start_time > NUM_TIMEOUT_CYCLES/2 && !approaching_timeout_synced_tokens_printed) { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | RDMA SENDER COORDINATOR - approaching timeout | dst_rdma_rank: %d | synced_num_tokens_to_send: %d", rank, warp_id, channel_id, dst_rdma_rank, synced_num_tokens_to_send); + approaching_timeout_synced_tokens_printed = true; + } + __syncwarp();//for the debug print +#endif + if (synced_num_tokens_to_send == 0) + continue; + + // Read the latest progress + // NOTES: `rdma_send_channel_tail` does not need to be protected by lock + auto processed_tail = + __shfl_sync(0xffffffff, ld_acquire_cta(const_cast(rdma_send_channel_tail + dst_rdma_rank)), 0); + auto synced_last_issued_tail = __shfl_sync(0xffffffff, last_issued_tail, dst_rdma_rank); + auto num_tokens_processed = processed_tail - synced_last_issued_tail; +#ifdef ENABLE_DEBUG_LOGS + if (clock64() - start_time > NUM_TIMEOUT_CYCLES/2 && !approaching_timeout_num_tokens_processed_printed) { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | RDMA SENDER COORDINATOR - approaching timeout | dst_rdma_rank: %d | num_tokens_processed: %d |synced_num_tokens_to_send: %d |num_max_rdma_chunked_send_tokens: %d", rank, warp_id, channel_id, dst_rdma_rank, num_tokens_processed, synced_num_tokens_to_send, num_max_rdma_chunked_send_tokens); + approaching_timeout_num_tokens_processed_printed = true; + } + __syncwarp();//for the debug print +#endif + if (num_tokens_processed != synced_num_tokens_to_send and num_tokens_processed < num_max_rdma_chunked_send_tokens) + continue; + + // Issue RDMA send + auto num_tokens_to_issue = min(num_tokens_processed, num_max_rdma_chunked_send_tokens); + EP_DEVICE_ASSERT(num_tokens_to_issue >= 0 and num_tokens_to_issue <= synced_num_tokens_to_send); + if (dst_rdma_rank != rdma_rank) { + auto dst_slot_idx = synced_last_issued_tail % num_max_rdma_chunked_recv_tokens; + EP_DEVICE_ASSERT(dst_slot_idx + num_tokens_to_issue <= num_max_rdma_chunked_recv_tokens); + const size_t num_bytes_per_msg = num_bytes_per_token * num_tokens_to_issue; + const auto dst_ptr = + reinterpret_cast(rdma_channel_data.recv_buffer(rdma_rank) + dst_slot_idx * num_bytes_per_token); + const auto src_ptr = + reinterpret_cast(rdma_channel_data.send_buffer(dst_rdma_rank) + dst_slot_idx * num_bytes_per_token); + size_t src_offset = nixl_ctx.batch_offset_get(src_ptr); + size_t dst_offset = nixl_ctx.batch_offset_get(dst_ptr); + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d channel %d | RDMA SENDER COORDINATOR | dst_rdma_rank: %d | SENDING num_tokens_to_issue: %d |last_issued_tail: %d", rank, warp_id, channel_id, dst_rdma_rank, num_tokens_to_issue, synced_last_issued_tail); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(dst_rdma_rank, nvl_rank)); + + EP_DEVICE_ASSERT(nixlGpuPostSingleWriteXferReq( + batch_req, 0, src_offset, dst_offset, num_bytes_per_msg, channel_id, true) == + NIXL_IN_PROG); + // Increment tail counter on the receiver + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d channel %d | RDMA SENDER COORDINATOR - done | dst_rdma_rank: %d | SENDING num_tokens_to_issue: %d |last_issued_tail: %d", rank, warp_id, channel_id, dst_rdma_rank, num_tokens_to_issue, synced_last_issued_tail); + } else { + // Lighter fence for local RDMA rank + memory_fence(); + } + __syncwarp(); + + // Update tails + if (lane_id == dst_rdma_rank) { + last_issued_tail += num_tokens_to_issue; + num_tokens_to_send -= num_tokens_to_issue; + if (dst_rdma_rank == rdma_rank) { + DEVICE_LOG_DEBUG_LANE(0, "rank %d warp %d | channel %d | dst rdma rank %d | RDMA SENDER COORDINATOR | TAIL UPDATED (local), num_tokens_to_issue: %d, num_tokens_to_send: %d | last_issued_tail: %d | local counter pointer: %p", rank, warp_id, channel_id, dst_rdma_rank, num_tokens_to_issue, num_tokens_to_send, last_issued_tail, (void*)rdma_channel_tail.buffer(dst_rdma_rank)); + atomicAdd(reinterpret_cast(rdma_channel_tail.buffer(dst_rdma_rank)), static_cast(num_tokens_to_issue)); + DEVICE_LOG_DEBUG_LANE(0, "rank %d warp %d | channel %d | dst rdma rank %d | RDMA SENDER COORDINATOR | TAIL UPDATED (local) - done, num_tokens_to_issue: %d, num_tokens_to_send: %d | last_issued_tail: %d | local counter pointer: %p", rank, warp_id, channel_id, dst_rdma_rank, num_tokens_to_issue, num_tokens_to_send, last_issued_tail, (void*)rdma_channel_tail.buffer(dst_rdma_rank)); + } else { + size_t tail_counter_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_channel_tail.buffer(rdma_rank))); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(dst_rdma_rank, nvl_rank)); + EP_DEVICE_ASSERT(nixlGpuPostSignalXferReq( + batch_req, 0, num_tokens_to_issue, tail_counter_offset, channel_id, true) == + NIXL_IN_PROG); + // TODO_Roey: is it commented out because the RDMA WRITE xfer already includes the tail update? + // DEVICE_LOG_DEBUG_LANE(0, "[DEBUG] rank %d warp %d | channel %d | dst rdma rank %d | RDMA SENDER COORDINATOR | TAIL UPDATED (remote) - signal, num_tokens_to_issue: %d, num_tokens_to_send: %d | last_issued_tail: %d", rank, warp_id, channel_id, dst_rdma_rank, num_tokens_to_issue, num_tokens_to_send, last_issued_tail); + // nixlPostPartialGpuXferReq(channel_data_requests_handles[dst_rdma_rank], num_tokens_to_issue, 0, nullptr, nullptr, nullptr, nullptr, nullptr); + } + } + __syncwarp(); + } + } +#ifdef ENABLE_DEBUG_LOGS + if (lane_id < kNumRDMARanks) { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | RDMA SENDER COORDINATOR FINISHED | rank: %d, last_issued_tail: %d", rank, warp_id, channel_id, lane_id, last_issued_tail); + } +#endif + } else if (warp_role == WarpRole::kRDMAAndNVLForwarder) { + // RDMA consumers and NVL producers + const auto dst_nvl_rank = target_rank; + + // Wait counters to arrive + int num_tokens_to_recv_from_rdma = 0, src_rdma_channel_prefix = 0; + EP_DEVICE_ASSERT(kNumRDMARanks <= 32); + auto start_time = clock64(); + if (lane_id < kNumRDMARanks) { +#ifdef ENABLE_DEBUG_LOGS + int print_initial_num_tokens = 0; +#endif + while (true) { + auto meta_0 = ld_volatile_global(rdma_channel_meta.recv_buffer(lane_id) + dst_nvl_rank); + auto meta_1 = ld_volatile_global(rdma_channel_meta.recv_buffer(lane_id) + NUM_MAX_NVL_PEERS + dst_nvl_rank); + auto meta_2 = ld_volatile_global(rdma_channel_meta.recv_buffer(lane_id) + NUM_MAX_NVL_PEERS * 2); + auto meta_3 = ld_volatile_global(rdma_channel_meta.recv_buffer(lane_id) + NUM_MAX_NVL_PEERS * 2 + 1); + if (meta_0 < 0 and meta_1 < 0 and meta_2 < 0 and meta_3 < 0) { + // Notify NVL ranks + int start_sum = -meta_0 - 1, end_sum = -meta_1 - 1; + EP_DEVICE_ASSERT(start_sum >= 0 and end_sum >= 0 and end_sum >= start_sum); + st_relaxed_sys_global(nvl_channel_prefix_start.buffer() + lane_id, -start_sum - 1); + st_relaxed_sys_global(nvl_channel_prefix_end.buffer() + lane_id, -end_sum - 1); + + // Save RDMA channel received token count + src_rdma_channel_prefix = -meta_2 - 1; + auto src_rdma_channel_prefix_1 = -meta_3 - 1; + num_tokens_to_recv_from_rdma = src_rdma_channel_prefix_1 - src_rdma_channel_prefix; +#ifdef ENABLE_DEBUG_LOGS + if (print_initial_num_tokens == 0) { + DEVICE_LOG_DEBUG_LANE(lane_id, "rank %d warp %d channel %d | recevie from rdma rank %d initial num_tokens_to_recv_from_rdma: %d over channel %d", rank, warp_id, channel_id, lane_id, num_tokens_to_recv_from_rdma, channel_id); + print_initial_num_tokens = 1; + } +#endif + if (not kCachedMode) + recv_rdma_channel_prefix_matrix[lane_id * num_channels + channel_id] = src_rdma_channel_prefix_1; + src_rdma_channel_prefix += lane_id == 0 ? 0 : recv_rdma_rank_prefix_sum[lane_id - 1]; + EP_DEVICE_ASSERT(num_tokens_to_recv_from_rdma >= 0); + break; + } + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES) { + printf( + "NixlEP dispatch forwarder timeout (RDMA meta), channel: %d, RDMA: %d, nvl: %d, src RDMA lane: %d, dst NVL: %d, meta: %d, %d, %d, %d\n", + channel_id, + rdma_rank, + nvl_rank, + lane_id, + dst_nvl_rank, + meta_0, + meta_1, + meta_2, + meta_3); + trap(); + } + } + } + __syncwarp(); + + // Shift cached head + send_nvl_head += src_rdma_channel_prefix * NUM_MAX_NVL_PEERS + dst_nvl_rank; + + // Wait shared memory to be cleaned + sync_forwarder_smem(); + + // Forward tokens from RDMA buffer + // NOTES: always start from the local rank + int src_rdma_rank = sm_id % kNumRDMARanks; + int cached_rdma_channel_head = 0, cached_rdma_channel_tail = 0; + int cached_nvl_channel_head = 0, cached_nvl_channel_tail = 0, rdma_nvl_token_idx = 0; + while (__any_sync(0xffffffff, num_tokens_to_recv_from_rdma > 0)) { + // Check destination queue emptiness, or wait a buffer to be released + start_time = clock64(); + while (true) { + const int num_used_slots = cached_nvl_channel_tail - cached_nvl_channel_head; + if (num_max_nvl_chunked_recv_tokens - num_used_slots >= num_max_nvl_chunked_send_tokens) + break; + cached_nvl_channel_head = __shfl_sync(0xffffffffu, ld_volatile_global(nvl_channel_head.buffer()), 0); + + // Timeout check + if (elect_one_sync() and clock64() - start_time > NUM_TIMEOUT_CYCLES) { + printf( + "NixlEP dispatch forwarder timeout (NVL check), channel: %d, RDMA: %d, nvl: %d, dst NVL: %d, head: %d, tail: %d\n", + channel_id, + rdma_rank, + nvl_rank, + dst_nvl_rank, + ld_volatile_global(nvl_channel_head.buffer()), + cached_nvl_channel_tail); + trap(); + } + } + + // Find next source RDMA rank (round-robin) + start_time = clock64(); +#ifdef ENABLE_DEBUG_LOGS + int polling_counter = 1; +#endif + while (true) { + src_rdma_rank = (src_rdma_rank + 1) % kNumRDMARanks; + if (__shfl_sync(0xffffffff, num_tokens_to_recv_from_rdma, src_rdma_rank) > 0) { + if (lane_id == src_rdma_rank and cached_rdma_channel_head == cached_rdma_channel_tail) + { + cached_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(src_rdma_rank))); +#ifdef ENABLE_DEBUG_LOGS + if( num_tokens_to_recv_from_rdma > 0 && polling_counter % 50000000 == 0) {//we can print only from one warp + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | poll counter %d |src_rdma_rank: %d, cached_rdma_channel_tail: %d, src_rdma_rank tail pointer: %p num_tokens_to_recv_from_rdma: %d",\ + rank, warp_id, channel_id, polling_counter, src_rdma_rank, cached_rdma_channel_tail, (void*)rdma_channel_tail.buffer(src_rdma_rank), num_tokens_to_recv_from_rdma); + } + polling_counter++; +#endif + } + if (__shfl_sync(0xffffffff, cached_rdma_channel_tail > cached_rdma_channel_head, src_rdma_rank)) + { + DEVICE_LOG_DEBUG_LANE(src_rdma_rank, "rank %d warp %d channel %d src_rdma_rank %d | RDMA CHANNEL TAIL UPDATED | num_tokens_to_recv_from_rdma: %d cached_rdma_channel_tail: %d cached_rdma_channel_head: %d tail pointer: %p", rank, warp_id, channel_id, src_rdma_rank, num_tokens_to_recv_from_rdma, cached_rdma_channel_tail, cached_rdma_channel_head, (void*)rdma_channel_tail.buffer(src_rdma_rank)); + break; + } + } + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES and lane_id < kNumRDMARanks) { + printf( + "NixlEP dispatch forwarder timeout (RDMA check), channel: %d, RDMA: %d, nvl: %d, dst NVL: %d, src RDMA lane: %d, " + "head: %d, tail: %d, expected: %d\n", + channel_id, + rdma_rank, + nvl_rank, + dst_nvl_rank, + lane_id, + cached_rdma_channel_head, + cached_rdma_channel_tail, + num_tokens_to_recv_from_rdma); + trap(); + } + } + auto src_rdma_head = __shfl_sync(0xffffffff, cached_rdma_channel_head, src_rdma_rank); + auto src_rdma_tail = __shfl_sync(0xffffffff, cached_rdma_channel_tail, src_rdma_rank); + + // Iterate over every token from the RDMA buffer + for (int i = src_rdma_head, num_tokens_sent = 0; i < src_rdma_tail; ++i) { + auto rdma_slot_idx = i % num_max_rdma_chunked_recv_tokens; + auto shifted = rdma_channel_data.recv_buffer(src_rdma_rank) + rdma_slot_idx * num_bytes_per_token; + auto src_meta = ld_nc_global(reinterpret_cast(shifted + hidden_bytes + scale_bytes)); + lane_id == src_rdma_rank ? (num_tokens_to_recv_from_rdma -= 1) : 0; + DEVICE_LOG_DEBUG_LANE_SYNC(src_rdma_rank, "rank %d warp %d channel %d | RDMA RECEIVED TOKEN | src_rdma_rank: %d | num_tokens_to_recv_from_rdma: %d, src_meta: %d", rank, warp_id, channel_id, src_rdma_rank, num_tokens_to_recv_from_rdma, src_meta.is_token_in_nvl_rank(dst_nvl_rank)); + bool is_in_dst_nvl_rank = src_meta.is_token_in_nvl_rank(dst_nvl_rank); + if (lane_id == src_rdma_rank) { + auto cached_head = is_in_dst_nvl_rank ? rdma_nvl_token_idx : -1; + rdma_nvl_token_idx += is_in_dst_nvl_rank; + if (not kCachedMode) + send_nvl_head[i * NUM_MAX_NVL_PEERS] = cached_head; + } + if (not is_in_dst_nvl_rank) + continue; + + // Get an empty slot + int dst_slot_idx = (cached_nvl_channel_tail++) % num_max_nvl_chunked_recv_tokens; + auto dst_shifted = nvl_channel_x.buffer() + dst_slot_idx * num_bytes_per_token; + + // Copy data + if (elect_one_sync()) { + tma_load_1d(tma_buffer, shifted, tma_mbarrier, num_bytes_per_token, false); + mbarrier_arrive_and_expect_tx(tma_mbarrier, num_bytes_per_token); + } + __syncwarp(); + mbarrier_wait(tma_mbarrier, tma_phase); + if (elect_one_sync()) + tma_store_1d(tma_buffer, dst_shifted, num_bytes_per_token); + __syncwarp(); + + // In case of insufficient NVL buffers, early stopping + if ((++num_tokens_sent) == num_max_nvl_chunked_send_tokens) + src_rdma_tail = i + 1; + + // Wait TMA to be finished + tma_store_wait<0>(); + __syncwarp(); + } + + // Sync head index + if (lane_id == src_rdma_rank) + { + forward_channel_head[dst_nvl_rank][src_rdma_rank] = (cached_rdma_channel_head = src_rdma_tail); + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | dst nvl rank %d | NVL HEAD SYNC | forward_channel_head[%d][%d]: %d", rank, warp_id, channel_id, dst_nvl_rank, dst_nvl_rank, src_rdma_rank, forward_channel_head[dst_nvl_rank][src_rdma_rank]); + } + + // Move tail index + __syncwarp(); + if (elect_one_sync()) + st_release_sys_global(nvl_channel_tail.buffer(), cached_nvl_channel_tail); + } + + // Retired + __syncwarp(); + if (elect_one_sync()) + { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | dst nvl rank %d | NVL RETIRED | forward_channel_retired[%d]: %d", rank, warp_id, channel_id, dst_nvl_rank, dst_nvl_rank, forward_channel_retired[dst_nvl_rank]); + forward_channel_retired[dst_nvl_rank] = true; + } +#ifdef ENABLE_DEBUG_LOGS + if (lane_id < kNumRDMARanks) { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | NVL RETIRED | last cached_rdma_channel_tail: %d", rank, warp_id, channel_id, cached_rdma_channel_tail); + } +#endif + } else if (warp_role == WarpRole::kForwarderCoordinator) { + // Extra warps for forwarder coordinator should exit directly + if (target_rank > 0) + return; + + // Forward warp coordinator + EP_STATIC_ASSERT(kNumRDMARanks <= 32, "Invalid number of RDMA peers"); + + // Clean shared memory + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS <= 32, "Invalid number of NVL peers"); + #pragma unroll + for (int i = lane_id; i < kNumRDMARanks * NUM_MAX_NVL_PEERS; i += 32) + forward_channel_head[i % NUM_MAX_NVL_PEERS][i / NUM_MAX_NVL_PEERS] = 0; + if (lane_id < NUM_MAX_NVL_PEERS) + forward_channel_retired[lane_id] = false; + sync_forwarder_smem(); + + int last_head = 0, target_rdma = lane_id < kNumRDMARanks ? lane_id : 0; +#ifdef ENABLE_DEBUG_LOGS + int forward_head_poll_counter = 1; + int timeout_check_counter = 0; + auto start_time = clock64(); +#endif + while (true) { + // Find minimum head + int min_head = std::numeric_limits::max(); + #pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) + if (not forward_channel_retired[i]) + min_head = min(min_head, forward_channel_head[i][target_rdma]); +#ifdef ENABLE_DEBUG_LOGS + if (clock64() - start_time > NUM_TIMEOUT_CYCLES/2 && timeout_check_counter%100000 == 0) { + for(int i = 0; i < NUM_MAX_NVL_PEERS; i++) { + DEVICE_LOG_DEBUG("rank %d warp %d | channel %d | lane_id: %d | KForwarderCoordinator - waiting for min head update | retired: %d | forward_channel_head[%d][%d]: %d", rank, warp_id, channel_id, lane_id, forward_channel_retired[i], i, target_rdma, ld_acquire_cta(&forward_channel_head[i][target_rdma])); + } + DEVICE_LOG_DEBUG("rank %d warp %d | channel %d | KForwarderCoordinator - waiting for min head update | forward head poll counter: %d | min_head: %d | lane_id: %d", rank, warp_id, channel_id, forward_head_poll_counter, min_head, lane_id); timeout_check_counter++; + } + forward_head_poll_counter++; + __syncwarp(); +#endif + if (__all_sync(0xffffffff, min_head == std::numeric_limits::max())) + break; + + // Update remote head + if (min_head != std::numeric_limits::max() and min_head >= last_head + num_max_rdma_chunked_send_tokens and + lane_id < kNumRDMARanks) { + if(lane_id == rdma_rank){ + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | KForwarderCoordinator | LOCAL HEAD UPDATE | head change: %d | new_head: %lu| local counter pointer: %p", rank, warp_id, channel_id, min_head - last_head, ld_acquire_sys_global(rdma_channel_head.buffer(lane_id)), (void*)rdma_channel_head.buffer(lane_id)); + atomicAdd(reinterpret_cast(rdma_channel_head.buffer(rdma_rank)), static_cast(min_head - last_head)); + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | KForwarderCoordinator | LOCAL HEAD UPDATE - done | head change: %d | new_head: %lu| local counter pointer: %p", rank, warp_id, channel_id, min_head - last_head, ld_acquire_sys_global(rdma_channel_head.buffer(lane_id)), (void*)rdma_channel_head.buffer(lane_id)); + }else{ + size_t head_counter_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_channel_head.buffer(rdma_rank))); + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | KForwarderCoordinator REMOTE HEAD UPDATE | dst_rank: %d| head change: %d| signal_off: %lu", rank, warp_id, channel_id, lane_id, min_head - last_head, head_counter_offset); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(lane_id, nvl_rank)); + nixl_status_t status = nixlGpuPostSignalXferReq( + batch_req, 0, min_head - last_head, head_counter_offset); + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | KForwarderCoordinator REMOTE HEAD UPDATE - posted request | dst_rank: %d| head change: %d", rank, warp_id, channel_id, lane_id, min_head - last_head); + if (status != NIXL_IN_PROG) { + DEVICE_LOG_DEBUG("NixlEP dispatch forwarder (RDMA increment) failed, channel: %d, RDMA: %d, nvl: %d, dst RDMA: %d, head: %d, last: %d, chunked: %d", + channel_id, rdma_rank, nvl_rank, lane_id, min_head, last_head, num_max_rdma_chunked_send_tokens); + trap(); + } + } + last_head = min_head; + } + + // Nanosleep and let other warps work + __nanosleep(NUM_WAIT_NANOSECONDS); + } + DEVICE_LOG_DEBUG("rank %d warp %d channel %d lane_id: %d | KForwarderCoordinator - FINISHED", rank, warp_id, channel_id, lane_id); + } else { + // NVL consumers + // Retrieve rank offset from barrier results (each lane's register stores an RDMA rank) + int src_nvl_rank = target_rank, total_offset = 0; + const int local_expert_begin = rank * (num_experts / num_ranks); + const int local_expert_end = local_expert_begin + (num_experts / num_ranks); + + EP_STATIC_ASSERT(kNumRDMARanks <= 32, "Invalid number of RDMA peers"); + if (lane_id < kNumRDMARanks and lane_id * NUM_MAX_NVL_PEERS + src_nvl_rank > 0) + total_offset = recv_gbl_rank_prefix_sum[lane_id * NUM_MAX_NVL_PEERS + src_nvl_rank - 1]; + + // Receive channel offsets + int start_offset = 0, end_offset = 0, num_tokens_to_recv; + auto start_time = clock64(); + while (lane_id < kNumRDMARanks) { + start_offset = ld_volatile_global(nvl_channel_prefix_start.buffer() + lane_id); + end_offset = ld_volatile_global(nvl_channel_prefix_end.buffer() + lane_id); + if (start_offset < 0 and end_offset < 0) { + start_offset = -start_offset - 1, end_offset = -end_offset - 1; + total_offset += start_offset; + break; + } + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES) { + printf( + "NixlEP dispatch NVL receiver timeout, channel: %d, RDMA: %d, nvl: %d, src RDMA: %d, src nvl: %d, start: %d, end: %d\n", + channel_id, + rdma_rank, + nvl_rank, + lane_id, + src_nvl_rank, + start_offset, + end_offset); + trap(); + } + } + num_tokens_to_recv = warp_reduce_sum(end_offset - start_offset); + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d | channel %d | NVL RECEIVER | num_tokens_to_recv: %d", rank, warp_id, channel_id, num_tokens_to_recv); + + // Save for combine usage + if (lane_id < kNumRDMARanks and not kCachedMode) + recv_gbl_channel_prefix_matrix[(lane_id * NUM_MAX_NVL_PEERS + src_nvl_rank) * num_channels + channel_id] = total_offset; + __syncwarp(); + + int cached_channel_head_idx = 0, cached_channel_tail_idx = 0; + while (num_tokens_to_recv > 0) { + // Check channel status by lane 0 + start_time = clock64(); + while (true) { + // Ready to copy + if (cached_channel_head_idx != cached_channel_tail_idx) + break; + cached_channel_tail_idx = __shfl_sync(0xffffffff, ld_acquire_sys_global(nvl_channel_tail.buffer()), 0); + + // Timeout check + if (elect_one_sync() and clock64() - start_time > NUM_TIMEOUT_CYCLES) { + printf("NixlEP dispatch NVL receiver timeout, channel: %d, RDMA: %d, nvl: %d, src NVL: %d, head: %d, tail: %d\n", + channel_id, + rdma_rank, + nvl_rank, + src_nvl_rank, + cached_channel_head_idx, + cached_channel_tail_idx); + trap(); + } + } + + // Copy data + int num_recv_tokens = cached_channel_tail_idx - cached_channel_head_idx; + for (int chunk_idx = 0; chunk_idx < num_recv_tokens; ++chunk_idx, --num_tokens_to_recv) { + int token_idx_in_buffer = (cached_channel_head_idx++) % num_max_nvl_chunked_recv_tokens; + auto shifted = nvl_channel_x.buffer() + token_idx_in_buffer * num_bytes_per_token; + auto meta = ld_nc_global(reinterpret_cast(shifted + hidden_bytes + scale_bytes)); + int64_t recv_token_idx = __shfl_sync(0xffffffff, total_offset, meta.src_rdma_rank); + (lane_id == meta.src_rdma_rank) ? (total_offset += 1) : 0; + + bool scale_aligned = (scale_bytes % 16 == 0); + auto tma_load_bytes = hidden_bytes + (scale_aligned ? scale_bytes : 0); + + // Copy data + if (elect_one_sync()) { + tma_load_1d(tma_buffer, shifted, tma_mbarrier, tma_load_bytes); + mbarrier_arrive_and_expect_tx(tma_mbarrier, tma_load_bytes); + } + __syncwarp(); + mbarrier_wait(tma_mbarrier, tma_phase); + if (elect_one_sync()) { + tma_store_1d(tma_buffer, recv_x + recv_token_idx * hidden_int4, hidden_bytes, false); + if (scale_aligned) + tma_store_1d(tma_buffer + hidden_bytes, recv_x_scales + recv_token_idx * num_scales, scale_bytes, false); + } + __syncwarp(); + shifted += hidden_bytes; + + // Copy scales + // TODO: make it as templated + if (not scale_aligned) { + UNROLLED_WARP_COPY(1, + lane_id, + num_scales, + recv_x_scales + recv_token_idx * num_scales, + reinterpret_cast(shifted), + ld_nc_global, + st_na_global); + } + shifted += scale_bytes; + + // Copy source meta + if (not kCachedMode and elect_one_sync()) + st_na_global(recv_src_meta + recv_token_idx, meta); + shifted += sizeof(SourceMeta); + + // Copy `topk_idx` and `topk_weights` + if (lane_id < num_topk) { + // Read + auto idx_value = static_cast(ld_nc_global(reinterpret_cast(shifted) + lane_id)); + auto weight_value = ld_nc_global(reinterpret_cast(shifted + sizeof(int) * num_topk) + lane_id); + auto recv_idx = recv_token_idx * num_topk + lane_id; + + // Transform and write + idx_value = (idx_value >= local_expert_begin and idx_value < local_expert_end) ? idx_value - local_expert_begin : -1; + weight_value = idx_value >= 0 ? weight_value : 0.0f; + st_na_global(recv_topk_idx + recv_idx, idx_value); + st_na_global(recv_topk_weights + recv_idx, weight_value); + } + + // Wait TMA to be finished + tma_store_wait<0>(); + __syncwarp(); + } + + // Move queue + if (elect_one_sync()) + st_relaxed_sys_global(nvl_channel_head.buffer(), cached_channel_head_idx); + } + } +} +void dispatch(void* recv_x, + float* recv_x_scales, + topk_idx_t* recv_topk_idx, + float* recv_topk_weights, + void* recv_src_meta, + const void* x, + const float* x_scales, + const topk_idx_t* topk_idx, + const float* topk_weights, + int* send_rdma_head, + int* send_nvl_head, + int* recv_rdma_channel_prefix_matrix, + int* recv_gbl_channel_prefix_matrix, + const int* rdma_channel_prefix_matrix, + const int* recv_rdma_rank_prefix_sum, + const int* gbl_channel_prefix_matrix, + const int* recv_gbl_rank_prefix_sum, + const bool* is_token_in_rank, + int num_tokens, + int hidden_int4, + int num_scales, + int num_topk, + int num_experts, + int scale_token_stride, + int scale_hidden_stride, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_send_tokens, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_send_tokens, + int num_max_nvl_chunked_recv_tokens, + int rank, + int num_ranks, + bool is_cached_dispatch, + cudaStream_t stream, + int num_channels, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx) { + constexpr int kNumDispatchRDMASenderWarps = 7; + constexpr int kNumTMABytesPerWarp = 16384; + constexpr int smem_size = kNumTMABytesPerWarp * NUM_MAX_NVL_PEERS; + + // Make sure never OOB + EP_HOST_ASSERT(static_cast(num_scales) * scale_hidden_stride < std::numeric_limits::max()); + +#define DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ + { \ + auto dispatch_func = low_latency_mode \ + ? (is_cached_dispatch ? dispatch \ + : dispatch) \ + : (is_cached_dispatch ? dispatch \ + : dispatch); \ + SET_SHARED_MEMORY_FOR_TMA(dispatch_func); \ + LAUNCH_KERNEL(&cfg, \ + dispatch_func, \ + reinterpret_cast(recv_x), \ + recv_x_scales, \ + recv_topk_idx, \ + recv_topk_weights, \ + reinterpret_cast(recv_src_meta), \ + reinterpret_cast(x), \ + x_scales, \ + topk_idx, \ + topk_weights, \ + send_rdma_head, \ + send_nvl_head, \ + recv_rdma_channel_prefix_matrix, \ + recv_gbl_channel_prefix_matrix, \ + rdma_channel_prefix_matrix, \ + recv_rdma_rank_prefix_sum, \ + gbl_channel_prefix_matrix, \ + recv_gbl_rank_prefix_sum, \ + is_token_in_rank, \ + num_tokens, \ + hidden_int4, \ + num_scales, \ + num_topk, \ + num_experts, \ + scale_token_stride, \ + scale_hidden_stride, \ + rdma_buffer_ptr, \ + num_max_rdma_chunked_send_tokens, \ + num_max_rdma_chunked_recv_tokens, \ + buffer_ptrs, \ + num_max_nvl_chunked_send_tokens, \ + num_max_nvl_chunked_recv_tokens, \ + rank, \ + num_ranks, \ + nixl_ctx); \ + } \ + break + + EP_HOST_ASSERT((topk_idx == nullptr) == (topk_weights == nullptr)); + EP_HOST_ASSERT((recv_topk_idx == nullptr) == (recv_topk_weights == nullptr)); + + SETUP_LAUNCH_CONFIG(num_channels * 2, (kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NVL_PEERS) * 32, stream); + SWITCH_RDMA_RANKS(DISPATCH_LAUNCH_CASE); +#undef DISPATCH_LAUNCH_CASE +} + +template +__global__ void cached_notify(const int rdma_clean_offset, + const int rdma_num_int_clean, + const int nvl_clean_offset, + const int nvl_num_int_clean, + int* combined_rdma_head, + int num_combined_tokens, + int num_channels, + const int* rdma_channel_prefix_matrix, + const int* rdma_rank_prefix_sum, + int* combined_nvl_head, + void* rdma_buffer_ptr, + void** buffer_ptrs, + int** barrier_signal_ptrs, + int rank, + int num_ranks, + bool is_cached_dispatch, + gpu_nixl_ctx nixl_ctx) { + auto sm_id = static_cast(blockIdx.x); + auto thread_id = static_cast(threadIdx.x); + auto num_threads = static_cast(blockDim.x); + auto num_warps = num_threads / 32; + auto warp_id = thread_id / 32; + auto lane_id = get_lane_id(); + + auto nvl_rank = rank % NUM_MAX_NVL_PEERS; + auto num_rdma_ranks = num_ranks / NUM_MAX_NVL_PEERS; + + // Using two SMs, which clean the RDMA/NVL buffer respectively + if (sm_id == 0) { + __syncthreads(); + + // Barrier for RDMA + if (thread_id == 32) { + DEVICE_LOG_DEBUG("rank %d warp %d | CACHED NOTIFY | RDMA BARRIER | num_rdma_ranks: %d, num_channels: %d, rank: %d, rdma_rank: %d", rank, warp_id, num_rdma_ranks, num_channels, rank, rank / NUM_MAX_NVL_PEERS); + nixl_barrier(nixl_ctx, num_channels); + } + + // Barrier for NVL + barrier_block(barrier_signal_ptrs, nvl_rank); + + // Clean RDMA buffer + auto rdma_buffer_ptr_int = static_cast(rdma_buffer_ptr); + #pragma unroll + for (int i = thread_id; i < rdma_num_int_clean; i += num_threads) + rdma_buffer_ptr_int[rdma_clean_offset + i] = 0; + + for (int i = 0; i < num_channels; ++ i) { + auto rdma_channel_head = SymBuffer(rdma_buffer_ptr, 1, num_rdma_ranks, i, num_channels); + auto rdma_channel_tail = SymBuffer(rdma_buffer_ptr, 1, num_rdma_ranks, i, num_channels); + if (thread_id < num_rdma_ranks) { + rdma_channel_head.buffer()[thread_id] = 0; + rdma_channel_tail.buffer()[thread_id] = 0; + } + } + // Clean NVL buffer + auto nvl_buffer_ptr_int = static_cast(buffer_ptrs[nvl_rank]); + #pragma unroll + for (int i = thread_id; i < nvl_num_int_clean; i += num_threads) + nvl_buffer_ptr_int[nvl_clean_offset + i] = 0; + __syncthreads(); + + // Barrier again + if (thread_id == 32) + nixl_barrier(nixl_ctx, num_channels); + barrier_block(barrier_signal_ptrs, nvl_rank); + } else if (sm_id == 1) { + if (is_cached_dispatch) + return; + + EP_DEVICE_ASSERT(num_warps >= num_channels); + EP_DEVICE_ASSERT(num_rdma_ranks <= 32); + + // Iterate in reverse order + if (lane_id < num_rdma_ranks and warp_id < num_channels) { + int token_start_idx, token_end_idx; + get_channel_task_range(num_combined_tokens, num_channels, warp_id, token_start_idx, token_end_idx); + + // NOTES: `1 << 25` is a heuristic large number + int last_head = 1 << 25; + for (int token_idx = token_end_idx - 1; token_idx >= token_start_idx; --token_idx) { + auto current_head = __ldg(combined_rdma_head + token_idx * num_rdma_ranks + lane_id); + if (current_head < 0) { + combined_rdma_head[token_idx * num_rdma_ranks + lane_id] = -last_head - 1; + } else { + last_head = current_head; + } + } + } + } else { + if (is_cached_dispatch) + return; + + EP_DEVICE_ASSERT(num_warps >= num_channels); + EP_DEVICE_ASSERT(rdma_channel_prefix_matrix != nullptr and rdma_rank_prefix_sum != nullptr); + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS <= 32, "Too many NVL peers"); + + if (warp_id < num_channels) { + constexpr int tma_batch_size = kNumTMABytesPerWarp - sizeof(uint64_t); + constexpr int num_bytes_per_token = sizeof(int) * NUM_MAX_NVL_PEERS; + constexpr int num_tokens_per_batch = tma_batch_size / num_bytes_per_token; + EP_STATIC_ASSERT(num_bytes_per_token % 16 == 0, "num_bytes_per_token should be divisible by 16"); + + // TMA stuffs + extern __shared__ __align__(1024) uint8_t smem_tma_buffer[]; + auto tma_buffer = smem_tma_buffer + warp_id * kNumTMABytesPerWarp; + auto tma_mbarrier = reinterpret_cast(tma_buffer + tma_batch_size); + uint32_t tma_phase = 0; + if (elect_one_sync()) { + mbarrier_init(tma_mbarrier, 1); + fence_barrier_init(); + } + __syncwarp(); + + for (int dst_rdma_rank = sm_id - 2; dst_rdma_rank < num_rdma_ranks; dst_rdma_rank += num_channels * 2 - 2) { + // Iterate in reverse order + int token_start_idx = warp_id == 0 ? 0 : rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + warp_id - 1]; + int token_end_idx = rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + warp_id]; + int shift = dst_rdma_rank == 0 ? 0 : rdma_rank_prefix_sum[dst_rdma_rank - 1]; + token_start_idx += shift, token_end_idx += shift; + + // NOTES: `1 << 25` is a heuristic large number + int last_head = 1 << 25; + for (int batch_end_idx = token_end_idx; batch_end_idx > token_start_idx; batch_end_idx -= num_tokens_per_batch) { + auto batch_start_idx = max(token_start_idx, batch_end_idx - num_tokens_per_batch); + + if (elect_one_sync()) { + tma_load_1d(tma_buffer, + combined_nvl_head + batch_start_idx * NUM_MAX_NVL_PEERS, + tma_mbarrier, + (batch_end_idx - batch_start_idx) * num_bytes_per_token); + mbarrier_arrive_and_expect_tx(tma_mbarrier, (batch_end_idx - batch_start_idx) * num_bytes_per_token); + } + mbarrier_wait(tma_mbarrier, tma_phase); + __syncwarp(); + + for (int token_idx = batch_end_idx - 1; token_idx >= batch_start_idx; --token_idx) { + if (lane_id < NUM_MAX_NVL_PEERS) { + auto current_head = + reinterpret_cast(tma_buffer)[(token_idx - batch_start_idx) * NUM_MAX_NVL_PEERS + lane_id]; + if (current_head < 0) { + reinterpret_cast(tma_buffer)[(token_idx - batch_start_idx) * NUM_MAX_NVL_PEERS + lane_id] = + -last_head - 1; + } else { + last_head = current_head; + } + } + } + tma_store_fence(); + __syncwarp(); + + if (elect_one_sync()) + tma_store_1d(tma_buffer, + combined_nvl_head + batch_start_idx * NUM_MAX_NVL_PEERS, + (batch_end_idx - batch_start_idx) * num_bytes_per_token); + tma_store_wait<0>(); + __syncwarp(); + } + } + } + } +} + +void cached_notify(int hidden_int4, + int num_scales, + int num_topk_idx, + int num_topk_weights, + int num_ranks, + int num_channels, + int num_combined_tokens, + int* combined_rdma_head, + const int* rdma_channel_prefix_matrix, + const int* rdma_rank_prefix_sum, + int* combined_nvl_head, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_recv_tokens, + int** barrier_signal_ptrs, + int rank, + cudaStream_t stream, + int64_t num_rdma_bytes, + int64_t num_nvl_bytes, + bool is_cached_dispatch, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx) { + const int num_threads = std::max(128, 32 * num_channels); + const int num_warps = num_threads / 32; + const auto num_rdma_ranks = num_ranks / NUM_MAX_NVL_PEERS; + const int kNumTMABytesPerWarp = 8192; + const int smem_size = kNumTMABytesPerWarp * num_warps; + + // Get clean meta + auto rdma_clean_meta = get_rdma_clean_meta( + hidden_int4, num_scales, num_topk_idx, num_topk_weights, num_rdma_ranks, num_max_rdma_chunked_recv_tokens, num_channels); + auto nvl_clean_meta = get_nvl_clean_meta(hidden_int4, + num_scales, + num_topk_idx, + num_topk_weights, + num_rdma_ranks, + NUM_MAX_NVL_PEERS, + num_max_nvl_chunked_recv_tokens, + num_channels, + is_cached_dispatch); + EP_HOST_ASSERT((rdma_clean_meta.first + rdma_clean_meta.second) * sizeof(int) <= num_rdma_bytes); + EP_HOST_ASSERT((nvl_clean_meta.first + nvl_clean_meta.second) * sizeof(int) <= num_nvl_bytes); + EP_HOST_ASSERT(num_rdma_bytes < std::numeric_limits::max()); + EP_HOST_ASSERT(num_nvl_bytes < std::numeric_limits::max()); + EP_HOST_ASSERT(num_channels * 2 > 3); + + // Launch kernel + auto cached_notify_func = low_latency_mode ? cached_notify : cached_notify; + SETUP_LAUNCH_CONFIG(num_channels * 2, num_threads, stream); + SET_SHARED_MEMORY_FOR_TMA(cached_notify_func); + LAUNCH_KERNEL(&cfg, + cached_notify_func, + rdma_clean_meta.first, + rdma_clean_meta.second, + nvl_clean_meta.first, + nvl_clean_meta.second, + combined_rdma_head, + num_combined_tokens, + num_channels, + rdma_channel_prefix_matrix, + rdma_rank_prefix_sum, + combined_nvl_head, + rdma_buffer_ptr, + buffer_ptrs, + barrier_signal_ptrs, + rank, + num_ranks, + is_cached_dispatch, + nixl_ctx); +} + +template +__device__ int combine_token(bool is_token_in_rank, + int head_idx, + int lane_id, + int hidden_int4, + int num_topk, + int4* combined_row, + float* combined_topk_weights, + const int4* bias_0_int4, + const int4* bias_1_int4, + int num_max_recv_tokens, + const GetAddrFn& get_addr_fn, + const ReceiveTWFn& recv_tw_fn, + uint8_t* smem_ptr, + uint32_t (&tma_phase)[kNumStages]) { + constexpr auto kDtypePerInt4 = sizeof(int4) / sizeof(dtype_t); + + // Broadcast current heads + // Lane `i` holds the head of rank `i` and `is_token_in_rank` + EP_STATIC_ASSERT(kMaxNumRanks <= 32, "Too many ranks"); + int num_topk_ranks = 0, topk_ranks[kMaxNumRanks], slot_indices[kMaxNumRanks]; + #pragma unroll + for (int i = 0; i < kNumRanks; ++i) + if (__shfl_sync(0xffffffff, is_token_in_rank, i)) { + slot_indices[num_topk_ranks] = __shfl_sync(0xffffffff, head_idx, i) % num_max_recv_tokens; + topk_ranks[num_topk_ranks++] = i; + } + EP_DEVICE_ASSERT(num_topk_ranks <= kMaxNumRanks); + EP_STATIC_ASSERT(not(kUseTMA and kMaybeWithBias), "TMA cannot be used by receiver warps"); + EP_STATIC_ASSERT(kNumStages == 2, "Only support 2 stages now"); + + // Reduce data + if constexpr (kUseTMA) { + constexpr int kNumTMABufferBytesPerStage = kNumTMALoadBytes * (NUM_MAX_NVL_PEERS + 1) + 16; + EP_DEVICE_ASSERT(hidden_int4 % 32 == 0); + + auto tma_load_buffer = [=](const int& i, const int& j) -> int4* { + return reinterpret_cast(smem_ptr + i * kNumTMABufferBytesPerStage + j * kNumTMALoadBytes); + }; + auto tma_store_buffer = [=](const int& i) -> int4* { + return reinterpret_cast(smem_ptr + i * kNumTMABufferBytesPerStage + NUM_MAX_NVL_PEERS * kNumTMALoadBytes); + }; + auto tma_mbarrier = [=](const int& i) -> uint64_t* { + return reinterpret_cast(smem_ptr + i * kNumTMABufferBytesPerStage + (NUM_MAX_NVL_PEERS + 1) * kNumTMALoadBytes); + }; + + // Prefetch + if (lane_id < num_topk_ranks) + tma_load_1d( + tma_load_buffer(0, lane_id), get_addr_fn(topk_ranks[lane_id], slot_indices[lane_id], 0), tma_mbarrier(0), kNumTMALoadBytes); + mbarrier_arrive_and_expect_tx(tma_mbarrier(0), lane_id < num_topk_ranks ? kNumTMALoadBytes : 0); + __syncwarp(); + + for (int shifted = 0, iter = 0; shifted < hidden_int4; shifted += 32, iter += 1) { + const int stage_idx = iter % kNumStages; + const int next_stage_idx = (iter + 1) % kNumStages; + + // Prefetch next stage + if (shifted + 32 < hidden_int4) { + if (lane_id < num_topk_ranks) + tma_load_1d(tma_load_buffer(next_stage_idx, lane_id), + get_addr_fn(topk_ranks[lane_id], slot_indices[lane_id], shifted + 32), + tma_mbarrier(next_stage_idx), + kNumTMALoadBytes); + mbarrier_arrive_and_expect_tx(tma_mbarrier(next_stage_idx), lane_id < num_topk_ranks ? kNumTMALoadBytes : 0); + __syncwarp(); + } + + mbarrier_wait(tma_mbarrier(stage_idx), tma_phase[stage_idx]); + float values[kDtypePerInt4] = {0}; + #pragma unroll + for (int j = 0; j < num_topk_ranks; ++j) { + auto recv_value_dtypes = reinterpret_cast(tma_load_buffer(stage_idx, j) + lane_id); + #pragma unroll + for (int k = 0; k < kDtypePerInt4; ++k) + values[k] += static_cast(recv_value_dtypes[k]); + } + + // Wait shared memory to be released + tma_store_wait(); + + // Copy into shared and issue TMA + auto out_dtypes = reinterpret_cast(tma_store_buffer(stage_idx) + lane_id); + #pragma unroll + for (int j = 0; j < kDtypePerInt4; ++j) + out_dtypes[j] = static_cast(values[j]); + tma_store_fence(); + __syncwarp(); + + if (elect_one_sync()) + tma_store_1d(tma_store_buffer(stage_idx), combined_row + shifted, kNumTMALoadBytes); + __syncwarp(); + } + + // Flush all writes + tma_store_wait<0>(); + } else { + #pragma unroll + for (int i = lane_id; i < hidden_int4; i += 32) { + // Read bias + // TODO: make it as a finer-grained template + int4 bias_0_value_int4, bias_1_value_int4; + if constexpr (kMaybeWithBias) { + bias_0_value_int4 = bias_0_int4 != nullptr ? ld_nc_global(bias_0_int4 + i) : make_int4(0, 0, 0, 0); + bias_1_value_int4 = bias_1_int4 != nullptr ? ld_nc_global(bias_1_int4 + i) : make_int4(0, 0, 0, 0); + } + + // Read buffers + // TODO: maybe too many registers here + int4 recv_value_int4[kMaxNumRanks]; + #pragma unroll + for (int j = 0; j < num_topk_ranks; ++j) + recv_value_int4[j] = ld_nc_global(get_addr_fn(topk_ranks[j], slot_indices[j], i)); + + // Clean + // Reduce bias + float values[kDtypePerInt4] = {0}; + if constexpr (kMaybeWithBias) { + auto bias_0_values = reinterpret_cast(&bias_0_value_int4); + auto bias_1_values = reinterpret_cast(&bias_1_value_int4); + #pragma unroll + for (int j = 0; j < kDtypePerInt4; ++j) + values[j] = static_cast(bias_0_values[j]) + static_cast(bias_1_values[j]); + } + + // Reduce all-to-all results + #pragma unroll + for (int j = 0; j < num_topk_ranks; ++j) { + auto recv_value_dtypes = reinterpret_cast(&recv_value_int4[j]); + #pragma unroll + for (int k = 0; k < kDtypePerInt4; ++k) + values[k] += static_cast(recv_value_dtypes[k]); + } + + // Cast back to `dtype_t` and write + int4 out_int4; + auto out_dtypes = reinterpret_cast(&out_int4); + #pragma unroll + for (int j = 0; j < kDtypePerInt4; ++j) + out_dtypes[j] = static_cast(values[j]); + st_na_global(combined_row + i, out_int4); + } + } + + // Reduce `topk_weights` + if (lane_id < num_topk) { + float value = 0; + #pragma unroll + for (int i = 0; i < num_topk_ranks; ++i) + value += recv_tw_fn(topk_ranks[i], slot_indices[i], lane_id); + st_na_global(combined_topk_weights + lane_id, value); + } + + // Return the minimum top-k rank + return topk_ranks[0]; +} + +template 0) ? kNumCombineForwarderWarps / kNumRDMARanks : 1, + int kNumForwarders = kNumRDMARanks * kNumWarpsPerForwarder, + int kNumRDMAReceivers = kNumForwarders - NUM_MAX_NVL_PEERS> +__global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* combined_x, + float* combined_topk_weights, + const bool* is_combined_token_in_rank, + const int4* x, + const float* topk_weights, + const int4* bias_0, + const int4* bias_1, + const int* combined_rdma_head, + const int* combined_nvl_head, + const SourceMeta* src_meta, + const int* rdma_channel_prefix_matrix, + const int* rdma_rank_prefix_sum, + const int* gbl_channel_prefix_matrix, + int num_tokens, + int num_combined_tokens, + int hidden, + int num_topk, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_send_tokens, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_send_tokens, + int num_max_nvl_chunked_recv_tokens, + int rank, + int num_ranks, + gpu_nixl_ctx nixl_ctx) { + enum class WarpRole { kNVLSender, kNVLAndRDMAForwarder, kRDMAReceiver, kCoordinator }; + + const auto sm_id = static_cast(blockIdx.x); + const auto num_threads = static_cast(blockDim.x), num_warps = num_threads / 32; + const auto thread_id = static_cast(threadIdx.x), lane_id = get_lane_id(); + const auto num_channels = static_cast(gridDim.x) / 2, channel_id = sm_id / 2; + const bool is_forwarder_sm = sm_id % 2 == 1; + + EP_DEVICE_ASSERT(num_topk <= 32); + EP_DEVICE_ASSERT(hidden % (sizeof(int4) / sizeof(dtype_t)) == 0); + const auto hidden_int4 = hidden / (sizeof(int4) / sizeof(dtype_t)); + const auto hidden_bytes = hidden_int4 * sizeof(int4); + const auto num_bytes_per_token = get_num_bytes_per_token(hidden_int4, 0, 0, num_topk); + + // NOTES: we decouple a channel into 2 SMs + const auto rdma_rank = rank / NUM_MAX_NVL_PEERS, nvl_rank = rank % NUM_MAX_NVL_PEERS; + auto role_meta = [=]() -> std::pair { + auto warp_id = thread_id / 32; + if (not is_forwarder_sm) { + if (warp_id < NUM_MAX_NVL_PEERS) { + auto shuffled_warp_id = warp_id; + shuffled_warp_id = (shuffled_warp_id + channel_id) % NUM_MAX_NVL_PEERS; + return {WarpRole::kNVLSender, shuffled_warp_id}; + } else if (warp_id < kNumForwarders) { + return {WarpRole::kRDMAReceiver, warp_id - NUM_MAX_NVL_PEERS}; + } else { + return {WarpRole::kCoordinator, 0}; + } + } else { + if (warp_id < kNumForwarders) { + auto shuffled_warp_id = (warp_id + channel_id) % kNumForwarders; + return {WarpRole::kNVLAndRDMAForwarder, shuffled_warp_id}; + } else { + return {WarpRole::kCoordinator, 0}; + } + } + }(); + auto warp_role = role_meta.first; + auto warp_id = role_meta.second; + + EP_DEVICE_ASSERT(num_warps == kNumForwarders + 1); + auto num_max_nvl_chunked_recv_tokens_per_rdma = num_max_nvl_chunked_recv_tokens / kNumRDMARanks; + + if (warp_role == WarpRole::kNVLSender) { + // NVL producers + const auto dst_nvl_rank = warp_id; + + // NVL layouts + // NOTES: to avoid deadlocks, we use separate NVL buffers for different RDMA sources + auto dst_buffer_ptr = buffer_ptrs[dst_nvl_rank], local_buffer_ptr = buffer_ptrs[nvl_rank]; + auto nvl_channel_x = AsymBuffer(dst_buffer_ptr, + num_max_nvl_chunked_recv_tokens * num_bytes_per_token, + NUM_MAX_NVL_PEERS, + channel_id, + num_channels, + nvl_rank) + .advance_also(local_buffer_ptr); + auto nvl_channel_head = AsymBuffer(local_buffer_ptr, kNumRDMARanks, NUM_MAX_NVL_PEERS, channel_id, num_channels, dst_nvl_rank) + .advance_also(dst_buffer_ptr); + auto nvl_channel_tail = AsymBuffer(dst_buffer_ptr, kNumRDMARanks, NUM_MAX_NVL_PEERS, channel_id, num_channels, nvl_rank) + .advance_also(local_buffer_ptr); + + // TMA stuffs + extern __shared__ __align__(1024) uint8_t smem_tma_buffer[]; + auto tma_buffer = smem_tma_buffer + dst_nvl_rank * kNumTMABytesPerSenderWarp; + auto tma_mbarrier = reinterpret_cast(tma_buffer + num_bytes_per_token); + uint32_t tma_phase = 0; + if (elect_one_sync()) { + mbarrier_init(tma_mbarrier, 1); + fence_barrier_init(); + EP_DEVICE_ASSERT(num_bytes_per_token + sizeof(uint64_t) <= kNumTMABytesPerSenderWarp); + } + __syncwarp(); + + // Get tasks for each RDMA lane + int token_start_idx = 0, token_end_idx = 0; + if (lane_id < kNumRDMARanks) { + int prefix_idx = (lane_id * NUM_MAX_NVL_PEERS + dst_nvl_rank) * num_channels + channel_id; + token_start_idx = gbl_channel_prefix_matrix[prefix_idx]; + token_end_idx = (prefix_idx == num_channels * num_ranks - 1) ? num_tokens : gbl_channel_prefix_matrix[prefix_idx + 1]; + } + __syncwarp(); + + // NOTES: here the cached value of each lane is only responsible for a single RDMA buffer + int cached_channel_head_idx = 0, cached_channel_tail_idx = 0; + EP_STATIC_ASSERT(kNumRDMARanks <= 32, "Invalid number of RDMA peers"); + + // Iterate over all tokens and send by chunks + int current_rdma_idx = channel_id % kNumRDMARanks; + while (true) { + // Exit if possible + if (__all_sync(0xffffffff, token_start_idx >= token_end_idx)) + break; + + // Decide the next RDMA buffer to send + bool is_lane_ready = false; + auto start_time = clock64(); + while (true) { + int num_used_slots = cached_channel_tail_idx - cached_channel_head_idx; + is_lane_ready = lane_id < kNumRDMARanks and token_start_idx < token_end_idx and + num_max_nvl_chunked_recv_tokens_per_rdma - num_used_slots >= num_max_nvl_chunked_send_tokens; + if (__any_sync(0xffffffff, is_lane_ready)) + break; + + // Retry + if (lane_id < kNumRDMARanks and token_start_idx < token_end_idx) + cached_channel_head_idx = ld_volatile_global(nvl_channel_head.buffer() + lane_id); + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES and lane_id < kNumRDMARanks) { + printf( + "NixlEP combine NVL sender timeout, channel: %d, RDMA: %d, nvl: %d, dst NVL: %d, RDMA lane: %d, head: %d, tail: " + "%d, start: %d, end: %d\n", + channel_id, + rdma_rank, + nvl_rank, + dst_nvl_rank, + lane_id, + ld_volatile_global(nvl_channel_head.buffer() + lane_id), + cached_channel_tail_idx, + token_start_idx, + token_end_idx); + trap(); + } + } + + // Sync token start index and count + for (int i = 0; i < kNumRDMARanks; ++i) { + current_rdma_idx = (current_rdma_idx + 1) % kNumRDMARanks; + if (__shfl_sync(0xffffffff, (token_start_idx >= token_end_idx) or (not is_lane_ready), current_rdma_idx)) + continue; + + // Sync token start index + auto token_idx = static_cast(__shfl_sync(0xffffffff, token_start_idx, current_rdma_idx)); + int num_tokens_in_chunk = + __shfl_sync(0xffffffff, min(num_max_nvl_chunked_send_tokens, token_end_idx - token_start_idx), current_rdma_idx); + + // Send by chunk + for (int chunk_idx = 0; chunk_idx < num_tokens_in_chunk; ++chunk_idx, ++token_idx) { + // Get an empty slot + int dst_slot_idx = 0; + if (lane_id == current_rdma_idx) { + dst_slot_idx = (cached_channel_tail_idx++) % num_max_nvl_chunked_recv_tokens_per_rdma; + dst_slot_idx = current_rdma_idx * num_max_nvl_chunked_recv_tokens_per_rdma + dst_slot_idx; + } + dst_slot_idx = __shfl_sync(0xffffffff, dst_slot_idx, current_rdma_idx); + + // Load data + auto shifted_x_buffers = nvl_channel_x.buffer() + dst_slot_idx * num_bytes_per_token; + auto shifted_x = x + token_idx * hidden_int4; + tma_store_wait<0>(); + if (elect_one_sync()) { + tma_load_1d(tma_buffer, shifted_x, tma_mbarrier, hidden_bytes); + mbarrier_arrive_and_expect_tx(tma_mbarrier, hidden_bytes); + } + __syncwarp(); + mbarrier_wait(tma_mbarrier, tma_phase); + + // Load source meta + if (lane_id == num_topk) + *reinterpret_cast(tma_buffer + hidden_bytes) = ld_nc_global(src_meta + token_idx); + + // Load `topk_weights` + if (lane_id < num_topk) + *reinterpret_cast(tma_buffer + hidden_bytes + sizeof(SourceMeta) + lane_id * sizeof(float)) = + ld_nc_global(topk_weights + token_idx * num_topk + lane_id); + + // Issue TMA store + tma_store_fence(); + __syncwarp(); + if (elect_one_sync()) + tma_store_1d(tma_buffer, shifted_x_buffers, num_bytes_per_token, false); + } + lane_id == current_rdma_idx ? (token_start_idx = static_cast(token_idx)) : 0; + } + + // Move queue tail + tma_store_wait<0>(); + __syncwarp(); + if (lane_id < kNumRDMARanks and is_lane_ready) + st_release_sys_global(nvl_channel_tail.buffer() + lane_id, cached_channel_tail_idx); + } + DEVICE_LOG_DEBUG_LANE(0,"rank %d warp %d | channel %d | RDMA SENDER DONE |token %d", rank, warp_id, channel_id, token_start_idx); + } else { + // Combiners and coordinators + // RDMA symmetric layout + auto rdma_channel_data = SymBuffer( + rdma_buffer_ptr, num_max_rdma_chunked_recv_tokens * num_bytes_per_token, kNumRDMARanks, channel_id, num_channels); + auto rdma_channel_head = SymBuffer(rdma_buffer_ptr, 1, kNumRDMARanks, channel_id, num_channels); + auto rdma_channel_tail = SymBuffer(rdma_buffer_ptr, 1, kNumRDMARanks, channel_id, num_channels); + + // NVL layouts + void* local_nvl_buffer = buffer_ptrs[nvl_rank]; + void* nvl_buffers[NUM_MAX_NVL_PEERS]; + #pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) + nvl_buffers[i] = buffer_ptrs[i]; + auto nvl_channel_x = + AsymBuffer( + local_nvl_buffer, num_max_nvl_chunked_recv_tokens * num_bytes_per_token, NUM_MAX_NVL_PEERS, channel_id, num_channels) + .advance_also(nvl_buffers); + auto nvl_channel_head = + AsymBuffer(nvl_buffers, kNumRDMARanks, NUM_MAX_NVL_PEERS, channel_id, num_channels, nvl_rank) + .advance_also(local_nvl_buffer); + auto nvl_channel_tail = AsymBuffer(local_nvl_buffer, kNumRDMARanks, NUM_MAX_NVL_PEERS, channel_id, num_channels) + .advance_also(nvl_buffers); + + // Combiner warp synchronization + __shared__ volatile int forwarder_nvl_head[kNumForwarders][NUM_MAX_NVL_PEERS]; + __shared__ volatile bool forwarder_retired[kNumForwarders]; + __shared__ volatile int rdma_receiver_rdma_head[kNumRDMAReceivers][kNumRDMARanks]; + __shared__ volatile bool rdma_receiver_retired[kNumRDMAReceivers]; + auto sync_forwarder_smem = [=]() { asm volatile("barrier.sync 0, %0;" ::"r"((kNumForwarders + 1) * 32)); }; + auto sync_rdma_receiver_smem = [=]() { asm volatile("barrier.sync 1, %0;" ::"r"((kNumRDMAReceivers + 1) * 32)); }; + + if (warp_role == WarpRole::kNVLAndRDMAForwarder) { + // Receive from NVL ranks and forward to RDMA ranks + // NOTES: this part is using "large warps" for each RDMA ranks + const auto dst_rdma_rank = warp_id / kNumWarpsPerForwarder; + const auto sub_warp_id = warp_id % kNumWarpsPerForwarder; + auto send_buffer = + dst_rdma_rank == rdma_rank ? rdma_channel_data.recv_buffer(dst_rdma_rank) : rdma_channel_data.send_buffer(dst_rdma_rank); + auto sync_large_warp = [=]() { + if (kNumWarpsPerForwarder == 1) { + __syncwarp(); + } else { +#ifdef ENABLE_DEBUG_LOGS + if (lane_id == 0 && warp_id % kNumWarpsPerForwarder == 0) { + DEVICE_LOG_DEBUG("rank %d warp %d | sync_large_warp barrier %d with %d threads", rank, warp_id, dst_rdma_rank + 2, kNumWarpsPerForwarder * 32); + } +#endif + asm volatile("bar.sync %0, %1;" ::"r"(dst_rdma_rank + 2), "r"(kNumWarpsPerForwarder * 32)); + } + }; + EP_STATIC_ASSERT(kNumWarpsPerForwarder == 1 or kNumRDMARanks + 2 <= 16, "Barriers are not enough"); + + // TMA stuffs + constexpr int kNumStages = 2; + constexpr int kNumTMALoadBytes = sizeof(int4) * 32; + constexpr int kNumTMABufferBytesPerStage = kNumTMALoadBytes * (NUM_MAX_NVL_PEERS + 1) + 16; + EP_STATIC_ASSERT(kNumTMABufferBytesPerStage * kNumStages <= kNumTMABytesPerForwarderWarp, "TMA buffer is not larger enough"); + + extern __shared__ __align__(1024) uint8_t smem_buffer[]; + auto smem_ptr = smem_buffer + warp_id * kNumStages * kNumTMABufferBytesPerStage; + auto tma_mbarrier = [=](const int& i) { + return reinterpret_cast(smem_ptr + i * kNumTMABufferBytesPerStage + kNumTMALoadBytes * (NUM_MAX_NVL_PEERS + 1)); + }; + uint32_t tma_phase[kNumStages] = {0}; + if (lane_id < kNumStages) { + mbarrier_init(tma_mbarrier(lane_id), 32); + fence_barrier_init(); + } + __syncwarp(); + + // Advance to the corresponding NVL buffer + nvl_channel_x.advance(dst_rdma_rank * num_max_nvl_chunked_recv_tokens_per_rdma * num_bytes_per_token); + nvl_channel_head.advance(dst_rdma_rank); + nvl_channel_tail.advance(dst_rdma_rank); + + // Clean shared memory and sync + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS <= 32, "Invalid number of NVL peers"); + lane_id < NUM_MAX_NVL_PEERS ? (forwarder_nvl_head[warp_id][lane_id] = 0) : 0; + lane_id == 0 ? (forwarder_retired[warp_id] = false) : false; + sync_forwarder_smem(); + + // Get count and cached head + int cached_nvl_channel_tail_idx = 0; + int num_tokens_to_combine = rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id]; + int num_tokens_prefix = channel_id == 0 ? 0 : rdma_channel_prefix_matrix[dst_rdma_rank * num_channels + channel_id - 1]; + num_tokens_to_combine -= num_tokens_prefix; + num_tokens_prefix += dst_rdma_rank == 0 ? 0 : rdma_rank_prefix_sum[dst_rdma_rank - 1]; + combined_nvl_head += num_tokens_prefix * NUM_MAX_NVL_PEERS; + + // Iterate over all tokens and combine by chunks + for (int token_start_idx = 0; token_start_idx < num_tokens_to_combine; token_start_idx += num_max_rdma_chunked_send_tokens) { + // Check destination queue emptiness, or wait a buffer to be released + auto token_end_idx = min(token_start_idx + num_max_rdma_chunked_send_tokens, num_tokens_to_combine); + auto num_chunked_tokens = token_end_idx - token_start_idx; + auto start_time = clock64(); + while (sub_warp_id == 0 and lane_id == 0) { + // Inequality: `num_max_rdma_chunked_recv_tokens - (tail - head) >= num_chunked_tokens` + // Here, `token_start_idx` is the actual tail + int num_used_slots = token_start_idx - ld_volatile_global(rdma_channel_head.buffer(dst_rdma_rank)); + if (num_max_rdma_chunked_recv_tokens - num_used_slots >= num_chunked_tokens) + break; + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES) { + printf( + "NixlEP combine forwarder (RDMA check) timeout, channel: %d, RDMA: %d, nvl: %d, dst RDMA: %d, head: %ld, tail: " + "%d, chunked: %d\n", + channel_id, + rdma_rank, + nvl_rank, + dst_rdma_rank, + ld_volatile_global(rdma_channel_head.buffer(dst_rdma_rank)), + token_start_idx, + num_chunked_tokens); + trap(); + } + } + sync_large_warp(); + + // Combine and write to the RDMA buffer + for (int token_idx = token_start_idx + sub_warp_id; token_idx < token_end_idx; token_idx += kNumWarpsPerForwarder) { + // Read expected head + EP_STATIC_ASSERT(kNumRDMARanks <= 32, "Invalid number of RDMA peers"); + int expected_head = -1; + if (lane_id < NUM_MAX_NVL_PEERS) { + expected_head = ld_nc_global(combined_nvl_head + token_idx * NUM_MAX_NVL_PEERS + lane_id); + expected_head < 0 ? (forwarder_nvl_head[warp_id][lane_id] = -expected_head - 1) + : (forwarder_nvl_head[warp_id][lane_id] = expected_head); + } + + // Wait lanes to be ready + start_time = clock64(); + while (cached_nvl_channel_tail_idx <= expected_head) { + cached_nvl_channel_tail_idx = ld_acquire_sys_global(nvl_channel_tail.buffer(lane_id)); + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES and lane_id < NUM_MAX_NVL_PEERS) { + printf( + "NixlEP combine forwarder (NVL check) timeout, channel: %d, RDMA: %d, nvl: %d, src NVL: %d, dst RDMA: %d, " + "tail: %d, waiting: %d, total: %d, sub: %d, large: %d, expected: %d\n", + channel_id, + rdma_rank, + nvl_rank, + lane_id, + dst_rdma_rank, + cached_nvl_channel_tail_idx, + token_idx, + num_tokens_to_combine, + sub_warp_id, + kNumWarpsPerForwarder, + expected_head); + trap(); + } + } + + // Combine current token + auto rdma_slot_idx = token_idx % num_max_rdma_chunked_recv_tokens; + void* shifted = send_buffer + rdma_slot_idx * num_bytes_per_token; + auto get_addr_fn = [&](int src_nvl_rank, int slot_idx, int hidden_int4_idx) -> int4* { + return reinterpret_cast(nvl_channel_x.buffer(src_nvl_rank) + slot_idx * num_bytes_per_token) + + hidden_int4_idx; + }; + auto recv_tw_fn = [&](int src_nvl_rank, int slot_idx, int topk_idx) -> float { + return ld_nc_global(reinterpret_cast(nvl_channel_x.buffer(src_nvl_rank) + slot_idx * num_bytes_per_token + + hidden_bytes + sizeof(SourceMeta)) + + topk_idx); + }; + combine_token( + expected_head >= 0, + expected_head, + lane_id, + hidden_int4, + num_topk, + static_cast(shifted), + reinterpret_cast(static_cast(shifted) + hidden_bytes + sizeof(SourceMeta)), + nullptr, + nullptr, + num_max_nvl_chunked_recv_tokens_per_rdma, + get_addr_fn, + recv_tw_fn, + smem_ptr, + tma_phase); + + // Update head + if (lane_id < NUM_MAX_NVL_PEERS) + expected_head < 0 ? (forwarder_nvl_head[warp_id][lane_id] = -expected_head - 1) + : (forwarder_nvl_head[warp_id][lane_id] = expected_head + 1); + } + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d | NVL FORWARDER about to sync_large_warp, dst_rdma_rank: %d, num_chunked_tokens: %d", rank, warp_id, dst_rdma_rank, num_chunked_tokens); + sync_large_warp(); + + DEVICE_LOG_DEBUG_LANE_SYNC(0,"rank %d warp %d | NVL FORWARDER sync_large_warp done, lane_id: %d, num_chunked_tokens: %d", rank, warp_id, lane_id, num_chunked_tokens); + // Issue RDMA send + if (sub_warp_id == kNumWarpsPerForwarder - 1) { + if (dst_rdma_rank != rdma_rank) { + auto rdma_slot_idx = token_start_idx % num_max_rdma_chunked_recv_tokens; + const size_t num_bytes_per_msg = num_chunked_tokens * num_bytes_per_token; + const auto dst_ptr = + reinterpret_cast(rdma_channel_data.recv_buffer(rdma_rank) + rdma_slot_idx * num_bytes_per_token); + const auto src_ptr = + reinterpret_cast(rdma_channel_data.send_buffer(dst_rdma_rank) + rdma_slot_idx * num_bytes_per_token); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(dst_rdma_rank, nvl_rank)); + EP_DEVICE_ASSERT(nixlGpuPostSingleWriteXferReq( + batch_req, 0, nixl_ctx.batch_offset_get(src_ptr), nixl_ctx.batch_offset_get(dst_ptr), num_bytes_per_msg, channel_id, true) == + NIXL_IN_PROG); + } else { + memory_fence(); + } + + // Write new RDMA tail + __syncwarp(); + if (elect_one_sync()) { + auto tail_ptr = reinterpret_cast(rdma_channel_tail.buffer(rdma_rank)); + if(dst_rdma_rank == rdma_rank){ + DEVICE_LOG_DEBUG("rank %d warp %d | RDMA FORWARDER | LOCAL TAIL UPDATED | dst_rdma_rank: %d, channel_id: %d, lane_id: %d, num_chunked_tokens: %d | local counter pointer: %p", rank, warp_id, dst_rdma_rank, channel_id, lane_id, num_chunked_tokens, (void*)rdma_channel_tail.buffer(rdma_rank)); + atomicAdd(reinterpret_cast(tail_ptr), static_cast(num_chunked_tokens)); + }else{ + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(dst_rdma_rank, nvl_rank)); + EP_DEVICE_ASSERT(nixlGpuPostSignalXferReq( + batch_req, 0, num_chunked_tokens, nixl_ctx.batch_offset_get(tail_ptr), channel_id) == + NIXL_IN_PROG); + } + } + } + } + + // Retired + __syncwarp(); + if (elect_one_sync()) + forwarder_retired[warp_id] = true; +#ifdef ENABLE_DEBUG_LOGS + if (lane_id == 0 and sub_warp_id == kNumWarpsPerForwarder - 1) { + DEVICE_LOG_DEBUG("rank %d warp %d | RDMA FORWARDER | sent %d tokens to RDMA rank %d FINISHED", rank, warp_id, num_tokens_to_combine, dst_rdma_rank); + } +#endif + } else if (warp_role == WarpRole::kRDMAReceiver) { + // Receive from RDMA ranks and write to the output tensor + // Clean shared memory and sync + EP_DEVICE_ASSERT(kNumRDMARanks <= 32); + lane_id < kNumRDMARanks ? (rdma_receiver_rdma_head[warp_id][lane_id] = 0) : 0; + lane_id == 0 ? (rdma_receiver_retired[warp_id] = false) : 0; + sync_rdma_receiver_smem(); + + // The same tokens as the dispatch process + int token_start_idx, token_end_idx; + get_channel_task_range(num_combined_tokens, num_channels, channel_id, token_start_idx, token_end_idx); + + // Iterate over all tokens and combine + int cached_channel_tail_idx = 0; + for (int64_t token_idx = token_start_idx + warp_id; token_idx < token_end_idx; token_idx += kNumRDMAReceivers) { + // Read expected head + EP_STATIC_ASSERT(kNumRDMARanks <= 32, "Invalid number of RDMA peers"); + int expected_head = -1; + if (lane_id < kNumRDMARanks) { + expected_head = ld_nc_global(combined_rdma_head + token_idx * kNumRDMARanks + lane_id); + (expected_head < 0) ? (rdma_receiver_rdma_head[warp_id][lane_id] = -expected_head - 1) + : (rdma_receiver_rdma_head[warp_id][lane_id] = expected_head); + } + + // Wait lanes to be ready + auto start_time = clock64(); +#ifdef ENABLE_DEBUG_LOGS + int nvl_recevier_poll_counter = 1; +#endif + while (cached_channel_tail_idx <= expected_head) { + cached_channel_tail_idx = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); +#ifdef ENABLE_DEBUG_LOGS + if (nvl_recevier_poll_counter % 100000 == 0) { + DEVICE_LOG_DEBUG_LANE(lane_id, "rank %d warp %d | RDMA RECEIVER, lane_id: %d, cached_channel_tail_idx: %d, expected_head: %d, print_count: %d", rank, warp_id, lane_id, cached_channel_tail_idx, expected_head, nvl_recevier_poll_counter); + } + nvl_recevier_poll_counter++; +#endif + + // Timeout check + if (clock64() - start_time > NUM_TIMEOUT_CYCLES) { + printf( + "NixlEP combine RDMA receiver timeout, channel: %d, RDMA: %d, nvl: %d, src RDMA: %d, tail: %d, waiting: %ld, " + "expect: %d\n", + channel_id, + rdma_rank, + nvl_rank, + lane_id, + cached_channel_tail_idx, + token_idx, + expected_head); + trap(); + } + } + __syncwarp(); + + // Combine current token + auto get_addr_fn = [&](int src_rdma_rank, int slot_idx, int hidden_int4_idx) -> int4* { + return reinterpret_cast(rdma_channel_data.recv_buffer(src_rdma_rank) + slot_idx * num_bytes_per_token) + + hidden_int4_idx; + }; + auto recv_tw_fn = [&](int src_rdma_rank, int slot_idx, int topk_idx) -> float { + return ld_nc_global(reinterpret_cast(rdma_channel_data.recv_buffer(src_rdma_rank) + + slot_idx * num_bytes_per_token + hidden_bytes + sizeof(SourceMeta)) + + topk_idx); + }; + uint32_t dummy_tma_phases[2]; + combine_token( + expected_head >= 0, + expected_head, + lane_id, + hidden_int4, + num_topk, + combined_x + token_idx * hidden_int4, + combined_topk_weights + token_idx * num_topk, + bias_0 == nullptr ? nullptr : bias_0 + token_idx * hidden_int4, + bias_1 == nullptr ? nullptr : bias_1 + token_idx * hidden_int4, + num_max_rdma_chunked_recv_tokens, + get_addr_fn, + recv_tw_fn, + nullptr, + dummy_tma_phases); + } + + // Retired + __syncwarp(); + if (elect_one_sync()) + rdma_receiver_retired[warp_id] = true; +#ifdef ENABLE_DEBUG_LOGS + if (lane_id < kNumRDMARanks) { + DEVICE_LOG_DEBUG("rank %d warp %d | RDMA RECEIVER FINISHED | last cached_channel_tail_idx: %d", rank, warp_id, cached_channel_tail_idx); + } +#endif + } else { + // Coordinator + // Sync shared memory status + is_forwarder_sm ? sync_forwarder_smem() : sync_rdma_receiver_smem(); + const auto num_warps_per_rdma_rank = kNumForwarders / kNumRDMARanks; + + int last_rdma_head = 0; + int last_nvl_head[kNumRDMARanks] = {0}; + int dst_rdma_rank = lane_id < kNumRDMARanks ? lane_id : 0; + int dst_nvl_rank = lane_id < NUM_MAX_NVL_PEERS ? lane_id : 0; + EP_STATIC_ASSERT(kNumCombineForwarderWarps <= 32, "Invalid number of forwarder warps"); + while (true) { + // Retired + if (not is_forwarder_sm and __all_sync(0xffffffff, lane_id >= kNumRDMAReceivers or rdma_receiver_retired[lane_id])) + break; + if (is_forwarder_sm and __all_sync(0xffffffff, lane_id >= kNumForwarders or forwarder_retired[lane_id])) + break; + + // Find minimum head for RDMA ranks + if (not is_forwarder_sm) { + int min_head = std::numeric_limits::max(); + #pragma unroll + for (int i = 0; i < kNumRDMAReceivers; ++i) + if (not rdma_receiver_retired[i]) + min_head = min(min_head, rdma_receiver_rdma_head[i][dst_rdma_rank]); + if (min_head != std::numeric_limits::max() and min_head >= last_rdma_head + num_max_rdma_chunked_send_tokens and + lane_id < kNumRDMARanks) { + if (dst_rdma_rank == rdma_rank) { + atomicAdd(reinterpret_cast(rdma_channel_head.buffer(rdma_rank)), static_cast(min_head - last_rdma_head)); + } else { + size_t head_counter_offset = nixl_ctx.batch_offset_get(reinterpret_cast(rdma_channel_head.buffer(rdma_rank))); + nixlGpuXferReqH batch_req = nixl_ctx.batch_get(translate_dst_rdma_rank(dst_rdma_rank, nvl_rank)); + EP_DEVICE_ASSERT(nixlGpuPostSignalXferReq(batch_req, 0, min_head - last_rdma_head, head_counter_offset) == NIXL_IN_PROG); + } + last_rdma_head = min_head; + } + } else { + // Find minimum head for NVL ranks + #pragma unroll + for (int i = 0; i < kNumRDMARanks; ++i) { + int min_head = std::numeric_limits::max(); + #pragma unroll + for (int j = 0; j < num_warps_per_rdma_rank; ++j) + if (not forwarder_retired[i * num_warps_per_rdma_rank + j]) + min_head = min(min_head, forwarder_nvl_head[i * num_warps_per_rdma_rank + j][dst_nvl_rank]); + if (min_head != std::numeric_limits::max() and min_head > last_nvl_head[i] and lane_id < NUM_MAX_NVL_PEERS) + st_relaxed_sys_global(nvl_channel_head.buffer_by(dst_nvl_rank) + i, last_nvl_head[i] = min_head); + } + } + + // Nanosleep and let other warps work + __nanosleep(NUM_WAIT_NANOSECONDS); + } +#ifdef ENABLE_DEBUG_LOGS + if (lane_id < kNumRDMARanks) { + DEVICE_LOG_DEBUG("rank %d warp %d channel %d | COORDINATOR FINISHED| last_rdma_head: %d, last_nvl_head: %d", rank, warp_id, channel_id, last_rdma_head, last_nvl_head[dst_nvl_rank]); + } +#endif + } + } +} + +void combine(cudaDataType_t type, + void* combined_x, + float* combined_topk_weights, + const bool* is_combined_token_in_rank, + const void* x, + const float* topk_weights, + const void* bias_0, + const void* bias_1, + const int* combined_rdma_head, + const int* combined_nvl_head, + const void* src_meta, + const int* rdma_channel_prefix_matrix, + const int* rdma_rank_prefix_sum, + const int* gbl_channel_prefix_matrix, + int num_tokens, + int num_combined_tokens, + int hidden, + int num_topk, + void* rdma_buffer_ptr, + int num_max_rdma_chunked_send_tokens, + int num_max_rdma_chunked_recv_tokens, + void** buffer_ptrs, + int num_max_nvl_chunked_send_tokens, + int num_max_nvl_chunked_recv_tokens, + int rank, + int num_ranks, + cudaStream_t stream, + int num_channels, + bool low_latency_mode, + gpu_nixl_ctx nixl_ctx) { + constexpr int kNumCombineForwarderWarps = 24; + constexpr int kNumTMABytesPerSenderWarp = 16384; + constexpr int kNumTMABytesPerForwarderWarp = 9248; + constexpr int smem_size = + std::max(kNumTMABytesPerSenderWarp * NUM_MAX_NVL_PEERS, kNumTMABytesPerForwarderWarp * kNumCombineForwarderWarps); + +#define COMBINE_LAUNCH_CASE(num_rdma_ranks) \ + { \ + auto combine_func = low_latency_mode ? combine \ + : combine; \ + SET_SHARED_MEMORY_FOR_TMA(combine_func); \ + LAUNCH_KERNEL(&cfg, \ + combine_func, \ + reinterpret_cast(combined_x), \ + combined_topk_weights, \ + is_combined_token_in_rank, \ + reinterpret_cast(x), \ + topk_weights, \ + reinterpret_cast(bias_0), \ + reinterpret_cast(bias_1), \ + combined_rdma_head, \ + combined_nvl_head, \ + reinterpret_cast(src_meta), \ + rdma_channel_prefix_matrix, \ + rdma_rank_prefix_sum, \ + gbl_channel_prefix_matrix, \ + num_tokens, \ + num_combined_tokens, \ + hidden, \ + num_topk, \ + rdma_buffer_ptr, \ + num_max_rdma_chunked_send_tokens, \ + num_max_rdma_chunked_recv_tokens, \ + buffer_ptrs, \ + num_max_nvl_chunked_send_tokens, \ + num_max_nvl_chunked_recv_tokens, \ + rank, \ + num_ranks, \ + nixl_ctx); \ + } \ + break + + int num_rdma_ranks = num_ranks / NUM_MAX_NVL_PEERS; + auto num_warps_per_forwarder = std::max(kNumCombineForwarderWarps / num_rdma_ranks, 1); + int num_forwarder_warps = num_rdma_ranks * num_warps_per_forwarder; + EP_HOST_ASSERT(num_rdma_ranks <= kNumCombineForwarderWarps); + EP_HOST_ASSERT(num_forwarder_warps > NUM_MAX_NVL_PEERS and num_forwarder_warps % num_rdma_ranks == 0); + EP_HOST_ASSERT(num_max_nvl_chunked_recv_tokens % num_rdma_ranks == 0); + EP_HOST_ASSERT(num_max_nvl_chunked_recv_tokens / num_rdma_ranks > + std::max(num_max_rdma_chunked_send_tokens, num_max_nvl_chunked_send_tokens)); + EP_HOST_ASSERT(num_max_nvl_chunked_recv_tokens / num_rdma_ranks - num_warps_per_forwarder >= num_max_nvl_chunked_send_tokens); + EP_HOST_ASSERT(num_max_rdma_chunked_send_tokens >= num_warps_per_forwarder); + EP_HOST_ASSERT(type == CUDA_R_16BF); + + SETUP_LAUNCH_CONFIG(num_channels * 2, (num_forwarder_warps + 1) * 32, stream); + SWITCH_RDMA_RANKS(COMBINE_LAUNCH_CASE); +#undef COMBINE_LAUNCH_CASE +} + +} // namespace internode + +} // namespace nixl_ep diff --git a/examples/device/ep/csrc/kernels/launch.cuh b/examples/device/ep/csrc/kernels/launch.cuh index 6c044a45cc..22090a4016 100644 --- a/examples/device/ep/csrc/kernels/launch.cuh +++ b/examples/device/ep/csrc/kernels/launch.cuh @@ -73,6 +73,30 @@ cfg.dynamicSmemBytes = smem_size; #endif #endif +#define SWITCH_NVL_RANKS(case_macro) \ + switch (num_nvl_ranks) { \ + case 2: \ + case_macro(2); \ + case 4: \ + case_macro(4); \ + case 8: \ + case_macro(8); \ + default: \ + EP_HOST_ASSERT(false and "Unsupported NVL ranks"); \ + } \ + while (false) + +#define SWITCH_RDMA_RANKS(case_macro) \ + switch (num_ranks / NUM_MAX_NVL_PEERS) { \ + case 2: case_macro(2); \ + case 4: case_macro(4); \ + case 8: case_macro(8); \ + case 16: case_macro(16); \ + case 18: case_macro(18); \ + case 20: case_macro(20); \ + default: EP_HOST_ASSERT(false and "Unsupported RDMA ranks"); \ + } while (false) + #define SWITCH_HIDDEN(case_macro) \ switch (hidden) { \ case 2048: case_macro(2048); \ diff --git a/examples/device/ep/csrc/kernels/layout.cu b/examples/device/ep/csrc/kernels/layout.cu new file mode 100644 index 0000000000..e4a80f5f05 --- /dev/null +++ b/examples/device/ep/csrc/kernels/layout.cu @@ -0,0 +1,135 @@ +#include "configs.cuh" +#include "exception.cuh" +#include "launch.cuh" + +namespace nixl_ep { + +namespace layout { + +template +__global__ void get_dispatch_layout(const topk_idx_t* topk_idx, + int* num_tokens_per_rank, int* num_tokens_per_rdma_rank, + int* num_tokens_per_expert, bool* is_token_in_rank, + int num_tokens, int num_topk, int num_ranks, int num_experts) { + auto sm_id = static_cast(blockIdx.x); + auto thread_id = static_cast(threadIdx.x); + + // Count expert statistics + __shared__ int num_tokens_per_expert_per_thread[kNumThreads][kNumExpertsPerSM]; + int expert_begin_idx = sm_id * kNumExpertsPerSM, expert_end_idx = min(expert_begin_idx + kNumExpertsPerSM, num_experts); + if (expert_begin_idx < expert_end_idx) { + // Per-thread count + #pragma unroll + for (int i = 0; i < kNumExpertsPerSM; ++ i) + num_tokens_per_expert_per_thread[thread_id][i] = 0; + #pragma unroll + for (int i = thread_id; i < num_tokens; i += kNumThreads) { + auto shifted_topk_idx = topk_idx + i * num_topk; + #pragma unroll + for (int j = 0, expert_idx; j < num_topk; ++ j) { + expert_idx = static_cast(shifted_topk_idx[j]); + if (expert_begin_idx <= expert_idx and expert_idx < expert_end_idx) + ++ num_tokens_per_expert_per_thread[thread_id][expert_idx - expert_begin_idx]; + } + } + __syncthreads(); + + // Sum up + EP_STATIC_ASSERT(kNumExpertsPerSM <= kNumThreads, "Too many experts per SM"); + if (expert_begin_idx + thread_id < expert_end_idx) { + int sum = 0; + #pragma unroll + for (int i = 0; i < kNumThreads; ++ i) + sum += num_tokens_per_expert_per_thread[i][thread_id]; + num_tokens_per_expert[expert_begin_idx + thread_id] = sum; + } + return; + } + + if (num_tokens_per_rdma_rank != nullptr) + EP_DEVICE_ASSERT(num_ranks % NUM_MAX_NVL_PEERS == 0 and num_ranks > NUM_MAX_NVL_PEERS); + + // Count rank statistics + constexpr int kNumRDMARanksPerSM = kNumRanksPerSM / NUM_MAX_NVL_PEERS; + __shared__ int num_tokens_per_rank_per_thread[kNumThreads][kNumRanksPerSM]; + __shared__ int num_tokens_per_rdma_rank_per_thread[kNumThreads][kNumRDMARanksPerSM]; + auto sm_begin = (num_experts + kNumExpertsPerSM - 1) / kNumExpertsPerSM; + int rank_begin_idx = (sm_id - sm_begin) * kNumRanksPerSM, rank_end_idx = min(rank_begin_idx + kNumRanksPerSM, num_ranks); + int rdma_rank_begin_idx = rank_begin_idx / NUM_MAX_NVL_PEERS, rdma_rank_end_idx = rank_end_idx / NUM_MAX_NVL_PEERS; + if (rank_begin_idx < rank_end_idx) { + const auto num_expert_per_rank = num_experts / num_ranks; + auto expert_begin = rank_begin_idx * num_expert_per_rank; + auto expert_end = rank_end_idx * num_expert_per_rank; + + // Per-thread count + #pragma unroll + for (int i = 0; i < kNumRanksPerSM; ++ i) + num_tokens_per_rank_per_thread[thread_id][i] = 0; + #pragma unroll + for (int i = 0; i < kNumRDMARanksPerSM; ++ i) + num_tokens_per_rdma_rank_per_thread[thread_id][i] = 0; + #pragma unroll + for (int i = thread_id; i < num_tokens; i += kNumThreads) { + auto shifted_topk_idx = topk_idx + i * num_topk; + int is_in_rank[kNumRanksPerSM] = {0}, is_in_rdma_rank[kNumRDMARanksPerSM] = {0}; + #pragma unroll + for (int j = 0, expert_idx, rank_idx; j < num_topk; ++j) { + expert_idx = static_cast(shifted_topk_idx[j]); + if (expert_begin <= expert_idx and expert_idx < expert_end) { + // Count single rank + rank_idx = expert_idx / num_expert_per_rank - rank_begin_idx; + is_in_rank[rank_idx] ++, is_in_rdma_rank[rank_idx / NUM_MAX_NVL_PEERS] ++; + } + } + + auto shifted_is_token_in_rank = is_token_in_rank + i * num_ranks; + #pragma unroll + for (int j = 0; j + rank_begin_idx < rank_end_idx; ++ j) { + shifted_is_token_in_rank[j + rank_begin_idx] = (is_in_rank[j] > 0); + num_tokens_per_rank_per_thread[thread_id][j] += (is_in_rank[j] > 0); + } + + #pragma unroll + for (int j = 0; j + rdma_rank_begin_idx < rdma_rank_end_idx; ++ j) + num_tokens_per_rdma_rank_per_thread[thread_id][j] += (is_in_rdma_rank[j] > 0); + } + __syncthreads(); + + // Sum up + EP_STATIC_ASSERT(kNumRanksPerSM <= kNumThreads, "Too many ranks per SM"); + if (rank_begin_idx + thread_id < rank_end_idx) { + int sum = 0; + #pragma unroll + for (int i = 0; i < kNumThreads; ++ i) + sum += num_tokens_per_rank_per_thread[i][thread_id]; + num_tokens_per_rank[rank_begin_idx + thread_id] = sum; + } + + if (num_tokens_per_rdma_rank != nullptr and rdma_rank_begin_idx + thread_id < rdma_rank_end_idx) { + int sum = 0; + #pragma unroll + for (int i = 0; i < kNumThreads; ++ i) + sum += num_tokens_per_rdma_rank_per_thread[i][thread_id]; + num_tokens_per_rdma_rank[rdma_rank_begin_idx + thread_id] = sum; + } + } +} + +void get_dispatch_layout(const topk_idx_t* topk_idx, + int* num_tokens_per_rank, int* num_tokens_per_rdma_rank, + int* num_tokens_per_expert, bool* is_token_in_rank, + int num_tokens, int num_topk, int num_ranks, int num_experts, + cudaStream_t stream) { + constexpr int kNumThreads = 256, kNumExpertsPerSM = 4, kNumRanksPerSM = 8; + int num_sms = ((num_experts + kNumExpertsPerSM - 1) / kNumExpertsPerSM) + (num_ranks + kNumRanksPerSM - 1) / kNumRanksPerSM; + EP_STATIC_ASSERT(kNumRanksPerSM % NUM_MAX_NVL_PEERS == 0, "Invalid number of ranks per SM"); + + SETUP_LAUNCH_CONFIG(num_sms, kNumThreads, stream); + LAUNCH_KERNEL(&cfg, (get_dispatch_layout), + topk_idx, num_tokens_per_rank, num_tokens_per_rdma_rank, num_tokens_per_expert, is_token_in_rank, + num_tokens, num_topk, num_ranks, num_experts); +} + +} // namespace layout + +} // namespace nixl_ep diff --git a/examples/device/ep/csrc/kernels/log.cuh b/examples/device/ep/csrc/kernels/log.cuh new file mode 100644 index 0000000000..2006a7b059 --- /dev/null +++ b/examples/device/ep/csrc/kernels/log.cuh @@ -0,0 +1,23 @@ +#pragma once +#ifdef ENABLE_DEBUG_LOGS +#define HOST_LOG_DEBUG(fmt, ...) printf("[DEBUG][%s] " fmt "\n", __func__, ##__VA_ARGS__) +#define DEVICE_LOG_DEBUG(fmt, ...) printf("[DEBUG][%s] " fmt "\n", __func__, ##__VA_ARGS__) + +#define _DEVICE_LOG_DEBUG_LANE_IMPL(lane, fmt, ...) do { \ + if (lane_id == (lane)) { \ + printf("[DEBUG][%s] " fmt "\n", __func__, ##__VA_ARGS__); \ + } \ +} while(0) + +#define DEVICE_LOG_DEBUG_LANE(lane, fmt, ...) _DEVICE_LOG_DEBUG_LANE_IMPL(lane, fmt, ##__VA_ARGS__) + +#define DEVICE_LOG_DEBUG_LANE_SYNC(lane, fmt, ...) do { \ + _DEVICE_LOG_DEBUG_LANE_IMPL(lane, fmt, ##__VA_ARGS__); \ + __syncwarp(); \ +} while(0) +#else +#define HOST_LOG_DEBUG(...) +#define DEVICE_LOG_DEBUG(...) +#define DEVICE_LOG_DEBUG_LANE(...) +#define DEVICE_LOG_DEBUG_LANE_SYNC(...) +#endif diff --git a/examples/device/ep/csrc/kernels/nixl_ep.cu b/examples/device/ep/csrc/kernels/nixl_ep.cu index b1b5dce442..d9291c54a2 100644 --- a/examples/device/ep/csrc/kernels/nixl_ep.cu +++ b/examples/device/ep/csrc/kernels/nixl_ep.cu @@ -61,7 +61,7 @@ dispatch(void* packed_recv_x, void* packed_recv_x_scales, int num_tokens, int num_max_dispatch_tokens_per_rank, int num_topk, int num_experts, int rank, int num_ranks, int num_warp_groups, int num_warps_per_group, - bool round_scale, int phases, ep_kernels::gpu_nixl_ctx nixl_ctx) { + bool round_scale, int phases, nixl_ep::gpu_nixl_ctx nixl_ctx) { const auto sm_id = static_cast(blockIdx.x); const auto thread_id = static_cast(threadIdx.x); const auto warp_id = thread_id / 32, lane_id = get_lane_id(); @@ -383,7 +383,7 @@ void dispatch(void* packed_recv_x, void* packed_recv_x_scales, int num_topk, int num_experts, int rank, int num_ranks, bool use_fp8, bool round_scale, bool use_ue8m0, void* workspace, int num_device_sms, - cudaStream_t stream, int phases, ep_kernels::gpu_nixl_ctx nixl_ctx) { + cudaStream_t stream, int phases, nixl_ep::gpu_nixl_ctx nixl_ctx) { constexpr int kNumMaxTopK = 11; const int num_warp_groups = ceil_div(num_experts, num_device_sms); const int num_warps_per_group = 32 / num_warp_groups; @@ -603,7 +603,7 @@ combine(void* combined_x, int num_max_dispatch_tokens_per_rank, int num_experts, int rank, int num_ranks, int num_warp_groups, int num_warps_per_group, - int phases, bool zero_copy, ep_kernels::gpu_nixl_ctx nixl_ctx) { + int phases, bool zero_copy, nixl_ep::gpu_nixl_ctx nixl_ctx) { const auto sm_id = __shfl_sync(0xffffffff, static_cast(blockIdx.x), 0); const auto num_sms = __shfl_sync(0xffffffff, static_cast(gridDim.x), 0); const auto thread_id = static_cast(threadIdx.x); @@ -998,7 +998,7 @@ void combine(void* combined_x, int num_topk, int num_experts, int rank, int num_ranks, bool use_logfmt, void* workspace, int num_device_sms, - cudaStream_t stream, int phases, bool zero_copy, ep_kernels::gpu_nixl_ctx nixl_ctx) { + cudaStream_t stream, int phases, bool zero_copy, nixl_ep::gpu_nixl_ctx nixl_ctx) { constexpr int kNumMaxTopk = 11; const int num_warp_groups = ceil_div(num_experts, num_device_sms); const int num_warps_per_group = 32 / num_warp_groups; @@ -1110,7 +1110,7 @@ void clean_mask_buffer(int* mask_buffer_ptr, int num_ranks, cudaStream_t stream) template __forceinline__ __device__ void barrier(int thread_id, int rank, int num_ranks, - int* mask_buffer_ptr, int* sync_buffer_ptr, ep_kernels::gpu_nixl_ctx nixl_ctx) { + int* mask_buffer_ptr, int* sync_buffer_ptr, nixl_ep::gpu_nixl_ctx nixl_ctx) { EP_DEVICE_ASSERT(kNumThreads >= num_ranks); if (thread_id < num_ranks && thread_id != rank) { @@ -1139,12 +1139,12 @@ __forceinline__ __device__ void barrier(int thread_id, int rank, int num_ranks, } template -__global__ void barrier_kernel(int* mask_buffer_ptr, int* sync_buffer_ptr, ep_kernels::gpu_nixl_ctx nixl_ctx) { +__global__ void barrier_kernel(int* mask_buffer_ptr, int* sync_buffer_ptr, nixl_ep::gpu_nixl_ctx nixl_ctx) { const auto thread_id = static_cast(threadIdx.x); barrier(thread_id, nixl_ctx.rank, nixl_ctx.num_ranks, mask_buffer_ptr, sync_buffer_ptr, nixl_ctx); } -void barrier(ep_kernels::gpu_nixl_ctx nixl_ctx, int* mask_buffer_ptr, int* sync_buffer_ptr, cudaStream_t stream) { +void barrier(nixl_ep::gpu_nixl_ctx nixl_ctx, int* mask_buffer_ptr, int* sync_buffer_ptr, cudaStream_t stream) { constexpr int kNumThreads = 32; SETUP_LAUNCH_CONFIG(1, kNumThreads, stream); LAUNCH_KERNEL(&cfg, barrier_kernel, mask_buffer_ptr, sync_buffer_ptr, nixl_ctx); diff --git a/examples/device/ep/csrc/kernels/runtime.cu b/examples/device/ep/csrc/kernels/runtime.cu new file mode 100644 index 0000000000..f387baf7a8 --- /dev/null +++ b/examples/device/ep/csrc/kernels/runtime.cu @@ -0,0 +1,46 @@ +#include +#include + +#include "configs.cuh" +#include "exception.cuh" +#include "launch.cuh" +#include "utils.cuh" + +#include + +namespace nixl_ep { + +namespace intranode { + +template +__global__ void barrier(int** barrier_signal_ptrs, int rank) { + barrier_block(barrier_signal_ptrs, rank); +} + +void barrier(int** barrier_signal_ptrs, int rank, int num_nvl_ranks, cudaStream_t stream) { +#define BARRIER_LAUNCH_CASE(ranks) \ + LAUNCH_KERNEL(&cfg, barrier, barrier_signal_ptrs, rank); \ + break + + SETUP_LAUNCH_CONFIG(1, 32, stream); + SWITCH_NVL_RANKS(BARRIER_LAUNCH_CASE); +#undef BARRIER_LAUNCH_CASE +} + +} // namespace intranode + +namespace internode { + +void* alloc(size_t size, size_t alignment) { + void *ptr; + CUDA_CHECK(cudaMalloc(&ptr, size)); + return ptr; +} + +void free(void* ptr) { + CUDA_CHECK(cudaFree(ptr)); +} + +} // namespace internode + +} // namespace nixl_ep diff --git a/examples/device/ep/csrc/kernels/utils.cuh b/examples/device/ep/csrc/kernels/utils.cuh index 7649326b94..10bce4ca5a 100644 --- a/examples/device/ep/csrc/kernels/utils.cuh +++ b/examples/device/ep/csrc/kernels/utils.cuh @@ -85,6 +85,9 @@ __device__ __forceinline__ void trap() { __device__ __forceinline__ void memory_fence() { asm volatile("fence.acq_rel.sys;":: : "memory"); } +__device__ __forceinline__ void st_relaxed_sys_global(const int *ptr, int val) { + asm volatile("st.relaxed.sys.global.s32 [%0], %1;"::"l"(ptr), "r"(val) : "memory"); +} __device__ __forceinline__ void st_release_sys_global(const int *ptr, int val) { asm volatile("st.release.sys.global.s32 [%0], %1;"::"l"(ptr), "r"(val) : "memory"); @@ -137,6 +140,41 @@ __device__ __forceinline__ int atomic_add_release_global(const uint64_t* ptr, ui asm volatile("atom.add.release.gpu.global.u64 %0, [%1], %2;" : "=l"(ret) : "l"(ptr), "l"(value)); return ret; } +__device__ __forceinline__ int ld_acquire_cta(const int *ptr) { + int ret; + asm volatile("ld.acquire.cta.s32 %0, [%1];" : "=r"(ret) : "l"(ptr)); + return ret; +} + +__device__ __forceinline__ int ld_acquire_cta(const volatile int *ptr) { + int ret; + asm volatile("ld.acquire.cta.s32 %0, [%1];" : "=r"(ret) : "l"(ptr)); + return ret; +} + +__device__ __forceinline__ int ld_volatile_global(const int *ptr) { + int ret; + asm volatile("ld.volatile.global.s32 %0, [%1];" : "=r"(ret) : "l"(ptr)); + return ret; +} + +__device__ __forceinline__ float ld_volatile_global(const float *ptr) { + float ret; + asm volatile("ld.volatile.global.f32 %0, [%1];" : "=f"(ret) : "l"(ptr)); + return ret; +} + +__device__ __forceinline__ int64_t ld_volatile_global(const int64_t *ptr) { + int64_t ret; + asm volatile("ld.volatile.global.s64 %0, [%1];" : "=l"(ret) : "l"(ptr)); + return ret; +} + +__device__ __forceinline__ int64_t ld_volatile_global(const uint64_t *ptr) { + int64_t ret; + asm volatile("ld.volatile.global.u64 %0, [%1];" : "=l"(ret) : "l"(ptr)); + return ret; +} #ifndef DISABLE_AGGRESSIVE_PTX_INSTRS #define LD_NC_FUNC "ld.global.nc.L1::no_allocate.L2::256B" @@ -266,6 +304,11 @@ __device__ __forceinline__ uint32_t elect_one_sync() { #endif } +__device__ __forceinline__ void fence_view_async_shared() { + asm volatile("fence.proxy.async.shared::cta; \n" :: ); +} + + // TMA PTX instructions #ifndef DISABLE_SM90_FEATURES __device__ __forceinline__ void fence_barrier_init() { @@ -332,7 +375,7 @@ __device__ __forceinline__ void tma_store_1d(const void* smem_ptr, const void* g asm volatile("cp.async.bulk.commit_group;"); } -template +template __device__ __forceinline__ void tma_store_wait() { asm volatile("cp.async.bulk.wait_group.read %0;" :: "n"(N) : "memory"); } @@ -421,6 +464,62 @@ __forceinline__ __device__ out_dtype_t extract_required_scale_format(float value } } +template +__forceinline__ __device__ void +barrier_block(int** barrier_signal_ptrs, int rank) { + auto thread_id = static_cast(threadIdx.x); + + // For non-sync-only cases, the memory operations by other threads in the block must be visible to the `sys` scope + if constexpr (not kSyncOnly) { + memory_fence(); + __syncthreads(); + } + + // Add self-ranks, sub other ranks + if (thread_id < kNumRanks) { + atomicAdd_system(barrier_signal_ptrs[rank] + thread_id, FINISHED_SUM_TAG); + atomicSub_system(barrier_signal_ptrs[thread_id] + rank, FINISHED_SUM_TAG); + } + EP_DEVICE_ASSERT(kNumRanks <= blockDim.x); + + // Check timeout + auto start_time = clock64(); + while (true) { + auto value = thread_id < kNumRanks ? ld_volatile_global(barrier_signal_ptrs[rank] + thread_id) : 0; + if (__all_sync(0xffffffff, value <= 0)) + break; + + if (clock64() - start_time > NUM_TIMEOUT_CYCLES and thread_id < kNumRanks) { + printf("NixlEP timeout check failed: rank = %d, thread = %d, value = %d)\n", rank, thread_id, value); + trap(); + } + } + __syncthreads(); +} + +__forceinline__ __device__ int atomic_cas_cta_acquire(int* addr, int x, int y) { + int ret; + asm volatile("atom.acquire.cta.shared::cta.cas.b32 %0, [%1], %2, %3;" : "=r"(ret) : "l"(addr), "r"(x), "r"(y) : "memory"); + return ret; +} + +__forceinline__ __device__ int atomic_exch_cta_release(int* addr, int x) { + int ret; + asm volatile("atom.release.cta.shared::cta.exch.b32 %0, [%1], %2;" : "=r"(ret) : "l"(addr), "r"(x) : "memory"); + return ret; +} + +__forceinline__ __device__ void acquire_lock(int* mutex) { + // To make later memory operations valid, we must use `acquire` for memory semantics + while (atomic_cas_cta_acquire(mutex, 0, 1) != 0); +} + +__forceinline__ __device__ void release_lock(int* mutex) { + // To make previous memory operations visible to other threads, we must use `release` for memory semantics + atomic_exch_cta_release(mutex, 0); +} + + // Operation functors template struct ReduceSum { __device__ T operator()(T a, T b) const { return a + b; } }; template struct ReduceMax { __device__ T operator()(T a, T b) const { return a > b ? a : b; } }; diff --git a/examples/device/ep/csrc/nixl_ep.cpp b/examples/device/ep/csrc/nixl_ep.cpp index 0761bf890e..2bde10d1f1 100644 --- a/examples/device/ep/csrc/nixl_ep.cpp +++ b/examples/device/ep/csrc/nixl_ep.cpp @@ -41,6 +41,7 @@ #include #include #include "kernels/exception.cuh" +#include "kernels/log.cuh" #include "nixl.h" #include #include @@ -51,12 +52,6 @@ #define NIXL_ETCD_WATCH_TIMEOUT std::chrono::microseconds(1000000000) // 1000 seconds -#ifdef ENABLE_DEBUG_LOGS -#define HOST_LOG_DEBUG(fmt, ...) printf("[DEBUG] " fmt "\n", ##__VA_ARGS__) -#else -#define HOST_LOG_DEBUG(...) -#endif - namespace nixl_ep { static void sleep_ms(int milliseconds) { @@ -112,37 +107,50 @@ static ino_t ipc_namespace_inode_get() { return st.st_ino; } -void Buffer::update_memory_buffers(int num_ranks, int num_experts_per_rank, int64_t num_rdma_bytes) +void Buffer::update_memory_buffers(int num_ranks, int num_experts_per_rank, int64_t num_nvl_bytes, int64_t num_rdma_bytes) { if (!available) { - init(num_ranks, num_experts_per_rank, num_rdma_bytes); + init(num_ranks, num_experts_per_rank, num_nvl_bytes, num_rdma_bytes); available = true; } else { throw std::runtime_error("Multiple calls to update_memory_buffers are not supported"); } } -Buffer::Buffer(int rank, bool explicitly_destroy, bool enable_shrink): +Buffer::Buffer(int rank, bool low_latency_mode, bool explicitly_destroy, bool enable_shrink): + low_latency_mode(low_latency_mode), rank(rank), num_ranks(1), explicitly_destroy(explicitly_destroy), comm_stream(at::cuda::getStreamFromPool(true)), enable_shrink(enable_shrink) {} -void Buffer::init(int num_ranks, int num_experts_per_rank, int64_t num_rdma_bytes) +void Buffer::init(int num_ranks, int num_experts_per_rank, int64_t num_nvl_bytes, int64_t num_rdma_bytes) { // Update buffer attributes this->max_num_ranks = num_ranks; this->max_experts_per_rank = num_experts_per_rank; + this->num_nvl_bytes = num_nvl_bytes; this->num_rdma_bytes = num_rdma_bytes; + // Metadata memory + int64_t barrier_signal_bytes = NUM_MAX_NVL_PEERS * sizeof(int); + int64_t buffer_ptr_bytes = NUM_MAX_NVL_PEERS * sizeof(void*); + int64_t barrier_signal_ptr_bytes = NUM_MAX_NVL_PEERS * sizeof(int*); + // Common checks EP_STATIC_ASSERT(NUM_BUFFER_ALIGNMENT_BYTES % sizeof(int4) == 0, "Invalid alignment"); - EP_HOST_ASSERT(num_rdma_bytes % NUM_BUFFER_ALIGNMENT_BYTES == 0); - EP_HOST_ASSERT(num_rdma_bytes / sizeof(int4) < std::numeric_limits::max()); - EP_HOST_ASSERT(0 <= rank and rank < num_ranks); + EP_HOST_ASSERT(num_nvl_bytes % NUM_BUFFER_ALIGNMENT_BYTES == 0 and (num_nvl_bytes <= std::numeric_limits::max() or num_rdma_bytes == 0)); + EP_HOST_ASSERT(num_rdma_bytes % NUM_BUFFER_ALIGNMENT_BYTES == 0 and (low_latency_mode or num_rdma_bytes <= std::numeric_limits::max())); + EP_HOST_ASSERT(0 <= rank and rank < num_ranks and (num_ranks <= NUM_MAX_NVL_PEERS * NUM_MAX_RDMA_PEERS or low_latency_mode)); + EP_HOST_ASSERT(num_ranks < NUM_MAX_NVL_PEERS or num_ranks % NUM_MAX_NVL_PEERS == 0); + if (num_rdma_bytes > 0) + EP_HOST_ASSERT(num_ranks > NUM_MAX_NVL_PEERS or low_latency_mode); // Get ranks CUDA_CHECK(cudaGetDevice(&device_id)); + rdma_rank = rank / NUM_MAX_NVL_PEERS, nvl_rank = rank % NUM_MAX_NVL_PEERS; + num_rdma_ranks = std::max(1, num_ranks / NUM_MAX_NVL_PEERS), num_nvl_ranks = std::min(num_ranks, NUM_MAX_NVL_PEERS); + // Get device info cudaDeviceProp device_prop = {}; CUDA_CHECK(cudaGetDeviceProperties(&device_prop, device_id)); @@ -151,10 +159,41 @@ void Buffer::init(int num_ranks, int num_experts_per_rank, int64_t num_rdma_byte auto per_channel_bytes = ceil_div(num_rdma_bytes, denom_sms); EP_HOST_ASSERT(per_channel_bytes < std::numeric_limits::max()); + if (num_nvl_bytes > 0) { + // Local IPC: alloc local memory and set local IPC handles + CUDA_CHECK(cudaMalloc(&buffer_ptrs[nvl_rank], num_nvl_bytes + barrier_signal_bytes + buffer_ptr_bytes + barrier_signal_ptr_bytes)); + CUDA_CHECK(cudaIpcGetMemHandle(&ipc_handles[nvl_rank], buffer_ptrs[nvl_rank])); + buffer_ptrs_gpu = reinterpret_cast(static_cast(buffer_ptrs[nvl_rank]) + num_nvl_bytes + barrier_signal_bytes); + + // Set barrier signals + barrier_signal_ptrs[nvl_rank] = reinterpret_cast(static_cast(buffer_ptrs[nvl_rank]) + num_nvl_bytes); + barrier_signal_ptrs_gpu = reinterpret_cast(static_cast(buffer_ptrs[nvl_rank]) + num_nvl_bytes + barrier_signal_bytes + buffer_ptr_bytes); + + // No need to synchronize, will do a full device sync during `sync` + CUDA_CHECK(cudaMemsetAsync(barrier_signal_ptrs[nvl_rank], 0, barrier_signal_bytes, comm_stream)); + } + // Create 32 MiB workspace CUDA_CHECK(cudaMalloc(&workspace, NUM_WORKSPACE_BYTES)); CUDA_CHECK(cudaMemsetAsync(workspace, 0, NUM_WORKSPACE_BYTES, comm_stream)); + // MoE counter + CUDA_CHECK(cudaMallocHost(&moe_recv_counter, sizeof(int64_t), cudaHostAllocMapped)); + CUDA_CHECK(cudaHostGetDevicePointer(&moe_recv_counter_mapped, const_cast(moe_recv_counter), 0)); + *moe_recv_counter = -1; + + // MoE expert-level counter + CUDA_CHECK(cudaMallocHost(&moe_recv_expert_counter, sizeof(int) * NUM_MAX_LOCAL_EXPERTS, cudaHostAllocMapped)); + CUDA_CHECK(cudaHostGetDevicePointer(&moe_recv_expert_counter_mapped, const_cast(moe_recv_expert_counter), 0)); + for (int i = 0; i < NUM_MAX_LOCAL_EXPERTS; ++ i) + moe_recv_expert_counter[i] = -1; + + // MoE RDMA-level counter + if (num_rdma_ranks > 0) { + CUDA_CHECK(cudaMallocHost(&moe_recv_rdma_counter, sizeof(int), cudaHostAllocMapped)); + CUDA_CHECK(cudaHostGetDevicePointer(&moe_recv_rdma_counter_mapped, const_cast(moe_recv_rdma_counter), 0)); + *moe_recv_rdma_counter = -1; + } EP_HOST_ASSERT(max_experts_per_rank > 0); CUDA_CHECK(cudaMalloc(&rdma_buffer_ptr, num_rdma_bytes)); CUDA_CHECK(cudaMemset(rdma_buffer_ptr, 0, num_rdma_bytes)); @@ -173,6 +212,10 @@ void Buffer::init(int num_ranks, int num_experts_per_rank, int64_t num_rdma_byte CUDA_CHECK(cudaMemset(sync_buffer_ptr, 0, num_sync_buffer_bytes)); CUDA_CHECK(cudaMalloc(&local_barrier_cnt_ptr, num_sync_buffer_bytes)); CUDA_CHECK(cudaMemset(local_barrier_cnt_ptr, 0, num_sync_buffer_bytes)); + + // Allocate internode barrier counter (for high-throughput mode) + CUDA_CHECK(cudaMalloc(&local_barrier_counter, sizeof(uint64_t))); + CUDA_CHECK(cudaMemset(local_barrier_counter, 0, sizeof(uint64_t))); CUDA_CHECK(cudaDeviceSynchronize()); strncpy(my_peer_info.ip, _get_local_ip().c_str(), MAX_IP_LENGTH - 1); @@ -180,6 +223,7 @@ void Buffer::init(int num_ranks, int num_experts_per_rank, int64_t num_rdma_byte my_peer_info.rdma_buffer_ptr = rdma_buffer_ptr; my_peer_info.device_id = get_local_device_id(); my_peer_info.sync_buffer_ptr = sync_buffer_ptr; + my_peer_info.barrier_ptr = local_barrier_counter; // For internode barrier my_peer_info.rank = rank; // Create an IPC handle for the rdma buffer @@ -210,15 +254,36 @@ bool Buffer::is_available() const { return available; } +bool Buffer::is_internode_available() const { + return is_available() and num_ranks > NUM_MAX_NVL_PEERS; +} + +int Buffer::get_num_rdma_ranks() const { + return num_rdma_ranks; +} + +int Buffer::get_rdma_rank() const { + return rdma_rank; +} + +int Buffer::get_root_rdma_rank(bool global) const { + return global ? nvl_rank : 0; +} + int Buffer::get_local_device_id() const { return device_id; } -torch::Tensor Buffer::get_local_buffer_tensor(const pybind11::object& dtype, int64_t offset) const { +pybind11::bytearray Buffer::get_local_ipc_handle() const { + return {ipc_handles[nvl_rank].reserved, CUDA_IPC_HANDLE_SIZE}; +} + +torch::Tensor Buffer::get_local_buffer_tensor(const pybind11::object& dtype, int64_t offset, bool use_rdma_buffer) const { torch::ScalarType casted_dtype = torch::python::detail::py_object_to_dtype(dtype); auto element_bytes = static_cast(elementSize(casted_dtype)); - auto base_ptr = static_cast(rdma_buffer_ptr) + offset; - return torch::from_blob(base_ptr, num_rdma_bytes / element_bytes, torch::TensorOptions().dtype(casted_dtype).device(at::kCUDA)); + auto base_ptr = static_cast(use_rdma_buffer ? rdma_buffer_ptr : buffer_ptrs[nvl_rank]) + offset; + auto num_bytes = use_rdma_buffer ? num_rdma_bytes : num_nvl_bytes; + return torch::from_blob(base_ptr, num_bytes / element_bytes, torch::TensorOptions().dtype(casted_dtype).device(at::kCUDA)); } torch::Stream Buffer::get_comm_stream() const { @@ -231,6 +296,19 @@ void Buffer::destroy() { // Synchronize CUDA_CHECK(cudaDeviceSynchronize()); + if (num_nvl_bytes > 0) { + intranode::barrier(barrier_signal_ptrs_gpu, nvl_rank, num_nvl_ranks, comm_stream); + CUDA_CHECK(cudaDeviceSynchronize()); + + // Close remote IPC + if (is_available()) { + for (int i = 0; i < num_nvl_ranks; ++ i) if (i != nvl_rank) + CUDA_CHECK(cudaIpcCloseMemHandle(buffer_ptrs[i])); + } + + // Free local buffer and error flag + CUDA_CHECK(cudaFree(buffer_ptrs[nvl_rank])); + } cudaFree(rdma_buffer_ptr); if (nixl_agent_info and nixl_agent_info->agent != nullptr and getenv("NIXL_ETCD_ENDPOINTS")) { @@ -245,9 +323,20 @@ void Buffer::destroy() { cudaFree(sync_buffer_ptr); cudaFree(local_barrier_cnt_ptr); + cudaFree(local_barrier_counter); + + // Free internode barrier epoch counter if allocated + if (nixl_ctx && nixl_ctx->gpu.last_barrier_counter) { + cudaFree(nixl_ctx->gpu.last_barrier_counter); + nixl_ctx->gpu.last_barrier_counter = nullptr; + } // Free workspace CUDA_CHECK(cudaFree(workspace)); + CUDA_CHECK(cudaFreeHost(const_cast(moe_recv_counter))); + + // Free chunked mode staffs + CUDA_CHECK(cudaFreeHost(const_cast(moe_recv_expert_counter))); destroyed = true; available = false; @@ -255,7 +344,7 @@ void Buffer::destroy() { void Buffer::barrier() { auto compute_stream = at::cuda::getCurrentCUDAStream(); - ep_kernels::barrier(nixl_ctx->gpu,mask_buffer_ptr, sync_buffer_ptr, compute_stream); + ep_kernels::barrier(nixl_ctx->gpu, mask_buffer_ptr, sync_buffer_ptr, compute_stream); } void Buffer::_nixl_agents_connect(const std::vector& ranks, const std::vector& remote_mds) { @@ -324,7 +413,31 @@ void Buffer::_nixl_agents_peer_info_gather(std::vector& ranks) { } } -void Buffer::connect_ranks(const std::vector& remote_ranks_list, const std::optional>& remote_mds) { +void Buffer::_ipc_handles_sync(const std::vector> &all_gathered_handles = {}) { + if (num_nvl_bytes > 0) { + EP_HOST_ASSERT(all_gathered_handles.size() == max_num_ranks); + for (int i = 0, offset = rdma_rank * num_nvl_ranks; i < num_nvl_ranks; ++ i) { + EP_HOST_ASSERT(all_gathered_handles[offset + i].has_value()); + auto handle_str = std::string(all_gathered_handles[offset + i].value()); + EP_HOST_ASSERT(handle_str.size() == CUDA_IPC_HANDLE_SIZE); + if (offset + i != rank) { + std::memcpy(ipc_handles[i].reserved, handle_str.c_str(), CUDA_IPC_HANDLE_SIZE); + CUDA_CHECK(cudaIpcOpenMemHandle(&buffer_ptrs[i], ipc_handles[i], cudaIpcMemLazyEnablePeerAccess)); + barrier_signal_ptrs[i] = reinterpret_cast(static_cast(buffer_ptrs[i]) + num_nvl_bytes); + } else { + EP_HOST_ASSERT(std::memcmp(ipc_handles[i].reserved, handle_str.c_str(), CUDA_IPC_HANDLE_SIZE) == 0); + } + } + + // Copy all buffer and barrier signal pointers to GPU + CUDA_CHECK(cudaMemcpy(buffer_ptrs_gpu, buffer_ptrs, sizeof(void*) * NUM_MAX_NVL_PEERS, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(barrier_signal_ptrs_gpu, barrier_signal_ptrs, sizeof(int*) * NUM_MAX_NVL_PEERS, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaDeviceSynchronize()); + } +} + +void Buffer::connect_ranks(const std::vector& remote_ranks_list, const std::optional>& remote_mds, + const std::vector> &all_gathered_handles) { EP_HOST_ASSERT(!remote_ranks_list.empty()); EP_HOST_ASSERT(!remote_mds.has_value() || remote_mds->size() == remote_ranks_list.size()); @@ -333,6 +446,9 @@ void Buffer::connect_ranks(const std::vector& remote_ranks_list, const std: int max_added_rank = std::max(rank, *std::max_element(remote_ranks_list.begin(), remote_ranks_list.end())); num_ranks = std::max(num_ranks, max_added_rank + 1); + if (all_gathered_handles.size() > 0) + _ipc_handles_sync(all_gathered_handles); + for (size_t i = 0; i < remote_ranks_list.size(); i++) { int remote_rank = remote_ranks_list[i]; // Skip self and ranks we are already connected to @@ -340,7 +456,9 @@ void Buffer::connect_ranks(const std::vector& remote_ranks_list, const std: continue; new_ranks.push_back(remote_rank); - CUDA_CHECK(cudaMemset(mask_buffer_ptr + remote_rank, 0, sizeof(int))); + if (enable_shrink) { + CUDA_CHECK(cudaMemset(mask_buffer_ptr + remote_rank, 0, sizeof(int))); + } CUDA_CHECK(cudaMemset(local_barrier_cnt_ptr + remote_rank, 0, sizeof(int))); CUDA_CHECK(cudaMemset(sync_buffer_ptr + remote_rank, 0, sizeof(int))); @@ -396,6 +514,481 @@ void Buffer::disconnect_ranks(const std::vector& remote_ranks_list) { num_ranks = max_rank + 1; // Sparse indexing maintained } +std::tuple, torch::Tensor, torch::Tensor, std::optional> +Buffer::get_dispatch_layout(const torch::Tensor& topk_idx, int num_experts, + std::optional& previous_event, bool async, bool allocate_on_comm_stream) { + EP_HOST_ASSERT(topk_idx.dim() == 2); + EP_HOST_ASSERT(topk_idx.is_contiguous()); + EP_HOST_ASSERT(num_experts > 0); + + // Allocate all tensors on comm stream if set + // NOTES: do not allocate tensors upfront! + auto compute_stream = at::cuda::getCurrentCUDAStream(); + if (allocate_on_comm_stream) { + EP_HOST_ASSERT(previous_event.has_value() and async); + at::cuda::setCurrentCUDAStream(comm_stream); + } + + // Wait previous tasks to be finished + if (previous_event.has_value()) { + stream_wait(comm_stream, previous_event.value()); + } else { + stream_wait(comm_stream, compute_stream); + } + + auto num_tokens = static_cast(topk_idx.size(0)), num_topk = static_cast(topk_idx.size(1)); + auto num_tokens_per_rank = torch::empty({num_ranks}, dtype(torch::kInt32).device(torch::kCUDA)); + auto num_tokens_per_rdma_rank = std::optional(); + auto num_tokens_per_expert = torch::empty({num_experts}, dtype(torch::kInt32).device(torch::kCUDA)); + auto is_token_in_rank = torch::empty({num_tokens, num_ranks}, dtype(torch::kBool).device(torch::kCUDA)); + if (is_internode_available()) + num_tokens_per_rdma_rank = torch::empty({num_rdma_ranks}, dtype(torch::kInt32).device(torch::kCUDA)); + + layout::get_dispatch_layout(topk_idx.data_ptr(), + num_tokens_per_rank.data_ptr(), + num_tokens_per_rdma_rank.has_value() ? num_tokens_per_rdma_rank.value().data_ptr() : nullptr, + num_tokens_per_expert.data_ptr(), + is_token_in_rank.data_ptr(), + num_tokens, num_topk, num_ranks, num_experts, + comm_stream); + + // Wait streams + std::optional event; + if (async) { + event = EventHandle(comm_stream); + for (auto& t: {topk_idx, num_tokens_per_rank, num_tokens_per_expert, is_token_in_rank}) { + t.record_stream(comm_stream); + if (allocate_on_comm_stream) + t.record_stream(compute_stream); + } + for (auto& to: {num_tokens_per_rdma_rank}) { + to.has_value() ? to->record_stream(comm_stream) : void(); + if (allocate_on_comm_stream) + to.has_value() ? to->record_stream(compute_stream) : void(); + } + } else { + stream_wait(compute_stream, comm_stream); + } + + // Switch back compute stream + if (allocate_on_comm_stream) + at::cuda::setCurrentCUDAStream(compute_stream); + + return {num_tokens_per_rank, num_tokens_per_rdma_rank, num_tokens_per_expert, is_token_in_rank, event}; +} + +std::tuple, std::optional, std::optional, std::vector, torch::Tensor, torch::Tensor, std::optional, torch::Tensor, std::optional, torch::Tensor, std::optional, std::optional, std::optional, std::optional> +Buffer::internode_dispatch(const torch::Tensor& x, const std::optional& x_scales, + const std::optional& topk_idx, const std::optional& topk_weights, + const std::optional& num_tokens_per_rank, const std::optional& num_tokens_per_rdma_rank, + const torch::Tensor& is_token_in_rank, const std::optional& num_tokens_per_expert, + int cached_num_recv_tokens, int cached_num_rdma_recv_tokens, + const std::optional& cached_rdma_channel_prefix_matrix, const std::optional& cached_recv_rdma_rank_prefix_sum, + const std::optional& cached_gbl_channel_prefix_matrix, const std::optional& cached_recv_gbl_rank_prefix_sum, + int expert_alignment, const Config& config, std::optional& previous_event, bool async, bool allocate_on_comm_stream) { + // In dispatch, CPU will busy-wait until GPU receive tensor size metadata from other ranks, which can be quite long. + // If users of DeepEP need to execute other Python code on other threads, such as KV transfer, their code will get stuck due to GIL + // unless we release GIL here. + pybind11::gil_scoped_release release; + HOST_LOG_DEBUG("internode_dispatch"); + + const int num_channels = config.num_sms / 2; + EP_HOST_ASSERT(config.num_sms % 2 == 0); + EP_HOST_ASSERT(0 < get_num_rdma_ranks() and get_num_rdma_ranks() <= NUM_MAX_RDMA_PEERS); + + bool cached_mode = cached_rdma_channel_prefix_matrix.has_value(); + if (cached_mode) { + EP_HOST_ASSERT(cached_rdma_channel_prefix_matrix.has_value()); + EP_HOST_ASSERT(cached_recv_rdma_rank_prefix_sum.has_value()); + EP_HOST_ASSERT(cached_gbl_channel_prefix_matrix.has_value()); + EP_HOST_ASSERT(cached_recv_gbl_rank_prefix_sum.has_value()); + } else { + EP_HOST_ASSERT(num_tokens_per_rank.has_value()); + EP_HOST_ASSERT(num_tokens_per_rdma_rank.has_value()); + EP_HOST_ASSERT(num_tokens_per_expert.has_value()); + } + + // Type checks + if (cached_mode) { + EP_HOST_ASSERT(cached_rdma_channel_prefix_matrix->scalar_type() == torch::kInt32); + EP_HOST_ASSERT(cached_recv_rdma_rank_prefix_sum->scalar_type() == torch::kInt32); + EP_HOST_ASSERT(cached_gbl_channel_prefix_matrix->scalar_type() == torch::kInt32); + EP_HOST_ASSERT(cached_recv_gbl_rank_prefix_sum->scalar_type() == torch::kInt32); + } else { + EP_HOST_ASSERT(num_tokens_per_rank->scalar_type() == torch::kInt32); + EP_HOST_ASSERT(num_tokens_per_rdma_rank->scalar_type() == torch::kInt32); + EP_HOST_ASSERT(num_tokens_per_expert->scalar_type() == torch::kInt32); + } + + // Shape and contiguous checks + EP_HOST_ASSERT(x.dim() == 2 and x.is_contiguous()); + EP_HOST_ASSERT((x.size(1) * x.element_size()) % sizeof(int4) == 0); + if (cached_mode) { + EP_HOST_ASSERT(cached_rdma_channel_prefix_matrix->dim() == 2 and cached_rdma_channel_prefix_matrix->is_contiguous()); + EP_HOST_ASSERT(cached_rdma_channel_prefix_matrix->size(0) == num_rdma_ranks and cached_rdma_channel_prefix_matrix->size(1) == num_channels); + EP_HOST_ASSERT(cached_recv_rdma_rank_prefix_sum->dim() == 1 and cached_recv_rdma_rank_prefix_sum->is_contiguous()); + EP_HOST_ASSERT(cached_recv_rdma_rank_prefix_sum->size(0) == num_rdma_ranks); + EP_HOST_ASSERT(cached_gbl_channel_prefix_matrix->dim() == 2 and cached_gbl_channel_prefix_matrix->is_contiguous()); + EP_HOST_ASSERT(cached_gbl_channel_prefix_matrix->size(0) == num_ranks and cached_gbl_channel_prefix_matrix->size(1) == num_channels); + EP_HOST_ASSERT(cached_recv_gbl_rank_prefix_sum->dim() == 1 and cached_recv_gbl_rank_prefix_sum->is_contiguous()); + EP_HOST_ASSERT(cached_recv_gbl_rank_prefix_sum->size(0) == num_ranks); + } else { + EP_HOST_ASSERT(num_tokens_per_rank->dim() == 1 and num_tokens_per_rank->is_contiguous()); + EP_HOST_ASSERT(num_tokens_per_rdma_rank->dim() == 1 and num_tokens_per_rdma_rank->is_contiguous()); + EP_HOST_ASSERT(num_tokens_per_expert->dim() == 1 and num_tokens_per_expert->is_contiguous()); + EP_HOST_ASSERT(num_tokens_per_rank->size(0) == num_ranks); + EP_HOST_ASSERT(num_tokens_per_rdma_rank->size(0) == num_rdma_ranks); + EP_HOST_ASSERT(num_tokens_per_expert->size(0) % num_ranks == 0); + EP_HOST_ASSERT(num_tokens_per_expert->size(0) / num_ranks <= NUM_MAX_LOCAL_EXPERTS); + } + + auto num_tokens = static_cast(x.size(0)), hidden = static_cast(x.size(1)), hidden_int4 = static_cast(x.size(1) * x.element_size() / sizeof(int4)); + auto num_experts = cached_mode ? 0 : static_cast(num_tokens_per_expert->size(0)), num_local_experts = num_experts / num_ranks; + + // Top-k checks + int num_topk = 0; + topk_idx_t* topk_idx_ptr = nullptr; + float* topk_weights_ptr = nullptr; + EP_HOST_ASSERT(topk_idx.has_value() == topk_weights.has_value()); + if (topk_idx.has_value()) { + num_topk = static_cast(topk_idx->size(1)); + EP_HOST_ASSERT(num_experts > 0); + EP_HOST_ASSERT(topk_idx->dim() == 2 and topk_idx->is_contiguous()); + EP_HOST_ASSERT(topk_weights->dim() == 2 and topk_weights->is_contiguous()); + EP_HOST_ASSERT(num_tokens == topk_idx->size(0) and num_tokens == topk_weights->size(0)); + EP_HOST_ASSERT(num_topk == topk_weights->size(1)); + EP_HOST_ASSERT(topk_weights->scalar_type() == torch::kFloat32); + topk_idx_ptr = topk_idx->data_ptr(); + topk_weights_ptr = topk_weights->data_ptr(); + } + + // FP8 scales checks + float* x_scales_ptr = nullptr; + int num_scales = 0, scale_token_stride = 0, scale_hidden_stride = 0; + if (x_scales.has_value()) { + EP_HOST_ASSERT(x.element_size() == 1); + EP_HOST_ASSERT(x_scales->scalar_type() == torch::kFloat32 or x_scales->scalar_type() == torch::kInt); + EP_HOST_ASSERT(x_scales->dim() == 2); + EP_HOST_ASSERT(x_scales->size(0) == num_tokens); + num_scales = x_scales->dim() == 1 ? 1 : static_cast(x_scales->size(1)); + x_scales_ptr = static_cast(x_scales->data_ptr()); + scale_token_stride = static_cast(x_scales->stride(0)); + scale_hidden_stride = static_cast(x_scales->stride(1)); + } + + // Allocate all tensors on comm stream if set + // NOTES: do not allocate tensors upfront! + auto compute_stream = at::cuda::getCurrentCUDAStream(); + if (allocate_on_comm_stream) { + EP_HOST_ASSERT(previous_event.has_value() and async); + at::cuda::setCurrentCUDAStream(comm_stream); + } + + // Wait previous tasks to be finished + if (previous_event.has_value()) { + stream_wait(comm_stream, previous_event.value()); + } else { + stream_wait(comm_stream, compute_stream); + } + + // Create handles (only return for non-cached mode) + int num_recv_tokens = -1, num_rdma_recv_tokens = -1; + auto rdma_channel_prefix_matrix = torch::Tensor(); + auto recv_rdma_rank_prefix_sum = torch::Tensor(); + auto gbl_channel_prefix_matrix = torch::Tensor(); + auto recv_gbl_rank_prefix_sum = torch::Tensor(); + std::vector num_recv_tokens_per_expert_list; + + // Barrier or send sizes + if (cached_mode) { + num_recv_tokens = cached_num_recv_tokens; + num_rdma_recv_tokens = cached_num_rdma_recv_tokens; + rdma_channel_prefix_matrix = cached_rdma_channel_prefix_matrix.value(); + recv_rdma_rank_prefix_sum = cached_recv_rdma_rank_prefix_sum.value(); + gbl_channel_prefix_matrix = cached_gbl_channel_prefix_matrix.value(); + recv_gbl_rank_prefix_sum = cached_recv_gbl_rank_prefix_sum.value(); + + // Just a barrier and clean flags + internode::cached_notify(hidden_int4, num_scales, num_topk, num_topk, + num_ranks, num_channels, 0, nullptr, + nullptr, nullptr, nullptr, + rdma_buffer_ptr, config.num_max_rdma_chunked_recv_tokens, + buffer_ptrs_gpu, config.num_max_nvl_chunked_recv_tokens, + barrier_signal_ptrs_gpu, rank, comm_stream, + config.get_rdma_buffer_size_hint(hidden_int4 * sizeof(int4), num_ranks), + num_nvl_bytes, true, low_latency_mode, nixl_ctx->gpu); + } else { + rdma_channel_prefix_matrix = torch::empty({num_rdma_ranks, num_channels}, dtype(torch::kInt32).device(torch::kCUDA)); + recv_rdma_rank_prefix_sum = torch::empty({num_rdma_ranks}, dtype(torch::kInt32).device(torch::kCUDA)); + gbl_channel_prefix_matrix = torch::empty({num_ranks, num_channels}, dtype(torch::kInt32).device(torch::kCUDA)); + recv_gbl_rank_prefix_sum = torch::empty({num_ranks}, dtype(torch::kInt32).device(torch::kCUDA)); + + // Send sizes + *moe_recv_counter = -1, *moe_recv_rdma_counter = -1; + for (int i = 0; i < num_local_experts; ++ i) + moe_recv_expert_counter[i] = -1; + internode::notify_dispatch(num_tokens_per_rank->data_ptr(), moe_recv_counter_mapped, num_ranks, + num_tokens_per_rdma_rank->data_ptr(), moe_recv_rdma_counter_mapped, + num_tokens_per_expert->data_ptr(), moe_recv_expert_counter_mapped, num_experts, + is_token_in_rank.data_ptr(), num_tokens, num_channels, + hidden_int4, num_scales, num_topk, expert_alignment, + rdma_channel_prefix_matrix.data_ptr(), recv_rdma_rank_prefix_sum.data_ptr(), + gbl_channel_prefix_matrix.data_ptr(), recv_gbl_rank_prefix_sum.data_ptr(), + rdma_buffer_ptr, config.num_max_rdma_chunked_recv_tokens, + buffer_ptrs_gpu, config.num_max_nvl_chunked_recv_tokens, + barrier_signal_ptrs_gpu, rank, comm_stream, + config.get_rdma_buffer_size_hint(hidden_int4 * sizeof(int4), num_ranks), + num_nvl_bytes, low_latency_mode, nixl_ctx->gpu); + + // Synchronize total received tokens and tokens per expert + auto start_time = std::chrono::high_resolution_clock::now(); + while (true) { + // Read total count + num_recv_tokens = static_cast(*moe_recv_counter); + num_rdma_recv_tokens = static_cast(*moe_recv_rdma_counter); + + // Read per-expert count + bool ready = (num_recv_tokens >= 0) and (num_rdma_recv_tokens >= 0); + for (int i = 0; i < num_local_experts and ready; ++ i) + ready &= moe_recv_expert_counter[i] >= 0; + + if (ready) + break; + + // Timeout check + if (std::chrono::duration_cast(std::chrono::high_resolution_clock::now() - start_time).count() > NUM_CPU_TIMEOUT_SECS) { + // printf("Global rank: %d, num_recv_tokens: %d, num_rdma_recv_tokens: %d\n", rank, num_recv_tokens, num_rdma_recv_tokens); + for (int i = 0; i < num_local_experts; ++ i) + printf("moe_recv_expert_counter[%d]: %d\n", i, moe_recv_expert_counter[i]); + throw std::runtime_error("NixlEP error: timeout (dispatch CPU)"); + } + } + num_recv_tokens_per_expert_list = std::vector(moe_recv_expert_counter, moe_recv_expert_counter + num_local_experts); + } + + // Allocate new tensors + auto recv_x = torch::empty({num_recv_tokens, hidden}, x.options()); + auto recv_topk_idx = std::optional(), recv_topk_weights = std::optional(), recv_x_scales = std::optional(); + auto recv_src_meta = std::optional(); + auto recv_rdma_channel_prefix_matrix = std::optional(); + auto recv_gbl_channel_prefix_matrix = std::optional(); + auto send_rdma_head = std::optional(); + auto send_nvl_head = std::optional(); + if (not cached_mode) { + recv_src_meta = torch::empty({num_recv_tokens, internode::get_source_meta_bytes()}, dtype(torch::kByte).device(torch::kCUDA)); + recv_rdma_channel_prefix_matrix = torch::empty({num_rdma_ranks, num_channels}, dtype(torch::kInt32).device(torch::kCUDA)); + recv_gbl_channel_prefix_matrix = torch::empty({num_ranks, num_channels}, dtype(torch::kInt32).device(torch::kCUDA)); + send_rdma_head = torch::empty({num_tokens, num_rdma_ranks}, dtype(torch::kInt32).device(torch::kCUDA)); + send_nvl_head = torch::empty({num_rdma_recv_tokens, NUM_MAX_NVL_PEERS}, dtype(torch::kInt32).device(torch::kCUDA)); + } + + // Assign pointers + topk_idx_t* recv_topk_idx_ptr = nullptr; + float* recv_topk_weights_ptr = nullptr; + float* recv_x_scales_ptr = nullptr; + if (topk_idx.has_value()) { + recv_topk_idx = torch::empty({num_recv_tokens, num_topk}, topk_idx->options()); + recv_topk_weights = torch::empty({num_recv_tokens, num_topk}, topk_weights->options()); + recv_topk_idx_ptr = recv_topk_idx->data_ptr(); + recv_topk_weights_ptr = recv_topk_weights->data_ptr(); + } + if (x_scales.has_value()) { + recv_x_scales = x_scales->dim() == 1 ? + torch::empty({num_recv_tokens}, x_scales->options()) : + torch::empty({num_recv_tokens, num_scales}, x_scales->options()); + recv_x_scales_ptr = static_cast(recv_x_scales->data_ptr()); + } + + // Launch data dispatch + // NOTES: the buffer size checks are moved into the `.cu` file + internode::dispatch(recv_x.data_ptr(), recv_x_scales_ptr, recv_topk_idx_ptr, recv_topk_weights_ptr, + cached_mode ? nullptr : recv_src_meta->data_ptr(), + x.data_ptr(), x_scales_ptr, topk_idx_ptr, topk_weights_ptr, + cached_mode ? nullptr : send_rdma_head->data_ptr(), cached_mode ? nullptr : send_nvl_head->data_ptr(), + cached_mode ? nullptr : recv_rdma_channel_prefix_matrix->data_ptr(), + cached_mode ? nullptr : recv_gbl_channel_prefix_matrix->data_ptr(), + rdma_channel_prefix_matrix.data_ptr(), recv_rdma_rank_prefix_sum.data_ptr(), + gbl_channel_prefix_matrix.data_ptr(), recv_gbl_rank_prefix_sum.data_ptr(), + is_token_in_rank.data_ptr(), + num_tokens, hidden_int4, num_scales, num_topk, num_experts, + scale_token_stride, scale_hidden_stride, + rdma_buffer_ptr, config.num_max_rdma_chunked_send_tokens, config.num_max_rdma_chunked_recv_tokens, + buffer_ptrs_gpu, config.num_max_nvl_chunked_send_tokens, config.num_max_nvl_chunked_recv_tokens, + rank, num_ranks, cached_mode, + comm_stream, num_channels, low_latency_mode, nixl_ctx->gpu); + + HOST_LOG_DEBUG("internode_dispatch finished"); + + // Wait streams + std::optional event; + if (async) { + event = EventHandle(comm_stream); + for (auto& t: {x, is_token_in_rank, recv_x, + rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum}) { + t.record_stream(comm_stream); + if (allocate_on_comm_stream) + t.record_stream(compute_stream); + } + for (auto& to: {x_scales, topk_idx, topk_weights, + num_tokens_per_rank, num_tokens_per_rdma_rank, num_tokens_per_expert, + cached_rdma_channel_prefix_matrix, cached_recv_rdma_rank_prefix_sum, + cached_gbl_channel_prefix_matrix, cached_recv_gbl_rank_prefix_sum, + recv_topk_idx, recv_topk_weights, recv_x_scales, + recv_rdma_channel_prefix_matrix, recv_gbl_channel_prefix_matrix, send_rdma_head, send_nvl_head, + recv_src_meta}) { + to.has_value() ? to->record_stream(comm_stream) : void(); + if (allocate_on_comm_stream) + to.has_value() ? to->record_stream(compute_stream) : void(); + } + } else { + stream_wait(compute_stream, comm_stream); + } + + // Switch back compute stream + if (allocate_on_comm_stream) + at::cuda::setCurrentCUDAStream(compute_stream); + + // Return values + return {recv_x, recv_x_scales, recv_topk_idx, recv_topk_weights, num_recv_tokens_per_expert_list, + rdma_channel_prefix_matrix, gbl_channel_prefix_matrix, + recv_rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, + recv_gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum, + recv_src_meta, send_rdma_head, send_nvl_head, event}; +} + +std::tuple, std::optional> +Buffer::internode_combine(const torch::Tensor& x, const std::optional& topk_weights, + const std::optional& bias_0, const std::optional& bias_1, + const torch::Tensor& src_meta, const torch::Tensor& is_combined_token_in_rank, + const torch::Tensor& rdma_channel_prefix_matrix, const torch::Tensor& rdma_rank_prefix_sum, const torch::Tensor& gbl_channel_prefix_matrix, + const torch::Tensor& combined_rdma_head, const torch::Tensor& combined_nvl_head, + const Config& config, std::optional& previous_event, bool async, bool allocate_on_comm_stream) { + const int num_channels = config.num_sms / 2; + HOST_LOG_DEBUG("internode_combine | num_channels: %d", num_channels); + EP_HOST_ASSERT(config.num_sms % 2 == 0); + + // Shape and contiguous checks + EP_HOST_ASSERT(x.dim() == 2 and x.is_contiguous()); + EP_HOST_ASSERT(src_meta.dim() == 2 and src_meta.is_contiguous() and src_meta.scalar_type() == torch::kByte); + EP_HOST_ASSERT(is_combined_token_in_rank.dim() == 2 and is_combined_token_in_rank.is_contiguous() and is_combined_token_in_rank.scalar_type() == torch::kBool); + EP_HOST_ASSERT(rdma_channel_prefix_matrix.dim() == 2 and rdma_channel_prefix_matrix.is_contiguous() and rdma_channel_prefix_matrix.scalar_type() == torch::kInt32); + EP_HOST_ASSERT(rdma_rank_prefix_sum.dim() == 1 and rdma_rank_prefix_sum.is_contiguous() and rdma_rank_prefix_sum.scalar_type() == torch::kInt32); + EP_HOST_ASSERT(gbl_channel_prefix_matrix.dim() == 2 and gbl_channel_prefix_matrix.is_contiguous() and gbl_channel_prefix_matrix.scalar_type() == torch::kInt32); + EP_HOST_ASSERT(combined_rdma_head.dim() == 2 and combined_rdma_head.is_contiguous() and combined_rdma_head.scalar_type() == torch::kInt32); + EP_HOST_ASSERT(combined_nvl_head.dim() == 2 and combined_nvl_head.is_contiguous() and combined_nvl_head.scalar_type() == torch::kInt32); + + auto num_tokens = static_cast(x.size(0)), hidden = static_cast(x.size(1)), hidden_int4 = static_cast(x.size(1) * x.element_size() / sizeof(int4)); + auto num_combined_tokens = static_cast(is_combined_token_in_rank.size(0)); + EP_HOST_ASSERT((hidden * x.element_size()) % sizeof(int4) == 0); + EP_HOST_ASSERT(src_meta.size(1) == internode::get_source_meta_bytes()); + EP_HOST_ASSERT(is_combined_token_in_rank.size(1) == num_ranks); + EP_HOST_ASSERT(rdma_channel_prefix_matrix.size(0) == num_rdma_ranks and rdma_channel_prefix_matrix.size(1) == num_channels); + EP_HOST_ASSERT(rdma_rank_prefix_sum.size(0) == num_rdma_ranks); + EP_HOST_ASSERT(gbl_channel_prefix_matrix.size(0) == num_ranks and gbl_channel_prefix_matrix.size(1) == num_channels); + EP_HOST_ASSERT(combined_rdma_head.dim() == 2 and combined_rdma_head.size(0) == num_combined_tokens and combined_rdma_head.size(1) == num_rdma_ranks); + EP_HOST_ASSERT(combined_nvl_head.dim() == 2 and combined_nvl_head.size(1) == NUM_MAX_NVL_PEERS); + + // Allocate all tensors on comm stream if set + // NOTES: do not allocate tensors upfront! + auto compute_stream = at::cuda::getCurrentCUDAStream(); + if (allocate_on_comm_stream) { + EP_HOST_ASSERT(previous_event.has_value() and async); + at::cuda::setCurrentCUDAStream(comm_stream); + } + + // Wait previous tasks to be finished + if (previous_event.has_value()) { + stream_wait(comm_stream, previous_event.value()); + } else { + stream_wait(comm_stream, compute_stream); + } + + // Top-k checks + int num_topk = 0; + auto combined_topk_weights = std::optional(); + float* topk_weights_ptr = nullptr; + float* combined_topk_weights_ptr = nullptr; + if (topk_weights.has_value()) { + EP_HOST_ASSERT(topk_weights->dim() == 2 and topk_weights->is_contiguous()); + EP_HOST_ASSERT(topk_weights->size(0) == num_tokens); + EP_HOST_ASSERT(topk_weights->scalar_type() == torch::kFloat32); + num_topk = static_cast(topk_weights->size(1)); + topk_weights_ptr = topk_weights->data_ptr(); + combined_topk_weights = torch::empty({num_combined_tokens, num_topk}, topk_weights->options()); + combined_topk_weights_ptr = combined_topk_weights->data_ptr(); + } + + // Extra check for avoid-dead-lock design + EP_HOST_ASSERT(config.num_max_nvl_chunked_recv_tokens % num_rdma_ranks == 0); + EP_HOST_ASSERT(config.num_max_nvl_chunked_send_tokens <= config.num_max_nvl_chunked_recv_tokens / num_rdma_ranks); + + // Launch barrier and reset queue head and tail + internode::cached_notify(hidden_int4, 0, 0, num_topk, + num_ranks, num_channels, + num_combined_tokens, combined_rdma_head.data_ptr(), + rdma_channel_prefix_matrix.data_ptr(), rdma_rank_prefix_sum.data_ptr(), combined_nvl_head.data_ptr(), + rdma_buffer_ptr, config.num_max_rdma_chunked_recv_tokens, + buffer_ptrs_gpu, config.num_max_nvl_chunked_recv_tokens, + barrier_signal_ptrs_gpu, rank, comm_stream, + config.get_rdma_buffer_size_hint(hidden_int4 * sizeof(int4), num_ranks), + num_nvl_bytes, false, low_latency_mode, nixl_ctx->gpu); + + // Assign bias pointers + auto bias_opts = std::vector>({bias_0, bias_1}); + void* bias_ptrs[2] = {nullptr, nullptr}; + for (int i = 0; i < 2; ++ i) if (bias_opts[i].has_value()) { + auto bias = bias_opts[i].value(); + EP_HOST_ASSERT(bias.dim() == 2 and bias.is_contiguous()); + EP_HOST_ASSERT(bias.scalar_type() == x.scalar_type()); + EP_HOST_ASSERT(bias.size(0) == num_combined_tokens and bias.size(1) == hidden); + bias_ptrs[i] = bias.data_ptr(); + } + + // Launch data combine + auto combined_x = torch::empty({num_combined_tokens, hidden}, x.options()); + internode::combine(at::cuda::ScalarTypeToCudaDataType(x.scalar_type()), + combined_x.data_ptr(), combined_topk_weights_ptr, + is_combined_token_in_rank.data_ptr(), + x.data_ptr(), topk_weights_ptr, bias_ptrs[0], bias_ptrs[1], + combined_rdma_head.data_ptr(), combined_nvl_head.data_ptr(), + src_meta.data_ptr(), rdma_channel_prefix_matrix.data_ptr(), rdma_rank_prefix_sum.data_ptr(), gbl_channel_prefix_matrix.data_ptr(), + num_tokens, num_combined_tokens, hidden, num_topk, + rdma_buffer_ptr, config.num_max_rdma_chunked_send_tokens, config.num_max_rdma_chunked_recv_tokens, + buffer_ptrs_gpu, config.num_max_nvl_chunked_send_tokens, config.num_max_nvl_chunked_recv_tokens, + rank, num_ranks, comm_stream, num_channels, low_latency_mode, nixl_ctx->gpu); + + HOST_LOG_DEBUG("internode_combine finished"); + + // Wait streams + std::optional event; + if (async) { + event = EventHandle(comm_stream); + for (auto& t: {x, src_meta, + is_combined_token_in_rank, rdma_channel_prefix_matrix, rdma_rank_prefix_sum, gbl_channel_prefix_matrix, + combined_x, combined_rdma_head, combined_nvl_head}) { + t.record_stream(comm_stream); + if (allocate_on_comm_stream) + t.record_stream(compute_stream); + } + for (auto& to: {topk_weights, combined_topk_weights, bias_0, bias_1}) { + to.has_value() ? to->record_stream(comm_stream) : void(); + if (allocate_on_comm_stream) + to.has_value() ? to->record_stream(compute_stream) : void(); + } + } else { + stream_wait(compute_stream, comm_stream); + } + + // Switch back compute stream + if (allocate_on_comm_stream) + at::cuda::setCurrentCUDAStream(compute_stream); + + // Return values + return {combined_x, combined_topk_weights, event}; +} + + + std::tuple, torch::Tensor, torch::Tensor, torch::Tensor, std::optional, std::optional>> Buffer::dispatch(const torch::Tensor& x, const torch::Tensor& topk_idx, const std::optional& cumulative_local_expert_recv_stats, @@ -683,7 +1276,21 @@ void Buffer::_nixl_ep_gpu_ctx_update() { nixl_ctx->gpu.local_barrier_buffer = sync_buffer_ptr; nixl_ctx->gpu.local_barrier_cnt = local_barrier_cnt_ptr; nixl_ctx->gpu.num_ranks = num_ranks; + nixl_ctx->gpu.num_rdma_ranks = num_rdma_ranks; nixl_ctx->gpu.rank = rank; + + /* Initialize internode barrier counters for high-throughput mode */ + if (!low_latency_mode) { + // local_barrier_counter_ptr points to our local_barrier_counter which remote peers write to + nixl_ctx->gpu.local_barrier_counter_ptr = local_barrier_counter; + + // last_barrier_counter tracks barrier epochs (allocate once) + if (nixl_ctx->gpu.last_barrier_counter == nullptr) { + CUDA_CHECK(cudaMalloc(&nixl_ctx->gpu.last_barrier_counter, sizeof(uint64_t))); + uint64_t zero = 0; + CUDA_CHECK(cudaMemcpy(nixl_ctx->gpu.last_barrier_counter, &zero, sizeof(uint64_t), cudaMemcpyHostToDevice)); + } + } } void Buffer::_nixl_ep_context_init() { @@ -747,6 +1354,11 @@ void Buffer::_nixl_agent_init() { barrier_cnt_dlist.addDesc(nixlBlobDesc((uintptr_t)(local_barrier_cnt_ptr), max_num_ranks * sizeof(int), get_local_device_id(), "")); EP_HOST_ASSERT(agent->registerMem(barrier_cnt_dlist) == NIXL_SUCCESS); + /* Register internode barrier counter (for high-throughput mode) */ + nixl_reg_dlist_t internode_barrier_dlist(VRAM_SEG); + internode_barrier_dlist.addDesc(nixlBlobDesc((uintptr_t)(local_barrier_counter), sizeof(uint64_t), get_local_device_id(), "")); + EP_HOST_ASSERT(agent->registerMem(internode_barrier_dlist) == NIXL_SUCCESS); + size_t signal_size = 0; EP_HOST_ASSERT(nixl_agent_info->agent->getGpuSignalSize(signal_size, &nixl_agent_info->extra_params) == NIXL_SUCCESS); EP_HOST_ASSERT(signal_size == sizeof(uint64_t)); @@ -763,25 +1375,72 @@ void Buffer::_nixl_agent_init() { void Buffer::_nixl_ep_batches_prepare(const std::vector& ranks) { nixl_status_t status; - for (int j : ranks) { - if (j == rank) continue; // Skip self - if (nixl_ctx->gpu_batch_reqs[j]) continue; // Skip if already exported - nixl_xfer_dlist_t src_vram(VRAM_SEG); - src_vram.addDesc(nixlBlobDesc((uintptr_t)(rdma_buffer_ptr), num_rdma_bytes, get_local_device_id(), "")); - nixl_xfer_dlist_t dst_vram(VRAM_SEG); - dst_vram.addDesc(nixlBlobDesc((uintptr_t)(nixl_peer_info[j].rdma_buffer_ptr), num_rdma_bytes, nixl_peer_info[j].device_id, "")); - nixl_opt_args_t extra_params = {}; - extra_params.backends.push_back(nixl_agent_info->backend); - status = nixl_agent_info->agent->createXferReq(NIXL_WRITE, src_vram, dst_vram, nixl_agent_info->remote_agent_names[j], nixl_ctx->cpu_batch_reqs[j], &extra_params); - EP_HOST_ASSERT(status == NIXL_SUCCESS); - EP_HOST_ASSERT(nixl_agent_info->agent->createGpuXferReq(*nixl_ctx->cpu_batch_reqs[j], nixl_ctx->gpu_batch_reqs[j]) == NIXL_SUCCESS); + if (low_latency_mode) { + // Low-latency mode: create requests for each global rank + for (int j : ranks) { + if (j == rank) continue; // Skip self + if (nixl_ctx->gpu_batch_reqs[j]) continue; // Skip if already exported + + nixl_xfer_dlist_t src_vram(VRAM_SEG); + src_vram.addDesc(nixlBlobDesc((uintptr_t)(rdma_buffer_ptr), num_rdma_bytes, get_local_device_id(), "")); + nixl_xfer_dlist_t dst_vram(VRAM_SEG); + dst_vram.addDesc(nixlBlobDesc((uintptr_t)(nixl_peer_info[j].rdma_buffer_ptr), num_rdma_bytes, nixl_peer_info[j].device_id, "")); + nixl_opt_args_t extra_params = {}; + extra_params.backends.push_back(nixl_agent_info->backend); + status = nixl_agent_info->agent->createXferReq(NIXL_WRITE, src_vram, dst_vram, nixl_agent_info->remote_agent_names[j], nixl_ctx->cpu_batch_reqs[j], &extra_params); + EP_HOST_ASSERT(status == NIXL_SUCCESS); + EP_HOST_ASSERT(nixl_agent_info->agent->createGpuXferReq(*nixl_ctx->cpu_batch_reqs[j], nixl_ctx->gpu_batch_reqs[j]) == NIXL_SUCCESS); + + // Low-latency barrier: write to sync_buffer_ptr (int array), indexed by global rank + nixl_xfer_dlist_t src_vram_ll(VRAM_SEG); + src_vram_ll.addDesc(nixlBlobDesc((uintptr_t)(local_barrier_cnt_ptr), max_num_ranks * sizeof(int), get_local_device_id(), "")); + nixl_xfer_dlist_t dst_vram_ll(VRAM_SEG); + dst_vram_ll.addDesc(nixlBlobDesc((uintptr_t)(nixl_peer_info[j].sync_buffer_ptr), max_num_ranks * sizeof(int), nixl_peer_info[j].device_id, "")); + EP_HOST_ASSERT(nixl_agent_info->agent->createXferReq(NIXL_WRITE, src_vram_ll, dst_vram_ll, nixl_agent_info->remote_agent_names[j], nixl_ctx->cpu_barrier_reqs[j], &extra_params) == NIXL_SUCCESS); + EP_HOST_ASSERT(nixl_agent_info->agent->createGpuXferReq(*nixl_ctx->cpu_barrier_reqs[j], nixl_ctx->gpu_barrier_reqs[j]) == NIXL_SUCCESS); + } + } else { + // Internode mode: create requests indexed by RDMA rank, targeting peer (same nvl_rank on remote node) + for (int remote_rdma_rank = 0; remote_rdma_rank < num_rdma_ranks; remote_rdma_rank++) { + if (remote_rdma_rank == rdma_rank) continue; // Skip self RDMA rank + + int remote_rank = nvl_rank + remote_rdma_rank * NUM_MAX_NVL_PEERS; + + // Check if we have peer info for this rank + if (remote_rank >= static_cast(nixl_peer_info.size()) || nixl_peer_info[remote_rank].barrier_ptr == nullptr) { + continue; + } + + nixl_opt_args_t extra_params = {}; + extra_params.backends.push_back(nixl_agent_info->backend); + + // Create batch request (indexed by RDMA rank) + if (!nixl_ctx->gpu_batch_reqs[remote_rdma_rank]) { + nixl_xfer_dlist_t src_vram(VRAM_SEG); + src_vram.addDesc(nixlBlobDesc((uintptr_t)(rdma_buffer_ptr), num_rdma_bytes, get_local_device_id(), "")); + nixl_xfer_dlist_t dst_vram(VRAM_SEG); + dst_vram.addDesc(nixlBlobDesc((uintptr_t)(nixl_peer_info[remote_rank].rdma_buffer_ptr), num_rdma_bytes, nixl_peer_info[remote_rank].device_id, "")); + + status = nixl_agent_info->agent->createXferReq(NIXL_WRITE, src_vram, dst_vram, + nixl_agent_info->remote_agent_names[remote_rank], nixl_ctx->cpu_batch_reqs[remote_rdma_rank], &extra_params); + EP_HOST_ASSERT(status == NIXL_SUCCESS); + EP_HOST_ASSERT(nixl_agent_info->agent->createGpuXferReq(*nixl_ctx->cpu_batch_reqs[remote_rdma_rank], + nixl_ctx->gpu_batch_reqs[remote_rdma_rank]) == NIXL_SUCCESS); + } - nixl_xfer_dlist_t src_vram_ll(VRAM_SEG); - src_vram_ll.addDesc(nixlBlobDesc((uintptr_t)(local_barrier_cnt_ptr), max_num_ranks * sizeof(int), get_local_device_id(), "")); - nixl_xfer_dlist_t dst_vram_ll(VRAM_SEG); - dst_vram_ll.addDesc(nixlBlobDesc((uintptr_t)(nixl_peer_info[j].sync_buffer_ptr), max_num_ranks * sizeof(int), nixl_peer_info[j].device_id, "")); - EP_HOST_ASSERT(nixl_agent_info->agent->createXferReq(NIXL_WRITE, src_vram_ll, dst_vram_ll, nixl_agent_info->remote_agent_names[j], nixl_ctx->cpu_barrier_reqs[j], &extra_params) == NIXL_SUCCESS); - EP_HOST_ASSERT(nixl_agent_info->agent->createGpuXferReq(*nixl_ctx->cpu_barrier_reqs[j], nixl_ctx->gpu_barrier_reqs[j]) == NIXL_SUCCESS); + // Create barrier request (indexed by RDMA rank) + if (!nixl_ctx->gpu_barrier_reqs[remote_rdma_rank]) { + nixl_xfer_dlist_t dummy_src(VRAM_SEG); + dummy_src.addDesc(nixlBlobDesc((uintptr_t)(local_barrier_counter), sizeof(uint64_t), get_local_device_id(), "")); + nixl_xfer_dlist_t barrier_dst(VRAM_SEG); + barrier_dst.addDesc(nixlBlobDesc((uintptr_t)(nixl_peer_info[remote_rank].barrier_ptr), sizeof(uint64_t), nixl_peer_info[remote_rank].device_id, "")); + + EP_HOST_ASSERT(nixl_agent_info->agent->createXferReq(NIXL_WRITE, dummy_src, barrier_dst, + nixl_agent_info->remote_agent_names[remote_rank], nixl_ctx->cpu_barrier_reqs[remote_rdma_rank], &extra_params) == NIXL_SUCCESS); + EP_HOST_ASSERT(nixl_agent_info->agent->createGpuXferReq(*nixl_ctx->cpu_barrier_reqs[remote_rdma_rank], + nixl_ctx->gpu_barrier_reqs[remote_rdma_rank]) == NIXL_SUCCESS); + } + } } } @@ -888,25 +1547,40 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.doc() = "NIXL_EP: an efficient expert-parallel communication library"; m.def("get_rdma_size_hint", &nixl_ep::get_rdma_size_hint); + pybind11::class_(m, "Config") + .def(pybind11::init(), + py::arg("num_sms") = 20, + py::arg("num_max_nvl_chunked_send_tokens") = 6, py::arg("num_max_nvl_chunked_recv_tokens") = 256, + py::arg("num_max_rdma_chunked_send_tokens") = 6, py::arg("num_max_rdma_chunked_recv_tokens") = 256) + .def("get_nvl_buffer_size_hint", &nixl_ep::Config::get_nvl_buffer_size_hint) + .def("get_rdma_buffer_size_hint", &nixl_ep::Config::get_rdma_buffer_size_hint); + pybind11::class_(m, "EventHandle") .def(pybind11::init<>()) .def("current_stream_wait", &nixl_ep::EventHandle::current_stream_wait); pybind11::class_(m, "Buffer") - .def(pybind11::init()) + .def(pybind11::init()) .def("update_memory_buffers", &nixl_ep::Buffer::update_memory_buffers) .def("barrier", &nixl_ep::Buffer::barrier) - .def("connect_ranks", [](nixl_ep::Buffer &buffer, const std::vector& remote_ranks, const std::optional>& remote_mds) { - buffer.connect_ranks(remote_ranks, nixl_ep::convert_mds(remote_mds)); - }, py::arg("remote_ranks"), py::arg("remote_mds") = std::nullopt) + .def("connect_ranks", [](nixl_ep::Buffer &buffer, const std::vector& remote_ranks, const std::optional>& remote_mds, const std::vector> &all_gathered_handles) { + buffer.connect_ranks(remote_ranks, nixl_ep::convert_mds(remote_mds), all_gathered_handles); + }, py::arg("remote_ranks"), py::arg("remote_mds") = std::nullopt, py::arg("ipc_handles") = std::vector>{}) .def("disconnect_ranks", &nixl_ep::Buffer::disconnect_ranks) .def("is_available", &nixl_ep::Buffer::is_available) + .def("get_num_rdma_ranks", &nixl_ep::Buffer::get_num_rdma_ranks) + .def("get_rdma_rank", &nixl_ep::Buffer::get_rdma_rank) + .def("get_root_rdma_rank", &nixl_ep::Buffer::get_root_rdma_rank) .def("get_local_device_id", &nixl_ep::Buffer::get_local_device_id) + .def("get_local_ipc_handle", &nixl_ep::Buffer::get_local_ipc_handle) .def("get_local_buffer_tensor", &nixl_ep::Buffer::get_local_buffer_tensor) .def("get_comm_stream", &nixl_ep::Buffer::get_comm_stream) .def("destroy", &nixl_ep::Buffer::destroy) + .def("get_dispatch_layout", &nixl_ep::Buffer::get_dispatch_layout) .def("dispatch", &nixl_ep::Buffer::dispatch) .def("combine", &nixl_ep::Buffer::combine) + .def("internode_dispatch", &nixl_ep::Buffer::internode_dispatch) + .def("internode_combine", &nixl_ep::Buffer::internode_combine) .def("update_mask_buffer", &nixl_ep::Buffer::update_mask_buffer) .def("query_mask_buffer", &nixl_ep::Buffer::query_mask_buffer) .def("clean_mask_buffer", &nixl_ep::Buffer::clean_mask_buffer) diff --git a/examples/device/ep/csrc/nixl_ep.hpp b/examples/device/ep/csrc/nixl_ep.hpp index 02e9e04916..6848d2e898 100644 --- a/examples/device/ep/csrc/nixl_ep.hpp +++ b/examples/device/ep/csrc/nixl_ep.hpp @@ -61,6 +61,7 @@ struct NixlPeerInfo { void *rdma_buffer_ptr; cudaIpcMemHandle_t rdma_ipc_handle; int* sync_buffer_ptr; + uint64_t* barrier_ptr; // For internode barrier (high-throughput mode) int device_id; int rank; }; @@ -86,12 +87,20 @@ struct nixl_ep_ctx { std::vector cpu_barrier_reqs; // [num_peers] std::vector gpu_barrier_reqs; // [num_peers] std::vector rdma_p2p_ptrs; // [num_ranks] - ep_kernels::gpu_nixl_ctx gpu; + nixl_ep::gpu_nixl_ctx gpu; }; struct Buffer { + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS == 8, "The number of maximum NVLink peers must be 8"); + private: int buffer_idx = 0; // Double buffering index + bool low_latency_mode = false; + + // NVLink Buffer + int64_t num_nvl_bytes; + void* buffer_ptrs[NUM_MAX_NVL_PEERS] = {nullptr}; + void** buffer_ptrs_gpu = nullptr; // RDMA Buffer int64_t num_rdma_bytes; @@ -106,9 +115,10 @@ struct Buffer { // Device info and communication int device_id; int num_device_sms; - int rank; - int num_ranks; + int rank, rdma_rank, nvl_rank; + int num_ranks, num_rdma_ranks, num_nvl_ranks; std::vector remote_ranks; /* global ranks */ + cudaIpcMemHandle_t ipc_handles[NUM_MAX_NVL_PEERS]; // Stream for communication at::cuda::CUDAStream comm_stream; @@ -121,14 +131,33 @@ struct Buffer { // After `destroy()` be called, this flag will be true bool destroyed = false; + // Barrier signals + int* barrier_signal_ptrs[NUM_MAX_NVL_PEERS] = {nullptr}; + int** barrier_signal_ptrs_gpu = nullptr; + // Workspace void* workspace = nullptr; + // Host-side MoE info + volatile int* moe_recv_counter = nullptr; + int* moe_recv_counter_mapped = nullptr; + + // Host-side expert-level MoE info + volatile int* moe_recv_expert_counter = nullptr; + int* moe_recv_expert_counter_mapped = nullptr; + + // Host-side RDMA-level MoE info + volatile int* moe_recv_rdma_counter = nullptr; + int* moe_recv_rdma_counter_mapped = nullptr; + std::unique_ptr nixl_agent_info; std::vector nixl_peer_info; NixlPeerInfo my_peer_info; uint64_t max_num_ranks; int max_experts_per_rank; + uint64_t* last_barrier_counter = nullptr; + uint64_t* local_barrier_counter = nullptr; + std::unique_ptr nixl_ctx = nullptr; /* Common private funcs */ @@ -150,30 +179,74 @@ struct Buffer { void _nixl_ep_batches_cleanup(const std::vector& ranks_to_remove); void _nixl_ep_p2p_ptrs_cleanup(const std::vector& ranks_to_remove); + /* Internode mode private funcs */ + void _nixl_internode_init(); + void _nixl_internode_local_data_init(); + void _nixl_remote_counters_prepare(); + void _nixl_internode_batches_prepare(); + void _nixl_kernels_params_free(); + void _nixl_establish_new_connections(); + void _ipc_handles_sync(const std::vector> &all_gathered_handles); + public: - Buffer(int rank, bool explicitly_destroy, bool enable_shrink); + Buffer(int rank, bool low_latency_mode, bool explicitly_destroy, bool enable_shrink); - void update_memory_buffers(int num_ranks, int max_experts_per_rank, int64_t num_rdma_bytes); + void update_memory_buffers(int num_ranks, int max_experts_per_rank, int64_t num_nvl_bytes, int64_t num_rdma_bytes); - void connect_ranks(const std::vector& remote_ranks_list, const std::optional>& remote_mds = std::nullopt); + void connect_ranks(const std::vector& remote_ranks_list, const std::optional>& remote_mds = std::nullopt, const std::vector>& all_gathered_handles = {}); void disconnect_ranks(const std::vector& remote_ranks_list); - void init(int num_ranks, int max_experts_per_rank, int64_t num_rdma_bytes); + void init(int num_ranks, int max_experts_per_rank, int64_t num_nvl_bytes, int64_t num_rdma_bytes); ~Buffer() noexcept(false); bool is_available() const; + bool is_internode_available() const; + + int get_num_rdma_ranks() const; + + int get_rdma_rank() const; + + int get_root_rdma_rank(bool global) const; + int get_local_device_id() const; - torch::Tensor get_local_buffer_tensor(const pybind11::object& dtype, int64_t offset) const; + pybind11::bytearray get_local_ipc_handle() const; + + torch::Tensor get_local_buffer_tensor(const pybind11::object& dtype, int64_t offset, bool use_rdma_buffer = false) const; torch::Stream get_comm_stream() const; + void sync(const std::vector& device_ids, const std::vector>& all_gathered_handles, const std::optional& root_unique_id_opt); + void destroy(); + std::tuple, torch::Tensor, torch::Tensor, std::optional> + get_dispatch_layout(const torch::Tensor& topk_idx, int num_experts, std::optional& previous_event, + bool async, bool allocate_on_comm_stream); + + std::tuple, std::optional, std::optional, std::vector, torch::Tensor, torch::Tensor, std::optional, torch::Tensor, std::optional, torch::Tensor, std::optional, std::optional, std::optional, std::optional> + internode_dispatch(const torch::Tensor& x, const std::optional& x_scales, + const std::optional& topk_idx, const std::optional& topk_weights, + const std::optional& num_tokens_per_rank, const std::optional& num_tokens_per_rdma_rank, + const torch::Tensor& is_token_in_rank, const std::optional& num_tokens_per_expert, + int cached_num_recv_tokens, int cached_num_rdma_recv_tokens, + const std::optional& cached_rdma_channel_prefix_matrix, const std::optional& cached_recv_rdma_rank_prefix_sum, + const std::optional& cached_gbl_channel_prefix_matrix, const std::optional& cached_recv_gbl_rank_prefix_sum, + int expert_alignment, const Config& config, std::optional& previous_event, bool async, bool allocate_on_comm_stream); + + std::tuple, std::optional> + internode_combine(const torch::Tensor& x, const std::optional& topk_weights, + const std::optional& bias_0, const std::optional& bias_1, + const torch::Tensor& src_meta, const torch::Tensor& is_combined_token_in_rank, + const torch::Tensor& rdma_channel_prefix_matrix, const torch::Tensor& rdma_rank_prefix_sum, const torch::Tensor& gbl_channel_prefix_matrix, + const torch::Tensor& combined_rdma_head, const torch::Tensor& combined_nvl_head, + const Config& config, std::optional& previous_event, bool async, bool allocate_on_comm_stream); + void clean_buffer(int num_max_dispatch_tokens_per_rank, int hidden, int num_experts); + std::tuple, torch::Tensor, torch::Tensor, torch::Tensor, std::optional, std::optional>> dispatch(const torch::Tensor& x, const torch::Tensor& topk_idx, diff --git a/examples/device/ep/meson.build b/examples/device/ep/meson.build index 1e82cb2a8e..189cbe64eb 100644 --- a/examples/device/ep/meson.build +++ b/examples/device/ep/meson.build @@ -64,6 +64,9 @@ endif nixl_ep_sources = [ 'csrc/nixl_ep.cpp', 'csrc/kernels/nixl_ep.cu', + 'csrc/kernels/internode.cu', + 'csrc/kernels/layout.cu', + 'csrc/kernels/runtime.cu', ] nixl_ep_inc_dirs = [ @@ -107,6 +110,13 @@ topk_idx_bits = meson.get_external_property('topk_idx_bits', '64') nixl_ep_cpp_args += ['-DTOPK_IDX_BITS=' + topk_idx_bits] nixl_ep_cuda_args += ['-DTOPK_IDX_BITS=' + topk_idx_bits] +# Enable debug logs via ENABLE_DEBUG_LOGS environment variable +env_debug_logs = run_command('sh', '-c', 'echo $ENABLE_DEBUG_LOGS', check: false).stdout().strip() +if env_debug_logs != '' + nixl_ep_cpp_args += ['-DENABLE_DEBUG_LOGS'] + nixl_ep_cuda_args += ['-DENABLE_DEBUG_LOGS'] +endif + # NVCC 12.9 workaround for bug (https://nvbugspro.nvidia.com/bug/5595631) # This workaround is needed because NVCC 12.9 has a bug which fails UCX compilation nvcc = meson.get_compiler('cuda') diff --git a/examples/device/ep/nixl_ep/__init__.py b/examples/device/ep/nixl_ep/__init__.py index 719e149b2f..488ad2fc29 100644 --- a/examples/device/ep/nixl_ep/__init__.py +++ b/examples/device/ep/nixl_ep/__init__.py @@ -25,5 +25,6 @@ from .utils import EventOverlap topk_idx_t = getattr(_nixl_ep_cpp, "topk_idx_t", torch.int64) +Config = _nixl_ep_cpp.Config -__all__ = ["Buffer", "EventOverlap"] +__all__ = ["Buffer", "EventOverlap", "Config"] diff --git a/examples/device/ep/nixl_ep/buffer.py b/examples/device/ep/nixl_ep/buffer.py index 2102525571..15a9aec3ed 100644 --- a/examples/device/ep/nixl_ep/buffer.py +++ b/examples/device/ep/nixl_ep/buffer.py @@ -30,7 +30,8 @@ from . import nixl_ep_cpp # noinspection PyUnresolvedReferences -from .nixl_ep_cpp import EventHandle +from .nixl_ep_cpp import Config, EventHandle +from .utils import check_nvlink_connections from .utils import EventOverlap if TYPE_CHECKING: @@ -55,6 +56,7 @@ def __init__( nvlink_backend: Literal["nixl", "ipc", "none"] = "nixl", explicitly_destroy: bool = False, rank: int = 0, + low_latency_mode: bool = False, enable_shrink: bool = False, group: Optional[dist.ProcessGroup] = None, comm: Optional["mpi4py.MPI.Comm"] = None, @@ -69,12 +71,15 @@ def __init__( otherwise, the resources will be released by the destructor. Note: Releasing resources in the destructor may cause Python's exception handling process to hang. rank: the rank number. + low_latency_mode: whether to enable low-latency mode. group: the communication group (optional). comm: the mpi4py.MPI.Comm communicator to use in case the group parameter is absent (optional). tcp_store_group: TCPStore for metadata exchange (optional). """ self.rank = rank self.group_size = 0 # Will be updated by `update_memory_buffers` + self.low_latency_mode = low_latency_mode + self.explicitly_destroy = explicitly_destroy self.group = group self.comm = comm @@ -88,7 +93,10 @@ def __init__( if nvlink_backend != "nixl": os.environ["UCX_TLS"] = "^cuda_ipc" - self.runtime = nixl_ep_cpp.Buffer(self.rank, explicitly_destroy, enable_shrink) + if self.group is not None: + check_nvlink_connections(self.group) + + self.runtime = nixl_ep_cpp.Buffer(self.rank, low_latency_mode, explicitly_destroy, enable_shrink) def destroy(self): """ @@ -163,7 +171,7 @@ def get_comm_stream(self) -> torch.Stream: ) def get_local_buffer_tensor( - self, dtype: torch.dtype, size: Optional[torch.Size] = None, offset: int = 0 + self, dtype: torch.dtype, size: Optional[torch.Size] = None, offset: int = 0, use_rdma_buffer: bool = False ) -> torch.Tensor: """ Get the raw buffer (slice supported) as a PyTorch tensor. @@ -172,8 +180,9 @@ def get_local_buffer_tensor( dtype: the data type (PyTorch `dtype`) for the tensor. size: the slice size (by elements) to get from the buffer. offset: the offset of the beginning element. + use_rdma_buffer: whether to return the RDMA buffer. """ - tensor = self.runtime.get_local_buffer_tensor(dtype, offset) + tensor = self.runtime.get_local_buffer_tensor(dtype, offset, use_rdma_buffer) if size is None: return tensor @@ -190,6 +199,91 @@ def _unpack_bias(bias: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]): bias_0, bias_1 = bias return bias_0, bias_1 + @staticmethod + def get_dispatch_config(num_ranks: int) -> Config: + """ + Get a recommended dispatch config. + + Argument: + num_ranks: the number of ranks. + + Returns: + config: the recommended config. + """ + + # TODO: automatically tune + config_map = { + 2: Config(Buffer.num_sms, 24, 256, 6, 128), + 4: Config(Buffer.num_sms, 6, 256, 6, 128), + 8: Config(Buffer.num_sms, 6, 256, 6, 128), + 16: Config(Buffer.num_sms, 36, 288, 20, 128), + 24: Config(Buffer.num_sms, 8, 288, 32, 128), + 32: Config(Buffer.num_sms, 32, 288, 32, 128), + 64: Config(Buffer.num_sms, 20, 288, 28, 128), + 128: Config(Buffer.num_sms, 20, 560, 32, 128), + 144: Config(Buffer.num_sms, 32, 720, 12, 128), + 160: Config(Buffer.num_sms, 28, 720, 12, 128), + } + assert num_ranks in config_map, f'Unsupported number of EP ranks: {num_ranks}' + return config_map[num_ranks] + + @staticmethod + def get_combine_config(num_ranks: int) -> Config: + """ + Get a recommended combine config. + + Argument: + num_ranks: the number of ranks. + + Returns: + config: the recommended config. + """ + + # TODO: automatically tune + config_map = { + 2: Config(Buffer.num_sms, 10, 256, 6, 128), + 4: Config(Buffer.num_sms, 9, 256, 6, 128), + 8: Config(Buffer.num_sms, 4, 256, 6, 128), + 16: Config(Buffer.num_sms, 4, 288, 12, 128), + 24: Config(Buffer.num_sms, 1, 288, 8, 128), + 32: Config(Buffer.num_sms, 1, 288, 8, 128), + 64: Config(Buffer.num_sms, 1, 288, 20, 128), + 128: Config(Buffer.num_sms, 1, 560, 12, 128), + 144: Config(Buffer.num_sms, 2, 720, 8, 128), + 160: Config(Buffer.num_sms, 2, 720, 8, 128), + } + assert num_ranks in config_map, f'Unsupported number of EP ranks: {num_ranks}' + return config_map[num_ranks] + + # noinspection PyTypeChecker + def get_dispatch_layout(self, topk_idx: torch.Tensor, num_experts: int, + previous_event: Optional[EventOverlap] = None, async_finish: bool = False, + allocate_on_comm_stream: bool = False) -> \ + Tuple[torch.Tensor, Optional[torch.Tensor], torch.Tensor, torch.Tensor, EventOverlap]: + """ + Calculate the layout required for later communication. + + Arguments: + topk_idx: `[num_tokens, num_topk]`, dtype must be `torch.int64`, the expert indices selected by each token, + `-1` means no selections. + num_experts: the number of experts. + previous_event: the event to wait before actually executing the kernel. + async_finish: the current stream will not wait for the communication kernels to be finished if set. + allocate_on_comm_stream: control whether all the allocated tensors' ownership to be on the communication stream. + + Returns: + num_tokens_per_rank: `[num_ranks]` with `torch.int`, the number of tokens to be sent to each rank. + num_tokens_per_rdma_rank: `[num_rdma_ranks]` with `torch.int`, the number of tokens to be sent to each RDMA + rank (with the same GPU index), return `None` for intranode settings. + num_tokens_per_expert: `[num_experts]` with `torch.int`, the number of tokens to be sent to each expert. + is_token_in_rank: `[num_tokens, num_ranks]` with `torch.bool`, whether a token be sent to a rank. + event: the event after executing the kernel (valid only if `async_finish` is set). + """ + num_tokens_per_rank, num_tokens_per_rdma_rank, num_tokens_per_expert, is_token_in_rank, event = \ + self.runtime.get_dispatch_layout(topk_idx, num_experts, getattr(previous_event, 'event', None), + async_finish, allocate_on_comm_stream) + return num_tokens_per_rank, num_tokens_per_rdma_rank, num_tokens_per_expert, is_token_in_rank, EventOverlap(event) + def clean_buffer( self, num_max_dispatch_tokens_per_rank: int, hidden: int, num_experts: int ) -> None: @@ -395,6 +489,90 @@ def combine( hook, ) + # noinspection PyTypeChecker + def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], + handle: Optional[Tuple] = None, + num_tokens_per_rank: Optional[torch.Tensor] = None, num_tokens_per_rdma_rank: Optional[torch.Tensor] = None, + is_token_in_rank: Optional[torch.Tensor] = None, num_tokens_per_expert: Optional[torch.Tensor] = None, + topk_idx: Optional[torch.Tensor] = None, topk_weights: Optional[torch.Tensor] = None, expert_alignment: int = 1, + config: Optional[Config] = None, + previous_event: Optional[EventOverlap] = None, async_finish: bool = False, + allocate_on_comm_stream: bool = False) -> \ + Tuple[Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor], Optional[torch.Tensor], + Optional[torch.Tensor], List[int], Tuple, EventOverlap]: + """ + Internode dispatch implementation, for more details, please refer to the `dispatch` docs. + Normally, you should not directly call this function. + """ + config = self.get_dispatch_config(self.group_size) if config is None else config + assert config is not None + + # Launch the kernel with cached or non-cached mode + x, x_scales = x if isinstance(x, tuple) else (x, None) + if handle is not None: + assert topk_idx is None and topk_weights is None + is_token_in_rank, \ + rdma_channel_prefix_matrix, gbl_channel_prefix_matrix, \ + recv_rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, recv_gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum, \ + recv_src_meta, send_rdma_head, send_nvl_head = handle + num_recv_tokens = recv_src_meta.size(0) + num_rdma_recv_tokens = send_nvl_head.size(0) + recv_x, recv_x_scales, _, _, _, _, _, _, _, _, _, _, _, _, event = self.runtime.internode_dispatch( + x, x_scales, topk_idx, topk_weights, + None, None, is_token_in_rank, None, + num_recv_tokens, num_rdma_recv_tokens, + rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum, + expert_alignment, config, getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream) + return (recv_x, recv_x_scales) if x_scales is not None else recv_x, None, None, None, None, EventOverlap(event) + else: + assert num_tokens_per_rank is not None and is_token_in_rank is not None and num_tokens_per_expert is not None + recv_x, recv_x_scales, recv_topk_idx, recv_topk_weights, num_recv_tokens_per_expert_list, \ + rdma_channel_prefix_matrix, gbl_channel_prefix_matrix, \ + recv_rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, \ + recv_gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum, \ + recv_src_meta, send_rdma_head, send_nvl_head, event = self.runtime.internode_dispatch( + x, x_scales, topk_idx, topk_weights, + num_tokens_per_rank, num_tokens_per_rdma_rank, is_token_in_rank, num_tokens_per_expert, + 0, 0, None, None, None, None, + expert_alignment, config, getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream) + handle = (is_token_in_rank, + rdma_channel_prefix_matrix, gbl_channel_prefix_matrix, + recv_rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, recv_gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum, + recv_src_meta, send_rdma_head, send_nvl_head) + return (recv_x, recv_x_scales) if x_scales is not None else recv_x, recv_topk_idx, recv_topk_weights, num_recv_tokens_per_expert_list, handle, EventOverlap(event) + + # noinspection PyTypeChecker + def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], + topk_weights: Optional[torch.Tensor] = None, + bias: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]] = None, + config: Optional[Config] = None, + previous_event: Optional[EventOverlap] = None, async_finish: bool = False, + allocate_on_comm_stream: bool = False) -> \ + Tuple[torch.Tensor, Optional[torch.Tensor], EventOverlap]: + """ + Internode combine implementation, for more details, please refer to the `combine` docs. + Normally, you should not directly call this function. + """ + config = self.get_combine_config(self.group_size) if config is None else config + assert config is not None + + # Unpack handle and bias + is_combined_token_in_rank, \ + _, _, \ + rdma_channel_prefix_matrix, rdma_rank_prefix_sum, gbl_channel_prefix_matrix, gbl_rank_prefix_sum, \ + src_meta, send_rdma_head, send_nvl_head = handle + bias_0, bias_1 = Buffer._unpack_bias(bias) + + # Launch the kernel + combined_x, combined_topk_weights, event = self.runtime.internode_combine( + x, topk_weights, bias_0, bias_1, + src_meta, is_combined_token_in_rank, + rdma_channel_prefix_matrix, rdma_rank_prefix_sum, gbl_channel_prefix_matrix, + send_rdma_head, send_nvl_head, config, getattr(previous_event, 'event', None), + async_finish, allocate_on_comm_stream) + return combined_x, combined_topk_weights, EventOverlap(event) + + def update_mask_buffer(self, rank_to_mask: int, mask: bool = False): """ Mask (unmask) a rank during communication (dispatch, combine, and clean) @@ -447,7 +625,7 @@ def get_next_combine_buffer( ) def update_memory_buffers( - self, num_ranks: int, num_experts_per_rank: int, num_rdma_bytes: int + self, num_ranks: int, num_experts_per_rank: int, num_nvl_bytes: int, num_rdma_bytes: int ): """ Allocate remote memory for the communication buffer. @@ -455,13 +633,15 @@ def update_memory_buffers( Arguments: num_ranks: the number of ranks. num_experts_per_rank: the number of experts per rank. + num_nvl_bytes: the buffer size for intranode NVLink communication. num_rdma_bytes: the buffer size for RDMA communication. """ self.group_size = num_ranks + self.num_nvl_bytes = num_nvl_bytes self.num_rdma_bytes = num_rdma_bytes os.environ.setdefault("UCX_RC_GDA_NUM_CHANNELS", str(num_experts_per_rank)) self.runtime.update_memory_buffers( - num_ranks, num_experts_per_rank, num_rdma_bytes + num_ranks, num_experts_per_rank, num_nvl_bytes,num_rdma_bytes ) def set_tcp_store_group(self, tcp_store_group: Optional[dist.TCPStore]) -> None: @@ -500,11 +680,35 @@ def connect_ranks(self, remote_ranks: List[int]) -> None: remote_ranks: List of remote rank IDs to establish connections with. The current rank will be automatically filtered out. """ - if self.tcp_store_group is not None: - with self._fetch_remote_metadata_from_tcp_store(remote_ranks) as remote_mds: - self.runtime.connect_ranks(remote_ranks, remote_mds) + if self.low_latency_mode: + if self.tcp_store_group is not None: + with self._fetch_remote_metadata_from_tcp_store(remote_ranks) as remote_mds: + self.runtime.connect_ranks(remote_ranks, remote_mds) + else: + self.runtime.connect_ranks(remote_ranks) else: - self.runtime.connect_ranks(remote_ranks) + # High-throughput internode mode: need group for IPC handles + if self.group is not None: + def all_gather_object(obj): + object_list = [None] * self.group_size + dist.all_gather_object(object_list, obj, self.group) + return object_list + elif self.comm is not None: + def all_gather_object(obj): + return self.comm.allgather(obj) + else: + raise ValueError("Either 'group' or 'comm' must be configured.") + + local_ipc_handle = self.runtime.get_local_ipc_handle() + ipc_handles = all_gather_object(local_ipc_handle) + + # Use TCPStore for NIXL metadata exchange if available, otherwise use ETCD + if self.tcp_store_group is not None: + with self._fetch_remote_metadata_from_tcp_store(remote_ranks) as remote_mds: + self.runtime.connect_ranks(remote_ranks, remote_mds, ipc_handles) + else: + self.runtime.connect_ranks(remote_ranks, None, ipc_handles) + def disconnect_ranks(self, remote_ranks: List[int]) -> None: """ diff --git a/examples/device/ep/nixl_ep/utils.py b/examples/device/ep/nixl_ep/utils.py index 69d6820825..d2c03f98b0 100644 --- a/examples/device/ep/nixl_ep/utils.py +++ b/examples/device/ep/nixl_ep/utils.py @@ -21,6 +21,7 @@ from typing import Any, Optional, Tuple import torch +import torch.distributed as dist # noinspection PyUnresolvedReferences from .nixl_ep_cpp import EventHandle @@ -82,3 +83,41 @@ def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: """ if self.event is not None: self.event.current_stream_wait() + + +def check_nvlink_connections(group: dist.ProcessGroup): + """ + Check NVLink connection between every pair of GPUs. + + Arguments: + group: the communication group. + """ + # Check NVLink connection + # NOTES: some A100 PCIE GPUs only have pairwise NVLink connection, so that we can only use EP2 + # TODO: check all cases, all local-node GPUs in the group should be connected via NVLink + if 'PCIE' in torch.cuda.get_device_name(): + assert group.size() <= 2, 'PCIe GPUs only have pairwise NVLink connections' + + # noinspection PyUnresolvedReferences + import pynvml + pynvml.nvmlInit() + + # noinspection PyTypeChecker + devices = os.environ.get('CUDA_VISIBLE_DEVICES', '0,1,2,3,4,5,6,7').strip(',').split(',') + physical_device_idx = int(devices[torch.cuda.current_device()]) + physical_device_indices = [0, ] * group.size() + dist.all_gather_object(physical_device_indices, physical_device_idx, group) + + # Check whether they are all connected via NVLink + # Reference: https://github.com/vllm-project/vllm/blob/b8e809a057765c574726a6077fd124db5077ce1f/vllm/platforms/cuda.py#L438 + handles = [pynvml.nvmlDeviceGetHandleByIndex(i) for i in physical_device_indices] + for i, handle in enumerate(handles): + for j, peer_handle in enumerate(handles): + if i >= j: + continue + status = pynvml.nvmlDeviceGetP2PStatus(handle, peer_handle, pynvml.NVML_P2P_CAPS_INDEX_NVLINK) + assert status == pynvml.NVML_P2P_STATUS_OK,\ + f'GPU {physical_device_indices[i]} and GPU {physical_device_indices[j]} are not connected via NVLink' + + # Close NVML + pynvml.nvmlShutdown() diff --git a/examples/device/ep/tests/elastic/elastic.py b/examples/device/ep/tests/elastic/elastic.py index ea2f4c0bf9..ff5c9f3843 100644 --- a/examples/device/ep/tests/elastic/elastic.py +++ b/examples/device/ep/tests/elastic/elastic.py @@ -492,10 +492,12 @@ def worker(torch_rank: int, args: argparse.Namespace): explicitly_destroy=True, enable_shrink=True, tcp_store_group=tcp_store, + low_latency_mode=True, ) buffer.update_memory_buffers( num_ranks=max_num_ranks, num_experts_per_rank=args.num_experts_per_rank, + num_nvl_bytes=0, num_rdma_bytes=num_rdma_bytes, ) signal.signal( diff --git a/examples/device/ep/tests/test_internode.py b/examples/device/ep/tests/test_internode.py new file mode 100644 index 0000000000..6b888a8884 --- /dev/null +++ b/examples/device/ep/tests/test_internode.py @@ -0,0 +1,325 @@ +import argparse +import os +import sys +import time + +# Add elastic subdirectory to path for store_group import +sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'elastic')) +import store_group + +import torch +import torch.distributed as dist + +# noinspection PyUnresolvedReferences +import nixl_ep +from utils import init_dist, bench, bench_kineto, calc_diff, create_grouped_scores, inplace_unique, per_token_cast_to_fp8, per_token_cast_back + +# Test compatibility with low latency functions +TCP_STORE_PORT = 9999 + + +# noinspection PyShadowingNames +def test_main(args: argparse.Namespace, num_sms: int, + local_rank: int, num_local_ranks: int, num_ranks: int, num_nodes: int, rank: int, + buffer: nixl_ep.Buffer, group: dist.ProcessGroup): + # Settings + num_tokens, hidden = args.num_tokens, args.hidden + num_topk_groups, num_topk, num_experts = args.num_topk_groups, args.num_topk, args.num_experts + + assert num_experts % num_ranks == 0 and num_local_ranks == 8 + if local_rank == 0: + print(f'[config] num_tokens={num_tokens}, hidden={hidden}, num_topk_groups={num_topk_groups}, num_topk={num_topk}', flush=True) + + # Random data + x = torch.ones((num_tokens, hidden), dtype=torch.bfloat16, device='cuda') * rank + x_pure_rand = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device='cuda') + x_e4m3 = per_token_cast_to_fp8(x) + x_e4m3 = (x_e4m3[0], x_e4m3[1].T.contiguous().T) + scores = torch.randn((num_tokens, num_experts), dtype=torch.float32, device='cuda').abs() + 1 + group_scores = scores.view(num_tokens, num_nodes, -1).amax(dim=-1) + group_idx = torch.topk(group_scores, k=num_topk_groups, dim=-1, sorted=False).indices + masked_scores = create_grouped_scores(scores, group_idx, num_nodes) + topk_idx = torch.topk(masked_scores, num_topk, dim=-1, largest=True, sorted=False)[1] + topk_idx = topk_idx.to(nixl_ep.topk_idx_t) + topk_weights = torch.ones((num_tokens, num_topk), dtype=torch.float32, device='cuda') * rank + topk_weights_pure_rand = torch.randn((num_tokens, num_topk), dtype=torch.float32, device='cuda') + rank_idx = topk_idx // (num_experts // num_ranks) + rank_idx = rank_idx.to(torch.int64) + rank_idx.masked_fill_(topk_idx == -1, -1) + inplace_unique(rank_idx, num_ranks) + rdma_rank_idx = rank_idx // num_local_ranks + rdma_rank_idx.masked_fill_(rank_idx == -1, -1) + inplace_unique(rdma_rank_idx, num_nodes) + + # RDMA dispatch counts + rdma_idx = topk_idx // (num_experts // num_nodes) + rdma_idx.masked_fill_(topk_idx == -1, -1) + inplace_unique(rdma_idx, num_nodes) + num_rdma_token_sent = rdma_idx.ne(-1).sum().item() + + # Expert meta + num_tokens_per_expert = torch.zeros((num_experts, ), dtype=torch.int, device='cuda') + for i in range(num_experts): + num_tokens_per_expert[i] = (topk_idx == i).sum() + gbl_num_tokens_per_expert = num_tokens_per_expert.clone() + dist.all_reduce(gbl_num_tokens_per_expert, group=group) + + # Rank layout meta + num_tokens_per_rank = torch.empty((num_ranks, ), dtype=torch.int, device='cuda') + num_tokens_per_rdma_rank = torch.empty((num_nodes, ), dtype=torch.int, device='cuda') + token_idx_in_rank = torch.full((num_ranks, num_tokens), -1, dtype=torch.long, device='cuda') + for i in range(num_ranks): + num_tokens_per_rank[i] = (rank_idx == i).sum() + token_sel = (rank_idx == i).max(dim=-1)[0] + count = token_sel.sum().item() + tokens = torch.sort(token_sel.to(torch.int), descending=True)[1] + tokens[:count] = torch.sort(tokens[:count])[0] + token_idx_in_rank[i][tokens[:count]] = torch.arange(count, dtype=torch.long, device='cuda') + for i in range(num_nodes): + num_tokens_per_rdma_rank[i] = (rdma_rank_idx == i).sum() + token_idx_in_rank = token_idx_in_rank.T.contiguous().to(torch.int) + is_token_in_rank = token_idx_in_rank >= 0 + gbl_num_tokens_per_rank = num_tokens_per_rank.clone() + dist.all_reduce(gbl_num_tokens_per_rank, group=group) + + ref_num_tokens_per_rank, ref_num_tokens_per_rdma_rank, ref_num_tokens_per_expert, ref_is_token_in_rank, _ = \ + buffer.get_dispatch_layout(topk_idx, num_experts) + assert torch.allclose(ref_num_tokens_per_rank, num_tokens_per_rank) + assert torch.allclose(ref_num_tokens_per_rdma_rank, num_tokens_per_rdma_rank) + assert torch.allclose(ref_num_tokens_per_expert, num_tokens_per_expert) + assert torch.allclose(ref_is_token_in_rank, is_token_in_rank) + t = bench(lambda: buffer.get_dispatch_layout(topk_idx, num_experts))[0] + if local_rank == 0: + print(f'[layout] Kernel performance: {t * 1000:.3f} ms', flush=True) + print('', flush=True) + group.barrier() + time.sleep(1) + + # Config + rdma_buffer_size, nvl_buffer_size = 128, (720 if num_ranks in (144, 160) else 512) + config = nixl_ep.Config(num_sms, 8, nvl_buffer_size, 16, rdma_buffer_size) + + # Test dispatch + # noinspection PyShadowingNames + def check_data(check_x, recv_gbl_rank_prefix_sum): + assert torch.allclose(check_x.amin(dim=1), check_x.amax(dim=1)) + check_start = 0 + for i in range(num_ranks): + check_end = recv_gbl_rank_prefix_sum[i].item() + assert (check_x[check_start:check_end, :].int() - i).sum().item() == 0 + check_start = check_end + + for previous_mode in (False, True): + for async_mode in (False, True): + for current_x in (x_pure_rand, x, x_e4m3): + for with_topk in (False, True): + if local_rank == 0: + print(f'[testing] Running with {"FP8" if isinstance(current_x, tuple) else "BF16"}, {"with" if with_topk else "without"} top-k (async={async_mode}, previous={previous_mode}) ...', flush=True, end='') + dispatch_args = {'x': current_x, 'num_tokens_per_rank': num_tokens_per_rank, 'num_tokens_per_rdma_rank': num_tokens_per_rdma_rank, 'is_token_in_rank': is_token_in_rank, + 'num_tokens_per_expert': num_tokens_per_expert, 'config': config, 'async_finish': async_mode} + if with_topk: + dispatch_args.update({'topk_idx': topk_idx, 'topk_weights': topk_weights_pure_rand if current_x is x_pure_rand else topk_weights}) + if previous_mode: + dispatch_args.update({'previous_event': buffer.capture()}) + recv_x, recv_topk_idx, recv_topk_weights, recv_num_tokens_per_expert_list, handle, event = buffer.internode_dispatch(**dispatch_args) + event.current_stream_wait() if async_mode else () + recv_x = per_token_cast_back(*recv_x) if isinstance(recv_x, tuple) else recv_x + + # Checks + recv_gbl_rank_prefix_sum = handle[-4] + assert gbl_num_tokens_per_rank[rank].item() == recv_x.size(0), f'{gbl_num_tokens_per_rank[rank].item()} != {recv_x.size(0)}' + assert gbl_num_tokens_per_expert.view(num_ranks, -1)[rank].tolist() == recv_num_tokens_per_expert_list + if current_x is not x_pure_rand: + check_data(recv_x, recv_gbl_rank_prefix_sum) + if with_topk: + # Check `topk_idx` + assert (recv_topk_idx.eq(-1) | ((recv_topk_idx >= 0) & (recv_topk_idx < (num_experts // num_ranks)))).sum().item() == recv_topk_idx.numel() + for i, count in enumerate(recv_num_tokens_per_expert_list): + assert recv_topk_idx.eq(i).sum().item() == count + + # Check `topk_weights` + if current_x is not x_pure_rand: + recv_topk_weights[recv_topk_idx.eq(-1)] = recv_topk_weights.amax(dim=1, keepdim=True).expand_as(recv_topk_weights)[recv_topk_idx.eq(-1)] + check_data(recv_topk_weights, recv_gbl_rank_prefix_sum) + + # Test cached dispatch (must without top-k staffs) + if not with_topk: + dispatch_args = {'x': current_x, 'handle': handle, 'config': config, 'async_finish': async_mode} + if previous_mode: + dispatch_args.update({'previous_event': buffer.capture()}) + recv_x_cached, _, _, _, _, event = buffer.internode_dispatch(**dispatch_args) + event.current_stream_wait() if async_mode else () + recv_x_cached = per_token_cast_back(*recv_x_cached) if isinstance(recv_x_cached, tuple) else recv_x_cached + + if current_x is not x_pure_rand: + check_data(recv_x_cached, recv_gbl_rank_prefix_sum) + + # Use cached result for combine + recv_x = recv_x_cached + + # Test combine + bias_0 = torch.ones((num_tokens, hidden), dtype=torch.bfloat16, device='cuda') + bias_1 = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device='cuda') + combine_args = {'x': recv_x, 'bias': (bias_0, bias_1), 'handle': handle, 'config': config, 'async_finish': async_mode} + if with_topk: + combine_args.update({'topk_weights': recv_topk_weights}) + if previous_mode: + combine_args.update({'previous_event': buffer.capture()}) + combined_x, combined_topk_weights, event = buffer.internode_combine(**combine_args) + event.current_stream_wait() if async_mode else () + + check_x = (combined_x.float() - bias_0.float() - bias_1.float()) / is_token_in_rank.sum(dim=1).unsqueeze(1) + ref_x = x_pure_rand if current_x is x_pure_rand else x + assert calc_diff(check_x, ref_x) < 5e-6 + if with_topk: + check_topk_weights = combined_topk_weights if (current_x is x_pure_rand) else (combined_topk_weights / is_token_in_rank.sum(dim=1).unsqueeze(1)) + ref_topk_weights = topk_weights_pure_rand if current_x is x_pure_rand else topk_weights + assert calc_diff(check_topk_weights, ref_topk_weights) < 1e-9 + + # For later tuning + dispatch_bf16_rdma_send_bytes = num_rdma_token_sent * hidden * 2 + dispatch_bf16_nvl_recv_bytes = recv_x.numel() * 2 + combine_bf16_nvl_send_bytes = dispatch_bf16_nvl_recv_bytes + combine_bf16_rdma_recv_bytes = dispatch_bf16_rdma_send_bytes + + # Sync all ranks before printing passed + group.barrier() + if local_rank == 0: + print(' passed', flush=True) + group.barrier() + if local_rank == 0: + print('', flush=True) + + # Tune dispatch performance + best_dispatch_results = None + fp8_factor = (1 + 4 / 128) / 2 + for current_x in (x_e4m3, x): + best_time, best_results = 1e10, None + rdma_send_bytes = (dispatch_bf16_rdma_send_bytes * fp8_factor) if isinstance(current_x, tuple) else dispatch_bf16_rdma_send_bytes + nvl_recv_bytes = (dispatch_bf16_nvl_recv_bytes * fp8_factor) if isinstance(current_x, tuple) else dispatch_bf16_nvl_recv_bytes + for nvl_chunk_size in range(4, 45, 4): + for rdma_chunk_size in range(4, 33, 4): + config = nixl_ep.Config(num_sms, nvl_chunk_size, nvl_buffer_size, rdma_chunk_size, rdma_buffer_size) + tune_args = {'x': current_x, 'handle': handle, 'config': config} + t, notify_t = bench_kineto(lambda: buffer.internode_dispatch(**tune_args), ('dispatch', 'notify')) + if t < best_time: + best_time, best_results = t, (num_sms, nvl_chunk_size, rdma_chunk_size, notify_t) + if local_rank == 0: + print(f'[tuning] SMs {num_sms}, NVL chunk {nvl_chunk_size}, RDMA chunk {rdma_chunk_size}, transmit: {t * 1e6:.2f} us, notify: {notify_t * 1e6:.2f} us, BW: {rdma_send_bytes / 1e9 / t:.2f} GB/s (RDMA), {nvl_recv_bytes / 1e9 / t:.2f} GB/s (NVL) ', flush=True) + if local_rank == 0: + print(f'[tuning] Best dispatch ({"FP8" if isinstance(current_x, tuple) else "BF16"}): SMs {best_results[0]}, NVL chunk {best_results[1]}, RDMA chunk {best_results[2]}, transmit: {best_time * 1e6:.2f} us, notify: {best_results[3] * 1e6:.2f} us, BW: {rdma_send_bytes / 1e9 / best_time:.2f} GB/s (RDMA), {nvl_recv_bytes / 1e9 / best_time:.2f} GB/s (NVL)', flush=True) + print('', flush=True) + + if isinstance(current_x, tuple): + # Gather FP8 the best config from rank 0 + best_dispatch_results = torch.tensor([best_results[0], best_results[1], best_results[2]], dtype=torch.int32, device='cuda') + all_best_fp8_results_list = [torch.zeros_like(best_dispatch_results) for _ in range(torch.distributed.get_world_size())] + dist.all_gather(all_best_fp8_results_list, best_dispatch_results, group=group) + best_dispatch_results = all_best_fp8_results_list[0].tolist() + dispatch_config = nixl_ep.Config(best_dispatch_results[0], best_dispatch_results[1], nvl_buffer_size, best_dispatch_results[2], rdma_buffer_size) + + dispatch_args = {'x': x, 'num_tokens_per_rank': num_tokens_per_rank, 'num_tokens_per_rdma_rank': num_tokens_per_rdma_rank, + 'is_token_in_rank': is_token_in_rank, 'num_tokens_per_expert': num_tokens_per_expert, + 'config': dispatch_config if dispatch_config is not None else config} + recv_x, _, _, _, handle, _ = buffer.internode_dispatch(**dispatch_args) + + # Tune combine performance + best_time, best_results = 1e10, None + for nvl_chunk_size in range(1, 8, 1): + for rdma_chunk_size in range(12 if num_nodes == 2 else 8, 33, 4): + config = nixl_ep.Config(num_sms, nvl_chunk_size, nvl_buffer_size, rdma_chunk_size, rdma_buffer_size) + tune_args = {'x': recv_x, 'handle': handle, 'config': config} + t, notify_t = bench_kineto(lambda: buffer.internode_combine(**tune_args), ('combine', 'notify')) + if local_rank == 0: + print(f'[tuning] SMs {num_sms}, NVL chunk {nvl_chunk_size}, RDMA chunk {rdma_chunk_size}, transmit: {t * 1e6:.2f} us, notify: {notify_t * 1e6:.2f} us, BW: {combine_bf16_rdma_recv_bytes / 1e9 / t:.2f} GB/s (RDMA), {combine_bf16_nvl_send_bytes / 1e9 / t:.2f} GB/s (NVL) ', flush=True) + if t < best_time: + best_time, best_results = t, (num_sms, nvl_chunk_size, rdma_chunk_size, notify_t) + + if local_rank == 0: + print(f'[tuning] Best combine: SMs {best_results[0]}, NVL chunk {best_results[1]}, RDMA chunk {best_results[2]}, transmit: {best_time * 1e6:.2f} us, notify: {best_results[3] * 1e6:.2f} us, BW: {combine_bf16_rdma_recv_bytes / 1e9 / best_time:.2f} GB/s (RDMA), {combine_bf16_nvl_send_bytes / 1e9 / best_time:.2f} GB/s (NVL)', flush=True) + print('', flush=True) + + +# noinspection PyUnboundLocalVariable,PyShadowingNames +def test_loop(local_rank: int, num_local_ranks: int, args: argparse.Namespace): + os.environ['CUDA_VISIBLE_DEVICES'] = str(local_rank) + + num_nodes = int(os.getenv('WORLD_SIZE', 1)) + + rank, num_ranks, group = init_dist(local_rank, num_local_ranks) + print(f"pid: {os.getpid()}, rank: {rank}, num_ranks: {num_ranks} ,local_rank: {local_rank}", flush=True) + if args.test_ll_compatibility: + ll_num_tokens, ll_hidden, ll_num_experts, ll_num_topk = 16, 5120, 256, 9 + + num_sms = 24 + num_qps_per_rank = max(num_sms // 2, ll_num_experts // num_ranks if args.test_ll_compatibility else 0) + + # Create TCPStore client for NIXL metadata exchange + tcp_server = args.tcp_server if args.tcp_server else "127.0.0.1" + tcp_store = store_group.create_client_store( + master_addr=tcp_server, + port=TCP_STORE_PORT, + ) + + # Initialize NIXL buffer with group (for IPC handles) and TCPStore (for NIXL metadata) + print(f"pid: {os.getpid()}, rank: {rank}, num_ranks: {num_ranks}, initializing buffer", flush=True) + buffer = nixl_ep.Buffer(rank=rank, low_latency_mode=False, explicitly_destroy=True, group=group, tcp_store_group=tcp_store, nvlink_backend="ipc") + buffer.update_memory_buffers(num_ranks=num_ranks, num_experts_per_rank=num_qps_per_rank, num_nvl_bytes=int(2e9), num_rdma_bytes=int(1e9)) + buffer.connect_ranks([i for i in range(num_ranks) if i != rank]) + + assert num_local_ranks == 8 and num_ranks > 8 + torch.manual_seed(rank) + + for i in (num_sms, ): + test_main(args, i, local_rank, num_local_ranks, num_ranks, num_nodes, rank, buffer, group) + if local_rank == 0: + print('', flush=True) + + + # Destroy the buffer runtime and communication group + buffer.destroy() + dist.barrier() + dist.destroy_process_group() + +def run_server(): + _store = store_group.create_master_store(port=TCP_STORE_PORT) # noqa: F841 + # Keep the server process alive while TCPStore serves requests + while True: + time.sleep(1) + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='Test internode EP kernels') + parser.add_argument('--num-processes', type=int, default=8, + help='Number of processes to spawn (default: 8)') + parser.add_argument('--num-tokens', type=int, default=4096, + help='Number of tokens (default: 4096)') + parser.add_argument('--hidden', type=int, default=7168, + help='Hidden dimension size (default: 7168)') + parser.add_argument('--num-topk-groups', type=int, default=None, + help='Number of top-k groups (default: `min(num_nodes, 4)`)') + parser.add_argument('--num-topk', type=int, default=8, + help='Number of top-k experts (default: 8)') + parser.add_argument('--num-experts', type=int, default=256, + help='Number of experts (default: 256') + parser.add_argument('--test-ll-compatibility', action='store_true', + help='whether to test compatibility with low-latency kernels') + parser.add_argument( + "--tcp-server", + type=str, + help="TCP server address (for both TCPStore and rank server). If not set, both will be started locally.", + ) + args = parser.parse_args() + + if not args.tcp_server: + print("Starting TCPStore and rank server locally", flush=True) + server_process = torch.multiprocessing.Process(target=run_server, daemon=True) + server_process.start() + time.sleep(0.5) + + # Set default `num_topk_groups` if not provided + if args.num_topk_groups is None: + num_nodes = int(os.getenv('WORLD_SIZE', 1)) + args.num_topk_groups = min(num_nodes, 4) + + num_processes = args.num_processes + torch.multiprocessing.spawn(test_loop, args=(num_processes, args), nprocs=num_processes)