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 function
ep_fused_dispatch_kernel_api
def ep_fused_dispatch_kernel_api[token_fmt_type: TokenFormat, dispatch_dtype: DType, //, n_experts: Int, max_token_per_rank: Int, n_gpus_per_node: Int, n_nodes: Int, fused_shared_expert: Bool, target: StringSlice[ImmStaticOrigin], input_scales_wrapper: Optional[def[dtype: DType](Int) capturing thin -> Scalar[dtype]] = None, skip_a2a: Bool = False, use_shmem: Bool = (n_nodes > Int(1)), allreduce_world_size: Int = Int(1)](token_handler: token_fmt_type, row_offsets: TileTensor[DType.uint32, address_space=row_offsets.address_space, linear_idx_type=row_offsets.linear_idx_type], expert_ids: TileTensor[DType.int32, address_space=expert_ids.address_space, linear_idx_type=expert_ids.linear_idx_type], src_info: TileTensor[DType.int32, address_space=src_info.address_space, linear_idx_type=src_info.linear_idx_type], atomic_counters: TileTensor[DType.int32, address_space=atomic_counters.address_space, linear_idx_type=atomic_counters.linear_idx_type], input_tokens: TileTensor[dispatch_dtype, address_space=input_tokens.address_space, linear_idx_type=input_tokens.linear_idx_type], topk_ids: TileTensor[DType.int32, address_space=topk_ids.address_space, linear_idx_type=topk_ids.linear_idx_type], send_ptrs: TileTensor[DType.uint64, address_space=send_ptrs.address_space, linear_idx_type=send_ptrs.linear_idx_type], recv_ptrs: TileTensor[DType.uint64, address_space=recv_ptrs.address_space, linear_idx_type=recv_ptrs.linear_idx_type], recv_count_ptrs: TileTensor[DType.uint64, address_space=recv_count_ptrs.address_space, linear_idx_type=recv_count_ptrs.linear_idx_type], context: DeviceContext)
Execute the fused Expert Parallelism dispatch kernel.
This function launches the fused dispatch_kernel from ep_comm.mojo that combines both dispatch_async and dispatch_wait functionality in a single kernel launch. It distributes input tokens to expert devices based on top-k routing decisions, then waits for all tokens to arrive and aggregates them for grouped matmul computation.
Arguments: token_handler: Token handler. Wrapper for the output token tensor. row_offsets: Row offsets for grouped matmul. expert_ids: Expert IDs for grouped matmul. src_info: Source routing information for combine phase. atomic_counters: EP kernel synchronization counters. input_tokens: Tokens to dispatch to experts. topk_ids: Expert assignments from router. send_ptrs: Send buffer pointers for each local GPU. recv_ptrs: Receive buffer pointers for each local GPU. recv_count_ptrs: Receive count buffer pointers for each local GPU. context: Device context pointer.
Parameters:
- βtoken_fmt_type (
TokenFormat): Token format type. - βdispatch_dtype (
DType): Data type of the dispatched tokens. - βn_experts (
Int): Total experts across all devices. - βmax_token_per_rank (
Int): Maximum tokens any device can send. - βn_gpus_per_node (
Int): GPUs per physical node. - βn_nodes (
Int): Number of physical nodes. - βfused_shared_expert (
Bool): Whether to pack shared expert inputs with routed experts' inputs. - βtarget (
StringSlice[ImmStaticOrigin]): Target. - βinput_scales_wrapper (
Optional[def[dtype: DType](Int) capturing thin -> Scalar[dtype]]): The wrapper for the input scales. - βskip_a2a (
Bool): Whether to skip the A2A communication. If true, we will only send tokens within the current device. - βuse_shmem (
Bool): Whether to enable SHMEM communication. - βallreduce_world_size (
Int): The world size of the allreduce operation. Only needed for skip_a2a. Used to calculate the workload distribution for the shared expert (if has one).