Skip to content

vllm_omni.diffusion.models.magi2.fused_moe_kernels

BF16 routed GEMMs for the MAGI-2 fast path.

The BF16 grouped GEMM is adapted from SGLang's Apache-2.0 implementation at d533d3bf464d280f180e18d4d396e613a5998502. This module implements only MAGI-2's interleaved gate/up GEMM with SwiGLU7 and routed down GEMM. The existing SandAI-derived reference kernel is retained below unchanged.

global_sort_routes

global_sort_routes(
    topk_probs: Tensor,
    topk_indices: Tensor,
    num_experts: int,
) -> tuple[Tensor, Tensor, Tensor]

Convert per-head routes into a stable flattened-expert CSR layout.

invoke_fused_moe_bf16

invoke_fused_moe_bf16(
    A: Tensor,
    B: Tensor,
    C: Tensor,
    topk_weights: Tensor,
    sorted_token_ids: Tensor,
    expert_ids: Tensor,
    num_tokens_post_padded: Tensor,
    *,
    top_k: int,
    config: dict[str, int],
    fuse_swiglu: bool,
) -> None

Launch one BF16 MAGI-2 GEMM in the original route order.

Gate/up calls use interleaved B=[experts, 2*intermediate, hidden] and retain FP32 accumulators through SwiGLU7 before writing BF16 C=[routes, intermediate]. Down calls use B=[experts, hidden, intermediate] and write C=[tokens, router_top_k, hidden]. The latter multiplies routing weights in FP32 before the BF16 output cast.

swiglu7_pair

swiglu7_pair(gate: Tensor, up: Tensor) -> Tensor

Reference SwiGLU7 activation used by the CPU correctness path.

torch_mh_moe_forward

torch_mh_moe_forward(
    x: Tensor,
    gather_ids: Tensor,
    probs: Tensor,
    expert_offsets: Tensor,
    w_gate: Tensor,
    w_up: Tensor,
    w_down: Tensor,
) -> Tensor

Small-shape correctness oracle for the fused expert kernel.

triton_mh_moe_forward

triton_mh_moe_forward(
    x: Tensor,
    gather_ids: Tensor,
    probs: Tensor,
    expert_offsets: Tensor,
    w_gate: Tensor,
    w_up: Tensor,
    w_down: Tensor,
    *,
    deterministic: bool = False,
) -> Tensor

Fused gather/expert/scatter kernel for released MAGI dimensions.