For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo struct
EPLocalSyncCounters
struct EPLocalSyncCounters[n_experts: Int]
Manages atomic counters for EP kernel synchronization within a device.
This struct provides dedicated atomic counter space for each of the four EP kernels: dispatch_async, dispatch_wait, combine_async, and combine_wait. Each kernel has its own memory region to avoid conflicts, except dispatch_wait and combine_async which must share memory since combine_async reads data that dispatch_wait writes.
The struct is used to synchronize between thread blocks within the same device.
Memory Layout (all sizes in Int32 elements):
- dispatch_async: 2 * n_experts + MAX_GPUS_PER_NODE
- dispatch_wait/combine_async: 8 * n_experts + 8
- combine_wait: MAX_SMS_PER_DEVICE
Fields
- ptr (
Pointer[Int32, MutUntrackedOrigin]): Base pointer to the allocated atomic counter memory.
Implemented traits
AnyType,
Copyable,
Deinitable,
DevicePassable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
device_type
comptime device_type = EPLocalSyncCounters[n_experts]
Methods
__init__
def __init__(ptr: Pointer[Int32, address_space=ptr.address_space]) -> Self
def __init__(mut buffer: DeviceBuffer[.int32]) -> Self
get_type_name
dispatch_async_size
static def dispatch_async_size() -> Int
Returns the size in Int32 elements needed by dispatch_async kernel.
Returns:
dispatch_wait_size
static def dispatch_wait_size() -> Int
Returns the size in Int32 elements needed by dispatch_wait kernel.
Layout (see EPDispatchKernel and EPCombineKernel for exact offset
constants):
Region A [0, 2n_experts): per expert-rank combine_async compat data
Region B [2n_experts, 3n_experts): within-expert rank prefix sums
Region C [3n_experts, 4n_experts): per-expert work counters
(only first n_local_experts entries used; rest unused)
Region D [4n_experts]: cleanup ref counter
Region E [4n_experts + 1]: global ready flag
Region F [4n_experts + 2]: send_buf_ready counter
Region G [4n_experts + 3]: shared_expert_started counter
Region H: the L1 virtual-slot ticket and the L2 pool cursors,
2n_local_experts + 3 words placed by
l1_vslot_ticket_offset; inside Region C's unused tail where
the expert-parallel degree leaves room, otherwise past Region G
Region I: the dispatch_wait SM-to-expert schedule
Region J: combine_async's blocks-finished counter and its
per-(expert, rank) completion counters
Region A will be used by combine_async kernel to track the number of tokens of each expert-rank pair. Regions D through J are reset to 0 by whichever kernel owns them, on its way out.
The returned size is the worst case over the expert-parallel degree, because the host allocates from n_experts alone.
Returns:
combine_async_size
static def combine_async_size() -> Int
Returns the size in Int32 elements needed by combine_async kernel.
Must match dispatch_wait_size() since combine_async reuses the same memory region.
Returns:
combine_wait_size
static def combine_wait_size() -> Int
Returns the size in Int32 elements needed by combine_wait kernel.
Returns:
total_size
static def total_size() -> Int
Returns the total size in Int32 elements needed for all counters.
Returns:
get_dispatch_async_ptr
def get_dispatch_async_ptr(self) -> Pointer[Int32, MutUntrackedOrigin]
Returns pointer to dispatch_async kernel atomic counters.
Layout: [0, n_experts): reserved counters per expert [n_experts, 2*n_experts): finished counters per expert
Returns:
get_dispatch_wait_ptr
def get_dispatch_wait_ptr(self) -> Pointer[Int32, MutUntrackedOrigin]
Returns pointer to dispatch_wait kernel atomic counters.
Returns:
get_combine_async_ptr
def get_combine_async_ptr(self) -> Pointer[Int32, MutUntrackedOrigin]
Returns pointer to combine_async kernel atomic counters.
Note: Returns the same pointer as get_dispatch_wait_ptr() because combine_async_kernel reads the offset/count data that dispatch_wait_kernel writes.
Returns:
get_combine_wait_ptr
def get_combine_wait_ptr(self) -> Pointer[Int32, MutUntrackedOrigin]
Returns pointer to combine_wait kernel atomic counters.
Returns: