Skip to content

vllm_omni.diffusion.models.magi2.layers

Checkpoint-compatible native layers for MAGI-2 Preview.

The grouped-linear, normalization, Fourier, and mHC formulas are adapted from SandAI's Apache-2.0 preview implementation. They have been modified to remove MagiCompiler and external Triton runtime dependencies.

MHCTensorTuple module-attribute

MHCTensorTuple = tuple[
    torch.Tensor, torch.Tensor, torch.Tensor
]

ElementWiseFourierEmbed

Bases: Module

Nine-coordinate Fourier embedding used as MAGI's element-wise RoPE.

bands instance-attribute

bands = nn.Parameter(bands)

dim instance-attribute

dim = dim

temperature instance-attribute

temperature = temperature

forward

forward(coords: Tensor) -> Tensor

MHCHandler

Exact four-stream manifold-constrained hyper-connection math.

dtype instance-attribute

dtype = dtype

hidden_size instance-attribute

hidden_size = hidden_size

matmul_scale instance-attribute

matmul_scale = 1.0 / math.sqrt(
    float(num_streams * hidden_size)
)

num_streams instance-attribute

num_streams = num_streams

sinkhorn_epsilon instance-attribute

sinkhorn_epsilon = sinkhorn_epsilon

sinkhorn_iterations instance-attribute

sinkhorn_iterations = sinkhorn_iterations

apply_pre

apply_pre(
    streams: Tensor,
    alpha_bias_logits: MHCTensorTuple,
    *,
    out_dtype: dtype | None = None,
) -> Tensor

compute_logits

compute_logits(
    flattened: Tensor,
    norm: Callable[[Tensor], Tensor],
    phi_fused: Tensor,
) -> MHCTensorTuple

compute_post_residual

compute_post_residual(
    post: MHCTensorTuple,
    residual: MHCTensorTuple,
    *,
    out_dtype: dtype,
) -> tuple[Tensor, Tensor]

flatten

flatten(tensor: Tensor) -> Tensor

hyper_connect

hyper_connect(
    residual_streams: Tensor,
    branch_output: Tensor,
    post_coefficients: Tensor,
    residual_matrix: Tensor,
) -> Tensor

Magi2GroupedLinear

Bases: Module

Modality-grouped linear with the released flattened weight layout.

bias instance-attribute

bias = nn.Parameter(
    torch.empty(
        num_experts * bias_features,
        dtype=dtype,
        device=device,
    )
)

in_features instance-attribute

in_features = in_features

local_in_features instance-attribute

local_in_features = in_features

local_out_features instance-attribute

local_out_features = out_features

num_experts instance-attribute

num_experts = num_experts

out_features instance-attribute

out_features = out_features

parallel_mode instance-attribute

parallel_mode = parallel_mode

qkv_splits instance-attribute

qkv_splits = qkv_splits

tp_group instance-attribute

tp_group = tp_group or get_magi2_tp_group()

weight instance-attribute

weight = nn.Parameter(
    torch.empty(
        num_experts * self.local_out_features,
        self.local_in_features,
        dtype=dtype,
        device=device,
    )
)

forward

forward(
    tensor: Tensor,
    modality_dispatcher: ModalityDispatcher | None = None,
) -> Tensor

shard_checkpoint_bias

shard_checkpoint_bias(checkpoint_tensor: Tensor) -> Tensor

shard_checkpoint_weight

shard_checkpoint_weight(
    checkpoint_tensor: Tensor,
) -> Tensor

Convert the released grouped weight into this TP rank's shard.

ModalityDispatcher

Precompute stable modality grouping and inverse permutation metadata.

cu_group_sizes instance-attribute

cu_group_sizes = F.pad(
    torch.cumsum(self.group_size, dim=0), (1, 0)
)

group_size instance-attribute

group_size = torch.bincount(
    self.permuted_modality_mapping.long(),
    minlength=num_modalities,
).to(torch.int32)

group_size_cpu instance-attribute

group_size_cpu = [
    int(value) for value in self.group_size.cpu().tolist()
]

inv_permute_mapping instance-attribute

inv_permute_mapping = torch.argsort(self.permute_mapping)

modality_mapping instance-attribute

modality_mapping = modality_mapping

num_modalities instance-attribute

num_modalities = num_modalities

permute_mapping instance-attribute

permute_mapping = torch.argsort(
    modality_mapping, stable=True
)

permuted_modality_mapping instance-attribute

permuted_modality_mapping = modality_mapping.index_select(
    0, self.permute_mapping
)

dispatch

dispatch(tensor: Tensor) -> list[Tensor]

inverse_permute

inverse_permute(tensor: Tensor) -> Tensor

permute

permute(tensor: Tensor) -> Tensor

undispatch staticmethod

undispatch(*groups: Tensor) -> Tensor

MultiModalityRMSNorm

Bases: Module

RMSNorm with independent modality and mHC-stream scales.

dim instance-attribute

dim = dim

eps instance-attribute

eps = eps

num_modality instance-attribute

num_modality = num_modality

num_patterns instance-attribute

num_patterns = num_patterns

out_dtype instance-attribute

out_dtype = out_dtype

weight instance-attribute

weight = nn.Parameter(
    torch.zeros(
        num_patterns * dim * num_modality,
        dtype=torch.float32,
        device=device,
    )
)

forward

forward(
    tensor: Tensor,
    modality_dispatcher: ModalityDispatcher | None = None,
) -> Tensor

make_grouped_linear

make_grouped_linear(
    in_features: int,
    out_features: int,
    *,
    num_experts: int = 1,
    bias: bool = False,
    dtype: dtype | None = None,
    parallel_mode: str | None = None,
    qkv_splits: tuple[int, int, int] | None = None,
) -> Magi2GroupedLinear

sinkhorn_knopp

sinkhorn_knopp(
    matrix_logits: Tensor, iterations: int, epsilon: float
) -> Tensor

swiglu7

swiglu7(
    x: Tensor,
    alpha: float = 1.702,
    limit: float = 7.0,
    out_dtype: dtype | None = None,
) -> Tensor

Released GPT-OSS-style clamped SwiGLU activation.