Skip to content

vllm_omni.diffusion.models.magi2.modeling_magi2

Native MAGI-2 Preview diffusion transformer.

The architecture and mHC/MoE equations are adapted from SandAI's Apache-2.0 MAGI-2 Preview implementation. This modified vLLM-Omni version removes the external inference package, MagiCompiler, and SandAI process manager while preserving the released checkpoint module names and tensor layouts.

Magi2Attention

Bases: Module

config instance-attribute

config = config

head_dim instance-attribute

head_dim = config.head_dim

k_norm instance-attribute

k_norm = MultiModalityRMSNorm(
    self.head_dim,
    num_modality=num_modality,
    out_dtype=torch.float32,
)

kv_size instance-attribute

kv_size = self.num_heads_kv * self.head_dim

linear_g instance-attribute

linear_g = make_grouped_linear(
    config.hidden_size,
    config.num_heads_q,
    num_experts=num_modality,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="column",
)

linear_proj instance-attribute

linear_proj = make_grouped_linear(
    global_q_size,
    config.hidden_size,
    num_experts=num_modality,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="row",
)

linear_qkv instance-attribute

linear_qkv = make_grouped_linear(
    config.hidden_size,
    global_q_size + 2 * global_kv_size,
    num_experts=num_modality,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="column",
    qkv_splits=(
        global_q_size,
        global_kv_size,
        global_kv_size,
    ),
)

num_heads_kv instance-attribute

num_heads_kv = (
    config.num_heads_kv // self.tp_group.world_size
)

num_heads_q instance-attribute

num_heads_q = config.num_heads_q // self.tp_group.world_size

num_modality instance-attribute

num_modality = num_modality

packed_attention instance-attribute

packed_attention = Attention(
    num_heads=self.num_heads_q,
    num_kv_heads=self.num_heads_kv,
    head_size=self.head_dim,
    causal=False,
    softmax_scale=self.head_dim**-0.5,
    qkv_layout="THD",
    skip_sequence_parallel=True,
    disable_kv_quant=True,
    custom_attention=Magi2PackedAttentionKernel(
        config.attention_softcap
    ),
)

pre_norm instance-attribute

pre_norm = MultiModalityRMSNorm(
    config.hidden_size, num_modality=num_modality
)

q_norm instance-attribute

q_norm = MultiModalityRMSNorm(
    self.head_dim,
    num_modality=num_modality,
    out_dtype=torch.float32,
)

q_size instance-attribute

q_size = self.num_heads_q * self.head_dim

sinks instance-attribute

sinks = nn.Parameter(
    torch.empty(
        config.attention_sink_tokens,
        self.num_heads_q,
        dtype=torch.float32,
    )
)

tp_group instance-attribute

tp_group = self.linear_qkv.tp_group

attend

attend(
    q: Tensor,
    k: Tensor,
    v: Tensor,
    varlen_handler: VarlenHandler,
    cp_split_sizes: list[int] | Tensor,
) -> Tensor

output

output(
    attention: Tensor,
    gates: Tensor,
    modality_dispatcher: ModalityDispatcher,
) -> Tensor

project

project(
    hidden_states: Tensor,
    rope: Tensor,
    modality_dispatcher: ModalityDispatcher,
) -> tuple[Tensor, Tensor, Tensor, Tensor]

Magi2MLP

Bases: Module

down_proj instance-attribute

down_proj = make_grouped_linear(
    intermediate_size,
    config.hidden_size,
    num_experts=num_modality,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="row",
)

intermediate_size instance-attribute

intermediate_size = intermediate_size

pre_norm instance-attribute

pre_norm = MultiModalityRMSNorm(
    config.hidden_size, num_modality=num_modality
)

up_gate_proj instance-attribute

up_gate_proj = make_grouped_linear(
    config.hidden_size,
    2 * intermediate_size,
    num_experts=num_modality,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="column",
)

forward

forward(
    hidden_states: Tensor, dispatcher: ModalityDispatcher
) -> Tensor

Magi2MultiHeadMoELayer

Bases: Module

config instance-attribute

config = config

local_modality_shared_expert_intermediate_size instance-attribute

local_modality_shared_expert_intermediate_size = (
    moe.modality_shared_expert_intermediate_size // tp_size
)

local_shared_expert_intermediate_size instance-attribute

local_shared_expert_intermediate_size = (
    moe.shared_expert_intermediate_size // tp_size
)

merge_linear instance-attribute

merge_linear = make_grouped_linear(
    config.hidden_size,
    config.hidden_size,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="row",
)

modality_specific_shared_expert_fc1 instance-attribute

modality_specific_shared_expert_fc1 = make_grouped_linear(
    config.hidden_size,
    2 * moe.modality_shared_expert_intermediate_size,
    num_experts=3,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="column",
)

modality_specific_shared_expert_fc2 instance-attribute

modality_specific_shared_expert_fc2 = make_grouped_linear(
    moe.modality_shared_expert_intermediate_size,
    config.hidden_size,
    num_experts=3,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="row",
)

moe_mlp instance-attribute

moe_mlp = Magi2MultiHeadMoE(
    Magi2MultiHeadMoEConfig(
        hidden_size=config.hidden_size,
        num_heads=moe.num_heads,
        num_experts=moe.num_experts,
        top_k=moe.top_k,
        expert_intermediate_size=moe.expert_intermediate_size,
        params_dtype=config.params_dtype,
        score_func=moe.score_function,
        route_norm=moe.normalize_routing_weights,
        route_scale=moe.routing_scale,
    )
)

pre_norm instance-attribute

pre_norm = MultiModalityRMSNorm(
    config.hidden_size, num_modality=3
)

shared_expert_fc1 instance-attribute

shared_expert_fc1 = make_grouped_linear(
    config.hidden_size,
    2 * moe.shared_expert_intermediate_size,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="column",
)

shared_expert_fc2 instance-attribute

shared_expert_fc2 = make_grouped_linear(
    moe.shared_expert_intermediate_size,
    config.hidden_size,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="row",
)

split_linear instance-attribute

split_linear = make_grouped_linear(
    config.hidden_size,
    config.hidden_size,
    bias=False,
    dtype=config.params_dtype,
    parallel_mode="column",
)

combine

combine(
    normalized: Tensor,
    routed: Tensor,
    dispatcher: ModalityDispatcher,
) -> Tensor

route_input

route_input(
    hidden_states: Tensor, dispatcher: ModalityDispatcher
) -> tuple[Tensor, Tensor]

Magi2PostAdapter

Bases: Module

adapter_dim instance-attribute

adapter_dim = config.virtual_width

config instance-attribute

config = config

final_linear_audio instance-attribute

final_linear_audio = nn.Linear(
    self.adapter_dim,
    config.audio_in_channels,
    bias=False,
    dtype=torch.float32,
)

final_linear_video instance-attribute

final_linear_video = nn.Linear(
    self.adapter_dim,
    config.video_in_channels,
    bias=False,
    dtype=torch.float32,
)

final_norm_audio instance-attribute

final_norm_audio = MultiModalityRMSNorm(self.adapter_dim)

final_norm_video instance-attribute

final_norm_video = MultiModalityRMSNorm(self.adapter_dim)

final_out_dim instance-attribute

final_out_dim = max(
    config.video_in_channels, config.audio_in_channels
)

forward

forward(
    hidden_states: Tensor,
    video_indices: Tensor,
    audio_indices: Tensor,
) -> Tensor

Magi2PreAdapter

Bases: Module

adapter_dim instance-attribute

adapter_dim = config.virtual_width

audio_embedder instance-attribute

audio_embedder = nn.Linear(
    config.audio_in_channels,
    self.adapter_dim,
    bias=True,
    dtype=torch.float32,
)

