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 ¶
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.