From 699d17ba2767b119121b6a63f52aef6127607b63 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Tue, 28 Jul 2026 16:42:52 +0800 Subject: [PATCH 01/16] fix(build): locate CUDA 13 CCCL headers --- setup.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) 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'] From 5867e09bf50700cd8a09ffe51336009734cbafbe Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Tue, 4 Aug 2026 09:59:52 +0800 Subject: [PATCH 02/16] feat(deepxtrace): validate and bind normal stats contract --- deep_ep/buffer.py | 45 ++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 44 insertions(+), 1 deletion(-) diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 8327abded..f1d6fcc60 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -1,7 +1,7 @@ import os import torch import torch.distributed as dist -from typing import Callable, List, Tuple, Optional, Union +from typing import Callable, Dict, List, Tuple, Optional, Union # noinspection PyUnresolvedReferences import deep_ep_cpp @@ -9,6 +9,35 @@ from deep_ep_cpp import Config, EventHandle from .utils import EventOverlap, check_nvlink_connections +_REQUIRED_NORMAL_STATS_SCHEMA = ( + "normal_notify_dispatch_full_kernel_duration_ns_stats", + "normal_notify_dispatch_full_kernel_count_stats", + "normal_cached_notify_dispatch_full_kernel_duration_ns_stats", + "normal_cached_notify_dispatch_full_kernel_count_stats", + "normal_cached_notify_combine_full_kernel_duration_ns_stats", + "normal_cached_notify_combine_full_kernel_count_stats", + "normal_dispatch_final_completion_cost_stats", + "normal_dispatch_final_completion_sample_count_stats", + "normal_dispatch_final_completion_token_count_stats", + "normal_dispatch_rdma_recv_completion_cost_stats", + "normal_dispatch_rdma_recv_completion_sample_count_stats", + "normal_dispatch_rdma_recv_completion_token_count_stats", + "normal_combine_logical_recv_completion_cost_stats", + "normal_combine_logical_recv_completion_sample_count_stats", + "normal_combine_logical_recv_completion_token_count_stats", +) + + +def _validate_deepxtrace_normal_stats_schema(diagnose_module) -> None: + actual_schema = getattr( + diagnose_module, "NORMAL_STATS_SCHEMA", None) + if tuple(actual_schema or ()) != _REQUIRED_NORMAL_STATS_SCHEMA: + raise RuntimeError( + "Incompatible deepxtrace normal-stats schema: DeepEP requires " + "deepxtrace>=0.2.0,<0.3.0 with fields ordered as " + "notify -> dispatch -> combine, but found schema " + f"{actual_schema!r}") + class Buffer: """ @@ -84,6 +113,10 @@ def all_gather_object(obj): return comm.allgather(obj) else: raise ValueError("Either 'group' or 'comm' must be provided.") + self.diagnose = None + self._normal_diagnose_stats = { + name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA + } self.num_nvl_bytes = num_nvl_bytes self.num_rdma_bytes = num_rdma_bytes self.low_latency_mode = low_latency_mode @@ -139,6 +172,16 @@ def all_gather_object(obj): self.runtime.sync(device_ids, ipc_handles, root_unique_id) assert self.runtime.is_available() + def _get_normal_diagnose_stats( + self) -> Dict[str, Optional[torch.Tensor]]: + tensors = self.diagnose.get_stats_normal_stats_tensor() + if len(tensors) != len(_REQUIRED_NORMAL_STATS_SCHEMA): + raise RuntimeError( + "DeepXTrace normal-stats tensor count does not match its " + "schema: " + f"{len(tensors)} != {len(_REQUIRED_NORMAL_STATS_SCHEMA)}") + return dict(zip(_REQUIRED_NORMAL_STATS_SCHEMA, tensors)) + @staticmethod def disable_ll_layered() -> bool: disable_ll_layered = False From 4365d7702371890351f1a27dcb01118434228858 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 27 Jul 2026 15:15:38 +0800 Subject: [PATCH 03/16] feat(normal): add notify full-kernel duration probes --- csrc/deep_ep.cpp | 83 ++++++++++++- csrc/deep_ep.hpp | 13 ++- csrc/kernels/api.cuh | 10 +- csrc/kernels/internode.cu | 240 ++++++++++++++++++++++++-------------- deep_ep/buffer.py | 81 ++++++++++++- 5 files changed, 329 insertions(+), 98 deletions(-) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index 714774c81..d844e694c 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -946,7 +946,13 @@ 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 std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats, + const std::optional& normal_notify_dispatch_full_kernel_count_stats, + const std::optional& normal_notify_dispatch_full_kernel_timer_state, + const std::optional& normal_cached_notify_dispatch_full_kernel_duration_ns_stats, + const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats, + const std::optional& normal_cached_notify_dispatch_full_kernel_timer_state) { #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 +1010,26 @@ 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); } + auto check_normal_notify_stats = [](const std::optional& duration_ns_stats, + const std::optional& count_stats, + const std::optional& timer_state) { + const bool enabled = duration_ns_stats.has_value(); + EP_HOST_ASSERT(count_stats.has_value() == enabled and timer_state.has_value() == enabled); + if (enabled) { + EP_HOST_ASSERT(duration_ns_stats->scalar_type() == torch::kInt64 and count_stats->scalar_type() == torch::kInt64 and + timer_state->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(duration_ns_stats->dim() == 1 and duration_ns_stats->numel() == 1 and + duration_ns_stats->is_contiguous()); + EP_HOST_ASSERT(count_stats->dim() == 1 and count_stats->numel() == 1 and count_stats->is_contiguous()); + EP_HOST_ASSERT(timer_state->dim() == 1 and timer_state->numel() == 2 and timer_state->is_contiguous()); + } + }; + check_normal_notify_stats(normal_notify_dispatch_full_kernel_duration_ns_stats, + normal_notify_dispatch_full_kernel_count_stats, + normal_notify_dispatch_full_kernel_timer_state); + check_normal_notify_stats(normal_cached_notify_dispatch_full_kernel_duration_ns_stats, + normal_cached_notify_dispatch_full_kernel_count_stats, + normal_cached_notify_dispatch_full_kernel_timer_state); 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 +1120,16 @@ 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, + normal_cached_notify_dispatch_full_kernel_duration_ns_stats.has_value() + ? normal_cached_notify_dispatch_full_kernel_duration_ns_stats->data_ptr() + : nullptr, + normal_cached_notify_dispatch_full_kernel_count_stats.has_value() + ? normal_cached_notify_dispatch_full_kernel_count_stats->data_ptr() + : nullptr, + normal_cached_notify_dispatch_full_kernel_timer_state.has_value() + ? normal_cached_notify_dispatch_full_kernel_timer_state->data_ptr() + : nullptr); } 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 +1169,16 @@ 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, + normal_notify_dispatch_full_kernel_duration_ns_stats.has_value() + ? normal_notify_dispatch_full_kernel_duration_ns_stats->data_ptr() + : nullptr, + normal_notify_dispatch_full_kernel_count_stats.has_value() + ? normal_notify_dispatch_full_kernel_count_stats->data_ptr() + : nullptr, + normal_notify_dispatch_full_kernel_timer_state.has_value() + ? normal_notify_dispatch_full_kernel_timer_state->data_ptr() + : nullptr); // Synchronize total received tokens and tokens per expert if (num_worst_tokens > 0) { @@ -1324,7 +1368,10 @@ std::tuple, std::optional& previous_event, bool async, - bool allocate_on_comm_stream) { + bool allocate_on_comm_stream, + const std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats, + const std::optional& normal_cached_notify_combine_full_kernel_count_stats, + const std::optional& normal_cached_notify_combine_full_kernel_timer_state) { #ifndef DISABLE_NVSHMEM const int num_channels = config.num_sms / 2; EP_HOST_ASSERT(config.num_sms % 2 == 0); @@ -1356,6 +1403,23 @@ std::tuple, std::optionalscalar_type() == torch::kInt64 and + normal_cached_notify_combine_full_kernel_count_stats->scalar_type() == torch::kInt64 and + normal_cached_notify_combine_full_kernel_timer_state->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_duration_ns_stats->dim() == 1 and + normal_cached_notify_combine_full_kernel_duration_ns_stats->numel() == 1 and + normal_cached_notify_combine_full_kernel_duration_ns_stats->is_contiguous()); + EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_count_stats->dim() == 1 and + normal_cached_notify_combine_full_kernel_count_stats->numel() == 1 and + normal_cached_notify_combine_full_kernel_count_stats->is_contiguous()); + EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_timer_state->dim() == 1 and + normal_cached_notify_combine_full_kernel_timer_state->numel() == 2 and + normal_cached_notify_combine_full_kernel_timer_state->is_contiguous()); + } // Allocate all tensors on comm stream if set // NOTES: do not allocate tensors upfront! @@ -1413,7 +1477,16 @@ std::tuple, std::optionaldata_ptr() + : nullptr, + normal_cached_notify_combine_enabled + ? normal_cached_notify_combine_full_kernel_count_stats->data_ptr() + : nullptr, + normal_cached_notify_combine_enabled + ? normal_cached_notify_combine_full_kernel_timer_state->data_ptr() + : nullptr); // Assign bias pointers auto bias_opts = std::vector>({bias_0, bias_1}); diff --git a/csrc/deep_ep.hpp b/csrc/deep_ep.hpp index 090e5a4f1..de8ce50d5 100644 --- a/csrc/deep_ep.hpp +++ b/csrc/deep_ep.hpp @@ -235,7 +235,13 @@ struct Buffer { const Config& config, std::optional& previous_event, bool async, - bool allocate_on_comm_stream); + bool allocate_on_comm_stream, + const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, + const std::optional& normal_notify_dispatch_full_kernel_count_stats = std::nullopt, + const std::optional& normal_notify_dispatch_full_kernel_timer_state = std::nullopt, + const std::optional& normal_cached_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, + const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats = std::nullopt, + const std::optional& normal_cached_notify_dispatch_full_kernel_timer_state = std::nullopt); std::tuple, std::optional> internode_combine( const torch::Tensor& x, @@ -252,7 +258,10 @@ struct Buffer { const Config& config, std::optional& previous_event, bool async, - bool allocate_on_comm_stream); + bool allocate_on_comm_stream, + const std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats = std::nullopt, + const std::optional& normal_cached_notify_combine_full_kernel_count_stats = std::nullopt, + const std::optional& normal_cached_notify_combine_full_kernel_timer_state = std::nullopt); 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..b7938e2b9 100644 --- a/csrc/kernels/api.cuh +++ b/csrc/kernels/api.cuh @@ -173,7 +173,10 @@ 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, + int64_t* normal_notify_dispatch_full_kernel_duration_ns_stats, + int64_t* normal_notify_dispatch_full_kernel_count_stats, + int64_t* normal_notify_dispatch_full_kernel_timer_state); void dispatch(void* recv_x, float* recv_x_scales, @@ -235,7 +238,10 @@ 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, + int64_t* normal_cached_notify_full_kernel_duration_ns_stats, + int64_t* normal_cached_notify_full_kernel_count_stats, + int64_t* normal_cached_notify_full_kernel_timer_state); void combine(cudaDataType_t type, void* combined_x, diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index 48c6c0018..ac43c3378 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -14,6 +14,42 @@ 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(int64_t* timer_state) { + if (timer_state == nullptr) + return; + + if (threadIdx.x == 0) { + auto state = reinterpret_cast(timer_state); + atomicMax(state, ~read_globaltimer_ns()); + } + __syncthreads(); +} + +__device__ __forceinline__ void notify_full_kernel_timer_end(int64_t* timer_state, + int64_t* duration_ns_stats, + int64_t* count_stats) { + if (timer_state == nullptr) + return; + + __syncthreads(); + if (threadIdx.x == 0) { + auto state = reinterpret_cast(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(duration_ns_stats), duration_ns); + atomicAdd(reinterpret_cast(count_stats), 1ull); + } + } +} + struct SourceMeta { int src_rdma_rank, is_token_in_nvl_rank_bits; @@ -115,7 +151,10 @@ __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, + int64_t* normal_notify_dispatch_full_kernel_duration_ns_stats, + int64_t* normal_notify_dispatch_full_kernel_count_stats, + int64_t* normal_notify_dispatch_full_kernel_timer_state) { 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 +162,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_dispatch_full_kernel_timer_state); + if (sm_id == 0) { // Communication with others // Global barrier: the first warp does intra-node sync, the second warp does internode sync @@ -341,6 +382,10 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, prefix_row[i] += prefix_row[i - 1]; } } + + notify_full_kernel_timer_end(normal_notify_dispatch_full_kernel_timer_state, + normal_notify_dispatch_full_kernel_duration_ns_stats, + normal_notify_dispatch_full_kernel_count_stats); } void notify_dispatch(const int* num_tokens_per_rank, @@ -372,7 +417,10 @@ 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, + int64_t* normal_notify_dispatch_full_kernel_duration_ns_stats, + int64_t* normal_notify_dispatch_full_kernel_count_stats, + int64_t* normal_notify_dispatch_full_kernel_timer_state) { #define NOTIFY_DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ { \ auto notify_dispatch_func = low_latency_mode ? notify_dispatch : notify_dispatch; \ @@ -403,7 +451,10 @@ void notify_dispatch(const int* num_tokens_per_rank, buffer_ptrs, \ barrier_signal_ptrs, \ rank, \ - cpu_rdma_team); \ + cpu_rdma_team, \ + normal_notify_dispatch_full_kernel_duration_ns_stats, \ + normal_notify_dispatch_full_kernel_count_stats, \ + normal_notify_dispatch_full_kernel_timer_state); \ } \ break @@ -429,6 +480,8 @@ void notify_dispatch(const int* num_tokens_per_rank, // Launch kernel SETUP_LAUNCH_CONFIG(1 + num_rdma_ranks, kNumThreads, stream); + if (normal_notify_dispatch_full_kernel_timer_state != nullptr) + CUDA_CHECK(cudaMemsetAsync(normal_notify_dispatch_full_kernel_timer_state, 0, 2 * sizeof(int64_t), stream)); SWITCH_RDMA_RANKS(NOTIFY_DISPATCH_LAUNCH_CASE); #undef NOTIFY_DISPATCH_LAUNCH_CASE } @@ -1324,7 +1377,10 @@ __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, + int64_t* normal_cached_notify_full_kernel_duration_ns_stats, + int64_t* normal_cached_notify_full_kernel_count_stats, + int64_t* normal_cached_notify_full_kernel_timer_state) { auto sm_id = static_cast(blockIdx.x); auto thread_id = static_cast(threadIdx.x); auto num_threads = static_cast(blockDim.x); @@ -1336,6 +1392,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_cached_notify_full_kernel_timer_state); + // 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; @@ -1371,100 +1429,102 @@ __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_cached_notify_full_kernel_timer_state, + normal_cached_notify_full_kernel_duration_ns_stats, + normal_cached_notify_full_kernel_count_stats); } void cached_notify(int hidden_int4, @@ -1488,7 +1548,10 @@ 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, + int64_t* normal_cached_notify_full_kernel_duration_ns_stats, + int64_t* normal_cached_notify_full_kernel_count_stats, + int64_t* normal_cached_notify_full_kernel_timer_state) { 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 +1580,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_cached_notify_full_kernel_timer_state != nullptr) + CUDA_CHECK(cudaMemsetAsync(normal_cached_notify_full_kernel_timer_state, 0, 2 * sizeof(int64_t), stream)); LAUNCH_KERNEL(&cfg, cached_notify_func, rdma_clean_meta.first, @@ -1535,7 +1600,10 @@ void cached_notify(int hidden_int4, rank, num_ranks, is_cached_dispatch, - cpu_rdma_team); + cpu_rdma_team, + normal_cached_notify_full_kernel_duration_ns_stats, + normal_cached_notify_full_kernel_count_stats, + normal_cached_notify_full_kernel_timer_state); } template None: + """Create persistent per-launch notify timer scratch owned by DeepEP. + + Each row is ``[inverted earliest start, completed blocks]``. C++ resets + the selected row on the communication stream immediately before the + corresponding kernel launch, so the pointers remain CUDA-graph stable. + """ + self._normal_notify_full_kernel_timer_states = torch.zeros( + (3, 2), dtype=torch.int64, device="cuda") + self._normal_notify_dispatch_full_kernel_timer_state = \ + self._normal_notify_full_kernel_timer_states[0] + self._normal_cached_notify_dispatch_full_kernel_timer_state = \ + self._normal_notify_full_kernel_timer_states[1] + self._normal_cached_notify_combine_full_kernel_timer_state = \ + self._normal_notify_full_kernel_timer_states[2] + + @staticmethod + def _select_normal_notify_timer_state( + duration_stats: Optional[torch.Tensor], + count_stats: Optional[torch.Tensor], + timer_state: Optional[torch.Tensor]) -> Optional[torch.Tensor]: + if (duration_stats is None) != (count_stats is None): + raise RuntimeError( + "DeepXTrace notify duration/count tensors must be enabled " + "or disabled together") + return timer_state if duration_stats is not None else None + def _get_normal_diagnose_stats( self) -> Dict[str, Optional[torch.Tensor]]: tensors = self.diagnose.get_stats_normal_stats_tensor() @@ -523,6 +554,25 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te # Launch the kernel with cached or non-cached mode x, x_scales = x if isinstance(x, tuple) else (x, None) + normal_stats = self._normal_diagnose_stats + normal_notify_dispatch_full_kernel_duration_ns_stats = normal_stats[ + "normal_notify_dispatch_full_kernel_duration_ns_stats"] + normal_notify_dispatch_full_kernel_count_stats = normal_stats[ + "normal_notify_dispatch_full_kernel_count_stats"] + normal_notify_dispatch_full_kernel_timer_state = \ + self._select_normal_notify_timer_state( + normal_notify_dispatch_full_kernel_duration_ns_stats, + normal_notify_dispatch_full_kernel_count_stats, + self._normal_notify_dispatch_full_kernel_timer_state) + normal_cached_notify_dispatch_full_kernel_duration_ns_stats = normal_stats[ + "normal_cached_notify_dispatch_full_kernel_duration_ns_stats"] + normal_cached_notify_dispatch_full_kernel_count_stats = normal_stats[ + "normal_cached_notify_dispatch_full_kernel_count_stats"] + normal_cached_notify_dispatch_full_kernel_timer_state = \ + self._select_normal_notify_timer_state( + normal_cached_notify_dispatch_full_kernel_duration_ns_stats, + normal_cached_notify_dispatch_full_kernel_count_stats, + self._normal_cached_notify_dispatch_full_kernel_timer_state) if handle is not None: assert topk_idx is None and topk_weights is None is_token_in_rank, \ @@ -534,7 +584,13 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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) + expert_alignment, num_worst_tokens, config, getattr(previous_event, 'event', None), async_finish, allocate_on_comm_stream, + normal_notify_dispatch_full_kernel_duration_ns_stats, + normal_notify_dispatch_full_kernel_count_stats, + normal_notify_dispatch_full_kernel_timer_state, + normal_cached_notify_dispatch_full_kernel_duration_ns_stats, + normal_cached_notify_dispatch_full_kernel_count_stats, + normal_cached_notify_dispatch_full_kernel_timer_state) 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 @@ -546,7 +602,13 @@ 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_notify_dispatch_full_kernel_duration_ns_stats, + normal_notify_dispatch_full_kernel_count_stats, + normal_notify_dispatch_full_kernel_timer_state, + normal_cached_notify_dispatch_full_kernel_duration_ns_stats, + normal_cached_notify_dispatch_full_kernel_count_stats, + normal_cached_notify_dispatch_full_kernel_timer_state) 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) @@ -575,6 +637,16 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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) + normal_stats = self._normal_diagnose_stats + normal_cached_notify_combine_full_kernel_duration_ns_stats = normal_stats[ + "normal_cached_notify_combine_full_kernel_duration_ns_stats"] + normal_cached_notify_combine_full_kernel_count_stats = normal_stats[ + "normal_cached_notify_combine_full_kernel_count_stats"] + normal_cached_notify_combine_full_kernel_timer_state = \ + self._select_normal_notify_timer_state( + normal_cached_notify_combine_full_kernel_duration_ns_stats, + normal_cached_notify_combine_full_kernel_count_stats, + self._normal_cached_notify_combine_full_kernel_timer_state) # Launch the kernel combined_x, combined_topk_weights, event = self.runtime.internode_combine(x, topk_weights, bias_0, bias_1, src_meta, @@ -582,7 +654,10 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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) + None), async_finish, allocate_on_comm_stream, + normal_cached_notify_combine_full_kernel_duration_ns_stats, + normal_cached_notify_combine_full_kernel_count_stats, + normal_cached_notify_combine_full_kernel_timer_state) 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: From f46f3ff223732cfa4c4192c241b3f06b6eaa9259 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 27 Jul 2026 15:19:14 +0800 Subject: [PATCH 04/16] feat(normal): add dispatch final completion probes --- csrc/deep_ep.cpp | 25 +++++++++++++++++++++++++ csrc/deep_ep.hpp | 3 +++ csrc/kernels/api.cuh | 3 +++ csrc/kernels/internode.cu | 35 +++++++++++++++++++++++++++++++++++ deep_ep/buffer.py | 12 ++++++++++++ 5 files changed, 78 insertions(+) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index d844e694c..c01c41a31 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -947,6 +947,9 @@ Buffer::internode_dispatch(const torch::Tensor& x, std::optional& previous_event, bool async, bool allocate_on_comm_stream, + const std::optional& normal_dispatch_final_completion_cost_stats, + const std::optional& normal_dispatch_final_completion_sample_count_stats, + const std::optional& normal_dispatch_final_completion_token_count_stats, const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats, const std::optional& normal_notify_dispatch_full_kernel_count_stats, const std::optional& normal_notify_dispatch_full_kernel_timer_state, @@ -1010,6 +1013,19 @@ 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); } + auto check_normal_dispatch_stat_tensor = [=](const std::optional& stats) { + if (stats.has_value()) { + EP_HOST_ASSERT(stats->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(stats->dim() == 1 and stats->is_contiguous()); + EP_HOST_ASSERT(stats->size(0) == num_ranks); + } + }; + const bool enable_normal_dispatch_final_completion_stats = normal_dispatch_final_completion_cost_stats.has_value(); + EP_HOST_ASSERT(normal_dispatch_final_completion_sample_count_stats.has_value() == enable_normal_dispatch_final_completion_stats); + EP_HOST_ASSERT(normal_dispatch_final_completion_token_count_stats.has_value() == enable_normal_dispatch_final_completion_stats); + check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_cost_stats); + check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_sample_count_stats); + check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_token_count_stats); auto check_normal_notify_stats = [](const std::optional& duration_ns_stats, const std::optional& count_stats, const std::optional& timer_state) { @@ -1279,6 +1295,15 @@ Buffer::internode_dispatch(const torch::Tensor& x, buffer_ptrs_gpu, config.num_max_nvl_chunked_send_tokens, config.num_max_nvl_chunked_recv_tokens, + normal_dispatch_final_completion_cost_stats.has_value() + ? normal_dispatch_final_completion_cost_stats->data_ptr() + : nullptr, + normal_dispatch_final_completion_sample_count_stats.has_value() + ? normal_dispatch_final_completion_sample_count_stats->data_ptr() + : nullptr, + normal_dispatch_final_completion_token_count_stats.has_value() + ? normal_dispatch_final_completion_token_count_stats->data_ptr() + : nullptr, rank, num_ranks, cached_mode, diff --git a/csrc/deep_ep.hpp b/csrc/deep_ep.hpp index de8ce50d5..833928126 100644 --- a/csrc/deep_ep.hpp +++ b/csrc/deep_ep.hpp @@ -236,6 +236,9 @@ struct Buffer { std::optional& previous_event, bool async, bool allocate_on_comm_stream, + const std::optional& normal_dispatch_final_completion_cost_stats = std::nullopt, + const std::optional& normal_dispatch_final_completion_sample_count_stats = std::nullopt, + const std::optional& normal_dispatch_final_completion_token_count_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_count_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_timer_state = std::nullopt, diff --git a/csrc/kernels/api.cuh b/csrc/kernels/api.cuh index b7938e2b9..848c32774 100644 --- a/csrc/kernels/api.cuh +++ b/csrc/kernels/api.cuh @@ -210,6 +210,9 @@ void dispatch(void* recv_x, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + int64_t* normal_dispatch_final_completion_cost_stats, + int64_t* normal_dispatch_final_completion_sample_count_stats, + int64_t* normal_dispatch_final_completion_token_count_stats, int rank, int num_ranks, bool is_cached_dispatch, diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index ac43c3378..4afec9f17 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -530,6 +530,9 @@ __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, + int64_t* normal_dispatch_final_completion_cost_stats, + int64_t* normal_dispatch_final_completion_sample_count_stats, + int64_t* normal_dispatch_final_completion_token_count_stats, int rank, int num_ranks) { enum class WarpRole { kRDMASender, kRDMASenderCoordinator, kRDMAAndNVLForwarder, kForwarderCoordinator, kNVLReceivers }; @@ -541,6 +544,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_dispatch_final_completion_cost_stats != nullptr; EP_DEVICE_ASSERT(ibgda_get_state()->num_rc_per_pe == num_channels or ibgda_get_state()->num_rc_per_pe >= num_sms); @@ -1144,6 +1148,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) @@ -1181,6 +1188,15 @@ __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); @@ -1235,6 +1251,16 @@ __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_dispatch_final_completion_cost_stats + completion_stats_idx), + clock64() - completion_start_time); + atomicAdd( + reinterpret_cast(normal_dispatch_final_completion_sample_count_stats + completion_stats_idx), + 1); + atomicAdd( + reinterpret_cast(normal_dispatch_final_completion_token_count_stats + completion_stats_idx), + static_cast(completion_expected_count)); + } } // Move queue @@ -1292,6 +1318,9 @@ void dispatch(void* recv_x, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + int64_t* normal_dispatch_final_completion_cost_stats, + int64_t* normal_dispatch_final_completion_sample_count_stats, + int64_t* normal_dispatch_final_completion_token_count_stats, int rank, int num_ranks, bool is_cached_dispatch, @@ -1304,6 +1333,9 @@ 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_dispatch_final_completion_cost_stats != nullptr; + EP_HOST_ASSERT((normal_dispatch_final_completion_sample_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); + EP_HOST_ASSERT((normal_dispatch_final_completion_token_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); #define DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ { \ @@ -1347,6 +1379,9 @@ void dispatch(void* recv_x, buffer_ptrs, \ num_max_nvl_chunked_send_tokens, \ num_max_nvl_chunked_recv_tokens, \ + normal_dispatch_final_completion_cost_stats, \ + normal_dispatch_final_completion_sample_count_stats, \ + normal_dispatch_final_completion_token_count_stats, \ rank, \ num_ranks); \ } \ diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 96da10e1e..6d6d92493 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -573,6 +573,12 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te normal_cached_notify_dispatch_full_kernel_duration_ns_stats, normal_cached_notify_dispatch_full_kernel_count_stats, self._normal_cached_notify_dispatch_full_kernel_timer_state) + normal_dispatch_final_completion_cost_stats = normal_stats[ + "normal_dispatch_final_completion_cost_stats"] + normal_dispatch_final_completion_sample_count_stats = normal_stats[ + "normal_dispatch_final_completion_sample_count_stats"] + normal_dispatch_final_completion_token_count_stats = normal_stats[ + "normal_dispatch_final_completion_token_count_stats"] if handle is not None: assert topk_idx is None and topk_weights is None is_token_in_rank, \ @@ -585,6 +591,9 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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, + normal_dispatch_final_completion_cost_stats, + normal_dispatch_final_completion_sample_count_stats, + normal_dispatch_final_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, normal_notify_dispatch_full_kernel_count_stats, normal_notify_dispatch_full_kernel_timer_state, @@ -603,6 +612,9 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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, + normal_dispatch_final_completion_cost_stats, + normal_dispatch_final_completion_sample_count_stats, + normal_dispatch_final_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, normal_notify_dispatch_full_kernel_count_stats, normal_notify_dispatch_full_kernel_timer_state, From b10f01d9adf57d716202ae49e4793d3c49104966 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Tue, 4 Aug 2026 10:00:52 +0800 Subject: [PATCH 05/16] feat(normal): add dispatch RDMA receive completion probes --- csrc/deep_ep.cpp | 29 +++++++++++++ csrc/deep_ep.hpp | 3 ++ csrc/kernels/api.cuh | 3 ++ csrc/kernels/internode.cu | 90 +++++++++++++++++++++++++++++++++++++++ deep_ep/buffer.py | 12 ++++++ 5 files changed, 137 insertions(+) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index c01c41a31..8f0fc164a 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -950,6 +950,9 @@ Buffer::internode_dispatch(const torch::Tensor& x, const std::optional& normal_dispatch_final_completion_cost_stats, const std::optional& normal_dispatch_final_completion_sample_count_stats, const std::optional& normal_dispatch_final_completion_token_count_stats, + const std::optional& normal_dispatch_rdma_recv_completion_cost_stats, + const std::optional& normal_dispatch_rdma_recv_completion_sample_count_stats, + const std::optional& normal_dispatch_rdma_recv_completion_token_count_stats, const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats, const std::optional& normal_notify_dispatch_full_kernel_count_stats, const std::optional& normal_notify_dispatch_full_kernel_timer_state, @@ -1013,6 +1016,20 @@ 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); } + auto check_normal_dispatch_stat_pair = [=](const std::optional& cost_stats, + const std::optional& count_stats) { + if (cost_stats.has_value()) { + EP_HOST_ASSERT(cost_stats->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(cost_stats->dim() == 1 and cost_stats->is_contiguous()); + EP_HOST_ASSERT(cost_stats->size(0) == num_ranks); + EP_HOST_ASSERT(count_stats.has_value()); + EP_HOST_ASSERT(count_stats->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(count_stats->dim() == 1 and count_stats->is_contiguous()); + EP_HOST_ASSERT(count_stats->size(0) == num_ranks); + } else { + EP_HOST_ASSERT(not count_stats.has_value()); + } + }; auto check_normal_dispatch_stat_tensor = [=](const std::optional& stats) { if (stats.has_value()) { EP_HOST_ASSERT(stats->scalar_type() == torch::kInt64); @@ -1026,6 +1043,9 @@ Buffer::internode_dispatch(const torch::Tensor& x, check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_cost_stats); check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_sample_count_stats); check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_token_count_stats); + check_normal_dispatch_stat_pair(normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats); + check_normal_dispatch_stat_tensor(normal_dispatch_rdma_recv_completion_token_count_stats); auto check_normal_notify_stats = [](const std::optional& duration_ns_stats, const std::optional& count_stats, const std::optional& timer_state) { @@ -1304,6 +1324,15 @@ Buffer::internode_dispatch(const torch::Tensor& x, normal_dispatch_final_completion_token_count_stats.has_value() ? normal_dispatch_final_completion_token_count_stats->data_ptr() : nullptr, + normal_dispatch_rdma_recv_completion_cost_stats.has_value() + ? normal_dispatch_rdma_recv_completion_cost_stats->data_ptr() + : nullptr, + normal_dispatch_rdma_recv_completion_sample_count_stats.has_value() + ? normal_dispatch_rdma_recv_completion_sample_count_stats->data_ptr() + : nullptr, + normal_dispatch_rdma_recv_completion_token_count_stats.has_value() + ? normal_dispatch_rdma_recv_completion_token_count_stats->data_ptr() + : nullptr, rank, num_ranks, cached_mode, diff --git a/csrc/deep_ep.hpp b/csrc/deep_ep.hpp index 833928126..8fecf02ea 100644 --- a/csrc/deep_ep.hpp +++ b/csrc/deep_ep.hpp @@ -239,6 +239,9 @@ struct Buffer { const std::optional& normal_dispatch_final_completion_cost_stats = std::nullopt, const std::optional& normal_dispatch_final_completion_sample_count_stats = std::nullopt, const std::optional& normal_dispatch_final_completion_token_count_stats = std::nullopt, + const std::optional& normal_dispatch_rdma_recv_completion_cost_stats = std::nullopt, + const std::optional& normal_dispatch_rdma_recv_completion_sample_count_stats = std::nullopt, + const std::optional& normal_dispatch_rdma_recv_completion_token_count_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_count_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_timer_state = std::nullopt, diff --git a/csrc/kernels/api.cuh b/csrc/kernels/api.cuh index 848c32774..3a74b237a 100644 --- a/csrc/kernels/api.cuh +++ b/csrc/kernels/api.cuh @@ -213,6 +213,9 @@ void dispatch(void* recv_x, int64_t* normal_dispatch_final_completion_cost_stats, int64_t* normal_dispatch_final_completion_sample_count_stats, int64_t* normal_dispatch_final_completion_token_count_stats, + int64_t* normal_dispatch_rdma_recv_completion_cost_stats, + int64_t* normal_dispatch_rdma_recv_completion_sample_count_stats, + int64_t* normal_dispatch_rdma_recv_completion_token_count_stats, int rank, int num_ranks, bool is_cached_dispatch, diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index 4afec9f17..730541ef1 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -533,6 +533,9 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV int64_t* normal_dispatch_final_completion_cost_stats, int64_t* normal_dispatch_final_completion_sample_count_stats, int64_t* normal_dispatch_final_completion_token_count_stats, + int64_t* normal_dispatch_rdma_recv_completion_cost_stats, + int64_t* normal_dispatch_rdma_recv_completion_sample_count_stats, + int64_t* normal_dispatch_rdma_recv_completion_token_count_stats, int rank, int num_ranks) { enum class WarpRole { kRDMASender, kRDMASenderCoordinator, kRDMAAndNVLForwarder, kForwarderCoordinator, kNVLReceivers }; @@ -904,6 +907,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) { @@ -927,6 +933,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; } @@ -949,6 +957,31 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV } } __syncwarp(); + if (((normal_dispatch_rdma_recv_completion_cost_stats != nullptr and + normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) or + normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) and + dst_nvl_rank == 0 and lane_id < kNumRDMARanks and rdma_recv_completion_expected_count > 0 and + not rdma_recv_completion_recorded) { + const auto observed_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); + if (observed_rdma_channel_tail >= rdma_recv_completion_expected_count) { + const auto src_gateway_rank = lane_id * NUM_MAX_NVL_PEERS + nvl_rank; + const auto stats_idx = src_gateway_rank; + if (normal_dispatch_rdma_recv_completion_cost_stats != nullptr and + normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) { + atomicAdd(reinterpret_cast(normal_dispatch_rdma_recv_completion_cost_stats + stats_idx), + clock64() - rdma_recv_completion_start_time); + atomicAdd( + reinterpret_cast(normal_dispatch_rdma_recv_completion_sample_count_stats + stats_idx), + 1); + } + if (normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) { + atomicAdd( + reinterpret_cast(normal_dispatch_rdma_recv_completion_token_count_stats + stats_idx), + rdma_recv_completion_expected_count); + } + rdma_recv_completion_recorded = true; + } + } // Shift cached head send_nvl_head += src_rdma_channel_prefix * NUM_MAX_NVL_PEERS + dst_nvl_rank; @@ -962,6 +995,32 @@ __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 (((normal_dispatch_rdma_recv_completion_cost_stats != nullptr and + normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) or + normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) and + dst_nvl_rank == 0 and lane_id < kNumRDMARanks and rdma_recv_completion_expected_count > 0 and + not rdma_recv_completion_recorded) { + const auto observed_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); + if (observed_rdma_channel_tail >= rdma_recv_completion_expected_count) { + const auto src_gateway_rank = lane_id * NUM_MAX_NVL_PEERS + nvl_rank; + const auto stats_idx = src_gateway_rank; + if (normal_dispatch_rdma_recv_completion_cost_stats != nullptr and + normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) { + atomicAdd( + reinterpret_cast(normal_dispatch_rdma_recv_completion_cost_stats + stats_idx), + clock64() - rdma_recv_completion_start_time); + atomicAdd(reinterpret_cast( + normal_dispatch_rdma_recv_completion_sample_count_stats + stats_idx), + 1); + } + if (normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) { + atomicAdd(reinterpret_cast( + normal_dispatch_rdma_recv_completion_token_count_stats + stats_idx), + rdma_recv_completion_expected_count); + } + rdma_recv_completion_recorded = true; + } + } // Check destination queue emptiness, or wait a buffer to be released start_time = clock64(); while (true) { @@ -1063,6 +1122,31 @@ __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 (((normal_dispatch_rdma_recv_completion_cost_stats != nullptr and + normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) or + normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) and + dst_nvl_rank == 0 and lane_id < kNumRDMARanks and rdma_recv_completion_expected_count > 0 and + not rdma_recv_completion_recorded) { + const auto observed_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); + if (observed_rdma_channel_tail >= rdma_recv_completion_expected_count) { + const auto src_gateway_rank = lane_id * NUM_MAX_NVL_PEERS + nvl_rank; + const auto stats_idx = src_gateway_rank; + if (normal_dispatch_rdma_recv_completion_cost_stats != nullptr and + normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) { + atomicAdd(reinterpret_cast(normal_dispatch_rdma_recv_completion_cost_stats + stats_idx), + clock64() - rdma_recv_completion_start_time); + atomicAdd( + reinterpret_cast(normal_dispatch_rdma_recv_completion_sample_count_stats + stats_idx), + 1); + } + if (normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) { + atomicAdd( + reinterpret_cast(normal_dispatch_rdma_recv_completion_token_count_stats + stats_idx), + rdma_recv_completion_expected_count); + } + rdma_recv_completion_recorded = true; + } + } // Retired __syncwarp(); @@ -1321,6 +1405,9 @@ void dispatch(void* recv_x, int64_t* normal_dispatch_final_completion_cost_stats, int64_t* normal_dispatch_final_completion_sample_count_stats, int64_t* normal_dispatch_final_completion_token_count_stats, + int64_t* normal_dispatch_rdma_recv_completion_cost_stats, + int64_t* normal_dispatch_rdma_recv_completion_sample_count_stats, + int64_t* normal_dispatch_rdma_recv_completion_token_count_stats, int rank, int num_ranks, bool is_cached_dispatch, @@ -1382,6 +1469,9 @@ void dispatch(void* recv_x, normal_dispatch_final_completion_cost_stats, \ normal_dispatch_final_completion_sample_count_stats, \ normal_dispatch_final_completion_token_count_stats, \ + normal_dispatch_rdma_recv_completion_cost_stats, \ + normal_dispatch_rdma_recv_completion_sample_count_stats, \ + normal_dispatch_rdma_recv_completion_token_count_stats, \ rank, \ num_ranks); \ } \ diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 6d6d92493..413c6c6ce 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -579,6 +579,12 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te "normal_dispatch_final_completion_sample_count_stats"] normal_dispatch_final_completion_token_count_stats = normal_stats[ "normal_dispatch_final_completion_token_count_stats"] + normal_dispatch_rdma_recv_completion_cost_stats = normal_stats[ + "normal_dispatch_rdma_recv_completion_cost_stats"] + normal_dispatch_rdma_recv_completion_sample_count_stats = normal_stats[ + "normal_dispatch_rdma_recv_completion_sample_count_stats"] + normal_dispatch_rdma_recv_completion_token_count_stats = normal_stats[ + "normal_dispatch_rdma_recv_completion_token_count_stats"] if handle is not None: assert topk_idx is None and topk_weights is None is_token_in_rank, \ @@ -594,6 +600,9 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te normal_dispatch_final_completion_cost_stats, normal_dispatch_final_completion_sample_count_stats, normal_dispatch_final_completion_token_count_stats, + normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, normal_notify_dispatch_full_kernel_count_stats, normal_notify_dispatch_full_kernel_timer_state, @@ -615,6 +624,9 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te normal_dispatch_final_completion_cost_stats, normal_dispatch_final_completion_sample_count_stats, normal_dispatch_final_completion_token_count_stats, + normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, normal_notify_dispatch_full_kernel_count_stats, normal_notify_dispatch_full_kernel_timer_state, From 7ab8aa8d51a7c9957739e39d957469514dc17986 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 27 Jul 2026 15:27:36 +0800 Subject: [PATCH 06/16] feat(normal): add combine logical receive completion probes --- csrc/deep_ep.cpp | 29 +++++++++++++++- csrc/deep_ep.hpp | 5 ++- csrc/kernels/api.cuh | 3 ++ csrc/kernels/internode.cu | 72 +++++++++++++++++++++++++++++++++++++++ deep_ep/buffer.py | 11 +++++- 5 files changed, 117 insertions(+), 3 deletions(-) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index 8f0fc164a..3fc1a8181 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -1425,7 +1425,10 @@ std::tuple, std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats, const std::optional& normal_cached_notify_combine_full_kernel_count_stats, - const std::optional& normal_cached_notify_combine_full_kernel_timer_state) { + const std::optional& normal_cached_notify_combine_full_kernel_timer_state, + const std::optional& normal_combine_logical_recv_completion_cost_stats, + const std::optional& normal_combine_logical_recv_completion_sample_count_stats, + const std::optional& normal_combine_logical_recv_completion_token_count_stats) { #ifndef DISABLE_NVSHMEM const int num_channels = config.num_sms / 2; EP_HOST_ASSERT(config.num_sms % 2 == 0); @@ -1474,6 +1477,21 @@ std::tuple, std::optionalnumel() == 2 and normal_cached_notify_combine_full_kernel_timer_state->is_contiguous()); } + const bool enable_normal_combine_logical_recv_completion_stats = + normal_combine_logical_recv_completion_cost_stats.has_value(); + EP_HOST_ASSERT(normal_combine_logical_recv_completion_sample_count_stats.has_value() == + enable_normal_combine_logical_recv_completion_stats); + EP_HOST_ASSERT(normal_combine_logical_recv_completion_token_count_stats.has_value() == + enable_normal_combine_logical_recv_completion_stats); + for (const auto& stats : {normal_combine_logical_recv_completion_cost_stats, + normal_combine_logical_recv_completion_sample_count_stats, + normal_combine_logical_recv_completion_token_count_stats}) { + if (stats.has_value()) { + EP_HOST_ASSERT(stats->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(stats->dim() == 1 and stats->is_contiguous()); + EP_HOST_ASSERT(stats->numel() == num_ranks); + } + } // Allocate all tensors on comm stream if set // NOTES: do not allocate tensors upfront! @@ -1580,6 +1598,15 @@ std::tuple, std::optionaldata_ptr() + : nullptr, + normal_combine_logical_recv_completion_sample_count_stats.has_value() + ? normal_combine_logical_recv_completion_sample_count_stats->data_ptr() + : nullptr, + normal_combine_logical_recv_completion_token_count_stats.has_value() + ? normal_combine_logical_recv_completion_token_count_stats->data_ptr() + : nullptr, rank, num_ranks, comm_stream, diff --git a/csrc/deep_ep.hpp b/csrc/deep_ep.hpp index 8fecf02ea..a12b454d8 100644 --- a/csrc/deep_ep.hpp +++ b/csrc/deep_ep.hpp @@ -267,7 +267,10 @@ struct Buffer { bool allocate_on_comm_stream, const std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats = std::nullopt, const std::optional& normal_cached_notify_combine_full_kernel_count_stats = std::nullopt, - const std::optional& normal_cached_notify_combine_full_kernel_timer_state = std::nullopt); + const std::optional& normal_cached_notify_combine_full_kernel_timer_state = std::nullopt, + const std::optional& normal_combine_logical_recv_completion_cost_stats = std::nullopt, + const std::optional& normal_combine_logical_recv_completion_sample_count_stats = std::nullopt, + const std::optional& normal_combine_logical_recv_completion_token_count_stats = std::nullopt); 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 3a74b237a..58368ca15 100644 --- a/csrc/kernels/api.cuh +++ b/csrc/kernels/api.cuh @@ -273,6 +273,9 @@ void combine(cudaDataType_t type, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + int64_t* normal_combine_logical_recv_completion_cost_stats, + int64_t* normal_combine_logical_recv_completion_sample_count_stats, + int64_t* normal_combine_logical_recv_completion_token_count_stats, int rank, int num_ranks, cudaStream_t stream, diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index 730541ef1..f24f5b174 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -1929,6 +1929,9 @@ __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, + int64_t* normal_combine_logical_recv_completion_cost_stats, + int64_t* normal_combine_logical_recv_completion_sample_count_stats, + int64_t* normal_combine_logical_recv_completion_token_count_stats, int rank, int num_ranks) { enum class WarpRole { kNVLSender, kNVLAndRDMAForwarder, kRDMAReceiver, kCoordinator }; @@ -2415,6 +2418,43 @@ __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_combine_logical_recv_completion_cost_stats != nullptr and + normal_combine_logical_recv_completion_sample_count_stats != nullptr and + normal_combine_logical_recv_completion_token_count_stats != 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) { + 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) { + 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); + EP_DEVICE_ASSERT(expected_head >= 0); + #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; @@ -2425,6 +2465,32 @@ __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_combine_logical_recv_completion_cost_stats + src_global_rank), + completion_end_time - logical_completion_start_time); + atomicAdd(reinterpret_cast( + normal_combine_logical_recv_completion_sample_count_stats + src_global_rank), + 1); + atomicAdd(reinterpret_cast( + normal_combine_logical_recv_completion_token_count_stats + 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; @@ -2492,6 +2558,9 @@ void combine(cudaDataType_t type, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, + int64_t* normal_combine_logical_recv_completion_cost_stats, + int64_t* normal_combine_logical_recv_completion_sample_count_stats, + int64_t* normal_combine_logical_recv_completion_token_count_stats, int rank, int num_ranks, cudaStream_t stream, @@ -2543,6 +2612,9 @@ void combine(cudaDataType_t type, buffer_ptrs, \ num_max_nvl_chunked_send_tokens, \ num_max_nvl_chunked_recv_tokens, \ + normal_combine_logical_recv_completion_cost_stats, \ + normal_combine_logical_recv_completion_sample_count_stats, \ + normal_combine_logical_recv_completion_token_count_stats, \ rank, \ num_ranks); \ } \ diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 413c6c6ce..5180d5991 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -671,6 +671,12 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], normal_cached_notify_combine_full_kernel_duration_ns_stats, normal_cached_notify_combine_full_kernel_count_stats, self._normal_cached_notify_combine_full_kernel_timer_state) + normal_combine_logical_recv_completion_cost_stats = normal_stats[ + "normal_combine_logical_recv_completion_cost_stats"] + normal_combine_logical_recv_completion_sample_count_stats = normal_stats[ + "normal_combine_logical_recv_completion_sample_count_stats"] + normal_combine_logical_recv_completion_token_count_stats = normal_stats[ + "normal_combine_logical_recv_completion_token_count_stats"] # Launch the kernel combined_x, combined_topk_weights, event = self.runtime.internode_combine(x, topk_weights, bias_0, bias_1, src_meta, @@ -681,7 +687,10 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], None), async_finish, allocate_on_comm_stream, normal_cached_notify_combine_full_kernel_duration_ns_stats, normal_cached_notify_combine_full_kernel_count_stats, - normal_cached_notify_combine_full_kernel_timer_state) + normal_cached_notify_combine_full_kernel_timer_state, + normal_combine_logical_recv_completion_cost_stats, + normal_combine_logical_recv_completion_sample_count_stats, + normal_combine_logical_recv_completion_token_count_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: From 95727078476ecfaca923ab6b926e25ccf46dd4c9 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Tue, 28 Jul 2026 15:45:17 +0800 Subject: [PATCH 07/16] feat(deepxtrace): make normal diagnosis optional and configurable --- deep_ep/buffer.py | 123 +++++++++++++++++++++++++++++++++++++++++++++- setup.py | 7 ++- 2 files changed, 127 insertions(+), 3 deletions(-) diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 5180d5991..a3bba6948 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -1,4 +1,6 @@ +import importlib import os +import warnings import torch import torch.distributed as dist from typing import Callable, Dict, List, Tuple, Optional, Union @@ -39,6 +41,24 @@ def _validate_deepxtrace_normal_stats_schema(diagnose_module) -> None: f"{actual_schema!r}") +def _load_deepxtrace(enable_deepxtrace: bool, has_torch_group: bool): + if not isinstance(enable_deepxtrace, bool): + raise TypeError("`enable_deepxtrace` must be a bool") + if not enable_deepxtrace: + return None, "disabled by configuration" + if not has_torch_group: + return None, ( + "the current DeepXTrace integration requires a " + "torch.distributed process group") + + try: + diagnose_module = importlib.import_module("deepxtrace.diagnose") + _validate_deepxtrace_normal_stats_schema(diagnose_module) + except (ImportError, RuntimeError) as exc: + return None, str(exc) + return diagnose_module, None + + class Buffer: """ The core expert-parallel (EP) communication buffers for Mixture of Experts (MoE) model, which supports: @@ -62,14 +82,16 @@ def __init__(self, group: Optional[dist.ProcessGroup], num_nvl_bytes: int = 0, num_rdma_bytes: int = 0, - low_latency_mode: bool = False, + low_latency_mode: bool = True, num_qps_per_rank: int = 24, allow_nvlink_for_low_latency_mode: bool = True, allow_mnnvl: bool = False, use_fabric: bool = False, explicitly_destroy: bool = False, enable_shrink: bool = False, - comm: Optional["mpi4py.MPI.Comm"] = None) -> None: # noqa: F821 + comm: Optional["mpi4py.MPI.Comm"] = None, # noqa: F821 + enable_deepxtrace: bool = True, + enable_deepxtrace_async: bool = True) -> None: """ Initialize the communication buffer. @@ -91,8 +113,19 @@ def __init__(self, otherwise, the resources will be released by the destructor. Note: Releasing resources in the destructor may cause Python's exception handling process to hang. comm: the `mpi4py.MPI.Comm` communicator to use in case the group parameter is absent. + enable_deepxtrace: automatically enable DeepXTrace when every rank has a compatible installation. + If disabled, DeepEP does not import or initialize DeepXTrace. + enable_deepxtrace_async: whether to run DeepXTrace collection in + asynchronous mode. If enabled, the periodic background + collector uses ``DEEPEP_DIAGNOSE_INTERVAL``. If disabled, all + EP ranks must call + :meth:`diagnose_normal_sync` at the same logical step, and + ``DEEPEP_DIAGNOSE_SYNC_STEP`` controls the collection cadence. """ check_nvlink_connections(group) + if not isinstance(enable_deepxtrace_async, bool): + raise TypeError("`enable_deepxtrace_async` must be a bool") + self.enable_deepxtrace_async = enable_deepxtrace_async # Initialize the CPP runtime if group is not None: @@ -113,6 +146,48 @@ def all_gather_object(obj): return comm.allgather(obj) else: raise ValueError("Either 'group' or 'comm' must be provided.") + + diagnose_module, deepxtrace_error = _load_deepxtrace( + enable_deepxtrace, group is not None) + local_deepxtrace_status = ( + enable_deepxtrace, + diagnose_module is not None, + deepxtrace_error, + self.enable_deepxtrace_async, + ) + + deepxtrace_statuses = all_gather_object(local_deepxtrace_status) + all_deepxtrace_requested = all( + status[0] for status in deepxtrace_statuses) + all_deepxtrace_ready = all( + status[1] for status in deepxtrace_statuses) + if all_deepxtrace_requested: + configured_async_modes = { + status[3] for status in deepxtrace_statuses + } + if len(configured_async_modes) != 1: + raise RuntimeError( + "All EP ranks must use the same " + "`enable_deepxtrace_async`, but found " + f"{sorted(configured_async_modes)!r}") + self.deepxtrace_enabled = ( + all_deepxtrace_requested and all_deepxtrace_ready) + if (any(status[0] for status in deepxtrace_statuses) and + not self.deepxtrace_enabled and self.rank == 0): + unavailable = [ + f"rank {rank}: {status[2]}" + for rank, status in enumerate(deepxtrace_statuses) + if not status[1] + ] + preview = "; ".join(unavailable[:8]) + if len(unavailable) > 8: + preview += f"; ... and {len(unavailable) - 8} more ranks" + warnings.warn( + "DeepXTrace diagnosis is disabled for all ranks because " + f"the integration is not uniformly available: {preview}", + RuntimeWarning, + stacklevel=2) + self.diagnose = None self._normal_diagnose_stats = { name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA @@ -121,6 +196,7 @@ def all_gather_object(obj): self._normal_notify_dispatch_full_kernel_timer_state = None self._normal_cached_notify_dispatch_full_kernel_timer_state = None self._normal_cached_notify_combine_full_kernel_timer_state = None + self.num_nvl_bytes = num_nvl_bytes self.num_rdma_bytes = num_rdma_bytes self.low_latency_mode = low_latency_mode @@ -172,9 +248,26 @@ def all_gather_object(obj): nvshmem_unique_ids = all_gather_object(root_unique_id) root_unique_id = nvshmem_unique_ids[0 if low_latency_mode else self.runtime.get_root_rdma_rank(True)] + # Start DeepXtrace + # LL diagnosis is intentionally disabled by DeepEP for now. Normal + # diagnostics instrument the normal dispatch/combine APIs and remain + # available when the runtime uses the low-latency NVSHMEM topology. + if self.deepxtrace_enabled: + self._initialize_normal_notify_timer_states() + self.diagnose = diagnose_module.Diagnose( + group=group, + enable_ll_diagnose=False, + enable_normal_diagnose=True, + enable_async=self.enable_deepxtrace_async, + snapshot_stream=self.get_comm_stream()) + self._normal_diagnose_stats = self._get_normal_diagnose_stats() + # End DeepXtrace + # Make CPP runtime available self.runtime.sync(device_ids, ipc_handles, root_unique_id) assert self.runtime.is_available() + if self.diagnose is not None and self.enable_deepxtrace_async: + self.diagnose.start_async_diagnose() def _initialize_normal_notify_timer_states(self) -> None: """Create persistent per-launch notify timer scratch owned by DeepEP. @@ -213,6 +306,24 @@ def _get_normal_diagnose_stats( f"{len(tensors)} != {len(_REQUIRED_NORMAL_STATS_SCHEMA)}") return dict(zip(_REQUIRED_NORMAL_STATS_SCHEMA, tensors)) + def diagnose_normal_sync(self, diagnose_step: int = 0): + """Collect normal DeepXTrace statistics at a caller-owned step boundary. + + Every rank in the EP group must call this method at the same logical + location and with the same cadence. ``DEEPEP_DIAGNOSE_SYNC_STEP`` is + used when ``diagnose_step`` is zero; a nonzero value overrides it. + + Returns ``None`` when DeepXTrace is disabled. In async mode, collection + is owned by the background thread and this method raises. + """ + if self.diagnose is None: + return None + if self.enable_deepxtrace_async: + raise RuntimeError( + "diagnose_normal_sync() requires " + "`enable_deepxtrace_async=False`") + return self.diagnose.diagnose_normal_sync(diagnose_step) + @staticmethod def disable_ll_layered() -> bool: disable_ll_layered = False @@ -763,6 +874,10 @@ def low_latency_dispatch(self, x: torch.Tensor, topk_idx: torch.Tensor, hook: the receiving hook function (valid only if `return_recv_hook` is set). """ assert self.nvshmem_qp_depth >= (num_max_dispatch_tokens_per_rank + 1) * 2 + if (dispatch_wait_recv_cost_stats is None and + self.diagnose is not None): + dispatch_wait_recv_cost_stats = \ + self.diagnose.get_stats_ll_stats_tensor()[0] packed_recv_x, packed_recv_x_scales, packed_recv_count, packed_recv_src_info, packed_recv_layout_range, event, hook = \ self.runtime.low_latency_dispatch(x, topk_idx, cumulative_local_expert_recv_stats, @@ -830,6 +945,10 @@ def low_latency_combine(self, x: torch.Tensor, topk_idx: torch.Tensor, topk_weig """ src_info, layout_range, num_max_dispatch_tokens_per_rank, hidden, num_experts = handle assert self.nvshmem_qp_depth >= (num_max_dispatch_tokens_per_rank + 1) * 2 + if (combine_wait_recv_cost_stats is None and + self.diagnose is not None): + combine_wait_recv_cost_stats = \ + self.diagnose.get_stats_ll_stats_tensor()[1] combined_x, event, hook = self.runtime.low_latency_combine(x, topk_idx, topk_weights, src_info, layout_range, overlap, packed_recv_count, comp_signal, block_m, threshold, num_sms, combine_wait_recv_cost_stats, num_max_dispatch_tokens_per_rank, diff --git a/setup.py b/setup.py index f135107bd..58b3bb369 100644 --- a/setup.py +++ b/setup.py @@ -115,10 +115,15 @@ def get_nvshmem_host_lib_name(base_dir): revision = '+' + subprocess.check_output(cmd).decode('ascii').rstrip() except Exception as _: revision = '' - setuptools.setup(name='deep_ep', version='1.2.1' + revision, packages=setuptools.find_packages(include=['deep_ep']), + install_requires=[], + extras_require={ + 'deepxtrace': [ + 'deepxtrace>=0.2.0,<0.3.0', + ], + }, ext_modules=[ CUDAExtension(name='deep_ep_cpp', include_dirs=include_dirs, From fd9182a8d4028664158b39853cdf7e1d5cdc3a6d Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 10 Aug 2026 20:52:41 +0800 Subject: [PATCH 08/16] fix(buffer): preserve backward-compatible optional defaults --- deep_ep/buffer.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index a3bba6948..cddaf5bd8 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -82,7 +82,7 @@ def __init__(self, group: Optional[dist.ProcessGroup], num_nvl_bytes: int = 0, num_rdma_bytes: int = 0, - low_latency_mode: bool = True, + low_latency_mode: bool = False, num_qps_per_rank: int = 24, allow_nvlink_for_low_latency_mode: bool = True, allow_mnnvl: bool = False, @@ -90,7 +90,7 @@ def __init__(self, explicitly_destroy: bool = False, enable_shrink: bool = False, comm: Optional["mpi4py.MPI.Comm"] = None, # noqa: F821 - enable_deepxtrace: bool = True, + enable_deepxtrace: bool = False, enable_deepxtrace_async: bool = True) -> None: """ Initialize the communication buffer. @@ -113,8 +113,8 @@ def __init__(self, otherwise, the resources will be released by the destructor. Note: Releasing resources in the destructor may cause Python's exception handling process to hang. comm: the `mpi4py.MPI.Comm` communicator to use in case the group parameter is absent. - enable_deepxtrace: automatically enable DeepXTrace when every rank has a compatible installation. - If disabled, DeepEP does not import or initialize DeepXTrace. + enable_deepxtrace: whether to enable the optional DeepXTrace integration. DeepEP does not import or initialize + DeepXTrace unless this option is explicitly enabled on every EP rank. enable_deepxtrace_async: whether to run DeepXTrace collection in asynchronous mode. If enabled, the periodic background collector uses ``DEEPEP_DIAGNOSE_INTERVAL``. If disabled, all From 5caf04e9ae871ffdabf62f0dd3e087f9910b3ca0 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 10 Aug 2026 20:53:44 +0800 Subject: [PATCH 09/16] fix(normal): harden optional DeepXTrace integration --- deep_ep/buffer.py | 49 +++++++++++++++++++++++++++++++++++------------ 1 file changed, 37 insertions(+), 12 deletions(-) diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index cddaf5bd8..11c2ee30f 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -1,6 +1,7 @@ import importlib import os import warnings +import weakref import torch import torch.distributed as dist from typing import Callable, Dict, List, Tuple, Optional, Union @@ -54,8 +55,8 @@ def _load_deepxtrace(enable_deepxtrace: bool, has_torch_group: bool): try: diagnose_module = importlib.import_module("deepxtrace.diagnose") _validate_deepxtrace_normal_stats_schema(diagnose_module) - except (ImportError, RuntimeError) as exc: - return None, str(exc) + except Exception as exc: + return None, f"{type(exc).__name__}: {exc}" return diagnose_module, None @@ -189,6 +190,7 @@ def all_gather_object(obj): stacklevel=2) self.diagnose = None + self._deepxtrace_finalizer = None self._normal_diagnose_stats = { name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA } @@ -248,7 +250,11 @@ def all_gather_object(obj): nvshmem_unique_ids = all_gather_object(root_unique_id) root_unique_id = nvshmem_unique_ids[0 if low_latency_mode else self.runtime.get_root_rdma_rank(True)] - # Start DeepXtrace + # Make CPP runtime available + self.runtime.sync(device_ids, ipc_handles, root_unique_id) + assert self.runtime.is_available() + + # Start DeepXTrace # LL diagnosis is intentionally disabled by DeepEP for now. Normal # diagnostics instrument the normal dispatch/combine APIs and remain # available when the runtime uses the low-latency NVSHMEM topology. @@ -261,13 +267,11 @@ def all_gather_object(obj): enable_async=self.enable_deepxtrace_async, snapshot_stream=self.get_comm_stream()) self._normal_diagnose_stats = self._get_normal_diagnose_stats() - # End DeepXtrace - - # Make CPP runtime available - self.runtime.sync(device_ids, ipc_handles, root_unique_id) - assert self.runtime.is_available() - if self.diagnose is not None and self.enable_deepxtrace_async: - self.diagnose.start_async_diagnose() + if self.enable_deepxtrace_async: + self.diagnose.start_async_diagnose() + self._deepxtrace_finalizer = weakref.finalize( + self, self.diagnose.stop_async_diagnose) + # End DeepXTrace def _initialize_normal_notify_timer_states(self) -> None: """Create persistent per-launch notify timer scratch owned by DeepEP. @@ -310,8 +314,10 @@ def diagnose_normal_sync(self, diagnose_step: int = 0): """Collect normal DeepXTrace statistics at a caller-owned step boundary. Every rank in the EP group must call this method at the same logical - location and with the same cadence. ``DEEPEP_DIAGNOSE_SYNC_STEP`` is - used when ``diagnose_step`` is zero; a nonzero value overrides it. + location and with the same cadence. Mismatched calls can deadlock the + diagnostic collectives or assign samples to different windows. + ``DEEPEP_DIAGNOSE_SYNC_STEP`` is used when ``diagnose_step`` is zero; + a nonzero value overrides it. Returns ``None`` when DeepXTrace is disabled. In async mode, collection is owned by the background thread and this method raises. @@ -339,9 +345,28 @@ def destroy(self): assert self.explicitly_destroy, '`explicitly_destroy` flag must be set' + self._stop_deepxtrace() self.runtime.destroy() self.runtime = None + def _stop_deepxtrace(self) -> None: + """Stop the optional background collector before runtime teardown.""" + finalizer = getattr(self, "_deepxtrace_finalizer", None) + if finalizer is not None and finalizer.alive: + finalizer() + elif (getattr(self, "diagnose", None) is not None and + getattr(self, "enable_deepxtrace_async", False)): + self.diagnose.stop_async_diagnose() + self._deepxtrace_finalizer = None + self.diagnose = None + self._normal_diagnose_stats = { + name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA + } + self._normal_notify_full_kernel_timer_states = None + self._normal_notify_dispatch_full_kernel_timer_state = None + self._normal_cached_notify_dispatch_full_kernel_timer_state = None + self._normal_cached_notify_combine_full_kernel_timer_state = None + @staticmethod def is_sm90_compiled(): return deep_ep_cpp.is_sm90_compiled() From 77977ed0d7cd81b268d422648731032575e8ed28 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 10 Aug 2026 20:53:56 +0800 Subject: [PATCH 10/16] fix(normal): isolate diagnostics from low-latency APIs --- deep_ep/buffer.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 11c2ee30f..224e7a287 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -899,10 +899,6 @@ def low_latency_dispatch(self, x: torch.Tensor, topk_idx: torch.Tensor, hook: the receiving hook function (valid only if `return_recv_hook` is set). """ assert self.nvshmem_qp_depth >= (num_max_dispatch_tokens_per_rank + 1) * 2 - if (dispatch_wait_recv_cost_stats is None and - self.diagnose is not None): - dispatch_wait_recv_cost_stats = \ - self.diagnose.get_stats_ll_stats_tensor()[0] packed_recv_x, packed_recv_x_scales, packed_recv_count, packed_recv_src_info, packed_recv_layout_range, event, hook = \ self.runtime.low_latency_dispatch(x, topk_idx, cumulative_local_expert_recv_stats, @@ -970,10 +966,6 @@ def low_latency_combine(self, x: torch.Tensor, topk_idx: torch.Tensor, topk_weig """ src_info, layout_range, num_max_dispatch_tokens_per_rank, hidden, num_experts = handle assert self.nvshmem_qp_depth >= (num_max_dispatch_tokens_per_rank + 1) * 2 - if (combine_wait_recv_cost_stats is None and - self.diagnose is not None): - combine_wait_recv_cost_stats = \ - self.diagnose.get_stats_ll_stats_tensor()[1] combined_x, event, hook = self.runtime.low_latency_combine(x, topk_idx, topk_weights, src_info, layout_range, overlap, packed_recv_count, comp_signal, block_m, threshold, num_sms, combine_wait_recv_cost_stats, num_max_dispatch_tokens_per_rank, From fc64467798e0c77c2145c607c27d7a2db0412fb9 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 10 Aug 2026 20:56:53 +0800 Subject: [PATCH 11/16] fix(normal): harden probe contracts and collection --- csrc/deep_ep.cpp | 98 ++++++++++++++++++------- csrc/kernels/api.cuh | 4 ++ csrc/kernels/internode.cu | 148 ++++++++++++++++++-------------------- 3 files changed, 148 insertions(+), 102 deletions(-) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index 3fc1a8181..26ca8eae0 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -1016,20 +1016,6 @@ 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); } - auto check_normal_dispatch_stat_pair = [=](const std::optional& cost_stats, - const std::optional& count_stats) { - if (cost_stats.has_value()) { - EP_HOST_ASSERT(cost_stats->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(cost_stats->dim() == 1 and cost_stats->is_contiguous()); - EP_HOST_ASSERT(cost_stats->size(0) == num_ranks); - EP_HOST_ASSERT(count_stats.has_value()); - EP_HOST_ASSERT(count_stats->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(count_stats->dim() == 1 and count_stats->is_contiguous()); - EP_HOST_ASSERT(count_stats->size(0) == num_ranks); - } else { - EP_HOST_ASSERT(not count_stats.has_value()); - } - }; auto check_normal_dispatch_stat_tensor = [=](const std::optional& stats) { if (stats.has_value()) { EP_HOST_ASSERT(stats->scalar_type() == torch::kInt64); @@ -1037,15 +1023,22 @@ Buffer::internode_dispatch(const torch::Tensor& x, EP_HOST_ASSERT(stats->size(0) == num_ranks); } }; - const bool enable_normal_dispatch_final_completion_stats = normal_dispatch_final_completion_cost_stats.has_value(); - EP_HOST_ASSERT(normal_dispatch_final_completion_sample_count_stats.has_value() == enable_normal_dispatch_final_completion_stats); - EP_HOST_ASSERT(normal_dispatch_final_completion_token_count_stats.has_value() == enable_normal_dispatch_final_completion_stats); - check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_cost_stats); - check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_sample_count_stats); - check_normal_dispatch_stat_tensor(normal_dispatch_final_completion_token_count_stats); - check_normal_dispatch_stat_pair(normal_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats); - check_normal_dispatch_stat_tensor(normal_dispatch_rdma_recv_completion_token_count_stats); + auto check_normal_dispatch_stat_triplet = [&](const std::optional& cost_stats, + const std::optional& sample_count_stats, + const std::optional& token_count_stats) { + const bool enabled = cost_stats.has_value(); + EP_HOST_ASSERT(sample_count_stats.has_value() == enabled); + EP_HOST_ASSERT(token_count_stats.has_value() == enabled); + check_normal_dispatch_stat_tensor(cost_stats); + check_normal_dispatch_stat_tensor(sample_count_stats); + check_normal_dispatch_stat_tensor(token_count_stats); + }; + check_normal_dispatch_stat_triplet(normal_dispatch_final_completion_cost_stats, + normal_dispatch_final_completion_sample_count_stats, + normal_dispatch_final_completion_token_count_stats); + check_normal_dispatch_stat_triplet(normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_stats); auto check_normal_notify_stats = [](const std::optional& duration_ns_stats, const std::optional& count_stats, const std::optional& timer_state) { @@ -2054,8 +2047,63 @@ 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_final_completion_cost_stats") = py::none(), + py::arg("normal_dispatch_final_completion_sample_count_stats") = py::none(), + py::arg("normal_dispatch_final_completion_token_count_stats") = py::none(), + py::arg("normal_dispatch_rdma_recv_completion_cost_stats") = py::none(), + py::arg("normal_dispatch_rdma_recv_completion_sample_count_stats") = py::none(), + py::arg("normal_dispatch_rdma_recv_completion_token_count_stats") = py::none(), + py::arg("normal_notify_dispatch_full_kernel_duration_ns_stats") = py::none(), + py::arg("normal_notify_dispatch_full_kernel_count_stats") = py::none(), + py::arg("normal_notify_dispatch_full_kernel_timer_state") = py::none(), + py::arg("normal_cached_notify_dispatch_full_kernel_duration_ns_stats") = py::none(), + py::arg("normal_cached_notify_dispatch_full_kernel_count_stats") = py::none(), + py::arg("normal_cached_notify_dispatch_full_kernel_timer_state") = 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_cached_notify_combine_full_kernel_duration_ns_stats") = py::none(), + py::arg("normal_cached_notify_combine_full_kernel_count_stats") = py::none(), + py::arg("normal_cached_notify_combine_full_kernel_timer_state") = py::none(), + py::arg("normal_combine_logical_recv_completion_cost_stats") = py::none(), + py::arg("normal_combine_logical_recv_completion_sample_count_stats") = py::none(), + py::arg("normal_combine_logical_recv_completion_token_count_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/kernels/api.cuh b/csrc/kernels/api.cuh index 58368ca15..5529f9986 100644 --- a/csrc/kernels/api.cuh +++ b/csrc/kernels/api.cuh @@ -210,6 +210,8 @@ 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. int64_t* normal_dispatch_final_completion_cost_stats, int64_t* normal_dispatch_final_completion_sample_count_stats, int64_t* normal_dispatch_final_completion_token_count_stats, @@ -273,6 +275,8 @@ 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. int64_t* normal_combine_logical_recv_completion_cost_stats, int64_t* normal_combine_logical_recv_completion_sample_count_stats, int64_t* normal_combine_logical_recv_completion_token_count_stats, diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index f24f5b174..401f2de54 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -50,6 +50,35 @@ __device__ __forceinline__ void notify_full_kernel_timer_end(int64_t* timer_stat } } +__device__ __forceinline__ void try_record_dispatch_rdma_recv_completion( + int64_t* cost_stats, + int64_t* sample_count_stats, + int64_t* token_count_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 (cost_stats == 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(cost_stats + src_gateway_proxy_rank), clock64() - start_time); + atomicAdd(reinterpret_cast(sample_count_stats + src_gateway_proxy_rank), 1ull); + atomicAdd(reinterpret_cast(token_count_stats + src_gateway_proxy_rank), + static_cast(expected_count)); + recorded = true; +} + struct SourceMeta { int src_rdma_rank, is_token_in_nvl_rank_bits; @@ -957,31 +986,16 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV } } __syncwarp(); - if (((normal_dispatch_rdma_recv_completion_cost_stats != nullptr and - normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) or - normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) and - dst_nvl_rank == 0 and lane_id < kNumRDMARanks and rdma_recv_completion_expected_count > 0 and - not rdma_recv_completion_recorded) { - const auto observed_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); - if (observed_rdma_channel_tail >= rdma_recv_completion_expected_count) { - const auto src_gateway_rank = lane_id * NUM_MAX_NVL_PEERS + nvl_rank; - const auto stats_idx = src_gateway_rank; - if (normal_dispatch_rdma_recv_completion_cost_stats != nullptr and - normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) { - atomicAdd(reinterpret_cast(normal_dispatch_rdma_recv_completion_cost_stats + stats_idx), - clock64() - rdma_recv_completion_start_time); - atomicAdd( - reinterpret_cast(normal_dispatch_rdma_recv_completion_sample_count_stats + stats_idx), - 1); - } - if (normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) { - atomicAdd( - reinterpret_cast(normal_dispatch_rdma_recv_completion_token_count_stats + stats_idx), - rdma_recv_completion_expected_count); - } - rdma_recv_completion_recorded = true; - } - } + if (dst_nvl_rank == 0 and lane_id < kNumRDMARanks) + try_record_dispatch_rdma_recv_completion(normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_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; @@ -995,32 +1009,16 @@ __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 (((normal_dispatch_rdma_recv_completion_cost_stats != nullptr and - normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) or - normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) and - dst_nvl_rank == 0 and lane_id < kNumRDMARanks and rdma_recv_completion_expected_count > 0 and - not rdma_recv_completion_recorded) { - const auto observed_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); - if (observed_rdma_channel_tail >= rdma_recv_completion_expected_count) { - const auto src_gateway_rank = lane_id * NUM_MAX_NVL_PEERS + nvl_rank; - const auto stats_idx = src_gateway_rank; - if (normal_dispatch_rdma_recv_completion_cost_stats != nullptr and - normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) { - atomicAdd( - reinterpret_cast(normal_dispatch_rdma_recv_completion_cost_stats + stats_idx), - clock64() - rdma_recv_completion_start_time); - atomicAdd(reinterpret_cast( - normal_dispatch_rdma_recv_completion_sample_count_stats + stats_idx), - 1); - } - if (normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) { - atomicAdd(reinterpret_cast( - normal_dispatch_rdma_recv_completion_token_count_stats + stats_idx), - rdma_recv_completion_expected_count); - } - rdma_recv_completion_recorded = true; - } - } + if (dst_nvl_rank == 0 and lane_id < kNumRDMARanks) + try_record_dispatch_rdma_recv_completion(normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_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) { @@ -1122,31 +1120,16 @@ __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 (((normal_dispatch_rdma_recv_completion_cost_stats != nullptr and - normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) or - normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) and - dst_nvl_rank == 0 and lane_id < kNumRDMARanks and rdma_recv_completion_expected_count > 0 and - not rdma_recv_completion_recorded) { - const auto observed_rdma_channel_tail = static_cast(ld_acquire_sys_global(rdma_channel_tail.buffer(lane_id))); - if (observed_rdma_channel_tail >= rdma_recv_completion_expected_count) { - const auto src_gateway_rank = lane_id * NUM_MAX_NVL_PEERS + nvl_rank; - const auto stats_idx = src_gateway_rank; - if (normal_dispatch_rdma_recv_completion_cost_stats != nullptr and - normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) { - atomicAdd(reinterpret_cast(normal_dispatch_rdma_recv_completion_cost_stats + stats_idx), - clock64() - rdma_recv_completion_start_time); - atomicAdd( - reinterpret_cast(normal_dispatch_rdma_recv_completion_sample_count_stats + stats_idx), - 1); - } - if (normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) { - atomicAdd( - reinterpret_cast(normal_dispatch_rdma_recv_completion_token_count_stats + stats_idx), - rdma_recv_completion_expected_count); - } - rdma_recv_completion_recorded = true; - } - } + if (dst_nvl_rank == 0 and lane_id < kNumRDMARanks) + try_record_dispatch_rdma_recv_completion(normal_dispatch_rdma_recv_completion_cost_stats, + normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_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(); @@ -1423,6 +1406,12 @@ void dispatch(void* recv_x, const bool enable_normal_dispatch_final_completion_stats = normal_dispatch_final_completion_cost_stats != nullptr; EP_HOST_ASSERT((normal_dispatch_final_completion_sample_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); EP_HOST_ASSERT((normal_dispatch_final_completion_token_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); + const bool enable_normal_dispatch_rdma_recv_completion_stats = + normal_dispatch_rdma_recv_completion_cost_stats != nullptr; + EP_HOST_ASSERT((normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) == + enable_normal_dispatch_rdma_recv_completion_stats); + EP_HOST_ASSERT((normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) == + enable_normal_dispatch_rdma_recv_completion_stats); #define DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ { \ @@ -2432,15 +2421,20 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co // 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); - EP_DEVICE_ASSERT(expected_head >= 0); + // 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; From 7d672b0eeaf4263b187e9f8f37cf9ebb7c4b1414 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 10 Aug 2026 21:01:22 +0800 Subject: [PATCH 12/16] test(normal): cover optional DeepXTrace integration --- tests/test_deepxtrace_integration.py | 251 +++++++++++++++++++++++++++ 1 file changed, 251 insertions(+) create mode 100644 tests/test_deepxtrace_integration.py diff --git a/tests/test_deepxtrace_integration.py b/tests/test_deepxtrace_integration.py new file mode 100644 index 000000000..b0e8ea939 --- /dev/null +++ b/tests/test_deepxtrace_integration.py @@ -0,0 +1,251 @@ +import gc +import importlib.util +import inspect +import sys +import types +import unittest +import weakref +from pathlib import Path +from unittest.mock import MagicMock, patch + + +def _load_buffer_module(): + """Load buffer.py with host-only stubs for torch and the CUDA extension.""" + torch_module = types.ModuleType("torch") + torch_module.Tensor = type("Tensor", (), {}) + torch_module.Stream = type("Stream", (), {}) + torch_module.Size = tuple + torch_module.dtype = type("dtype", (), {}) + torch_module.cuda = types.SimpleNamespace() + + dist_module = types.ModuleType("torch.distributed") + dist_module.ProcessGroup = type("ProcessGroup", (), {}) + dist_module.all_gather_object = lambda outputs, value, group: None + torch_module.distributed = dist_module + + cpp_module = types.ModuleType("deep_ep_cpp") + cpp_module.Buffer = type("Buffer", (), {}) + cpp_module.Config = type("Config", (), {}) + cpp_module.EventHandle = type("EventHandle", (), {}) + + package_module = types.ModuleType("deep_ep") + package_module.__path__ = [] + utils_module = types.ModuleType("deep_ep.utils") + + class EventOverlap: + + def __init__(self, *args): + self.args = args + + utils_module.EventOverlap = EventOverlap + utils_module.check_nvlink_connections = lambda group: None + + module_name = "deep_ep._buffer_host_test" + module_path = Path(__file__).resolve().parents[1] / "deep_ep" / "buffer.py" + spec = importlib.util.spec_from_file_location(module_name, module_path) + module = importlib.util.module_from_spec(spec) + stubs = { + "torch": torch_module, + "torch.distributed": dist_module, + "deep_ep_cpp": cpp_module, + "deep_ep": package_module, + "deep_ep.utils": utils_module, + module_name: module, + } + with patch.dict(sys.modules, stubs): + spec.loader.exec_module(module) + return module + + +buffer_module = _load_buffer_module() + + +class _FakeGroup: + + def rank(self): + return 0 + + def size(self): + return 2 + + +class _FakeRuntime: + + def __init__(self, *args): + self.args = args + self.synced = False + + def get_local_device_id(self): + return 0 + + def get_local_ipc_handle(self): + return bytearray() + + def get_num_rdma_ranks(self): + return 1 + + def sync(self, *args): + self.synced = True + + def is_available(self): + return self.synced + + +class TestDeepXTraceIntegration(unittest.TestCase): + + @staticmethod + def compatible_module(): + return types.SimpleNamespace(NORMAL_STATS_SCHEMA=buffer_module._REQUIRED_NORMAL_STATS_SCHEMA) + + def test_optional_defaults_are_backward_compatible(self): + parameters = inspect.signature(buffer_module.Buffer.__init__).parameters + self.assertIs(parameters["low_latency_mode"].default, False) + self.assertIs(parameters["enable_deepxtrace"].default, False) + + def test_optional_dependency_loading_is_fail_closed(self): + with patch.object(buffer_module.importlib, "import_module") as import_module: + diagnose_module, error = buffer_module._load_deepxtrace(False, True) + import_module.assert_not_called() + self.assertIsNone(diagnose_module) + self.assertEqual(error, "disabled by configuration") + + with patch.object(buffer_module.importlib, "import_module", return_value=self.compatible_module()): + diagnose_module, error = buffer_module._load_deepxtrace(True, True) + self.assertIsNotNone(diagnose_module) + self.assertIsNone(error) + + for exception in ( + ModuleNotFoundError("No module named 'deepxtrace'"), + OSError("broken optional dependency"), + AttributeError("incomplete installation"), + ): + with self.subTest(exception=exception): + with patch.object(buffer_module.importlib, "import_module", side_effect=exception): + diagnose_module, error = \ + buffer_module._load_deepxtrace(True, True) + self.assertIsNone(diagnose_module) + self.assertIn(type(exception).__name__, error) + + def test_schema_mismatch_is_reported_as_unavailable(self): + incompatible_module = types.SimpleNamespace(NORMAL_STATS_SCHEMA=("legacy", )) + with self.assertRaisesRegex(RuntimeError, "deepxtrace>=0.2.0"): + buffer_module._validate_deepxtrace_normal_stats_schema(incompatible_module) + + with patch.object(buffer_module.importlib, "import_module", return_value=incompatible_module): + diagnose_module, error = buffer_module._load_deepxtrace(True, True) + self.assertIsNone(diagnose_module) + self.assertIn("RuntimeError", error) + + def test_rank_wide_mixed_enablement_disables_diagnosis(self): + + def all_gather_object(outputs, value, group): + if isinstance(value, tuple) and len(value) == 4: + outputs[:] = [ + value, + (False, False, "disabled by configuration", value[3]), + ] + else: + outputs[:] = [value, value] + + with patch.object(buffer_module.deep_ep_cpp, "Buffer", _FakeRuntime), \ + patch.object(buffer_module.dist, "all_gather_object", + side_effect=all_gather_object), \ + patch.object(buffer_module, "_load_deepxtrace", + return_value=(self.compatible_module(), None)), \ + self.assertWarnsRegex(RuntimeWarning, "disabled for all ranks"): + buffer = buffer_module.Buffer(_FakeGroup(), enable_deepxtrace=True) + self.assertFalse(buffer.deepxtrace_enabled) + self.assertTrue(buffer.runtime.is_available()) + + def test_rank_wide_async_mode_mismatch_raises(self): + + def all_gather_object(outputs, value, group): + outputs[:] = [value, (True, True, None, not value[3])] + + with patch.object(buffer_module.dist, "all_gather_object", + side_effect=all_gather_object), \ + patch.object(buffer_module, "_load_deepxtrace", + return_value=(self.compatible_module(), None)), \ + self.assertRaisesRegex(RuntimeError, + "enable_deepxtrace_async"): + buffer_module.Buffer(_FakeGroup(), enable_deepxtrace=True) + + def test_sync_collection_contract(self): + buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) + buffer.diagnose = None + buffer.enable_deepxtrace_async = False + self.assertIsNone(buffer.diagnose_normal_sync()) + + buffer.diagnose = MagicMock() + buffer.enable_deepxtrace_async = True + with self.assertRaisesRegex(RuntimeError, "enable_deepxtrace_async=False"): + buffer.diagnose_normal_sync() + + expected = [{"probe": "normal", "status": "ok"}] + buffer.enable_deepxtrace_async = False + buffer.diagnose.diagnose_normal_sync.return_value = expected + self.assertIs(buffer.diagnose_normal_sync(17), expected) + buffer.diagnose.diagnose_normal_sync.assert_called_once_with(17) + + def test_destroy_stops_async_diagnosis_before_runtime(self): + order = [] + buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) + buffer.explicitly_destroy = True + buffer.enable_deepxtrace_async = True + buffer.diagnose = MagicMock() + buffer.diagnose.stop_async_diagnose.side_effect = \ + lambda: order.append("diagnose") + buffer.runtime = MagicMock() + buffer.runtime.destroy.side_effect = lambda: order.append("runtime") + buffer._deepxtrace_finalizer = weakref.finalize(buffer, buffer.diagnose.stop_async_diagnose) + buffer._normal_diagnose_stats = {"probe": object()} + buffer._normal_notify_full_kernel_timer_states = object() + buffer._normal_notify_dispatch_full_kernel_timer_state = object() + buffer._normal_cached_notify_dispatch_full_kernel_timer_state = object() + buffer._normal_cached_notify_combine_full_kernel_timer_state = object() + + buffer.destroy() + + self.assertEqual(order, ["diagnose", "runtime"]) + self.assertIsNone(buffer.runtime) + self.assertIsNone(buffer.diagnose) + self.assertFalse(buffer._deepxtrace_finalizer) + self.assertTrue(all(value is None for value in buffer._normal_diagnose_stats.values())) + + def test_finalizer_stops_async_diagnosis_on_implicit_cleanup(self): + stop_async_diagnose = MagicMock() + buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) + buffer._deepxtrace_finalizer = weakref.finalize(buffer, stop_async_diagnose) + buffer_ref = weakref.ref(buffer) + + del buffer + gc.collect() + + self.assertIsNone(buffer_ref()) + stop_async_diagnose.assert_called_once_with() + + def test_low_latency_stats_remain_caller_owned(self): + buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) + buffer.nvshmem_qp_depth = 1024 + buffer.diagnose = MagicMock() + buffer.runtime = MagicMock() + buffer.runtime.low_latency_dispatch.return_value = (object(), object(), object(), object(), object(), object(), object()) + x = MagicMock() + x.size.return_value = 64 + + buffer.low_latency_dispatch(x, MagicMock(), 1, 1) + + dispatch_args = buffer.runtime.low_latency_dispatch.call_args.args + self.assertIsNone(dispatch_args[3]) + + buffer.runtime.low_latency_combine.return_value = (object(), object(), object()) + handle = (object(), object(), 1, 64, 1) + buffer.low_latency_combine(MagicMock(), MagicMock(), MagicMock(), handle) + + combine_args = buffer.runtime.low_latency_combine.call_args.args + self.assertIsNone(combine_args[11]) + buffer.diagnose.get_stats_ll_stats_tensor.assert_not_called() + + +if __name__ == "__main__": + unittest.main() From 21a249929ace89a3fdb3293bdc98e5bd53949a2e Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Mon, 10 Aug 2026 21:01:51 +0800 Subject: [PATCH 13/16] style(normal): apply repository formatters --- csrc/deep_ep.cpp | 251 ++++++++++++++++++-------------------- csrc/kernels/internode.cu | 60 ++++----- deep_ep/buffer.py | 208 ++++++++++++------------------- 3 files changed, 227 insertions(+), 292 deletions(-) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index 26ca8eae0..81771ff82 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -1047,8 +1047,7 @@ Buffer::internode_dispatch(const torch::Tensor& x, if (enabled) { EP_HOST_ASSERT(duration_ns_stats->scalar_type() == torch::kInt64 and count_stats->scalar_type() == torch::kInt64 and timer_state->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(duration_ns_stats->dim() == 1 and duration_ns_stats->numel() == 1 and - duration_ns_stats->is_contiguous()); + EP_HOST_ASSERT(duration_ns_stats->dim() == 1 and duration_ns_stats->numel() == 1 and duration_ns_stats->is_contiguous()); EP_HOST_ASSERT(count_stats->dim() == 1 and count_stats->numel() == 1 and count_stats->is_contiguous()); EP_HOST_ASSERT(timer_state->dim() == 1 and timer_state->numel() == 2 and timer_state->is_contiguous()); } @@ -1169,45 +1168,44 @@ Buffer::internode_dispatch(const torch::Tensor& x, *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_worst_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, - normal_notify_dispatch_full_kernel_duration_ns_stats.has_value() - ? normal_notify_dispatch_full_kernel_duration_ns_stats->data_ptr() - : nullptr, - normal_notify_dispatch_full_kernel_count_stats.has_value() - ? normal_notify_dispatch_full_kernel_count_stats->data_ptr() - : nullptr, - normal_notify_dispatch_full_kernel_timer_state.has_value() - ? normal_notify_dispatch_full_kernel_timer_state->data_ptr() - : nullptr); + 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_worst_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, + normal_notify_dispatch_full_kernel_duration_ns_stats.has_value() + ? normal_notify_dispatch_full_kernel_duration_ns_stats->data_ptr() + : nullptr, + normal_notify_dispatch_full_kernel_count_stats.has_value() ? normal_notify_dispatch_full_kernel_count_stats->data_ptr() + : nullptr, + normal_notify_dispatch_full_kernel_timer_state.has_value() ? normal_notify_dispatch_full_kernel_timer_state->data_ptr() + : nullptr); // Synchronize total received tokens and tokens per expert if (num_worst_tokens > 0) { @@ -1276,62 +1274,61 @@ Buffer::internode_dispatch(const torch::Tensor& x, // 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, - num_worst_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, - normal_dispatch_final_completion_cost_stats.has_value() - ? normal_dispatch_final_completion_cost_stats->data_ptr() - : nullptr, - normal_dispatch_final_completion_sample_count_stats.has_value() - ? normal_dispatch_final_completion_sample_count_stats->data_ptr() - : nullptr, - normal_dispatch_final_completion_token_count_stats.has_value() - ? normal_dispatch_final_completion_token_count_stats->data_ptr() - : nullptr, - normal_dispatch_rdma_recv_completion_cost_stats.has_value() - ? normal_dispatch_rdma_recv_completion_cost_stats->data_ptr() - : nullptr, - normal_dispatch_rdma_recv_completion_sample_count_stats.has_value() - ? normal_dispatch_rdma_recv_completion_sample_count_stats->data_ptr() - : nullptr, - normal_dispatch_rdma_recv_completion_token_count_stats.has_value() - ? normal_dispatch_rdma_recv_completion_token_count_stats->data_ptr() - : nullptr, - rank, - num_ranks, - cached_mode, - comm_stream, - num_channels, - low_latency_mode); + 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, + num_worst_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, + normal_dispatch_final_completion_cost_stats.has_value() ? normal_dispatch_final_completion_cost_stats->data_ptr() + : nullptr, + normal_dispatch_final_completion_sample_count_stats.has_value() + ? normal_dispatch_final_completion_sample_count_stats->data_ptr() + : nullptr, + normal_dispatch_final_completion_token_count_stats.has_value() + ? normal_dispatch_final_completion_token_count_stats->data_ptr() + : nullptr, + normal_dispatch_rdma_recv_completion_cost_stats.has_value() ? normal_dispatch_rdma_recv_completion_cost_stats->data_ptr() + : nullptr, + normal_dispatch_rdma_recv_completion_sample_count_stats.has_value() + ? normal_dispatch_rdma_recv_completion_sample_count_stats->data_ptr() + : nullptr, + normal_dispatch_rdma_recv_completion_token_count_stats.has_value() + ? normal_dispatch_rdma_recv_completion_token_count_stats->data_ptr() + : nullptr, + rank, + num_ranks, + cached_mode, + comm_stream, + num_channels, + low_latency_mode); // Wait streams std::optional event; @@ -1470,8 +1467,7 @@ std::tuple, std::optionalnumel() == 2 and normal_cached_notify_combine_full_kernel_timer_state->is_contiguous()); } - const bool enable_normal_combine_logical_recv_completion_stats = - normal_combine_logical_recv_completion_cost_stats.has_value(); + const bool enable_normal_combine_logical_recv_completion_stats = normal_combine_logical_recv_completion_cost_stats.has_value(); EP_HOST_ASSERT(normal_combine_logical_recv_completion_sample_count_stats.has_value() == enable_normal_combine_logical_recv_completion_stats); EP_HOST_ASSERT(normal_combine_logical_recv_completion_token_count_stats.has_value() == @@ -1521,37 +1517,32 @@ std::tuple, std::optional(), - 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, - normal_cached_notify_combine_enabled - ? normal_cached_notify_combine_full_kernel_duration_ns_stats->data_ptr() - : nullptr, - normal_cached_notify_combine_enabled - ? normal_cached_notify_combine_full_kernel_count_stats->data_ptr() - : nullptr, - normal_cached_notify_combine_enabled - ? normal_cached_notify_combine_full_kernel_timer_state->data_ptr() - : nullptr); + 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, + normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_duration_ns_stats->data_ptr() : nullptr, + normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_count_stats->data_ptr() : nullptr, + normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_timer_state->data_ptr() : nullptr); // Assign bias pointers auto bias_opts = std::vector>({bias_0, bias_1}); diff --git a/csrc/kernels/internode.cu b/csrc/kernels/internode.cu index 401f2de54..5b8ef9562 100644 --- a/csrc/kernels/internode.cu +++ b/csrc/kernels/internode.cu @@ -31,9 +31,7 @@ __device__ __forceinline__ void notify_full_kernel_timer_begin(int64_t* timer_st __syncthreads(); } -__device__ __forceinline__ void notify_full_kernel_timer_end(int64_t* timer_state, - int64_t* duration_ns_stats, - int64_t* count_stats) { +__device__ __forceinline__ void notify_full_kernel_timer_end(int64_t* timer_state, int64_t* duration_ns_stats, int64_t* count_stats) { if (timer_state == nullptr) return; @@ -50,16 +48,15 @@ __device__ __forceinline__ void notify_full_kernel_timer_end(int64_t* timer_stat } } -__device__ __forceinline__ void try_record_dispatch_rdma_recv_completion( - int64_t* cost_stats, - int64_t* sample_count_stats, - int64_t* token_count_stats, - const uint64_t* rdma_channel_tail, - int src_rdma_rank, - int gateway_nvl_rank, - int expected_count, - uint64_t start_time, - bool& recorded) { +__device__ __forceinline__ void try_record_dispatch_rdma_recv_completion(int64_t* cost_stats, + int64_t* sample_count_stats, + int64_t* token_count_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 (cost_stats == nullptr or expected_count <= 0 or recorded) return; @@ -1256,8 +1253,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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) { + 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; @@ -1406,8 +1402,7 @@ void dispatch(void* recv_x, const bool enable_normal_dispatch_final_completion_stats = normal_dispatch_final_completion_cost_stats != nullptr; EP_HOST_ASSERT((normal_dispatch_final_completion_sample_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); EP_HOST_ASSERT((normal_dispatch_final_completion_token_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); - const bool enable_normal_dispatch_rdma_recv_completion_stats = - normal_dispatch_rdma_recv_completion_cost_stats != nullptr; + const bool enable_normal_dispatch_rdma_recv_completion_stats = normal_dispatch_rdma_recv_completion_cost_stats != nullptr; EP_HOST_ASSERT((normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) == enable_normal_dispatch_rdma_recv_completion_stats); EP_HOST_ASSERT((normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) == @@ -1458,9 +1453,9 @@ void dispatch(void* recv_x, normal_dispatch_final_completion_cost_stats, \ normal_dispatch_final_completion_sample_count_stats, \ normal_dispatch_final_completion_token_count_stats, \ - normal_dispatch_rdma_recv_completion_cost_stats, \ - normal_dispatch_rdma_recv_completion_sample_count_stats, \ - normal_dispatch_rdma_recv_completion_token_count_stats, \ + normal_dispatch_rdma_recv_completion_cost_stats, \ + normal_dispatch_rdma_recv_completion_sample_count_stats, \ + normal_dispatch_rdma_recv_completion_token_count_stats, \ rank, \ num_ranks); \ } \ @@ -2408,8 +2403,8 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co } else { // Coordinator const bool enable_logical_completion = normal_combine_logical_recv_completion_cost_stats != nullptr and - normal_combine_logical_recv_completion_sample_count_stats != nullptr and - normal_combine_logical_recv_completion_token_count_stats != nullptr; + normal_combine_logical_recv_completion_sample_count_stats != nullptr and + normal_combine_logical_recv_completion_token_count_stats != 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; @@ -2427,8 +2422,8 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co 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)); + 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); @@ -2440,8 +2435,7 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co 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_last_expected_head[src_nvl_rank] = max(logical_last_expected_head[src_nvl_rank], expected_head); logical_expected_mask |= src_bit; } } @@ -2471,14 +2465,14 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co 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_combine_logical_recv_completion_cost_stats + src_global_rank), - completion_end_time - logical_completion_start_time); - atomicAdd(reinterpret_cast( - normal_combine_logical_recv_completion_sample_count_stats + src_global_rank), + atomicAdd( + reinterpret_cast(normal_combine_logical_recv_completion_cost_stats + src_global_rank), + completion_end_time - logical_completion_start_time); + atomicAdd(reinterpret_cast(normal_combine_logical_recv_completion_sample_count_stats + + src_global_rank), 1); - atomicAdd(reinterpret_cast( - normal_combine_logical_recv_completion_token_count_stats + src_global_rank), + atomicAdd(reinterpret_cast(normal_combine_logical_recv_completion_token_count_stats + + src_global_rank), static_cast(logical_expected_count[src_nvl_rank])); logical_recorded_mask |= src_bit; } diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index 224e7a287..c9b5be8e9 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -32,14 +32,12 @@ def _validate_deepxtrace_normal_stats_schema(diagnose_module) -> None: - actual_schema = getattr( - diagnose_module, "NORMAL_STATS_SCHEMA", None) + actual_schema = getattr(diagnose_module, "NORMAL_STATS_SCHEMA", None) if tuple(actual_schema or ()) != _REQUIRED_NORMAL_STATS_SCHEMA: - raise RuntimeError( - "Incompatible deepxtrace normal-stats schema: DeepEP requires " - "deepxtrace>=0.2.0,<0.3.0 with fields ordered as " - "notify -> dispatch -> combine, but found schema " - f"{actual_schema!r}") + raise RuntimeError("Incompatible deepxtrace normal-stats schema: DeepEP requires " + "deepxtrace>=0.2.0,<0.3.0 with fields ordered as " + "notify -> dispatch -> combine, but found schema " + f"{actual_schema!r}") def _load_deepxtrace(enable_deepxtrace: bool, has_torch_group: bool): @@ -48,9 +46,8 @@ def _load_deepxtrace(enable_deepxtrace: bool, has_torch_group: bool): if not enable_deepxtrace: return None, "disabled by configuration" if not has_torch_group: - return None, ( - "the current DeepXTrace integration requires a " - "torch.distributed process group") + return None, ("the current DeepXTrace integration requires a " + "torch.distributed process group") try: diagnose_module = importlib.import_module("deepxtrace.diagnose") @@ -79,20 +76,21 @@ class Buffer: num_sms: int = 20 - def __init__(self, - group: Optional[dist.ProcessGroup], - num_nvl_bytes: int = 0, - num_rdma_bytes: int = 0, - low_latency_mode: bool = False, - num_qps_per_rank: int = 24, - allow_nvlink_for_low_latency_mode: bool = True, - allow_mnnvl: bool = False, - use_fabric: bool = False, - explicitly_destroy: bool = False, - enable_shrink: bool = False, - comm: Optional["mpi4py.MPI.Comm"] = None, # noqa: F821 - enable_deepxtrace: bool = False, - enable_deepxtrace_async: bool = True) -> None: + def __init__( + self, + group: Optional[dist.ProcessGroup], + num_nvl_bytes: int = 0, + num_rdma_bytes: int = 0, + low_latency_mode: bool = False, + num_qps_per_rank: int = 24, + allow_nvlink_for_low_latency_mode: bool = True, + allow_mnnvl: bool = False, + use_fabric: bool = False, + explicitly_destroy: bool = False, + enable_shrink: bool = False, + comm: Optional["mpi4py.MPI.Comm"] = None, # noqa: F821 + enable_deepxtrace: bool = False, + enable_deepxtrace_async: bool = True) -> None: """ Initialize the communication buffer. @@ -148,8 +146,7 @@ def all_gather_object(obj): else: raise ValueError("Either 'group' or 'comm' must be provided.") - diagnose_module, deepxtrace_error = _load_deepxtrace( - enable_deepxtrace, group is not None) + diagnose_module, deepxtrace_error = _load_deepxtrace(enable_deepxtrace, group is not None) local_deepxtrace_status = ( enable_deepxtrace, diagnose_module is not None, @@ -158,28 +155,17 @@ def all_gather_object(obj): ) deepxtrace_statuses = all_gather_object(local_deepxtrace_status) - all_deepxtrace_requested = all( - status[0] for status in deepxtrace_statuses) - all_deepxtrace_ready = all( - status[1] for status in deepxtrace_statuses) + all_deepxtrace_requested = all(status[0] for status in deepxtrace_statuses) + all_deepxtrace_ready = all(status[1] for status in deepxtrace_statuses) if all_deepxtrace_requested: - configured_async_modes = { - status[3] for status in deepxtrace_statuses - } + configured_async_modes = {status[3] for status in deepxtrace_statuses} if len(configured_async_modes) != 1: - raise RuntimeError( - "All EP ranks must use the same " - "`enable_deepxtrace_async`, but found " - f"{sorted(configured_async_modes)!r}") - self.deepxtrace_enabled = ( - all_deepxtrace_requested and all_deepxtrace_ready) - if (any(status[0] for status in deepxtrace_statuses) and - not self.deepxtrace_enabled and self.rank == 0): - unavailable = [ - f"rank {rank}: {status[2]}" - for rank, status in enumerate(deepxtrace_statuses) - if not status[1] - ] + raise RuntimeError("All EP ranks must use the same " + "`enable_deepxtrace_async`, but found " + f"{sorted(configured_async_modes)!r}") + self.deepxtrace_enabled = (all_deepxtrace_requested and all_deepxtrace_ready) + if (any(status[0] for status in deepxtrace_statuses) and not self.deepxtrace_enabled and self.rank == 0): + unavailable = [f"rank {rank}: {status[2]}" for rank, status in enumerate(deepxtrace_statuses) if not status[1]] preview = "; ".join(unavailable[:8]) if len(unavailable) > 8: preview += f"; ... and {len(unavailable) - 8} more ranks" @@ -191,9 +177,7 @@ def all_gather_object(obj): self.diagnose = None self._deepxtrace_finalizer = None - self._normal_diagnose_stats = { - name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA - } + self._normal_diagnose_stats = {name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA} self._normal_notify_full_kernel_timer_states = None self._normal_notify_dispatch_full_kernel_timer_state = None self._normal_cached_notify_dispatch_full_kernel_timer_state = None @@ -260,17 +244,15 @@ def all_gather_object(obj): # available when the runtime uses the low-latency NVSHMEM topology. if self.deepxtrace_enabled: self._initialize_normal_notify_timer_states() - self.diagnose = diagnose_module.Diagnose( - group=group, - enable_ll_diagnose=False, - enable_normal_diagnose=True, - enable_async=self.enable_deepxtrace_async, - snapshot_stream=self.get_comm_stream()) + self.diagnose = diagnose_module.Diagnose(group=group, + enable_ll_diagnose=False, + enable_normal_diagnose=True, + enable_async=self.enable_deepxtrace_async, + snapshot_stream=self.get_comm_stream()) self._normal_diagnose_stats = self._get_normal_diagnose_stats() if self.enable_deepxtrace_async: self.diagnose.start_async_diagnose() - self._deepxtrace_finalizer = weakref.finalize( - self, self.diagnose.stop_async_diagnose) + self._deepxtrace_finalizer = weakref.finalize(self, self.diagnose.stop_async_diagnose) # End DeepXTrace def _initialize_normal_notify_timer_states(self) -> None: @@ -280,8 +262,7 @@ def _initialize_normal_notify_timer_states(self) -> None: the selected row on the communication stream immediately before the corresponding kernel launch, so the pointers remain CUDA-graph stable. """ - self._normal_notify_full_kernel_timer_states = torch.zeros( - (3, 2), dtype=torch.int64, device="cuda") + self._normal_notify_full_kernel_timer_states = torch.zeros((3, 2), dtype=torch.int64, device="cuda") self._normal_notify_dispatch_full_kernel_timer_state = \ self._normal_notify_full_kernel_timer_states[0] self._normal_cached_notify_dispatch_full_kernel_timer_state = \ @@ -290,24 +271,19 @@ def _initialize_normal_notify_timer_states(self) -> None: self._normal_notify_full_kernel_timer_states[2] @staticmethod - def _select_normal_notify_timer_state( - duration_stats: Optional[torch.Tensor], - count_stats: Optional[torch.Tensor], - timer_state: Optional[torch.Tensor]) -> Optional[torch.Tensor]: + def _select_normal_notify_timer_state(duration_stats: Optional[torch.Tensor], count_stats: Optional[torch.Tensor], + timer_state: Optional[torch.Tensor]) -> Optional[torch.Tensor]: if (duration_stats is None) != (count_stats is None): - raise RuntimeError( - "DeepXTrace notify duration/count tensors must be enabled " - "or disabled together") + raise RuntimeError("DeepXTrace notify duration/count tensors must be enabled " + "or disabled together") return timer_state if duration_stats is not None else None - def _get_normal_diagnose_stats( - self) -> Dict[str, Optional[torch.Tensor]]: + def _get_normal_diagnose_stats(self) -> Dict[str, Optional[torch.Tensor]]: tensors = self.diagnose.get_stats_normal_stats_tensor() if len(tensors) != len(_REQUIRED_NORMAL_STATS_SCHEMA): - raise RuntimeError( - "DeepXTrace normal-stats tensor count does not match its " - "schema: " - f"{len(tensors)} != {len(_REQUIRED_NORMAL_STATS_SCHEMA)}") + raise RuntimeError("DeepXTrace normal-stats tensor count does not match its " + "schema: " + f"{len(tensors)} != {len(_REQUIRED_NORMAL_STATS_SCHEMA)}") return dict(zip(_REQUIRED_NORMAL_STATS_SCHEMA, tensors)) def diagnose_normal_sync(self, diagnose_step: int = 0): @@ -325,9 +301,8 @@ def diagnose_normal_sync(self, diagnose_step: int = 0): if self.diagnose is None: return None if self.enable_deepxtrace_async: - raise RuntimeError( - "diagnose_normal_sync() requires " - "`enable_deepxtrace_async=False`") + raise RuntimeError("diagnose_normal_sync() requires " + "`enable_deepxtrace_async=False`") return self.diagnose.diagnose_normal_sync(diagnose_step) @staticmethod @@ -354,14 +329,11 @@ def _stop_deepxtrace(self) -> None: finalizer = getattr(self, "_deepxtrace_finalizer", None) if finalizer is not None and finalizer.alive: finalizer() - elif (getattr(self, "diagnose", None) is not None and - getattr(self, "enable_deepxtrace_async", False)): + elif (getattr(self, "diagnose", None) is not None and getattr(self, "enable_deepxtrace_async", False)): self.diagnose.stop_async_diagnose() self._deepxtrace_finalizer = None self.diagnose = None - self._normal_diagnose_stats = { - name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA - } + self._normal_diagnose_stats = {name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA} self._normal_notify_full_kernel_timer_states = None self._normal_notify_dispatch_full_kernel_timer_state = None self._normal_cached_notify_dispatch_full_kernel_timer_state = None @@ -691,10 +663,8 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te # Launch the kernel with cached or non-cached mode x, x_scales = x if isinstance(x, tuple) else (x, None) normal_stats = self._normal_diagnose_stats - normal_notify_dispatch_full_kernel_duration_ns_stats = normal_stats[ - "normal_notify_dispatch_full_kernel_duration_ns_stats"] - normal_notify_dispatch_full_kernel_count_stats = normal_stats[ - "normal_notify_dispatch_full_kernel_count_stats"] + normal_notify_dispatch_full_kernel_duration_ns_stats = normal_stats["normal_notify_dispatch_full_kernel_duration_ns_stats"] + normal_notify_dispatch_full_kernel_count_stats = normal_stats["normal_notify_dispatch_full_kernel_count_stats"] normal_notify_dispatch_full_kernel_timer_state = \ self._select_normal_notify_timer_state( normal_notify_dispatch_full_kernel_duration_ns_stats, @@ -702,25 +672,18 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te self._normal_notify_dispatch_full_kernel_timer_state) normal_cached_notify_dispatch_full_kernel_duration_ns_stats = normal_stats[ "normal_cached_notify_dispatch_full_kernel_duration_ns_stats"] - normal_cached_notify_dispatch_full_kernel_count_stats = normal_stats[ - "normal_cached_notify_dispatch_full_kernel_count_stats"] + normal_cached_notify_dispatch_full_kernel_count_stats = normal_stats["normal_cached_notify_dispatch_full_kernel_count_stats"] normal_cached_notify_dispatch_full_kernel_timer_state = \ self._select_normal_notify_timer_state( normal_cached_notify_dispatch_full_kernel_duration_ns_stats, normal_cached_notify_dispatch_full_kernel_count_stats, self._normal_cached_notify_dispatch_full_kernel_timer_state) - normal_dispatch_final_completion_cost_stats = normal_stats[ - "normal_dispatch_final_completion_cost_stats"] - normal_dispatch_final_completion_sample_count_stats = normal_stats[ - "normal_dispatch_final_completion_sample_count_stats"] - normal_dispatch_final_completion_token_count_stats = normal_stats[ - "normal_dispatch_final_completion_token_count_stats"] - normal_dispatch_rdma_recv_completion_cost_stats = normal_stats[ - "normal_dispatch_rdma_recv_completion_cost_stats"] - normal_dispatch_rdma_recv_completion_sample_count_stats = normal_stats[ - "normal_dispatch_rdma_recv_completion_sample_count_stats"] - normal_dispatch_rdma_recv_completion_token_count_stats = normal_stats[ - "normal_dispatch_rdma_recv_completion_token_count_stats"] + normal_dispatch_final_completion_cost_stats = normal_stats["normal_dispatch_final_completion_cost_stats"] + normal_dispatch_final_completion_sample_count_stats = normal_stats["normal_dispatch_final_completion_sample_count_stats"] + normal_dispatch_final_completion_token_count_stats = normal_stats["normal_dispatch_final_completion_token_count_stats"] + normal_dispatch_rdma_recv_completion_cost_stats = normal_stats["normal_dispatch_rdma_recv_completion_cost_stats"] + normal_dispatch_rdma_recv_completion_sample_count_stats = normal_stats["normal_dispatch_rdma_recv_completion_sample_count_stats"] + normal_dispatch_rdma_recv_completion_token_count_stats = normal_stats["normal_dispatch_rdma_recv_completion_token_count_stats"] if handle is not None: assert topk_idx is None and topk_weights is None is_token_in_rank, \ @@ -731,19 +694,14 @@ 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, - normal_dispatch_final_completion_cost_stats, - normal_dispatch_final_completion_sample_count_stats, - normal_dispatch_final_completion_token_count_stats, - normal_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats, - normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats, - normal_notify_dispatch_full_kernel_timer_state, - normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - normal_cached_notify_dispatch_full_kernel_count_stats, + 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_final_completion_cost_stats, + normal_dispatch_final_completion_sample_count_stats, normal_dispatch_final_completion_token_count_stats, + normal_dispatch_rdma_recv_completion_cost_stats, normal_dispatch_rdma_recv_completion_sample_count_stats, + normal_dispatch_rdma_recv_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, + normal_notify_dispatch_full_kernel_count_stats, normal_notify_dispatch_full_kernel_timer_state, + normal_cached_notify_dispatch_full_kernel_duration_ns_stats, normal_cached_notify_dispatch_full_kernel_count_stats, normal_cached_notify_dispatch_full_kernel_timer_state) return (recv_x, recv_x_scales) if x_scales is not None else recv_x, None, None, None, None, EventOverlap(event) else: @@ -800,33 +758,25 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], normal_stats = self._normal_diagnose_stats normal_cached_notify_combine_full_kernel_duration_ns_stats = normal_stats[ "normal_cached_notify_combine_full_kernel_duration_ns_stats"] - normal_cached_notify_combine_full_kernel_count_stats = normal_stats[ - "normal_cached_notify_combine_full_kernel_count_stats"] + normal_cached_notify_combine_full_kernel_count_stats = normal_stats["normal_cached_notify_combine_full_kernel_count_stats"] normal_cached_notify_combine_full_kernel_timer_state = \ self._select_normal_notify_timer_state( normal_cached_notify_combine_full_kernel_duration_ns_stats, normal_cached_notify_combine_full_kernel_count_stats, self._normal_cached_notify_combine_full_kernel_timer_state) - normal_combine_logical_recv_completion_cost_stats = normal_stats[ - "normal_combine_logical_recv_completion_cost_stats"] + normal_combine_logical_recv_completion_cost_stats = normal_stats["normal_combine_logical_recv_completion_cost_stats"] normal_combine_logical_recv_completion_sample_count_stats = normal_stats[ "normal_combine_logical_recv_completion_sample_count_stats"] - normal_combine_logical_recv_completion_token_count_stats = normal_stats[ - "normal_combine_logical_recv_completion_token_count_stats"] + normal_combine_logical_recv_completion_token_count_stats = normal_stats["normal_combine_logical_recv_completion_token_count_stats"] # 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, - normal_cached_notify_combine_full_kernel_duration_ns_stats, - normal_cached_notify_combine_full_kernel_count_stats, - normal_cached_notify_combine_full_kernel_timer_state, - normal_combine_logical_recv_completion_cost_stats, - normal_combine_logical_recv_completion_sample_count_stats, - normal_combine_logical_recv_completion_token_count_stats) + 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_cached_notify_combine_full_kernel_duration_ns_stats, normal_cached_notify_combine_full_kernel_count_stats, + normal_cached_notify_combine_full_kernel_timer_state, normal_combine_logical_recv_completion_cost_stats, + normal_combine_logical_recv_completion_sample_count_stats, normal_combine_logical_recv_completion_token_count_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: From d66d866a3a3e3e368c70f732b0e4d39ba1ceece2 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Tue, 11 Aug 2026 16:24:18 +0800 Subject: [PATCH 14/16] refactor(deepxtrace): decouple diagnosis lifecycle from DeepEP --- csrc/deep_ep.cpp | 52 ++--- csrc/deep_ep.hpp | 8 +- deep_ep/buffer.py | 282 +++++---------------------- setup.py | 7 +- tests/test_deepxtrace_integration.py | 251 ------------------------ 5 files changed, 80 insertions(+), 520 deletions(-) delete mode 100644 tests/test_deepxtrace_integration.py diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index 81771ff82..a91db9e0a 100644 --- a/csrc/deep_ep.cpp +++ b/csrc/deep_ep.cpp @@ -198,6 +198,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 +317,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 @@ -955,10 +957,8 @@ Buffer::internode_dispatch(const torch::Tensor& x, const std::optional& normal_dispatch_rdma_recv_completion_token_count_stats, const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats, const std::optional& normal_notify_dispatch_full_kernel_count_stats, - const std::optional& normal_notify_dispatch_full_kernel_timer_state, const std::optional& normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats, - const std::optional& normal_cached_notify_dispatch_full_kernel_timer_state) { + const std::optional& normal_cached_notify_dispatch_full_kernel_count_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 @@ -1040,24 +1040,24 @@ Buffer::internode_dispatch(const torch::Tensor& x, normal_dispatch_rdma_recv_completion_sample_count_stats, normal_dispatch_rdma_recv_completion_token_count_stats); auto check_normal_notify_stats = [](const std::optional& duration_ns_stats, - const std::optional& count_stats, - const std::optional& timer_state) { + const std::optional& count_stats) { const bool enabled = duration_ns_stats.has_value(); - EP_HOST_ASSERT(count_stats.has_value() == enabled and timer_state.has_value() == enabled); + EP_HOST_ASSERT(count_stats.has_value() == enabled); if (enabled) { - EP_HOST_ASSERT(duration_ns_stats->scalar_type() == torch::kInt64 and count_stats->scalar_type() == torch::kInt64 and - timer_state->scalar_type() == torch::kInt64); + EP_HOST_ASSERT(duration_ns_stats->scalar_type() == torch::kInt64 and count_stats->scalar_type() == torch::kInt64); EP_HOST_ASSERT(duration_ns_stats->dim() == 1 and duration_ns_stats->numel() == 1 and duration_ns_stats->is_contiguous()); EP_HOST_ASSERT(count_stats->dim() == 1 and count_stats->numel() == 1 and count_stats->is_contiguous()); - EP_HOST_ASSERT(timer_state->dim() == 1 and timer_state->numel() == 2 and timer_state->is_contiguous()); } }; check_normal_notify_stats(normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats, - normal_notify_dispatch_full_kernel_timer_state); + normal_notify_dispatch_full_kernel_count_stats); check_normal_notify_stats(normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - normal_cached_notify_dispatch_full_kernel_count_stats, - normal_cached_notify_dispatch_full_kernel_timer_state); + normal_cached_notify_dispatch_full_kernel_count_stats); + + auto* normal_notify_dispatch_full_kernel_timer_state = + normal_notify_dispatch_full_kernel_duration_ns_stats.has_value() ? normal_notify_full_kernel_timer_states : nullptr; + auto* normal_cached_notify_dispatch_full_kernel_timer_state = + normal_cached_notify_dispatch_full_kernel_duration_ns_stats.has_value() ? normal_notify_full_kernel_timer_states + 2 : nullptr; 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)); @@ -1155,9 +1155,7 @@ Buffer::internode_dispatch(const torch::Tensor& x, normal_cached_notify_dispatch_full_kernel_count_stats.has_value() ? normal_cached_notify_dispatch_full_kernel_count_stats->data_ptr() : nullptr, - normal_cached_notify_dispatch_full_kernel_timer_state.has_value() - ? normal_cached_notify_dispatch_full_kernel_timer_state->data_ptr() - : nullptr); + normal_cached_notify_dispatch_full_kernel_timer_state); } 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)); @@ -1204,8 +1202,7 @@ Buffer::internode_dispatch(const torch::Tensor& x, : nullptr, normal_notify_dispatch_full_kernel_count_stats.has_value() ? normal_notify_dispatch_full_kernel_count_stats->data_ptr() : nullptr, - normal_notify_dispatch_full_kernel_timer_state.has_value() ? normal_notify_dispatch_full_kernel_timer_state->data_ptr() - : nullptr); + normal_notify_dispatch_full_kernel_timer_state); // Synchronize total received tokens and tokens per expert if (num_worst_tokens > 0) { @@ -1415,7 +1412,6 @@ std::tuple, std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats, const std::optional& normal_cached_notify_combine_full_kernel_count_stats, - const std::optional& normal_cached_notify_combine_full_kernel_timer_state, const std::optional& normal_combine_logical_recv_completion_cost_stats, const std::optional& normal_combine_logical_recv_completion_sample_count_stats, const std::optional& normal_combine_logical_recv_completion_token_count_stats) { @@ -1451,22 +1447,19 @@ std::tuple, std::optionalscalar_type() == torch::kInt64 and - normal_cached_notify_combine_full_kernel_count_stats->scalar_type() == torch::kInt64 and - normal_cached_notify_combine_full_kernel_timer_state->scalar_type() == torch::kInt64); + normal_cached_notify_combine_full_kernel_count_stats->scalar_type() == torch::kInt64); EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_duration_ns_stats->dim() == 1 and normal_cached_notify_combine_full_kernel_duration_ns_stats->numel() == 1 and normal_cached_notify_combine_full_kernel_duration_ns_stats->is_contiguous()); EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_count_stats->dim() == 1 and normal_cached_notify_combine_full_kernel_count_stats->numel() == 1 and normal_cached_notify_combine_full_kernel_count_stats->is_contiguous()); - EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_timer_state->dim() == 1 and - normal_cached_notify_combine_full_kernel_timer_state->numel() == 2 and - normal_cached_notify_combine_full_kernel_timer_state->is_contiguous()); } + auto* normal_cached_notify_combine_full_kernel_timer_state = + normal_cached_notify_combine_enabled ? normal_notify_full_kernel_timer_states + 4 : nullptr; const bool enable_normal_combine_logical_recv_completion_stats = normal_combine_logical_recv_completion_cost_stats.has_value(); EP_HOST_ASSERT(normal_combine_logical_recv_completion_sample_count_stats.has_value() == enable_normal_combine_logical_recv_completion_stats); @@ -1542,7 +1535,7 @@ std::tuple, std::optionaldata_ptr() : nullptr, normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_count_stats->data_ptr() : nullptr, - normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_timer_state->data_ptr() : nullptr); + normal_cached_notify_combine_full_kernel_timer_state); // Assign bias pointers auto bias_opts = std::vector>({bias_0, bias_1}); @@ -2068,10 +2061,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("normal_dispatch_rdma_recv_completion_token_count_stats") = py::none(), py::arg("normal_notify_dispatch_full_kernel_duration_ns_stats") = py::none(), py::arg("normal_notify_dispatch_full_kernel_count_stats") = py::none(), - py::arg("normal_notify_dispatch_full_kernel_timer_state") = py::none(), py::arg("normal_cached_notify_dispatch_full_kernel_duration_ns_stats") = py::none(), - py::arg("normal_cached_notify_dispatch_full_kernel_count_stats") = py::none(), - py::arg("normal_cached_notify_dispatch_full_kernel_timer_state") = py::none()) + py::arg("normal_cached_notify_dispatch_full_kernel_count_stats") = py::none()) .def("internode_combine", &deep_ep::Buffer::internode_combine, py::arg("x"), @@ -2091,7 +2082,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("allocate_on_comm_stream"), py::arg("normal_cached_notify_combine_full_kernel_duration_ns_stats") = py::none(), py::arg("normal_cached_notify_combine_full_kernel_count_stats") = py::none(), - py::arg("normal_cached_notify_combine_full_kernel_timer_state") = py::none(), py::arg("normal_combine_logical_recv_completion_cost_stats") = py::none(), py::arg("normal_combine_logical_recv_completion_sample_count_stats") = py::none(), py::arg("normal_combine_logical_recv_completion_token_count_stats") = py::none()) diff --git a/csrc/deep_ep.hpp b/csrc/deep_ep.hpp index a12b454d8..4061b1c98 100644 --- a/csrc/deep_ep.hpp +++ b/csrc/deep_ep.hpp @@ -99,6 +99,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; @@ -244,10 +247,8 @@ struct Buffer { const std::optional& normal_dispatch_rdma_recv_completion_token_count_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, const std::optional& normal_notify_dispatch_full_kernel_count_stats = std::nullopt, - const std::optional& normal_notify_dispatch_full_kernel_timer_state = std::nullopt, const std::optional& normal_cached_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, - const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats = std::nullopt, - const std::optional& normal_cached_notify_dispatch_full_kernel_timer_state = std::nullopt); + const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats = std::nullopt); std::tuple, std::optional> internode_combine( const torch::Tensor& x, @@ -267,7 +268,6 @@ struct Buffer { bool allocate_on_comm_stream, const std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats = std::nullopt, const std::optional& normal_cached_notify_combine_full_kernel_count_stats = std::nullopt, - const std::optional& normal_cached_notify_combine_full_kernel_timer_state = std::nullopt, const std::optional& normal_combine_logical_recv_completion_cost_stats = std::nullopt, const std::optional& normal_combine_logical_recv_completion_sample_count_stats = std::nullopt, const std::optional& normal_combine_logical_recv_completion_token_count_stats = std::nullopt); diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index c9b5be8e9..f0400e897 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -1,10 +1,7 @@ -import importlib import os -import warnings -import weakref import torch import torch.distributed as dist -from typing import Callable, Dict, List, Tuple, Optional, Union +from typing import Callable, List, Tuple, Optional, Union # noinspection PyUnresolvedReferences import deep_ep_cpp @@ -12,50 +9,6 @@ from deep_ep_cpp import Config, EventHandle from .utils import EventOverlap, check_nvlink_connections -_REQUIRED_NORMAL_STATS_SCHEMA = ( - "normal_notify_dispatch_full_kernel_duration_ns_stats", - "normal_notify_dispatch_full_kernel_count_stats", - "normal_cached_notify_dispatch_full_kernel_duration_ns_stats", - "normal_cached_notify_dispatch_full_kernel_count_stats", - "normal_cached_notify_combine_full_kernel_duration_ns_stats", - "normal_cached_notify_combine_full_kernel_count_stats", - "normal_dispatch_final_completion_cost_stats", - "normal_dispatch_final_completion_sample_count_stats", - "normal_dispatch_final_completion_token_count_stats", - "normal_dispatch_rdma_recv_completion_cost_stats", - "normal_dispatch_rdma_recv_completion_sample_count_stats", - "normal_dispatch_rdma_recv_completion_token_count_stats", - "normal_combine_logical_recv_completion_cost_stats", - "normal_combine_logical_recv_completion_sample_count_stats", - "normal_combine_logical_recv_completion_token_count_stats", -) - - -def _validate_deepxtrace_normal_stats_schema(diagnose_module) -> None: - actual_schema = getattr(diagnose_module, "NORMAL_STATS_SCHEMA", None) - if tuple(actual_schema or ()) != _REQUIRED_NORMAL_STATS_SCHEMA: - raise RuntimeError("Incompatible deepxtrace normal-stats schema: DeepEP requires " - "deepxtrace>=0.2.0,<0.3.0 with fields ordered as " - "notify -> dispatch -> combine, but found schema " - f"{actual_schema!r}") - - -def _load_deepxtrace(enable_deepxtrace: bool, has_torch_group: bool): - if not isinstance(enable_deepxtrace, bool): - raise TypeError("`enable_deepxtrace` must be a bool") - if not enable_deepxtrace: - return None, "disabled by configuration" - if not has_torch_group: - return None, ("the current DeepXTrace integration requires a " - "torch.distributed process group") - - try: - diagnose_module = importlib.import_module("deepxtrace.diagnose") - _validate_deepxtrace_normal_stats_schema(diagnose_module) - except Exception as exc: - return None, f"{type(exc).__name__}: {exc}" - return diagnose_module, None - class Buffer: """ @@ -88,9 +41,7 @@ def __init__( use_fabric: bool = False, explicitly_destroy: bool = False, enable_shrink: bool = False, - comm: Optional["mpi4py.MPI.Comm"] = None, # noqa: F821 - enable_deepxtrace: bool = False, - enable_deepxtrace_async: bool = True) -> None: + comm: Optional["mpi4py.MPI.Comm"] = None) -> None: # noqa: F821 """ Initialize the communication buffer. @@ -112,19 +63,8 @@ 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. comm: the `mpi4py.MPI.Comm` communicator to use in case the group parameter is absent. - enable_deepxtrace: whether to enable the optional DeepXTrace integration. DeepEP does not import or initialize - DeepXTrace unless this option is explicitly enabled on every EP rank. - enable_deepxtrace_async: whether to run DeepXTrace collection in - asynchronous mode. If enabled, the periodic background - collector uses ``DEEPEP_DIAGNOSE_INTERVAL``. If disabled, all - EP ranks must call - :meth:`diagnose_normal_sync` at the same logical step, and - ``DEEPEP_DIAGNOSE_SYNC_STEP`` controls the collection cadence. """ check_nvlink_connections(group) - if not isinstance(enable_deepxtrace_async, bool): - raise TypeError("`enable_deepxtrace_async` must be a bool") - self.enable_deepxtrace_async = enable_deepxtrace_async # Initialize the CPP runtime if group is not None: @@ -146,43 +86,6 @@ def all_gather_object(obj): else: raise ValueError("Either 'group' or 'comm' must be provided.") - diagnose_module, deepxtrace_error = _load_deepxtrace(enable_deepxtrace, group is not None) - local_deepxtrace_status = ( - enable_deepxtrace, - diagnose_module is not None, - deepxtrace_error, - self.enable_deepxtrace_async, - ) - - deepxtrace_statuses = all_gather_object(local_deepxtrace_status) - all_deepxtrace_requested = all(status[0] for status in deepxtrace_statuses) - all_deepxtrace_ready = all(status[1] for status in deepxtrace_statuses) - if all_deepxtrace_requested: - configured_async_modes = {status[3] for status in deepxtrace_statuses} - if len(configured_async_modes) != 1: - raise RuntimeError("All EP ranks must use the same " - "`enable_deepxtrace_async`, but found " - f"{sorted(configured_async_modes)!r}") - self.deepxtrace_enabled = (all_deepxtrace_requested and all_deepxtrace_ready) - if (any(status[0] for status in deepxtrace_statuses) and not self.deepxtrace_enabled and self.rank == 0): - unavailable = [f"rank {rank}: {status[2]}" for rank, status in enumerate(deepxtrace_statuses) if not status[1]] - preview = "; ".join(unavailable[:8]) - if len(unavailable) > 8: - preview += f"; ... and {len(unavailable) - 8} more ranks" - warnings.warn( - "DeepXTrace diagnosis is disabled for all ranks because " - f"the integration is not uniformly available: {preview}", - RuntimeWarning, - stacklevel=2) - - self.diagnose = None - self._deepxtrace_finalizer = None - self._normal_diagnose_stats = {name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA} - self._normal_notify_full_kernel_timer_states = None - self._normal_notify_dispatch_full_kernel_timer_state = None - self._normal_cached_notify_dispatch_full_kernel_timer_state = None - self._normal_cached_notify_combine_full_kernel_timer_state = None - self.num_nvl_bytes = num_nvl_bytes self.num_rdma_bytes = num_rdma_bytes self.low_latency_mode = low_latency_mode @@ -238,72 +141,18 @@ def all_gather_object(obj): self.runtime.sync(device_ids, ipc_handles, root_unique_id) assert self.runtime.is_available() - # Start DeepXTrace - # LL diagnosis is intentionally disabled by DeepEP for now. Normal - # diagnostics instrument the normal dispatch/combine APIs and remain - # available when the runtime uses the low-latency NVSHMEM topology. - if self.deepxtrace_enabled: - self._initialize_normal_notify_timer_states() - self.diagnose = diagnose_module.Diagnose(group=group, - enable_ll_diagnose=False, - enable_normal_diagnose=True, - enable_async=self.enable_deepxtrace_async, - snapshot_stream=self.get_comm_stream()) - self._normal_diagnose_stats = self._get_normal_diagnose_stats() - if self.enable_deepxtrace_async: - self.diagnose.start_async_diagnose() - self._deepxtrace_finalizer = weakref.finalize(self, self.diagnose.stop_async_diagnose) - # End DeepXTrace - - def _initialize_normal_notify_timer_states(self) -> None: - """Create persistent per-launch notify timer scratch owned by DeepEP. - - Each row is ``[inverted earliest start, completed blocks]``. C++ resets - the selected row on the communication stream immediately before the - corresponding kernel launch, so the pointers remain CUDA-graph stable. - """ - self._normal_notify_full_kernel_timer_states = torch.zeros((3, 2), dtype=torch.int64, device="cuda") - self._normal_notify_dispatch_full_kernel_timer_state = \ - self._normal_notify_full_kernel_timer_states[0] - self._normal_cached_notify_dispatch_full_kernel_timer_state = \ - self._normal_notify_full_kernel_timer_states[1] - self._normal_cached_notify_combine_full_kernel_timer_state = \ - self._normal_notify_full_kernel_timer_states[2] - @staticmethod - def _select_normal_notify_timer_state(duration_stats: Optional[torch.Tensor], count_stats: Optional[torch.Tensor], - timer_state: Optional[torch.Tensor]) -> Optional[torch.Tensor]: - if (duration_stats is None) != (count_stats is None): - raise RuntimeError("DeepXTrace notify duration/count tensors must be enabled " - "or disabled together") - return timer_state if duration_stats is not None else None - - def _get_normal_diagnose_stats(self) -> Dict[str, Optional[torch.Tensor]]: - tensors = self.diagnose.get_stats_normal_stats_tensor() - if len(tensors) != len(_REQUIRED_NORMAL_STATS_SCHEMA): - raise RuntimeError("DeepXTrace normal-stats tensor count does not match its " - "schema: " - f"{len(tensors)} != {len(_REQUIRED_NORMAL_STATS_SCHEMA)}") - return dict(zip(_REQUIRED_NORMAL_STATS_SCHEMA, tensors)) - - def diagnose_normal_sync(self, diagnose_step: int = 0): - """Collect normal DeepXTrace statistics at a caller-owned step boundary. - - Every rank in the EP group must call this method at the same logical - location and with the same cadence. Mismatched calls can deadlock the - diagnostic collectives or assign samples to different windows. - ``DEEPEP_DIAGNOSE_SYNC_STEP`` is used when ``diagnose_step`` is zero; - a nonzero value overrides it. - - Returns ``None`` when DeepXTrace is disabled. In async mode, collection - is owned by the background thread and this method raises. - """ - if self.diagnose is None: - return None - if self.enable_deepxtrace_async: - raise RuntimeError("diagnose_normal_sync() requires " - "`enable_deepxtrace_async=False`") - return self.diagnose.diagnose_normal_sync(diagnose_step) + def _unpack_normal_stats( + stats: Optional[Tuple[Optional[torch.Tensor], ...]], + expected_count: int, + argument_name: str) -> Tuple[Optional[torch.Tensor], ...]: + if stats is None: + return (None,) * expected_count + if len(stats) != expected_count: + raise ValueError( + f"`{argument_name}` must contain {expected_count} tensors, " + f"but got {len(stats)}") + return stats @staticmethod def disable_ll_layered() -> bool: @@ -320,25 +169,9 @@ def destroy(self): assert self.explicitly_destroy, '`explicitly_destroy` flag must be set' - self._stop_deepxtrace() self.runtime.destroy() self.runtime = None - def _stop_deepxtrace(self) -> None: - """Stop the optional background collector before runtime teardown.""" - finalizer = getattr(self, "_deepxtrace_finalizer", None) - if finalizer is not None and finalizer.alive: - finalizer() - elif (getattr(self, "diagnose", None) is not None and getattr(self, "enable_deepxtrace_async", False)): - self.diagnose.stop_async_diagnose() - self._deepxtrace_finalizer = None - self.diagnose = None - self._normal_diagnose_stats = {name: None for name in _REQUIRED_NORMAL_STATS_SCHEMA} - self._normal_notify_full_kernel_timer_states = None - self._normal_notify_dispatch_full_kernel_timer_state = None - self._normal_cached_notify_dispatch_full_kernel_timer_state = None - self._normal_cached_notify_combine_full_kernel_timer_state = None - @staticmethod def is_sm90_compiled(): return deep_ep_cpp.is_sm90_compiled() @@ -521,7 +354,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]: """ @@ -551,6 +385,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 @@ -570,7 +407,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) @@ -601,7 +438,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 @@ -619,6 +457,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. @@ -630,7 +471,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 @@ -651,7 +493,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]: """ @@ -662,28 +505,18 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te # Launch the kernel with cached or non-cached mode x, x_scales = x if isinstance(x, tuple) else (x, None) - normal_stats = self._normal_diagnose_stats - normal_notify_dispatch_full_kernel_duration_ns_stats = normal_stats["normal_notify_dispatch_full_kernel_duration_ns_stats"] - normal_notify_dispatch_full_kernel_count_stats = normal_stats["normal_notify_dispatch_full_kernel_count_stats"] - normal_notify_dispatch_full_kernel_timer_state = \ - self._select_normal_notify_timer_state( - normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats, - self._normal_notify_dispatch_full_kernel_timer_state) - normal_cached_notify_dispatch_full_kernel_duration_ns_stats = normal_stats[ - "normal_cached_notify_dispatch_full_kernel_duration_ns_stats"] - normal_cached_notify_dispatch_full_kernel_count_stats = normal_stats["normal_cached_notify_dispatch_full_kernel_count_stats"] - normal_cached_notify_dispatch_full_kernel_timer_state = \ - self._select_normal_notify_timer_state( - normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - normal_cached_notify_dispatch_full_kernel_count_stats, - self._normal_cached_notify_dispatch_full_kernel_timer_state) - normal_dispatch_final_completion_cost_stats = normal_stats["normal_dispatch_final_completion_cost_stats"] - normal_dispatch_final_completion_sample_count_stats = normal_stats["normal_dispatch_final_completion_sample_count_stats"] - normal_dispatch_final_completion_token_count_stats = normal_stats["normal_dispatch_final_completion_token_count_stats"] - normal_dispatch_rdma_recv_completion_cost_stats = normal_stats["normal_dispatch_rdma_recv_completion_cost_stats"] - normal_dispatch_rdma_recv_completion_sample_count_stats = normal_stats["normal_dispatch_rdma_recv_completion_sample_count_stats"] - normal_dispatch_rdma_recv_completion_token_count_stats = normal_stats["normal_dispatch_rdma_recv_completion_token_count_stats"] + normal_notify_dispatch_full_kernel_duration_ns_stats, \ + normal_notify_dispatch_full_kernel_count_stats, \ + normal_cached_notify_dispatch_full_kernel_duration_ns_stats, \ + normal_cached_notify_dispatch_full_kernel_count_stats, \ + normal_dispatch_final_completion_cost_stats, \ + normal_dispatch_final_completion_sample_count_stats, \ + normal_dispatch_final_completion_token_count_stats, \ + normal_dispatch_rdma_recv_completion_cost_stats, \ + normal_dispatch_rdma_recv_completion_sample_count_stats, \ + normal_dispatch_rdma_recv_completion_token_count_stats = \ + self._unpack_normal_stats( + normal_dispatch_stats, 10, "normal_dispatch_stats") if handle is not None: assert topk_idx is None and topk_weights is None is_token_in_rank, \ @@ -700,9 +533,9 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te normal_dispatch_final_completion_sample_count_stats, normal_dispatch_final_completion_token_count_stats, normal_dispatch_rdma_recv_completion_cost_stats, normal_dispatch_rdma_recv_completion_sample_count_stats, normal_dispatch_rdma_recv_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats, normal_notify_dispatch_full_kernel_timer_state, + normal_notify_dispatch_full_kernel_count_stats, normal_cached_notify_dispatch_full_kernel_duration_ns_stats, normal_cached_notify_dispatch_full_kernel_count_stats, - normal_cached_notify_dispatch_full_kernel_timer_state) + ) 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 @@ -723,10 +556,8 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te normal_dispatch_rdma_recv_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, normal_notify_dispatch_full_kernel_count_stats, - normal_notify_dispatch_full_kernel_timer_state, normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - normal_cached_notify_dispatch_full_kernel_count_stats, - normal_cached_notify_dispatch_full_kernel_timer_state) + normal_cached_notify_dispatch_full_kernel_count_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) @@ -741,7 +572,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. @@ -755,19 +587,13 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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) - normal_stats = self._normal_diagnose_stats - normal_cached_notify_combine_full_kernel_duration_ns_stats = normal_stats[ - "normal_cached_notify_combine_full_kernel_duration_ns_stats"] - normal_cached_notify_combine_full_kernel_count_stats = normal_stats["normal_cached_notify_combine_full_kernel_count_stats"] - normal_cached_notify_combine_full_kernel_timer_state = \ - self._select_normal_notify_timer_state( - normal_cached_notify_combine_full_kernel_duration_ns_stats, - normal_cached_notify_combine_full_kernel_count_stats, - self._normal_cached_notify_combine_full_kernel_timer_state) - normal_combine_logical_recv_completion_cost_stats = normal_stats["normal_combine_logical_recv_completion_cost_stats"] - normal_combine_logical_recv_completion_sample_count_stats = normal_stats[ - "normal_combine_logical_recv_completion_sample_count_stats"] - normal_combine_logical_recv_completion_token_count_stats = normal_stats["normal_combine_logical_recv_completion_token_count_stats"] + normal_cached_notify_combine_full_kernel_duration_ns_stats, \ + normal_cached_notify_combine_full_kernel_count_stats, \ + normal_combine_logical_recv_completion_cost_stats, \ + normal_combine_logical_recv_completion_sample_count_stats, \ + normal_combine_logical_recv_completion_token_count_stats = \ + self._unpack_normal_stats( + normal_combine_stats, 5, "normal_combine_stats") # Launch the kernel combined_x, combined_topk_weights, event = self.runtime.internode_combine( @@ -775,8 +601,8 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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_cached_notify_combine_full_kernel_duration_ns_stats, normal_cached_notify_combine_full_kernel_count_stats, - normal_cached_notify_combine_full_kernel_timer_state, normal_combine_logical_recv_completion_cost_stats, - normal_combine_logical_recv_completion_sample_count_stats, normal_combine_logical_recv_completion_token_count_stats) + normal_combine_logical_recv_completion_cost_stats, normal_combine_logical_recv_completion_sample_count_stats, + normal_combine_logical_recv_completion_token_count_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 58b3bb369..f135107bd 100644 --- a/setup.py +++ b/setup.py @@ -115,15 +115,10 @@ def get_nvshmem_host_lib_name(base_dir): revision = '+' + subprocess.check_output(cmd).decode('ascii').rstrip() except Exception as _: revision = '' + setuptools.setup(name='deep_ep', version='1.2.1' + revision, packages=setuptools.find_packages(include=['deep_ep']), - install_requires=[], - extras_require={ - 'deepxtrace': [ - 'deepxtrace>=0.2.0,<0.3.0', - ], - }, ext_modules=[ CUDAExtension(name='deep_ep_cpp', include_dirs=include_dirs, diff --git a/tests/test_deepxtrace_integration.py b/tests/test_deepxtrace_integration.py deleted file mode 100644 index b0e8ea939..000000000 --- a/tests/test_deepxtrace_integration.py +++ /dev/null @@ -1,251 +0,0 @@ -import gc -import importlib.util -import inspect -import sys -import types -import unittest -import weakref -from pathlib import Path -from unittest.mock import MagicMock, patch - - -def _load_buffer_module(): - """Load buffer.py with host-only stubs for torch and the CUDA extension.""" - torch_module = types.ModuleType("torch") - torch_module.Tensor = type("Tensor", (), {}) - torch_module.Stream = type("Stream", (), {}) - torch_module.Size = tuple - torch_module.dtype = type("dtype", (), {}) - torch_module.cuda = types.SimpleNamespace() - - dist_module = types.ModuleType("torch.distributed") - dist_module.ProcessGroup = type("ProcessGroup", (), {}) - dist_module.all_gather_object = lambda outputs, value, group: None - torch_module.distributed = dist_module - - cpp_module = types.ModuleType("deep_ep_cpp") - cpp_module.Buffer = type("Buffer", (), {}) - cpp_module.Config = type("Config", (), {}) - cpp_module.EventHandle = type("EventHandle", (), {}) - - package_module = types.ModuleType("deep_ep") - package_module.__path__ = [] - utils_module = types.ModuleType("deep_ep.utils") - - class EventOverlap: - - def __init__(self, *args): - self.args = args - - utils_module.EventOverlap = EventOverlap - utils_module.check_nvlink_connections = lambda group: None - - module_name = "deep_ep._buffer_host_test" - module_path = Path(__file__).resolve().parents[1] / "deep_ep" / "buffer.py" - spec = importlib.util.spec_from_file_location(module_name, module_path) - module = importlib.util.module_from_spec(spec) - stubs = { - "torch": torch_module, - "torch.distributed": dist_module, - "deep_ep_cpp": cpp_module, - "deep_ep": package_module, - "deep_ep.utils": utils_module, - module_name: module, - } - with patch.dict(sys.modules, stubs): - spec.loader.exec_module(module) - return module - - -buffer_module = _load_buffer_module() - - -class _FakeGroup: - - def rank(self): - return 0 - - def size(self): - return 2 - - -class _FakeRuntime: - - def __init__(self, *args): - self.args = args - self.synced = False - - def get_local_device_id(self): - return 0 - - def get_local_ipc_handle(self): - return bytearray() - - def get_num_rdma_ranks(self): - return 1 - - def sync(self, *args): - self.synced = True - - def is_available(self): - return self.synced - - -class TestDeepXTraceIntegration(unittest.TestCase): - - @staticmethod - def compatible_module(): - return types.SimpleNamespace(NORMAL_STATS_SCHEMA=buffer_module._REQUIRED_NORMAL_STATS_SCHEMA) - - def test_optional_defaults_are_backward_compatible(self): - parameters = inspect.signature(buffer_module.Buffer.__init__).parameters - self.assertIs(parameters["low_latency_mode"].default, False) - self.assertIs(parameters["enable_deepxtrace"].default, False) - - def test_optional_dependency_loading_is_fail_closed(self): - with patch.object(buffer_module.importlib, "import_module") as import_module: - diagnose_module, error = buffer_module._load_deepxtrace(False, True) - import_module.assert_not_called() - self.assertIsNone(diagnose_module) - self.assertEqual(error, "disabled by configuration") - - with patch.object(buffer_module.importlib, "import_module", return_value=self.compatible_module()): - diagnose_module, error = buffer_module._load_deepxtrace(True, True) - self.assertIsNotNone(diagnose_module) - self.assertIsNone(error) - - for exception in ( - ModuleNotFoundError("No module named 'deepxtrace'"), - OSError("broken optional dependency"), - AttributeError("incomplete installation"), - ): - with self.subTest(exception=exception): - with patch.object(buffer_module.importlib, "import_module", side_effect=exception): - diagnose_module, error = \ - buffer_module._load_deepxtrace(True, True) - self.assertIsNone(diagnose_module) - self.assertIn(type(exception).__name__, error) - - def test_schema_mismatch_is_reported_as_unavailable(self): - incompatible_module = types.SimpleNamespace(NORMAL_STATS_SCHEMA=("legacy", )) - with self.assertRaisesRegex(RuntimeError, "deepxtrace>=0.2.0"): - buffer_module._validate_deepxtrace_normal_stats_schema(incompatible_module) - - with patch.object(buffer_module.importlib, "import_module", return_value=incompatible_module): - diagnose_module, error = buffer_module._load_deepxtrace(True, True) - self.assertIsNone(diagnose_module) - self.assertIn("RuntimeError", error) - - def test_rank_wide_mixed_enablement_disables_diagnosis(self): - - def all_gather_object(outputs, value, group): - if isinstance(value, tuple) and len(value) == 4: - outputs[:] = [ - value, - (False, False, "disabled by configuration", value[3]), - ] - else: - outputs[:] = [value, value] - - with patch.object(buffer_module.deep_ep_cpp, "Buffer", _FakeRuntime), \ - patch.object(buffer_module.dist, "all_gather_object", - side_effect=all_gather_object), \ - patch.object(buffer_module, "_load_deepxtrace", - return_value=(self.compatible_module(), None)), \ - self.assertWarnsRegex(RuntimeWarning, "disabled for all ranks"): - buffer = buffer_module.Buffer(_FakeGroup(), enable_deepxtrace=True) - self.assertFalse(buffer.deepxtrace_enabled) - self.assertTrue(buffer.runtime.is_available()) - - def test_rank_wide_async_mode_mismatch_raises(self): - - def all_gather_object(outputs, value, group): - outputs[:] = [value, (True, True, None, not value[3])] - - with patch.object(buffer_module.dist, "all_gather_object", - side_effect=all_gather_object), \ - patch.object(buffer_module, "_load_deepxtrace", - return_value=(self.compatible_module(), None)), \ - self.assertRaisesRegex(RuntimeError, - "enable_deepxtrace_async"): - buffer_module.Buffer(_FakeGroup(), enable_deepxtrace=True) - - def test_sync_collection_contract(self): - buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) - buffer.diagnose = None - buffer.enable_deepxtrace_async = False - self.assertIsNone(buffer.diagnose_normal_sync()) - - buffer.diagnose = MagicMock() - buffer.enable_deepxtrace_async = True - with self.assertRaisesRegex(RuntimeError, "enable_deepxtrace_async=False"): - buffer.diagnose_normal_sync() - - expected = [{"probe": "normal", "status": "ok"}] - buffer.enable_deepxtrace_async = False - buffer.diagnose.diagnose_normal_sync.return_value = expected - self.assertIs(buffer.diagnose_normal_sync(17), expected) - buffer.diagnose.diagnose_normal_sync.assert_called_once_with(17) - - def test_destroy_stops_async_diagnosis_before_runtime(self): - order = [] - buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) - buffer.explicitly_destroy = True - buffer.enable_deepxtrace_async = True - buffer.diagnose = MagicMock() - buffer.diagnose.stop_async_diagnose.side_effect = \ - lambda: order.append("diagnose") - buffer.runtime = MagicMock() - buffer.runtime.destroy.side_effect = lambda: order.append("runtime") - buffer._deepxtrace_finalizer = weakref.finalize(buffer, buffer.diagnose.stop_async_diagnose) - buffer._normal_diagnose_stats = {"probe": object()} - buffer._normal_notify_full_kernel_timer_states = object() - buffer._normal_notify_dispatch_full_kernel_timer_state = object() - buffer._normal_cached_notify_dispatch_full_kernel_timer_state = object() - buffer._normal_cached_notify_combine_full_kernel_timer_state = object() - - buffer.destroy() - - self.assertEqual(order, ["diagnose", "runtime"]) - self.assertIsNone(buffer.runtime) - self.assertIsNone(buffer.diagnose) - self.assertFalse(buffer._deepxtrace_finalizer) - self.assertTrue(all(value is None for value in buffer._normal_diagnose_stats.values())) - - def test_finalizer_stops_async_diagnosis_on_implicit_cleanup(self): - stop_async_diagnose = MagicMock() - buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) - buffer._deepxtrace_finalizer = weakref.finalize(buffer, stop_async_diagnose) - buffer_ref = weakref.ref(buffer) - - del buffer - gc.collect() - - self.assertIsNone(buffer_ref()) - stop_async_diagnose.assert_called_once_with() - - def test_low_latency_stats_remain_caller_owned(self): - buffer = buffer_module.Buffer.__new__(buffer_module.Buffer) - buffer.nvshmem_qp_depth = 1024 - buffer.diagnose = MagicMock() - buffer.runtime = MagicMock() - buffer.runtime.low_latency_dispatch.return_value = (object(), object(), object(), object(), object(), object(), object()) - x = MagicMock() - x.size.return_value = 64 - - buffer.low_latency_dispatch(x, MagicMock(), 1, 1) - - dispatch_args = buffer.runtime.low_latency_dispatch.call_args.args - self.assertIsNone(dispatch_args[3]) - - buffer.runtime.low_latency_combine.return_value = (object(), object(), object()) - handle = (object(), object(), 1, 64, 1) - buffer.low_latency_combine(MagicMock(), MagicMock(), MagicMock(), handle) - - combine_args = buffer.runtime.low_latency_combine.call_args.args - self.assertIsNone(combine_args[11]) - buffer.diagnose.get_stats_ll_stats_tensor.assert_not_called() - - -if __name__ == "__main__": - unittest.main() From 706415440c2d0205647064e295cbe9d0408cf7dc Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Wed, 12 Aug 2026 18:42:26 +0800 Subject: [PATCH 15/16] test(normal): add optional DeepXTrace normal-mode test --- deep_ep/buffer.py | 26 +++++++++++------------ tests/test_internode.py | 46 +++++++++++++++++++++++++++++++++-------- 2 files changed, 49 insertions(+), 23 deletions(-) diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index f0400e897..a7b27a6ce 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -29,19 +29,18 @@ class Buffer: num_sms: int = 20 - def __init__( - self, - group: Optional[dist.ProcessGroup], - num_nvl_bytes: int = 0, - num_rdma_bytes: int = 0, - low_latency_mode: bool = False, - num_qps_per_rank: int = 24, - allow_nvlink_for_low_latency_mode: bool = True, - allow_mnnvl: bool = False, - use_fabric: bool = False, - explicitly_destroy: bool = False, - enable_shrink: bool = False, - comm: Optional["mpi4py.MPI.Comm"] = None) -> None: # noqa: F821 + def __init__(self, + group: Optional[dist.ProcessGroup], + num_nvl_bytes: int = 0, + num_rdma_bytes: int = 0, + low_latency_mode: bool = False, + num_qps_per_rank: int = 24, + allow_nvlink_for_low_latency_mode: bool = True, + allow_mnnvl: bool = False, + use_fabric: bool = False, + explicitly_destroy: bool = False, + enable_shrink: bool = False, + comm: Optional["mpi4py.MPI.Comm"] = None) -> None: # noqa: F821 """ Initialize the communication buffer. @@ -85,7 +84,6 @@ def all_gather_object(obj): return comm.allgather(obj) else: raise ValueError("Either 'group' or 'comm' must be provided.") - self.num_nvl_bytes = num_nvl_bytes self.num_rdma_bytes = num_rdma_bytes self.low_latency_mode = low_latency_mode 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 From 8a2d98f2956642d14952dd09fc3acac75f82c6e3 Mon Sep 17 00:00:00 2001 From: Xing Yuming Date: Wed, 26 Aug 2026 14:41:12 +0800 Subject: [PATCH 16/16] refactor(normal): group diagnostic probe arguments --- csrc/deep_ep.cpp | 448 ++++++++++++++++++-------------------- csrc/deep_ep.hpp | 40 ++-- csrc/kernels/api.cuh | 32 +-- csrc/kernels/internode.cu | 277 ++++++++++------------- deep_ep/buffer.py | 56 +---- 5 files changed, 374 insertions(+), 479 deletions(-) diff --git a/csrc/deep_ep.cpp b/csrc/deep_ep.cpp index a91db9e0a..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, @@ -949,16 +992,7 @@ Buffer::internode_dispatch(const torch::Tensor& x, std::optional& previous_event, bool async, bool allocate_on_comm_stream, - const std::optional& normal_dispatch_final_completion_cost_stats, - const std::optional& normal_dispatch_final_completion_sample_count_stats, - const std::optional& normal_dispatch_final_completion_token_count_stats, - const std::optional& normal_dispatch_rdma_recv_completion_cost_stats, - const std::optional& normal_dispatch_rdma_recv_completion_sample_count_stats, - const std::optional& normal_dispatch_rdma_recv_completion_token_count_stats, - const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats, - const std::optional& normal_notify_dispatch_full_kernel_count_stats, - const std::optional& normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats) { + 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 @@ -1016,48 +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); } - auto check_normal_dispatch_stat_tensor = [=](const std::optional& stats) { - if (stats.has_value()) { - EP_HOST_ASSERT(stats->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(stats->dim() == 1 and stats->is_contiguous()); - EP_HOST_ASSERT(stats->size(0) == num_ranks); - } - }; - auto check_normal_dispatch_stat_triplet = [&](const std::optional& cost_stats, - const std::optional& sample_count_stats, - const std::optional& token_count_stats) { - const bool enabled = cost_stats.has_value(); - EP_HOST_ASSERT(sample_count_stats.has_value() == enabled); - EP_HOST_ASSERT(token_count_stats.has_value() == enabled); - check_normal_dispatch_stat_tensor(cost_stats); - check_normal_dispatch_stat_tensor(sample_count_stats); - check_normal_dispatch_stat_tensor(token_count_stats); - }; - check_normal_dispatch_stat_triplet(normal_dispatch_final_completion_cost_stats, - normal_dispatch_final_completion_sample_count_stats, - normal_dispatch_final_completion_token_count_stats); - check_normal_dispatch_stat_triplet(normal_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats); - auto check_normal_notify_stats = [](const std::optional& duration_ns_stats, - const std::optional& count_stats) { - const bool enabled = duration_ns_stats.has_value(); - EP_HOST_ASSERT(count_stats.has_value() == enabled); - if (enabled) { - EP_HOST_ASSERT(duration_ns_stats->scalar_type() == torch::kInt64 and count_stats->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(duration_ns_stats->dim() == 1 and duration_ns_stats->numel() == 1 and duration_ns_stats->is_contiguous()); - EP_HOST_ASSERT(count_stats->dim() == 1 and count_stats->numel() == 1 and count_stats->is_contiguous()); - } - }; - check_normal_notify_stats(normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats); - check_normal_notify_stats(normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - normal_cached_notify_dispatch_full_kernel_count_stats); - - auto* normal_notify_dispatch_full_kernel_timer_state = - normal_notify_dispatch_full_kernel_duration_ns_stats.has_value() ? normal_notify_full_kernel_timer_states : nullptr; - auto* normal_cached_notify_dispatch_full_kernel_timer_state = - normal_cached_notify_dispatch_full_kernel_duration_ns_stats.has_value() ? normal_notify_full_kernel_timer_states + 2 : nullptr; + 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)); @@ -1149,13 +1145,7 @@ Buffer::internode_dispatch(const torch::Tensor& x, num_nvl_bytes, true, low_latency_mode, - normal_cached_notify_dispatch_full_kernel_duration_ns_stats.has_value() - ? normal_cached_notify_dispatch_full_kernel_duration_ns_stats->data_ptr() - : nullptr, - normal_cached_notify_dispatch_full_kernel_count_stats.has_value() - ? normal_cached_notify_dispatch_full_kernel_count_stats->data_ptr() - : nullptr, - normal_cached_notify_dispatch_full_kernel_timer_state); + 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)); @@ -1166,43 +1156,37 @@ Buffer::internode_dispatch(const torch::Tensor& x, *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_worst_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, - normal_notify_dispatch_full_kernel_duration_ns_stats.has_value() - ? normal_notify_dispatch_full_kernel_duration_ns_stats->data_ptr() - : nullptr, - normal_notify_dispatch_full_kernel_count_stats.has_value() ? normal_notify_dispatch_full_kernel_count_stats->data_ptr() - : nullptr, - normal_notify_dispatch_full_kernel_timer_state); + 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_worst_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, + notify_stats); // Synchronize total received tokens and tokens per expert if (num_worst_tokens > 0) { @@ -1271,61 +1255,46 @@ Buffer::internode_dispatch(const torch::Tensor& x, // 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, - num_worst_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, - normal_dispatch_final_completion_cost_stats.has_value() ? normal_dispatch_final_completion_cost_stats->data_ptr() - : nullptr, - normal_dispatch_final_completion_sample_count_stats.has_value() - ? normal_dispatch_final_completion_sample_count_stats->data_ptr() - : nullptr, - normal_dispatch_final_completion_token_count_stats.has_value() - ? normal_dispatch_final_completion_token_count_stats->data_ptr() - : nullptr, - normal_dispatch_rdma_recv_completion_cost_stats.has_value() ? normal_dispatch_rdma_recv_completion_cost_stats->data_ptr() - : nullptr, - normal_dispatch_rdma_recv_completion_sample_count_stats.has_value() - ? normal_dispatch_rdma_recv_completion_sample_count_stats->data_ptr() - : nullptr, - normal_dispatch_rdma_recv_completion_token_count_stats.has_value() - ? normal_dispatch_rdma_recv_completion_token_count_stats->data_ptr() - : nullptr, - rank, - num_ranks, - cached_mode, - comm_stream, - num_channels, - low_latency_mode); + 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, + num_worst_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, + final_completion_stats, + rdma_recv_completion_stats, + rank, + num_ranks, + cached_mode, + comm_stream, + num_channels, + low_latency_mode); // Wait streams std::optional event; @@ -1410,11 +1379,7 @@ std::tuple, std::optional& previous_event, bool async, bool allocate_on_comm_stream, - const std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats, - const std::optional& normal_cached_notify_combine_full_kernel_count_stats, - const std::optional& normal_combine_logical_recv_completion_cost_stats, - const std::optional& normal_combine_logical_recv_completion_sample_count_stats, - const std::optional& normal_combine_logical_recv_completion_token_count_stats) { + const NormalCombineStats& normal_stats) { #ifndef DISABLE_NVSHMEM const int num_channels = config.num_sms / 2; EP_HOST_ASSERT(config.num_sms % 2 == 0); @@ -1446,34 +1411,8 @@ std::tuple, std::optionalscalar_type() == torch::kInt64 and - normal_cached_notify_combine_full_kernel_count_stats->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_duration_ns_stats->dim() == 1 and - normal_cached_notify_combine_full_kernel_duration_ns_stats->numel() == 1 and - normal_cached_notify_combine_full_kernel_duration_ns_stats->is_contiguous()); - EP_HOST_ASSERT(normal_cached_notify_combine_full_kernel_count_stats->dim() == 1 and - normal_cached_notify_combine_full_kernel_count_stats->numel() == 1 and - normal_cached_notify_combine_full_kernel_count_stats->is_contiguous()); - } - auto* normal_cached_notify_combine_full_kernel_timer_state = - normal_cached_notify_combine_enabled ? normal_notify_full_kernel_timer_states + 4 : nullptr; - const bool enable_normal_combine_logical_recv_completion_stats = normal_combine_logical_recv_completion_cost_stats.has_value(); - EP_HOST_ASSERT(normal_combine_logical_recv_completion_sample_count_stats.has_value() == - enable_normal_combine_logical_recv_completion_stats); - EP_HOST_ASSERT(normal_combine_logical_recv_completion_token_count_stats.has_value() == - enable_normal_combine_logical_recv_completion_stats); - for (const auto& stats : {normal_combine_logical_recv_completion_cost_stats, - normal_combine_logical_recv_completion_sample_count_stats, - normal_combine_logical_recv_completion_token_count_stats}) { - if (stats.has_value()) { - EP_HOST_ASSERT(stats->scalar_type() == torch::kInt64); - EP_HOST_ASSERT(stats->dim() == 1 and stats->is_contiguous()); - EP_HOST_ASSERT(stats->numel() == num_ranks); - } - } + const auto cached_notify_stats = prepare_normal_notify_stats(normal_stats.cached_notify, normal_notify_full_kernel_timer_states + 4); + const auto logical_recv_completion_stats = prepare_normal_completion_stats(normal_stats.logical_recv_completion, num_ranks); // Allocate all tensors on comm stream if set // NOTES: do not allocate tensors upfront! @@ -1510,32 +1449,29 @@ std::tuple, std::optional(), - 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, - normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_duration_ns_stats->data_ptr() : nullptr, - normal_cached_notify_combine_enabled ? normal_cached_notify_combine_full_kernel_count_stats->data_ptr() : nullptr, - normal_cached_notify_combine_full_kernel_timer_state); + 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, + cached_notify_stats); // Assign bias pointers auto bias_opts = std::vector>({bias_0, bias_1}); @@ -1575,15 +1511,7 @@ std::tuple, std::optionaldata_ptr() - : nullptr, - normal_combine_logical_recv_completion_sample_count_stats.has_value() - ? normal_combine_logical_recv_completion_sample_count_stats->data_ptr() - : nullptr, - normal_combine_logical_recv_completion_token_count_stats.has_value() - ? normal_combine_logical_recv_completion_token_count_stats->data_ptr() - : nullptr, + logical_recv_completion_stats, rank, num_ranks, comm_stream, @@ -1997,6 +1925,71 @@ void Buffer::low_latency_clean_mask_buffer() { } // namespace deep_ep +namespace { + +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"; @@ -2053,16 +2046,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("previous_event"), py::arg("async"), py::arg("allocate_on_comm_stream"), - py::arg("normal_dispatch_final_completion_cost_stats") = py::none(), - py::arg("normal_dispatch_final_completion_sample_count_stats") = py::none(), - py::arg("normal_dispatch_final_completion_token_count_stats") = py::none(), - py::arg("normal_dispatch_rdma_recv_completion_cost_stats") = py::none(), - py::arg("normal_dispatch_rdma_recv_completion_sample_count_stats") = py::none(), - py::arg("normal_dispatch_rdma_recv_completion_token_count_stats") = py::none(), - py::arg("normal_notify_dispatch_full_kernel_duration_ns_stats") = py::none(), - py::arg("normal_notify_dispatch_full_kernel_count_stats") = py::none(), - py::arg("normal_cached_notify_dispatch_full_kernel_duration_ns_stats") = py::none(), - py::arg("normal_cached_notify_dispatch_full_kernel_count_stats") = py::none()) + py::arg("normal_dispatch_stats") = py::none()) .def("internode_combine", &deep_ep::Buffer::internode_combine, py::arg("x"), @@ -2080,11 +2064,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("previous_event"), py::arg("async"), py::arg("allocate_on_comm_stream"), - py::arg("normal_cached_notify_combine_full_kernel_duration_ns_stats") = py::none(), - py::arg("normal_cached_notify_combine_full_kernel_count_stats") = py::none(), - py::arg("normal_combine_logical_recv_completion_cost_stats") = py::none(), - py::arg("normal_combine_logical_recv_completion_sample_count_stats") = py::none(), - py::arg("normal_combine_logical_recv_completion_token_count_stats") = py::none()) + 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 4061b1c98..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"); @@ -239,16 +262,7 @@ struct Buffer { std::optional& previous_event, bool async, bool allocate_on_comm_stream, - const std::optional& normal_dispatch_final_completion_cost_stats = std::nullopt, - const std::optional& normal_dispatch_final_completion_sample_count_stats = std::nullopt, - const std::optional& normal_dispatch_final_completion_token_count_stats = std::nullopt, - const std::optional& normal_dispatch_rdma_recv_completion_cost_stats = std::nullopt, - const std::optional& normal_dispatch_rdma_recv_completion_sample_count_stats = std::nullopt, - const std::optional& normal_dispatch_rdma_recv_completion_token_count_stats = std::nullopt, - const std::optional& normal_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, - const std::optional& normal_notify_dispatch_full_kernel_count_stats = std::nullopt, - const std::optional& normal_cached_notify_dispatch_full_kernel_duration_ns_stats = std::nullopt, - const std::optional& normal_cached_notify_dispatch_full_kernel_count_stats = std::nullopt); + const NormalDispatchStats& normal_stats = {}); std::tuple, std::optional> internode_combine( const torch::Tensor& x, @@ -266,11 +280,7 @@ struct Buffer { std::optional& previous_event, bool async, bool allocate_on_comm_stream, - const std::optional& normal_cached_notify_combine_full_kernel_duration_ns_stats = std::nullopt, - const std::optional& normal_cached_notify_combine_full_kernel_count_stats = std::nullopt, - const std::optional& normal_combine_logical_recv_completion_cost_stats = std::nullopt, - const std::optional& normal_combine_logical_recv_completion_sample_count_stats = std::nullopt, - const std::optional& normal_combine_logical_recv_completion_token_count_stats = std::nullopt); + 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 5529f9986..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, @@ -174,9 +186,7 @@ void notify_dispatch(const int* num_tokens_per_rank, int64_t num_rdma_bytes, int64_t num_nvl_bytes, bool low_latency_mode, - int64_t* normal_notify_dispatch_full_kernel_duration_ns_stats, - int64_t* normal_notify_dispatch_full_kernel_count_stats, - int64_t* normal_notify_dispatch_full_kernel_timer_state); + NormalNotifyStats normal_notify_stats); void dispatch(void* recv_x, float* recv_x_scales, @@ -212,12 +222,8 @@ void dispatch(void* recv_x, int num_max_nvl_chunked_recv_tokens, // Completion cost tensors accumulate clock64() SM cycles; // Notify duration tensors accumulate %globaltimer nanoseconds. - int64_t* normal_dispatch_final_completion_cost_stats, - int64_t* normal_dispatch_final_completion_sample_count_stats, - int64_t* normal_dispatch_final_completion_token_count_stats, - int64_t* normal_dispatch_rdma_recv_completion_cost_stats, - int64_t* normal_dispatch_rdma_recv_completion_sample_count_stats, - int64_t* normal_dispatch_rdma_recv_completion_token_count_stats, + NormalCompletionStats normal_final_completion_stats, + NormalCompletionStats normal_rdma_recv_completion_stats, int rank, int num_ranks, bool is_cached_dispatch, @@ -247,9 +253,7 @@ void cached_notify(int hidden_int4, int64_t num_nvl_bytes, bool is_cached_dispatch, bool low_latency_mode, - int64_t* normal_cached_notify_full_kernel_duration_ns_stats, - int64_t* normal_cached_notify_full_kernel_count_stats, - int64_t* normal_cached_notify_full_kernel_timer_state); + NormalNotifyStats normal_notify_stats); void combine(cudaDataType_t type, void* combined_x, @@ -277,9 +281,7 @@ void combine(cudaDataType_t type, int num_max_nvl_chunked_recv_tokens, // Completion cost tensors accumulate clock64() SM cycles; // Notify duration tensors accumulate %globaltimer nanoseconds. - int64_t* normal_combine_logical_recv_completion_cost_stats, - int64_t* normal_combine_logical_recv_completion_sample_count_stats, - int64_t* normal_combine_logical_recv_completion_token_count_stats, + 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 5b8ef9562..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" @@ -20,44 +21,44 @@ __device__ __forceinline__ uint64_t read_globaltimer_ns() { return value; } -__device__ __forceinline__ void notify_full_kernel_timer_begin(int64_t* timer_state) { - if (timer_state == nullptr) +__device__ __forceinline__ void notify_full_kernel_timer_begin(NormalNotifyStats stats) { + if (stats.timer_state == nullptr) return; if (threadIdx.x == 0) { - auto state = reinterpret_cast(timer_state); + 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()); } - __syncthreads(); } -__device__ __forceinline__ void notify_full_kernel_timer_end(int64_t* timer_state, int64_t* duration_ns_stats, int64_t* count_stats) { - if (timer_state == nullptr) +__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(timer_state); + 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(duration_ns_stats), duration_ns); - atomicAdd(reinterpret_cast(count_stats), 1ull); + atomicAdd(reinterpret_cast(stats.duration_ns), duration_ns); + atomicAdd(reinterpret_cast(stats.count), 1ull); } } } -__device__ __forceinline__ void try_record_dispatch_rdma_recv_completion(int64_t* cost_stats, - int64_t* sample_count_stats, - int64_t* token_count_stats, +__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 (cost_stats == nullptr or expected_count <= 0 or 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)); @@ -69,9 +70,9 @@ __device__ __forceinline__ void try_record_dispatch_rdma_recv_completion(int64_t // 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(cost_stats + src_gateway_proxy_rank), clock64() - start_time); - atomicAdd(reinterpret_cast(sample_count_stats + src_gateway_proxy_rank), 1ull); - atomicAdd(reinterpret_cast(token_count_stats + src_gateway_proxy_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; } @@ -87,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; } @@ -178,9 +179,7 @@ __global__ void notify_dispatch(const int* num_tokens_per_rank, int** barrier_signal_ptrs, int rank, const nvshmem_team_t rdma_team, - int64_t* normal_notify_dispatch_full_kernel_duration_ns_stats, - int64_t* normal_notify_dispatch_full_kernel_count_stats, - int64_t* normal_notify_dispatch_full_kernel_timer_state) { + 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; @@ -188,7 +187,7 @@ __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_dispatch_full_kernel_timer_state); + notify_full_kernel_timer_begin(normal_notify_stats); if (sm_id == 0) { // Communication with others @@ -216,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]; @@ -280,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; @@ -289,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; @@ -299,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; @@ -314,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]; } @@ -327,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]; @@ -341,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; @@ -370,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); @@ -378,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; @@ -395,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]; } @@ -403,15 +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_dispatch_full_kernel_timer_state, - normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats); + notify_full_kernel_timer_end(normal_notify_stats); } void notify_dispatch(const int* num_tokens_per_rank, @@ -444,9 +441,7 @@ void notify_dispatch(const int* num_tokens_per_rank, int64_t num_rdma_bytes, int64_t num_nvl_bytes, bool low_latency_mode, - int64_t* normal_notify_dispatch_full_kernel_duration_ns_stats, - int64_t* normal_notify_dispatch_full_kernel_count_stats, - int64_t* normal_notify_dispatch_full_kernel_timer_state) { + NormalNotifyStats normal_notify_stats) { #define NOTIFY_DISPATCH_LAUNCH_CASE(num_rdma_ranks) \ { \ auto notify_dispatch_func = low_latency_mode ? notify_dispatch : notify_dispatch; \ @@ -478,9 +473,7 @@ void notify_dispatch(const int* num_tokens_per_rank, barrier_signal_ptrs, \ rank, \ cpu_rdma_team, \ - normal_notify_dispatch_full_kernel_duration_ns_stats, \ - normal_notify_dispatch_full_kernel_count_stats, \ - normal_notify_dispatch_full_kernel_timer_state); \ + normal_notify_stats); \ } \ break @@ -506,8 +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_dispatch_full_kernel_timer_state != nullptr) - CUDA_CHECK(cudaMemsetAsync(normal_notify_dispatch_full_kernel_timer_state, 0, 2 * sizeof(int64_t), 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 } @@ -556,12 +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, - int64_t* normal_dispatch_final_completion_cost_stats, - int64_t* normal_dispatch_final_completion_sample_count_stats, - int64_t* normal_dispatch_final_completion_token_count_stats, - int64_t* normal_dispatch_rdma_recv_completion_cost_stats, - int64_t* normal_dispatch_rdma_recv_completion_sample_count_stats, - int64_t* normal_dispatch_rdma_recv_completion_token_count_stats, + NormalCompletionStats normal_final_completion_stats, + NormalCompletionStats normal_rdma_recv_completion_stats, int rank, int num_ranks) { enum class WarpRole { kRDMASender, kRDMASenderCoordinator, kRDMAAndNVLForwarder, kForwarderCoordinator, kNVLReceivers }; @@ -573,7 +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_dispatch_final_completion_cost_stats != nullptr; + 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); @@ -752,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; @@ -768,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)); @@ -984,9 +973,7 @@ __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_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats, + try_record_dispatch_rdma_recv_completion(normal_rdma_recv_completion_stats, rdma_channel_tail.buffer(lane_id), lane_id, nvl_rank, @@ -1007,9 +994,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats, + try_record_dispatch_rdma_recv_completion(normal_rdma_recv_completion_stats, rdma_channel_tail.buffer(lane_id), lane_id, nvl_rank, @@ -1118,9 +1103,7 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV 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_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats, + try_record_dispatch_rdma_recv_completion(normal_rdma_recv_completion_stats, rdma_channel_tail.buffer(lane_id), lane_id, nvl_rank, @@ -1142,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) @@ -1153,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]); @@ -1315,14 +1298,11 @@ __global__ void __launch_bounds__(((kNumDispatchRDMASenderWarps + 1 + NUM_MAX_NV tma_store_wait<0>(); __syncwarp(); if (completion_stats_idx >= 0) { - atomicAdd(reinterpret_cast(normal_dispatch_final_completion_cost_stats + completion_stats_idx), + atomicAdd(reinterpret_cast(normal_final_completion_stats.cost + completion_stats_idx), clock64() - completion_start_time); - atomicAdd( - reinterpret_cast(normal_dispatch_final_completion_sample_count_stats + completion_stats_idx), - 1); - atomicAdd( - reinterpret_cast(normal_dispatch_final_completion_token_count_stats + completion_stats_idx), - static_cast(completion_expected_count)); + 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)); } } @@ -1343,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; } @@ -1381,12 +1361,8 @@ void dispatch(void* recv_x, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, - int64_t* normal_dispatch_final_completion_cost_stats, - int64_t* normal_dispatch_final_completion_sample_count_stats, - int64_t* normal_dispatch_final_completion_token_count_stats, - int64_t* normal_dispatch_rdma_recv_completion_cost_stats, - int64_t* normal_dispatch_rdma_recv_completion_sample_count_stats, - int64_t* normal_dispatch_rdma_recv_completion_token_count_stats, + NormalCompletionStats normal_final_completion_stats, + NormalCompletionStats normal_rdma_recv_completion_stats, int rank, int num_ranks, bool is_cached_dispatch, @@ -1399,14 +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_dispatch_final_completion_cost_stats != nullptr; - EP_HOST_ASSERT((normal_dispatch_final_completion_sample_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); - EP_HOST_ASSERT((normal_dispatch_final_completion_token_count_stats != nullptr) == enable_normal_dispatch_final_completion_stats); - const bool enable_normal_dispatch_rdma_recv_completion_stats = normal_dispatch_rdma_recv_completion_cost_stats != nullptr; - EP_HOST_ASSERT((normal_dispatch_rdma_recv_completion_sample_count_stats != nullptr) == - enable_normal_dispatch_rdma_recv_completion_stats); - EP_HOST_ASSERT((normal_dispatch_rdma_recv_completion_token_count_stats != nullptr) == - enable_normal_dispatch_rdma_recv_completion_stats); + 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) \ { \ @@ -1450,12 +1424,8 @@ void dispatch(void* recv_x, buffer_ptrs, \ num_max_nvl_chunked_send_tokens, \ num_max_nvl_chunked_recv_tokens, \ - normal_dispatch_final_completion_cost_stats, \ - normal_dispatch_final_completion_sample_count_stats, \ - normal_dispatch_final_completion_token_count_stats, \ - normal_dispatch_rdma_recv_completion_cost_stats, \ - normal_dispatch_rdma_recv_completion_sample_count_stats, \ - normal_dispatch_rdma_recv_completion_token_count_stats, \ + normal_final_completion_stats, \ + normal_rdma_recv_completion_stats, \ rank, \ num_ranks); \ } \ @@ -1487,9 +1457,7 @@ __global__ void cached_notify(const int rdma_clean_offset, int num_ranks, bool is_cached_dispatch, const nvshmem_team_t rdma_team, - int64_t* normal_cached_notify_full_kernel_duration_ns_stats, - int64_t* normal_cached_notify_full_kernel_count_stats, - int64_t* normal_cached_notify_full_kernel_timer_state) { + 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); @@ -1501,7 +1469,7 @@ __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_cached_notify_full_kernel_timer_state); + notify_full_kernel_timer_begin(normal_notify_stats); // Using two SMs, which clean the RDMA/NVL buffer respectively if (sm_id == 0) { @@ -1522,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(); @@ -1631,9 +1599,7 @@ __global__ void cached_notify(const int rdma_clean_offset, } } - notify_full_kernel_timer_end(normal_cached_notify_full_kernel_timer_state, - normal_cached_notify_full_kernel_duration_ns_stats, - normal_cached_notify_full_kernel_count_stats); + notify_full_kernel_timer_end(normal_notify_stats); } void cached_notify(int hidden_int4, @@ -1658,9 +1624,7 @@ void cached_notify(int hidden_int4, int64_t num_nvl_bytes, bool is_cached_dispatch, bool low_latency_mode, - int64_t* normal_cached_notify_full_kernel_duration_ns_stats, - int64_t* normal_cached_notify_full_kernel_count_stats, - int64_t* normal_cached_notify_full_kernel_timer_state) { + 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; @@ -1689,8 +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_cached_notify_full_kernel_timer_state != nullptr) - CUDA_CHECK(cudaMemsetAsync(normal_cached_notify_full_kernel_timer_state, 0, 2 * sizeof(int64_t), stream)); + 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, @@ -1710,9 +1674,7 @@ void cached_notify(int hidden_int4, num_ranks, is_cached_dispatch, cpu_rdma_team, - normal_cached_notify_full_kernel_duration_ns_stats, - normal_cached_notify_full_kernel_count_stats, - normal_cached_notify_full_kernel_timer_state); + 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]); } @@ -1806,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(); @@ -1820,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 @@ -1833,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)); @@ -1843,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]); } @@ -1860,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); @@ -1870,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); @@ -1888,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, @@ -1913,9 +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, - int64_t* normal_combine_logical_recv_completion_cost_stats, - int64_t* normal_combine_logical_recv_completion_sample_count_stats, - int64_t* normal_combine_logical_recv_completion_token_count_stats, + NormalCompletionStats normal_logical_recv_completion_stats, int rank, int num_ranks) { enum class WarpRole { kNVLSender, kNVLAndRDMAForwarder, kRDMAReceiver, kCoordinator }; @@ -2111,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 = @@ -2402,14 +2362,12 @@ __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_combine_logical_recv_completion_cost_stats != nullptr and - normal_combine_logical_recv_completion_sample_count_stats != nullptr and - normal_combine_logical_recv_completion_token_count_stats != nullptr; + 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 +#pragma unroll for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) logical_last_expected_head[i] = -1; @@ -2430,7 +2388,7 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co // Probe reconstruction must remain fail-open if queue metadata is incomplete. if (expected_head < 0) continue; - #pragma unroll +#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) { @@ -2459,21 +2417,20 @@ __global__ void __launch_bounds__((kNumForwarders + 1) * 32, 1) combine(int4* co 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 +#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_combine_logical_recv_completion_cost_stats + src_global_rank), - completion_end_time - logical_completion_start_time); - atomicAdd(reinterpret_cast(normal_combine_logical_recv_completion_sample_count_stats + - src_global_rank), - 1); - atomicAdd(reinterpret_cast(normal_combine_logical_recv_completion_token_count_stats + - src_global_rank), - static_cast(logical_expected_count[src_nvl_rank])); + 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; } } @@ -2488,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]); @@ -2502,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]); @@ -2546,9 +2503,7 @@ void combine(cudaDataType_t type, void** buffer_ptrs, int num_max_nvl_chunked_send_tokens, int num_max_nvl_chunked_recv_tokens, - int64_t* normal_combine_logical_recv_completion_cost_stats, - int64_t* normal_combine_logical_recv_completion_sample_count_stats, - int64_t* normal_combine_logical_recv_completion_token_count_stats, + NormalCompletionStats normal_logical_recv_completion_stats, int rank, int num_ranks, cudaStream_t stream, @@ -2600,9 +2555,7 @@ void combine(cudaDataType_t type, buffer_ptrs, \ num_max_nvl_chunked_send_tokens, \ num_max_nvl_chunked_recv_tokens, \ - normal_combine_logical_recv_completion_cost_stats, \ - normal_combine_logical_recv_completion_sample_count_stats, \ - normal_combine_logical_recv_completion_token_count_stats, \ + normal_logical_recv_completion_stats, \ rank, \ num_ranks); \ } \ diff --git a/deep_ep/buffer.py b/deep_ep/buffer.py index a7b27a6ce..0f551dabd 100644 --- a/deep_ep/buffer.py +++ b/deep_ep/buffer.py @@ -139,19 +139,6 @@ def all_gather_object(obj): self.runtime.sync(device_ids, ipc_handles, root_unique_id) assert self.runtime.is_available() - @staticmethod - def _unpack_normal_stats( - stats: Optional[Tuple[Optional[torch.Tensor], ...]], - expected_count: int, - argument_name: str) -> Tuple[Optional[torch.Tensor], ...]: - if stats is None: - return (None,) * expected_count - if len(stats) != expected_count: - raise ValueError( - f"`{argument_name}` must contain {expected_count} tensors, " - f"but got {len(stats)}") - return stats - @staticmethod def disable_ll_layered() -> bool: disable_ll_layered = False @@ -503,18 +490,6 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te # Launch the kernel with cached or non-cached mode x, x_scales = x if isinstance(x, tuple) else (x, None) - normal_notify_dispatch_full_kernel_duration_ns_stats, \ - normal_notify_dispatch_full_kernel_count_stats, \ - normal_cached_notify_dispatch_full_kernel_duration_ns_stats, \ - normal_cached_notify_dispatch_full_kernel_count_stats, \ - normal_dispatch_final_completion_cost_stats, \ - normal_dispatch_final_completion_sample_count_stats, \ - normal_dispatch_final_completion_token_count_stats, \ - normal_dispatch_rdma_recv_completion_cost_stats, \ - normal_dispatch_rdma_recv_completion_sample_count_stats, \ - normal_dispatch_rdma_recv_completion_token_count_stats = \ - self._unpack_normal_stats( - normal_dispatch_stats, 10, "normal_dispatch_stats") if handle is not None: assert topk_idx is None and topk_weights is None is_token_in_rank, \ @@ -527,13 +502,7 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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, normal_dispatch_final_completion_cost_stats, - normal_dispatch_final_completion_sample_count_stats, normal_dispatch_final_completion_token_count_stats, - normal_dispatch_rdma_recv_completion_cost_stats, normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats, normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats, - normal_cached_notify_dispatch_full_kernel_duration_ns_stats, normal_cached_notify_dispatch_full_kernel_count_stats, - ) + 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 @@ -546,16 +515,7 @@ def internode_dispatch(self, x: Union[torch.Tensor, Tuple[torch.Tensor, torch.Te 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, - normal_dispatch_final_completion_cost_stats, - normal_dispatch_final_completion_sample_count_stats, - normal_dispatch_final_completion_token_count_stats, - normal_dispatch_rdma_recv_completion_cost_stats, - normal_dispatch_rdma_recv_completion_sample_count_stats, - normal_dispatch_rdma_recv_completion_token_count_stats, - normal_notify_dispatch_full_kernel_duration_ns_stats, - normal_notify_dispatch_full_kernel_count_stats, - normal_cached_notify_dispatch_full_kernel_duration_ns_stats, - normal_cached_notify_dispatch_full_kernel_count_stats) + 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) @@ -585,22 +545,12 @@ def internode_combine(self, x: torch.Tensor, handle: Union[tuple, list], 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) - normal_cached_notify_combine_full_kernel_duration_ns_stats, \ - normal_cached_notify_combine_full_kernel_count_stats, \ - normal_combine_logical_recv_completion_cost_stats, \ - normal_combine_logical_recv_completion_sample_count_stats, \ - normal_combine_logical_recv_completion_token_count_stats = \ - self._unpack_normal_stats( - normal_combine_stats, 5, "normal_combine_stats") # 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, - normal_cached_notify_combine_full_kernel_duration_ns_stats, normal_cached_notify_combine_full_kernel_count_stats, - normal_combine_logical_recv_completion_cost_stats, normal_combine_logical_recv_completion_sample_count_stats, - normal_combine_logical_recv_completion_token_count_stats) + 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: