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
EPCombineKernel
struct EPCombineKernel[num_threads: Int, n_sms: Int, top_k: Int, n_experts: Int, n_ranks: Int, msg_bytes: Int, max_tokens_per_rank: Int, p2p_world_size: Int, use_shmem: Bool = True, fused_shared_expert: Bool = False, skip_a2a: Bool = False]
Implements combine_async and combine_wait kernel logic for Expert Parallelism.
This struct encapsulates the token combine operations used in MoE (Mixture of Experts) models with expert parallelism. It provides methods for:
-
Async Combine:
send_tokens_back: Send processed tokens back to their original ranks.
-
Wait for Arrivals:
wait_for_all_arrivals: Aux SMs wait for all tokens to arrive.reduce_and_copy_to_output: Comm SMs reduce and copy tokens to output.
Parameters
- num_threads (
Int): The number of threads per block. - n_sms (
Int): The total number of SMs in the device. - top_k (
Int): The number of selected experts per token. - n_experts (
Int): The total number of experts in the model. - n_ranks (
Int): The number of devices participating in communication. - msg_bytes (
Int): The number of bytes per token message. - max_tokens_per_rank (
Int): The maximum number of tokens per rank. - p2p_world_size (
Int): Size of a high-speed GPU interconnect group. - use_shmem (
Bool): Whether to use the SHMEM API for communication. - fused_shared_expert (
Bool): Whether to filter out the shared expert's outputs. - skip_a2a (
Bool): Whether to skip the A2A communication. If true, we will only receive tokens from the current device.
Implemented traits
comptime members
blocks_done_offset
comptime blocks_done_offset = ((Int(7) * n_experts) + Int(7))
n_local_experts
comptime n_local_experts = (n_experts // n_ranks)
n_reduce_sms
comptime n_reduce_sms = (n_sms - Int(1))
n_wait_sms
comptime n_wait_sms = 1
n_warps
comptime n_warps = (num_threads // _resolve_warp_size())
pair_done_offset
comptime pair_done_offset = (EPCombineKernel[num_threads, n_sms, top_k, n_experts, n_ranks, msg_bytes, max_tokens_per_rank, p2p_world_size, use_shmem, fused_shared_expert, skip_a2a].blocks_done_offset + Int(1))
use_balanced_send
comptime use_balanced_send = not skip_a2a if not use_shmem else not use_shmem and (p2p_world_size == n_ranks)
Methods
send_buf_layout
static def send_buf_layout[out_dtype: DType = _get_index_type[Layout[TypeList[ComptimeInt[((EPCombineKernel[num_threads, n_sms, top_k, n_experts, n_ranks, msg_bytes, max_tokens_per_rank, p2p_world_size, use_shmem, fused_shared_expert, skip_a2a].n_local_experts * n_ranks) * max_tokens_per_rank)], ComptimeInt[msg_bytes]](), TypeList[ComptimeInt[msg_bytes], ComptimeInt[Int(1)]]()]](AddressSpace.GENERIC)](coord: Coord) -> Scalar[out_dtype]
Returns:
recv_buf_layout
recv_count_layout
copy_shared_expert_outputs
static def copy_shared_expert_outputs[input_type: DType, //](input_tokens: TileTensor[input_type, address_space=input_tokens.address_space, linear_idx_type=input_tokens.linear_idx_type], output_tokens: TileTensor[input_type, address_space=output_tokens.address_space, linear_idx_type=output_tokens.linear_idx_type])
Copies shared expert outputs to the output tensor.
This method copies the shared expert's output tokens from the input tensor to the output tensor when fused_shared_expert is enabled.
Args:
- input_tokens (
TileTensor[input_type, address_space=input_tokens.address_space, linear_idx_type=input_tokens.linear_idx_type]): The input tokens containing shared expert outputs. - output_tokens (
TileTensor[input_type, address_space=output_tokens.address_space, linear_idx_type=output_tokens.linear_idx_type]): The output tensor to copy shared expert outputs to.
send_tokens_back
static def send_tokens_back[input_type: DType, //, gate_ffn_done: Bool = False, ffn_done_sentinel: UInt32 = UInt32(4294967295), ffn_done_max_spin: Int = Int(16777216), ffn_done_m_block: Int = Int(0), ffn_done_tiles_per_m_block: Int = Int(0), thread_base: Int = Int(0), warp_base: Int = Int(0), comm_barrier_id: Int = Int(-1)](input_tokens: TileTensor[input_type, address_space=input_tokens.address_space, linear_idx_type=input_tokens.linear_idx_type], src_info: TileTensor[.int32, address_space=src_info.address_space, linear_idx_type=src_info.linear_idx_type], send_buf_p: Pointer[UInt8, MutUntrackedOrigin], recv_buf_ptrs: Array[Pointer[UInt8, MutUntrackedOrigin], p2p_world_size], recv_count_ptrs: Array[Pointer[UInt64, MutUntrackedOrigin], p2p_world_size], atomic_counter: Pointer[Int32, MutUntrackedOrigin], rank_completion_counter: Pointer[Int32, MutUntrackedOrigin], my_rank: Int32, send_sm_id: Int, n_send_sms: Int, ffn_done_ptr: Optional[Pointer[UInt32, MutUntrackedOrigin]] = None, row_offsets_ptr: Optional[Pointer[UInt32, MutUntrackedOrigin]] = None)
Send processed tokens back to their original ranks.
Each SM handles one expert-rank pair, sending all tokens for that pair back to the original rank. 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 tokens to be sent back. - src_info (
TileTensor[.int32, address_space=src_info.address_space, linear_idx_type=src_info.linear_idx_type]): Source token info (original position and top-k ID). - 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. - recv_count_ptrs (
Array[Pointer[UInt64, MutUntrackedOrigin], p2p_world_size]): Array of pointers to receive count buffers. - atomic_counter (
Pointer[Int32, MutUntrackedOrigin]): Atomic counter for synchronization. - rank_completion_counter (
Pointer[Int32, MutUntrackedOrigin]): Counter for per-rank completion tracking. - my_rank (
Int32): The rank of the current device. - send_sm_id (
Int): The role-local SM index (block_idx.xin the standalone kernels). Passed in so the same phase can run under a different SM mapping inside a fused persistent kernel. - n_send_sms (
Int): Stride over experts across the sending SMs (grid_dim.xin the standalone kernels). - ffn_done_ptr (
Optional[Pointer[UInt32, MutUntrackedOrigin]]): Per-expert FFN completion counters. Undergate_ffn_doneone elected lane spins on this until the producer's count reaches the expert's target, then a block barrier fans release-visibility out to every thread. - row_offsets_ptr (
Optional[Pointer[UInt32, MutUntrackedOrigin]]): Per-expert output row offsets, used to size thegate_ffn_donewait target. Required whenffn_done_tiles_per_m_blockis set.
wait_for_all_arrivals
static def wait_for_all_arrivals(recv_count_p: Pointer[UInt64, MutUntrackedOrigin], atomic_counter: Pointer[Int32, MutUntrackedOrigin], n_active_reduce_sms: Int)
Auxiliary SM logic for combine_wait_kernel.
Waits for all tokens to arrive from all ranks, then signals other SMs that they can start copying tokens to the output tensor.
Args:
- recv_count_p (
Pointer[UInt64, MutUntrackedOrigin]): Pointer to the receive count buffer. - atomic_counter (
Pointer[Int32, MutUntrackedOrigin]): Atomic counter for synchronization. - n_active_reduce_sms (
Int): Count of reduce communication SMs to seed with the data-ready flag (grid_dim.x - n_wait_smsin the standalone kernels).
reduce_and_copy_to_output
static def reduce_and_copy_to_output[output_type: DType, router_weights_wrapper: Optional[def[width: Int](token_idx: Int, topk_id: Int) capturing thin -> SIMD[.float32, width]] = None, elementwise_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None] = None](output_tokens: TileTensor[output_type, address_space=output_tokens.address_space, linear_idx_type=output_tokens.linear_idx_type], recv_buf_p: Pointer[UInt8, MutUntrackedOrigin], atomic_counter: Pointer[Int32, MutUntrackedOrigin], my_rank: Int32, reduce_sm_id: Int, n_active_reduce_sms: Int, topk_ids_p: Optional[Pointer[Int32, ImmUntrackedOrigin]] = None)
Communication SM logic for combine_wait_kernel.
Copies received tokens to the output tensor, optionally applying router weights and reduction across top-k experts.
Args:
- output_tokens (
TileTensor[output_type, address_space=output_tokens.address_space, linear_idx_type=output_tokens.linear_idx_type]): The tensor to store the output tokens. - 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. - reduce_sm_id (
Int): The role-local SM index (block_idx.xin the standalone kernels). Passed in so the same phase can run under a different SM mapping inside a fused persistent kernel. - n_active_reduce_sms (
Int): Count of reduce communication SMs (grid_dim.x - n_wait_smsin the standalone kernels). - topk_ids_p (
Optional[Pointer[Int32, ImmUntrackedOrigin]]): Pointer to the top-k IDs for each token, only required if skip_a2a is True.