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 function

moe_finalize

def moe_finalize[down_type: DType, weight_type: DType, out_type: DType, //, target: StringSpan[ImmStaticOrigin]](output: TileTensor[out_type, Engine=output.Engine, address_space=output.address_space, linear_idx_type=output.linear_idx_type], down: TileTensor[down_type, Engine=down.Engine, address_space=down.address_space, linear_idx_type=down.linear_idx_type], restore_order: TileTensor[.uint32, Engine=restore_order.Engine, address_space=restore_order.address_space, linear_idx_type=restore_order.linear_idx_type], router_weight: TileTensor[weight_type, Engine=router_weight.Engine, address_space=router_weight.address_space, linear_idx_type=router_weight.linear_idx_type], context: DeviceContext)

Fuses the MoE unpermute gather with the top-k weighted row sum.

Each output element reads its token's num_experts_per_token rows of down through restore_order, scales each by its router weight and accumulates in fp32, so the [num_tokens, num_experts_per_token, hidden] unpermuted tensor never materializes.

Parameters:

  • ​down_type (DType): DType of the permuted expert outputs.
  • ​weight_type (DType): DType of the router weights.
  • ​out_type (DType): DType of the combined output.
  • ​target (StringSpan[ImmStaticOrigin]): The target device to run the kernel on.

Args:

Was this page helpful?