diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index 714774c81..24ace74a7 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -125,6 +125,49 @@ void SharedMemoryAllocator::close_mem_handle(void* ptr) { namespace deep_ep { +namespace { + +internode::NormalNotifyStats prepare_normal_notify_stats(const NormalNotifyStats& stats, int64_t* timer_state) { + const bool enabled = stats.duration_ns.has_value(); + EP_HOST_ASSERT(stats.count.has_value() == enabled); + if (not enabled) + return {}; + + EP_HOST_ASSERT(stats.duration_ns->is_cuda() and stats.count->is_cuda()); + EP_HOST_ASSERT(stats.duration_ns->scalar_type() == torch::kInt64 and stats.count->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(stats.duration_ns->dim() == 1 and stats.duration_ns->numel() == 1 and stats.duration_ns->is_contiguous()); + EP_HOST_ASSERT(stats.count->dim() == 1 and stats.count->numel() == 1 and stats.count->is_contiguous()); + + return { + stats.duration_ns->data_ptr(), + stats.count->data_ptr(), + timer_state, + }; +} + +internode::NormalCompletionStats prepare_normal_completion_stats(const NormalCompletionStats& stats, int num_ranks) { + const bool enabled = stats.cost.has_value(); + EP_HOST_ASSERT(stats.sample_count.has_value() == enabled); + EP_HOST_ASSERT(stats.token_count.has_value() == enabled); + if (not enabled) + return {}; + + for (const auto& tensor : {stats.cost, stats.sample_count, stats.token_count}) { + EP_HOST_ASSERT(tensor->is_cuda()); + EP_HOST_ASSERT(tensor->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(tensor->dim() == 1 and tensor->is_contiguous()); + EP_HOST_ASSERT(tensor->numel() == num_ranks); + } + + return { + stats.cost->data_ptr(), + stats.sample_count->data_ptr(), + stats.token_count->data_ptr(), + }; +} + +} // namespace + Buffer::Buffer(int rank, int num_ranks, int64_t num_nvl_bytes, @@ -198,6 +241,7 @@ Buffer::Buffer(int rank, // Create 32 MiB workspace CUDA_CHECK(cudaMalloc(&workspace, NUM_WORKSPACE_BYTES)); CUDA_CHECK(cudaMemsetAsync(workspace, 0, NUM_WORKSPACE_BYTES, comm_stream)); + CUDA_CHECK(cudaMalloc(&normal_notify_full_kernel_timer_states, 3 * 2 * sizeof(int64_t))); // MoE counter CUDA_CHECK(cudaMallocHost(&moe_recv_counter, sizeof(int64_t), cudaHostAllocMapped)); @@ -316,6 +360,7 @@ void Buffer::destroy() { // Free workspace and MoE counter CUDA_CHECK(cudaFree(workspace)); + CUDA_CHECK(cudaFree(normal_notify_full_kernel_timer_states)); CUDA_CHECK(cudaFreeHost(const_cast(moe_recv_counter))); // Free chunked mode staffs @@ -946,7 +991,8 @@ Buffer::internode_dispatch(const torch::Tensor& x, const Config& config, std::optional& previous_event, bool async, - bool allocate_on_comm_stream) { + bool allocate_on_comm_stream, + const NormalDispatchStats& normal_stats) { #ifndef DISABLE_NVSHMEM // 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 @@ -1004,6 +1050,10 @@ Buffer::internode_dispatch(const torch::Tensor& x, 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); } + const auto notify_stats = prepare_normal_notify_stats(normal_stats.notify, normal_notify_full_kernel_timer_states); + const auto cached_notify_stats = prepare_normal_notify_stats(normal_stats.cached_notify, normal_notify_full_kernel_timer_states + 2); + const auto final_completion_stats = prepare_normal_completion_stats(normal_stats.final_completion, num_ranks); + const auto rdma_recv_completion_stats = prepare_normal_completion_stats(normal_stats.rdma_recv_completion, num_ranks); 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)); @@ -1094,7 +1144,8 @@ Buffer::internode_dispatch(const torch::Tensor& x, config.get_rdma_buffer_size_hint(hidden_int4 * sizeof(int4), num_ranks), num_nvl_bytes, true, - low_latency_mode); + low_latency_mode, + cached_notify_stats); } 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)); @@ -1134,7 +1185,8 @@ Buffer::internode_dispatch(const torch::Tensor& x, comm_stream, config.get_rdma_buffer_size_hint(hidden_int4 * sizeof(int4), num_ranks), num_nvl_bytes, - low_latency_mode); + low_latency_mode, + notify_stats); // Synchronize total received tokens and tokens per expert if (num_worst_tokens > 0) { @@ -1235,6 +1287,8 @@ Buffer::internode_dispatch(const torch::Tensor& x, buffer_ptrs_gpu, config.num_max_nvl_chunked_send_tokens, config.num_max_nvl_chunked_recv_tokens, + final_completion_stats, + rdma_recv_completion_stats, rank, num_ranks, cached_mode, @@ -1324,7 +1378,8 @@ std::tuple, std::optional& previous_event, bool async, - bool allocate_on_comm_stream) { + bool allocate_on_comm_stream, + const NormalCombineStats& normal_stats) { #ifndef DISABLE_NVSHMEM const int num_channels = config.num_sms / 2; EP_HOST_ASSERT(config.num_sms % 2 == 0); @@ -1356,6 +1411,8 @@ std::tuple, std::optional, std::optional>({bias_0, bias_1}); @@ -1453,6 +1511,7 @@ std::tuple, std::optional cast_optional_tensor(pybind11::handle value) { + return value.is_none() ? std::nullopt : std::make_optional(pybind11::cast(value)); +} + +pybind11::sequence require_normal_stats_sequence(pybind11::handle source, size_t expected_size, const char* argument_name) { + if (not pybind11::isinstance(source) and not pybind11::isinstance(source)) + throw pybind11::value_error(std::string("`") + argument_name + "` must be a tuple or list"); + + auto sequence = pybind11::reinterpret_borrow(source); + if (static_cast(sequence.size()) != expected_size) + throw pybind11::value_error(std::string("`") + argument_name + "` must contain " + std::to_string(expected_size) + + " tensors, but got " + std::to_string(sequence.size())); + return sequence; +} + +} // namespace + +namespace pybind11::detail { + +template <> +struct type_caster { +public: + PYBIND11_TYPE_CASTER(deep_ep::NormalDispatchStats, const_name("NormalDispatchStats")); + + bool load(handle source, bool) { + if (source.is_none()) { + value = {}; + return true; + } + + auto stats = require_normal_stats_sequence(source, 10, "normal_dispatch_stats"); + value = { + {cast_optional_tensor(stats[0]), cast_optional_tensor(stats[1])}, + {cast_optional_tensor(stats[2]), cast_optional_tensor(stats[3])}, + {cast_optional_tensor(stats[4]), cast_optional_tensor(stats[5]), cast_optional_tensor(stats[6])}, + {cast_optional_tensor(stats[7]), cast_optional_tensor(stats[8]), cast_optional_tensor(stats[9])}, + }; + return true; + } +}; + +template <> +struct type_caster { +public: + PYBIND11_TYPE_CASTER(deep_ep::NormalCombineStats, const_name("NormalCombineStats")); + + bool load(handle source, bool) { + if (source.is_none()) { + value = {}; + return true; + } + + auto stats = require_normal_stats_sequence(source, 5, "normal_combine_stats"); + value = { + {cast_optional_tensor(stats[0]), cast_optional_tensor(stats[1])}, + {cast_optional_tensor(stats[2]), cast_optional_tensor(stats[3]), cast_optional_tensor(stats[4])}, + }; + return true; + } +}; + +} // namespace pybind11::detail + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.doc() = "DeepEP: an efficient expert-parallel communication library"; @@ -1900,8 +2024,47 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { .def("get_dispatch_layout", &deep_ep::Buffer::get_dispatch_layout) .def("intranode_dispatch", &deep_ep::Buffer::intranode_dispatch) .def("intranode_combine", &deep_ep::Buffer::intranode_combine) - .def("internode_dispatch", &deep_ep::Buffer::internode_dispatch) - .def("internode_combine", &deep_ep::Buffer::internode_combine) + .def("internode_dispatch", + &deep_ep::Buffer::internode_dispatch, + py::arg("x"), + py::arg("x_scales"), + py::arg("topk_idx"), + py::arg("topk_weights"), + py::arg("num_tokens_per_rank"), + py::arg("num_tokens_per_rdma_rank"), + py::arg("is_token_in_rank"), + py::arg("num_tokens_per_expert"), + py::arg("cached_num_recv_tokens"), + py::arg("cached_num_rdma_recv_tokens"), + py::arg("cached_rdma_channel_prefix_matrix"), + py::arg("cached_recv_rdma_rank_prefix_sum"), + py::arg("cached_gbl_channel_prefix_matrix"), + py::arg("cached_recv_gbl_rank_prefix_sum"), + py::arg("expert_alignment"), + py::arg("num_worst_tokens"), + py::arg("config"), + py::arg("previous_event"), + py::arg("async"), + py::arg("allocate_on_comm_stream"), + py::arg("normal_dispatch_stats") = py::none()) + .def("internode_combine", + &deep_ep::Buffer::internode_combine, + py::arg("x"), + py::arg("topk_weights"), + py::arg("bias_0"), + py::arg("bias_1"), + py::arg("src_meta"), + py::arg("is_combined_token_in_rank"), + py::arg("rdma_channel_prefix_matrix"), + py::arg("rdma_rank_prefix_sum"), + py::arg("gbl_channel_prefix_matrix"), + py::arg("combined_rdma_head"), + py::arg("combined_nvl_head"), + py::arg("config"), + py::arg("previous_event"), + py::arg("async"), + py::arg("allocate_on_comm_stream"), + py::arg("normal_combine_stats") = py::none()) .def("clean_low_latency_buffer", &deep_ep::Buffer::clean_low_latency_buffer) .def("low_latency_dispatch", &deep_ep::Buffer::low_latency_dispatch) .def("low_latency_combine", &deep_ep::Buffer::low_latency_combine) diff --git a/csrc/deep_ep.hpp b/csrc/deep_ep.hpp index 090e5a4f1..eba63f98a 100644 --- a/csrc/deep_ep.hpp +++ b/csrc/deep_ep.hpp @@ -51,6 +51,29 @@ class SharedMemoryAllocator { namespace deep_ep { +struct NormalNotifyStats { + std::optional duration_ns; + std::optional count; +}; + +struct NormalCompletionStats { + std::optional cost; + std::optional sample_count; + std::optional token_count; +}; + +struct NormalDispatchStats { + NormalNotifyStats notify; + NormalNotifyStats cached_notify; + NormalCompletionStats final_completion; + NormalCompletionStats rdma_recv_completion; +}; + +struct NormalCombineStats { + NormalNotifyStats cached_notify; + NormalCompletionStats logical_recv_completion; +}; + struct Buffer { EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS == 8, "The number of maximum NVLink peers must be 8"); @@ -99,6 +122,9 @@ struct Buffer { // Workspace void* workspace = nullptr; + // Per-launch scratch for the three normal-mode Notify timers + int64_t* normal_notify_full_kernel_timer_states = nullptr; + // Host-side MoE info volatile int* moe_recv_counter = nullptr; int* moe_recv_counter_mapped = nullptr; @@ -235,7 +261,8 @@ struct Buffer { const Config& config, std::optional& previous_event, bool async, - bool allocate_on_comm_stream); + bool allocate_on_comm_stream, + const NormalDispatchStats& normal_stats = {}); std::tuple, std::optional> internode_combine( const torch::Tensor& x, @@ -252,7 +279,8 @@ struct Buffer { const Config& config, std::optional& previous_event, bool async, - bool allocate_on_comm_stream); + bool allocate_on_comm_stream, + const NormalCombineStats& normal_stats = {}); void clean_low_latency_buffer(int num_max_dispatch_tokens_per_rank, int hidden, int num_experts); diff --git a/csrc/kernels/api.cuh b/csrc/kernels/api.cuh index c43dd5ecf..c4e015951 100644 --- a/csrc/kernels/api.cuh +++ b/csrc/kernels/api.cuh @@ -144,6 +144,18 @@ namespace internode { int get_source_meta_bytes(); +struct NormalNotifyStats { + int64_t* duration_ns = nullptr; + int64_t* count = nullptr; + int64_t* timer_state = nullptr; +}; + +struct NormalCompletionStats { + int64_t* cost = nullptr; + int64_t* sample_count = nullptr; + int64_t* token_count = nullptr; +}; + void notify_dispatch(const int* num_tokens_per_rank, int* moe_recv_counter_mapped, int num_ranks, @@ -173,7 +185,8 @@ void notify_dispatch(const int* num_tokens_per_rank, cudaStream_t stream, int64_t num_rdma_bytes, int64_t num_nvl_bytes, - bool low_latency_mode); + bool low_latency_mode, + NormalNotifyStats normal_notify_stats); void dispatch(void* recv_x, float* recv_x_scales, @@ -207,6 +220,10 @@ void dispatch(void* recv_x, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + // Completion cost tensors accumulate clock64() SM cycles; + // Notify duration tensors accumulate %globaltimer nanoseconds. + NormalCompletionStats normal_final_completion_stats, + NormalCompletionStats normal_rdma_recv_completion_stats, int rank, int num_ranks, bool is_cached_dispatch, @@ -235,7 +252,8 @@ void cached_notify(int hidden_int4, int64_t num_rdma_bytes, int64_t num_nvl_bytes, bool is_cached_dispatch, - bool low_latency_mode); + bool low_latency_mode, + NormalNotifyStats normal_notify_stats); void combine(cudaDataType_t type, void* combined_x, @@ -261,6 +279,9 @@ void combine(cudaDataType_t type, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + // Completion cost tensors accumulate clock64() SM cycles; + // Notify duration tensors accumulate %globaltimer nanoseconds. + NormalCompletionStats normal_logical_recv_completion_stats, int rank, int num_ranks, cudaStream_t stream, diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index 48c6c0018..53520edbb 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -1,6 +1,7 @@ #include #include +#include "api.cuh" #include "buffer.cuh" #include "configs.cuh" #include "exception.cuh" @@ -14,6 +15,68 @@ namespace internode { extern nvshmem_team_t cpu_rdma_team; +__device__ __forceinline__ uint64_t read_globaltimer_ns() { + uint64_t value; + asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(value)); + return value; +} + +__device__ __forceinline__ void notify_full_kernel_timer_begin(NormalNotifyStats stats) { + if (stats.timer_state == nullptr) + return; + + if (threadIdx.x == 0) { + auto state = reinterpret_cast(stats.timer_state); + // The state is zero-initialized before every launch. Complementing the + // timestamp lets atomicMax select the earliest block start without a + // separate UINT64_MAX initialization for this field. + atomicMax(state, ~read_globaltimer_ns()); + } +} + +__device__ __forceinline__ void notify_full_kernel_timer_end(NormalNotifyStats stats) { + if (stats.timer_state == nullptr) + return; + + __syncthreads(); + if (threadIdx.x == 0) { + auto state = reinterpret_cast(stats.timer_state); + auto completed_blocks = atomicAdd(state + 1, 1ull); + if (completed_blocks + 1 == gridDim.x) { + auto inverted_start = atomicAdd(state, 0ull); + auto duration_ns = read_globaltimer_ns() - ~inverted_start; + atomicAdd(reinterpret_cast(stats.duration_ns), duration_ns); + atomicAdd(reinterpret_cast(stats.count), 1ull); + } + } +} + +__device__ __forceinline__ void try_record_dispatch_rdma_recv_completion(NormalCompletionStats stats, + const uint64_t* rdma_channel_tail, + int src_rdma_rank, + int gateway_nvl_rank, + int expected_count, + uint64_t start_time, + bool& recorded) { + if (stats.cost == nullptr or expected_count <= 0 or recorded) + return; + + const auto observed_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail)); + if (observed_tail < expected_count) + return; + + // An RDMA lane aggregates one source node. DeepEP's symmetric PE mapping + // attributes that node-level completion to its gateway PE, whose local NVL + // rank matches this receiver. This is a gateway proxy, not per-source-GPU + // completion timing. + const auto src_gateway_proxy_rank = src_rdma_rank * NUM_MAX_NVL_PEERS + gateway_nvl_rank; + atomicAdd(reinterpret_cast(stats.cost + src_gateway_proxy_rank), clock64() - start_time); + atomicAdd(reinterpret_cast(stats.sample_count + src_gateway_proxy_rank), 1ull); + atomicAdd(reinterpret_cast(stats.token_count + src_gateway_proxy_rank), + static_cast(expected_count)); + recorded = true; +} + struct SourceMeta { int src_rdma_rank, is_token_in_nvl_rank_bits; @@ -25,7 +88,7 @@ struct SourceMeta { __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 +#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; } @@ -115,7 +178,8 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, void** buffer_ptrs, int** barrier_signal_ptrs, int rank, - const nvshmem_team_t rdma_team) { + const nvshmem_team_t rdma_team, + NormalNotifyStats normal_notify_stats) { 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; @@ -123,6 +187,8 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, 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; + notify_full_kernel_timer_begin(normal_notify_stats); + if (sm_id == 0) { // Communication with others // Global barrier: the first warp does intra-node sync, the second warp does internode sync @@ -149,15 +215,15 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, // Clean up for later data dispatch EP_DEVICE_ASSERT(rdma_recv_num_tokens_mixed.total_bytes <= rdma_clean_offset * sizeof(int)); - #pragma unroll +#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 +// 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 +#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]; @@ -213,7 +279,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_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 +#pragma unroll for (int i = thread_id; i < nvl_num_int_clean; i += num_threads) nvl_buffer_ptr_int[nvl_clean_offset + i] = 0; @@ -222,7 +288,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, EP_DEVICE_ASSERT(num_rdma_experts <= num_threads); if (thread_id < num_rdma_experts) { int sum = 0; - #pragma unroll +#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; @@ -232,7 +298,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, // Reduce RDMA received tokens if (thread_id == 0) { int sum = 0; - #pragma unroll +#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]; recv_rdma_rank_prefix_sum[i] = sum; @@ -247,10 +313,10 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, // 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 +#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 +#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]; } @@ -260,7 +326,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, EP_DEVICE_ASSERT(num_nvl_experts <= num_threads); if (thread_id == 0) { int sum = 0; - #pragma unroll +#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]; @@ -274,7 +340,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, } if (thread_id < num_nvl_experts) { int sum = 0; - #pragma unroll +#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; @@ -303,7 +369,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, 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 +#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); @@ -311,13 +377,13 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, // Warp reduce total_count = warp_reduce_sum(total_count); - #pragma unroll +#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 +#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; @@ -328,7 +394,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, __syncthreads(); if (thread_id == 0) { auto prefix_row = rdma_channel_prefix_matrix + dst_rdma_rank * num_channels; - #pragma unroll +#pragma unroll for (int i = 1; i < num_channels; ++i) prefix_row[i] += prefix_row[i - 1]; } @@ -336,11 +402,13 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, 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 +#pragma unroll for (int i = 1; i < num_channels; ++i) prefix_row[i] += prefix_row[i - 1]; } } + + notify_full_kernel_timer_end(normal_notify_stats); } void notify_dispatch(const int* num_tokens_per_rank, @@ -372,7 +440,8 @@ void notify_dispatch(const int* num_tokens_per_rank, cudaStream_t stream, int64_t num_rdma_bytes, int64_t num_nvl_bytes, - bool low_latency_mode) { + bool low_latency_mode, + NormalNotifyStats normal_notify_stats) { #define NOTIFY_DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ { \ auto notify_dispatch_func = low_latency_mode ? notify_dispatch : notify_dispatch; \ @@ -403,7 +472,8 @@ void notify_dispatch(const int* num_tokens_per_rank, buffer_ptrs, \ barrier_signal_ptrs, \ rank, \ - cpu_rdma_team); \ + cpu_rdma_team, \ + normal_notify_stats); \ } \ break @@ -429,6 +499,8 @@ void notify_dispatch(const int* num_tokens_per_rank, // Launch kernel SETUP_LAUNCH_CONFIG(1 + num_rdma_ranks, kNumThreads, stream); + if (normal_notify_stats.timer_state != nullptr) + CUDA_CHECK(cudaMemsetAsync(normal_notify_stats.timer_state, 0, 2 * sizeof(int64_t), stream)); SWITCH_RDMA_RANKS(NOTIFY_DISPATCH_LAUNCH_CASE); #undef NOTIFY_DISPATCH_LAUNCH_CASE } @@ -477,6 +549,8 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + NormalCompletionStats normal_final_completion_stats, + NormalCompletionStats normal_rdma_recv_completion_stats, int rank, int num_ranks) { enum class WarpRole { kRDMASender, kRDMASenderCoordinator, kRDMAAndNVLForwarder, kForwarderCoordinator, kNVLReceivers }; @@ -488,6 +562,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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; + const bool enable_normal_dispatch_final_completion_stats = normal_final_completion_stats.cost != nullptr; EP_DEVICE_ASSERT(ibgda_get_state()->num_rc_per_pe == num_channels or ibgda_get_state()->num_rc_per_pe >= num_sms); @@ -666,7 +741,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV SourceMeta src_meta; int num_topk_ranks = 0, topk_ranks[kNumTopkRDMARanks]; void* dst_send_buffers[kNumTopkRDMARanks]; - #pragma unroll +#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; @@ -682,37 +757,37 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV // Copy `x` into symmetric send buffer auto st_broadcast = [=](const int key, const int4& value) { - #pragma unroll +#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 +#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 +// 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 +#pragma unroll for (int j = 0; j < num_topk_ranks; ++j) st_na_global(reinterpret_cast(dst_send_buffers[j]) + i, value); } - #pragma unroll +#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 +#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 +// 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)); @@ -847,6 +922,9 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV // Wait counters to arrive int num_tokens_to_recv_from_rdma = 0, src_rdma_channel_prefix = 0; + int rdma_recv_completion_expected_count = 0; + uint64_t rdma_recv_completion_start_time = 0; + bool rdma_recv_completion_recorded = false; EP_DEVICE_ASSERT(kNumRDMARanks <= 32); auto start_time = clock64(); if (lane_id < kNumRDMARanks) { @@ -870,6 +948,8 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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); + rdma_recv_completion_expected_count = num_tokens_to_recv_from_rdma; + rdma_recv_completion_start_time = clock64(); break; } @@ -892,6 +972,14 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV } } __syncwarp(); + if (dst_nvl_rank == 0 and lane_id < kNumRDMARanks) + try_record_dispatch_rdma_recv_completion(normal_rdma_recv_completion_stats, + rdma_channel_tail.buffer(lane_id), + lane_id, + nvl_rank, + rdma_recv_completion_expected_count, + rdma_recv_completion_start_time, + rdma_recv_completion_recorded); // Shift cached head send_nvl_head += src_rdma_channel_prefix * NUM_MAX_NVL_PEERS + dst_nvl_rank; @@ -905,6 +993,14 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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)) { + if (dst_nvl_rank == 0 and lane_id < kNumRDMARanks) + try_record_dispatch_rdma_recv_completion(normal_rdma_recv_completion_stats, + rdma_channel_tail.buffer(lane_id), + lane_id, + nvl_rank, + rdma_recv_completion_expected_count, + rdma_recv_completion_start_time, + rdma_recv_completion_recorded); // Check destination queue emptiness, or wait a buffer to be released start_time = clock64(); while (true) { @@ -1006,6 +1102,14 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV if (elect_one_sync()) st_release_sys_global(nvl_channel_tail.buffer(), cached_nvl_channel_tail); } + if (dst_nvl_rank == 0 and lane_id < kNumRDMARanks) + try_record_dispatch_rdma_recv_completion(normal_rdma_recv_completion_stats, + rdma_channel_tail.buffer(lane_id), + lane_id, + nvl_rank, + rdma_recv_completion_expected_count, + rdma_recv_completion_start_time, + rdma_recv_completion_recorded); // Retired __syncwarp(); @@ -1021,7 +1125,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV // Clean shared memory EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS <= 32, "Invalid number of NVL peers"); - #pragma unroll +#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) @@ -1032,7 +1136,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV while (true) { // Find minimum head int min_head = std::numeric_limits::max(); - #pragma unroll +#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]); @@ -1091,6 +1195,9 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV } } num_tokens_to_recv = warp_reduce_sum(end_offset - start_offset); + int completion_expected_count = end_offset - start_offset; + int completion_recv_count = 0; + uint64_t completion_start_time = clock64(); // Save for combine usage if (lane_id < kNumRDMARanks and not kCachedMode) @@ -1128,6 +1235,14 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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; + int completion_stats_idx = -1; + if (enable_normal_dispatch_final_completion_stats and lane_id == meta.src_rdma_rank and completion_expected_count > 0) { + completion_recv_count += 1; + if (completion_recv_count == completion_expected_count) { + const auto src_global_rank = lane_id * NUM_MAX_NVL_PEERS + src_nvl_rank; + completion_stats_idx = src_global_rank; + } + } bool scale_aligned = (scale_bytes % 16 == 0); auto tma_load_bytes = hidden_bytes + (scale_aligned ? scale_bytes : 0); @@ -1182,6 +1297,13 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV // Wait TMA to be finished tma_store_wait<0>(); __syncwarp(); + if (completion_stats_idx >= 0) { + atomicAdd(reinterpret_cast(normal_final_completion_stats.cost + completion_stats_idx), + clock64() - completion_start_time); + atomicAdd(reinterpret_cast(normal_final_completion_stats.sample_count + completion_stats_idx), 1); + atomicAdd(reinterpret_cast(normal_final_completion_stats.token_count + completion_stats_idx), + static_cast(completion_expected_count)); + } } // Move queue @@ -1201,7 +1323,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV const auto clean_start = num_recv_tokens * num_topk + channel_id * num_threads; const auto clean_end = num_worst_tokens * num_topk; const auto clean_stride = num_channels * num_threads; - #pragma unroll +#pragma unroll for (int i = clean_start + thread_id; i < clean_end; i += clean_stride) recv_topk_idx[i] = -1; } @@ -1239,6 +1361,8 @@ void dispatch(void* recv_x, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + NormalCompletionStats normal_final_completion_stats, + NormalCompletionStats normal_rdma_recv_completion_stats, int rank, int num_ranks, bool is_cached_dispatch, @@ -1251,6 +1375,12 @@ void dispatch(void* recv_x, // Make sure never OOB EP_HOST_ASSERT(static_cast(num_scales) * scale_hidden_stride < std::numeric_limits::max()); + const bool enable_normal_dispatch_final_completion_stats = normal_final_completion_stats.cost != nullptr; + EP_HOST_ASSERT((normal_final_completion_stats.sample_count != nullptr) == enable_normal_dispatch_final_completion_stats); + EP_HOST_ASSERT((normal_final_completion_stats.token_count != nullptr) == enable_normal_dispatch_final_completion_stats); + const bool enable_normal_dispatch_rdma_recv_completion_stats = normal_rdma_recv_completion_stats.cost != nullptr; + EP_HOST_ASSERT((normal_rdma_recv_completion_stats.sample_count != nullptr) == enable_normal_dispatch_rdma_recv_completion_stats); + EP_HOST_ASSERT((normal_rdma_recv_completion_stats.token_count != nullptr) == enable_normal_dispatch_rdma_recv_completion_stats); #define DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ { \ @@ -1294,6 +1424,8 @@ void dispatch(void* recv_x, buffer_ptrs, \ num_max_nvl_chunked_send_tokens, \ num_max_nvl_chunked_recv_tokens, \ + normal_final_completion_stats, \ + normal_rdma_recv_completion_stats, \ rank, \ num_ranks); \ } \ @@ -1324,7 +1456,8 @@ __global__ void cached_notify(const int rdma_clean_offset, int rank, int num_ranks, bool is_cached_dispatch, - const nvshmem_team_t rdma_team) { + const nvshmem_team_t rdma_team, + NormalNotifyStats normal_notify_stats) { auto sm_id = static_cast(blockIdx.x); auto thread_id = static_cast(threadIdx.x); auto num_threads = static_cast(blockDim.x); @@ -1336,6 +1469,8 @@ __global__ void cached_notify(const int rdma_clean_offset, auto num_rdma_ranks = num_ranks / NUM_MAX_NVL_PEERS; auto rdma_rank = rank / NUM_MAX_NVL_PEERS; + notify_full_kernel_timer_begin(normal_notify_stats); + // Using two SMs, which clean the RDMA/NVL buffer respectively if (sm_id == 0) { auto qps_per_rdma_rank = ibgda_get_state()->num_rc_per_pe * ibgda_get_state()->num_devices_initialized; @@ -1355,13 +1490,13 @@ __global__ void cached_notify(const int rdma_clean_offset, // Clean RDMA buffer auto rdma_buffer_ptr_int = static_cast(rdma_buffer_ptr); - #pragma unroll +#pragma unroll for (int i = thread_id; i < rdma_num_int_clean; i += num_threads) rdma_buffer_ptr_int[rdma_clean_offset + i] = 0; // Clean NVL buffer auto nvl_buffer_ptr_int = static_cast(buffer_ptrs[nvl_rank]); - #pragma unroll +#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(); @@ -1371,100 +1506,100 @@ __global__ void cached_notify(const int rdma_clean_offset, nvshmem_sync_with_same_gpu_idx(rdma_team); barrier_block(barrier_signal_ptrs, nvl_rank); } else if (sm_id == 1) { - if (is_cached_dispatch) - return; + if (not is_cached_dispatch) { + EP_DEVICE_ASSERT(num_warps >= num_channels); + EP_DEVICE_ASSERT(num_rdma_ranks <= 32); - 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); - // 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; + // 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(); + if (not is_cached_dispatch) { + 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 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; + 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(); } - 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(); } } } } + + notify_full_kernel_timer_end(normal_notify_stats); } void cached_notify(int hidden_int4, @@ -1488,7 +1623,8 @@ void cached_notify(int hidden_int4, int64_t num_rdma_bytes, int64_t num_nvl_bytes, bool is_cached_dispatch, - bool low_latency_mode) { + bool low_latency_mode, + NormalNotifyStats normal_notify_stats) { 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; @@ -1517,6 +1653,8 @@ void cached_notify(int hidden_int4, 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); + if (normal_notify_stats.timer_state != nullptr) + CUDA_CHECK(cudaMemsetAsync(normal_notify_stats.timer_state, 0, 2 * sizeof(int64_t), stream)); LAUNCH_KERNEL(&cfg, cached_notify_func, rdma_clean_meta.first, @@ -1535,7 +1673,8 @@ void cached_notify(int hidden_int4, rank, num_ranks, is_cached_dispatch, - cpu_rdma_team); + cpu_rdma_team, + normal_notify_stats); } template (tma_load_buffer(stage_idx, j) + lane_id); - #pragma unroll +#pragma unroll for (int k = 0; k < kDtypePerInt4; ++k) values[k] += static_cast(recv_value_dtypes[k]); } @@ -1629,7 +1768,7 @@ __device__ int combine_token(bool is_token_in_rank, // Copy into shared and issue TMA auto out_dtypes = reinterpret_cast(tma_store_buffer(stage_idx) + lane_id); - #pragma unroll +#pragma unroll for (int j = 0; j < kDtypePerInt4; ++j) out_dtypes[j] = static_cast(values[j]); tma_store_fence(); @@ -1643,7 +1782,7 @@ __device__ int combine_token(bool is_token_in_rank, // Flush all writes tma_store_wait<0>(); } else { - #pragma unroll +#pragma unroll for (int i = lane_id; i < hidden_int4; i += 32) { // Read bias // TODO: make it as a finer-grained template @@ -1656,7 +1795,7 @@ __device__ int combine_token(bool is_token_in_rank, // Read buffers // TODO: maybe too many registers here int4 recv_value_int4[kMaxNumRanks]; - #pragma unroll +#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)); @@ -1666,16 +1805,16 @@ __device__ int combine_token(bool is_token_in_rank, 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 +#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 +// 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 +#pragma unroll for (int k = 0; k < kDtypePerInt4; ++k) values[k] += static_cast(recv_value_dtypes[k]); } @@ -1683,7 +1822,7 @@ __device__ int combine_token(bool is_token_in_rank, // Cast back to `dtype_t` and write int4 out_int4; auto out_dtypes = reinterpret_cast(&out_int4); - #pragma unroll +#pragma unroll for (int j = 0; j < kDtypePerInt4; ++j) out_dtypes[j] = static_cast(values[j]); st_na_global(combined_row + i, out_int4); @@ -1693,7 +1832,7 @@ __device__ int combine_token(bool is_token_in_rank, // Reduce `topk_weights` if (lane_id < num_topk) { float value = 0; - #pragma unroll +#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); @@ -1711,7 +1850,7 @@ template 0) ? kNumCombineForwarderWarps / kNumRDMARanks : 1, - int kNumForwarders = kNumRDMARanks* kNumWarpsPerForwarder, + 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, @@ -1736,6 +1875,7 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + NormalCompletionStats normal_logical_recv_completion_stats, int rank, int num_ranks) { enum class WarpRole { kNVLSender, kNVLAndRDMAForwarder, kRDMAReceiver, kCoordinator }; @@ -1931,7 +2071,7 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co // NVL layouts void* local_nvl_buffer = buffer_ptrs[nvl_rank]; void* nvl_buffers[NUM_MAX_NVL_PEERS]; - #pragma unroll +#pragma unroll for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) nvl_buffers[i] = buffer_ptrs[i]; auto nvl_channel_x = @@ -2222,6 +2362,45 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co rdma_receiver_retired[warp_id] = true; } else { // Coordinator + const bool enable_logical_completion = normal_logical_recv_completion_stats.cost != nullptr; + int logical_expected_count[NUM_MAX_NVL_PEERS] = {0}; + int logical_last_expected_head[NUM_MAX_NVL_PEERS]; + uint32_t logical_expected_mask = 0, logical_recorded_mask = 0; + uint64_t logical_completion_start_time = 0; +#pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) + logical_last_expected_head[i] = -1; + + // Recover logical global-rank dependencies from the dispatch routing bitmap. + // One coordinator lane owns one source RDMA rank and its eight NVL ranks. + if (not is_forwarder_sm and enable_logical_completion and lane_id < kNumRDMARanks) { + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * sizeof(bool) == sizeof(uint64_t), + "The combine routing bitmap must be readable in aligned 64-bit groups"); + int token_start_idx, token_end_idx; + get_channel_task_range(num_combined_tokens, num_channels, channel_id, token_start_idx, token_end_idx); + for (int token_idx = token_start_idx; token_idx < token_end_idx; ++token_idx) { + // Torch allocations are sufficiently aligned, and both the row stride and lane offset are multiples of 8 bytes. + const auto src_rank_mask = __ldg( + reinterpret_cast(is_combined_token_in_rank + token_idx * num_ranks + lane_id * NUM_MAX_NVL_PEERS)); + if (src_rank_mask == 0) + continue; + const auto expected_head = ld_nc_global(combined_rdma_head + token_idx * kNumRDMARanks + lane_id); + // Probe reconstruction must remain fail-open if queue metadata is incomplete. + if (expected_head < 0) + continue; +#pragma unroll + for (int src_nvl_rank = 0; src_nvl_rank < NUM_MAX_NVL_PEERS; ++src_nvl_rank) { + const auto src_bit = 1u << src_nvl_rank; + if (((src_rank_mask >> (src_nvl_rank * 8)) & 0xffu) != 0) { + logical_expected_count[src_nvl_rank] += 1; + logical_last_expected_head[src_nvl_rank] = max(logical_last_expected_head[src_nvl_rank], expected_head); + logical_expected_mask |= src_bit; + } + } + } + logical_completion_start_time = clock64(); + } + // Sync shared memory status is_forwarder_sm ? sync_forwarder_smem() : sync_rdma_receiver_smem(); const auto num_warps_per_rdma_rank = kNumForwarders / kNumRDMARanks; @@ -2232,6 +2411,31 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co 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) { + // A logical source is complete once the node-level RDMA queue has + // published the last row that can contain that source's contribution. + if (not is_forwarder_sm and enable_logical_completion and lane_id < kNumRDMARanks and + logical_recorded_mask != logical_expected_mask) { + const auto observed_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); + const auto completion_end_time = clock64(); +#pragma unroll + for (int src_nvl_rank = 0; src_nvl_rank < NUM_MAX_NVL_PEERS; ++src_nvl_rank) { + const auto src_bit = 1u << src_nvl_rank; + if ((logical_expected_mask & src_bit) != 0 and (logical_recorded_mask & src_bit) == 0 and + observed_tail > logical_last_expected_head[src_nvl_rank]) { + const auto src_global_rank = lane_id * NUM_MAX_NVL_PEERS + src_nvl_rank; + atomicAdd(reinterpret_cast(normal_logical_recv_completion_stats.cost + src_global_rank), + completion_end_time - logical_completion_start_time); + atomicAdd( + reinterpret_cast(normal_logical_recv_completion_stats.sample_count + src_global_rank), + 1); + atomicAdd( + reinterpret_cast(normal_logical_recv_completion_stats.token_count + src_global_rank), + static_cast(logical_expected_count[src_nvl_rank])); + logical_recorded_mask |= src_bit; + } + } + } + // Retired if (not is_forwarder_sm and __all_sync(0xffffffff, lane_id >= kNumRDMAReceivers or rdma_receiver_retired[lane_id])) break; @@ -2241,7 +2445,7 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co // Find minimum head for RDMA ranks if (not is_forwarder_sm) { int min_head = std::numeric_limits::max(); - #pragma unroll +#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]); @@ -2255,11 +2459,11 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co last_rdma_head = min_head; } } else { - // Find minimum head for NVL ranks - #pragma unroll +// Find minimum head for NVL ranks +#pragma unroll for (int i = 0; i < kNumRDMARanks; ++i) { int min_head = std::numeric_limits::max(); - #pragma unroll +#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]); @@ -2299,6 +2503,7 @@ void combine(cudaDataType_t type, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + NormalCompletionStats normal_logical_recv_completion_stats, int rank, int num_ranks, cudaStream_t stream, @@ -2350,6 +2555,7 @@ void combine(cudaDataType_t type, buffer_ptrs, \ num_max_nvl_chunked_send_tokens, \ num_max_nvl_chunked_recv_tokens, \ + normal_logical_recv_completion_stats, \ rank, \ num_ranks); \ } \ diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 8327abded..0f551dabd 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -339,7 +339,8 @@ def dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], expert_alignment: int = 1, num_worst_tokens: int = 0, config: Optional[Config] = None, previous_event: Optional[EventOverlap] = None, async_finish: bool = False, - allocate_on_comm_stream: bool = False) -> \ + allocate_on_comm_stream: bool = False, + normal_dispatch_stats: Optional[Tuple[Optional[torch.Tensor], ...]] = None) -> \ Tuple[Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor], List[int], Tuple, EventOverlap]: """ @@ -369,6 +370,9 @@ def dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], 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. + normal_dispatch_stats: optional opaque bundle of ten cumulative + normal-mode probe tensors. Pass the value returned by the + external diagnostics provider directly. Internode only. Returns: recv_x: received tokens, the same type and tuple as the input `x`, but the number of tokens equals to the @@ -388,7 +392,7 @@ def dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], if self.runtime.get_num_rdma_ranks() > 1: return self.internode_dispatch(x, handle, num_tokens_per_rank, num_tokens_per_rdma_rank, is_token_in_rank, num_tokens_per_expert, topk_idx, topk_weights, expert_alignment, num_worst_tokens, config, - previous_event, async_finish, allocate_on_comm_stream) + previous_event, async_finish, allocate_on_comm_stream, normal_dispatch_stats) # Launch the kernel with cached or non-cached mode x, x_scales = x if isinstance(x, tuple) else (x, None) @@ -419,7 +423,8 @@ def combine(self, x: torch.Tensor, handle: Tuple, 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) -> \ + allocate_on_comm_stream: bool = False, + normal_combine_stats: Optional[Tuple[Optional[torch.Tensor], ...]] = None) -> \ Tuple[torch.Tensor, Optional[torch.Tensor], EventOverlap]: """ Combine (reduce) tokens (addition **without** weights) from different ranks, both intranode and internode @@ -437,6 +442,9 @@ def combine(self, x: torch.Tensor, handle: Tuple, 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. + normal_combine_stats: optional opaque bundle of five cumulative + normal-mode probe tensors. Pass the value returned by the + external diagnostics provider directly. Internode only. Returns: recv_x: the reduced token from its dispatched ranks. @@ -448,7 +456,8 @@ def combine(self, x: torch.Tensor, handle: Tuple, # Internode if self.runtime.get_num_rdma_ranks() > 1: - return self.internode_combine(x, handle, topk_weights, bias, config, previous_event, async_finish, allocate_on_comm_stream) + return self.internode_combine(x, handle, topk_weights, bias, config, previous_event, async_finish, + allocate_on_comm_stream, normal_combine_stats) # NOTES: the second `_` is for the sending side, so we should use the third one rank_prefix_matrix, _, channel_prefix_matrix, src_idx, is_recv_token_in_rank, send_head = handle @@ -469,7 +478,8 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te topk_idx: Optional[torch.Tensor] = None, topk_weights: Optional[torch.Tensor] = None, expert_alignment: int = 1, num_worst_tokens: int = 0, config: Optional[Config] = None, previous_event: Optional[EventOverlap] = None, async_finish: bool = False, - allocate_on_comm_stream: bool = False) -> \ + allocate_on_comm_stream: bool = False, + normal_dispatch_stats: Optional[Tuple[Optional[torch.Tensor], ...]] = None) -> \ Tuple[Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor], List[int], Tuple, EventOverlap]: """ @@ -490,8 +500,9 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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, num_worst_tokens, config, getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream) + rdma_channel_prefix_matrix, recv_rdma_rank_prefix_sum, + gbl_channel_prefix_matrix, recv_gbl_rank_prefix_sum, expert_alignment, num_worst_tokens, config, + getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream, normal_dispatch_stats) 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 @@ -503,7 +514,8 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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, num_worst_tokens, config, getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream) + expert_alignment, num_worst_tokens, config, getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream, + normal_dispatch_stats) 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) @@ -518,7 +530,8 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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) -> \ + allocate_on_comm_stream: bool = False, + normal_combine_stats: Optional[Tuple[Optional[torch.Tensor], ...]] = None) -> \ Tuple[torch.Tensor, Optional[torch.Tensor], EventOverlap]: """ Internode combine implementation, for more details, please refer to the `combine` docs. @@ -534,12 +547,10 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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) + 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, normal_combine_stats) return combined_x, combined_topk_weights, EventOverlap(event) def clean_low_latency_buffer(self, num_max_dispatch_tokens_per_rank: int, hidden: int, num_experts: int) -> None: diff --git a/setup.py b/setup.py index e89a9dbb4..f135107bd 100644 --- a/setup.py +++ b/setup.py @@ -4,7 +4,7 @@ import importlib from pathlib import Path -from torch.utils.cpp_extension import BuildExtension, CUDAExtension +from torch.utils.cpp_extension import BuildExtension, CUDAExtension, CUDA_HOME # Wheel specific: the wheels only include the soname of the host library `libnvshmem_host.so.X` @@ -39,6 +39,9 @@ def get_nvshmem_host_lib_name(base_dir): nvcc_flags = ['-O3', '-Xcompiler', '-O3'] sources = ['csrc/deep_ep.cpp', 'csrc/kernels/runtime.cu', 'csrc/kernels/layout.cu', 'csrc/kernels/intranode.cu'] include_dirs = ['csrc/'] + cuda_cccl_include = Path(CUDA_HOME, 'include', 'cccl') if CUDA_HOME else None + if cuda_cccl_include and cuda_cccl_include.exists(): + include_dirs.append(str(cuda_cccl_include)) library_dirs = [] nvcc_dlink = [] extra_link_args = ['-lcuda'] diff --git a/tests/test_internode.py b/tests/test_internode.py index 6530669da..8555a6d0a 100644 --- a/tests/test_internode.py +++ b/tests/test_internode.py @@ -22,7 +22,9 @@ def test_main(args: argparse.Namespace, rank: int, buffer: deep_ep.Buffer, group: dist.ProcessGroup, - skip_benchmark: bool = False): + skip_benchmark: bool = False, + normal_dispatch_stats=None, + normal_combine_stats=None): # 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 @@ -129,7 +131,8 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): 'is_token_in_rank': is_token_in_rank, 'num_tokens_per_expert': num_tokens_per_expert, 'config': config, - 'async_finish': async_mode + 'async_finish': async_mode, + 'normal_dispatch_stats': normal_dispatch_stats } if with_topk: dispatch_args.update({'topk_idx': topk_idx, 'topk_weights': topk_weights_pure_rand if is_rand else topk_weights}) @@ -188,7 +191,7 @@ def check_data(check_x, 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} + dispatch_args = {'x': current_x, 'handle': handle, 'config': config, 'async_finish': async_mode, 'normal_dispatch_stats': normal_dispatch_stats} if previous_mode: dispatch_args.update({'previous_event': buffer.capture()}) recv_x, _, _, _, _, event = buffer.dispatch(**dispatch_args) @@ -200,7 +203,7 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): # 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} + combine_args = {'x': recv_x, 'bias': (bias_0, bias_1), 'handle': handle, 'config': config, 'async_finish': async_mode, 'normal_combine_stats': normal_combine_stats} if with_topk: combine_args.update({'topk_weights': recv_topk_weights}) if previous_mode: @@ -242,7 +245,7 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): for nvl_chunk_size in range(4, 45, 4): for rdma_chunk_size in range(4, 33, 4): config = deep_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} + tune_args = {'x': current_x, 'handle': handle, 'config': config, 'normal_dispatch_stats': normal_dispatch_stats} t, notify_t = bench_kineto( lambda: buffer.dispatch(**tune_args), # noqa: B023 ('dispatch', 'notify'), @@ -278,7 +281,8 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): '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 + 'config': dispatch_config if dispatch_config is not None else config, + 'normal_dispatch_stats': normal_dispatch_stats } recv_x, _, _, _, handle, _ = buffer.dispatch(**dispatch_args) @@ -287,7 +291,7 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): 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 = deep_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} + tune_args = {'x': recv_x, 'handle': handle, 'config': config, 'normal_combine_stats': normal_combine_stats} t, notify_t = bench_kineto( lambda: buffer.combine(**tune_args), # noqa: B023 ('combine', 'notify'), @@ -330,6 +334,23 @@ def test_loop(local_rank: int, num_local_ranks: int, args: argparse.Namespace): explicitly_destroy=True) assert num_local_ranks == 8 and num_ranks > 8 + diagnose = None + normal_dispatch_stats = None + normal_combine_stats = None + if args.test_normal_deepxtrace: + os.environ['DEEPEP_DIAGNOSE_ENABLE'] = '1' + from deepxtrace import diagnose as ds + + diagnose = ds.Diagnose.create_from_env( + group=group, + enable_ll_diagnose=False, + enable_normal_diagnose=True, + snapshot_stream=buffer.get_comm_stream()) + assert diagnose is not None + normal_dispatch_stats = diagnose.get_normal_dispatch_stats_tensor() + normal_combine_stats = diagnose.get_normal_combine_stats_tensor() + diagnose.start() + for seed in range(int(1e9)): if local_rank == 0: print(f'Testing with seed {seed} ...', flush=True) @@ -337,7 +358,9 @@ def test_loop(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ref_hash = 0 for i in (num_sms, ): ref_hash += test_main(args, i, local_rank, num_local_ranks, num_ranks, num_nodes, rank, buffer, group, - args.pressure_test_mode == 1) + args.pressure_test_mode == 1, normal_dispatch_stats, normal_combine_stats) + if diagnose is not None and not diagnose.enable_async: + diagnose.diagnose_normal_sync(diagnose_step=1) if local_rank == 0: print('', flush=True) if args.pressure_test_mode == 0: @@ -352,7 +375,9 @@ def test_loop(local_rank: int, num_local_ranks: int, args: argparse.Namespace): current_hash = 0 for i in (num_sms, ): current_hash += test_main(args, i, local_rank, num_local_ranks, num_ranks, num_nodes, rank, buffer, group, - args.pressure_test_mode == 1) + args.pressure_test_mode == 1, normal_dispatch_stats, normal_combine_stats) + if diagnose is not None and not diagnose.enable_async: + diagnose.diagnose_normal_sync(diagnose_step=1) if local_rank == 0: print('', flush=True) assert current_hash == ref_hash @@ -363,6 +388,8 @@ def test_loop(local_rank: int, num_local_ranks: int, args: argparse.Namespace): test_low_latency.test_main(ll_num_tokens, ll_hidden, ll_num_experts, ll_num_topk, rank, num_ranks, group, buffer, seed=1) # Destroy the buffer runtime and communication group + if diagnose is not None: + diagnose.stop() buffer.destroy() dist.barrier() dist.destroy_process_group() @@ -382,6 +409,7 @@ def test_loop(local_rank: int, num_local_ranks: int, args: argparse.Namespace): help='Pressure test mode. 0: don\'t do pressure test, 1: do pressure test without benchmarks, 2: do pressure test with benchmarks') 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('--test-normal-deepxtrace', action='store_true', help='whether to test DeepXTrace with normal kernels') args = parser.parse_args() # Set default `num_topk_groups` if not provided