From 2578e50a33fd35094ae0cb7de6a779400c4e7c16 Mon Sep 17 00:00:00 2001 From: hyw Date: Wed, 30 Sep 2026 13:37:34 +0800 Subject: [PATCH 1/5] Add opt-in borrowed receive view for compact BF16 dispatch --- README.md | 10 ++ csrc/buffers/ep.hpp | 32 ++++- csrc/kernels/ep/dispatch.hpp | 12 +- deep_ep/__init__.py | 1 + deep_ep/buffers/ep.py | 73 +++++++--- deep_ep/buffers/recv_view.py | 57 ++++++++ .../impls/ep/dispatch_copy_epilogue.cuh | 56 +++++--- tests/ep/test_recv_view.py | 133 ++++++++++++++++++ 8 files changed, 323 insertions(+), 51 deletions(-) create mode 100644 deep_ep/buffers/recv_view.py create mode 100644 tests/ep/test_recv_view.py diff --git a/README.md b/README.md index 25a50df3..4fc015b3 100644 --- a/README.md +++ b/README.md @@ -276,6 +276,16 @@ The expanded `recv_topk_weights` is one-dimensional, with one value per expert r Training saves the forward `handle` for `combine_backward` and `dispatch_backward`. The backward helpers above return tensors and an event; wait on that event before using the tensors. Cached dispatch replays the saved expanded layout without another receive-count CPU synchronization. Keep the original `topk_idx` unchanged until all uses of the handle finish. +### Borrowed receive storage + +For a fresh compact BF16 dispatch within one physical NVLink domain, `borrow_recv=True` returns a `DispatchRecvView` in place of `recv_x`. Logical receive row `r` is stored at `view.slab[view.row_indices[r]]`. Pass the slab and row map directly to an indexed consumer to avoid the activation copy in the dispatch epilogue. This mode requires exact CPU receive counts and `defer_epilogue=False`. + +Wait on the returned dispatch event before accessing the view. After enqueueing all consumers, call `view.release()` on the consuming stream, or pass every consuming CUDA stream to `view.release(*streams)`. Release orders subsequent communication after those consumers. Dispatch, combine, load-balancing calls, and explicit buffer destruction reject an active view. Keep host API calls serialized and stop using raw slab references after release. + +Borrowed storage is intended for forward consumers; retaining it for backward is unsupported. After release, the returned handle can be used for ordinary materialized cached replay. Expanded layouts, FP8, cached borrowed dispatch, deferred epilogues, and multiple NVLink domains are outside this initial API. The option defaults to `False`; whether it reduces latency depends on the consumer and token count. + +See `tests/ep/test_recv_view.py` for event ordering, cross-stream lifetime, replay, combine, and layout checks. + ### Deferred epilogues Dispatch and combine accept `defer_epilogue=True` together with `async_with_compute_stream=True`. In this mode, the call returns an `EventOverlap` directly. Calling `.wait()` runs the deferred epilogue on the current stream and returns `(recv_x, recv_topk_idx, recv_topk_weights, handle)` for dispatch, or `(combined_x, combined_topk_weights)` for combine. diff --git a/csrc/buffers/ep.hpp b/csrc/buffers/ep.hpp index 04f4dbfd..b643f5d9 100644 --- a/csrc/buffers/ep.hpp +++ b/csrc/buffers/ep.hpp @@ -307,7 +307,10 @@ class EPBuffer: public BufferBase { const bool& do_cpu_sync, const bool& do_expand, const bool& do_zero_padding, const bool& use_tma_aligned_col_major_sf, - const bool& defer_epilogue) const { + const bool& defer_epilogue, const bool& materialize_recv_x) const { + EP_HOST_ASSERT(materialize_recv_x or (context->num_scaleout_ranks == 1 and context->num_rdma_ranks == 1 and not do_expand and + not sf.has_value() and x.scalar_type() == torch::kBFloat16 and + do_cpu_sync and not cached_num_recv_tokens.has_value() and not defer_epilogue)); // Check SM count EP_HOST_ASSERT(num_sms > 0 and num_sms <= jit->device.get_num_sms()); EP_HOST_ASSERT((num_sms > 1 or context->num_scaleout_ranks == 1 or context->num_scaleup_ranks == 1) and @@ -671,7 +674,11 @@ class EPBuffer: public BufferBase { // Allocate received tensors // `recv_src_metadata` includes source token indices and buffer slot indices const auto num_allocated_tokens = do_expand ? num_expanded_tokens : num_recv_tokens; - auto recv_x = torch::empty({num_allocated_tokens, hidden}, x.options()); + std::optional recv_x, recv_row_indices; + if (materialize_recv_x) + recv_x = torch::empty({num_allocated_tokens, hidden}, x.options()); + else + recv_row_indices = torch::empty({num_recv_tokens}, x.options().dtype(torch::kInt64)); auto recv_sf = std::optional(); auto recv_topk_idx = std::optional(); auto recv_topk_weights = std::optional(); @@ -724,9 +731,10 @@ class EPBuffer: public BufferBase { launch_dispatch_copy_epilogue(context->buffer, context->workspace, psum_num_recv_tokens_per_scaleup_rank.data_ptr(), psum_num_recv_tokens_per_expert.data_ptr(), - recv_x.data_ptr(), recv_sf_ptr, + recv_x.has_value() ? recv_x->data_ptr() : nullptr, recv_sf_ptr, recv_topk_idx_ptr, recv_topk_weights_ptr, recv_src_metadata.data_ptr(), + recv_row_indices.has_value() ? recv_row_indices->data_ptr() : nullptr, channel_linked_list_ptr, num_unaligned_recv_tokens_per_expert_ptr, num_recv_tokens, num_max_tokens_per_rank, @@ -739,7 +747,7 @@ class EPBuffer: public BufferBase { jit->device.get_num_smem_bytes(), num_channels, do_expand, cached_mode, - do_zero_padding, + do_zero_padding, materialize_recv_x, stream); auto result = pybind11::make_tuple( @@ -753,12 +761,13 @@ class EPBuffer: public BufferBase { recv_src_metadata, dst_buffer_slot_idx, token_metadata_at_forward, - channel_linked_list); + channel_linked_list, recv_row_indices); // For non-deferring tensor recording if (tensors_to_record_opt.has_value()) { auto& tensors = tensors_to_record_opt->get(); tensors.push_back(recv_x); + tensors.push_back(recv_row_indices); tensors.push_back(recv_sf); tensors.push_back(recv_topk_idx); tensors.push_back(recv_topk_weights); @@ -786,6 +795,18 @@ class EPBuffer: public BufferBase { return pybind11::make_tuple(result, event, pybind11::none()); } + torch::Tensor get_dispatch_recv_slab(const torch::Tensor& x, const int& num_topk, + const int& num_max_tokens_per_rank) const { + EP_HOST_ASSERT(not destroyed and context->num_scaleout_ranks == 1 and context->num_rdma_ranks == 1); + EP_HOST_ASSERT(x.is_cuda() and x.scalar_type() == torch::kBFloat16 and x.dim() == 2); + const auto token_layout = layout::TokenLayout(x.size(1) * x.element_size(), 0, num_topk, true); + const auto slots = static_cast(context->num_scaleup_ranks) * num_max_tokens_per_rank; + const auto stride_bytes = token_layout.get_num_bytes(); + EP_HOST_ASSERT(slots * stride_bytes <= context->num_gpu_buffer_bytes); + return torch::from_blob(context->buffer, {slots, x.size(1)}, + {stride_bytes / static_cast(x.element_size()), 1}, x.options()); + } + pybind11::tuple combine(const torch::Tensor& x, const std::optional& topk_weights, @@ -1087,6 +1108,7 @@ static void register_apis(pybind11::module_& m) { .def_readonly("context", &EPBuffer::context) .def_readonly("lb_storage", &EPBuffer::lb_storage) .def("dispatch", &EPBuffer::dispatch) + .def("get_dispatch_recv_slab", &EPBuffer::get_dispatch_recv_slab) .def("combine", &EPBuffer::combine) .def("lb_prefetch_weights", &EPBuffer::lb_prefetch_weights) .def("lb_reduce_grads", &EPBuffer::lb_reduce_grads); diff --git a/csrc/kernels/ep/dispatch.hpp b/csrc/kernels/ep/dispatch.hpp index 953af9f7..50c5bf3f 100644 --- a/csrc/kernels/ep/dispatch.hpp +++ b/csrc/kernels/ep/dispatch.hpp @@ -170,7 +170,7 @@ static void launch_dispatch_copy_epilogue(void* buffer, void* workspace, int* psum_num_recv_tokens_per_expert, void* recv_x, void* recv_sf, topk_idx_t* recv_topk_idx, float* recv_topk_weights, - int* recv_src_metadata, + int* recv_src_metadata, int64_t* recv_row_indices, int* channel_linked_list, int* num_unaligned_recv_tokens_per_expert, const int& num_recv_tokens, const int& num_max_tokens_per_rank, @@ -182,7 +182,7 @@ static void launch_dispatch_copy_epilogue(void* buffer, void* workspace, const int& num_sms, const int& num_smem_bytes, const int& num_channels, const bool& do_expand, const bool& cached_mode, - const bool& do_zero_padding, + const bool& do_zero_padding, const bool& materialize_recv_x, const at::cuda::CUDAStream& stream) { // Maximize shared memory utilization const auto token_layout = layout::TokenLayout(num_hidden_bytes, num_sf_packs * sizeof(sf_pack_t), num_topk, true); @@ -194,9 +194,9 @@ static void launch_dispatch_copy_epilogue(void* buffer, void* workspace, #include static void __instantiate_kernel() {{ - auto ptr = reinterpret_cast(&deep_ep::ep::dispatch_copy_epilogue_impl<{}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}>); + auto ptr = reinterpret_cast(&deep_ep::ep::dispatch_copy_epilogue_impl<{}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}>); }} -)", do_expand, cached_mode, do_zero_padding, +)", do_expand, cached_mode, do_zero_padding, materialize_recv_x, num_sms, num_channels, num_warps, num_scaleout_ranks, num_scaleup_ranks, num_hidden_bytes, num_sf_packs, @@ -207,7 +207,7 @@ static void __instantiate_kernel() {{ jit->launch( kernel, { .stream = stream.stream(), - .num_smem_bytes = num_smem_bytes, + .num_smem_bytes = materialize_recv_x ? num_smem_bytes : 0, .grid_dim = dim3(num_sms, 1, 1), .block_dim = dim3(num_threads, 1, 1), .enable_pdl = true, @@ -216,7 +216,7 @@ static void __instantiate_kernel() {{ psum_num_recv_tokens_per_scaleup_rank, psum_num_recv_tokens_per_expert, recv_x, recv_sf, recv_topk_idx, recv_topk_weights, - recv_src_metadata, + recv_src_metadata, recv_row_indices, channel_linked_list, num_unaligned_recv_tokens_per_expert, num_recv_tokens, diff --git a/deep_ep/__init__.py b/deep_ep/__init__.py index 2cb581ac..83597c74 100644 --- a/deep_ep/__init__.py +++ b/deep_ep/__init__.py @@ -62,6 +62,7 @@ def init_jit(): from .buffers.allocator import BufferAllocator from .buffers.base import BufferBase from .buffers.ep import EPBuffer, EPHandle +from .buffers.recv_view import DispatchRecvView from .buffers.engram import EngramBuffer from .buffers.bucket import BucketBuffer, BucketSession from .buffers.pp import PPBuffer diff --git a/deep_ep/buffers/ep.py b/deep_ep/buffers/ep.py index 9b8a9536..944e9271 100644 --- a/deep_ep/buffers/ep.py +++ b/deep_ep/buffers/ep.py @@ -1,4 +1,3 @@ -import functools import os import math import torch @@ -12,6 +11,7 @@ from .allocator import BufferAllocator from .base import BufferBase +from .recv_view import DispatchRecvView from .. import comm from ..utils.event import EventOverlap from ..utils.math import align @@ -113,11 +113,12 @@ def topk_idx(self) -> torch.Tensor: def deterministic_sort(self, do_cpu_sync: bool, is_cached_dispatch: bool, - recv_x: torch.Tensor, + recv_x: Optional[torch.Tensor], recv_sf: Optional[torch.Tensor], recv_topk_idx: torch.Tensor, recv_topk_weights: torch.Tensor, - channel_linked_list: Optional[torch.Tensor]): + channel_linked_list: Optional[torch.Tensor], + recv_row_indices: Optional[torch.Tensor] = None): """ Sort received tokens to guarantee deterministic dispatch output. The principle: @@ -160,6 +161,7 @@ def permute(tensor: Optional[torch.Tensor], orig_indices: torch.Tensor): # Non-expand mode # If cached dispatch is enabled, the `dispatch` kernel stores values according to `dst_buffer_slot_idx`, and the `dispatch_copy_epilogue_impl` kernel writes the info of token i into the i-th slot permute(recv_x, orig_indices) + permute(recv_row_indices, orig_indices) permute(recv_sf, orig_indices) permute(recv_topk_weights, orig_indices) permute(recv_topk_idx, orig_indices) @@ -294,7 +296,7 @@ def __init__(self, self.allow_multiple_reduction = allow_multiple_reduction self.prefer_overlap_with_compute = prefer_overlap_with_compute self.deterministic = deterministic - + # Create NCCL comm handle self.nccl_comm_handle = comm.get_nccl_comm_handle(group) @@ -362,10 +364,15 @@ def __init__(self, group.barrier() torch.cuda.synchronize() + def _check_recv_view_released(self) -> None: + if getattr(self, '_active_recv_view', None) is not None: + raise RuntimeError('Release the dispatch receive view before reusing or destroying this buffer') + def destroy(self) -> None: """ Destroy the C++ runtime and release resources. Requires `explicitly_destroy=True` at construction. """ + self._check_recv_view_released() super().destroy() self.context = None self.nccl_comm_handle = None @@ -577,8 +584,9 @@ def dispatch(self, do_expand: bool = False, do_zero_padding: bool = False, use_tma_aligned_col_major_sf: bool = False, - defer_epilogue: bool = False) \ - -> Union[Tuple[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], + defer_epilogue: bool = False, + borrow_recv: bool = False) \ + -> Union[Tuple[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor], DispatchRecvView], Optional[torch.Tensor], Optional[torch.Tensor], EPHandle, EventOverlap], EventOverlap]: """ @@ -620,11 +628,17 @@ def dispatch(self, do_zero_padding: whether to zero out the alignment padding slots in the expanded output. Only valid when `do_expand` is True. Ensures alignment gaps between experts are zeroed. use_tma_aligned_col_major_sf: whether to use TMA-aligned column-major layout for scale factors. + borrow_recv: return a DispatchRecvView instead of a materialized receive tensor. + Requires fresh compact BF16 dispatch in one NVLink domain, exact CPU counts, and + defer_epilogue=False. Wait the dispatch event, enqueue indexed consumers, then + release the view with every consuming stream before reusing or destroying this buffer. + The view cannot be retained for backward; a cached materialized replay is supported. defer_epilogue: whether to defer the CPU receive-count wait and copy epilogue until `event.current_stream_wait()` is called. This requires `async_with_compute_stream=True`. Returns: recv_x: received tokens, the same type and tuple as the input `x`. + With `borrow_recv=True`, a DispatchRecvView owning a lease on the receive buffer. Only returned when `defer_epilogue=False`. recv_topk_idx: received expert indices. Only returned when `defer_epilogue=False`. recv_topk_weights: received expert weights (`None` if `topk_weights` was not provided). @@ -635,6 +649,11 @@ def dispatch(self, of the five-item tuple. Call `event.current_stream_wait()` to run the copy epilogue and obtain `(recv_x, recv_topk_idx, recv_topk_weights, handle)`. """ + self._check_recv_view_released() + if borrow_recv and (handle is not None or do_expand or defer_epilogue or do_cpu_sync is False or + not isinstance(x, torch.Tensor) or x.dtype != torch.bfloat16 or + self.num_scaleout_ranks != 1 or self.num_rdma_ranks != 1): + raise ValueError('borrow_recv requires fresh compact BF16 dispatch, CPU counts, one NVLink domain, and no deferred epilogue') assert not do_handle_copy, '`do_handle_copy` must be False; handle copying is no longer supported' check_torch_deterministic() @@ -642,7 +661,7 @@ def dispatch(self, num_topk = (handle.topk_idx if topk_idx is None else topk_idx).shape[1] num_sms = self.get_theoretical_num_sms(num_experts, num_topk) if num_sms == 0 else align(num_sms, 2) num_qps = self.get_theoretical_num_qps(num_sms) if num_qps == 0 else num_qps - assert num_qps <= self.num_allocated_qps, f'Allocated QPs are not enough' + assert num_qps <= self.num_allocated_qps, 'Allocated QPs are not enough' # Unpack SF x, sf = x if isinstance(x, tuple) else (x, None) @@ -696,7 +715,7 @@ def dispatch(self, do_cpu_sync, do_expand, do_zero_padding, use_tma_aligned_col_major_sf, - defer_epilogue) + defer_epilogue, not borrow_recv) event_overlap = EventOverlap(event) def finalize_dispatch(dispatch_result: tuple, deterministic_by_hook: bool): @@ -710,7 +729,7 @@ def finalize_dispatch(dispatch_result: tuple, deterministic_by_hook: bool): recv_src_metadata, dst_buffer_slot_idx, token_metadata_at_forward, - channel_linked_list) = dispatch_result + channel_linked_list, recv_row_indices) = dispatch_result # Create handle if not cached nonlocal handle @@ -730,21 +749,28 @@ def finalize_dispatch(dispatch_result: tuple, deterministic_by_hook: bool): token_metadata_at_forward, channel_linked_list) if handle is None else handle - # Do deterministic - if self.deterministic: - deterministic_epilogue = functools.partial( - handle.deterministic_sort, - do_cpu_sync, is_cached_dispatch, - recv_x, recv_sf, recv_topk_idx, recv_topk_weights, channel_linked_list - ) + recv_view = None + if borrow_recv: + slab = self.runtime.get_dispatch_recv_slab(x, num_topk, num_max_tokens_per_rank) + recv_view = DispatchRecvView(self, slab, recv_row_indices) + self._active_recv_view = recv_view + + def complete_dispatch(): + if self.deterministic: + handle.deterministic_sort( + do_cpu_sync, is_cached_dispatch, recv_x, recv_sf, + recv_topk_idx, recv_topk_weights, channel_linked_list, recv_row_indices) + if recv_view is not None: + recv_view._mark_ready() + + if self.deterministic or recv_view is not None: if deterministic_by_hook: - event_overlap.register_hook_after_wait(deterministic_epilogue) + event_overlap.register_hook_after_wait(complete_dispatch) else: - deterministic_epilogue() + complete_dispatch() - # Return values - recv_x = (recv_x, recv_sf) if recv_sf is not None else recv_x - return recv_x, recv_topk_idx, recv_topk_weights, handle + received = recv_view if recv_view is not None else ((recv_x, recv_sf) if recv_sf is not None else recv_x) + return received, recv_topk_idx, recv_topk_weights, handle # Just launch the dispatch if deferred_epilogue is not None: @@ -810,12 +836,13 @@ def combine(self, of the three-item tuple. Call `event.current_stream_wait()` to run the reduce epilogue and obtain `(combined_x, combined_topk_weights)`. """ + self._check_recv_view_released() check_torch_deterministic() # Automatic decide SM and QP count num_sms = handle.num_sms if num_sms == 0 else align(num_sms, 2) num_qps = self.get_theoretical_num_qps(num_sms) if num_qps == 0 else num_qps - assert num_qps <= self.num_allocated_qps, f'Allocated QPs are not enough' + assert num_qps <= self.num_allocated_qps, 'Allocated QPs are not enough' bias_0, bias_1 = EPBuffer._unpack_bias(bias) result, event, deferred_epilogue = self.runtime.combine(x, topk_weights, @@ -875,6 +902,7 @@ def lb_prefetch_weights(self, num_sms: the number of SMs to use; 0 uses `lb_get_theoretical_num_sms()`. previous_event: the event to wait for before communication; defaults to waiting for the current stream """ + self._check_recv_view_released() redundant_expert_weights = ([redundant_expert_weights] if isinstance(redundant_expert_weights, torch.Tensor) else list(redundant_expert_weights)) expert_weights = [expert_weights] if isinstance(expert_weights, torch.Tensor) else list(expert_weights) @@ -901,6 +929,7 @@ def lb_reduce_grads(self, num_sms: the number of SMs to use; 0 uses `lb_get_theoretical_num_sms()`. previous_event: the event to wait for before communication; defaults to waiting for the current stream """ + self._check_recv_view_released() num_sms = self.lb_get_theoretical_num_sms() if num_sms == 0 else align(num_sms, 2) return EventOverlap(self.runtime.lb_reduce_grads( redundant_expert_grads, expert_grads, redundancy_mapping, num_sms, previous_event)) diff --git a/deep_ep/buffers/recv_view.py b/deep_ep/buffers/recv_view.py new file mode 100644 index 00000000..f9c45749 --- /dev/null +++ b/deep_ep/buffers/recv_view.py @@ -0,0 +1,57 @@ +"""An exclusive, forward-only view of compact BF16 dispatch payloads.""" +import torch + +from .. import comm + + +class DispatchRecvView: + """Read logical row ``r`` as ``slab[row_indices[r]]`` after dispatch wait. + + The payload belongs to the communication buffer. Keep host calls serialized + and release this view after enqueueing all consumers, passing every consuming + CUDA stream. Copies of ``slab`` references are invalid after release. Saving + this view for backward is unsupported; use a materialized cached replay. + """ + + def __init__(self, owner, slab, row_indices): + self._owner = owner + self._slab = slab + self._row_indices = row_indices + self._ready_event = None + self._released = False + + def _mark_ready(self): + self._ready_event = torch.cuda.Event() + self._ready_event.record(torch.cuda.current_stream(self._slab.device)) + + def wait(self): + if self._released: + raise RuntimeError('The dispatch receive view has been released') + if self._ready_event is None: + raise RuntimeError('Wait the dispatch event before consuming the receive view') + torch.cuda.current_stream(self._slab.device).wait_event(self._ready_event) + + @property + def slab(self): + self.wait() + return self._slab + + @property + def row_indices(self): + self.wait() + return self._row_indices + + def release(self, *consumer_streams): + self.wait() + if not consumer_streams: + consumer_streams = (torch.cuda.current_stream(self._slab.device),) + comm_stream = comm.get_comm_stream(self._owner) + for stream in consumer_streams: + if stream.device != self._slab.device: + raise ValueError('Consumer streams must belong to the receive-view device') + stream.wait_event(self._ready_event) + self._row_indices.record_stream(stream) + comm_stream.wait_stream(stream) + self._owner._active_recv_view = None + self._released = True + self._owner = self._slab = self._row_indices = self._ready_event = None diff --git a/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh b/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh index b9c76ec8..d1fc818b 100644 --- a/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh +++ b/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh @@ -9,7 +9,7 @@ namespace deep_ep::ep { -template (blockIdx.x), thread_idx = static_cast(threadIdx.x); const auto warp_idx = ptx::get_warp_idx(), lane_idx = ptx::get_lane_idx(); @@ -52,8 +54,10 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // Init TMA ptx::arrival_phase phase = 0; const auto mbarrier_ptr = tma_buffer.get_mbarrier_ptr(); - if (ptx::elect_one_sync()) - ptx::mbarrier_init_with_fence(mbarrier_ptr, 1); + if constexpr (kMaterializeRecvX) { + if (ptx::elect_one_sync()) + ptx::mbarrier_init_with_fence(mbarrier_ptr, 1); + } __syncwarp(); // Will block until the main dispatch kernel has finished and all data are visible @@ -80,17 +84,25 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, current_rank_end = ptx::exchange(stored_psum_num_recv_tokens, stored_lane_idx); } const auto buffer_token = scaleup_buffer.get_rank_buffer(current_rank_idx).get_token_buffer(i - current_rank_start); + if constexpr (not kMaterializeRecvX) { + if (ptx::elect_one_sync()) + recv_row_indices[i] = static_cast(current_rank_idx) * kNumMaxTokensPerRank + i - current_rank_start; + } + // Wait buffer releases - ptx::tma_store_wait(); + if constexpr (kMaterializeRecvX) + ptx::tma_store_wait(); __syncwarp(); // Issue TMA loads // Including all stuffs: data, SF, top-k metadata - if (ptx::elect_one_sync()) { - ptx::tma_load_1d(tma_buffer.get_base_ptr(), buffer_token.get_base_ptr(), - mbarrier_ptr, tma_buffer.get_num_bytes()); - ptx::mbarrier_arrive_and_set_tx(mbarrier_ptr, tma_buffer.get_num_bytes()); + if constexpr (kMaterializeRecvX) { + if (ptx::elect_one_sync()) { + ptx::tma_load_1d(tma_buffer.get_base_ptr(), buffer_token.get_base_ptr(), + mbarrier_ptr, tma_buffer.get_num_bytes()); + ptx::mbarrier_arrive_and_set_tx(mbarrier_ptr, tma_buffer.get_num_bytes()); + } } __syncwarp(); @@ -124,10 +136,16 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, __syncwarp(); // Wait for TMA arrival - if (ptx::elect_one_sync()) - ptx::mbarrier_wait_and_flip_phase(mbarrier_ptr, phase); + if constexpr (kMaterializeRecvX) { + if (ptx::elect_one_sync()) + ptx::mbarrier_wait_and_flip_phase(mbarrier_ptr, phase); + } __syncwarp(); + // Metadata-only mode must skip the activation LOAD as well as its + // store. Read the small metadata fields directly after the PDL fence. + const auto metadata_token = kMaterializeRecvX ? tma_buffer : buffer_token; + // Maintain linked list if constexpr (kDoCreateLinkedList) { if (ptx::elect_one_sync()) @@ -136,10 +154,12 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, } // Issue TMA stores for data - if (kDoExpand ? (dst_tensor_idx >= 0) : ptx::elect_one_sync()) { - ptx::tma_store_1d(math::advance_ptr(recv_x, static_cast(dst_tensor_idx) * kNumHiddenBytes), - tma_buffer.get_hidden_ptr(), kNumHiddenBytes); - ptx::tma_store_commit(); + if constexpr (kMaterializeRecvX) { + if (kDoExpand ? (dst_tensor_idx >= 0) : ptx::elect_one_sync()) { + ptx::tma_store_1d(math::advance_ptr(recv_x, static_cast(dst_tensor_idx) * kNumHiddenBytes), + tma_buffer.get_hidden_ptr(), kNumHiddenBytes); + ptx::tma_store_commit(); + } } __syncwarp(); @@ -179,10 +199,10 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // Store the top-k weights if (kDoExpand and recv_topk_weights != nullptr and dst_tensor_idx >= 0) { - recv_topk_weights[dst_tensor_idx] = tma_buffer.get_topk_weights_ptr()[lane_idx]; + recv_topk_weights[dst_tensor_idx] = metadata_token.get_topk_weights_ptr()[lane_idx]; } else if (not kDoExpand and recv_topk_weights != nullptr and lane_idx < kNumTopk) { // For backward, weights are optional - recv_topk_weights[i * kNumTopk + lane_idx] = tma_buffer.get_topk_weights_ptr()[lane_idx]; + recv_topk_weights[i * kNumTopk + lane_idx] = metadata_token.get_topk_weights_ptr()[lane_idx]; } __syncwarp(); @@ -192,7 +212,7 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // - Hybrid mode: the slot index and master top-k lane index if constexpr (not kCachedMode) { if (ptx::elect_one_sync()) { - recv_src_metadata[i * kMetadataStride + 0] = *tma_buffer.get_src_token_global_idx_ptr(); + recv_src_metadata[i * kMetadataStride + 0] = *metadata_token.get_src_token_global_idx_ptr(); if constexpr (kNumScaleoutRanks == 1) { recv_src_metadata[i * kMetadataStride + 1] = current_rank_idx * kNumTopk + master_src_topk_idx; } else { diff --git a/tests/ep/test_recv_view.py b/tests/ep/test_recv_view.py new file mode 100644 index 00000000..ea7ae1e1 --- /dev/null +++ b/tests/ep/test_recv_view.py @@ -0,0 +1,133 @@ +"""Run with torchrun --standalone --nproc-per-node=8 tests/ep/test_recv_view.py.""" +import os + +import torch +import torch.distributed as dist + + +def exact(actual, expected): + assert actual.shape == expected.shape and actual.dtype == expected.dtype + assert torch.equal(actual.contiguous().view(torch.uint8), expected.contiguous().view(torch.uint8)) + + +def rejected(fn, error): + try: + fn() + except error: + return + raise AssertionError(f'Expected {error.__name__}') + + +def test_case(deep_ep, deterministic, asynchronous, with_weights, skew): + rank, world = dist.get_rank(), dist.get_world_size() + torch.manual_seed(42 + rank) + capacity, hidden, topk, local_experts = 128, 256, 2, 2 + tokens = 65 - rank + full_x = torch.randn((capacity, hidden), device='cuda', dtype=torch.bfloat16) + x = full_x[:tokens] + indices = torch.rand((tokens, local_experts * world), device='cuda').argsort(dim=1)[:, :topk] + indices = indices.to(deep_ep.topk_idx_t).contiguous() + if skew: + indices = torch.arange(topk, device='cuda', dtype=deep_ep.topk_idx_t).repeat(tokens, 1) + indices[::7] = -1 + weights = torch.randn((tokens, topk), device='cuda', dtype=torch.float32) if with_weights else None + gathered = [torch.empty_like(full_x) for _ in range(world)] + dist.all_gather(gathered, full_x) + all_x = torch.stack(gathered) + buffer = deep_ep.EPBuffer(dist.group.WORLD, num_max_tokens_per_rank=capacity, + hidden=hidden, num_topk=topk, deterministic=deterministic, + explicitly_destroy=True) + kwargs = dict(topk_idx=indices, topk_weights=weights, num_experts=local_experts * world, + num_sms=8, async_with_compute_stream=asynchronous, + allocate_on_comm_stream=asynchronous) + + def dispatch(borrow, value=x): + result = buffer.dispatch(value, borrow_recv=borrow, **kwargs) + if asynchronous: + result[-1].current_stream_wait() + return result[:4] + + materialized, ref_ids, ref_weights, ref_handle = dispatch(False) + view, recv_ids, recv_weights, handle = dispatch(True) + assert isinstance(view, deep_ep.DispatchRecvView) + payload = view.slab.index_select(0, view.row_indices) + source = handle.recv_src_metadata[:, 0].long() + exact(payload, all_x[source // capacity, source % capacity]) + order = source.argsort() + ref_order = ref_handle.recv_src_metadata[:, 0].argsort() + exact(payload[order], materialized[ref_order]) + exact(recv_ids[order], ref_ids[ref_order]) + if with_weights: + exact(recv_weights[order], ref_weights[ref_order]) + else: + assert recv_weights is None + rejected(lambda: buffer.dispatch(x, **kwargs), RuntimeError) + rejected(lambda: buffer.combine(payload, handle), RuntimeError) + rejected(buffer.destroy, RuntimeError) + view.release() + rejected(lambda: view.slab, RuntimeError) + rejected(view.release, RuntimeError) + + replay, _, _, _, event = buffer.dispatch(x, handle=handle, num_sms=8, + async_with_compute_stream=True) + event.current_stream_wait() + exact(replay, payload) + # Integer sums are exact even with an unspecified receive order. + combined, _, event = buffer.combine(torch.ones_like(replay), handle, + async_with_compute_stream=True) + event.current_stream_wait() + expected_count = torch.stack([ + ((indices >= 0) & (indices // local_experts == peer)).any(dim=1) + for peer in range(world) + ]).sum(dim=0).to(x.dtype) + exact(combined, expected_count[:, None].expand_as(x)) + + # Consumers can move to another stream after the dispatch event is waited. + view, _, _, delayed_handle = dispatch(True) + delayed_source = delayed_handle.recv_src_metadata[:, 0].long() + delayed_expected = all_x[delayed_source // capacity, delayed_source % capacity] + side = torch.cuda.Stream() + with torch.cuda.stream(side): + torch.cuda._sleep(4_000_000) + delayed = view.slab.index_select(0, view.row_indices) + view.release(side) + # Overwrite the receive buffer before synchronizing the delayed consumer. + dispatch(False, x + 1) + side.synchronize() + exact(delayed, delayed_expected) + + for extra in [dict(do_expand=True), dict(defer_epilogue=True), dict(do_cpu_sync=False), dict(handle=handle)]: + rejected(lambda extra=extra: buffer.dispatch(x, borrow_recv=True, **(kwargs | extra)), ValueError) + rejected(lambda: buffer.dispatch(x.float(), borrow_recv=True, **kwargs), ValueError) + # One logical scaleup domain is insufficient when hybrid mode is disabled. + original_rdma_ranks = buffer.num_rdma_ranks + try: + buffer.num_rdma_ranks = 2 + rejected(lambda: buffer.dispatch(x, borrow_recv=True, **kwargs), ValueError) + finally: + buffer.num_rdma_ranks = original_rdma_ranks + dist.barrier() + buffer.destroy() + + +def main(): + rank = int(os.environ['LOCAL_RANK']) + torch.cuda.set_device(rank) + torch.use_deterministic_algorithms(True) + torch.utils.deterministic.fill_uninitialized_memory = False + import deep_ep + + dist.init_process_group('nccl', device_id=torch.device('cuda', rank)) + for deterministic in [False, True]: + for asynchronous in [False, True]: + for with_weights in [False, True]: + for skew in [False, True]: + test_case(deep_ep, deterministic, asynchronous, with_weights, skew) + if rank == 0: + print(f'PASS {deterministic=} {asynchronous=} {with_weights=} {skew=}', flush=True) + deep_ep.destroy_all_managed_nccl_comm() + dist.destroy_process_group() + + +if __name__ == '__main__': + main() From 59abc161bb70509994375e96e3efa6b0e0748e36 Mon Sep 17 00:00:00 2001 From: hyw Date: Wed, 30 Sep 2026 13:37:34 +0800 Subject: [PATCH 2/5] Apply upstream format and lint rules to changed files --- csrc/buffers/ep.hpp | 671 ++++++++++-------- csrc/kernels/ep/dispatch.hpp | 296 +++++--- deep_ep/__init__.py | 30 +- deep_ep/buffers/ep.py | 247 +++---- deep_ep/buffers/recv_view.py | 2 +- .../impls/ep/dispatch_copy_epilogue.cuh | 104 +-- tests/ep/test_recv_view.py | 26 +- 7 files changed, 717 insertions(+), 659 deletions(-) diff --git a/csrc/buffers/ep.hpp b/csrc/buffers/ep.hpp index b643f5d9..5cacea44 100644 --- a/csrc/buffers/ep.hpp +++ b/csrc/buffers/ep.hpp @@ -2,25 +2,25 @@ #include #include -#include -#include #include #include +#include #include #include -#include +#include +#include -#include "base.hpp" #include "../comm/api.hpp" #include "../kernels/ep/api.hpp" #include "../runtime/jit.hpp" #include "../utils/event.hpp" #include "../utils/tensor.hpp" +#include "base.hpp" namespace deep_ep::ep { -class EPBuffer: public BufferBase { +class EPBuffer : public BufferBase { // Buffer bytes exclude workspace // Memory layout: [Workspace, buffer] int64_t num_buffer_bytes; @@ -48,42 +48,55 @@ class EPBuffer: public BufferBase { // For load balance torch::Tensor lb_storage; - EPBuffer(const int& rank_idx, const int& num_ranks, + EPBuffer(const int& rank_idx, + const int& num_ranks, const int64_t& nccl_comm, const int64_t& num_buffer_bytes, const int64_t& num_lb_buffer_bytes, const bool& allow_hybrid_mode, const bool& allow_multiple_reduction, const bool& prefer_overlap_with_compute, - const std::optional& sl_idx, const int& num_allocated_qps, - const int& num_cpu_timeout_secs, const int& num_gpu_timeout_secs, - const bool& explicitly_destroy): - BufferBase(explicitly_destroy), - num_buffer_bytes(num_buffer_bytes), - allow_hybrid_mode(allow_hybrid_mode), - allow_multiple_reduction(allow_multiple_reduction), - prefer_overlap_with_compute(prefer_overlap_with_compute) { + const std::optional& sl_idx, + const int& num_allocated_qps, + const int& num_cpu_timeout_secs, + const int& num_gpu_timeout_secs, + const bool& explicitly_destroy) + : BufferBase(explicitly_destroy), + num_buffer_bytes(num_buffer_bytes), + allow_hybrid_mode(allow_hybrid_mode), + allow_multiple_reduction(allow_multiple_reduction), + prefer_overlap_with_compute(prefer_overlap_with_compute) { EP_HOST_ASSERT(num_buffer_bytes > 0 and num_buffer_bytes % kNumAllocationAlignmentBytes == 0); EP_HOST_ASSERT(num_lb_buffer_bytes >= 0 and num_lb_buffer_bytes % kNumAllocationAlignmentBytes == 0); // Workspace is aligned to 2 MB so that it sits cleanly at the front of the GPU segment - const auto num_workspace_bytes = math::align( - layout::EPWorkspaceLayout::get_num_bytes(), kNumAllocationAlignmentBytes); - - context = std::make_shared( - nccl_comm, symmetric::shared_comm_t{}, num_ranks, rank_idx, - num_workspace_bytes, num_buffer_bytes + num_lb_buffer_bytes, 0, true, - allow_hybrid_mode, sl_idx, num_allocated_qps, - 0, num_cpu_timeout_secs, num_gpu_timeout_secs); + const auto num_workspace_bytes = math::align(layout::EPWorkspaceLayout::get_num_bytes(), kNumAllocationAlignmentBytes); + + context = std::make_shared(nccl_comm, + symmetric::shared_comm_t{}, + num_ranks, + rank_idx, + num_workspace_bytes, + num_buffer_bytes + num_lb_buffer_bytes, + 0, + true, + allow_hybrid_mode, + sl_idx, + num_allocated_qps, + 0, + num_cpu_timeout_secs, + num_gpu_timeout_secs); main_context = context; auto& workspace = *static_cast(context->workspace); context->set_barrier_signals(&workspace.barrier_signals); // Expose only the LB region; the preceding bytes remain reserved for EP communication. lb_storage = torch::from_blob( - context->buffer, {num_buffer_bytes + num_lb_buffer_bytes}, [context = context](void*) {}, - torch::TensorOptions().dtype(torch::kByte).device(torch::kCUDA) - ).narrow(0, num_buffer_bytes, num_lb_buffer_bytes); + context->buffer, + {num_buffer_bytes + num_lb_buffer_bytes}, + [context = context](void*) {}, + torch::TensorOptions().dtype(torch::kByte).device(torch::kCUDA)) + .narrow(0, num_buffer_bytes, num_lb_buffer_bytes); // Allocate host workspaces CUDA_RUNTIME_CHECK(cudaMallocHost(&host_workspace, layout::EPWorkspaceLayout::get_num_bytes(), cudaHostAllocMapped)); @@ -95,16 +108,13 @@ class EPBuffer: public BufferBase { // NOTES: do not call our barrier, as the workspace is not ready yet } - ~EPBuffer() noexcept(false) override { - destroy_on_destruction("EP"); - } + ~EPBuffer() noexcept(false) override { destroy_on_destruction("EP"); } void destroy() override { EP_HOST_ASSERT(not destroyed); // Finish all works on all GPUs - comm::barrier(*context, context->barrier_signals, at::cuda::getCurrentCUDAStream(), - context->num_gpu_timeout_cycles, true, true); + comm::barrier(*context, context->barrier_signals, at::cuda::getCurrentCUDAStream(), context->num_gpu_timeout_cycles, true, true); // Deallocate host workspaces CUDA_RUNTIME_CHECK(cudaFreeHost(host_workspace)); @@ -163,10 +173,11 @@ class EPBuffer: public BufferBase { if (get_env("EP_AVOID_RECORD_STREAM", 0)) { event->tensors_to_record = tensors; } else { - for (auto& t: tensors) if (t.has_value()) { - t->record_stream(compute_stream); - t->record_stream(comm_stream); - } + for (auto& t : tensors) + if (t.has_value()) { + t->record_stream(compute_stream); + t->record_stream(comm_stream); + } } } else { comm::stream_wait(compute_stream, comm_stream); @@ -181,37 +192,39 @@ class EPBuffer: public BufferBase { } static int64_t get_dispatch_buffer_size(const int& num_max_tokens_per_rank, - const int& hidden, const int& num_sf_packs, const int& num_topk, + const int& hidden, + const int& num_sf_packs, + const int& num_topk, const int& elem_size, - const int& num_scaleout_ranks, const int& num_scaleup_ranks, + const int& num_scaleout_ranks, + const int& num_scaleup_ranks, const bool& is_scaleup_nvlink) { const auto num_ranks = num_scaleup_ranks * num_scaleout_ranks; const auto token_layout = get_dispatch_token_layout(hidden, elem_size, num_sf_packs, num_topk); if (num_scaleout_ranks == 1) { // Direct dispatch - const auto send_buffer_layout = layout::BufferLayout( - token_layout, is_scaleup_nvlink ? 0 : 1, num_max_tokens_per_rank); - const auto recv_buffer_layout = layout::BufferLayout( - token_layout, num_ranks, num_max_tokens_per_rank); + const auto send_buffer_layout = layout::BufferLayout(token_layout, is_scaleup_nvlink ? 0 : 1, num_max_tokens_per_rank); + const auto recv_buffer_layout = layout::BufferLayout(token_layout, num_ranks, num_max_tokens_per_rank); return send_buffer_layout.get_num_bytes() + recv_buffer_layout.get_num_bytes(); } else { // Hybrid dispatch - const auto scaleup_recv_buffer = layout::BufferLayout( - token_layout, num_scaleup_ranks, num_scaleout_ranks * num_max_tokens_per_rank); - const auto scaleout_send_buffer = layout::BufferLayout( - token_layout, 1, num_max_tokens_per_rank); - const auto scaleout_recv_buffer = layout::BufferLayout( - token_layout, num_scaleout_ranks, - /* kNumChannels * kNumMaxTokensPerChannel */ num_max_tokens_per_rank + kNumMaxChannels); - return scaleup_recv_buffer.get_num_bytes() + - scaleout_send_buffer.get_num_bytes() + - scaleout_recv_buffer.get_num_bytes(); + const auto scaleup_recv_buffer = + layout::BufferLayout(token_layout, num_scaleup_ranks, num_scaleout_ranks * num_max_tokens_per_rank); + const auto scaleout_send_buffer = layout::BufferLayout(token_layout, 1, num_max_tokens_per_rank); + const auto scaleout_recv_buffer = + layout::BufferLayout(token_layout, + num_scaleout_ranks, + /* kNumChannels * kNumMaxTokensPerChannel */ num_max_tokens_per_rank + kNumMaxChannels); + return scaleup_recv_buffer.get_num_bytes() + scaleout_send_buffer.get_num_bytes() + scaleout_recv_buffer.get_num_bytes(); } } - static int64_t get_combine_buffer_size(const int& num_max_tokens_per_rank, const int& hidden, const int& num_topk, - const int& num_scaleout_ranks, const int& num_scaleup_ranks, + static int64_t get_combine_buffer_size(const int& num_max_tokens_per_rank, + const int& hidden, + const int& num_topk, + const int& num_scaleout_ranks, + const int& num_scaleup_ranks, const bool& is_scaleup_nvlink, const bool& allow_multiple_reduction) { const auto num_ranks = num_scaleup_ranks * num_scaleout_ranks; @@ -220,35 +233,35 @@ class EPBuffer: public BufferBase { if (num_scaleout_ranks == 1) { // Direct combine const auto num_tokens_in_layout = allow_multiple_reduction ? std::min(num_ranks, num_topk) : num_topk; - const auto send_buffer_layout = layout::BufferLayout( - token_layout, is_scaleup_nvlink ? 0 : num_ranks, - // For single reduction cases, the maximum number of received tokens is - // `num_ranks * num_topk * num_max_tokens_per_rank` (we assume the bad case of `do_expand=True`) - num_max_tokens_per_rank * (allow_multiple_reduction ? 1 : num_topk)); - const auto recv_buffer_layout = layout::BufferLayout( - token_layout, num_tokens_in_layout, num_max_tokens_per_rank); + const auto send_buffer_layout = + layout::BufferLayout(token_layout, + is_scaleup_nvlink ? 0 : num_ranks, + // For single reduction cases, the maximum number of received tokens is + // `num_ranks * num_topk * num_max_tokens_per_rank` (we assume the bad case of `do_expand=True`) + num_max_tokens_per_rank * (allow_multiple_reduction ? 1 : num_topk)); + const auto recv_buffer_layout = layout::BufferLayout(token_layout, num_tokens_in_layout, num_max_tokens_per_rank); return send_buffer_layout.get_num_bytes() + recv_buffer_layout.get_num_bytes(); } else { // Hybrid combine const int num_tokens_in_scaleup_layout = allow_multiple_reduction ? std::min(num_scaleup_ranks, num_topk) : num_topk; const int num_tokens_in_scaleout_layout = allow_multiple_reduction ? std::min(num_scaleout_ranks, num_topk) : num_topk; - const auto scaleup_recv_buffer = layout::BufferLayout( - token_layout, num_tokens_in_scaleup_layout, num_scaleout_ranks * num_max_tokens_per_rank); - const auto scaleout_recv_buffer = layout::BufferLayout( - token_layout, num_tokens_in_scaleout_layout, num_max_tokens_per_rank); - const auto scaleout_send_buffer = layout::BufferLayout( - token_layout, allow_multiple_reduction ? 1 : num_topk, - /* kNumChannels * num_scaleout_ranks * kNumMaxTokensPerChannel */ - num_scaleout_ranks * (num_max_tokens_per_rank + kNumMaxChannels)); - return scaleup_recv_buffer.get_num_bytes() + - scaleout_send_buffer.get_num_bytes() + - scaleout_recv_buffer.get_num_bytes(); + const auto scaleup_recv_buffer = + layout::BufferLayout(token_layout, num_tokens_in_scaleup_layout, num_scaleout_ranks * num_max_tokens_per_rank); + const auto scaleout_recv_buffer = + layout::BufferLayout(token_layout, num_tokens_in_scaleout_layout, num_max_tokens_per_rank); + const auto scaleout_send_buffer = layout::BufferLayout(token_layout, + allow_multiple_reduction ? 1 : num_topk, + /* kNumChannels * num_scaleout_ranks * kNumMaxTokensPerChannel */ + num_scaleout_ranks * (num_max_tokens_per_rank + kNumMaxChannels)); + return scaleup_recv_buffer.get_num_bytes() + scaleout_send_buffer.get_num_bytes() + scaleout_recv_buffer.get_num_bytes(); } } static int64_t calculate_buffer_size(const int64_t& nccl_comm, - const int& num_max_tokens_per_rank, const int& hidden, - int num_topk, const bool& use_fp8_dispatch, + const int& num_max_tokens_per_rank, + const int& hidden, + int num_topk, + const bool& use_fp8_dispatch, const bool& allow_hybrid_mode, const bool& allow_multiple_reduction) { EP_HOST_ASSERT(num_max_tokens_per_rank > 0 and hidden > 0); @@ -266,51 +279,51 @@ class EPBuffer: public BufferBase { // Dispatch size const auto elem_size = use_fp8_dispatch ? sizeof(__nv_fp8_e4m3) : sizeof(nv_bfloat16); - const auto num_sf_packs = use_fp8_dispatch ? math::ceil_div(hidden, 32) : 0; // An approximation for number of SF packs + const auto num_sf_packs = use_fp8_dispatch ? math::ceil_div(hidden, 32) : 0; // An approximation for number of SF packs const auto num_dispatch_bytes = get_dispatch_buffer_size( - num_max_tokens_per_rank, hidden, num_sf_packs, num_topk, elem_size, - num_scaleout_ranks, num_scaleup_ranks, - is_scaleup_nvlink); + num_max_tokens_per_rank, hidden, num_sf_packs, num_topk, elem_size, num_scaleout_ranks, num_scaleup_ranks, is_scaleup_nvlink); // Combine layout const auto num_combine_bytes = get_combine_buffer_size( - num_max_tokens_per_rank, hidden, num_topk, - num_scaleout_ranks, num_scaleup_ranks, - is_scaleup_nvlink, allow_multiple_reduction); + num_max_tokens_per_rank, hidden, num_topk, num_scaleout_ranks, num_scaleup_ranks, is_scaleup_nvlink, allow_multiple_reduction); // Return the maximum of those layouts, aligned to 2 MB return math::align(std::max(num_dispatch_bytes, num_combine_bytes), kNumAllocationAlignmentBytes); } - pybind11::tuple - dispatch(const torch::Tensor& x, - const std::optional& sf, - const torch::Tensor& topk_idx, - const std::optional& topk_weights, - const std::optional& cumulative_local_expert_recv_stats, - const std::optional& cached_num_recv_tokens, - const std::optional& cached_num_expanded_tokens, - const std::optional>& cached_num_recv_tokens_per_expert_list, - const std::optional& cached_psum_num_recv_tokens_per_scaleup_rank, - const std::optional& cached_psum_num_recv_tokens_per_expert, - const std::optional& cached_num_unaligned_recv_tokens_per_expert, - const std::optional& cached_dst_buffer_slot_idx, - const std::optional& cached_token_metadata_at_forward, - const std::optional& cached_recv_src_metadata, - const std::optional& cached_channel_linked_list, - const int& num_max_tokens_per_rank, - const int& num_experts, const int& expert_alignment, - const int& num_sms, const int& num_qps, - const std::optional& previous_event, - const bool& async_with_compute_stream, - const bool& allocate_on_comm_stream, - const bool& do_cpu_sync, - const bool& do_expand, const bool& do_zero_padding, - const bool& use_tma_aligned_col_major_sf, - const bool& defer_epilogue, const bool& materialize_recv_x) const { - EP_HOST_ASSERT(materialize_recv_x or (context->num_scaleout_ranks == 1 and context->num_rdma_ranks == 1 and not do_expand and - not sf.has_value() and x.scalar_type() == torch::kBFloat16 and - do_cpu_sync and not cached_num_recv_tokens.has_value() and not defer_epilogue)); + pybind11::tuple dispatch(const torch::Tensor& x, + const std::optional& sf, + const torch::Tensor& topk_idx, + const std::optional& topk_weights, + const std::optional& cumulative_local_expert_recv_stats, + const std::optional& cached_num_recv_tokens, + const std::optional& cached_num_expanded_tokens, + const std::optional>& cached_num_recv_tokens_per_expert_list, + const std::optional& cached_psum_num_recv_tokens_per_scaleup_rank, + const std::optional& cached_psum_num_recv_tokens_per_expert, + const std::optional& cached_num_unaligned_recv_tokens_per_expert, + const std::optional& cached_dst_buffer_slot_idx, + const std::optional& cached_token_metadata_at_forward, + const std::optional& cached_recv_src_metadata, + const std::optional& cached_channel_linked_list, + const int& num_max_tokens_per_rank, + const int& num_experts, + const int& expert_alignment, + const int& num_sms, + const int& num_qps, + const std::optional& previous_event, + const bool& async_with_compute_stream, + const bool& allocate_on_comm_stream, + const bool& do_cpu_sync, + const bool& do_expand, + const bool& do_zero_padding, + const bool& use_tma_aligned_col_major_sf, + const bool& defer_epilogue, + const bool& materialize_recv_x) const { + EP_HOST_ASSERT(materialize_recv_x or + (context->num_scaleout_ranks == 1 and context->num_rdma_ranks == 1 and not do_expand and not sf.has_value() and + x.scalar_type() == torch::kBFloat16 and do_cpu_sync and not cached_num_recv_tokens.has_value() and + not defer_epilogue)); // Check SM count EP_HOST_ASSERT(num_sms > 0 and num_sms <= jit->device.get_num_sms()); EP_HOST_ASSERT((num_sms > 1 or context->num_scaleout_ranks == 1 or context->num_scaleup_ranks == 1) and @@ -381,8 +394,7 @@ class EPBuffer: public BufferBase { int* cumulative_local_expert_recv_stats_ptr = nullptr; if (cumulative_local_expert_recv_stats.has_value()) { const auto [num_local_experts_] = get_shape<1>(cumulative_local_expert_recv_stats.value()); - EP_HOST_ASSERT(cumulative_local_expert_recv_stats->is_cuda() and - cumulative_local_expert_recv_stats->is_contiguous()); + EP_HOST_ASSERT(cumulative_local_expert_recv_stats->is_cuda() and cumulative_local_expert_recv_stats->is_contiguous()); EP_HOST_ASSERT(num_local_experts == num_local_experts_); cumulative_local_expert_recv_stats_ptr = cumulative_local_expert_recv_stats->data_ptr(); } @@ -403,8 +415,7 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(psum_num_recv_tokens_per_expert.scalar_type() == torch::kInt); } else { // NOTES: for expand mode, the input is exclusive prefix sum, while for non-expand, it is inclusive - psum_num_recv_tokens_per_expert = torch::empty( - {num_local_experts + 1}, at::TensorOptions(torch::kCUDA).dtype(torch::kInt)); + psum_num_recv_tokens_per_expert = torch::empty({num_local_experts + 1}, at::TensorOptions(torch::kCUDA).dtype(torch::kInt)); } // The unaligned (actual) number of received tokens per expert @@ -417,8 +428,7 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(num_unaligned_recv_tokens_per_expert.is_cuda() and num_unaligned_recv_tokens_per_expert.is_contiguous()); EP_HOST_ASSERT(num_unaligned_recv_tokens_per_expert.scalar_type() == torch::kInt); } else { - num_unaligned_recv_tokens_per_expert = torch::empty( - {num_local_experts}, at::TensorOptions(torch::kCUDA).dtype(torch::kInt)); + num_unaligned_recv_tokens_per_expert = torch::empty({num_local_experts}, at::TensorOptions(torch::kCUDA).dtype(torch::kInt)); } num_unaligned_recv_tokens_per_expert_ptr = num_unaligned_recv_tokens_per_expert.data_ptr(); @@ -431,8 +441,8 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(psum_num_recv_tokens_per_scaleup_rank.is_cuda() and psum_num_recv_tokens_per_scaleup_rank.is_contiguous()); EP_HOST_ASSERT(psum_num_recv_tokens_per_scaleup_rank.scalar_type() == torch::kInt); } else { - psum_num_recv_tokens_per_scaleup_rank = torch::empty( - {context->num_scaleup_ranks}, at::TensorOptions(torch::kCUDA).dtype(torch::kInt)); + psum_num_recv_tokens_per_scaleup_rank = + torch::empty({context->num_scaleup_ranks}, at::TensorOptions(torch::kCUDA).dtype(torch::kInt)); } // Decide number of channels by shared memory consumption @@ -446,9 +456,7 @@ class EPBuffer: public BufferBase { num_channels_per_sm = std::min( (num_smem_bytes - get_num_notify_smem_bytes(context->num_ranks, num_experts)) / dispatch_token_layout.get_num_bytes(), 32 - kNumNotifyWarps); - num_channels_per_sm = std::min( - num_smem_bytes / combine_token_layout.get_num_bytes(), - num_channels_per_sm); + num_channels_per_sm = std::min(num_smem_bytes / combine_token_layout.get_num_bytes(), num_channels_per_sm); num_channels_per_sm = std::min( /* 2 kinds of warps */ num_channels_per_sm / 2, kNumMaxChannelsPerSM); if (not prefer_overlap_with_compute) @@ -468,8 +476,7 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(dst_buffer_slot_idx.scalar_type() == torch::kInt); } else { // Allocate a new tensor - dst_buffer_slot_idx = torch::empty( - {num_tokens, num_topk}, torch::TensorOptions(torch::kCUDA).dtype(torch::kInt)); + dst_buffer_slot_idx = torch::empty({num_tokens, num_topk}, torch::TensorOptions(torch::kCUDA).dtype(torch::kInt)); } } @@ -483,17 +490,14 @@ class EPBuffer: public BufferBase { // TODO: May make it a linked list to remove the redundant info in `token_metadata_at_forward` const auto num_max_tokens_per_channel = math::ceil_div(num_max_tokens_per_rank, num_channels); if (cached_mode) { - const auto [num_channels_, num_scaleout_ranks_, num_max_tokens_per_channel_, num_topk_] = - get_shape<4>(dst_buffer_slot_idx); + const auto [num_channels_, num_scaleout_ranks_, num_max_tokens_per_channel_, num_topk_] = get_shape<4>(dst_buffer_slot_idx); EP_HOST_ASSERT(num_channels == num_channels_ and context->num_scaleout_ranks == num_scaleout_ranks_ and num_max_tokens_per_channel == num_max_tokens_per_channel_ and num_topk == num_topk_); EP_HOST_ASSERT(dst_buffer_slot_idx.is_cuda() and dst_buffer_slot_idx.is_contiguous()); EP_HOST_ASSERT(dst_buffer_slot_idx.scalar_type() == torch::kInt); } else { - dst_buffer_slot_idx = torch::empty( - {num_channels, context->num_scaleout_ranks, num_max_tokens_per_channel, num_topk}, - torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt) - ); + dst_buffer_slot_idx = torch::empty({num_channels, context->num_scaleout_ranks, num_max_tokens_per_channel, num_topk}, + torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt)); } // The token metadata during forward @@ -507,16 +511,15 @@ class EPBuffer: public BufferBase { const auto num_forward_metadata_dims = 2 + num_topk * 2; if (cached_mode) { token_metadata_at_forward = cached_token_metadata_at_forward; - const auto [num_channels_, num_max_forwarded_tokens_, num_forward_metadata_dims_] = get_shape<3>(token_metadata_at_forward.value()); - EP_HOST_ASSERT(num_channels == num_channels_ and num_max_forwarded_tokens == num_max_forwarded_tokens_ - and num_forward_metadata_dims == num_forward_metadata_dims_); + const auto [num_channels_, num_max_forwarded_tokens_, num_forward_metadata_dims_] = + get_shape<3>(token_metadata_at_forward.value()); + EP_HOST_ASSERT(num_channels == num_channels_ and num_max_forwarded_tokens == num_max_forwarded_tokens_ and + num_forward_metadata_dims == num_forward_metadata_dims_); EP_HOST_ASSERT(token_metadata_at_forward->is_cuda() and token_metadata_at_forward->is_contiguous()); EP_HOST_ASSERT(token_metadata_at_forward->scalar_type() == torch::kInt); } else { - token_metadata_at_forward = torch::empty( - {num_channels, num_max_forwarded_tokens, num_forward_metadata_dims}, - torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt) - ); + token_metadata_at_forward = torch::empty({num_channels, num_max_forwarded_tokens, num_forward_metadata_dims}, + torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt)); } token_metadata_at_forward_ptr = token_metadata_at_forward->data_ptr(); @@ -534,66 +537,81 @@ class EPBuffer: public BufferBase { } else { channel_linked_list = torch::empty( // Index 0 of the list means the starting item - {num_channels, - context->num_scaleout_ranks * num_max_tokens_per_channel + 1, - context->num_scaleup_ranks}, - torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt) - ); + {num_channels, context->num_scaleout_ranks * num_max_tokens_per_channel + 1, context->num_scaleup_ranks}, + torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt)); } channel_linked_list_ptr = channel_linked_list->data_ptr(); } // Check buffer size - EP_HOST_ASSERT(get_dispatch_buffer_size( - num_max_tokens_per_rank, hidden, num_sf_packs, num_topk, x.element_size(), - context->num_scaleout_ranks, context->num_scaleup_ranks, - context->is_scaleup_nvlink) <= num_buffer_bytes); + EP_HOST_ASSERT(get_dispatch_buffer_size(num_max_tokens_per_rank, + hidden, + num_sf_packs, + num_topk, + x.element_size(), + context->num_scaleout_ranks, + context->num_scaleup_ranks, + context->is_scaleup_nvlink) <= num_buffer_bytes); // Ready and clean host workspace for this round - const auto host_workspace_layout = layout::EPWorkspaceLayout( - host_workspace, - context->num_scaleout_ranks, - context->num_scaleup_ranks, - num_experts); + const auto host_workspace_layout = + layout::EPWorkspaceLayout(host_workspace, context->num_scaleout_ranks, context->num_scaleup_ranks, num_experts); std::fill_n(host_workspace_layout.get_scaleup_rank_count_ptr(), context->num_scaleup_ranks, 0); std::fill_n(host_workspace_layout.get_scaleup_expert_count_ptr(), num_local_experts, 0); std::atomic_thread_fence(std::memory_order_seq_cst); // Do dispatch into the buffers (with SM limitation) - launch_dispatch(x.data_ptr(), sf_ptr, - topk_idx.data_ptr(), topk_weights_ptr, + launch_dispatch(x.data_ptr(), + sf_ptr, + topk_idx.data_ptr(), + topk_weights_ptr, cumulative_local_expert_recv_stats_ptr, psum_num_recv_tokens_per_scaleup_rank.data_ptr(), psum_num_recv_tokens_per_expert.data_ptr(), num_unaligned_recv_tokens_per_expert_ptr, dst_buffer_slot_idx.data_ptr(), token_metadata_at_forward_ptr, - num_tokens, num_max_tokens_per_rank, - hidden, x.element_size(), - num_sf_packs, sf_token_stride, sf_hidden_stride, - num_experts, num_topk, expert_alignment, - context->dev_comm, context->window, + num_tokens, + num_max_tokens_per_rank, + hidden, + x.element_size(), + num_sf_packs, + sf_token_stride, + sf_hidden_stride, + num_experts, + num_topk, + expert_alignment, + context->dev_comm, + context->window, context->buffer, - context->workspace, mapped_host_workspace, - context->scaleout_rank_idx, context->scaleup_rank_idx, - context->num_scaleout_ranks, context->num_scaleup_ranks, + context->workspace, + mapped_host_workspace, + context->scaleout_rank_idx, + context->scaleup_rank_idx, + context->num_scaleout_ranks, + context->num_scaleup_ranks, context->is_scaleup_nvlink, - num_sms, num_channels_per_sm, + num_sms, + num_channels_per_sm, num_smem_bytes, - num_qps, context->num_gpu_timeout_cycles, - cached_mode, do_cpu_sync, + num_qps, + context->num_gpu_timeout_cycles, + cached_mode, + do_cpu_sync, comm_stream); // For tensor recording - tensor_list_t tensors_to_record = { - x, sf, topk_idx, topk_weights, - cumulative_local_expert_recv_stats, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - num_unaligned_recv_tokens_per_expert, - dst_buffer_slot_idx, - token_metadata_at_forward, - channel_linked_list}; + tensor_list_t tensors_to_record = {x, + sf, + topk_idx, + topk_weights, + cumulative_local_expert_recv_stats, + psum_num_recv_tokens_per_scaleup_rank, + psum_num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + dst_buffer_slot_idx, + token_metadata_at_forward, + channel_linked_list}; // Epilogue can be deferred, so it is a lambda auto epilogue = [=, this](const at::cuda::CUDAStream& stream, @@ -612,8 +630,8 @@ class EPBuffer: public BufferBase { num_expanded_tokens = cached_num_expanded_tokens.value(); } else if (do_cpu_sync) { // In dispatch, CPU will busy-wait until GPU receive tensor size metadata from other ranks, which can be quite long. - // If users of DeepEP need to execute other Python code on other threads, such as KV transfer, their code will get stuck due to GIL - // unless we release GIL here. + // If users of DeepEP need to execute other Python code on other threads, such as KV transfer, their code will get stuck due + // to GIL unless we release GIL here. pybind11::gil_scoped_release release; // Non-cached mode with sync @@ -627,7 +645,7 @@ class EPBuffer: public BufferBase { host_workspace_layout.get_scaleup_rank_count_ptr()[counter_scaleup_rank_idx]); if ((ready = math::is_decoded_positive_ready(count))) { num_recv_tokens += count; - ++ counter_scaleup_rank_idx; + ++counter_scaleup_rank_idx; } } @@ -638,7 +656,7 @@ class EPBuffer: public BufferBase { if ((ready = math::is_decoded_positive_ready(count))) { num_recv_tokens_per_expert_list.push_back(count); num_expanded_tokens += count; - ++ counter_local_expert_idx; + ++counter_local_expert_idx; } } @@ -646,9 +664,9 @@ class EPBuffer: public BufferBase { const auto get_buffer_info = [&]() { std::stringstream ss; ss << "CPU side received count (scaleup: " << context->scaleup_rank_idx << "): "; - for (int i = 0; i < context->num_scaleup_ranks + num_local_experts; ++ i) { + for (int i = 0; i < context->num_scaleup_ranks + num_local_experts; ++i) { ss << host_workspace_layout.get_scaleup_rank_expert_count_ptr()[i]; - ss << (i == context->num_scaleup_ranks - 1 ? " # ": " "); + ss << (i == context->num_scaleup_ranks - 1 ? " # " : " "); } return ss.str(); }; @@ -682,10 +700,9 @@ class EPBuffer: public BufferBase { auto recv_sf = std::optional(); auto recv_topk_idx = std::optional(); auto recv_topk_weights = std::optional(); - auto recv_src_metadata = cached_mode ? - cached_recv_src_metadata.value() : - torch::empty({num_recv_tokens, num_topk + 2}, - torch::TensorOptions(torch::kCUDA).dtype(torch::kInt)); + auto recv_src_metadata = cached_mode + ? cached_recv_src_metadata.value() + : torch::empty({num_recv_tokens, num_topk + 2}, torch::TensorOptions(torch::kCUDA).dtype(torch::kInt)); // Optional tensors void* recv_sf_ptr = nullptr; @@ -699,9 +716,8 @@ class EPBuffer: public BufferBase { // TMA-aligned layout for the next GEMM input recv_sf_token_stride = 1, recv_sf_hidden_stride = math::align(num_allocated_tokens, kNumAlignedSFPacks); } - recv_sf = torch::empty_strided({num_allocated_tokens, num_sf_packs}, - {recv_sf_token_stride, recv_sf_hidden_stride}, - sf->options()); + recv_sf = torch::empty_strided( + {num_allocated_tokens, num_sf_packs}, {recv_sf_token_stride, recv_sf_hidden_stride}, sf->options()); recv_sf_ptr = recv_sf->data_ptr(); } if (not do_expand) { @@ -709,9 +725,8 @@ class EPBuffer: public BufferBase { recv_topk_idx_ptr = recv_topk_idx->data_ptr(); } if (topk_weights.has_value()) { - recv_topk_weights = do_expand ? - torch::empty({num_allocated_tokens}, topk_weights->options()) : - torch::empty({num_allocated_tokens, num_topk}, topk_weights->options()); + recv_topk_weights = do_expand ? torch::empty({num_allocated_tokens}, topk_weights->options()) + : torch::empty({num_allocated_tokens, num_topk}, topk_weights->options()); recv_topk_weights_ptr = recv_topk_weights->data_ptr(); } @@ -728,40 +743,55 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(psum_num_recv_tokens_per_expert.size(0) == num_local_experts); // Launch copy kernels with full SMs - launch_dispatch_copy_epilogue(context->buffer, context->workspace, + launch_dispatch_copy_epilogue(context->buffer, + context->workspace, psum_num_recv_tokens_per_scaleup_rank.data_ptr(), psum_num_recv_tokens_per_expert.data_ptr(), - recv_x.has_value() ? recv_x->data_ptr() : nullptr, recv_sf_ptr, - recv_topk_idx_ptr, recv_topk_weights_ptr, + recv_x.has_value() ? recv_x->data_ptr() : nullptr, + recv_sf_ptr, + recv_topk_idx_ptr, + recv_topk_weights_ptr, recv_src_metadata.data_ptr(), recv_row_indices.has_value() ? recv_row_indices->data_ptr() : nullptr, channel_linked_list_ptr, num_unaligned_recv_tokens_per_expert_ptr, - num_recv_tokens, num_max_tokens_per_rank, + num_recv_tokens, + num_max_tokens_per_rank, num_hidden_bytes, - num_sf_packs, recv_sf_token_stride, recv_sf_hidden_stride, - num_experts, num_topk, expert_alignment, - context->scaleout_rank_idx, context->scaleup_rank_idx, - context->num_scaleout_ranks, context->num_scaleup_ranks, + num_sf_packs, + recv_sf_token_stride, + recv_sf_hidden_stride, + num_experts, + num_topk, + expert_alignment, + context->scaleout_rank_idx, + context->scaleup_rank_idx, + context->num_scaleout_ranks, + context->num_scaleup_ranks, jit->device.get_num_sms(), jit->device.get_num_smem_bytes(), num_channels, - do_expand, cached_mode, - do_zero_padding, materialize_recv_x, + do_expand, + cached_mode, + do_zero_padding, + materialize_recv_x, stream); - auto result = pybind11::make_tuple( - recv_x, recv_sf, - recv_topk_idx, recv_topk_weights, - num_recv_tokens, num_expanded_tokens, - num_recv_tokens_per_expert_list, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - num_unaligned_recv_tokens_per_expert, - recv_src_metadata, - dst_buffer_slot_idx, - token_metadata_at_forward, - channel_linked_list, recv_row_indices); + auto result = pybind11::make_tuple(recv_x, + recv_sf, + recv_topk_idx, + recv_topk_weights, + num_recv_tokens, + num_expanded_tokens, + num_recv_tokens_per_expert_list, + psum_num_recv_tokens_per_scaleup_rank, + psum_num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + recv_src_metadata, + dst_buffer_slot_idx, + token_metadata_at_forward, + channel_linked_list, + recv_row_indices); // For non-deferring tensor recording if (tensors_to_record_opt.has_value()) { @@ -780,8 +810,7 @@ class EPBuffer: public BufferBase { // NOTES: CPU sync will be deferred, too if (defer_epilogue) { EP_HOST_ASSERT(async_with_compute_stream); - const auto event = stream_control_epilogue( - tensors_to_record, compute_stream, allocate_on_comm_stream, true); + const auto event = stream_control_epilogue(tensors_to_record, compute_stream, allocate_on_comm_stream, true); std::function epilogue_hook = [epilogue = std::move(epilogue)]() mutable { return epilogue(at::cuda::getCurrentCUDAStream(), std::nullopt); }; @@ -790,41 +819,39 @@ class EPBuffer: public BufferBase { // Do epilogue intermediately and record all tensors auto result = epilogue(comm_stream, std::ref(tensors_to_record)); - const auto event = stream_control_epilogue( - tensors_to_record, compute_stream, allocate_on_comm_stream, async_with_compute_stream); + const auto event = stream_control_epilogue(tensors_to_record, compute_stream, allocate_on_comm_stream, async_with_compute_stream); return pybind11::make_tuple(result, event, pybind11::none()); } - torch::Tensor get_dispatch_recv_slab(const torch::Tensor& x, const int& num_topk, - const int& num_max_tokens_per_rank) const { + torch::Tensor get_dispatch_recv_slab(const torch::Tensor& x, const int& num_topk, const int& num_max_tokens_per_rank) const { EP_HOST_ASSERT(not destroyed and context->num_scaleout_ranks == 1 and context->num_rdma_ranks == 1); EP_HOST_ASSERT(x.is_cuda() and x.scalar_type() == torch::kBFloat16 and x.dim() == 2); const auto token_layout = layout::TokenLayout(x.size(1) * x.element_size(), 0, num_topk, true); const auto slots = static_cast(context->num_scaleup_ranks) * num_max_tokens_per_rank; const auto stride_bytes = token_layout.get_num_bytes(); EP_HOST_ASSERT(slots * stride_bytes <= context->num_gpu_buffer_bytes); - return torch::from_blob(context->buffer, {slots, x.size(1)}, - {stride_bytes / static_cast(x.element_size()), 1}, x.options()); + return torch::from_blob( + context->buffer, {slots, x.size(1)}, {stride_bytes / static_cast(x.element_size()), 1}, x.options()); } - pybind11::tuple - combine(const torch::Tensor& x, - const std::optional& topk_weights, - const std::optional& bias_0, - const std::optional& bias_1, - const torch::Tensor& src_metadata, - const torch::Tensor& combined_topk_idx, - const torch::Tensor& psum_num_recv_tokens_per_scaleup_rank, - const std::optional& token_metadata_at_forward, - const std::optional& channel_linked_list, - const int& num_experts, - const int& num_max_tokens_per_rank, - const int& num_sms, const int& num_qps, - const std::optional& previous_event, - const bool& async_with_compute_stream, - const bool& allocate_on_comm_stream, - const bool& use_expanded_layout, - const bool& defer_epilogue) const { + pybind11::tuple combine(const torch::Tensor& x, + const std::optional& topk_weights, + const std::optional& bias_0, + const std::optional& bias_1, + const torch::Tensor& src_metadata, + const torch::Tensor& combined_topk_idx, + const torch::Tensor& psum_num_recv_tokens_per_scaleup_rank, + const std::optional& token_metadata_at_forward, + const std::optional& channel_linked_list, + const int& num_experts, + const int& num_max_tokens_per_rank, + const int& num_sms, + const int& num_qps, + const std::optional& previous_event, + const bool& async_with_compute_stream, + const bool& allocate_on_comm_stream, + const bool& use_expanded_layout, + const bool& defer_epilogue) const { // Check SM count EP_HOST_ASSERT(num_sms > 0 and num_sms <= jit->device.get_num_sms()); EP_HOST_ASSERT((num_sms > 1 or context->num_scaleout_ranks == 1 or context->num_scaleup_ranks == 1) and @@ -869,7 +896,7 @@ class EPBuffer: public BufferBase { } const auto bias_opts = std::vector({bias_0, bias_1}); - for (int i = 0; i < 2; ++ i) { + for (int i = 0; i < 2; ++i) { if (bias_opts[i].has_value()) { auto bias = bias_opts[i].value(); EP_HOST_ASSERT(bias.dim() == 2 and bias.is_cuda() and bias.is_contiguous()); @@ -884,9 +911,13 @@ class EPBuffer: public BufferBase { const auto comm_stream = comm::get_comm_stream(); // Check buffer size - EP_HOST_ASSERT(get_combine_buffer_size(num_max_tokens_per_rank, hidden, num_topk, - context->num_scaleout_ranks, context->num_scaleup_ranks, - context->is_scaleup_nvlink, allow_multiple_reduction) <= num_buffer_bytes); + EP_HOST_ASSERT(get_combine_buffer_size(num_max_tokens_per_rank, + hidden, + num_topk, + context->num_scaleout_ranks, + context->num_scaleup_ranks, + context->is_scaleup_nvlink, + allow_multiple_reduction) <= num_buffer_bytes); // Optional configs and metadata for hybrid combine int num_channels = 1; @@ -915,32 +946,45 @@ class EPBuffer: public BufferBase { // Push data into remote buffers // NOTES: we don't use `num_hidden_bytes` due to enable later quantization possibility - const auto reduce_buffer = launch_combine( - x.data_ptr(), - topk_weights.has_value() ? topk_weights->data_ptr() : nullptr, - src_metadata.data_ptr(), - psum_num_recv_tokens_per_scaleup_rank.data_ptr(), - token_metadata_at_forward_ptr, - channel_linked_list_ptr, - context->dev_comm, context->window, - context->buffer, context->workspace, - num_reduced_tokens, num_max_tokens_per_rank, - hidden, num_experts, num_topk, - num_qps, context->num_gpu_timeout_cycles, - context->num_scaleout_ranks, context->num_scaleup_ranks, - context->scaleout_rank_idx, context->scaleup_rank_idx, - context->is_scaleup_nvlink, - num_sms, jit->device.get_num_smem_bytes(), - num_channels, - use_expanded_layout, allow_multiple_reduction, - comm_stream); + const auto reduce_buffer = launch_combine(x.data_ptr(), + topk_weights.has_value() ? topk_weights->data_ptr() : nullptr, + src_metadata.data_ptr(), + psum_num_recv_tokens_per_scaleup_rank.data_ptr(), + token_metadata_at_forward_ptr, + channel_linked_list_ptr, + context->dev_comm, + context->window, + context->buffer, + context->workspace, + num_reduced_tokens, + num_max_tokens_per_rank, + hidden, + num_experts, + num_topk, + num_qps, + context->num_gpu_timeout_cycles, + context->num_scaleout_ranks, + context->num_scaleup_ranks, + context->scaleout_rank_idx, + context->scaleup_rank_idx, + context->is_scaleup_nvlink, + num_sms, + jit->device.get_num_smem_bytes(), + num_channels, + use_expanded_layout, + allow_multiple_reduction, + comm_stream); // For tensor recording - tensor_list_t tensors_to_record = { - x, topk_weights, bias_0, bias_1, - src_metadata, combined_topk_idx, - psum_num_recv_tokens_per_scaleup_rank, - token_metadata_at_forward, channel_linked_list}; + tensor_list_t tensors_to_record = {x, + topk_weights, + bias_0, + bias_1, + src_metadata, + combined_topk_idx, + psum_num_recv_tokens_per_scaleup_rank, + token_metadata_at_forward, + channel_linked_list}; // Epilogue can be deferred, so it is a lambda auto epilogue = [=, this](const at::cuda::CUDAStream& stream, @@ -955,25 +999,29 @@ class EPBuffer: public BufferBase { } // Resolve pointers here to retain bias tensors in a deferred epilogue - void* bias_ptrs[2] = { - bias_opts[0].has_value() ? bias_opts[0]->data_ptr() : nullptr, - bias_opts[1].has_value() ? bias_opts[1]->data_ptr() : nullptr - }; + void* bias_ptrs[2] = {bias_opts[0].has_value() ? bias_opts[0]->data_ptr() : nullptr, + bias_opts[1].has_value() ? bias_opts[1]->data_ptr() : nullptr}; // Combine pushed data launch_combine_reduce_epilogue(combined_x.data_ptr(), combined_topk_weights_ptr, combined_topk_idx.data_ptr(), - num_combined_tokens, num_max_tokens_per_rank, + num_combined_tokens, + num_max_tokens_per_rank, hidden, - num_experts, num_topk, + num_experts, + num_topk, reduce_buffer, - bias_ptrs[0], bias_ptrs[1], - context->num_scaleout_ranks, context->num_scaleup_ranks, - context->scaleout_rank_idx, context->scaleup_rank_idx, + bias_ptrs[0], + bias_ptrs[1], + context->num_scaleout_ranks, + context->num_scaleup_ranks, + context->scaleout_rank_idx, + context->scaleup_rank_idx, jit->device.get_num_sms(), jit->device.get_num_smem_bytes(), - use_expanded_layout, allow_multiple_reduction, + use_expanded_layout, + allow_multiple_reduction, stream); if (tensors_to_record_opt.has_value()) { @@ -986,8 +1034,7 @@ class EPBuffer: public BufferBase { // Defer epilogue if (defer_epilogue) { - const auto event = stream_control_epilogue( - tensors_to_record, compute_stream, allocate_on_comm_stream, true); + const auto event = stream_control_epilogue(tensors_to_record, compute_stream, allocate_on_comm_stream, true); std::function epilogue_hook = [epilogue = std::move(epilogue)]() { return epilogue(at::cuda::getCurrentCUDAStream(), std::nullopt); }; @@ -996,17 +1043,15 @@ class EPBuffer: public BufferBase { // Do epilogue auto result = epilogue(comm_stream, std::ref(tensors_to_record)); - const auto event = stream_control_epilogue( - tensors_to_record, compute_stream, allocate_on_comm_stream, async_with_compute_stream); + const auto event = stream_control_epilogue(tensors_to_record, compute_stream, allocate_on_comm_stream, async_with_compute_stream); return pybind11::make_tuple(result, event, pybind11::none()); } - std::optional - lb_prefetch_weights(const std::vector& redundant_expert_weights, - const std::vector& expert_weights, - const torch::Tensor& redundancy_mapping, - const int& num_sms, - const std::optional& previous_event) const { + std::optional lb_prefetch_weights(const std::vector& redundant_expert_weights, + const std::vector& expert_weights, + const torch::Tensor& redundancy_mapping, + const int& num_sms, + const std::optional& previous_event) const { // Checks EP_HOST_ASSERT(num_sms > 0 and num_sms <= jit->device.get_num_sms()); EP_HOST_ASSERT(redundant_expert_weights.size() == expert_weights.size()); @@ -1023,7 +1068,7 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(num_local_experts > 0); layout::EPWeightList weights; - for (int i = 0; i < num_weights; ++ i) { + for (int i = 0; i < num_weights; ++i) { const auto& redundant_expert_weight = redundant_expert_weights[i]; const auto& expert_weight = expert_weights[i]; EP_HOST_ASSERT(redundant_expert_weight.is_cuda() and redundant_expert_weight.is_contiguous()); @@ -1032,27 +1077,28 @@ class EPBuffer: public BufferBase { EP_HOST_ASSERT(expert_weight.dim() >= 1); EP_HOST_ASSERT(redundant_expert_weight.size(0) == num_redundant_experts); EP_HOST_ASSERT(expert_weight.size(0) == num_local_experts); - const int64_t num_bytes_per_expert = c10::multiply_integers(redundant_expert_weight.sizes().slice(1)) * - redundant_expert_weight.element_size(); - EP_HOST_ASSERT(num_bytes_per_expert == c10::multiply_integers(expert_weight.sizes().slice(1)) * - expert_weight.element_size()); + const int64_t num_bytes_per_expert = + c10::multiply_integers(redundant_expert_weight.sizes().slice(1)) * redundant_expert_weight.element_size(); + EP_HOST_ASSERT(num_bytes_per_expert == c10::multiply_integers(expert_weight.sizes().slice(1)) * expert_weight.element_size()); EP_HOST_ASSERT(num_bytes_per_expert % 16 == 0); - weights[i] = { - .redundant_expert_weights = redundant_expert_weight.data_ptr(), - .expert_weights = expert_weight.data_ptr(), - .num_bytes_per_expert = num_bytes_per_expert - }; + weights[i] = {.redundant_expert_weights = redundant_expert_weight.data_ptr(), + .expert_weights = expert_weight.data_ptr(), + .num_bytes_per_expert = num_bytes_per_expert}; } // Stream control const auto compute_stream = stream_control_prologue(previous_event); // Launch - launch_lb_prefetch_weights( - *context, num_weights, weights, redundancy_mapping.data_ptr(), - num_redundant_experts, num_local_experts, - num_sms, comm::get_comm_stream()); + launch_lb_prefetch_weights(*context, + num_weights, + weights, + redundancy_mapping.data_ptr(), + num_redundant_experts, + num_local_experts, + num_sms, + comm::get_comm_stream()); // Stream epilogue tensor_list_t tensors_to_record = {redundancy_mapping}; @@ -1061,12 +1107,11 @@ class EPBuffer: public BufferBase { return stream_control_epilogue(tensors_to_record, compute_stream, false, true); } - std::optional - lb_reduce_grads(const torch::Tensor& redundant_expert_grads, - const torch::Tensor& expert_grads, - const torch::Tensor& redundancy_mapping, - const int& num_sms, - const std::optional& previous_event) const { + std::optional lb_reduce_grads(const torch::Tensor& redundant_expert_grads, + const torch::Tensor& expert_grads, + const torch::Tensor& redundancy_mapping, + const int& num_sms, + const std::optional& previous_event) const { // Checks EP_HOST_ASSERT(num_sms > 0 and num_sms <= jit->device.get_num_sms()); EP_HOST_ASSERT(redundant_expert_grads.is_cuda() and redundant_expert_grads.is_contiguous()); @@ -1090,15 +1135,18 @@ class EPBuffer: public BufferBase { const auto compute_stream = stream_control_prologue(previous_event); // Launch: accumulate peers' redundant gradients into the local expert gradients - launch_lb_reduce_grads( - *context, redundant_expert_grads.data_ptr(), expert_grads.data_ptr(), - redundancy_mapping.data_ptr(), num_redundant_experts, num_local_experts, hidden, - num_sms, comm::get_comm_stream()); + launch_lb_reduce_grads(*context, + redundant_expert_grads.data_ptr(), + expert_grads.data_ptr(), + redundancy_mapping.data_ptr(), + num_redundant_experts, + num_local_experts, + hidden, + num_sms, + comm::get_comm_stream()); // Stream epilogue - return stream_control_epilogue( - {redundant_expert_grads, expert_grads, redundancy_mapping}, - compute_stream, false, true); + return stream_control_epilogue({redundant_expert_grads, expert_grads, redundancy_mapping}, compute_stream, false, true); } }; @@ -1113,10 +1161,7 @@ static void register_apis(pybind11::module_& m) { .def("lb_prefetch_weights", &EPBuffer::lb_prefetch_weights) .def("lb_reduce_grads", &EPBuffer::lb_reduce_grads); m.def("calculate_ep_buffer_size", &EPBuffer::calculate_buffer_size); - m.def("get_ep_buffer_alignment", [=]() { - return kNumAllocationAlignmentBytes; - }); - + m.def("get_ep_buffer_alignment", [=]() { return kNumAllocationAlignmentBytes; }); } } // namespace deep_ep::ep diff --git a/csrc/kernels/ep/dispatch.hpp b/csrc/kernels/ep/dispatch.hpp index 50c5bf3f..244a9fba 100644 --- a/csrc/kernels/ep/dispatch.hpp +++ b/csrc/kernels/ep/dispatch.hpp @@ -1,17 +1,16 @@ #pragma once -#include -#include -#include -#include - #include #include #include +#include +#include #include #include #include +#include +#include #include "../../runtime/jit.hpp" @@ -23,32 +22,48 @@ static int get_num_notify_smem_bytes(const int& num_ranks, const int& num_expert return math::align(num_ranks + num_experts, kNumNotifyWarps * 32) * sizeof(int); } -static layout::TokenLayout get_dispatch_token_layout( - const int& hidden, const int& elem_size, const int& num_sf_packs, const int& num_topk) { +static layout::TokenLayout get_dispatch_token_layout(const int& hidden, + const int& elem_size, + const int& num_sf_packs, + const int& num_topk) { return layout::TokenLayout(hidden * elem_size, num_sf_packs * sizeof(sf_pack_t), num_topk, true); } -static void launch_dispatch(void* x, void* sf, - topk_idx_t* topk_idx, float* topk_weights, +static void launch_dispatch(void* x, + void* sf, + topk_idx_t* topk_idx, + float* topk_weights, int* cumulative_local_expert_recv_stats, int* psum_num_recv_tokens_per_scaleup_rank, int* psum_num_recv_tokens_per_expert, int* num_unaligned_recv_tokens_per_expert, int* dst_buffer_slot_idx, int* token_metadata_at_forward, - const int& num_tokens, const int& num_max_tokens_per_rank, - const int& hidden, const int& elem_size, - const int& num_sf_packs, const int& sf_token_stride, const int& sf_hidden_stride, - const int& num_experts, const int& num_topk, const int& expert_alignment, - const deep_jit::NoRefPtr& nccl_dev_comm, const ncclWindow_t& nccl_window, + const int& num_tokens, + const int& num_max_tokens_per_rank, + const int& hidden, + const int& elem_size, + const int& num_sf_packs, + const int& sf_token_stride, + const int& sf_hidden_stride, + const int& num_experts, + const int& num_topk, + const int& expert_alignment, + const deep_jit::NoRefPtr& nccl_dev_comm, + const ncclWindow_t& nccl_window, void* buffer, - void* workspace, void* mapped_host_workspace, - const int& scaleout_rank_idx, const int& scaleup_rank_idx, - const int& num_scaleout_ranks, const int& num_scaleup_ranks, + void* workspace, + void* mapped_host_workspace, + const int& scaleout_rank_idx, + const int& scaleup_rank_idx, + const int& num_scaleout_ranks, + const int& num_scaleup_ranks, const bool& is_scaleup_nvlink, - const int& num_sms, const int& num_channels_per_sm, + const int& num_sms, + const int& num_channels_per_sm, const int& num_smem_bytes, - const int& num_qps, const int64_t& num_timeout_cycles, + const int& num_qps, + const int64_t& num_timeout_cycles, const bool& cached_mode, const bool& do_cpu_sync, const at::cuda::CUDAStream& stream) { @@ -74,8 +89,8 @@ static void launch_dispatch(void* x, void* sf, // Maximize shared memory utilization if (num_scaleout_ranks == 1) { const auto token_layout = get_dispatch_token_layout(hidden, elem_size, num_sf_packs, num_topk); - num_dispatch_warps = std::min( - (num_smem_bytes - num_notify_smem_bytes) / token_layout.get_num_bytes(), 32 - num_notify_warps); + num_dispatch_warps = + std::min((num_smem_bytes - num_notify_smem_bytes) / token_layout.get_num_bytes(), 32 - num_notify_warps); num_threads = (num_notify_warps + num_dispatch_warps) * 32; } else { // Hybrid kernels @@ -89,39 +104,54 @@ static void launch_dispatch(void* x, void* sf, if (num_scaleout_ranks == 1) { header_name = "dispatch"; func_name = std::format("dispatch_impl<{}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}>", - is_scaleup_nvlink, - do_cpu_sync, - reuse_slot_indices, - num_sms, - num_notify_warps, num_dispatch_warps, - num_scaleup_ranks, - hidden * elem_size, num_sf_packs, - num_max_tokens_per_rank, - num_experts, num_topk, expert_alignment, - num_qps, num_timeout_cycles); + is_scaleup_nvlink, + do_cpu_sync, + reuse_slot_indices, + num_sms, + num_notify_warps, + num_dispatch_warps, + num_scaleup_ranks, + hidden * elem_size, + num_sf_packs, + num_max_tokens_per_rank, + num_experts, + num_topk, + expert_alignment, + num_qps, + num_timeout_cycles); } else { header_name = "hybrid_dispatch"; func_name = std::format("hybrid_dispatch_impl<{}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}>", - do_cpu_sync, - reuse_slot_indices, - num_sms, - num_notify_warps, num_scaleout_warps, num_forward_warps, - num_scaleout_ranks, num_scaleup_ranks, - hidden * elem_size, num_sf_packs, - num_max_tokens_per_rank, - num_experts, num_topk, expert_alignment, - num_qps, num_timeout_cycles); + do_cpu_sync, + reuse_slot_indices, + num_sms, + num_notify_warps, + num_scaleout_warps, + num_forward_warps, + num_scaleout_ranks, + num_scaleup_ranks, + hidden * elem_size, + num_sf_packs, + num_max_tokens_per_rank, + num_experts, + num_topk, + expert_alignment, + num_qps, + num_timeout_cycles); } - const auto kernel = jit->compile("dispatch", std::format(R"( + const auto kernel = jit->compile("dispatch", + std::format(R"( #include static void __instantiate_kernel() {{ auto ptr = reinterpret_cast(&deep_ep::ep::{}); }} -)", header_name, func_name)); +)", + header_name, + func_name)); // Launch - const auto options = deep_jit::cuda::LaunchOptions { + const auto options = deep_jit::cuda::LaunchOptions{ .stream = stream.stream(), .num_smem_bytes = num_smem_bytes, .grid_dim = dim3(num_sms, 1, 1), @@ -130,59 +160,84 @@ static void __instantiate_kernel() {{ .cooperative = true, }; if (num_scaleout_ranks == 1) { - jit->launch( - kernel, options, - x, static_cast(sf), topk_idx, topk_weights, - cumulative_local_expert_recv_stats, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - num_unaligned_recv_tokens_per_expert, - dst_buffer_slot_idx, - num_tokens, - sf_token_stride, sf_hidden_stride, - nccl_dev_comm, nccl_window, - buffer, - workspace, mapped_host_workspace, - scaleup_rank_idx - ); + jit->launch(kernel, + options, + x, + static_cast(sf), + topk_idx, + topk_weights, + cumulative_local_expert_recv_stats, + psum_num_recv_tokens_per_scaleup_rank, + psum_num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + dst_buffer_slot_idx, + num_tokens, + sf_token_stride, + sf_hidden_stride, + nccl_dev_comm, + nccl_window, + buffer, + workspace, + mapped_host_workspace, + scaleup_rank_idx); } else { - jit->launch( - kernel, options, - x, static_cast(sf), topk_idx, topk_weights, - cumulative_local_expert_recv_stats, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - num_unaligned_recv_tokens_per_expert, - dst_buffer_slot_idx, - token_metadata_at_forward, - num_tokens, - sf_token_stride, sf_hidden_stride, - nccl_dev_comm, nccl_window, - buffer, - workspace, mapped_host_workspace, - scaleout_rank_idx, scaleup_rank_idx - ); + jit->launch(kernel, + options, + x, + static_cast(sf), + topk_idx, + topk_weights, + cumulative_local_expert_recv_stats, + psum_num_recv_tokens_per_scaleup_rank, + psum_num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + dst_buffer_slot_idx, + token_metadata_at_forward, + num_tokens, + sf_token_stride, + sf_hidden_stride, + nccl_dev_comm, + nccl_window, + buffer, + workspace, + mapped_host_workspace, + scaleout_rank_idx, + scaleup_rank_idx); } } -static void launch_dispatch_copy_epilogue(void* buffer, void* workspace, +static void launch_dispatch_copy_epilogue(void* buffer, + void* workspace, int* psum_num_recv_tokens_per_scaleup_rank, int* psum_num_recv_tokens_per_expert, - void* recv_x, void* recv_sf, - topk_idx_t* recv_topk_idx, float* recv_topk_weights, - int* recv_src_metadata, int64_t* recv_row_indices, + void* recv_x, + void* recv_sf, + topk_idx_t* recv_topk_idx, + float* recv_topk_weights, + int* recv_src_metadata, + int64_t* recv_row_indices, int* channel_linked_list, int* num_unaligned_recv_tokens_per_expert, - const int& num_recv_tokens, const int& num_max_tokens_per_rank, + const int& num_recv_tokens, + const int& num_max_tokens_per_rank, const int& num_hidden_bytes, - const int& num_sf_packs, const int& recv_sf_token_stride, const int& recv_sf_hidden_stride, - const int& num_experts, const int& num_topk, const int& expert_alignment, - const int& scaleout_rank_idx, const int& scaleup_rank_idx, - const int& num_scaleout_ranks, const int& num_scaleup_ranks, - const int& num_sms, const int& num_smem_bytes, + const int& num_sf_packs, + const int& recv_sf_token_stride, + const int& recv_sf_hidden_stride, + const int& num_experts, + const int& num_topk, + const int& expert_alignment, + const int& scaleout_rank_idx, + const int& scaleup_rank_idx, + const int& num_scaleout_ranks, + const int& num_scaleup_ranks, + const int& num_sms, + const int& num_smem_bytes, const int& num_channels, - const bool& do_expand, const bool& cached_mode, - const bool& do_zero_padding, const bool& materialize_recv_x, + const bool& do_expand, + const bool& cached_mode, + const bool& do_zero_padding, + const bool& materialize_recv_x, const at::cuda::CUDAStream& stream) { // Maximize shared memory utilization const auto token_layout = layout::TokenLayout(num_hidden_bytes, num_sf_packs * sizeof(sf_pack_t), num_topk, true); @@ -190,39 +245,56 @@ static void launch_dispatch_copy_epilogue(void* buffer, void* workspace, const auto num_threads = num_warps * 32; // Compile - const auto kernel = jit->compile("dispatch_copy_epilogue", std::format(R"( + const auto kernel = jit->compile("dispatch_copy_epilogue", + std::format(R"( #include static void __instantiate_kernel() {{ auto ptr = reinterpret_cast(&deep_ep::ep::dispatch_copy_epilogue_impl<{}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}, {}>); }} -)", do_expand, cached_mode, do_zero_padding, materialize_recv_x, - num_sms, num_channels, num_warps, - num_scaleout_ranks, num_scaleup_ranks, - num_hidden_bytes, num_sf_packs, - num_max_tokens_per_rank, - num_experts, num_topk, expert_alignment)); +)", + do_expand, + cached_mode, + do_zero_padding, + materialize_recv_x, + num_sms, + num_channels, + num_warps, + num_scaleout_ranks, + num_scaleup_ranks, + num_hidden_bytes, + num_sf_packs, + num_max_tokens_per_rank, + num_experts, + num_topk, + expert_alignment)); // Launch - jit->launch( - kernel, { - .stream = stream.stream(), - .num_smem_bytes = materialize_recv_x ? num_smem_bytes : 0, - .grid_dim = dim3(num_sms, 1, 1), - .block_dim = dim3(num_threads, 1, 1), - .enable_pdl = true, - }, - buffer, workspace, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - recv_x, recv_sf, recv_topk_idx, recv_topk_weights, - recv_src_metadata, recv_row_indices, - channel_linked_list, - num_unaligned_recv_tokens_per_expert, - num_recv_tokens, - recv_sf_token_stride, recv_sf_hidden_stride, - scaleout_rank_idx, scaleup_rank_idx - ); + jit->launch(kernel, + { + .stream = stream.stream(), + .num_smem_bytes = materialize_recv_x ? num_smem_bytes : 0, + .grid_dim = dim3(num_sms, 1, 1), + .block_dim = dim3(num_threads, 1, 1), + .enable_pdl = true, + }, + buffer, + workspace, + psum_num_recv_tokens_per_scaleup_rank, + psum_num_recv_tokens_per_expert, + recv_x, + recv_sf, + recv_topk_idx, + recv_topk_weights, + recv_src_metadata, + recv_row_indices, + channel_linked_list, + num_unaligned_recv_tokens_per_expert, + num_recv_tokens, + recv_sf_token_stride, + recv_sf_hidden_stride, + scaleout_rank_idx, + scaleup_rank_idx); } } // namespace deep_ep::ep diff --git a/deep_ep/__init__.py b/deep_ep/__init__.py index 83597c74..b2fc1218 100644 --- a/deep_ep/__init__.py +++ b/deep_ep/__init__.py @@ -1,7 +1,7 @@ import filecmp import glob import os -import torch +import torch # noqa: F401 -- Load PyTorch libraries before checking NCCL. from .utils.find_pkgs import find_nccl_root @@ -55,26 +55,20 @@ def init_jit(): check_nccl_so() init_jit() - # Import APIs after initialization -from . import comm -from .comm import destroy_all_managed_nccl_comm, get_physical_domain_size, get_logical_domain_size -from .buffers.allocator import BufferAllocator -from .buffers.base import BufferBase -from .buffers.ep import EPBuffer, EPHandle -from .buffers.recv_view import DispatchRecvView -from .buffers.engram import EngramBuffer -from .buffers.bucket import BucketBuffer, BucketSession -from .buffers.pp import PPBuffer +from . import comm as comm +from .comm import destroy_all_managed_nccl_comm as destroy_all_managed_nccl_comm, get_physical_domain_size as get_physical_domain_size, get_logical_domain_size as get_logical_domain_size +from .buffers.allocator import BufferAllocator as BufferAllocator +from .buffers.base import BufferBase as BufferBase +from .buffers.ep import EPBuffer as EPBuffer, EPHandle as EPHandle +from .buffers.recv_view import DispatchRecvView as DispatchRecvView +from .buffers.engram import EngramBuffer as EngramBuffer +from .buffers.bucket import BucketBuffer as BucketBuffer, BucketSession as BucketSession +from .buffers.pp import PPBuffer as PPBuffer # noinspection PyUnresolvedReferences -from .utils.event import EventOverlap, EventHandle +from .utils.event import EventOverlap as EventOverlap, EventHandle as EventHandle # noinspection PyUnresolvedReferences -from deep_ep._C import ( - get_num_allocation_alignment, - get_num_rdma_alignment, - get_num_tma_alignment, - topk_idx_t, -) +from deep_ep._C import get_num_allocation_alignment as get_num_allocation_alignment, get_num_rdma_alignment as get_num_rdma_alignment, get_num_tma_alignment as get_num_tma_alignment, topk_idx_t as topk_idx_t __version__ = '2.5.0' diff --git a/deep_ep/buffers/ep.py b/deep_ep/buffers/ep.py index 944e9271..2228e19f 100644 --- a/deep_ep/buffers/ep.py +++ b/deep_ep/buffers/ep.py @@ -16,12 +16,8 @@ from ..utils.event import EventOverlap from ..utils.math import align from ..utils.semantic import value_or, weak_lru -from ..utils.envs import ( - check_fast_rdma_atomic_support, - check_nvlink_connections, check_torch_deterministic, - get_nvlink_gbs, get_rdma_gbs, - get_sm_read_gbs, get_sm_write_gbs -) +from ..utils.envs import (check_fast_rdma_atomic_support, check_nvlink_connections, check_torch_deterministic, get_nvlink_gbs, get_rdma_gbs, + get_sm_read_gbs, get_sm_write_gbs) class EPHandle: @@ -61,22 +57,11 @@ class EPHandle: num_recv_tokens: the total number of received tokens. """ - def __init__(self, - do_expand: bool, - num_experts: int, expert_alignment: int, - num_max_tokens_per_rank: int, - num_sms: int, - topk_idx: torch.Tensor, - num_recv_tokens: int, - num_expanded_tokens: int, - num_recv_tokens_per_expert_list: list, - psum_num_recv_tokens_per_scaleup_rank: torch.Tensor, - psum_num_recv_tokens_per_expert: torch.Tensor, - num_unaligned_recv_tokens_per_expert: torch.Tensor, - recv_src_metadata: torch.Tensor, - dst_buffer_slot_idx: torch.Tensor, - token_metadata_at_forward: Optional[torch.Tensor], - channel_linked_list: Optional[torch.Tensor]): + def __init__(self, do_expand: bool, num_experts: int, expert_alignment: int, num_max_tokens_per_rank: int, num_sms: int, + topk_idx: torch.Tensor, num_recv_tokens: int, num_expanded_tokens: int, num_recv_tokens_per_expert_list: list, + psum_num_recv_tokens_per_scaleup_rank: torch.Tensor, psum_num_recv_tokens_per_expert: torch.Tensor, + num_unaligned_recv_tokens_per_expert: torch.Tensor, recv_src_metadata: torch.Tensor, dst_buffer_slot_idx: torch.Tensor, + token_metadata_at_forward: Optional[torch.Tensor], channel_linked_list: Optional[torch.Tensor]): assert topk_idx is not None self.do_expand = do_expand @@ -183,20 +168,24 @@ def permute(tensor: Optional[torch.Tensor], orig_indices: torch.Tensor): # - `expert_idx * src_token_global_index_max_x2`, for padding slots # This guarantees a two-key sort: first by expert, then by order within each expert. # Valid tokens precede padding tokens, and valid tokens are sorted by `src_token_global_idx`. - src_token_global_index_max_x2 = 10000000000 # 1e10 + src_token_global_index_max_x2 = 10000000000 # 1e10 tensor_dim0_after_expand = recv_x.shape[0] expert_token_idx_start = self.psum_num_recv_tokens_per_expert - self.num_unaligned_recv_tokens_per_expert token_idx2expert_idx = torch.bucketize(torch.arange(tensor_dim0_after_expand, device='cuda'), - expert_token_idx_start[1:], right=True, out_int32=False) + expert_token_idx_start[1:], + right=True, + out_int32=False) sort_keys_for_expanded_tensors = token_idx2expert_idx * src_token_global_index_max_x2 - slots = self.cached_recv_src_metadata_before_sort[:, 2:] # [num_recv_tokens, topk] + slots = self.cached_recv_src_metadata_before_sort[:, 2:] # [num_recv_tokens, topk] src_global_idx = self.cached_recv_src_metadata_before_sort[:, 0] valid_mask = slots >= 0 if not do_cpu_sync: valid_mask[oob_tokens_mask] = False - sort_keys_for_expanded_tensors.scatter_add_(0, slots[valid_mask], -src_token_global_index_max_x2//2 + src_global_idx.unsqueeze(1).expand_as(slots)[valid_mask].to(torch.int64)) + sort_keys_for_expanded_tensors.scatter_add_( + 0, slots[valid_mask], + -src_token_global_index_max_x2 // 2 + src_global_idx.unsqueeze(1).expand_as(slots)[valid_mask].to(torch.int64)) orig_indices_for_expanded_tensors = torch.sort(sort_keys_for_expanded_tensors, stable=True).indices.to(torch.int32) permute(recv_x, orig_indices_for_expanded_tensors) @@ -241,26 +230,28 @@ class EPBuffer(BufferBase): get_physical_domain_size = comm.get_physical_domain_size get_logical_domain_size = comm.get_logical_domain_size - def __init__(self, - group: dist.ProcessGroup, - # Provide `num_bytes` (excludes workspace) - num_bytes: Optional[int] = None, - # Or provide MoE settings (BF16 by default) - num_max_tokens_per_rank: int = 0, - hidden: int = 0, - num_topk: int = 0, - use_fp8_dispatch: bool = False, - # Load balance configs - lb_allocation_plan_or_num_bytes: Union[BufferAllocator, int] = 0, - # Configs - deterministic: bool = False, - allow_hybrid_mode: bool = True, - allow_multiple_reduction: bool = True, - prefer_overlap_with_compute: bool = True, - sl_idx: Optional[int] = None, - num_allocated_qps: int = 0, - num_cpu_timeout_secs: int = 300, num_gpu_timeout_secs: int = 100, - explicitly_destroy: bool = False): + def __init__( + self, + group: dist.ProcessGroup, + # Provide `num_bytes` (excludes workspace) + num_bytes: Optional[int] = None, + # Or provide MoE settings (BF16 by default) + num_max_tokens_per_rank: int = 0, + hidden: int = 0, + num_topk: int = 0, + use_fp8_dispatch: bool = False, + # Load balance configs + lb_allocation_plan_or_num_bytes: Union[BufferAllocator, int] = 0, + # Configs + deterministic: bool = False, + allow_hybrid_mode: bool = True, + allow_multiple_reduction: bool = True, + prefer_overlap_with_compute: bool = True, + sl_idx: Optional[int] = None, + num_allocated_qps: int = 0, + num_cpu_timeout_secs: int = 300, + num_gpu_timeout_secs: int = 100, + explicitly_destroy: bool = False): """ Initialize the EP communication buffer. @@ -303,10 +294,8 @@ def __init__(self, # Calculate buffer size (already 2 MB-aligned from hint functions / calculate_ep_buffer_size) if num_bytes is None: # NOTES: we allow `num_topk == 0`, as the buffer size can also be calculated by number of ranks (maybe bigger though) - num_bytes = _C.calculate_ep_buffer_size( - self.nccl_comm_handle.get(), - num_max_tokens_per_rank, hidden, num_topk, use_fp8_dispatch, - allow_hybrid_mode, allow_multiple_reduction) + num_bytes = _C.calculate_ep_buffer_size(self.nccl_comm_handle.get(), num_max_tokens_per_rank, hidden, num_topk, + use_fp8_dispatch, allow_hybrid_mode, allow_multiple_reduction) if os.environ.get('EP_BUFFER_DEBUG', 0): print(f'Initializing EP buffer with {num_bytes} bytes at rank EP {group.rank()}/{group.size()}') @@ -338,13 +327,9 @@ def __init__(self, # Create CPP handle super().__init__(explicitly_destroy) - self.runtime = _C.EPBuffer( - self.rank_idx, self.num_ranks, - self.nccl_comm_handle.get(), num_bytes, num_lb_bytes, - allow_hybrid_mode, allow_multiple_reduction, prefer_overlap_with_compute, - sl_idx, num_allocated_qps, - num_cpu_timeout_secs, num_gpu_timeout_secs, - self.explicitly_destroy) + self.runtime = _C.EPBuffer(self.rank_idx, self.num_ranks, self.nccl_comm_handle.get(), num_bytes, num_lb_bytes, allow_hybrid_mode, + allow_multiple_reduction, prefer_overlap_with_compute, sl_idx, num_allocated_qps, num_cpu_timeout_secs, + num_gpu_timeout_secs, self.explicitly_destroy) self.context = self.runtime.context # Materialize LB allocation plan @@ -379,8 +364,10 @@ def destroy(self) -> None: @staticmethod def get_buffer_size_hint(group: dist.ProcessGroup, - num_max_tokens_per_rank: int, hidden: int, - num_topk: int = 0, use_fp8_dispatch: bool = False, + num_max_tokens_per_rank: int, + hidden: int, + num_topk: int = 0, + use_fp8_dispatch: bool = False, allow_hybrid_mode: bool = True, allow_multiple_reduction: bool = True) -> int: """ @@ -401,9 +388,8 @@ def get_buffer_size_hint(group: dist.ProcessGroup, """ # NOTES: calculate_ep_buffer_size already returns 2 MB-aligned values return _C.calculate_ep_buffer_size( - comm.get_nccl_comm_handle(group).get(), - num_max_tokens_per_rank, hidden, num_topk, use_fp8_dispatch, - allow_hybrid_mode, allow_multiple_reduction) + comm.get_nccl_comm_handle(group).get(), num_max_tokens_per_rank, hidden, num_topk, use_fp8_dispatch, allow_hybrid_mode, + allow_multiple_reduction) @staticmethod def _unpack_handle(handle: Optional[EPHandle] = None) \ @@ -413,16 +399,10 @@ def _unpack_handle(handle: Optional[EPHandle] = None) \ Optional[torch.Tensor], Optional[torch.Tensor]]: if handle is None: return None, None, None, None, None, None, None, None, None, None - return (handle.num_recv_tokens, - handle.num_expanded_tokens, - handle.num_recv_tokens_per_expert_list, - handle.psum_num_recv_tokens_per_scaleup_rank, - handle.psum_num_recv_tokens_per_expert, - handle.num_unaligned_recv_tokens_per_expert, - handle.dst_buffer_slot_idx, - handle.token_metadata_at_forward, - handle.recv_src_metadata, - handle.channel_linked_list) + return (handle.num_recv_tokens, handle.num_expanded_tokens, handle.num_recv_tokens_per_expert_list, + handle.psum_num_recv_tokens_per_scaleup_rank, handle.psum_num_recv_tokens_per_expert, + handle.num_unaligned_recv_tokens_per_expert, handle.dst_buffer_slot_idx, handle.token_metadata_at_forward, + handle.recv_src_metadata, handle.channel_linked_list) @staticmethod def capture() -> EventHandle: @@ -435,10 +415,14 @@ def capture() -> EventHandle: return EventHandle() @weak_lru(maxsize=None) - def get_theoretical_num_sms(self, num_experts: int, num_topk: int, + def get_theoretical_num_sms(self, + num_experts: int, + num_topk: int, num_scaleout_topk: int = 0, - rdma_gbs: float = 0, nvlink_gbs: float = 0, - sm_read_gbs: float = 0, sm_write_gbs: float = 0) -> int: + rdma_gbs: float = 0, + nvlink_gbs: float = 0, + sm_read_gbs: float = 0, + sm_write_gbs: float = 0) -> int: """ Estimate the optimal number of SMs for dispatch/combine kernels based on bandwidth modeling. The result is cached. This assumes a balanced gate distribution. @@ -650,9 +634,8 @@ def dispatch(self, `(recv_x, recv_topk_idx, recv_topk_weights, handle)`. """ self._check_recv_view_released() - if borrow_recv and (handle is not None or do_expand or defer_epilogue or do_cpu_sync is False or - not isinstance(x, torch.Tensor) or x.dtype != torch.bfloat16 or - self.num_scaleout_ranks != 1 or self.num_rdma_ranks != 1): + if borrow_recv and (handle is not None or do_expand or defer_epilogue or do_cpu_sync is False or not isinstance(x, torch.Tensor) + or x.dtype != torch.bfloat16 or self.num_scaleout_ranks != 1 or self.num_rdma_ranks != 1): raise ValueError('borrow_recv requires fresh compact BF16 dispatch, CPU counts, one NVLink domain, and no deferred epilogue') assert not do_handle_copy, '`do_handle_copy` must be False; handle copying is no longer supported' check_torch_deterministic() @@ -680,13 +663,9 @@ def dispatch(self, # Should be aligned with the handle context assert (num_experts, expert_alignment, num_max_tokens_per_rank) == \ (handle.num_experts, handle.expert_alignment, handle.num_max_tokens_per_rank) - (cached_num_recv_tokens, cached_num_expanded_tokens, - cached_num_recv_tokens_per_expert_list, - cached_psum_num_recv_tokens_per_scaleup_rank, cached_psum_num_recv_tokens_per_expert, - cached_num_unaligned_recv_tokens_per_expert, - cached_dst_buffer_slot_idx, - cached_token_metadata_at_forward, - cached_recv_src_metadata, + (cached_num_recv_tokens, cached_num_expanded_tokens, cached_num_recv_tokens_per_expert_list, + cached_psum_num_recv_tokens_per_scaleup_rank, cached_psum_num_recv_tokens_per_expert, cached_num_unaligned_recv_tokens_per_expert, + cached_dst_buffer_slot_idx, cached_token_metadata_at_forward, cached_recv_src_metadata, cached_channel_linked_list) = self._unpack_handle(handle) # Some default values @@ -695,59 +674,27 @@ def dispatch(self, do_cpu_sync = value_or(do_cpu_sync, True) # Do dispatch - result, event, deferred_epilogue = self.runtime.dispatch(x, sf, topk_idx, topk_weights, - cumulative_local_expert_recv_stats, - cached_num_recv_tokens, - cached_num_expanded_tokens, - cached_num_recv_tokens_per_expert_list, - cached_psum_num_recv_tokens_per_scaleup_rank, - cached_psum_num_recv_tokens_per_expert, - cached_num_unaligned_recv_tokens_per_expert, - cached_dst_buffer_slot_idx, - cached_token_metadata_at_forward, - cached_recv_src_metadata, - cached_channel_linked_list, - num_max_tokens_per_rank, - num_experts, expert_alignment, - num_sms, num_qps, - previous_event, - async_with_compute_stream, allocate_on_comm_stream, - do_cpu_sync, do_expand, - do_zero_padding, - use_tma_aligned_col_major_sf, - defer_epilogue, not borrow_recv) + result, event, deferred_epilogue = self.runtime.dispatch( + x, sf, topk_idx, topk_weights, cumulative_local_expert_recv_stats, cached_num_recv_tokens, cached_num_expanded_tokens, + cached_num_recv_tokens_per_expert_list, cached_psum_num_recv_tokens_per_scaleup_rank, cached_psum_num_recv_tokens_per_expert, + cached_num_unaligned_recv_tokens_per_expert, cached_dst_buffer_slot_idx, cached_token_metadata_at_forward, + cached_recv_src_metadata, cached_channel_linked_list, num_max_tokens_per_rank, num_experts, expert_alignment, num_sms, num_qps, + previous_event, async_with_compute_stream, allocate_on_comm_stream, do_cpu_sync, do_expand, do_zero_padding, + use_tma_aligned_col_major_sf, defer_epilogue, not borrow_recv) event_overlap = EventOverlap(event) def finalize_dispatch(dispatch_result: tuple, deterministic_by_hook: bool): - (recv_x, recv_sf, - recv_topk_idx, recv_topk_weights, - num_recv_tokens, num_expanded_tokens, - num_recv_tokens_per_expert_list, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - num_unaligned_recv_tokens_per_expert, - recv_src_metadata, - dst_buffer_slot_idx, - token_metadata_at_forward, - channel_linked_list, recv_row_indices) = dispatch_result + (recv_x, recv_sf, recv_topk_idx, recv_topk_weights, num_recv_tokens, num_expanded_tokens, num_recv_tokens_per_expert_list, + psum_num_recv_tokens_per_scaleup_rank, psum_num_recv_tokens_per_expert, num_unaligned_recv_tokens_per_expert, + recv_src_metadata, dst_buffer_slot_idx, token_metadata_at_forward, channel_linked_list, recv_row_indices) = dispatch_result # Create handle if not cached nonlocal handle is_cached_dispatch = handle is not None - handle = EPHandle(do_expand, - num_experts, expert_alignment, - num_max_tokens_per_rank, - num_sms, - topk_idx, - num_recv_tokens, num_expanded_tokens, - num_recv_tokens_per_expert_list, - psum_num_recv_tokens_per_scaleup_rank, - psum_num_recv_tokens_per_expert, - num_unaligned_recv_tokens_per_expert, - recv_src_metadata, - dst_buffer_slot_idx, - token_metadata_at_forward, - channel_linked_list) if handle is None else handle + handle = EPHandle(do_expand, num_experts, expert_alignment, num_max_tokens_per_rank, num_sms, topk_idx, num_recv_tokens, + num_expanded_tokens, num_recv_tokens_per_expert_list, psum_num_recv_tokens_per_scaleup_rank, + psum_num_recv_tokens_per_expert, num_unaligned_recv_tokens_per_expert, recv_src_metadata, dst_buffer_slot_idx, + token_metadata_at_forward, channel_linked_list) if handle is None else handle recv_view = None if borrow_recv: @@ -757,9 +704,8 @@ def finalize_dispatch(dispatch_result: tuple, deterministic_by_hook: bool): def complete_dispatch(): if self.deterministic: - handle.deterministic_sort( - do_cpu_sync, is_cached_dispatch, recv_x, recv_sf, - recv_topk_idx, recv_topk_weights, channel_linked_list, recv_row_indices) + handle.deterministic_sort(do_cpu_sync, is_cached_dispatch, recv_x, recv_sf, recv_topk_idx, recv_topk_weights, + channel_linked_list, recv_row_indices) if recv_view is not None: recv_view._mark_ready() @@ -774,8 +720,7 @@ def complete_dispatch(): # Just launch the dispatch if deferred_epilogue is not None: - event_overlap.register_hook_after_wait( - lambda: finalize_dispatch(deferred_epilogue(), False)) + event_overlap.register_hook_after_wait(lambda: finalize_dispatch(deferred_epilogue(), False)) return event_overlap # Do epilogue ASAP @@ -845,21 +790,12 @@ def combine(self, assert num_qps <= self.num_allocated_qps, 'Allocated QPs are not enough' bias_0, bias_1 = EPBuffer._unpack_bias(bias) - result, event, deferred_epilogue = self.runtime.combine(x, topk_weights, - bias_0, bias_1, - handle.recv_src_metadata, - handle.topk_idx, + result, event, deferred_epilogue = self.runtime.combine(x, topk_weights, bias_0, bias_1, handle.recv_src_metadata, handle.topk_idx, handle.psum_num_recv_tokens_per_scaleup_rank, - handle.token_metadata_at_forward, - handle.channel_linked_list, - handle.num_experts, - handle.num_max_tokens_per_rank, - num_sms, num_qps, - previous_event, - async_with_compute_stream, - allocate_on_comm_stream, - handle.do_expand, - defer_epilogue) + handle.token_metadata_at_forward, handle.channel_linked_list, + handle.num_experts, handle.num_max_tokens_per_rank, num_sms, num_qps, + previous_event, async_with_compute_stream, allocate_on_comm_stream, + handle.do_expand, defer_epilogue) event_overlap = EventOverlap(event) if deferred_epilogue is not None: event_overlap.register_hook_after_wait(deferred_epilogue) @@ -903,13 +839,13 @@ def lb_prefetch_weights(self, previous_event: the event to wait for before communication; defaults to waiting for the current stream """ self._check_recv_view_released() - redundant_expert_weights = ([redundant_expert_weights] if isinstance(redundant_expert_weights, torch.Tensor) - else list(redundant_expert_weights)) + redundant_expert_weights = ([redundant_expert_weights] + if isinstance(redundant_expert_weights, torch.Tensor) else list(redundant_expert_weights)) expert_weights = [expert_weights] if isinstance(expert_weights, torch.Tensor) else list(expert_weights) num_sms = self.lb_get_theoretical_num_sms() if num_sms == 0 else align(num_sms, 2) - return EventOverlap(self.runtime.lb_prefetch_weights( - redundant_expert_weights, expert_weights, redundancy_mapping, num_sms, previous_event)) + return EventOverlap( + self.runtime.lb_prefetch_weights(redundant_expert_weights, expert_weights, redundancy_mapping, num_sms, previous_event)) def lb_reduce_grads(self, redundant_expert_grads: torch.Tensor, @@ -931,5 +867,4 @@ def lb_reduce_grads(self, """ self._check_recv_view_released() num_sms = self.lb_get_theoretical_num_sms() if num_sms == 0 else align(num_sms, 2) - return EventOverlap(self.runtime.lb_reduce_grads( - redundant_expert_grads, expert_grads, redundancy_mapping, num_sms, previous_event)) + return EventOverlap(self.runtime.lb_reduce_grads(redundant_expert_grads, expert_grads, redundancy_mapping, num_sms, previous_event)) diff --git a/deep_ep/buffers/recv_view.py b/deep_ep/buffers/recv_view.py index f9c45749..c848c58c 100644 --- a/deep_ep/buffers/recv_view.py +++ b/deep_ep/buffers/recv_view.py @@ -44,7 +44,7 @@ def row_indices(self): def release(self, *consumer_streams): self.wait() if not consumer_streams: - consumer_streams = (torch.cuda.current_stream(self._slab.device),) + consumer_streams = (torch.cuda.current_stream(self._slab.device), ) comm_stream = comm.get_comm_stream(self._owner) for stream in consumer_streams: if stream.device != self._slab.device: diff --git a/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh b/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh index d1fc818b..5d3ed618 100644 --- a/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh +++ b/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh @@ -6,32 +6,45 @@ #include #include - namespace deep_ep::ep { -template 1 and not kCachedMode)> -__global__ void __launch_bounds__(kNumThreads, 1) -dispatch_copy_epilogue_impl(void* buffer, void* workspace, - int* psum_num_recv_tokens_per_scaleup_rank, - int* psum_num_recv_tokens_per_expert, - void* recv_x, sf_pack_t* recv_sf, - topk_idx_t* recv_topk_idx, float* recv_topk_weights, - int* recv_src_metadata, int64_t* recv_row_indices, - int* channel_linked_list, - int* num_unaligned_recv_tokens_per_expert, - int num_recv_tokens, - const int recv_sf_token_stride, const int recv_sf_hidden_stride, - const int scaleout_rank_idx, const int scaleup_rank_idx) { +__global__ void __launch_bounds__(kNumThreads, 1) dispatch_copy_epilogue_impl(void* buffer, + void* workspace, + int* psum_num_recv_tokens_per_scaleup_rank, + int* psum_num_recv_tokens_per_expert, + void* recv_x, + sf_pack_t* recv_sf, + topk_idx_t* recv_topk_idx, + float* recv_topk_weights, + int* recv_src_metadata, + int64_t* recv_row_indices, + int* channel_linked_list, + int* num_unaligned_recv_tokens_per_expert, + int num_recv_tokens, + const int recv_sf_token_stride, + const int recv_sf_hidden_stride, + const int scaleout_rank_idx, + const int scaleup_rank_idx) { EP_STATIC_ASSERT(kMaterializeRecvX or (not kDoExpand and kNumSFPacks == 0 and kNumScaleoutRanks == 1), "Borrowed receive only supports compact, intra-node BF16 dispatch"); // Utils @@ -47,9 +60,9 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // Buffer layouts extern __shared__ __align__(kNumTMAAlignmentBytes) int8_t smem[]; const auto token_layout = layout::TokenLayout(kNumHiddenBytes, kNumSFPacks * sizeof(sf_pack_t), kNumTopk, true); - const auto tma_buffer = layout::BufferLayout(token_layout, kNumWarps, 1, smem) - .get_rank_buffer(warp_idx).get_token_buffer(0); - const auto scaleup_buffer = layout::BufferLayout(token_layout, kNumScaleupRanks, kNumScaleoutRanks * kNumMaxTokensPerRank, buffer); + const auto tma_buffer = layout::BufferLayout(token_layout, kNumWarps, 1, smem).get_rank_buffer(warp_idx).get_token_buffer(0); + const auto scaleup_buffer = + layout::BufferLayout(token_layout, kNumScaleupRanks, kNumScaleoutRanks * kNumMaxTokensPerRank, buffer); // Init TMA ptx::arrival_phase phase = 0; @@ -89,7 +102,6 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, recv_row_indices[i] = static_cast(current_rank_idx) * kNumMaxTokensPerRank + i - current_rank_start; } - // Wait buffer releases if constexpr (kMaterializeRecvX) ptx::tma_store_wait(); @@ -99,8 +111,7 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // Including all stuffs: data, SF, top-k metadata if constexpr (kMaterializeRecvX) { if (ptx::elect_one_sync()) { - ptx::tma_load_1d(tma_buffer.get_base_ptr(), buffer_token.get_base_ptr(), - mbarrier_ptr, tma_buffer.get_num_bytes()); + ptx::tma_load_1d(tma_buffer.get_base_ptr(), buffer_token.get_base_ptr(), mbarrier_ptr, tma_buffer.get_num_bytes()); ptx::mbarrier_arrive_and_set_tx(mbarrier_ptr, tma_buffer.get_num_bytes()); } } @@ -157,7 +168,8 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, if constexpr (kMaterializeRecvX) { if (kDoExpand ? (dst_tensor_idx >= 0) : ptx::elect_one_sync()) { ptx::tma_store_1d(math::advance_ptr(recv_x, static_cast(dst_tensor_idx) * kNumHiddenBytes), - tma_buffer.get_hidden_ptr(), kNumHiddenBytes); + tma_buffer.get_hidden_ptr(), + kNumHiddenBytes); ptx::tma_store_commit(); } } @@ -173,7 +185,7 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, const auto smem_src_ptr = tma_buffer.get_sf_ptr(); sf_pack_t reg_src[kNumFullIters + 1]; #pragma unroll - for (int k = 0; k < kNumFullIters; ++ k) + for (int k = 0; k < kNumFullIters; ++k) reg_src[k] = smem_src_ptr[k * 32 + lane_idx]; if (do_last_iter) reg_src[kNumFullIters] = smem_src_ptr[kNumFullIters * 32 + lane_idx]; @@ -186,10 +198,10 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, auto mask = kDoExpand ? ptx::gather(dst_tensor_idx >= 0) : 1; while (mask) { const int valid_lane_idx = __ffs(mask) - 1; - const auto gmem_dst = math::advance_ptr(recv_sf, - ptx::exchange(dst_tensor_idx, valid_lane_idx) * (recv_sf_token_stride_i64 * sizeof(sf_pack_t))); + const auto gmem_dst = math::advance_ptr( + recv_sf, ptx::exchange(dst_tensor_idx, valid_lane_idx) * (recv_sf_token_stride_i64 * sizeof(sf_pack_t))); #pragma unroll - for (int k = 0; k < kNumFullIters; ++ k) + for (int k = 0; k < kNumFullIters; ++k) gmem_dst[(k * 32 + lane_idx) * recv_sf_hidden_stride_i64] = reg_src[k]; if (do_last_iter) gmem_dst[(kNumFullIters * 32 + lane_idx) * recv_sf_hidden_stride_i64] = reg_src[kNumFullIters]; @@ -235,11 +247,9 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, const auto workspace_layout = layout::EPWorkspaceLayout(workspace, kNumScaleoutRanks, kNumScaleupRanks, kNumExperts); for (int i = global_warp_idx; i < kNumChannels; i += kNumSMs * kNumWarps) { #pragma unroll - for (int j = 0; j < kNumScaleupRanksPerLane; ++ j) { + for (int j = 0; j < kNumScaleupRanksPerLane; ++j) { if (const auto k = j * 32 + lane_idx; j < (kNumScaleupRanksPerLane - 1) or k < kNumScaleupRanks) { - channel_linked_list[ - *workspace_layout.get_channel_scaleup_tail_ptr(i, k) - ] = -1; + channel_linked_list[*workspace_layout.get_channel_scaleup_tail_ptr(i, k)] = -1; // Clean for combine usages *workspace_layout.get_channel_scaleup_tail_ptr(i, k) = 0; @@ -264,10 +274,9 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, constexpr int kNumExpertsPerLane = math::constexpr_ceil_div(kNumExpertsPerRank, 32); int num_experts_per_lane[kNumExpertsPerLane]; #pragma unroll - for (int i = 0; i < kNumExpertsPerLane; ++ i) { + for (int i = 0; i < kNumExpertsPerLane; ++i) { const int expert_idx = i * 32 + lane_idx; - num_experts_per_lane[i] = expert_idx < kNumExpertsPerRank ? - num_unaligned_recv_tokens_per_expert[expert_idx] : 0; + num_experts_per_lane[i] = expert_idx < kNumExpertsPerRank ? num_unaligned_recv_tokens_per_expert[expert_idx] : 0; } // Single while loop over all padding tokens @@ -279,7 +288,7 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // Assign current wave expert token count int wave_num_experts_per_lane; #pragma unroll - for (int i = 0; i < kNumExpertsPerLane; ++ i) + for (int i = 0; i < kNumExpertsPerLane; ++i) wave_num_experts_per_lane = i == wave_idx ? num_experts_per_lane[i] : wave_num_experts_per_lane; int wave_num_pads_per_lane = 0; if (wave_idx * 32 + lane_idx < kNumExpertsPerRank) { @@ -293,15 +302,15 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, const int local_pad_idx = pad_idx - wave_pad_psum; const int pad_psum = ptx::warp_inclusive_sum(wave_num_pads_per_lane, lane_idx); const int owner_lane_idx = __ffs(__ballot_sync(0xffffffff, pad_psum > local_pad_idx)) - 1; - dst_tensor_idx = wave_tensor_psum + ptx::exchange( - ptx::warp_exclusive_sum(wave_num_experts_per_lane + wave_num_pads_per_lane, lane_idx) + - wave_num_experts_per_lane + local_pad_idx - (pad_psum - wave_num_pads_per_lane), - owner_lane_idx); + dst_tensor_idx = wave_tensor_psum + + ptx::exchange(ptx::warp_exclusive_sum(wave_num_experts_per_lane + wave_num_pads_per_lane, lane_idx) + + wave_num_experts_per_lane + local_pad_idx - (pad_psum - wave_num_pads_per_lane), + owner_lane_idx); break; } // Move to the next wave - wave_idx ++; + wave_idx++; wave_pad_psum += wave_num_pads; wave_tensor_psum += ptx::reduce_add(wave_num_experts_per_lane + wave_num_pads_per_lane); } @@ -311,7 +320,8 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, // Zero data via TMA store if (ptx::elect_one_sync()) { ptx::tma_store_1d(math::advance_ptr(recv_x, static_cast(dst_tensor_idx) * kNumHiddenBytes), - tma_buffer.get_hidden_ptr(), kNumHiddenBytes); + tma_buffer.get_hidden_ptr(), + kNumHiddenBytes); ptx::tma_store_commit(); } __syncwarp(); @@ -326,11 +336,11 @@ dispatch_copy_epilogue_impl(void* buffer, void* workspace, const auto recv_sf_token_stride_i64 = static_cast(recv_sf_token_stride); const auto recv_sf_hidden_stride_i64 = static_cast(recv_sf_hidden_stride); constexpr sf_pack_t zero_sf_pack = {0}; - const auto gmem_dst = math::advance_ptr(recv_sf, - dst_tensor_idx * (recv_sf_token_stride_i64 * sizeof(sf_pack_t))); + const auto gmem_dst = + math::advance_ptr(recv_sf, dst_tensor_idx * (recv_sf_token_stride_i64 * sizeof(sf_pack_t))); constexpr auto kNumFullIters = kNumSFPacks / 32; #pragma unroll - for (int k = 0; k < kNumFullIters; ++ k) + for (int k = 0; k < kNumFullIters; ++k) gmem_dst[(k * 32 + lane_idx) * recv_sf_hidden_stride_i64] = zero_sf_pack; if constexpr (kNumSFPacks % 32 != 0) { if (kNumFullIters * 32 + lane_idx < kNumSFPacks) diff --git a/tests/ep/test_recv_view.py b/tests/ep/test_recv_view.py index ea7ae1e1..cb80f5ac 100644 --- a/tests/ep/test_recv_view.py +++ b/tests/ep/test_recv_view.py @@ -34,11 +34,17 @@ def test_case(deep_ep, deterministic, asynchronous, with_weights, skew): gathered = [torch.empty_like(full_x) for _ in range(world)] dist.all_gather(gathered, full_x) all_x = torch.stack(gathered) - buffer = deep_ep.EPBuffer(dist.group.WORLD, num_max_tokens_per_rank=capacity, - hidden=hidden, num_topk=topk, deterministic=deterministic, + buffer = deep_ep.EPBuffer(dist.group.WORLD, + num_max_tokens_per_rank=capacity, + hidden=hidden, + num_topk=topk, + deterministic=deterministic, explicitly_destroy=True) - kwargs = dict(topk_idx=indices, topk_weights=weights, num_experts=local_experts * world, - num_sms=8, async_with_compute_stream=asynchronous, + kwargs = dict(topk_idx=indices, + topk_weights=weights, + num_experts=local_experts * world, + num_sms=8, + async_with_compute_stream=asynchronous, allocate_on_comm_stream=asynchronous) def dispatch(borrow, value=x): @@ -68,18 +74,14 @@ def dispatch(borrow, value=x): rejected(lambda: view.slab, RuntimeError) rejected(view.release, RuntimeError) - replay, _, _, _, event = buffer.dispatch(x, handle=handle, num_sms=8, - async_with_compute_stream=True) + replay, _, _, _, event = buffer.dispatch(x, handle=handle, num_sms=8, async_with_compute_stream=True) event.current_stream_wait() exact(replay, payload) # Integer sums are exact even with an unspecified receive order. - combined, _, event = buffer.combine(torch.ones_like(replay), handle, - async_with_compute_stream=True) + combined, _, event = buffer.combine(torch.ones_like(replay), handle, async_with_compute_stream=True) event.current_stream_wait() - expected_count = torch.stack([ - ((indices >= 0) & (indices // local_experts == peer)).any(dim=1) - for peer in range(world) - ]).sum(dim=0).to(x.dtype) + expected_count = torch.stack([((indices >= 0) & (indices // local_experts == peer)).any(dim=1) + for peer in range(world)]).sum(dim=0).to(x.dtype) exact(combined, expected_count[:, None].expand_as(x)) # Consumers can move to another stream after the dispatch event is waited. From 227e434437d19a5c70081e130cf2febf265b9536 Mon Sep 17 00:00:00 2001 From: hyw Date: Wed, 30 Sep 2026 15:10:45 +0800 Subject: [PATCH 3/5] Clarify receive-view waits for synchronous dispatch --- README.md | 2 +- deep_ep/buffers/ep.py | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index 4fc015b3..2ab4f0a2 100644 --- a/README.md +++ b/README.md @@ -280,7 +280,7 @@ Training saves the forward `handle` for `combine_backward` and `dispatch_backwar For a fresh compact BF16 dispatch within one physical NVLink domain, `borrow_recv=True` returns a `DispatchRecvView` in place of `recv_x`. Logical receive row `r` is stored at `view.slab[view.row_indices[r]]`. Pass the slab and row map directly to an indexed consumer to avoid the activation copy in the dispatch epilogue. This mode requires exact CPU receive counts and `defer_epilogue=False`. -Wait on the returned dispatch event before accessing the view. After enqueueing all consumers, call `view.release()` on the consuming stream, or pass every consuming CUDA stream to `view.release(*streams)`. Release orders subsequent communication after those consumers. Dispatch, combine, load-balancing calls, and explicit buffer destruction reject an active view. Keep host API calls serialized and stop using raw slab references after release. +With `async_with_compute_stream=True`, call `event.current_stream_wait()` before accessing the view. The default synchronous dispatch establishes the stream dependency before returning and needs no event wait. After enqueueing all consumers, call `view.release()` on the consuming stream, or pass every consuming CUDA stream to `view.release(*streams)`. Release orders subsequent communication after those consumers. Dispatch, combine, load-balancing calls, and explicit buffer destruction reject an active view. Keep host API calls serialized and stop using raw slab references after release. Borrowed storage is intended for forward consumers; retaining it for backward is unsupported. After release, the returned handle can be used for ordinary materialized cached replay. Expanded layouts, FP8, cached borrowed dispatch, deferred epilogues, and multiple NVLink domains are outside this initial API. The option defaults to `False`; whether it reduces latency depends on the consumer and token count. diff --git a/deep_ep/buffers/ep.py b/deep_ep/buffers/ep.py index 2228e19f..5fb95479 100644 --- a/deep_ep/buffers/ep.py +++ b/deep_ep/buffers/ep.py @@ -614,7 +614,8 @@ def dispatch(self, use_tma_aligned_col_major_sf: whether to use TMA-aligned column-major layout for scale factors. borrow_recv: return a DispatchRecvView instead of a materialized receive tensor. Requires fresh compact BF16 dispatch in one NVLink domain, exact CPU counts, and - defer_epilogue=False. Wait the dispatch event, enqueue indexed consumers, then + defer_epilogue=False. Wait the dispatch event only for asynchronous dispatch, + enqueue indexed consumers, then release the view with every consuming stream before reusing or destroying this buffer. The view cannot be retained for backward; a cached materialized replay is supported. defer_epilogue: whether to defer the CPU receive-count wait and copy epilogue until From 4a2f5b3b95c4c0288110369b1971ab0063e6264c Mon Sep 17 00:00:00 2001 From: hyw Date: Wed, 30 Sep 2026 16:04:37 +0800 Subject: [PATCH 4/5] Trim borrowed receive documentation and redundant tests --- README.md | 2 -- deep_ep/buffers/ep.py | 10 ++++------ deep_ep/buffers/recv_view.py | 9 +++------ .../deep_ep/impls/ep/dispatch_copy_epilogue.cuh | 3 +-- tests/ep/test_recv_view.py | 12 +----------- 5 files changed, 9 insertions(+), 27 deletions(-) diff --git a/README.md b/README.md index 2ab4f0a2..04bc3f73 100644 --- a/README.md +++ b/README.md @@ -284,8 +284,6 @@ With `async_with_compute_stream=True`, call `event.current_stream_wait()` before Borrowed storage is intended for forward consumers; retaining it for backward is unsupported. After release, the returned handle can be used for ordinary materialized cached replay. Expanded layouts, FP8, cached borrowed dispatch, deferred epilogues, and multiple NVLink domains are outside this initial API. The option defaults to `False`; whether it reduces latency depends on the consumer and token count. -See `tests/ep/test_recv_view.py` for event ordering, cross-stream lifetime, replay, combine, and layout checks. - ### Deferred epilogues Dispatch and combine accept `defer_epilogue=True` together with `async_with_compute_stream=True`. In this mode, the call returns an `EventOverlap` directly. Calling `.wait()` runs the deferred epilogue on the current stream and returns `(recv_x, recv_topk_idx, recv_topk_weights, handle)` for dispatch, or `(combined_x, combined_topk_weights)` for combine. diff --git a/deep_ep/buffers/ep.py b/deep_ep/buffers/ep.py index 5fb95479..5aef5654 100644 --- a/deep_ep/buffers/ep.py +++ b/deep_ep/buffers/ep.py @@ -612,12 +612,10 @@ def dispatch(self, do_zero_padding: whether to zero out the alignment padding slots in the expanded output. Only valid when `do_expand` is True. Ensures alignment gaps between experts are zeroed. use_tma_aligned_col_major_sf: whether to use TMA-aligned column-major layout for scale factors. - borrow_recv: return a DispatchRecvView instead of a materialized receive tensor. - Requires fresh compact BF16 dispatch in one NVLink domain, exact CPU counts, and - defer_epilogue=False. Wait the dispatch event only for asynchronous dispatch, - enqueue indexed consumers, then - release the view with every consuming stream before reusing or destroying this buffer. - The view cannot be retained for backward; a cached materialized replay is supported. + borrow_recv: return a forward-only DispatchRecvView instead of a receive tensor. + Requires fresh compact BF16 dispatch, exact CPU counts, one physical NVLink domain, + and defer_epilogue=False. For asynchronous dispatch, wait the returned event first. + Release on all consuming streams before reusing or destroying this buffer. defer_epilogue: whether to defer the CPU receive-count wait and copy epilogue until `event.current_stream_wait()` is called. This requires `async_with_compute_stream=True`. diff --git a/deep_ep/buffers/recv_view.py b/deep_ep/buffers/recv_view.py index c848c58c..8ae9cbfe 100644 --- a/deep_ep/buffers/recv_view.py +++ b/deep_ep/buffers/recv_view.py @@ -1,16 +1,13 @@ -"""An exclusive, forward-only view of compact BF16 dispatch payloads.""" import torch from .. import comm class DispatchRecvView: - """Read logical row ``r`` as ``slab[row_indices[r]]`` after dispatch wait. + """Forward-only receive storage: logical row ``r`` is ``slab[row_indices[r]]``. - The payload belongs to the communication buffer. Keep host calls serialized - and release this view after enqueueing all consumers, passing every consuming - CUDA stream. Copies of ``slab`` references are invalid after release. Saving - this view for backward is unsupported; use a materialized cached replay. + Keep host calls serialized. After enqueueing all consumers, release with every + consuming CUDA stream. Raw slab references become invalid after release. """ def __init__(self, owner, slab, row_indices): diff --git a/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh b/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh index 5d3ed618..b4eb99f3 100644 --- a/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh +++ b/deep_ep/include/deep_ep/impls/ep/dispatch_copy_epilogue.cuh @@ -153,8 +153,7 @@ __global__ void __launch_bounds__(kNumThreads, 1) dispatch_copy_epilogue_impl(vo } __syncwarp(); - // Metadata-only mode must skip the activation LOAD as well as its - // store. Read the small metadata fields directly after the PDL fence. + // Borrowed mode reads metadata directly from the buffer after the PDL fence. const auto metadata_token = kMaterializeRecvX ? tma_buffer : buffer_token; // Maintain linked list diff --git a/tests/ep/test_recv_view.py b/tests/ep/test_recv_view.py index cb80f5ac..b5aa034d 100644 --- a/tests/ep/test_recv_view.py +++ b/tests/ep/test_recv_view.py @@ -55,7 +55,6 @@ def dispatch(borrow, value=x): materialized, ref_ids, ref_weights, ref_handle = dispatch(False) view, recv_ids, recv_weights, handle = dispatch(True) - assert isinstance(view, deep_ep.DispatchRecvView) payload = view.slab.index_select(0, view.row_indices) source = handle.recv_src_metadata[:, 0].long() exact(payload, all_x[source // capacity, source % capacity]) @@ -72,7 +71,6 @@ def dispatch(borrow, value=x): rejected(buffer.destroy, RuntimeError) view.release() rejected(lambda: view.slab, RuntimeError) - rejected(view.release, RuntimeError) replay, _, _, _, event = buffer.dispatch(x, handle=handle, num_sms=8, async_with_compute_stream=True) event.current_stream_wait() @@ -84,7 +82,7 @@ def dispatch(borrow, value=x): for peer in range(world)]).sum(dim=0).to(x.dtype) exact(combined, expected_count[:, None].expand_as(x)) - # Consumers can move to another stream after the dispatch event is waited. + # Release must order buffer reuse after the delayed reader. view, _, _, delayed_handle = dispatch(True) delayed_source = delayed_handle.recv_src_metadata[:, 0].long() delayed_expected = all_x[delayed_source // capacity, delayed_source % capacity] @@ -93,7 +91,6 @@ def dispatch(borrow, value=x): torch.cuda._sleep(4_000_000) delayed = view.slab.index_select(0, view.row_indices) view.release(side) - # Overwrite the receive buffer before synchronizing the delayed consumer. dispatch(False, x + 1) side.synchronize() exact(delayed, delayed_expected) @@ -101,13 +98,6 @@ def dispatch(borrow, value=x): for extra in [dict(do_expand=True), dict(defer_epilogue=True), dict(do_cpu_sync=False), dict(handle=handle)]: rejected(lambda extra=extra: buffer.dispatch(x, borrow_recv=True, **(kwargs | extra)), ValueError) rejected(lambda: buffer.dispatch(x.float(), borrow_recv=True, **kwargs), ValueError) - # One logical scaleup domain is insufficient when hybrid mode is disabled. - original_rdma_ranks = buffer.num_rdma_ranks - try: - buffer.num_rdma_ranks = 2 - rejected(lambda: buffer.dispatch(x, borrow_recv=True, **kwargs), ValueError) - finally: - buffer.num_rdma_ranks = original_rdma_ranks dist.barrier() buffer.destroy() From 0ba497cad188eb741147fadb242a35f750dec574 Mon Sep 17 00:00:00 2001 From: hyw Date: Wed, 30 Sep 2026 16:55:44 +0800 Subject: [PATCH 5/5] Test receive views against original routes and concurrent consumers --- README.md | 2 + tests/ep/test_recv_view.py | 84 +++++++++++++++++++++++++------------- 2 files changed, 57 insertions(+), 29 deletions(-) diff --git a/README.md b/README.md index 04bc3f73..f215cc75 100644 --- a/README.md +++ b/README.md @@ -284,6 +284,8 @@ With `async_with_compute_stream=True`, call `event.current_stream_wait()` before Borrowed storage is intended for forward consumers; retaining it for backward is unsupported. After release, the returned handle can be used for ordinary materialized cached replay. Expanded layouts, FP8, cached borrowed dispatch, deferred epilogues, and multiple NVLink domains are outside this initial API. The option defaults to `False`; whether it reduces latency depends on the consumer and token count. +The exclusive lease makes the next dispatch on the same buffer wait for its consumers; performance gains with pipelined communication and computation have not been measured. + ### Deferred epilogues Dispatch and combine accept `defer_epilogue=True` together with `async_with_compute_stream=True`. In this mode, the call returns an `EventOverlap` directly. Calling `.wait()` runs the deferred epilogue on the current stream and returns `(recv_x, recv_topk_idx, recv_topk_weights, handle)` for dispatch, or `(combined_x, combined_topk_weights)` for combine. diff --git a/tests/ep/test_recv_view.py b/tests/ep/test_recv_view.py index b5aa034d..05723ca3 100644 --- a/tests/ep/test_recv_view.py +++ b/tests/ep/test_recv_view.py @@ -7,7 +7,8 @@ def exact(actual, expected): assert actual.shape == expected.shape and actual.dtype == expected.dtype - assert torch.equal(actual.contiguous().view(torch.uint8), expected.contiguous().view(torch.uint8)) + if actual.numel(): + assert torch.equal(actual.contiguous().view(torch.uint8), expected.contiguous().view(torch.uint8)) def rejected(fn, error): @@ -25,15 +26,41 @@ def test_case(deep_ep, deterministic, asynchronous, with_weights, skew): tokens = 65 - rank full_x = torch.randn((capacity, hidden), device='cuda', dtype=torch.bfloat16) x = full_x[:tokens] - indices = torch.rand((tokens, local_experts * world), device='cuda').argsort(dim=1)[:, :topk] - indices = indices.to(deep_ep.topk_idx_t).contiguous() + full_indices = torch.full((capacity, topk), -1, device='cuda', dtype=deep_ep.topk_idx_t) + indices = full_indices[:tokens] + indices.copy_(torch.rand((tokens, local_experts * world), device='cuda').argsort(dim=1)[:, :topk]) if skew: - indices = torch.arange(topk, device='cuda', dtype=deep_ep.topk_idx_t).repeat(tokens, 1) + indices[:] = torch.arange(topk, device='cuda', dtype=deep_ep.topk_idx_t) indices[::7] = -1 - weights = torch.randn((tokens, topk), device='cuda', dtype=torch.float32) if with_weights else None - gathered = [torch.empty_like(full_x) for _ in range(world)] - dist.all_gather(gathered, full_x) - all_x = torch.stack(gathered) + full_weights = torch.randn((capacity, topk), device='cuda', dtype=torch.float32) if with_weights else None + weights = full_weights[:tokens] if with_weights else None + + def gather(tensor): + gathered = [torch.empty_like(tensor) for _ in range(world)] + dist.all_gather(gathered, tensor) + return torch.cat(gathered) + + all_x, all_indices = gather(full_x), gather(full_indices) + local_routes = (all_indices >= rank * local_experts) & (all_indices < (rank + 1) * local_experts) + expected_source = torch.arange(world * capacity, device='cuda')[local_routes.any(dim=1)] + expected_lane = torch.where(local_routes[expected_source], torch.arange(topk, device='cuda'), -1).amax(dim=1) + expected_peer_topk = expected_source // capacity * topk + expected_lane + expected_x = all_x[expected_source] + expected_ids = torch.where(local_routes, all_indices - rank * local_experts, -1)[expected_source] + expected_weights = gather(full_weights)[expected_source] if with_weights else None + + def check(payload, recv_ids, recv_weights, handle): + source = handle.recv_src_metadata[:, 0].long() + order = slice(None) if deterministic else source.argsort() + exact(source[order], expected_source) + exact(handle.recv_src_metadata[:, 1].long()[order], expected_peer_topk) + exact(payload[order], expected_x) + exact(recv_ids[order], expected_ids) + if with_weights: + exact(recv_weights[order], expected_weights) + else: + assert recv_weights is None + buffer = deep_ep.EPBuffer(dist.group.WORLD, num_max_tokens_per_rank=capacity, hidden=hidden, @@ -53,19 +80,10 @@ def dispatch(borrow, value=x): result[-1].current_stream_wait() return result[:4] - materialized, ref_ids, ref_weights, ref_handle = dispatch(False) + check(*dispatch(False)) view, recv_ids, recv_weights, handle = dispatch(True) payload = view.slab.index_select(0, view.row_indices) - source = handle.recv_src_metadata[:, 0].long() - exact(payload, all_x[source // capacity, source % capacity]) - order = source.argsort() - ref_order = ref_handle.recv_src_metadata[:, 0].argsort() - exact(payload[order], materialized[ref_order]) - exact(recv_ids[order], ref_ids[ref_order]) - if with_weights: - exact(recv_weights[order], ref_weights[ref_order]) - else: - assert recv_weights is None + check(payload, recv_ids, recv_weights, handle) rejected(lambda: buffer.dispatch(x, **kwargs), RuntimeError) rejected(lambda: buffer.combine(payload, handle), RuntimeError) rejected(buffer.destroy, RuntimeError) @@ -82,18 +100,26 @@ def dispatch(borrow, value=x): for peer in range(world)]).sum(dim=0).to(x.dtype) exact(combined, expected_count[:, None].expand_as(x)) - # Release must order buffer reuse after the delayed reader. - view, _, _, delayed_handle = dispatch(True) - delayed_source = delayed_handle.recv_src_metadata[:, 0].long() - delayed_expected = all_x[delayed_source // capacity, delayed_source % capacity] - side = torch.cuda.Stream() - with torch.cuda.stream(side): + # First wait on a third stream; release must protect both delayed readers. + view, delayed_ids, delayed_weights, delayed_handle, event = buffer.dispatch(x, + borrow_recv=True, + **(kwargs | dict(async_with_compute_stream=True))) + ready = torch.cuda.Stream() + with torch.cuda.stream(ready): torch.cuda._sleep(4_000_000) - delayed = view.slab.index_select(0, view.row_indices) - view.release(side) + event.current_stream_wait() + readers = [torch.cuda.Stream(), torch.cuda.Stream()] + delayed = [] + for reader, delay in zip(readers, (4_000_000, 64_000_000)): + with torch.cuda.stream(reader): + torch.cuda._sleep(delay) + delayed.append(view.slab.index_select(0, view.row_indices)) + view.release(*readers) + assert not readers[-1].query(), 'Delayed reader finished before buffer reuse was tested' dispatch(False, x + 1) - side.synchronize() - exact(delayed, delayed_expected) + for reader, received in zip(readers, delayed): + reader.synchronize() + check(received, delayed_ids, delayed_weights, delayed_handle) for extra in [dict(do_expand=True), dict(defer_epilogue=True), dict(do_cpu_sync=False), dict(handle=handle)]: rejected(lambda extra=extra: buffer.dispatch(x, borrow_recv=True, **(kwargs | extra)), ValueError)