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:
- output (
TileTensor[out_type, Engine=output.Engine, address_space=output.address_space, linear_idx_type=output.linear_idx_type]): One combined row per token. Shape: [num_tokens, hidden]. - down (
TileTensor[down_type, Engine=down.Engine, address_space=down.address_space, linear_idx_type=down.linear_idx_type]): Expert outputs in expert-permuted (token_expert_order) row order. Shape: [num_tokens * num_experts_per_token, hidden]. - restore_order (
TileTensor[.uint32, Engine=restore_order.Engine, address_space=restore_order.address_space, linear_idx_type=restore_order.linear_idx_type]): Maps token-major indexi = token * num_experts_per_token + kto its row indown. Shape: [num_tokens * num_experts_per_token]. - router_weight (
TileTensor[weight_type, Engine=router_weight.Engine, address_space=router_weight.address_space, linear_idx_type=router_weight.linear_idx_type]): Per-(token, expert) routing weight applied before the sum. Shape: [num_tokens, num_experts_per_token]. - context (
DeviceContext): The device context.