vllm_omni.diffusion.models.magi2.mh_moe ¶
Native multi-head MoE used by MAGI-2 Preview.
Adapted from SandAI's Apache-2.0 flash_mh_moe implementation and modified to use vLLM's existing expert-parallel group. MAGI's routing is unusual: each of twelve 256-wide hidden-state heads independently selects experts from its own 256-expert bank. It is therefore not representable by vLLM's conventional whole-token :class:FusedMoE primitive.
Magi2MultiHeadMoE ¶
Bases: Module
Checkpoint-compatible MAGI-2 head-routed expert layer.
W_down instance-attribute ¶
W_down = nn.Parameter(
torch.empty(
self.local_flatten_num_experts,
self.d_expert,
self.d_head,
dtype=config.params_dtype,
)
)
W_gate instance-attribute ¶
W_gate = nn.Parameter(
torch.empty(
self.local_flatten_num_experts,
self.d_head,
self.d_expert,
dtype=config.params_dtype,
)
)
W_up instance-attribute ¶
W_up = nn.Parameter(
torch.empty(
self.local_flatten_num_experts,
self.d_head,
self.d_expert,
dtype=config.params_dtype,
)
)
gate instance-attribute ¶
gate = nn.Parameter(
torch.empty(
self.local_flatten_num_experts,
self.d_head,
dtype=torch.float32,
)
)
local_flatten_num_experts instance-attribute ¶
local_num_heads instance-attribute ¶
padded_num_heads instance-attribute ¶
padded_num_heads = (
math.ceil(self.num_heads / self.ep_group.world_size)
* self.ep_group.world_size
)
ep_slice ¶
Slice flattened (head,expert) checkpoint rows for this rank.
Magi2MultiHeadMoEConfig dataclass ¶
compute_topk_probs_and_indices ¶
compute_topk_probs_and_indices(
router_logits: Tensor,
top_k: int,
*,
score_func: RoutingScore = "sigmoid",
expert_bias: Tensor | None = None,
route_norm: bool = True,
norm_eps: float = 1e-12,
) -> tuple[Tensor, Tensor]
Route independently for every [head, token] pair.
The auxiliary-free bias affects expert selection but deliberately does not affect the returned routing probability, matching the training recipe.
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.
swiglu7_pair ¶
Released clamped SwiGLU7 expert activation, evaluated in fp32.
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.