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 ¶
MHCHandler ¶
Exact four-stream manifold-constrained hyper-connection math.
matmul_scale instance-attribute ¶
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]
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 ¶
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_weight ¶
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 ¶
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()
]
permute_mapping instance-attribute ¶
permuted_modality_mapping instance-attribute ¶
MultiModalityRMSNorm ¶
Bases: Module
RMSNorm with independent modality and mHC-stream scales.
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