Skip to content

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.

RoutingScore module-attribute

RoutingScore = Literal['softmax', 'sigmoid']

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,
    )
)

config instance-attribute

config = config

d_expert instance-attribute

d_expert = config.expert_intermediate_size

d_head instance-attribute

d_head = config.hidden_size // config.num_heads

ep_group instance-attribute

ep_group = ep_group or get_magi2_ep_group()

ep_pad_heads instance-attribute

ep_pad_heads = self.padded_num_heads - self.num_heads

gate instance-attribute

gate = nn.Parameter(
    torch.empty(
        self.local_flatten_num_experts,
        self.d_head,
        dtype=torch.float32,
    )
)

has_real_moe_heads instance-attribute

has_real_moe_heads = self.local_head_start < self.num_heads

local_flatten_num_experts instance-attribute

local_flatten_num_experts = (
    self.local_num_heads * self.num_experts
)

local_head_start instance-attribute

local_head_start = self.ep_group.rank * self.local_num_heads

local_num_heads instance-attribute

local_num_heads = (
    self.padded_num_heads // self.ep_group.world_size
)

num_experts instance-attribute

num_experts = config.num_experts

num_heads instance-attribute

num_heads = config.num_heads

padded_num_heads instance-attribute

padded_num_heads = (
    math.ceil(self.num_heads / self.ep_group.world_size)
    * self.ep_group.world_size
)

router instance-attribute

router = nn.Module()

top_k instance-attribute

top_k = config.top_k

ep_slice

ep_slice(checkpoint_tensor: Tensor) -> Tensor

Slice flattened (head,expert) checkpoint rows for this rank.

forward

forward(x: Tensor) -> Tensor

Magi2MultiHeadMoEConfig dataclass

expert_intermediate_size instance-attribute

expert_intermediate_size: int

hidden_size instance-attribute

hidden_size: int

num_experts instance-attribute

num_experts: int

num_heads instance-attribute

num_heads: int

params_dtype instance-attribute

params_dtype: dtype

route_norm class-attribute instance-attribute

route_norm: bool = True

route_scale class-attribute instance-attribute

route_scale: float = 1.0

score_func class-attribute instance-attribute

score_func: RoutingScore = 'sigmoid'

top_k instance-attribute

top_k: int

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

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

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.

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.