Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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`.

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.

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.
Expand Down
679 changes: 373 additions & 306 deletions csrc/buffers/ep.hpp

Large diffs are not rendered by default.

294 changes: 183 additions & 111 deletions csrc/kernels/ep/dispatch.hpp

Large diffs are not rendered by default.

29 changes: 12 additions & 17 deletions deep_ep/__init__.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -55,25 +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.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'
305 changes: 134 additions & 171 deletions deep_ep/buffers/ep.py

Large diffs are not rendered by default.

54 changes: 54 additions & 0 deletions deep_ep/buffers/recv_view.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
import torch

from .. import comm


class DispatchRecvView:
"""Forward-only receive storage: logical row ``r`` is ``slab[row_indices[r]]``.

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):
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
Loading