config instance-attribute

config = config

rope instance-attribute

rope = ElementWiseFourierEmbed(
    config.head_dim, learnable=False
)

text_embedder instance-attribute

text_embedder = nn.Linear(
    config.text_in_channels,
    self.adapter_dim,
    bias=True,
    dtype=torch.float32,
)

video_embedder instance-attribute

video_embedder = nn.Linear(
    config.video_in_channels,
    self.adapter_dim,
    bias=True,
    dtype=torch.float32,
)

forward

forward(
    packed: Tensor,
    video_indices: Tensor,
    audio_indices: Tensor,
    text_indices: Tensor,
) -> Tensor

Magi2PreviewTransformer

Bases: Module

Native preview DiT with the released checkpoint hierarchy.

block instance-attribute

block = Magi2TransformerBlock(self.config)

config instance-attribute

config = config or Magi2PreviewConfig()

layers property

layers: ModuleList

Expose Preview layers for direct access and compatibility.

post_adapter instance-attribute

post_adapter = Magi2PostAdapter(self.config)

pre_adapter instance-attribute

pre_adapter = Magi2PreAdapter(self.config)

compile_regions

compile_regions(**compile_kwargs: Any) -> None

Compile the dense compute between the eager attention and MoE kernels.

forward

forward(
    x: Tensor,
    coords_mapping: Tensor,
    modality_mapping: Tensor,
    varlen_handler: VarlenHandler,
    time_token_sequence: Tensor | None = None,
) -> Tensor

load_weights

load_weights(
    weights: Iterable[tuple[str, Tensor]],
) -> set[str]

Strictly load and TP/MoE-slice the released Preview checkpoint.

This is the ordinary loader used by resident and rank-local DLO deployments. DLO+AllGather uses the pipeline's mmap mapping and the same per-parameter transforms before orthogonal DP sharding.

validate_loaded_weights

validate_loaded_weights(loaded_names: set[str]) -> None

Fail closed when the mmap loader misses a Preview tensor.

Magi2TransformerBlock

Bases: Module

config instance-attribute

config = config

layers instance-attribute

layers = nn.ModuleList(
    Magi2TransformerLayer(config, index)
    for index in range(config.num_layers)
)

forward

forward(
    hidden_states: Tensor,
    rope: Tensor,
    varlen_handler: VarlenHandler,
    modality_dispatcher: ModalityDispatcher,
    cp_split_sizes: list[int] | Tensor,
) -> Tensor

Magi2TransformerLayer

Bases: Module

attention instance-attribute

attention = Magi2Attention(
    config, num_modality=num_modality
)

config instance-attribute

config = config

mlp instance-attribute

mlp: Module

region_methods instance-attribute

region_methods = (
    "_attention_input",
    "_moe_input",
    "_moe_output",
)

forward

forward(
    hidden_states: Tensor,
    rope: Tensor,
    varlen_handler: VarlenHandler,
    modality_dispatcher: ModalityDispatcher,
    cp_split_sizes: list[int] | Tensor,
) -> Tensor

Modality

Bases: IntEnum

AUDIO class-attribute instance-attribute

AUDIO = 1

TEXT class-attribute instance-attribute

TEXT = 2

TIME class-attribute instance-attribute

TIME = 3

VIDEO class-attribute instance-attribute

VIDEO = 0

VarlenHandler dataclass

Packed-sequence metadata consumed by MAGI-2 attention.

cu_seqlens_k instance-attribute

cu_seqlens_k: Tensor | None

cu_seqlens_q instance-attribute

cu_seqlens_q: Tensor | None

max_seqlen_k class-attribute instance-attribute

max_seqlen_k: int | None = None

max_seqlen_q class-attribute instance-attribute

max_seqlen_q: int | None = None

resolved

resolved(
    q_tokens: int, k_tokens: int
) -> tuple[Tensor, Tensor, int, int]