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
EPDispatchKernel
struct EPDispatchKernel[num_threads: Int, n_sms: Int, n_experts: Int, n_ranks: Int, max_tokens_per_rank: Int, p2p_world_size: Int, token_fmt_type: TokenFormat, use_shmem: Bool = True, fused_shared_expert: Bool = False, skip_a2a: Bool = False, has_rank_flag: Bool = False, ep_ord_r: Bool = False, ep_final_layout: Bool = False, ep_prod_reserve: Bool = False, ep_copy_role_split: Bool = False, ep_send_join_named: Bool = False]
Implements dispatch_async and dispatch_wait kernel logic for Expert Parallelism.
This struct encapsulates the token dispatch operations used in MoE (Mixture of Experts) models with expert parallelism. It provides methods for:
-
Async Dispatch:
monitor_and_signal_completion: Aux SMs count tokens per expert and signal completion when all tokens for an expert have been sent.copy_and_send_tokens: Comm SMs copy tokens to send buffer and transfer them to destination ranks.
-
Wait for Arrivals:
wait_for_arrivals_and_compute_offsets: Aux SMs wait for token arrivals and compute output tensor offsets. Also signals other SMs to copy the tokens to the output tensor once data is ready.copy_received_tokens_to_output: Comm SMs copy received tokens to the output tensor.
Parameters
- num_threads (
Int): The number of threads per block. - n_sms (
Int): The total number of SMs in the device. - n_experts (
Int): The total number of experts in the model. - n_ranks (
Int): The number of devices participating in communication. - max_tokens_per_rank (
Int): The maximum number of tokens per rank. - p2p_world_size (
Int): Size of a high-speed GPU interconnect group. - token_fmt_type (
TokenFormat): Type conforming to TokenFormat trait. - use_shmem (
Bool): Whether to use the SHMEM API for communication. - fused_shared_expert (
Bool): Whether to pack the shared expert inputs with the routed experts' inputs. - skip_a2a (
Bool): Whether to skip the A2A communication. If true, we will only send tokens within the current device. - has_rank_flag (
Bool): Whether to reserve a dedicated rank-completion word per source rank in the receive-count buffer tail. Required by, and only meaningful with,ep_ord_r. - ep_ord_r (
Bool): Whether completion is published at rank level (one elected system-scope release per source rank) instead of the legacy per-expert scheme. A fused consumer that acquires once per source rank needs this; seeep_signal_completion. - ep_final_layout (
Bool): Whether source ranks write each routed row directly into its final contiguous per-expert row, so a destination reads one expert as a single dense range. Off by default, which keeps the rank-major staging layout and its per-source-rank gather. - ep_prod_reserve (
Bool): Whether a destination hands each source a reserved contiguous row range per non-empty expert block, published through a generation-tagged source-local mailbox. Off by default. Only meaningful withep_final_layout. - ep_copy_role_split (
Bool): Opt in to the split's publisher fan-out (the token format carries the matching parameter for the copy body). Off by default, which keeps the stock warp-strided fan-out. - ep_send_join_named (
Bool): Join the per-token send on the dedicatedNB_SENDhardware barrier id instead of the generic id-0barrier(). Off by default. NVIDIA only; AMD keepsbarrier()either way.
Implemented traits
comptime members
cleanup_counter_offset
comptime cleanup_counter_offset = (Int(4) * n_experts)
hid_dim
comptime hid_dim = token_fmt_type.hid_dim
l1_vslot_ticket_offset
comptime l1_vslot_ticket_offset = ((Int(3) * n_experts) + EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].n_local_experts) if ((Int(1) + EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].l2_pool_cursor_slots) <= (n_experts - EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].n_local_experts)) else ((Int(4) * n_experts) + Int(4))
l2_pool_cursor_offset
comptime l2_pool_cursor_offset = (EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].l1_vslot_ticket_offset + Int(1))
l2_pool_cursor_slots
comptime l2_pool_cursor_slots = ((Int(2) * EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].n_local_experts) + Int(2))
msg_bytes
comptime msg_bytes = token_fmt_type.msg_size()()
n_dispatch_async_comm_sms
comptime n_dispatch_async_comm_sms = (n_sms - ceildiv(n_experts, (num_threads // _resolve_warp_size())))
n_dispatch_wait_comm_sms
comptime n_dispatch_wait_comm_sms = (n_sms - Int(1))
n_local_experts
comptime n_local_experts = (n_experts // n_ranks)
n_offset_sms
comptime n_offset_sms = 1
n_signal_sms
comptime n_signal_sms = ceildiv(n_experts, (num_threads // _resolve_warp_size()))
n_warps
comptime n_warps = (num_threads // _resolve_warp_size())
NB_SEND
comptime NB_SEND = 3
rank_flag_base
comptime rank_flag_base = (EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].n_local_experts * n_ranks)
rank_prefix_offset
comptime rank_prefix_offset = (Int(2) * n_experts)
rc_cursor_base
comptime rc_cursor_base = (EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].rank_flag_base + n_ranks)
ready_flag_offset
comptime ready_flag_offset = ((Int(4) * n_experts) + Int(1))
role_split
comptime role_split = EPRoleSplit[num_threads, (token_fmt_type.hid_dim // Int(8)), ep_copy_role_split]
send_buf_ready_offset
comptime send_buf_ready_offset = ((Int(4) * n_experts) + Int(2))
shared_expert_started_offset
comptime shared_expert_started_offset = ((Int(4) * n_experts) + Int(3))
top_k
comptime top_k = token_fmt_type.top_k
work_counter_offset
comptime work_counter_offset = (Int(3) * n_experts)
Methods
assert_l1_vslot_ticket_layout
static def assert_l1_vslot_ticket_layout()
Static layout guarantees for the S1/S2 scheduler words.
rank_flag_offset
static def rank_flag_offset(src_rank: Int) -> Int32
Offset of src_rank's dedicated rank-completion flag.
Args:
- src_rank (
Int): The SOURCE rank whose completion the flag publishes.
Returns:
Int32: Element offset into the destination's receive-count buffer.
recv_count_size
static def recv_count_size() -> Int
Receive-count buffer element count, including both tails.
The per-(destination, local expert) reservation cursors sit directly after the ORD-R rank flags and are always reserved: the struct cannot see whether a given launch enables the production reservation, and under-allocating would put the cursors past the end of the buffer.
Returns:
recv_buf_layout
static def recv_buf_layout[out_dtype: DType = _get_index_type[Layout[TypeList[ComptimeInt[EPDispatchKernel[num_threads, n_sms, n_experts, n_ranks, max_tokens_per_rank, p2p_world_size, token_fmt_type, use_shmem, fused_shared_expert, skip_a2a, has_rank_flag, ep_ord_r, ep_final_layout, ep_prod_reserve, ep_copy_role_split, ep_send_join_named].n_local_experts], ComptimeInt[n_ranks], ComptimeInt[max_tokens_per_rank], ComptimeInt[token_fmt_type.msg_size()()]](), TypeList[ComptimeInt[Int((mul token_fmt_type.msg_size()(), max_tokens_per_rank, n_ranks))], ComptimeInt[Int((mul token_fmt_type.msg_size()(), max_tokens_per_rank))], ComptimeInt[token_fmt_type.msg_size()()], ComptimeInt[Int(1)]]()]](AddressSpace.GENERIC)](coord: Coord) -> Scalar[out_dtype]
Returns:
recv_count_layout
send_buf_layout
monitor_and_signal_completion
static def monitor_and_signal_completion(topk_ids: TileTensor[.int32, address_space=topk_ids.address_space, linear_idx_type=topk_ids.linear_idx_type], recv_count_ptrs: Array[Pointer[UInt64, MutUntrackedOrigin], p2p_world_size], expert_reserved_counter: Pointer[Int32, MutUntrackedOrigin], expert_finished_counter: Pointer[Int32, MutUntrackedOrigin], rank_completion_counter: Pointer[Int32, MutUntrackedOrigin], my_rank: Int32)
Auxiliary SM logic for dispatch_kernel.
Counts tokens per expert and signals completion when all tokens for an expert have been sent. Each warp handles one expert.
Args:
- topk_ids (
TileTensor[.int32, address_space=topk_ids.address_space, linear_idx_type=topk_ids.linear_idx_type]): The top-k expert IDs for each token. - recv_count_ptrs (
Array[Pointer[UInt64, MutUntrackedOrigin], p2p_world_size]): Array of pointers to receive count buffers. - expert_reserved_counter (
Pointer[Int32, MutUntrackedOrigin]): Counter for reserved slots per expert. - expert_finished_counter (
Pointer[Int32, MutUntrackedOrigin]): Counter for finished sends per expert. - rank_completion_counter (
Pointer[Int32, MutUntrackedOrigin]): Counter for per-rank completion tracking. - my_rank (
Int32): The rank of the current device.
copy_and_send_tokens
static def copy_and_send_tokens[input_type: DType, //, input_scales_wrapper: Optional[def[dtype: DType](Int) capturing thin -> Scalar[dtype]] = None](input_tokens: TileTensor[input_type, address_space=input_tokens.address_space, linear_idx_type=input_tokens.linear_idx_type], topk_ids: TileTensor[.int32, address_space=topk_ids.address_space, linear_idx_type=topk_ids.linear_idx_type], send_buf_p: Pointer[UInt8, MutUntrackedOrigin], recv_buf_ptrs: Array[Pointer[UInt8, MutUntrackedOrigin], p2p_world_size], expert_reserved_counter: Pointer[Int32, MutUntrackedOrigin], expert_finished_counter: Pointer[Int32, MutUntrackedOrigin], my_rank: Int32, prod_gen_p: Pointer[Int32, MutUntrackedOrigin] = Pointer(unsafe_from_address=16), prod_gen: Int32 = Int32(0))
Communication SM logic for dispatch_kernel.
Copies tokens to send buffer and transfers them to destination ranks. Uses direct P2P transfers for same-node destinations and SHMEM for cross-node destinations.
Args:
- input_tokens (
TileTensor[input_type, address_space=input_tokens.address_space, linear_idx_type=input_tokens.linear_idx_type]): The input tokens to be dispatched. - topk_ids (
TileTensor[.int32, address_space=topk_ids.address_space, linear_idx_type=topk_ids.linear_idx_type]): The top-k expert IDs for each token. - send_buf_p (
Pointer[UInt8, MutUntrackedOrigin]): Pointer to the send buffer. - recv_buf_ptrs (
Array[Pointer[UInt8, MutUntrackedOrigin], p2p_world_size]): Array of pointers to receive buffers. - expert_reserved_counter (
Pointer[Int32, MutUntrackedOrigin]): Counter for reserved slots per expert. - expert_finished_counter (
Pointer[Int32, MutUntrackedOrigin]): Counter for finished sends per expert. - my_rank (
Int32): The rank of the current device. - prod_gen_p (
Pointer[Int32, MutUntrackedOrigin]): Source-local block-base mailbox. Entries[0, n_experts)hold the reserved base for each expert and[n_experts, 2 * n_experts)hold the matching generation tags. Inert when defaulted. - prod_gen (
Int32): Generation tag this launch waits for before reading a base fromprod_gen_p. Inert when defaulted.
wait_for_arrivals_and_compute_offsets
static def wait_for_arrivals_and_compute_offsets(format_handler: token_fmt_type, row_offsets: TileTensor[.uint32, address_space=row_offsets.address_space, linear_idx_type=row_offsets.linear_idx_type], expert_ids: TileTensor[.int32, address_space=expert_ids.address_space, linear_idx_type=expert_ids.linear_idx_type], recv_count_p: Pointer[UInt64, MutUntrackedOrigin], atomic_counter: Pointer[Int32, MutUntrackedOrigin], my_rank: Int32, reserved_shared_expert_tokens: UInt32 = UInt32(0))
Auxiliary SM logic for dispatch_wait_kernel.
Waits for token arrivals from all ranks and computes the output tensor offsets for each local expert. Also signals other SMs to copy the tokens to the output tensor once data is ready.
Args:
- format_handler (
token_fmt_type): Instance of token_fmt_type for token decoding. - row_offsets (
TileTensor[.uint32, address_space=row_offsets.address_space, linear_idx_type=row_offsets.linear_idx_type]): Output row offsets for grouped matmul. - expert_ids (
TileTensor[.int32, address_space=expert_ids.address_space, linear_idx_type=expert_ids.linear_idx_type]): Output expert IDs for grouped matmul. - recv_count_p (
Pointer[UInt64, MutUntrackedOrigin]): Pointer to receive count buffer. - atomic_counter (
Pointer[Int32, MutUntrackedOrigin]): Atomic counter for synchronization. - my_rank (
Int32): The rank of the current device. - reserved_shared_expert_tokens (
UInt32): The number of tokens reserved for the shared expert.
copy_received_tokens_to_output
static def copy_received_tokens_to_output(format_handler: token_fmt_type, row_offsets: TileTensor[.uint32, address_space=row_offsets.address_space, linear_idx_type=row_offsets.linear_idx_type], src_info: TileTensor[.int32, address_space=src_info.address_space, linear_idx_type=src_info.linear_idx_type], recv_buf_p: Pointer[UInt8, MutUntrackedOrigin], atomic_counter: Pointer[Int32, MutUntrackedOrigin], my_rank: Int32)
Communication SM logic for dispatch_wait_kernel.
Copies received tokens from the receive buffer to the output tensor. Each SM is assigned to one local expert and dynamically claims tiles via per-expert atomic counters. Tokens within a tile may come from multiple source ranks; rank boundaries are resolved via the within-expert prefix sums written by the auxiliary SM.
Args:
- format_handler (
token_fmt_type): Instance of token_fmt_type for token decoding. - row_offsets (
TileTensor[.uint32, address_space=row_offsets.address_space, linear_idx_type=row_offsets.linear_idx_type]): Output row offsets for grouped matmul. - src_info (
TileTensor[.int32, address_space=src_info.address_space, linear_idx_type=src_info.linear_idx_type]): Output tensor for source token info. - recv_buf_p (
Pointer[UInt8, MutUntrackedOrigin]): Pointer to the receive buffer. - atomic_counter (
Pointer[Int32, MutUntrackedOrigin]): Atomic counter for synchronization. - my_rank (
Int32): The rank of the current device.
pack_shared_expert_inputs
static def pack_shared_expert_inputs(format_handler: token_fmt_type, send_buf_p: Pointer[UInt8, MutUntrackedOrigin], fused_se_counter: Pointer[Int32, MutUntrackedOrigin], shared_expert_token_count: Int)
Copies already-quantized shared expert tokens from send_buf to output.
Waits for dispatch_async signal SMs to indicate all tokens have been written to the send buffer, then uses tile-based copy via copy_msg_tile_to_output_tensor. Only SMs needed for the copy participate.
Args:
- format_handler (
token_fmt_type): Instance of token_fmt_type for token decoding. - send_buf_p (
Pointer[UInt8, MutUntrackedOrigin]): Pointer to the send buffer containing serialized tokens. - fused_se_counter (
Pointer[Int32, MutUntrackedOrigin]): Pointer to the two fused shared expert atomic counters (send_buf_ready at [0], started at [1]). - shared_expert_token_count (
Int): Number of shared expert tokens to copy.