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

static def token_size() -> Int

Returns the size of the (quantized) token in bytes.

Returns:

Int

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:

Int

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:

Int

msg_size​

static def msg_size() -> Int

Returns the size of the message in bytes.

Returns:

Int

src_info_offset​

static def src_info_offset() -> Int

Returns the offset of the source info in the message.

Returns:

Int

topk_info_offset​

static def topk_info_offset() -> Int

Returns the offset of the top-k info in the message.

Returns:

Int

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.

Was this page helpful?