IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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:

  1. 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.
  2. 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; see ep_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 with ep_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 dedicated NB_SEND hardware barrier id instead of the generic id-0 barrier(). Off by default. NVIDIA only; AMD keeps barrier() either way.

Implemented traits​

AnyType, Deinitable, Movable

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

sm_schedule_offset​

comptime sm_schedule_offset = ((Int(6) * n_experts) + Int(7))

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 ticket and pool-cursor 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:

Int

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:

Scalar[out_dtype]

recv_count_layout​

static def recv_count_layout(coord: Coord) -> Int32

Returns:

Int32

send_buf_layout​

static def send_buf_layout(coord: Coord) -> Int32

Returns:

Int32

monitor_and_signal_completion​

static def monitor_and_signal_completion[comm_warp_base: Int = Int(0)](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, sm_id: Int)

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.

Parameters:

  • ​comm_warp_base (Int): Absolute warp_id() of the first warp in the comm thread-class (0 in the standalone kernels; the FFN warp count in the fused megakernel). The per-expert warp index is the comm-local warp_id() - comm_warp_base.

Args:

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, comm_thread_base: Int = Int(0), n_comm_threads: Int = num_threads, 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], 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, sm_id: Int, n_active_send_sms: Int, 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.

Parameters:

  • ​input_type (DType): DType of the input token elements (inferred).
  • ​input_scales_wrapper (Optional[def[dtype: DType](Int) capturing thin -> Scalar[dtype]]): Optional wrapper supplying the input's block scaling factors; None for an unquantized input.
  • ​comm_thread_base (Int): Absolute thread_idx.x of the first thread in the comm thread-class (0 in the standalone kernels, the FFN thread count inside the fused megakernel). Comm-local indices are thread_idx.x - comm_thread_base.
  • ​n_comm_threads (Int): Number of threads in the comm thread-class (the block size in the standalone kernels). Sets the token-copy and cross-warp striding.
  • ​comm_barrier_id (Int): Hardware named-barrier id for the comm class, or negative (default) to use the block-wide barrier. See _comm_barrier.

Args:

wait_for_arrivals_and_compute_offsets​

static def wait_for_arrivals_and_compute_offsets[comm_thread_base: Int = Int(0), n_comm_threads: Int = num_threads, comm_barrier_id: Int = Int(-1)](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, n_active_offset_sms: Int, 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.

Parameters:

  • ​comm_thread_base (Int): Absolute thread_idx.x of the first thread in the comm thread-class (0 in the standalone kernels; the FFN thread count in the fused megakernel). All tid indices below are thread_idx.x - comm_thread_base.
  • ​n_comm_threads (Int): Number of threads in the comm thread-class (scopes the phase's block barriers when comm_barrier_id >= 0).
  • ​comm_barrier_id (Int): Named-barrier id for the comm class, or negative (the default) for the full-block barrier().

Args:

copy_received_tokens_to_output​

static def copy_received_tokens_to_output[comm_thread_base: Int = Int(0), n_comm_threads: Int = num_threads, comm_barrier_id: Int = Int(-1), comm_smem_base: Int = Int(0), emit_l1_release: Bool = False, l1_release_token_block: Int = Int(1), l1_release_atomic_pad: Int = Int(1), l1_release_delta: UInt32 = UInt32(1), trace_scatter_release: Bool = False, trace_rings_per_cta: Int = Int(1), trace_ring_capacity: Int = Int(1), trace_comm_ring_id: Int = Int(0), TraceBufT: TraceBuf = NullTrace](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, scatter_sm_id: Int, l1_arrival_ptr: Optional[Pointer[UInt32, MutAnyOrigin]] = None, trace_buf: TraceBufT = NullTrace(), trace_ring_base: Int = Int(0))

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.

Parameters:

  • ​comm_thread_base (Int): Absolute thread_idx.x of the first thread in the comm thread-class (0 in the standalone kernels, the FFN thread count inside the fused megakernel). Comm-local indices are thread_idx.x - comm_thread_base.
  • ​n_comm_threads (Int): Number of threads in the comm thread-class (the block size in the standalone kernels). Sets the tile-copy warp striding.
  • ​comm_barrier_id (Int): Hardware named-barrier id for the comm class, or negative (default) to use the block-wide barrier. See _comm_barrier.
  • ​comm_smem_base (Int): Byte offset into dynamic shared memory where the tile-copy staging begins (0 in the standalone kernels, the host kernel's SMEM size inside the fused megakernel). Forwarded to copy_msg_tile_to_output_tensor.
  • ​emit_l1_release (Bool): Publish the L1 arrival counter as tokens land, so a co-resident FFN can start before the whole scatter drains.
  • ​l1_release_token_block (Int): Token granularity of one L1 release.
  • ​l1_release_atomic_pad (Int): Padding between per-pool release counters, so they do not share a cache line.
  • ​l1_release_delta (UInt32): Increment applied per release.
  • ​trace_scatter_release (Bool): Stamp each L1 release into the trace ring.
  • ​trace_rings_per_cta (Int): Rings allocated per CTA.
  • ​trace_ring_capacity (Int): Events per ring.
  • ​trace_comm_ring_id (Int): Ring id this warp role owns.
  • ​TraceBufT (TraceBuf): Trace sink type; NullTrace compiles every stamp out.

Args:

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, pack_sm_id: Int, n_active_comm_sms: Int, n_send_sms: Int)

Copies already-quantized shared expert tokens from send_buf to output.

Waits for every send SM to publish its rows of 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.
  • ​pack_sm_id (Int): The role-local SM index (block_idx.x in the standalone kernels). Passed in so the same phase can run under a different SM mapping inside a fused persistent kernel.
  • ​n_active_comm_sms (Int): Count of communication SMs available for the copy (grid_dim.x - n_offset_sms in the standalone kernels).
  • ​n_send_sms (Int): Count of send SMs whose send_buf_ready bumps to wait for (grid_dim.x - n_signal_sms in the standalone kernels).

Was this page helpful?