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 trait
TokenFormat
Specifies the wire format for a single MoE token in EP dispatch/combine.
Implementors encode how a token's hidden-state vector is packed for cross-GPU transfer (quantization, scale placement, alignment) and how the received bytes are unpacked into the output tensor. The graph compiler selects a concrete implementation based on the model's quantization config.
All size and alignment properties are compile-time constants so the dispatch kernel can allocate receive buffers and issue vectorized copies without runtime branching.
Implemented traits
AnyType,
Deinitable,
DevicePassable
comptime members
alignment
comptime alignment
dispatch_smem_size
comptime dispatch_smem_size
dispatch_wait_tile_shape
comptime dispatch_wait_tile_shape
hid_dim
comptime hid_dim
top_k
comptime top_k
Required methods
token_size
copy_token_to_send_buf
static def copy_token_to_send_buf[src_type: DType, block_size: Int, buf_addr_space: AddressSpace = .GENERIC, thread_base: Int = Int(0)](buf_p: Pointer[UInt8, address_space=buf_addr_space], src_p: Pointer[Scalar[src_type], address_space=src_p.address_space], input_scale: Float32)
Copies the token to the send buffer, called by all comm threads.
thread_base (comptime) is the absolute thread_idx.x of the first
comm thread; it is subtracted from thread_idx.x so the copy can stripe
over a sub-range of a larger fused block. 0 (default) reproduces the
standalone all-threads-in-block striding.
copy_msg_to_output_tensor
def copy_msg_to_output_tensor[buf_addr_space: AddressSpace = .GENERIC](self, buf_p: Pointer[UInt8, address_space=buf_addr_space], token_index: Int, expert_slot: Int = Int(0), expert_start: Int = Int(0))
Copy the message to the output tensor. This function needs to be called by all threads in a warp.
expert_slot (= expert_id + shared_expert_offset) and expert_start
(the expert's first output row) are supplied by the tile loop and used
only by formats that fold the grouped-matmul scale preshuffle into this
copy (MXFP4 KS224); other formats ignore them.
Provided methods
src_info_size
static def src_info_size() -> Int
Returns the size of the source info in bytes. Currently, source info is a single int32 that stores a token's index in the original rank.
Returns:
topk_info_size
static def topk_info_size() -> Int
Returns the size of the top-k info in bytes. Currently, top-k info is an array of uint16 that stores a token's top-k expert IDs.
Returns:
msg_size
src_info_offset
static def src_info_offset() -> Int
Returns the offset of the source info in the message.
Returns:
topk_info_offset
static def topk_info_offset() -> Int
Returns the offset of the top-k info in the message.
Returns:
pad_expert_offsets
def pad_expert_offsets[n_groups: Int](self, row_offsets: Pointer[UInt32, address_space=row_offsets.address_space])
Pad the offsets to satisfy the grouped matmul alignment requirement.
scatter_row_scales
static def scatter_row_scales(recv_base: Pointer[UInt8, MutUntrackedOrigin], send_buf_p: Pointer[UInt8, MutUntrackedOrigin], dst_expert_local_idx: Int32, final_row: Int32, lane: Int)
Scatters this row's scale factors into the format's own arena.
Called by the sending warp once the destination row is fixed, before the warp joins and publishes readiness. Formats that carry their scales inside the message body -- every format but the block-scaled one, and the block-scaled one until an arena is configured -- do nothing here, so the dispatch kernel needs no knowledge of scale placement or of the target's scale-factor layout.
Args:
- recv_base (
Pointer[UInt8, MutUntrackedOrigin]): Base of the destination peer's receive allocation. The arena's offset inside it is the format's own business. - send_buf_p (
Pointer[UInt8, MutUntrackedOrigin]): This row's staged send-buffer message. - dst_expert_local_idx (
Int32): Destination-local expert index. - final_row (
Int32): The row this block already reserved in the final layout. - lane (
Int): Calling lane, used to spread the work across the warp.
init_smem_resources
def init_smem_resources[smem_base_offset: Int = Int(0), warp_base: Int = Int(0)](self)
Initialize the shared memory resources for the token format.
warp_base is the comm warp-class's first warp inside a fused
persistent kernel (0 in the standalone kernels); staging mbars index by
the comm-local warp warp_id() - warp_base.
copy_msg_tile_to_output_tensor
def copy_msg_tile_to_output_tensor[extract_topk_info_func: def(Pointer[UInt8, MutUntrackedOrigin], Int) -> None, recv_buf_ptr_func: def(Int) -> Pointer[UInt8, MutUntrackedOrigin], //, n_warps: Int, shared_expert_offset: Int = Int(0), warp_base: Int = Int(0), smem_base_offset: Int = Int(0)](self, expert_id: Int, expert_start_pos: Int, tile_id: Int, tile_end: Int, extract_topk_info_functor: extract_topk_info_func, recv_buf_ptr_functor: recv_buf_ptr_func)
Copy a tile of tokens from the receive buffer to the output tensor.